From 9b5b0199ba86a84041ab7eaccfa928e5db3a282a Mon Sep 17 00:00:00 2001 From: Sven Berger Date: Thu, 28 May 2026 11:31:22 +0200 Subject: [PATCH] Align SciML Dict tuple pairs --- src/styles/sciml/nest.jl | 163 ++++++++++++++++++++++++++++++++++++++- test/sciml_style.jl | 64 +++++++++++++++ 2 files changed, 226 insertions(+), 1 deletion(-) diff --git a/src/styles/sciml/nest.jl b/src/styles/sciml/nest.jl index 974832e66..55c1b9cc5 100644 --- a/src/styles/sciml/nest.jl +++ b/src/styles/sciml/nest.jl @@ -44,6 +44,162 @@ end # nest!(style, fst[end], s, lineage) # end +function _align_tuple_comments!(fst::FST) + for n in fst.nodes::Vector + n.typ === NOTCODE && (n.indent = fst.indent) + end +end + +function _is_dict_call(fst::FST) + fst.typ === Call && + length(fst.nodes::Vector) > 0 && + fst[1].typ === IDENTIFIER && + fst[1].val == "Dict" +end + +function _is_tuple_pair(fst::FST) + fst.typ === Binary && op_kind(fst) === K"=>" && fst[end].typ === TupleN +end + +function _align_dict_tuple_pair_arrows!(fst::FST) + pair_inds = findall(_is_tuple_pair, fst.nodes::Vector) + length(pair_inds) > 1 || return + + align_len = maximum(node_align_length(fst[i][1]) for i in pair_inds) + for i in pair_inds + pair = fst[i] + ws = align_len - node_align_length(pair[1]) + 1 + align_binaryopcall!(pair, ws) + end +end + +function _tuple_rhs_line_fits(rhs::FST, indent::Int, s::State) + nodes = rhs.nodes::Vector + idx = findfirst( + n -> + !(n.typ in (PUNCTUATION, WHITESPACE, PLACEHOLDER, NEWLINE, TRAILINGCOMMA)) && !is_closer(n), + nodes, + ) + isnothing(idx) && return true + + width, _ = length_to(rhs, (NEWLINE, PLACEHOLDER); start = idx) + return indent + width + rhs.extra_margin <= s.opts.margin +end + +function _align_pair_tuple_rhs!( + fst::FST, + s::State; + align_to_pair::Bool = true, + fallback_indent::Int = fst.indent + s.opts.indent, +) + rhs = fst[end] + rhs.typ === TupleN || return + + desired_indent = if !align_to_pair + fallback_indent + elseif length(rhs.nodes::Vector) > 1 && rhs[2].typ === NEWLINE + rhs.indent - 1 + else + fst.indent + node_align_length(fst[1:(end-1)]) + 1 + end + add_indent!(rhs, s, desired_indent - rhs.indent) + _align_tuple_comments!(rhs) +end + +function _unnest_pair_tuple_body!(style::AbstractStyle, rhs::FST, s::State) + lo = s.line_offset + s.line_offset = rhs.indent + for n in rhs.nodes::Vector + n.typ === NEWLINE && (s.line_offset = rhs.indent; continue) + walk(unnest!(style; dedent = false), n, s) + end + s.line_offset = lo +end + +function _preserve_pair_tuple_newlines!(fst::FST; closer_indent::Int = fst.indent) + rhs = fst[end] + nodes = rhs.nodes::Vector + length(nodes) > 2 || return + + nodes[2].typ !== NEWLINE && (nodes[2] = Newline(; length = nodes[2].len)) + if is_closer(nodes[end]) && nodes[end-1].typ !== NEWLINE + nodes[end-1] = Newline(; length = nodes[end-1].len) + nodes[end].indent = closer_indent + end +end + +function _preserve_multiline_closer!(fst::FST) + nodes = fst.nodes::Vector + length(nodes) > 2 && is_closer(nodes[end]) || return + + nodes[end-1].typ !== NEWLINE && (nodes[end-1] = Newline(; length = nodes[end-1].len)) +end + +function _unnest_short_binary_lines!(fst::FST, s::State) + is_leaf(fst) && return + + if fst.typ === Binary && !contains_comment(fst) && fst.line_offset >= 0 + nl_inds = findall(n -> n.typ === NEWLINE, fst.nodes::Vector) + if length(nl_inds) > 0 && + fst.line_offset + fst.extra_margin + node_align_length(fst) <= s.opts.margin + nl_to_ws!(fst, nl_inds) + end + end + + for n in fst.nodes::Vector + _unnest_short_binary_lines!(n, s) + end +end + +function n_binaryopcall!( + ss::SciMLStyle, + fst::FST, + s::State, + lineage::Vector{Tuple{FNode,Union{Nothing,Metadata}}}; + indent::Int = -1, +) + if op_kind(fst) === K"=>" && + fst[end].typ === TupleN && + fst[1].endline == fst[end].startline + style = getstyle(ss) + rhs_style = YASStyle(style) + nodes = fst.nodes::Vector + oplen = sum(length.(fst[2:end])) + line_offset = s.line_offset + fallback_indent = line_offset + s.opts.indent + nested = false + for (i, n) in enumerate(nodes) + if n.typ === NEWLINE + s.line_offset = fst.indent + elseif i == 1 + n.extra_margin = oplen + fst.extra_margin + nested |= nest!(style, n, s, lineage) + elseif i == length(nodes) + n.extra_margin = fst.extra_margin + align_to_pair = _tuple_rhs_line_fits(n, s.line_offset, s) + n.indent = align_to_pair ? s.line_offset : fallback_indent + nested |= nest!(rhs_style, n, s, lineage) + _align_pair_tuple_rhs!(fst, s; align_to_pair, fallback_indent) + if align_to_pair + lo = s.line_offset + walk(unnest!(rhs_style; dedent = false), n, s) + s.line_offset = lo + _unnest_short_binary_lines!(n, s) + else + _unnest_pair_tuple_body!(rhs_style, n, s) + _preserve_pair_tuple_newlines!(fst; closer_indent = line_offset) + end + _align_tuple_comments!(n) + else + nested |= nest!(style, n, s, lineage) + end + end + return nested + end + + n_binaryopcall!(DefaultStyle(getstyle(ss)), fst, s, lineage; indent) +end + function n_functiondef!( ss::SciMLStyle, fst::FST, @@ -103,8 +259,12 @@ function _n_tuple!( lineage::Vector{Tuple{FNode,Union{Nothing,Metadata}}}, ) style = getstyle(ss) - line_margin = s.line_offset + length(fst) + fst.extra_margin nodes = fst.nodes::Vector + _is_dict_call(fst) && + fst.startline != fst.endline && + _align_dict_tuple_pair_arrows!(fst) + + line_margin = s.line_offset + length(fst) + fst.extra_margin has_closer = is_closer(fst[end]) start_line_offset = s.line_offset @@ -224,6 +384,7 @@ function _n_tuple!( s.line_offset = start_line_offset walk(unnest!(style; dedent = false), fst, s) + _is_dict_call(fst) && fst.startline != fst.endline && _preserve_multiline_closer!(fst) s.line_offset = start_line_offset walk(increment_line_offset!, fst, s) diff --git a/test/sciml_style.jl b/test/sciml_style.jl index 8ed8b5799..1185264e2 100644 --- a/test/sciml_style.jl +++ b/test/sciml_style.jl @@ -621,6 +621,70 @@ """ @test format_text(str, SciMLStyle(); margin = 13) == fstr end + + @testset "dict tuple pair regression" begin + str = """ + function build_search_cases(item_spacing) + algorithm_cases = Dict( + "default search" => (), + "with static scheduler" => (backend=StaticBackend(),), + "with damping source term" => (source_terms=DampingTerm(coefficient=1e-4),), + "with lookup density" => (density_calculator=LookupDensity(), + clip_negative_pressure=true), + "with weighted viscosity" => ( + # from 0.02*10.0*1.2*0.05/8 + viscosity_model=WeightedViscosity(nu=0.0015),), + "with spline kernel" => (smoothing_length=1.1 * + item_spacing, + smoothing_kernel=SplineKernel{2}()) + ) + end + """ + + fstr = """ + function build_search_cases(item_spacing) + algorithm_cases = Dict( + "default search" => (), + "with static scheduler" => (backend=StaticBackend(),), + "with damping source term" => (source_terms=DampingTerm(coefficient=1e-4),), + "with lookup density" => (density_calculator=LookupDensity(), + clip_negative_pressure=true), + "with weighted viscosity" => ( + # from 0.02*10.0*1.2*0.05/8 + viscosity_model=WeightedViscosity(nu=0.0015),), + "with spline kernel" => (smoothing_length=1.1 * item_spacing, + smoothing_kernel=SplineKernel{2}()) + ) + end + """ + formatted = format_text(str, SciMLStyle(); whitespace_in_kwargs = false) + @test formatted == fstr + + formatted_lines = split(formatted, '\n') + comment_line = only(filter(line -> contains(line, "# from"), formatted_lines)) + option_line = + only(filter(line -> contains(line, "viscosity_model"), formatted_lines)) + @test length(comment_line) - length(lstrip(comment_line)) == + length(option_line) - length(lstrip(option_line)) + + str = """ + d = Dict( + "a moderately long key" => ( + x = some_long_function_name(alpha, beta), + ), + ) + """ + fstr = """ + d = Dict( + "a moderately long key" => ( + x = some_long_function_name(alpha, beta), + ), + ) + """ + formatted = format_text(str, SciMLStyle(); margin = 50) + @test formatted == fstr + @test all(length(line) <= 50 for line in split(formatted, '\n')) + end end str = raw"""