diff --git a/src/Blocks/Blocks.jl b/src/Blocks/Blocks.jl index 2670de16..e07ec5e8 100644 --- a/src/Blocks/Blocks.jl +++ b/src/Blocks/Blocks.jl @@ -5,6 +5,7 @@ module Blocks using ModelingToolkitBase: ModelingToolkitBase, @component, @connector, @named, @parameters, @unpack, System, compose, connect, extend, getdefault, t_nounits as t, D_nounits as D +using SymbolicUtils: @syms, symtype using Symbolics: Symbolics, @register_symbolic, @variables, Differential, Equation import Base: ifelse import ..@symcheck diff --git a/src/Blocks/sources.jl b/src/Blocks/sources.jl index 9e282d86..e7463a5d 100644 --- a/src/Blocks/sources.jl +++ b/src/Blocks/sources.jl @@ -864,7 +864,15 @@ end Base.nameof(::CachedInterpolation) = :CachedInterpolation -@register_symbolic (f::CachedInterpolation)(u::AbstractArray, x::AbstractArray, args::Tuple) +const _INTERPOLATOR_SYMTYPE = let + @syms interpolator(::Real)::Real + symtype(interpolator) +end + +# Calling the builder returns the interpolation function, not an interpolated scalar. +@register_symbolic (f::CachedInterpolation)( + u::AbstractArray, x::AbstractArray, args::Tuple +)::_INTERPOLATOR_SYMTYPE """ ParametrizedInterpolation(interp_type, u, x, args...; name, t = ModelingToolkitBase.t_nounits) diff --git a/test/sources.jl b/test/sources.jl index 31efad4d..780830e3 100644 --- a/test/sources.jl +++ b/test/sources.jl @@ -14,6 +14,7 @@ using SymbolicIndexingInterface using SciMLStructures: SciMLStructures, Tunable using ForwardDiff using ADTypes +using SymbolicUtils: symtype @testset "Constant" begin @named src = Constant(k = 2) @@ -650,6 +651,9 @@ end @testset "LinearInterpolation" begin @named i = ParametrizedInterpolation(LinearInterpolation, u, x) + interpolator = only(values(bindings(i))) + @test symtype(interpolator(0.0)) === Real + eqs = [i.input.u ~ t, D(y) ~ i.output.u] @named model = System(eqs, t, systems = [i])