diff --git a/src/plugins/variational_constraints/variational_constraints_engine.jl b/src/plugins/variational_constraints/variational_constraints_engine.jl index 11878392..f9f532f1 100644 --- a/src/plugins/variational_constraints/variational_constraints_engine.jl +++ b/src/plugins/variational_constraints/variational_constraints_engine.jl @@ -951,7 +951,7 @@ function apply_constraints!( applicable_nodes = unroll_nocreate(context[getvariables(marginal_constraint)]) for node in applicable_nodes if hasextra(model[node], VariationalConstraintsMarginalFormConstraintKey) - @warn lazy"Node $node already has functional form constraint $(opt[:q]) applied, therefore $constraint_data will not be applied" + @warn lazy"Node $node already has functional form constraint $(getextra(model[node], VariationalConstraintsMarginalFormConstraintKey)) applied, therefore $(getconstraint(marginal_constraint)) will not be applied" else setextra!(model[node], VariationalConstraintsMarginalFormConstraintKey, getconstraint(marginal_constraint)) end @@ -966,7 +966,7 @@ function apply_constraints!(model::Model, context::Context, message_constraint:: applicable_nodes = unroll_nocreate(context[getvariables(message_constraint)]) for node in applicable_nodes if hasextra(model[node], VariationalConstraintsMessagesFormConstraintKey) - @warn lazy"Node $node already has functional form constraint $(opt[:q]) applied, therefore $constraint_data will not be applied" + @warn lazy"Node $node already has functional form constraint $(getextra(model[node], VariationalConstraintsMessagesFormConstraintKey)) applied, therefore $(getconstraint(message_constraint)) will not be applied" else setextra!(model[node], VariationalConstraintsMessagesFormConstraintKey, getconstraint(message_constraint)) end diff --git a/test/plugins/variational_constraints/variational_constraints_engine_tests.jl b/test/plugins/variational_constraints/variational_constraints_engine_tests.jl index 1e7aae25..f64f37b0 100644 --- a/test/plugins/variational_constraints/variational_constraints_engine_tests.jl +++ b/test/plugins/variational_constraints/variational_constraints_engine_tests.jl @@ -613,6 +613,52 @@ end end end + +@testitem "Applying a second form constraint warns and preserves the original" begin + import GraphPPL: + create_model, + MarginalFormConstraint, + MessageFormConstraint, + IndexedVariable, + apply_constraints!, + getextra, + VariationalConstraintsMarginalFormConstraintKey, + VariationalConstraintsMessagesFormConstraintKey + + include("../../testutils.jl") + + using .TestUtils.ModelZoo + + struct FirstArbitraryFormConstraint end + struct SecondArbitraryFormConstraint end + + # A node that already carries a marginal form constraint must warn (not throw) and keep the first one + model = create_model(simple_model()) + context = GraphPPL.getcontext(model) + apply_constraints!(model, context, MarginalFormConstraint(IndexedVariable(:x, nothing), FirstArbitraryFormConstraint())) + + @test_logs (:warn, r"already has functional form constraint") match_mode = :any apply_constraints!( + model, context, MarginalFormConstraint(IndexedVariable(:x, nothing), SecondArbitraryFormConstraint()) + ) + + for node in filter(GraphPPL.as_variable(:x), model) + @test getextra(model[node], VariationalConstraintsMarginalFormConstraintKey) == FirstArbitraryFormConstraint() + end + + # ... and the same for message form constraints + model = create_model(simple_model()) + context = GraphPPL.getcontext(model) + apply_constraints!(model, context, MessageFormConstraint(IndexedVariable(:x, nothing), FirstArbitraryFormConstraint())) + + @test_logs (:warn, r"already has functional form constraint") match_mode = :any apply_constraints!( + model, context, MessageFormConstraint(IndexedVariable(:x, nothing), SecondArbitraryFormConstraint()) + ) + + for node in filter(GraphPPL.as_variable(:x), model) + @test getextra(model[node], VariationalConstraintsMessagesFormConstraintKey) == FirstArbitraryFormConstraint() + end +end + @testitem "save constraints with constants via `mean_field_constraint!`" begin using BitSetTuples import GraphPPL: