From 29360c3352ef250656c68da1a41943809d6f71cb Mon Sep 17 00:00:00 2001 From: Jake LoRocco Date: Wed, 26 Aug 2026 16:25:13 -0400 Subject: [PATCH 1/5] fix: validation input and context along with requirement wording Signed-off-by: Jake LoRocco Assisted-by: CLAUDE:OPUS --- docs/docs/advanced/lora-and-alora-adapters.md | 15 +++ docs/docs/concepts/requirements-system.md | 43 ++++++ docs/docs/how-to/write-custom-verifiers.md | 28 ++++ mellea/core/base.py | 9 +- mellea/core/requirement.py | 59 ++++++-- mellea/core/sampling.py | 3 +- mellea/stdlib/components/genstub.py | 65 +++++---- mellea/stdlib/functional.py | 57 +++++--- mellea/stdlib/requirements/requirement.py | 1 + mellea/stdlib/sampling/base.py | 15 ++- mellea/stdlib/sampling/budget_forcing.py | 12 +- mellea/stdlib/sampling/majority_voting.py | 3 +- mellea/stdlib/sampling/sofai.py | 10 +- mellea/stdlib/session.py | 18 +-- .../prompts/default/Requirement.jinja2 | 3 +- test/core/test_stream_validate.py | 6 +- test/stdlib/requirements/test_requirement.py | 127 +++++++++++++++++- .../sampling/test_sampling_base_unit.py | 97 +++++++++++++ test/stdlib/test_functional_unit.py | 94 ++++++++++++- 19 files changed, 564 insertions(+), 101 deletions(-) diff --git a/docs/docs/advanced/lora-and-alora-adapters.md b/docs/docs/advanced/lora-and-alora-adapters.md index 51ffaa6470..3deeb6baff 100644 --- a/docs/docs/advanced/lora-and-alora-adapters.md +++ b/docs/docs/advanced/lora-and-alora-adapters.md @@ -188,6 +188,21 @@ directly (bypassing the normal `validate()` call), use `ALoraRequirement` from but still requires a matching adapter to actually be registered. If none is found, Mellea logs a warning and falls back to regular generation rather than erroring. +### What the adapter sees + +The `requirement-check` adapter judges the last assistant turn of the conversation it is +given, so the routing above only produces useful verdicts if that conversation actually +reaches it. Mellea builds the adapter's message list from the validation context's +`view_for_generation()`, and validation runs over the post-generation context — the same +conversation the model generated into, with the generated output last. No extra setup is +needed for that; it is the default. + +One consequence is worth knowing: a context whose generation view is empty gives the +adapter nothing to judge. `SimpleContext` is the case to watch — it retains +`last_output()` but its `view_for_generation()` is always empty by design. Adapter-backed +requirements need a context that renders the assistant turn, such as `ChatContext`. +See [What the validator sees](../concepts/requirements-system.md#what-the-validator-sees). + ## Disable adapter validation To run without adapter validation (for benchmarking or debugging): diff --git a/docs/docs/concepts/requirements-system.md b/docs/docs/concepts/requirements-system.md index a998a2ef9f..0b3ffda6c9 100644 --- a/docs/docs/concepts/requirements-system.md +++ b/docs/docs/concepts/requirements-system.md @@ -296,6 +296,49 @@ reserve LLM-based requirements for subjective criteria that cannot be coded dire For a full walkthrough of using LLM-as-a-judge for output quality evaluation, see [Evaluate with LLM-as-a-Judge](../how-to/evaluate-with-llm-as-a-judge). +## What the validator sees + +Validation runs over the **post-generation context** — the same conversation the model +just generated into, with the generated output as its last entry. This matters for all +three validation approaches: + +- A `validation_fn` receives that context, so `ctx.last_output()` is the output under + judgement and `ctx.as_list()` is the conversation that produced it. +- An LLM-as-a-judge requirement sends that conversation to the model, followed by the + specific output to judge. The judge prompt scopes the verdict to that output, while + telling the model the conversation is there to help it interpret the requirement — + which is what makes conversation-dependent requirements ("must answer the question + that was asked") judgeable at all. +- An `ALoraRequirement` needs the conversation to reach the adapter at all — the + `requirement-check` adapter judges the last assistant turn of the conversation it is + given. See [LoRA and aLoRA Adapters](../advanced/lora-and-alora-adapters.md). + +The output under judgement is passed as the `ModelOutputThunk` itself, not a detached +string copy, so its `parsed_repr` and any structured representation survive into the +judge prompt. + +### Overriding the validation context + +`SamplingStrategy.sample()` takes an optional `validation_ctx` for validating over +something other than the post-generation context — for example, a trimmed context that +excludes long retrieved documents. When it is given, the sampled output is appended to it +so it is still the validation target; when it is `None`, each attempt is validated over its +own post-generation context. + +`instruct()`, `act()`, and their async counterparts always pass `None`, so the +post-generation context is what you get unless you drive a strategy's `sample()` directly. + +`validate()` and `avalidate()` follow the same rule at a lower level: they validate over +the `context` you hand them, in your own context type. Their `output` argument designates +*which* output is under judgement rather than replacing the context — it is appended only +when it is not already the context's last output. + +> **Deprecated:** the `input` parameter of `validate()` / `avalidate()` is deprecated and +> will be removed in a future release. Pass a context that already contains the input. + +Preconditions are the one case that never sees the conversation: `precondition_requirements` +are judged over the function arguments alone, in a fresh context. + ## Composing requirements Requirements are composable: mix strings, `req()`, `check()`, and `Requirement` diff --git a/docs/docs/how-to/write-custom-verifiers.md b/docs/docs/how-to/write-custom-verifiers.md index 62456e79b9..124aeee117 100644 --- a/docs/docs/how-to/write-custom-verifiers.md +++ b/docs/docs/how-to/write-custom-verifiers.md @@ -77,6 +77,34 @@ result = m.instruct( print(str(result)) ``` +### Which context your function receives + +The `ctx` argument is the **post-generation context**: the conversation the model just +generated into, with the output under judgement as its last entry. So `ctx.last_output()` +is the output to check, and `ctx.as_list()` gives you the turns that produced it — useful +when a verdict depends on what was asked, not just what was answered: + +```python +from mellea.core import Context, ValidationResult + + +def validate_answers_the_question(ctx: Context) -> ValidationResult: + """Fail if the output repeats the user's question back instead of answering it.""" + turns = ctx.as_list() + output = ctx.last_output() + text = output.value if output and output.value else "" + + prior = [str(turn) for turn in turns[:-1]] + if prior and text.strip() in prior[-1]: + return ValidationResult(False, reason="The response echoes the question.") + return ValidationResult(True) +``` + +Under `instruct()` and `act()` this is always the post-generation context. Driving a +sampling strategy's `sample()` directly lets you pass a `validation_ctx` to validate over +something else instead — see +[Overriding the validation context](../concepts/requirements-system.md#overriding-the-validation-context). + ## Common validation patterns ### JSON validity diff --git a/mellea/core/base.py b/mellea/core/base.py index 1080ef94b3..5cebb4170a 100644 --- a/mellea/core/base.py +++ b/mellea/core/base.py @@ -1976,7 +1976,14 @@ class TemplateRepresentation: obj: Any args: dict[ str, - str | Component | CBlock | Iterable | Mapping | TemplateRepresentation | None, + str + | Component + | CBlock + | ModelOutputThunk + | Iterable + | Mapping + | TemplateRepresentation + | None, ] tools: dict[str, AbstractMelleaTool] | None = ( None # the key must be the name of the function. diff --git a/mellea/core/requirement.py b/mellea/core/requirement.py index 592595b177..bf3788bfcb 100644 --- a/mellea/core/requirement.py +++ b/mellea/core/requirement.py @@ -243,8 +243,25 @@ def __init__( self.validation_fn = validation_fn self.check_only = check_only - # Used for validation. Do not manually populate. - self._output: str | None = None + # The span under judgement. Bound only on the transient copy created inside + # `validate` (see `_bind_validation_target`). Do not manually populate. + self._validation_target: Span | None = None + + def _bind_validation_target(self, target: Span) -> "Requirement": + """Return a shallow copy of this requirement bound to the span under judgement. + + Binding happens on a copy so that a `Requirement` object can be reused across + validation calls (and across sampling iterations) without accumulating state. + + Args: + target: The `Component`, `CBlock`, or `ModelOutputThunk` being validated. + + Returns: + Requirement: A copy of this requirement whose `_validation_target` is `target`. + """ + bound = copy(self) + bound._validation_target = target + return bound async def validate( self, @@ -287,10 +304,10 @@ async def validate( " Context has no appropriate last output" ) - # Create a copy of the requirement that holds the output - # and its template gets populated with the output correctly. - req_copy = copy(self) - req_copy._output = last_output.value + # Bind the output being judged to a copy of this requirement, so that the + # judgement request carries the output itself -- not a detached string copy + # of it -- and so that `self` is left unmodified. + req_copy = self._bind_validation_target(last_output) llm_as_a_judge_result, val_ctx = await backend.generate_from_context( req_copy, ctx, format=format, model_options=model_options ) @@ -367,29 +384,43 @@ async def stream_validate( def parts(self) -> list[Span]: """Returns all of the constituent parts of a Requirement. + Once a validation target has been bound (inside a `validate` call), that target is + the requirement's sole part. Exposing it here is what allows `generate_walk` to + await it if it has not been computed yet. + Returns: - List of constituent components. Empty by default; subclasses override - to expose their internal structure. + List of constituent components. Empty unless a validation target is bound. """ - return [] + return [] if self._validation_target is None else [self._validation_target] def format_for_llm(self) -> TemplateRepresentation | str: """Returns a `TemplateRepresentation` for LLM-as-a-Judge evaluation of this requirement. - Populates the template with the requirement's `description` and the stored model - `_output`. Must only be called from within a `validate` call for this same requirement, - after `_output` has been set. + Populates the template with the requirement's `description` and the bound validation + target. Must only be called from within a `validate` call for this same requirement, + after the target has been bound by `_bind_validation_target`. + + When the target is a `ModelOutputThunk` with a `Component` `parsed_repr`, the parsed + representation is used so that the judge sees the structured output rather than the + raw generated string. Returns: TemplateRepresentation | str: A `TemplateRepresentation` containing the description and the model output to be judged. """ - assert self._output is not None, ( + assert self._validation_target is not None, ( "Object protocol error: should never try to templatize a Requirement except inside of a validate call for that same requirement." ) + + target: Span = self._validation_target + if isinstance(target, ModelOutputThunk) and isinstance( + target.parsed_repr, Component + ): + target = target.parsed_repr + return TemplateRepresentation( obj=self, - args={"description": self.description, "output": self._output}, + args={"description": self.description, "output": target}, tools=None, template_order=["*", "Requirement"], ) diff --git a/mellea/core/sampling.py b/mellea/core/sampling.py index 7b2d84a140..24785487cb 100644 --- a/mellea/core/sampling.py +++ b/mellea/core/sampling.py @@ -148,7 +148,8 @@ async def sample( context: The context to be passed to the sampling strategy. backend: The backend used for generating samples. requirements: List of requirements to test against (merged with global requirements). - validation_ctx: Optional context to use for validation. If None, validation_ctx = ctx. + validation_ctx: Optional context to validate over. If None, each sample is validated + over its own post-generation context. format: output format for structured outputs. model_options: model options to pass to the backend during generation / validation. tool_calls: True if tool calls should be used during this sampling strategy. diff --git a/mellea/stdlib/components/genstub.py b/mellea/stdlib/components/genstub.py index 3ea461f66a..c831d8f3b6 100644 --- a/mellea/stdlib/components/genstub.py +++ b/mellea/stdlib/components/genstub.py @@ -29,6 +29,7 @@ ValidationResult, ) from ...helpers.annotation_helpers import resolve_signature_annotations +from ..context import ChatContext from ..requirements.requirement import reqify from ..session import MelleaSession @@ -625,23 +626,20 @@ def __call__(self, *args, **kwargs) -> tuple[R, Context] | R: # Do precondition validation first. if stub_copy._arguments is not None: - if extracted.m is not None: - val_results = extracted.m.validate( - reqs=stub_copy.precondition_requirements, - model_options=extracted.model_options, - output=ModelOutputThunk(stub_copy._arguments.value), - ) - else: - # We know these aren't None from the `extract_args_and_kwargs` function. - assert extracted.context is not None - assert extracted.backend is not None - val_results = mfuncs.validate( - reqs=stub_copy.precondition_requirements, - context=extracted.context, - backend=extracted.backend, - model_options=extracted.model_options, - output=ModelOutputThunk(stub_copy._arguments.value), - ) + # Preconditions are judged over the arguments alone, so they get a fresh + # context rather than the caller's/session's conversation history. + precondition_backend = ( + extracted.m.backend if extracted.m is not None else extracted.backend + ) + # We know this isn't None from the `extract_args_and_kwargs` function. + assert precondition_backend is not None + val_results = mfuncs.validate( + reqs=stub_copy.precondition_requirements, + context=ChatContext(), + backend=precondition_backend, + model_options=extracted.model_options, + output=ModelOutputThunk(stub_copy._arguments.value), + ) # No retries if precondition validation fails. if not all(bool(val_result) for val_result in val_results): @@ -771,23 +769,22 @@ async def __async_call__() -> tuple[R, Context] | R: # Do precondition validation first. if stub_copy._arguments is not None: - if extracted.m is not None: - val_results = await extracted.m.avalidate( - reqs=stub_copy.precondition_requirements, - model_options=extracted.model_options, - output=ModelOutputThunk(stub_copy._arguments.value), - ) - else: - # We know these aren't None from the `extract_args_and_kwargs` function. - assert extracted.context is not None - assert extracted.backend is not None - val_results = await mfuncs.avalidate( - reqs=stub_copy.precondition_requirements, - context=extracted.context, - backend=extracted.backend, - model_options=extracted.model_options, - output=ModelOutputThunk(stub_copy._arguments.value), - ) + # Preconditions are judged over the arguments alone, so they get a fresh + # context rather than the caller's/session's conversation history. + precondition_backend = ( + extracted.m.backend + if extracted.m is not None + else extracted.backend + ) + # We know this isn't None from the `extract_args_and_kwargs` function. + assert precondition_backend is not None + val_results = await mfuncs.avalidate( + reqs=stub_copy.precondition_requirements, + context=ChatContext(), + backend=precondition_backend, + model_options=extracted.model_options, + output=ModelOutputThunk(stub_copy._arguments.value), + ) # No retries if precondition validation fails. if not all(bool(val_result) for val_result in val_results): diff --git a/mellea/stdlib/functional.py b/mellea/stdlib/functional.py index c5fa4ba266..8654906e53 100644 --- a/mellea/stdlib/functional.py +++ b/mellea/stdlib/functional.py @@ -13,6 +13,7 @@ import asyncio import time import uuid +import warnings from collections.abc import Coroutine, Iterable from typing import Any, Literal, overload @@ -343,24 +344,31 @@ def validate( context: Context, backend: Backend, *, - output: CBlock | ModelOutputThunk | None = None, + output: ModelOutputThunk | None = None, format: type[BaseModelSubclass] | None = None, model_options: dict | None = None, generate_logs: list[GenerateLog] | None = None, # TODO: Can we get rid of gen logs here and in act? input: CBlock | ModelOutputThunk | None = None, ) -> list[ValidationResult]: - """Validates a set of requirements over the output (if provided) or the current context (if the output is not provided). + """Validates a set of requirements over the given context. + + Validation always runs over `context`, in the caller's own context type. `output` + designates *which* output is under judgement; see `avalidate` for details. Args: reqs: A single `Requirement` or a list of them to validate. - context: The current conversation context. + context: The context to validate over. backend: The backend used for LLM-as-a-judge requirements. - output: Optional model output to validate against instead of the context. + output: Optional model output designating the validation target. When `None`, the + context's last output is validated. format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Optional input to include alongside `output` when validating. + input: Deprecated. Optional input to prepend to the validation context. Pass a + context that already contains the input instead. Passing it raises a + `DeprecationWarning` from `avalidate`, so the reported source location is this + wrapper rather than the caller. Returns: List of `ValidationResult` objects, one per requirement. @@ -1023,23 +1031,30 @@ async def avalidate( context: Context, backend: Backend, *, - output: CBlock | ModelOutputThunk | None = None, + output: ModelOutputThunk | None = None, format: type[BaseModelSubclass] | None = None, model_options: dict | None = None, generate_logs: list[GenerateLog] | None = None, input: CBlock | ModelOutputThunk | None = None, ) -> list[ValidationResult]: - """Asynchronous version of .validate; validates a set of requirements over the output (if provided) or the current context (if the output is not provided). + """Asynchronous version of .validate; validates a set of requirements over the given context. + + Validation always runs over `context`, in the caller's own context type, so requirements + (including adapter-backed ones) see the same conversation the model saw. `output` + designates *which* output is under judgement: it is appended to `context` when it is not + already the context's last output, and the requirement then validates that output. Args: reqs: A single `Requirement` or a list of them to validate. - context: The current conversation context. + context: The context to validate over. backend: The backend used for LLM-as-a-judge requirements. - output: Optional model output to validate against instead of the context. + output: Optional model output designating the validation target. When `None`, the + context's last output is validated. format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Optional input to include alongside `output` when validating. + input: Deprecated. Optional input to prepend to the validation context. Pass a + context that already contains the input instead. Returns: List of `ValidationResult` objects, one per requirement. @@ -1050,14 +1065,22 @@ async def avalidate( validation_id = str(uuid.uuid4()) - if output is None: - validation_target_ctx = context - else: - validation_target_ctx = SimpleContext() + validation_target_ctx = context + + if input is not None: + warnings.warn( + "The `input` parameter of validate/avalidate is deprecated and will be removed " + "in a future release. Pass a validation context that already contains the input " + "instead.", + DeprecationWarning, + stacklevel=2, + ) + validation_target_ctx = validation_target_ctx.add(input) - # Add the input/output to the validation context - if input is not None: - validation_target_ctx = validation_target_ctx.add(input) + # `output` designates the validation target rather than replacing the context. Adding it + # is a no-op for the sampling path: ComputedModelOutputThunk reassigns __class__ in place, + # so the thunk passed here *is* the one already in the context. + if output is not None and validation_target_ctx.last_output() is not output: validation_target_ctx = validation_target_ctx.add(output) # --- validation_pre_check hook --- diff --git a/mellea/stdlib/requirements/requirement.py b/mellea/stdlib/requirements/requirement.py index fc072f2aed..9366072a47 100644 --- a/mellea/stdlib/requirements/requirement.py +++ b/mellea/stdlib/requirements/requirement.py @@ -13,6 +13,7 @@ MelleaLogger, ModelOutputThunk, Requirement, + TemplateRepresentation, ValidationResult, ) from ..components.intrinsic import Intrinsic diff --git a/mellea/stdlib/sampling/base.py b/mellea/stdlib/sampling/base.py index 133c45fb5a..d628b558ee 100644 --- a/mellea/stdlib/sampling/base.py +++ b/mellea/stdlib/sampling/base.py @@ -222,7 +222,8 @@ async def sample( context: The context to be passed to the sampling strategy. backend: The backend used for generating samples. requirements: List of requirements to test against (merged with global requirements). - validation_ctx: Optional context to use for validation. If None, validation_ctx = ctx. + validation_ctx: Optional context to validate over. If None, each sample is validated + over its own post-generation context. format: output format for structured outputs. model_options: model options to pass to the backend during generation / validation. tool_calls: True if tool calls should be used during this sampling strategy. @@ -240,8 +241,6 @@ async def sample( the first non-cancellation exception is re-raised (e.g. a backend error). """ - validation_ctx = validation_ctx if validation_ctx is not None else context - flog = MelleaLogger.get_logger() with log_context(strategy=type(self).__name__, loop_budget=self.loop_budget): @@ -490,7 +489,8 @@ async def _subsample_iteration( Yields a :class:`_SamplingResultSlice` per attempt and ends early after the first successful slice. `subsample_index` (0-based) identifies this subsample within the parent `sample()` call; it is used to derive a globally unique iteration counter for telemetry and hooks. `sampling_id` is passed through to the - iteration and repair hooks this subsample fires. + iteration and repair hooks this subsample fires. When `validation_ctx` is `None`, each attempt is + validated over its own post-generation context. """ flog = MelleaLogger.get_logger() sampled_results: list[ComputedModelOutputThunk[S]] = [] @@ -533,10 +533,13 @@ async def _subsample_iteration( else result.value # type: ignore[assignment] ) - # validation pass + # validation pass; validate over the caller's validation context if they + # supplied one, otherwise over this attempt's post-generation context. val_scores_co = mfuncs.avalidate( reqs=requirements, - context=result_ctx, + context=validation_ctx + if validation_ctx is not None + else result_ctx, backend=backend, output=result, format=None, diff --git a/mellea/stdlib/sampling/budget_forcing.py b/mellea/stdlib/sampling/budget_forcing.py index 2df57f40f8..0dd7cf5346 100644 --- a/mellea/stdlib/sampling/budget_forcing.py +++ b/mellea/stdlib/sampling/budget_forcing.py @@ -117,7 +117,8 @@ async def sample( context: The context to be passed to the sampling strategy. backend: The backend used for generating samples. requirements: List of requirements to test against (merged with global requirements). - validation_ctx: Optional context to use for validation. If None, validation_ctx = ctx. + validation_ctx: Optional context to validate over. If None, each sample is validated + over its own post-generation context. format: output format for structured outputs. model_options: model options to pass to the backend during generation / validation. tool_calls: True if tool calls should be used during this sampling strategy. @@ -129,8 +130,6 @@ async def sample( Raises: AssertionError: Asserts that all required components (repair, select_from_failure, validate, and generate) are provided before proceeding with the sampling. """ - validation_ctx = validation_ctx if validation_ctx is not None else context - flog = MelleaLogger.get_logger() with log_context(strategy=type(self).__name__, loop_budget=self.loop_budget): @@ -211,10 +210,13 @@ async def sample( else result.value ) - # validation pass + # validation pass; validate over the caller's validation context if they + # supplied one, otherwise over this attempt's post-generation context. val_scores_co = mfuncs.avalidate( reqs=reqs, - context=result_ctx, + context=validation_ctx + if validation_ctx is not None + else result_ctx, backend=backend, output=result, format=format, diff --git a/mellea/stdlib/sampling/majority_voting.py b/mellea/stdlib/sampling/majority_voting.py index 91fa2b6ca4..566150be3b 100644 --- a/mellea/stdlib/sampling/majority_voting.py +++ b/mellea/stdlib/sampling/majority_voting.py @@ -154,7 +154,8 @@ async def sample( context: The context to be passed to the sampling strategy. backend: The backend used for generating samples. requirements: List of requirements to test against (merged with global requirements). - validation_ctx: Optional context to use for validation. If None, validation_ctx = ctx. + validation_ctx: Optional context to validate over. If None, each sample is validated + over its own post-generation context. format: output format for structured outputs; ignored for this sampling strategy. model_options: model options to pass to the backend during generation / validation. tool_calls: True if tool calls should be used during this sampling strategy. diff --git a/mellea/stdlib/sampling/sofai.py b/mellea/stdlib/sampling/sofai.py index cb0e938215..710e773569 100644 --- a/mellea/stdlib/sampling/sofai.py +++ b/mellea/stdlib/sampling/sofai.py @@ -520,6 +520,7 @@ async def _generate_and_validate( format: type[BaseModelSubclass] | None, model_options: dict | None, tool_calls: bool, + validation_ctx: Context | None = None, ) -> tuple[ ComputedModelOutputThunk, Context, list[tuple[Requirement, ValidationResult]] ]: @@ -534,6 +535,8 @@ async def _generate_and_validate( format: Output format for structured outputs. model_options: Model options for generation. tool_calls: Whether to use tool calls. + validation_ctx: Optional context to validate over. If None, validation runs over + the post-generation context. Returns: Tuple of (result, result_context, validation_scores). @@ -560,7 +563,7 @@ async def _generate_and_validate( ) val_scores = await mfuncs.avalidate( reqs=reqs_for_validation, - context=result_ctx, + context=validation_ctx if validation_ctx is not None else result_ctx, backend=session_backend, output=computed_result, format=None, @@ -610,7 +613,8 @@ async def sample( context: The session context (must be ChatContext). backend: Session backend (used for validation fallback). requirements: Requirements to validate against. - validation_ctx: Optional separate validation context (unused). + validation_ctx: Optional context to validate over. If None, each sample is validated + over its own post-generation context. format: Output format for structured outputs. model_options: Model options to pass to backends. tool_calls: True if tool calls should be used. @@ -683,6 +687,7 @@ async def sample( format=format, model_options=model_options, tool_calls=tool_calls, + validation_ctx=validation_ctx, ) # Store attempt @@ -771,6 +776,7 @@ async def sample( format=format, model_options=model_options, tool_calls=tool_calls, + validation_ctx=validation_ctx, ) # Store S2 attempt diff --git a/mellea/stdlib/session.py b/mellea/stdlib/session.py index 303572aaeb..5f0503c17a 100644 --- a/mellea/stdlib/session.py +++ b/mellea/stdlib/session.py @@ -726,21 +726,22 @@ def validate( self, reqs: Requirement | list[Requirement], *, - output: CBlock | ModelOutputThunk | None = None, + output: ModelOutputThunk | None = None, format: type[BaseModelSubclass] | None = None, model_options: dict | None = None, generate_logs: list[GenerateLog] | None = None, input: CBlock | ModelOutputThunk | None = None, ) -> list[ValidationResult]: - """Validates a set of requirements over the output (if provided) or the current context (if the output is not provided). + """Validates a set of requirements over this session's context. Args: reqs: A single `Requirement` or a list of them to validate. - output: Optional model output to validate against instead of the context. + output: Optional model output designating the validation target. When `None`, + the context's last output is validated. format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Optional input to include alongside `output` when validating. + input: Deprecated. Optional input to prepend to the validation context. Returns: List of `ValidationResult` objects, one per requirement. @@ -1142,21 +1143,22 @@ async def avalidate( self, reqs: Requirement | list[Requirement], *, - output: CBlock | ModelOutputThunk | None = None, + output: ModelOutputThunk | None = None, format: type[BaseModelSubclass] | None = None, model_options: dict | None = None, generate_logs: list[GenerateLog] | None = None, input: CBlock | ModelOutputThunk | None = None, ) -> list[ValidationResult]: - """Validates a set of requirements over the output (if provided) or the current context (if the output is not provided). + """Validates a set of requirements over this session's context. Args: reqs: A single `Requirement` or a list of them to validate. - output: Optional model output to validate against instead of the context. + output: Optional model output designating the validation target. When `None`, + the context's last output is validated. format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Optional input to include alongside `output` when validating. + input: Deprecated. Optional input to prepend to the validation context. Returns: List of `ValidationResult` objects, one per requirement. diff --git a/mellea/templates/prompts/default/Requirement.jinja2 b/mellea/templates/prompts/default/Requirement.jinja2 index 21c895adfa..a39f2ae3e6 100644 --- a/mellea/templates/prompts/default/Requirement.jinja2 +++ b/mellea/templates/prompts/default/Requirement.jinja2 @@ -1,4 +1,5 @@ -Please check the following requirement against the following output. +Check whether the output below satisfies the requirement below. +The conversation is provided so that you can interpret the requirement; judge only the output. Reply with 'yes' if the requirement is satisfied and 'no' otherwise. Do not include any other text in your response. diff --git a/test/core/test_stream_validate.py b/test/core/test_stream_validate.py index 99a05e70bb..a4a7b8f2d5 100644 --- a/test/core/test_stream_validate.py +++ b/test/core/test_stream_validate.py @@ -66,13 +66,13 @@ async def stream_validate( async def test_does_not_mutate_requirement(): req = Requirement(description="original description") original_description = req.description - original_output = req._output + original_target = req._validation_target original_validation_fn = req.validation_fn await req.stream_validate("some chunk", backend=None, ctx=None) # type: ignore[arg-type] assert req.description == original_description - assert req._output == original_output + assert req._validation_target == original_target assert req.validation_fn == original_validation_fn @@ -83,7 +83,7 @@ async def test_stream_validate_idempotent(): result2 = await req.stream_validate("chunk two", backend=None, ctx=None) # type: ignore[arg-type] assert result1.success == "unknown" assert result2.success == "unknown" - assert req._output is None + assert req._validation_target is None @pytest.mark.asyncio diff --git a/test/stdlib/requirements/test_requirement.py b/test/stdlib/requirements/test_requirement.py index f8d4116896..449170a6bd 100644 --- a/test/stdlib/requirements/test_requirement.py +++ b/test/stdlib/requirements/test_requirement.py @@ -7,7 +7,9 @@ import pytest from mellea.backends.adapters import AdapterSchemaMismatchError -from mellea.core import ModelOutputThunk, Requirement +from mellea.core import ModelOutputThunk, Requirement, TemplateRepresentation +from mellea.formatters.template_formatter import TemplateFormatter +from mellea.stdlib.components import Message from mellea.stdlib.context import ChatContext from mellea.stdlib.requirements import LLMaJRequirement, simple_validate from mellea.stdlib.requirements.requirement import ( @@ -28,11 +30,11 @@ async def test_llmaj_validation_req_output_field(): m = start_session(ctx=ctx) req = Requirement("Must output test.") - assert req._output is None + assert req._validation_target is None _ = await req.validate(m.backend, ctx=ctx) - assert req._output is None, ( - "requirement's output shouldn't be updated during/after validation" + assert req._validation_target is None, ( + "requirement's validation target shouldn't be bound during/after validation" ) @@ -41,11 +43,11 @@ async def test_llmaj_validation_req_output_field(): async def test_llmaj_requirement_uses_requirement_template(): m = start_session(ctx=ctx) req = LLMaJRequirement("Must output test.") - assert req._output is None + assert req._validation_target is None _ = await req.validate(m.backend, ctx=ctx) - assert req._output is None, ( - "requirement's output shouldn't be updated during/after validation" + assert req._validation_target is None, ( + "requirement's validation target shouldn't be bound during/after validation" ) @@ -205,6 +207,117 @@ def test_simple_validate_none_output(): assert result.as_bool() is False +# --- validation target binding (issue #426) --- + + +def test_parts_empty_until_target_bound(): + r = Requirement("must mention Paris") + assert r.parts() == [], "an unbound requirement has no parts" + + +def test_bind_validation_target_leaves_original_unbound(): + """Binding happens on a copy so a requirement can be reused across validations.""" + r = Requirement("must mention Paris") + target = ModelOutputThunk("The capital of France is Paris.") + + bound = r._bind_validation_target(target) + + assert bound is not r + assert r._validation_target is None + assert bound._validation_target is target + assert bound.parts() == [target], ( + "the bound target must be exposed as a part so generate_walk can await it" + ) + + +def test_format_for_llm_carries_the_target_span(): + """The judge prompt gets the thunk itself, not a detached string copy of it.""" + r = Requirement("must mention Paris") + target = ModelOutputThunk("The capital of France is Paris.") + + representation = r._bind_validation_target(target).format_for_llm() + + assert isinstance(representation, TemplateRepresentation) + assert representation.args["output"] is target + assert representation.args["description"] == "must mention Paris" + + +def test_format_for_llm_prefers_component_parsed_repr(): + """A parsed `Component` repr beats the raw generated string.""" + r = Requirement("must be polite") + parsed = Message("assistant", "parsed message content") + target = ModelOutputThunk("raw string value") + target.parsed_repr = parsed + + representation = r._bind_validation_target(target).format_for_llm() + + assert isinstance(representation, TemplateRepresentation) + assert representation.args["output"] is parsed + + +def test_format_for_llm_without_bound_target_raises(): + with pytest.raises(AssertionError, match="Object protocol error"): + Requirement("must mention Paris").format_for_llm() + + +# --- judge prompt rendering --- + + +@pytest.fixture +def formatter(): + return TemplateFormatter(model_id="ibm-granite/granite-3.3-8b-instruct") + + +def test_requirement_prompt_inlines_the_output(formatter): + r = Requirement("must mention Paris") + target = ModelOutputThunk("The capital of France is Paris.") + + rendered = formatter.print(r._bind_validation_target(target)) + + assert "The capital of France is Paris." in rendered + assert "must mention Paris" in rendered + + +# --- the judgement request carries the conversation (issue #426, defect 4) --- + + +async def test_validate_hands_the_backend_a_context_that_renders_the_output(): + """The context reaching the backend must render the output under judgement. + + This is what makes adapter-backed requirement checking work at all: + `_generate_from_intrinsic` builds the `requirement-check` conversation from + `ctx.view_for_generation()`. Under the old throwaway `SimpleContext` that view was + empty, so the adapter was handed only the injected requirement-check message and + never saw the assistant turn it was supposed to judge. + """ + seen: list = [] + target = ModelOutputThunk("The capital of France is Paris.") + validation_ctx = ( + ChatContext().add(Message("user", "capital of France?")).add(target) + ) + + async def capture(action, ctx, **kwargs): + seen.append((action, ctx)) + # A thunk constructed with a value is already computed. + return ModelOutputThunk("yes"), ctx + + backend = MagicMock() + backend.generate_from_context = capture + + result = await Requirement("must mention Paris").validate(backend, validation_ctx) + + assert result.as_bool() is True + bound_req, passed_ctx = seen[0] + view = passed_ctx.view_for_generation() + assert view is not None and target in view, ( + "the judged output must be visible to the model, not merely present in as_list()" + ) + assert any("capital of France?" in str(node) for node in view), ( + "the conversation that produced the output must reach the judge too" + ) + assert bound_req._validation_target is target + + # --- LLMaJRequirement --- diff --git a/test/stdlib/sampling/test_sampling_base_unit.py b/test/stdlib/sampling/test_sampling_base_unit.py index c90c538a4f..93e189d295 100644 --- a/test/stdlib/sampling/test_sampling_base_unit.py +++ b/test/stdlib/sampling/test_sampling_base_unit.py @@ -258,6 +258,103 @@ async def test_multi_turn_strategy_with_concurrency(mocked_context_backend): ) +# --- validation_ctx is actually used (issues #426, #668) --- + + +def _ctx_capturing_requirement() -> tuple[Requirement, list[Context]]: + """A passing requirement that records every context it is validated over.""" + seen: list[Context] = [] + + def capture(ctx: Context) -> ValidationResult: + seen.append(ctx) + return ValidationResult(result=True) + + return Requirement("capture", validation_fn=capture), seen + + +async def test_explicit_validation_ctx_is_used_for_validation(mocked_context_backend): + """An explicit `validation_ctx` — not the generation context — is validated over. + + Regression test for #668: `validation_ctx` was accepted, defaulted, and threaded + through the strategy, then never passed to `avalidate`. + """ + req, seen = _ctx_capturing_requirement() + validation_ctx = ChatContext().add(Message("user", "VALIDATION MARKER")) + + await RejectionSamplingStrategy(loop_budget=1).sample( + action=Instruction(description="write something"), + context=ChatContext().add(Message("user", "GENERATION MARKER")), + backend=mocked_context_backend, + requirements=[req], + validation_ctx=validation_ctx, + ) + + assert len(seen) == 1 + rendered = [str(node) for node in seen[0].as_list()] + assert any("VALIDATION MARKER" in r for r in rendered), ( + "the caller's validation context must reach the requirement" + ) + assert not any("GENERATION MARKER" in r for r in rendered), ( + "the generation context must not leak into an explicit validation context" + ) + assert seen[0].last_output() is not None, ( + "the sampled output is appended so it is the validation target" + ) + + +async def test_validation_ctx_defaults_to_post_generation_context( + mocked_context_backend, +): + """Without `validation_ctx`, each attempt validates over its own post-generation context.""" + req, seen = _ctx_capturing_requirement() + + result = await RejectionSamplingStrategy(loop_budget=1).sample( + action=Instruction(description="write something"), + context=ChatContext().add(Message("user", "GENERATION MARKER")), + backend=mocked_context_backend, + requirements=[req], + ) + + assert len(seen) == 1 + # ComputedModelOutputThunk reassigns __class__ in place, so the sampled result *is* + # the thunk already in the post-generation context -- `avalidate` appends nothing. + assert seen[0] is result.sample_contexts[0], ( + "validation should run over the attempt's post-generation context unchanged" + ) + assert seen[0].last_output() is result.result, ( + "the validation target is the generation under judgement" + ) + rendered = [str(node) for node in seen[0].as_list()] + assert any("GENERATION MARKER" in r for r in rendered) + + +async def test_validation_ctx_sees_the_output_in_its_generation_view( + mocked_context_backend, +): + """The judged output must be visible to the model, not just present in `as_list()`. + + This is what makes adapter-backed requirement checking work: `_generate_from_intrinsic` + builds its conversation from `view_for_generation()`, which was empty under the old + throwaway `SimpleContext`. + """ + req, seen = _ctx_capturing_requirement() + + result = await RejectionSamplingStrategy(loop_budget=1).sample( + action=Instruction(description="write something"), + context=ChatContext(), + backend=mocked_context_backend, + requirements=[req], + ) + + view = seen[0].view_for_generation() + assert view is not None and len(view) > 0, ( + "the validation context must render something for the model" + ) + assert result.result in view, ( + "the generation under judgement must appear in the validation context's generation view" + ) + + # --- Constructor validation --- diff --git a/test/stdlib/test_functional_unit.py b/test/stdlib/test_functional_unit.py index e2afa346db..b5d2fb8eb2 100644 --- a/test/stdlib/test_functional_unit.py +++ b/test/stdlib/test_functional_unit.py @@ -13,7 +13,15 @@ import pytest from PIL import Image as PILImage -from mellea.core import AudioBlock, Context, ImageBlock, ModelToolCall +from mellea.core import ( + AudioBlock, + Context, + ImageBlock, + ModelOutputThunk, + ModelToolCall, + Requirement, + ValidationResult, +) from mellea.stdlib.components import ( Document, Instruction, @@ -27,6 +35,7 @@ aact, achat, ainstruct, + avalidate, chat, instruct, ) @@ -492,5 +501,88 @@ async def test_atransform_persists_chosen_tool_message_in_context( _assert_tool_message_persisted_after(ctx, new_ctx, [prior_message], tool_message) +# --- avalidate context handling (issue #426) --- + + +def _ctx_capturing_requirement() -> tuple[Requirement, list[Context]]: + """A passing requirement that records every context it is validated over.""" + seen: list[Context] = [] + + def capture(ctx: Context) -> ValidationResult: + seen.append(ctx) + return ValidationResult(result=True) + + return Requirement("capture", validation_fn=capture), seen + + +async def test_avalidate_validates_over_the_callers_context(): + """Validation runs over the caller's own context, not a throwaway `SimpleContext`. + + Before #426 was fixed, `avalidate` built a fresh `SimpleContext` whose + `view_for_generation()` is always empty, so nothing the caller passed ever reached + the model. + """ + req, seen = _ctx_capturing_requirement() + caller_ctx = ChatContext().add(Message("user", "hello")).add(ModelOutputThunk("hi")) + + await avalidate(reqs=[req], context=caller_ctx, backend=MagicMock()) + + assert seen == [caller_ctx], ( + "the caller's context object must be handed to the requirement untouched" + ) + assert isinstance(seen[0], ChatContext), ( + "validation must run in the caller's context type" + ) + + +async def test_avalidate_appends_output_when_not_already_last(): + """`output` designates the validation target and is appended when it is not last.""" + req, seen = _ctx_capturing_requirement() + caller_ctx = ChatContext().add(Message("user", "hello")) + target = ModelOutputThunk("the output under judgement") + + await avalidate(reqs=[req], context=caller_ctx, backend=MagicMock(), output=target) + + assert seen[0] is not caller_ctx + assert seen[0].last_output() is target + assert target in seen[0].view_for_generation() # type: ignore[operator] + assert any("hello" in str(node) for node in seen[0].as_list()), ( + "appending the target must not discard the caller's history" + ) + + +async def test_avalidate_does_not_double_add_the_context_last_output(): + """Passing the context's own last output as `output` is a no-op. + + This is the sampling path: `ComputedModelOutputThunk` reassigns `__class__` in + place, so the thunk handed to `avalidate` *is* the one already in the context. + """ + req, seen = _ctx_capturing_requirement() + target = ModelOutputThunk("already in the context") + caller_ctx = ChatContext().add(Message("user", "hello")).add(target) + + await avalidate(reqs=[req], context=caller_ctx, backend=MagicMock(), output=target) + + assert seen == [caller_ctx] + assert len(seen[0].as_list()) == 2, "the target must not be appended twice" + + +async def test_avalidate_input_is_deprecated(): + req, seen = _ctx_capturing_requirement() + caller_ctx = ChatContext().add(ModelOutputThunk("hi")) + + with pytest.warns(DeprecationWarning, match="`input` parameter"): + await avalidate( + reqs=[req], + context=caller_ctx, + backend=MagicMock(), + input=Message("user", "the deprecated input"), + ) + + assert any("the deprecated input" in str(node) for node in seen[0].as_list()), ( + "the deprecated `input` should still be honoured for one more release" + ) + + if __name__ == "__main__": pytest.main([__file__, "-v"]) From 91c81ec42a3ec12afb636fdd98c4f6ee353d10b5 Mon Sep 17 00:00:00 2001 From: Jake LoRocco Date: Tue, 1 Sep 2026 13:17:47 -0400 Subject: [PATCH 2/5] test: look into fixing validation without removing input by fixing simple context Signed-off-by: Jake LoRocco --- mellea/stdlib/functional.py | 37 ++++++++++++++++++----------- mellea/stdlib/session.py | 6 +++-- test/stdlib/test_functional_unit.py | 37 +++++++++++++++++++++++++---- 3 files changed, 59 insertions(+), 21 deletions(-) diff --git a/mellea/stdlib/functional.py b/mellea/stdlib/functional.py index 8654906e53..95708f5031 100644 --- a/mellea/stdlib/functional.py +++ b/mellea/stdlib/functional.py @@ -13,7 +13,6 @@ import asyncio import time import uuid -import warnings from collections.abc import Coroutine, Iterable from typing import Any, Literal, overload @@ -365,10 +364,9 @@ def validate( format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Deprecated. Optional input to prepend to the validation context. Pass a - context that already contains the input instead. Passing it raises a - `DeprecationWarning` from `avalidate`, so the reported source location is this - wrapper rather than the caller. + input: Optional input to prepend to the validation context, for judging an output + against a specific input rather than the whole conversation. See `avalidate` + for the visibility caveat on contexts that render no history. Returns: List of `ValidationResult` objects, one per requirement. @@ -1053,8 +1051,11 @@ async def avalidate( format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Deprecated. Optional input to prepend to the validation context. Pass a - context that already contains the input instead. + input: Optional input to prepend to the validation context, for judging an output + against a specific input rather than the whole conversation. It is added ahead + of `output`, so the judge sees the pair in order. It only reaches the model if + the context renders history: on a context whose `view_for_generation()` is empty + (`SimpleContext`), nothing added here is visible to the judge. Returns: List of `ValidationResult` objects, one per requirement. @@ -1068,13 +1069,6 @@ async def avalidate( validation_target_ctx = context if input is not None: - warnings.warn( - "The `input` parameter of validate/avalidate is deprecated and will be removed " - "in a future release. Pass a validation context that already contains the input " - "instead.", - DeprecationWarning, - stacklevel=2, - ) validation_target_ctx = validation_target_ctx.add(input) # `output` designates the validation target rather than replacing the context. Adding it @@ -1083,6 +1077,21 @@ async def avalidate( if output is not None and validation_target_ctx.last_output() is not output: validation_target_ctx = validation_target_ctx.add(output) + # A context that renders no history hands the judge nothing but the requirement and the + # inlined output, so anything conversational -- `input`, earlier turns, and the whole + # premise of adapter-backed requirement checking -- is silently dropped. Say so rather + # than letting the caller infer working validation from a plausible-looking verdict. + if ( + not validation_target_ctx.view_for_generation() + and validation_target_ctx.as_list() + ): + MelleaLogger.get_logger().warning( + f"validating over a {type(validation_target_ctx).__name__} whose" + " view_for_generation() is empty: the judge sees only the requirement and the" + " inlined output, not the conversation that produced it. Pass a context that" + " renders history (e.g. ChatContext) to validate against the input." + ) + # --- validation_pre_check hook --- if has_plugins(HookType.VALIDATION_PRE_CHECK): from ..plugins.hooks.validation import ValidationPreCheckPayload diff --git a/mellea/stdlib/session.py b/mellea/stdlib/session.py index 5f0503c17a..161ed7d076 100644 --- a/mellea/stdlib/session.py +++ b/mellea/stdlib/session.py @@ -741,7 +741,8 @@ def validate( format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Deprecated. Optional input to prepend to the validation context. + input: Optional input to prepend to the validation context, for judging an + output against a specific input rather than the whole conversation. Returns: List of `ValidationResult` objects, one per requirement. @@ -1158,7 +1159,8 @@ async def avalidate( format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Deprecated. Optional input to prepend to the validation context. + input: Optional input to prepend to the validation context, for judging an + output against a specific input rather than the whole conversation. Returns: List of `ValidationResult` objects, one per requirement. diff --git a/test/stdlib/test_functional_unit.py b/test/stdlib/test_functional_unit.py index b5d2fb8eb2..c9be215cbc 100644 --- a/test/stdlib/test_functional_unit.py +++ b/test/stdlib/test_functional_unit.py @@ -8,6 +8,7 @@ import base64 import io +import logging from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -15,6 +16,7 @@ from mellea.core import ( AudioBlock, + CBlock, Context, ImageBlock, ModelOutputThunk, @@ -567,21 +569,46 @@ async def test_avalidate_does_not_double_add_the_context_last_output(): assert len(seen[0].as_list()) == 2, "the target must not be appended twice" -async def test_avalidate_input_is_deprecated(): +async def test_avalidate_input_reaches_the_judges_view(): + """`input` is supported: it lands in the context the judge actually sees.""" req, seen = _ctx_capturing_requirement() caller_ctx = ChatContext().add(ModelOutputThunk("hi")) - with pytest.warns(DeprecationWarning, match="`input` parameter"): + await avalidate( + reqs=[req], + context=caller_ctx, + backend=MagicMock(), + input=CBlock("the specific input"), + ) + + view = seen[0].view_for_generation() + assert view is not None and any( + "the specific input" in str(node) for node in view + ), ( + "`input` must reach view_for_generation(), not merely as_list() -- only the" + " generation view is sent to the model" + ) + + +async def test_avalidate_warns_when_the_context_renders_nothing(caplog): + """A context with an empty generation view must not fail silently.""" + req, _seen = _ctx_capturing_requirement() + # SimpleContext retains as_list()/last_output() but renders no history. + caller_ctx = SimpleContext().add(ModelOutputThunk("hi")) + + with caplog.at_level(logging.WARNING): await avalidate( reqs=[req], context=caller_ctx, backend=MagicMock(), - input=Message("user", "the deprecated input"), + input=CBlock("invisible to the judge"), ) - assert any("the deprecated input" in str(node) for node in seen[0].as_list()), ( - "the deprecated `input` should still be honoured for one more release" + assert "view_for_generation() is empty" in caplog.text, ( + "validating over a context that renders nothing should warn, since `input` and the" + " conversation are dropped without any other signal" ) + assert "SimpleContext" in caplog.text, "the warning should name the context type" if __name__ == "__main__": From 26a33ad05977ef8760288df0f2a2c0b285653a13 Mon Sep 17 00:00:00 2001 From: Jake LoRocco Date: Thu, 3 Sep 2026 16:09:24 -0400 Subject: [PATCH 3/5] test: look into alora requirements to ensure working Signed-off-by: Jake LoRocco --- mellea/backends/huggingface.py | 29 +++++++ mellea/backends/openai.py | 29 +++++++ test/backends/test_openai_intrinsics_unit.py | 82 +++++++++++++++++++- 3 files changed, 139 insertions(+), 1 deletion(-) diff --git a/mellea/backends/huggingface.py b/mellea/backends/huggingface.py index e471cd423a..af0e901caf 100644 --- a/mellea/backends/huggingface.py +++ b/mellea/backends/huggingface.py @@ -505,6 +505,10 @@ async def _generate_from_context( tool_calls (bool): If `True`, expose available tools to the model and parse tool-call responses. + Raises: + ValueError: If `action` is an `ALoraRequirement` but `ctx` renders no + conversation, leaving the requirement-check adapter nothing to judge. + Returns: tuple[ModelOutputThunk[C], Context]: A thunk holding the (lazy) model output and an updated context that includes `action` and the new output. @@ -550,6 +554,31 @@ async def _generate_from_context( if issubclass(type(action), LLMaJRequirement): reroute_to_alora = False + # The requirement-check adapter judges the last assistant turn of the + # conversation it is given, and this path never renders the requirement + # template -- so unlike LLM-as-a-judge it has no inlined copy of the output + # to fall back on. A context that renders nothing leaves it nothing to judge. + if reroute_to_alora and not ctx.view_for_generation(): + if isinstance(action, ALoraRequirement): + raise ValueError( + f"cannot validate an ALoraRequirement over a " + f"{type(ctx).__name__}: it renders no conversation, so the " + f"{adapter_name} adapter has no assistant turn to judge. " + "Validate over a context that renders history (e.g. ChatContext), " + "or use a plain Requirement, which falls back to LLM-as-a-judge." + ) + # Auto-rerouting a plain Requirement is an optimisation, not a request: + # fall back to LLM-as-a-judge, which inlines the output and still works. + warn_key = f"alora_reroute_empty_view_{type(ctx).__name__}" + if warn_key not in self._warned_about: + self._warned_about.add(warn_key) + MelleaLogger.get_logger().warning( + f"not rerouting requirements to the {adapter_name} adapter: " + f"{type(ctx).__name__} renders no conversation for the adapter " + "to judge; using LLM-as-a-judge instead." + ) + reroute_to_alora = False + if reroute_to_alora: # Keep the alora requirement handling separate for now. mot = await self._generate_from_intrinsic( diff --git a/mellea/backends/openai.py b/mellea/backends/openai.py index 9d536f996f..0a2b35a225 100644 --- a/mellea/backends/openai.py +++ b/mellea/backends/openai.py @@ -591,6 +591,10 @@ async def _generate_from_context( tool_calls (bool): If `True`, expose available tools to the model and parse tool-call responses. + Raises: + ValueError: If `action` is an `ALoraRequirement` but `ctx` renders no + conversation, leaving the requirement-check adapter nothing to judge. + Returns: tuple[ModelOutputThunk[C], Context]: A thunk holding the (lazy) model output and an updated context that includes `action` and the new output. @@ -635,6 +639,31 @@ async def _generate_from_context( if issubclass(type(action), LLMaJRequirement): reroute_to_alora = False + # The requirement-check adapter judges the last assistant turn of the + # conversation it is given, and this path never renders the requirement + # template -- so unlike LLM-as-a-judge it has no inlined copy of the output + # to fall back on. A context that renders nothing leaves it nothing to judge. + if reroute_to_alora and not ctx.view_for_generation(): + if isinstance(action, ALoraRequirement): + raise ValueError( + f"cannot validate an ALoraRequirement over a " + f"{type(ctx).__name__}: it renders no conversation, so the " + f"{adapter_name} adapter has no assistant turn to judge. " + "Validate over a context that renders history (e.g. ChatContext), " + "or use a plain Requirement, which falls back to LLM-as-a-judge." + ) + # Auto-rerouting a plain Requirement is an optimisation, not a request: + # fall back to LLM-as-a-judge, which inlines the output and still works. + warn_key = f"alora_reroute_empty_view_{type(ctx).__name__}" + if warn_key not in self._warned_about: + self._warned_about.add(warn_key) + MelleaLogger.get_logger().warning( + f"not rerouting requirements to the {adapter_name} adapter: " + f"{type(ctx).__name__} renders no conversation for the adapter " + "to judge; using LLM-as-a-judge instead." + ) + reroute_to_alora = False + if reroute_to_alora: mot = await self._generate_from_intrinsic( alora_action, diff --git a/test/backends/test_openai_intrinsics_unit.py b/test/backends/test_openai_intrinsics_unit.py index 35a0934dd9..b4e1c6982f 100644 --- a/test/backends/test_openai_intrinsics_unit.py +++ b/test/backends/test_openai_intrinsics_unit.py @@ -13,6 +13,7 @@ """ import json +import logging import pathlib from copy import deepcopy from unittest.mock import AsyncMock, MagicMock, PropertyMock, patch @@ -30,9 +31,11 @@ from mellea.backends import ModelOption from mellea.backends.adapters.adapter import EmbeddedIntrinsicAdapter from mellea.backends.openai import OpenAIBackend +from mellea.core import ModelOutputThunk, Requirement from mellea.stdlib import functional as mfuncs from mellea.stdlib.components import Intrinsic, Message -from mellea.stdlib.context import ChatContext +from mellea.stdlib.context import ChatContext, SimpleContext +from mellea.stdlib.requirements import ALoraRequirement _TEST_DIR = pathlib.Path(__file__).parent _INTRINSICS_DATA = _TEST_DIR / "test_adapters" / "intrinsics-data" @@ -744,3 +747,80 @@ def get_temperature(location: str) -> int: assert tools is not None assert len(tools) == 1 assert tools[0]["function"]["name"] == "get_temperature" + + +# --------------------------------------------------------------------------- +# Requirement rerouting requires a context the adapter can actually judge +# --------------------------------------------------------------------------- + + +def _make_backend_with_requirement_adapter() -> OpenAIBackend: + """Return a backend whose `requirement-check` aLoRA would capture Requirements.""" + backend = OpenAIBackend( + model_id="granite-switch", + api_key="fake-key", + base_url="http://localhost:9999/v1", + ) + backend.add_adapter( + EmbeddedIntrinsicAdapter( + intrinsic_name="requirement-check", + config=deepcopy(_SIMPLE_CONFIG), + technology="alora", + ) + ) + return backend + + +async def test_alora_requirement_over_a_context_that_renders_nothing_raises(): + """An explicit `ALoraRequirement` cannot be honoured over an empty generation view. + + The `requirement-check` adapter judges the last assistant turn of the conversation it + is handed, and this path never renders the requirement template, so there is no inlined + output to fall back on. Asking for the adapter anyway is a caller error, not something + to silently downgrade. + """ + backend = _make_backend_with_requirement_adapter() + ctx = SimpleContext().add(ModelOutputThunk("The capital of France is Paris.")) + + with pytest.raises(ValueError, match="renders no conversation"): + await ALoraRequirement("must mention Paris").validate(backend, ctx) + + +async def test_plain_requirement_over_empty_view_falls_back_to_llmaj(caplog): + """Auto-rerouting is an optimisation, so a plain `Requirement` degrades instead. + + The reroute is mellea's choice rather than the caller's, so an unusable adapter must + not break validation: it falls through to LLM-as-a-judge, whose template inlines the + output and therefore still works over a context that renders nothing. + """ + backend = _make_backend_with_requirement_adapter() + ctx = SimpleContext().add(ModelOutputThunk("The capital of France is Paris.")) + + mock_create = AsyncMock(return_value=_simple_chat_completion("yes")) + mock_client = MagicMock() + mock_client.chat.completions.create = mock_create + + with ( + patch.object( + OpenAIBackend, + "_async_client", + new_callable=PropertyMock, + return_value=mock_client, + ), + patch.object( + OpenAIBackend, "_generate_from_intrinsic", new_callable=AsyncMock + ) as mock_intrinsic, + caplog.at_level(logging.WARNING), + ): + await Requirement("must mention Paris").validate(backend, ctx) + + mock_intrinsic.assert_not_called() + assert mock_create.await_count == 1, ( + "should fall through to ordinary chat generation" + ) + sent = mock_create.await_args.kwargs["messages"] + assert any("The capital of France is Paris." in str(m) for m in sent), ( + "the LLMaJ template must inline the output, which is what makes validation still" + " work over a context that renders no history" + ) + assert "not rerouting requirements" in caplog.text From bf02113838aac61c746b455b4a3c08c638910cc5 Mon Sep 17 00:00:00 2001 From: Jake LoRocco Date: Tue, 15 Sep 2026 13:37:23 -0400 Subject: [PATCH 4/5] fix: modify *validate docstrings for input Signed-off-by: Jake LoRocco --- mellea/stdlib/functional.py | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/mellea/stdlib/functional.py b/mellea/stdlib/functional.py index 95708f5031..5992d4237c 100644 --- a/mellea/stdlib/functional.py +++ b/mellea/stdlib/functional.py @@ -364,9 +364,10 @@ def validate( format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Optional input to prepend to the validation context, for judging an output - against a specific input rather than the whole conversation. See `avalidate` - for the visibility caveat on contexts that render no history. + input: Optional input to append to the validation context, for judging an output + against a specific input rather than the whole conversation. It is added + ahead of `output` if `output` is provided. See `avalidate` for the visibility + caveat on contexts that render no history. Returns: List of `ValidationResult` objects, one per requirement. @@ -1051,11 +1052,11 @@ async def avalidate( format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Optional input to prepend to the validation context, for judging an output + input: Optional input to append to the validation context, for judging an output against a specific input rather than the whole conversation. It is added ahead - of `output`, so the judge sees the pair in order. It only reaches the model if - the context renders history: on a context whose `view_for_generation()` is empty - (`SimpleContext`), nothing added here is visible to the judge. + of `output` if `output` is provided, so the judge sees the pair in order. It only reaches + the model if the context renders history: on a context whose `view_for_generation()` + is empty (`SimpleContext`), nothing added here is visible to the judge. Returns: List of `ValidationResult` objects, one per requirement. From fa68e8413806fe62fb24e87ce90139de493aec17 Mon Sep 17 00:00:00 2001 From: Jake LoRocco Date: Wed, 16 Sep 2026 15:13:04 -0400 Subject: [PATCH 5/5] fix: address pr comments Signed-off-by: Jake LoRocco --- docs/docs/advanced/lora-and-alora-adapters.md | 21 ++- docs/docs/concepts/requirements-system.md | 9 +- mellea/backends/huggingface.py | 22 ++- mellea/backends/openai.py | 22 ++- mellea/core/base.py | 14 +- mellea/stdlib/components/genstub.py | 14 +- mellea/stdlib/functional.py | 77 ++++++-- mellea/stdlib/requirements/requirement.py | 1 - mellea/stdlib/session.py | 4 +- test/core/test_base.py | 41 +++++ test/stdlib/components/test_genstub_unit.py | 44 ++++- test/stdlib/test_functional_unit.py | 171 ++++++++++++++++-- 12 files changed, 378 insertions(+), 62 deletions(-) diff --git a/docs/docs/advanced/lora-and-alora-adapters.md b/docs/docs/advanced/lora-and-alora-adapters.md index 648f9af9df..c3e3ee78b3 100644 --- a/docs/docs/advanced/lora-and-alora-adapters.md +++ b/docs/docs/advanced/lora-and-alora-adapters.md @@ -200,13 +200,22 @@ The `requirement-check` adapter judges the last assistant turn of the conversati given, so the routing above only produces useful verdicts if that conversation actually reaches it. Mellea builds the adapter's message list from the validation context's `view_for_generation()`, and validation runs over the post-generation context — the same -conversation the model generated into, with the generated output last. No extra setup is -needed for that; it is the default. +conversation the model generated into, with the generated output last. + +That only works on a context that renders history, which the default session does not +provide: `start_session()` returns a session backed by `SimpleContext`, which retains +`last_output()` but whose `view_for_generation()` is always empty by design. An +adapter-backed requirement has no conversation to judge there, so an explicit +`ALoraRequirement` raises `ValueError` and a plain `Requirement` logs a warning and falls +back to LLM-as-a-judge. Ask for a chat context explicitly: + +```python +import mellea + +# Adapter-backed requirements need a context that renders the assistant turn. +m = mellea.start_session(context_type="chat") +``` -One consequence is worth knowing: a context whose generation view is empty gives the -adapter nothing to judge. `SimpleContext` is the case to watch — it retains -`last_output()` but its `view_for_generation()` is always empty by design. Adapter-backed -requirements need a context that renders the assistant turn, such as `ChatContext`. See [What the validator sees](../concepts/requirements-system.md#what-the-validator-sees). ## Disable adapter validation diff --git a/docs/docs/concepts/requirements-system.md b/docs/docs/concepts/requirements-system.md index c8fee73344..48c7da0786 100644 --- a/docs/docs/concepts/requirements-system.md +++ b/docs/docs/concepts/requirements-system.md @@ -331,11 +331,10 @@ post-generation context is what you get unless you drive a strategy's `sample()` `validate()` and `avalidate()` follow the same rule at a lower level: they validate over the `context` you hand them, in your own context type. Their `output` argument designates -*which* output is under judgement rather than replacing the context — it is appended only -when it is not already the context's last output. - -> **Deprecated:** the `input` parameter of `validate()` / `avalidate()` is deprecated and -> will be removed in a future release. Pass a context that already contains the input. +*which* output is under judgement rather than replacing the context — it is appended unless +it is already the context's last entry. Passing an `output` from an earlier turn works too: +it is appended so that it becomes the target, and the judge still sees the conversation +around it. Preconditions are the one case that never sees the conversation: `precondition_requirements` are judged over the function arguments alone, in a fresh context. diff --git a/mellea/backends/huggingface.py b/mellea/backends/huggingface.py index 1cff71e28e..7b157301d0 100644 --- a/mellea/backends/huggingface.py +++ b/mellea/backends/huggingface.py @@ -698,24 +698,36 @@ async def _generate_from_context( # conversation it is given, and this path never renders the requirement # template -- so unlike LLM-as-a-judge it has no inlined copy of the output # to fall back on. A context that renders nothing leaves it nothing to judge. - if reroute_to_alora and not ctx.view_for_generation(): + adapter_view = ctx.view_for_generation() + if reroute_to_alora and not adapter_view: + # `None` means the history is non-linear and cannot be rendered at all; + # `[]` means it renders, but to nothing. Neither gives the adapter an + # assistant turn, so both are fatal here -- only the wording differs. + empty_view_reason = ( + "its history is non-linear, so no conversation can be rendered" + if adapter_view is None + else "it renders no conversation" + ) if isinstance(action, ALoraRequirement): raise ValueError( f"cannot validate an ALoraRequirement over a " - f"{type(ctx).__name__}: it renders no conversation, so the " + f"{type(ctx).__name__}: {empty_view_reason}, so the " f"{adapter_name} adapter has no assistant turn to judge. " "Validate over a context that renders history (e.g. ChatContext), " "or use a plain Requirement, which falls back to LLM-as-a-judge." ) # Auto-rerouting a plain Requirement is an optimisation, not a request: # fall back to LLM-as-a-judge, which inlines the output and still works. - warn_key = f"alora_reroute_empty_view_{type(ctx).__name__}" + warn_key = ( + f"alora_reroute_empty_view_{type(ctx).__name__}" + f"_{adapter_view is None}" + ) if warn_key not in self._warned_about: self._warned_about.add(warn_key) MelleaLogger.get_logger().warning( f"not rerouting requirements to the {adapter_name} adapter: " - f"{type(ctx).__name__} renders no conversation for the adapter " - "to judge; using LLM-as-a-judge instead." + f"{type(ctx).__name__} gives it nothing to judge -- " + f"{empty_view_reason}; using LLM-as-a-judge instead." ) reroute_to_alora = False diff --git a/mellea/backends/openai.py b/mellea/backends/openai.py index 85cd6178a4..b9546952cf 100644 --- a/mellea/backends/openai.py +++ b/mellea/backends/openai.py @@ -896,24 +896,36 @@ async def _generate_from_context( # conversation it is given, and this path never renders the requirement # template -- so unlike LLM-as-a-judge it has no inlined copy of the output # to fall back on. A context that renders nothing leaves it nothing to judge. - if reroute_to_alora and not ctx.view_for_generation(): + adapter_view = ctx.view_for_generation() + if reroute_to_alora and not adapter_view: + # `None` means the history is non-linear and cannot be rendered at all; + # `[]` means it renders, but to nothing. Neither gives the adapter an + # assistant turn, so both are fatal here -- only the wording differs. + empty_view_reason = ( + "its history is non-linear, so no conversation can be rendered" + if adapter_view is None + else "it renders no conversation" + ) if isinstance(action, ALoraRequirement): raise ValueError( f"cannot validate an ALoraRequirement over a " - f"{type(ctx).__name__}: it renders no conversation, so the " + f"{type(ctx).__name__}: {empty_view_reason}, so the " f"{adapter_name} adapter has no assistant turn to judge. " "Validate over a context that renders history (e.g. ChatContext), " "or use a plain Requirement, which falls back to LLM-as-a-judge." ) # Auto-rerouting a plain Requirement is an optimisation, not a request: # fall back to LLM-as-a-judge, which inlines the output and still works. - warn_key = f"alora_reroute_empty_view_{type(ctx).__name__}" + warn_key = ( + f"alora_reroute_empty_view_{type(ctx).__name__}" + f"_{adapter_view is None}" + ) if warn_key not in self._warned_about: self._warned_about.add(warn_key) MelleaLogger.get_logger().warning( f"not rerouting requirements to the {adapter_name} adapter: " - f"{type(ctx).__name__} renders no conversation for the adapter " - "to judge; using LLM-as-a-judge instead." + f"{type(ctx).__name__} gives it nothing to judge -- " + f"{empty_view_reason}; using LLM-as-a-judge instead." ) reroute_to_alora = False diff --git a/mellea/core/base.py b/mellea/core/base.py index fe668172f8..81347096cd 100644 --- a/mellea/core/base.py +++ b/mellea/core/base.py @@ -2180,6 +2180,11 @@ def as_list(self, last_n_components: int | None = None) -> list[Span]: If `last_n_components` is `None`, then all components are returned. + The same `Span` may legitimately appear more than once in the returned list: adding + an earlier output back onto a context to designate it as a validation target is the + motivating case. Only a repeated context *node* is a cycle, and that is what the + guard on this walk rejects. + Args: last_n_components (int | None): Maximum number of most-recent components to include. Pass `None` to return the full history. @@ -2189,16 +2194,19 @@ def as_list(self, last_n_components: int | None = None) -> list[Span]: """ context_list: list[Span] = [] current_context: Context = self + visited_nodes: set[int] = set() last_n_count = 0 while not current_context.is_root_node and ( last_n_components is None or last_n_count < last_n_components ): - data = current_context.node_data - assert data is not None, "Data cannot be None (except for root context)." - assert data not in context_list, ( + assert id(current_context) not in visited_nodes, ( "There might be a cycle in the context tree. That is not allowed." ) + visited_nodes.add(id(current_context)) + + data = current_context.node_data + assert data is not None, "Data cannot be None (except for root context)." context_list.append(data) last_n_count += 1 diff --git a/mellea/stdlib/components/genstub.py b/mellea/stdlib/components/genstub.py index c831d8f3b6..54807d0601 100644 --- a/mellea/stdlib/components/genstub.py +++ b/mellea/stdlib/components/genstub.py @@ -29,7 +29,7 @@ ValidationResult, ) from ...helpers.annotation_helpers import resolve_signature_annotations -from ..context import ChatContext +from ..context import SimpleContext from ..requirements.requirement import reqify from ..session import MelleaSession @@ -628,6 +628,10 @@ def __call__(self, *args, **kwargs) -> tuple[R, Context] | R: if stub_copy._arguments is not None: # Preconditions are judged over the arguments alone, so they get a fresh # context rather than the caller's/session's conversation history. + # `SimpleContext` renders no conversation, which is what keeps the arguments + # from reaching the judge twice: `ArgPreconditionRequirement`'s template + # already inlines them in its `Arguments:` block, and a context that rendered + # them would also emit them as an `assistant` turn. precondition_backend = ( extracted.m.backend if extracted.m is not None else extracted.backend ) @@ -635,7 +639,7 @@ def __call__(self, *args, **kwargs) -> tuple[R, Context] | R: assert precondition_backend is not None val_results = mfuncs.validate( reqs=stub_copy.precondition_requirements, - context=ChatContext(), + context=SimpleContext(), backend=precondition_backend, model_options=extracted.model_options, output=ModelOutputThunk(stub_copy._arguments.value), @@ -771,6 +775,10 @@ async def __async_call__() -> tuple[R, Context] | R: if stub_copy._arguments is not None: # Preconditions are judged over the arguments alone, so they get a fresh # context rather than the caller's/session's conversation history. + # `SimpleContext` renders no conversation, which is what keeps the + # arguments from reaching the judge twice: `ArgPreconditionRequirement`'s + # template already inlines them in its `Arguments:` block, and a context + # that rendered them would also emit them as an `assistant` turn. precondition_backend = ( extracted.m.backend if extracted.m is not None @@ -780,7 +788,7 @@ async def __async_call__() -> tuple[R, Context] | R: assert precondition_backend is not None val_results = await mfuncs.avalidate( reqs=stub_copy.precondition_requirements, - context=ChatContext(), + context=SimpleContext(), backend=precondition_backend, model_options=extracted.model_options, output=ModelOutputThunk(stub_copy._arguments.value), diff --git a/mellea/stdlib/functional.py b/mellea/stdlib/functional.py index 0d026ba213..36ef9cf97a 100644 --- a/mellea/stdlib/functional.py +++ b/mellea/stdlib/functional.py @@ -57,6 +57,11 @@ from .context import SimpleContext from .sampling import RejectionSamplingStrategy +# Context types whose empty generation view has already been warned about, so the warning +# in `avalidate` fires once per situation rather than on every call. Mirrors the +# `_warned_about` guards the backends keep for the same reason. +_validation_empty_view_warned: set[str] = set() + # Bound to Context so functions can return the same subtype they were given # (issue #1522): a `ChatContext` in yields a `ChatContext` out, statically. ContextT = TypeVar("ContextT", bound=Context) @@ -1639,8 +1644,10 @@ async def avalidate( Validation always runs over `context`, in the caller's own context type, so requirements (including adapter-backed ones) see the same conversation the model saw. `output` - designates *which* output is under judgement: it is appended to `context` when it is not - already the context's last output, and the requirement then validates that output. + designates *which* output is under judgement: it is appended to `context` unless it is + already the context's last entry, and the requirement then validates that output. An + `output` that appears earlier in `context` is appended too, so validating an output from + a previous turn works and the judge still sees the conversation around it. Args: reqs: A single `Requirement` or a list of them to validate. @@ -1671,26 +1678,46 @@ async def avalidate( if input is not None: validation_target_ctx = validation_target_ctx.add(input) - # `output` designates the validation target rather than replacing the context. Adding it - # is a no-op for the sampling path: ComputedModelOutputThunk reassigns __class__ in place, - # so the thunk passed here *is* the one already in the context. - if output is not None and validation_target_ctx.last_output() is not output: + # `output` designates the validation target rather than replacing the context, and + # `Requirement.validate` re-derives that target from `ctx.last_output()` -- so `output` + # has to end up last. It is added unless it is already the final entry, which is the + # sampling path: ComputedModelOutputThunk reassigns __class__ in place, so the thunk + # passed here *is* the one already at the tail, and re-adding it would only render the + # judged output twice. Anything else is added, including an `output` that appears + # earlier in the chain (validating an older output); `Context.as_list` tolerates a + # repeated span because its cycle guard tracks nodes rather than data. + if output is not None and validation_target_ctx.node_data is not output: validation_target_ctx = validation_target_ctx.add(output) # A context that renders no history hands the judge nothing but the requirement and the - # inlined output, so anything conversational -- `input`, earlier turns, and the whole - # premise of adapter-backed requirement checking -- is silently dropped. Say so rather - # than letting the caller infer working validation from a plausible-looking verdict. - if ( - not validation_target_ctx.view_for_generation() - and validation_target_ctx.as_list() - ): - MelleaLogger.get_logger().warning( - f"validating over a {type(validation_target_ctx).__name__} whose" - " view_for_generation() is empty: the judge sees only the requirement and the" - " inlined output, not the conversation that produced it. Pass a context that" - " renders history (e.g. ChatContext) to validate against the input." + # inlined output, and `input` is the only thing that costs. A plain LLM-as-a-judge + # requirement loses nothing (its template inlines the output either way), and an + # adapter-backed requirement is already reported by the backend, which knows whether the + # adapter was actually reached: openai/huggingface raise or fall back with a warning, + # and the other backends reject aLoRA outright. Warning here too would duplicate that, + # and would misfire in the case where the named adapter is absent -- the backend then + # never takes the adapter path at all and LLM-as-a-judge handles it correctly. + view = validation_target_ctx.view_for_generation() + if input is not None and not view and validation_target_ctx.as_list(): + ctx_type = type(validation_target_ctx).__name__ + # `view_for_generation()` returning None means the history is non-linear and + # cannot be rendered at all; [] means it renders, but to nothing. + detail = ( + f"view_for_generation() is None: {ctx_type} has a non-linear history, so no" + " conversation can be rendered" + if view is None + else f"view_for_generation() is empty: {ctx_type} renders no conversation" ) + # Deduped the way the backends' `_warned_about` guards are: without this, the + # warning fires on every single `avalidate` call. + warn_key = f"{ctx_type}:{view is None}" + if warn_key not in _validation_empty_view_warned: + _validation_empty_view_warned.add(warn_key) + MelleaLogger.get_logger().warning( + f"validating over a context whose {detail}, so `input` never reaches the" + " judge. Pass a context that renders history (e.g. ChatContext) to" + " validate against the input." + ) # --- validation_pre_check hook --- if has_plugins(HookType.VALIDATION_PRE_CHECK): @@ -1709,6 +1736,20 @@ async def avalidate( reqs = pre_payload.requirements model_options = pre_payload.model_options or model_options + # Compute the spans the requirements will read *before* fanning out. Every requirement + # binds the same validation target, `Requirement.parts()` exposes it, and each backend + # awaits the uncomputed leaves of its action -- so with two or more requirements the + # gather below would have two tasks calling `avalue()` on one uncomputed thunk at once. + # That deadlocks: `ModelOutputThunk.astream()` supports a single consumer, so one task + # takes the completion signal and the other waits forever for a chunk that never comes. + # Awaiting here, sequentially, means the fan-out only ever sees computed thunks. + if isinstance(input, ModelOutputThunk) and not input.is_computed(): + await input.avalue() + + validation_target = validation_target_ctx.last_output() + if validation_target is not None and not validation_target.is_computed(): + await validation_target.avalue() + rvs: list[ValidationResult] = [] coroutines: list[Coroutine[Any, Any, ValidationResult]] = [] diff --git a/mellea/stdlib/requirements/requirement.py b/mellea/stdlib/requirements/requirement.py index cd8b6d9acc..2e33fae678 100644 --- a/mellea/stdlib/requirements/requirement.py +++ b/mellea/stdlib/requirements/requirement.py @@ -13,7 +13,6 @@ MelleaLogger, ModelOutputThunk, Requirement, - TemplateRepresentation, ValidationResult, ) from ..components.intrinsic import Intrinsic diff --git a/mellea/stdlib/session.py b/mellea/stdlib/session.py index 3216a294b4..117257a9ad 100644 --- a/mellea/stdlib/session.py +++ b/mellea/stdlib/session.py @@ -1025,7 +1025,7 @@ def validate( format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Optional input to prepend to the validation context, for judging an + input: Optional input to append to the validation context, for judging an output against a specific input rather than the whole conversation. Returns: @@ -1466,7 +1466,7 @@ async def avalidate( format: Optional Pydantic model for constrained decoding. model_options: Additional model options to merge with backend defaults. generate_logs: Optional list to append generation logs to. - input: Optional input to prepend to the validation context, for judging an + input: Optional input to append to the validation context, for judging an output against a specific input rather than the whole conversation. Returns: diff --git a/test/core/test_base.py b/test/core/test_base.py index ff6842ba3a..a1f20dfa5c 100644 --- a/test/core/test_base.py +++ b/test/core/test_base.py @@ -1486,3 +1486,44 @@ def test_last_turn_tool_message_treated_as_input(): assert isinstance(turn.model_input, Message) assert turn.model_input.role == "tool" assert turn.output is None + + +# --- Context.as_list cycle guard (PR #1628 review) --- + + +def test_as_list_allows_the_same_span_twice(): + """A span may appear more than once in a chain without tripping the cycle guard. + + Re-adding an earlier output to designate it as a validation target is the motivating + case (`avalidate(..., output=an_earlier_output)`). The guard used to compare the data + each node carries, so this raised `AssertionError: There might be a cycle`. + """ + from mellea.stdlib.context import ChatContext + + repeated = ModelOutputThunk("first output") + ctx = ( + ChatContext() + .add(Message("user", "q1")) + .add(repeated) + .add(Message("user", "q2")) + .add(ModelOutputThunk("second output")) + .add(repeated) + ) + + spans = ctx.as_list() + + assert [s for s in spans if s is repeated] == [repeated, repeated], ( + "both occurrences of the repeated span must survive the walk" + ) + + +def test_as_list_still_detects_a_real_node_cycle(): + """A chain that revisits a *node* is a real cycle and must still be rejected.""" + from mellea.stdlib.context import ChatContext + + ctx = ChatContext().add(Message("user", "q1")).add(Message("user", "q2")) + # Splice the chain back onto itself: walking it would otherwise never terminate. + ctx.previous_node._previous = ctx # type: ignore[union-attr] + + with pytest.raises(AssertionError, match="cycle in the context tree"): + ctx.as_list() diff --git a/test/stdlib/components/test_genstub_unit.py b/test/stdlib/components/test_genstub_unit.py index ffd44c1eed..6ae36b51a0 100644 --- a/test/stdlib/components/test_genstub_unit.py +++ b/test/stdlib/components/test_genstub_unit.py @@ -9,11 +9,12 @@ import inspect from typing import Any, Literal, get_type_hints +from unittest.mock import MagicMock import pytest from mellea import generative -from mellea.core import TemplateRepresentation, ValidationResult +from mellea.core import Requirement, TemplateRepresentation, ValidationResult from mellea.stdlib.components.genstub import ( ArgPreconditionRequirement, Arguments, @@ -26,6 +27,7 @@ describe_function, get_argument, ) +from mellea.stdlib.context import SimpleContext from mellea.stdlib.requirements.requirement import reqify from test.stdlib.components._postponed_annotation_samples import ( extract_requirements, @@ -393,6 +395,46 @@ def test_arg_precondition_deepcopy(): assert cloned.description == req.description +# --- precondition validation context (PR #1628 review) --- + + +def test_precondition_validation_renders_no_conversation(): + """The arguments must reach the precondition judge exactly once. + + `ArgPreconditionRequirement.jinja2` already inlines them in its `Arguments:` block, so + validating over a context that renders history sends them a second time as an + `assistant` turn (any `ModelOutputThunk` maps to `assistant`). The precondition context + therefore has to be one that renders nothing. + """ + seen: list[Any] = [] + + def capture_and_fail(ctx): + seen.append(ctx) + return ValidationResult(result=False, reason="forced failure") + + @generative + def classify(text: str) -> str: ... + + with pytest.raises(PreconditionException): + classify( + context=SimpleContext(), + backend=MagicMock(), + text="hello", + precondition_requirements=[ + Requirement("forced failure", validation_fn=capture_and_fail) + ], + ) + + assert seen, "the precondition requirement must actually have been validated" + assert seen[0].view_for_generation() == [], ( + "the precondition context must render no conversation, or the arguments reach the" + " judge twice" + ) + assert "hello" in str(seen[0].last_output()), ( + "the arguments must still be the validation target, so the template inlines them" + ) + + # --- PreconditionException --- diff --git a/test/stdlib/test_functional_unit.py b/test/stdlib/test_functional_unit.py index c9be215cbc..0093ddaa81 100644 --- a/test/stdlib/test_functional_unit.py +++ b/test/stdlib/test_functional_unit.py @@ -6,6 +6,7 @@ Covers image preprocessing plus chat()/instruct() forwarding of multimodal inputs. """ +import asyncio import base64 import io import logging @@ -506,6 +507,16 @@ async def test_atransform_persists_chosen_tool_message_in_context( # --- avalidate context handling (issue #426) --- +@pytest.fixture +def reset_empty_view_warnings(): + """Clear `avalidate`'s dedup set so warning assertions do not depend on test order.""" + from mellea.stdlib import functional + + functional._validation_empty_view_warned.clear() + yield + functional._validation_empty_view_warned.clear() + + def _ctx_capturing_requirement() -> tuple[Requirement, list[Context]]: """A passing requirement that records every context it is validated over.""" seen: list[Context] = [] @@ -590,25 +601,159 @@ async def test_avalidate_input_reaches_the_judges_view(): ) -async def test_avalidate_warns_when_the_context_renders_nothing(caplog): - """A context with an empty generation view must not fail silently.""" +async def test_avalidate_accepts_an_output_from_an_earlier_turn(): + """`output=` may name an output further back than `last_output()` searches. + + `last_output()` only looks at the last 3 components, so an older output was reported as + "not last", appended a second time, and the whole call died on the cycle assertion in + `Context.as_list`. It is now appended so that it becomes the validation target, with the + conversation around it intact. + """ + req, seen = _ctx_capturing_requirement() + older = ModelOutputThunk("the output under judgement") + ctx = ( + ChatContext() + .add(Message("user", "q1")) + .add(older) + .add(Message("user", "q2")) + .add(ModelOutputThunk("a later output")) + ) + + await avalidate(reqs=[req], context=ctx, backend=MagicMock(), output=older) + + assert seen[0].last_output() is older, ( + "the requested output must be the target, since Requirement.validate re-derives it" + " from ctx.last_output()" + ) + assert any("q1" in str(node) for node in seen[0].as_list()), ( + "appending an older target must not discard the conversation around it" + ) + + +async def test_avalidate_computes_the_target_once_before_fanning_out(): + """Two requirements over an uncomputed thunk must not deadlock on `astream()`. + + `Requirement.parts()` exposes the bound target, so each requirement's backend call + awaits it — and two `astream()` consumers on one thunk hang forever. `avalidate` now + computes the target before the fan-out. The timeout is the assertion: on a regression + this test hangs rather than fails. + """ + from mellea.core import Backend + from mellea.core.base import GenerateType + + async def _process(mot: ModelOutputThunk, chunk) -> None: + if mot._underlying_value is None: + mot._underlying_value = "" + if chunk is not None: + mot._underlying_value += chunk + + async def _post_process(mot: ModelOutputThunk) -> None: + pass + + target = ModelOutputThunk(value=None) + target._call.action = CBlock("action") + target._gen.generate_type = GenerateType.ASYNC + target._gen.process = _process + target._gen.post_process = _post_process + target._gen.chunk_size = 0 + + async def produce_chunks_over_time(): + # Spacing the chunks is what exposes the race: a pre-filled queue drains before the + # second validation task ever starts. + for chunk in ("the ", "generated ", "answer"): + await asyncio.sleep(0.01) + target._gen.queue.put_nowait(chunk) + await asyncio.sleep(0.01) + target._gen.queue.put_nowait(None) + + class AwaitsItsActionBackend(Backend): + """Test double that awaits its action's uncomputed parts, as real backends do.""" + + _model_id = "test-model" + _provider = "test" + + async def _generate_from_context( + self, action, ctx, *, format=None, model_options=None, tool_calls=False + ): + await self.do_generate_walk(action) + return ModelOutputThunk("yes"), ctx + + async def _generate_from_raw(self, *args, **kwargs): + raise NotImplementedError + + ctx = ChatContext().add(Message("user", "q")).add(target) + producer = asyncio.create_task(produce_chunks_over_time()) + try: + results = await asyncio.wait_for( + avalidate( + reqs=[Requirement("first"), Requirement("second")], + context=ctx, + backend=AwaitsItsActionBackend(), + ), + timeout=10, + ) + finally: + producer.cancel() + + assert len(results) == 2, "both requirements must produce a verdict" + + +async def test_avalidate_computes_an_uncomputed_input_thunk(): + """An `input` that is still streaming is computed too, for the same reason.""" + from mellea.core.base import GenerateType + + async def _process(mot: ModelOutputThunk, chunk) -> None: + if mot._underlying_value is None: + mot._underlying_value = "" + if chunk is not None: + mot._underlying_value += chunk + + async def _post_process(mot: ModelOutputThunk) -> None: + pass + + streaming_input = ModelOutputThunk(value=None) + streaming_input._call.action = CBlock("action") + streaming_input._gen.generate_type = GenerateType.ASYNC + streaming_input._gen.process = _process + streaming_input._gen.post_process = _post_process + streaming_input._gen.chunk_size = 0 + streaming_input._gen.queue.put_nowait("the specific input") + streaming_input._gen.queue.put_nowait(None) + req, _seen = _ctx_capturing_requirement() - # SimpleContext retains as_list()/last_output() but renders no history. + ctx = ChatContext().add(ModelOutputThunk("hi")) + + await avalidate(reqs=[req], context=ctx, backend=MagicMock(), input=streaming_input) + + assert streaming_input.is_computed(), ( + "an uncomputed input must be resolved before the judge renders the context" + ) + + +async def test_avalidate_does_not_warn_for_an_adapter_backed_requirement( + caplog, reset_empty_view_warnings +): + """The empty-view warning is left to the backend for adapter-backed requirements. + + The backend is the only layer that knows whether the adapter was actually reached: + openai/huggingface raise for an explicit `ALoraRequirement` or fall back with their own + warning, and the remaining backends reject aLoRA outright. Warning here as well would + duplicate that, and would misfire when the named adapter is simply absent. + """ + from mellea.stdlib.requirements import ALoraRequirement + + # ALoraRequirement pins validation_fn to None; attaching one keeps this a unit test by + # taking the LLM-as-a-judge branch out of Requirement.validate. + req = ALoraRequirement("must be polite") + req.validation_fn = lambda ctx: ValidationResult(result=True) caller_ctx = SimpleContext().add(ModelOutputThunk("hi")) with caplog.at_level(logging.WARNING): - await avalidate( - reqs=[req], - context=caller_ctx, - backend=MagicMock(), - input=CBlock("invisible to the judge"), - ) + await avalidate(reqs=[req], context=caller_ctx, backend=MagicMock()) - assert "view_for_generation() is empty" in caplog.text, ( - "validating over a context that renders nothing should warn, since `input` and the" - " conversation are dropped without any other signal" + assert "view_for_generation()" not in caplog.text, ( + "avalidate must not pre-empt the backend's own adapter reporting" ) - assert "SimpleContext" in caplog.text, "the warning should name the context type" if __name__ == "__main__":