Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
58 changes: 34 additions & 24 deletions lib/ModelingToolkitBase/src/utils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -936,6 +936,16 @@ function collect_vars!(
return nothing
end

# Break the inference cycle between the mutually recursive collectors, which Julia
# 1.10 can otherwise miscompile after unrelated method additions. Dispatching in the
# latest world age also means a downstream `collect_vars!` method defined while this
# call is already running is still found; because the four-argument fallback returns
# `nothing`, missing it would silently drop parameters rather than error. The explicit
# keyword preserves metadata recursion's depth-zero semantics.
function _call_collect_vars!(unknowns, parameters, expr, iv)
return @invokelatest collect_vars!(unknowns, parameters, expr, iv; depth = 0)
end

"""
$(TYPEDSIGNATURES)

Expand Down Expand Up @@ -966,7 +976,7 @@ function collect_var!(unknowns::OrderedSet{SymbolicT}, parameters::OrderedSet{Sy
any(!SU.isconst, Iterators.drop(arguments(var), 1))
)
for arg in Iterators.drop(arguments(var), 1)
collect_vars!(unknowns, parameters, arg, iv)
_call_collect_vars!(unknowns, parameters, arg, iv)
end
var = arr
end
Expand All @@ -975,7 +985,7 @@ function collect_var!(unknowns::OrderedSet{SymbolicT}, parameters::OrderedSet{Sy
if iscalledparameter(var)
callable = getcalledparameter(var)
push!(parameters, callable)
collect_vars!(unknowns, parameters, arguments(var), iv)
_call_collect_vars!(unknowns, parameters, arguments(var), iv)
elseif isparameter(var) || (iscall(var) && isparameter(operation(var)))
push!(parameters, var)
else
Expand All @@ -984,55 +994,55 @@ function collect_var!(unknowns::OrderedSet{SymbolicT}, parameters::OrderedSet{Sy
# Add also any parameters that appear only as defaults in the var
if hasdefault(var) && (def = getdefault(var)) !== missing
if def isa SymbolicT
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa Num
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa Arr{Num, 1}
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa Arr{Num, 2}
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa CallAndWrap{Num}
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa CallAndWrap{Arr{Num, 1}}
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa CallAndWrap{Arr{Num, 2}}
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa Arr
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa CallAndWrap
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
else
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
end
end
# Add also any parameters that appear only in the bounds of the var
if hasbounds(var)
(lo, hi) = getbounds(var)
if lo isa SymbolicT
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
elseif lo isa Num
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
elseif lo isa Arr{Num, 1}
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
elseif lo isa Arr{Num, 2}
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
elseif lo isa Arr
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
else
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
end
if hi isa SymbolicT
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
elseif hi isa Num
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
elseif hi isa Arr{Num, 1}
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
elseif hi isa Arr{Num, 2}
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
elseif hi isa Arr
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
else
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
end
end
return nothing
Expand Down
101 changes: 81 additions & 20 deletions lib/ModelingToolkitBase/test/dq_units.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ using ModelingToolkitBase, OrdinaryDiffEq, JumpProcesses, DynamicQuantities
using Symbolics
import SymbolicUtils as SU
using Test
using DataStructures: OrderedSet
MT = ModelingToolkitBase
using ModelingToolkitBase: t, D
@parameters τ [unit = u"s"] γ
Expand Down Expand Up @@ -273,30 +274,90 @@ let
@test MT.get_unit(x_mat) == u"1"
end

# Issue #4211: collect_vars! must discover parameters in defaults with DynamicQuantities loaded
# On Julia 1.10, loading DynamicQuantities could cause collect_vars! to fail to discover
# parameters used in variable defaults due to a method invalidation bug.
@testset "Issue #4211: collect_vars! discovers parameters in defaults" begin
using DataStructures: OrderedSet
struct CollectorInvalidationString
value::String
end

struct CollectorInvalidationExpression
parameter::Symbolics.SymbolicT
end

struct CollectorSameWorldExpression
parameter::Symbolics.SymbolicT
end

@parameters X0_test
@variables X_test(t) = X0_test
function collect_default_parameters(var)
unknowns = OrderedSet{Symbolics.SymbolicT}()
parameters = OrderedSet{Symbolics.SymbolicT}()
MT.collect_vars!(
unknowns, parameters, Symbolics.unwrap(var), Symbolics.unwrap(t),
Symbolics.Operator; depth = 0
)
return parameters
end

@testset "collect_vars! survives late method invalidation" begin
@parameters x0_test y0_test scale_test
@variables x_test(t) = x0_test y_test(t) = scale_test * y0_test

us = OrderedSet{Symbolics.SymbolicT}()
ps = OrderedSet{Symbolics.SymbolicT}()
MT.collect_vars!(us, ps, Symbolics.unwrap(X_test), Symbolics.unwrap(t), Symbolics.Operator; depth = 0)
@test Symbolics.unwrap(x0_test) in collect_default_parameters(x_test)

# X0_test should be discovered in X_test's default value
@test Symbolics.unwrap(X0_test) in ps
# This late definition exercises the Julia 1.10 invalidation that caused
# parameters in defaults to disappear after loading unrelated packages.
@eval Base.convert(::Type{Symbol}, value::CollectorInvalidationString) =
Symbol(value.value)

# Test with expression in default
@parameters a_test b_test
@variables Y_test(t) = a_test + 2 * b_test
x_parameters = collect_default_parameters(x_test)
y_parameters = collect_default_parameters(y_test)

empty!(us)
empty!(ps)
MT.collect_vars!(us, ps, Symbolics.unwrap(Y_test), Symbolics.unwrap(t), Symbolics.Operator; depth = 0)
@test Symbolics.unwrap(x0_test) in x_parameters
@test Symbolics.unwrap(y0_test) in y_parameters
@test Symbolics.unwrap(scale_test) in y_parameters

custom_default = MT.setdefault(
x_test, CollectorInvalidationExpression(Symbolics.unwrap(scale_test))
)
@eval function MT.collect_vars!(
unknowns::OrderedSet{Symbolics.SymbolicT},
parameters::OrderedSet{Symbolics.SymbolicT},
expr::CollectorInvalidationExpression,
::Union{Symbolics.SymbolicT, Nothing};
depth = 0
)
push!(parameters, expr.parameter)
return nothing
end

@test Symbolics.unwrap(scale_test) in collect_default_parameters(custom_default)
@test Symbolics.unwrap(x0_test) in collect_default_parameters(x_test)
end

# Defines the downstream method from inside a running call, so the enclosing frame
# keeps its old world age while collecting.
function define_then_collect(var)
@eval function MT.collect_vars!(
unknowns::OrderedSet{Symbolics.SymbolicT},
parameters::OrderedSet{Symbolics.SymbolicT},
expr::CollectorSameWorldExpression,
::Union{Symbolics.SymbolicT, Nothing};
depth = 0
)
push!(parameters, expr.parameter)
return nothing
end
return collect_default_parameters(var)
end

@testset "collect_vars! dispatches to same-world-age downstream methods" begin
@parameters same_world_test
@variables same_world_var(t)

target = MT.setdefault(
same_world_var, CollectorSameWorldExpression(Symbolics.unwrap(same_world_test))
)

@test Symbolics.unwrap(a_test) in ps
@test Symbolics.unwrap(b_test) in ps
# The four-argument fallback returns `nothing`, so failing to see the method here
# would drop the parameter silently instead of raising a `MethodError`.
@test Symbolics.unwrap(same_world_test) in define_then_collect(target)
@test Symbolics.unwrap(same_world_test) in collect_default_parameters(target)
end
Loading