diff --git a/docs/docs/advanced/lora-and-alora-adapters.md b/docs/docs/advanced/lora-and-alora-adapters.md index c2f27f0d45..c3e3ee78b3 100644 --- a/docs/docs/advanced/lora-and-alora-adapters.md +++ b/docs/docs/advanced/lora-and-alora-adapters.md @@ -194,6 +194,30 @@ 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. + +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") +``` + +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 1fc060342c..48c7da0786 100644 --- a/docs/docs/concepts/requirements-system.md +++ b/docs/docs/concepts/requirements-system.md @@ -297,6 +297,48 @@ 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.md). +## 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 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. + ## 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 724ab78850..61cbf45fcf 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/backends/huggingface.py b/mellea/backends/huggingface.py index a756fefcb6..7b157301d0 100644 --- a/mellea/backends/huggingface.py +++ b/mellea/backends/huggingface.py @@ -635,6 +635,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. @@ -690,6 +694,43 @@ 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. + 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__}: {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__}" + 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__} gives it nothing to judge -- " + f"{empty_view_reason}; 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 817c31d032..b9546952cf 100644 --- a/mellea/backends/openai.py +++ b/mellea/backends/openai.py @@ -833,6 +833,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. @@ -888,6 +892,43 @@ 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. + 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__}: {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__}" + 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__} gives it nothing to judge -- " + f"{empty_view_reason}; 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/mellea/core/base.py b/mellea/core/base.py index 01276609b1..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 @@ -2386,7 +2394,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 0b0cf28962..35227e537d 100644 --- a/mellea/core/requirement.py +++ b/mellea/core/requirement.py @@ -329,12 +329,29 @@ def __init__( self.check_only = check_only self.chunking: ChunkingStrategy | None = resolve_chunking_strategy(chunking) - # 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 # Per-stream chunker for streaming validation, built lazily. self._chunker: Chunker | 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 + def __copy__(self) -> "Requirement": """Return a shallow copy with the live `_chunker` reset to `None`. @@ -391,10 +408,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 ) @@ -562,29 +579,43 @@ async def stream_flush( 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..54807d0601 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 SimpleContext from ..requirements.requirement import reqify from ..session import MelleaSession @@ -625,23 +626,24 @@ 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. + # `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 + ) + # 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=SimpleContext(), + 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 +773,26 @@ 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. + # `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 + ) + # 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=SimpleContext(), + 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 b48ad9ae27..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) @@ -529,24 +534,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: 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. @@ -1622,23 +1634,35 @@ 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` 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. - 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: 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, 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. @@ -1649,16 +1673,52 @@ 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 - # Add the input/output to the validation context - if input is not None: - validation_target_ctx = validation_target_ctx.add(input) + if input is not None: + validation_target_ctx = validation_target_ctx.add(input) + + # `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, 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): from ..plugins.hooks.validation import ValidationPreCheckPayload @@ -1676,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/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 fe21f17b4e..117257a9ad 100644 --- a/mellea/stdlib/session.py +++ b/mellea/stdlib/session.py @@ -1010,21 +1010,23 @@ 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: Optional input to append 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. @@ -1449,21 +1451,23 @@ 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: Optional input to append 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/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/backends/test_openai_intrinsics_unit.py b/test/backends/test_openai_intrinsics_unit.py index 6ac81f15df..c31fe44f8b 100644 --- a/test/backends/test_openai_intrinsics_unit.py +++ b/test/backends/test_openai_intrinsics_unit.py @@ -15,6 +15,7 @@ import asyncio import base64 import json +import logging import pathlib from copy import deepcopy from unittest.mock import AsyncMock, MagicMock, PropertyMock, patch @@ -32,9 +33,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" @@ -852,6 +855,83 @@ def get_temperature(location: str) -> int: 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 + + _WAV_B64 = base64.b64encode( b"RIFF$\x00\x00\x00WAVEfmt " b"\x10\x00\x00\x00\x01\x00\x01\x00" 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/core/test_stream_validate.py b/test/core/test_stream_validate.py index 92bdcbb039..33b0b4e746 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[0].success == "unknown" assert result2[0].success == "unknown" - assert req._output is None + assert req._validation_target is None @pytest.mark.asyncio 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/requirements/test_requirement.py b/test/stdlib/requirements/test_requirement.py index 3e9cca8753..1b702f6db9 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, AdapterType -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..0093ddaa81 100644 --- a/test/stdlib/test_functional_unit.py +++ b/test/stdlib/test_functional_unit.py @@ -6,14 +6,25 @@ Covers image preprocessing plus chat()/instruct() forwarding of multimodal inputs. """ +import asyncio import base64 import io +import logging from unittest.mock import AsyncMock, MagicMock, patch import pytest from PIL import Image as PILImage -from mellea.core import AudioBlock, Context, ImageBlock, ModelToolCall +from mellea.core import ( + AudioBlock, + CBlock, + Context, + ImageBlock, + ModelOutputThunk, + ModelToolCall, + Requirement, + ValidationResult, +) from mellea.stdlib.components import ( Document, Instruction, @@ -27,6 +38,7 @@ aact, achat, ainstruct, + avalidate, chat, instruct, ) @@ -492,5 +504,257 @@ 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) --- + + +@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] = [] + + 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_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")) + + 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_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() + 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()) + + assert "view_for_generation()" not in caplog.text, ( + "avalidate must not pre-empt the backend's own adapter reporting" + ) + + if __name__ == "__main__": pytest.main([__file__, "-v"])