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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
127 changes: 57 additions & 70 deletions protest/evals/wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
from typing import Annotated, Any, get_args, get_origin

from protest.di.hints import get_type_hints_compat
from protest.di.markers import From
from protest.entities.events import EvalPayload, EvalScoreEntry
from protest.evals.evaluator import (
EvalCase,
Expand All @@ -39,18 +40,23 @@ def make_eval_wrapper(
) -> Any:
"""Wrap a function to run evaluators on its return value."""

_validate_single_evalcase_param(func)
# Resolve which parameter carries the EvalCase once, statically. The name
# is then used for a direct kwargs lookup below - the case is never
# rediscovered by scanning values.
case_param = _resolve_case_param(func)
validate_evaluators(evaluators)

@functools.wraps(func)
async def eval_wrapper(**kwargs: Any) -> EvalPayload:
expected = _extract_expected(kwargs)
case_name = _extract_case_name(kwargs, func.__name__)
inputs = _extract_inputs(kwargs)
metadata = _extract_metadata(kwargs)
case = kwargs.get(case_param) if case_param is not None else None
expected = case.expected if case is not None else None
case_name = case.name if case is not None else func.__name__
inputs = case.inputs if case is not None else None
metadata = (case.metadata or None) if case is not None else None

all_evaluators = list(evaluators)
all_evaluators.extend(_extract_per_case_evaluators(kwargs))
if case is not None and case.evaluators:
all_evaluators.extend(case.evaluators)

# Both guards run before the task itself: the evaluator list is
# fully known from the kwargs alone, and the task is typically the
Expand Down Expand Up @@ -140,35 +146,59 @@ async def eval_wrapper(**kwargs: Any) -> EvalPayload:
# ---------------------------------------------------------------------------


def _validate_single_evalcase_param(func: Any) -> None:
"""Raise MultipleEvalCaseParamsError if `func` has > 1 EvalCase parameter.
def _resolve_case_param(func: Any) -> str | None:
"""Return the name of the single parameter that carries the EvalCase.

Runs at decorator time. The runtime contract (`_find_case`) silently
picks the first EvalCase in kwargs, which would drop the second one's
name/expected/inputs/metadata/per-case evaluators downstream. We catch
that here so the failure is loud and pinpoints the offending eval.
Resolved once at decoration time, from the signature alone. A parameter
is the case parameter when either signal holds:

Subclasses of EvalCase count: the runtime uses isinstance(_, EvalCase),
so any subclass would trigger the same silent drop.
- its declared type is an EvalCase (subclass) - ``case: EvalCase`` or
``Annotated[EvalCase, From(cases)]``; or
- it is bound via ``From(source)`` whose source yields EvalCase instances -
``Annotated[Any, From(cases)]``, where the type is deliberately loose.

The returned name drives a direct ``kwargs[name]`` lookup at runtime. This
is the whole point: the case is identified by *which parameter it is*, not
by scanning kwargs values for an EvalCase instance. A fixture that merely
returns an EvalCase on an unrelated parameter therefore cannot be mistaken
for the case (the silent misattribution the old isinstance scan allowed).

Raises MultipleEvalCaseParamsError if more than one parameter qualifies -
only one case per eval defines its identity (name, expected, inputs,
metadata, per-case evaluators). Returns None when no parameter qualifies
(a static eval, or one parametrized over non-EvalCase values).
"""
hints = get_type_hints_compat(func)
offending: list[str] = []
for param_name, annotation in hints.items():
if param_name == "return":
continue
underlying = (
get_args(annotation)[0]
if get_origin(annotation) is Annotated
else annotation
)
matches = [
name
for name, annotation in hints.items()
if name != "return" and _is_case_param(annotation)
]
if len(matches) > 1:
raise MultipleEvalCaseParamsError(func.__name__, matches)
return matches[0] if matches else None


def _is_case_param(annotation: Any) -> bool:
if get_origin(annotation) is Annotated:
args = get_args(annotation)
underlying = args[0]
if isinstance(underlying, type) and issubclass(underlying, EvalCase):
offending.append(param_name)
if len(offending) > 1:
raise MultipleEvalCaseParamsError(func.__name__, offending)
return True
return any(
isinstance(meta, From) and _source_yields_evalcase(meta.source)
for meta in args[1:]
)
return isinstance(annotation, type) and issubclass(annotation, EvalCase)


def _source_yields_evalcase(source: Any) -> bool:
"""True if a From source yields EvalCase instances (ForEach is non-empty)."""
return isinstance(next(iter(source), None), EvalCase)


# ---------------------------------------------------------------------------
# Extract helpers - pull EvalCase from kwargs
# Evaluator list helpers
# ---------------------------------------------------------------------------


Expand Down Expand Up @@ -204,49 +234,6 @@ def _flatten_evaluators(
return flat


def _find_case(kwargs: dict[str, Any]) -> EvalCase | None:
"""Find the EvalCase instance in kwargs."""
for v in kwargs.values():
if isinstance(v, EvalCase):
return v
return None


def _extract_expected(kwargs: dict[str, Any]) -> Any:
case = _find_case(kwargs)
if case is None:
return None
return case.expected


def _extract_case_name(kwargs: dict[str, Any], fallback: str) -> str:
case = _find_case(kwargs)
if case is None:
return fallback
return case.name


def _extract_inputs(kwargs: dict[str, Any]) -> Any:
case = _find_case(kwargs)
if case is None:
return None
return case.inputs


def _extract_metadata(kwargs: dict[str, Any]) -> Any:
case = _find_case(kwargs)
if case is None:
return None
return case.metadata or None


def _extract_per_case_evaluators(kwargs: dict[str, Any]) -> list[Any]:
case = _find_case(kwargs)
if case is None or not case.evaluators:
return []
return list(case.evaluators)


# ---------------------------------------------------------------------------
# Evaluator execution
# ---------------------------------------------------------------------------
Expand Down
100 changes: 94 additions & 6 deletions tests/evals/test_multiple_evalcase_params.py
Original file line number Diff line number Diff line change
@@ -1,24 +1,33 @@
"""Tests for `_validate_single_evalcase_param` - D1 registration-time check.
"""Tests for case-parameter resolution (`_resolve_case_param`).

The runtime contract (`_find_case`) picks the first `EvalCase` in kwargs and
silently drops any others. The wrapper detects > 1 EvalCase param at
registration and raises a clear error pointing at the offending parameters.
The wrapper identifies the EvalCase by *which parameter carries it*, resolved
once from the signature at decoration time, not by scanning kwargs values at
runtime. This file covers:

- single case parameter accepted (typed, subclass, or loose-typed via `From`);
- more than one case parameter rejected loudly at registration;
- a parameter that merely *holds* an EvalCase at runtime, without being the
declared case parameter, is ignored - not misattributed (issue #120).
"""

from __future__ import annotations

from typing import Annotated
import asyncio
from typing import Annotated, Any

import pytest

from protest import ForEach, From, ProTestSession
from protest.evals import EvalCase
from protest.evals import EvalCase, EvalContext, evaluator
from protest.evals.suite import EvalSuite
from protest.evals.wrapper import make_eval_wrapper
from protest.exceptions import MultipleEvalCaseParamsError

# Module-level case sources so `get_type_hints()` can resolve Annotated args.
_cases_a = ForEach([EvalCase(inputs="a", name="a1")])
_cases_b = ForEach([EvalCase(inputs="b", name="b1")])
_loose_cases = ForEach([EvalCase(inputs="loose", name="loose1")])
_dict_cases = ForEach([{"inputs": "d"}])


class _MyCase(EvalCase):
Expand All @@ -28,6 +37,11 @@ class _MyCase(EvalCase):
_subclass_cases = ForEach([_MyCase(inputs="x", name="x1")])


@evaluator
def _ok(ctx: EvalContext) -> bool:
return True


class TestSingleEvalCaseParamAccepted:
def test_one_evalcase_param_via_annotated_from(self) -> None:
session = ProTestSession()
Expand Down Expand Up @@ -63,6 +77,32 @@ def good(case: Annotated[_MyCase, From(_subclass_cases)]) -> str:
_ = good
session.add_suite(suite)

def test_loose_typed_from_source_is_recognized_as_case(self) -> None:
"""`Annotated[Any, From(cases)]` over EvalCase items is the case param.

The type is deliberately loose; resolution keys off the From source
yielding EvalCase instances, so case identity is still wired up.
"""

def task(case: Annotated[Any, From(_loose_cases)]) -> str:
return str(case.inputs)

wrapped = make_eval_wrapper(task, [_ok])
payload = asyncio.run(wrapped(case=EvalCase(inputs="z", name="z1")))
assert payload.case_name == "z1"

def test_from_source_of_non_evalcase_is_not_a_case_param(self) -> None:
"""A `From` over plain dicts is parametrization, not a case: the eval
falls back to the function name and carries no expected/inputs."""

def task(case: Annotated[dict, From(_dict_cases)]) -> str:
return str(case["inputs"])

wrapped = make_eval_wrapper(task, [_ok])
payload = asyncio.run(wrapped(case={"inputs": "d"}))
assert payload.case_name == "task"
assert payload.expected_output is None


class TestMultipleEvalCaseParamRejected:
def test_two_evalcase_params_raise(self) -> None:
Expand Down Expand Up @@ -97,3 +137,51 @@ def bad(

assert "case_a" in str(excinfo.value)
assert "case_b" in str(excinfo.value)

def test_loose_from_collides_with_typed_case(self) -> None:
"""A loose `From(EvalCase items)` param and a typed EvalCase param both
qualify - two case params, rejected at registration."""
suite = EvalSuite("evals")

with pytest.raises(MultipleEvalCaseParamsError) as excinfo:

@suite.eval()
def bad(
case_a: Annotated[EvalCase, From(_cases_a)],
case_b: Annotated[Any, From(_loose_cases)],
) -> str:
return str(case_a.inputs)

assert "case_a" in str(excinfo.value)
assert "case_b" in str(excinfo.value)


class TestRuntimeEvalCaseOnUnrelatedParamIgnored:
"""Issue #120, resolved structurally. The case is identified by parameter
name (resolved at decoration), not by scanning values. A fixture that
returns an EvalCase on an unrelated parameter is therefore simply ignored -
it cannot shadow or be misattributed as the declared case.
"""

def test_evalcase_on_unrelated_param_is_ignored(self) -> None:
def task(case: EvalCase, sneaky: object) -> str:
return str(case.inputs)

wrapped = make_eval_wrapper(task, [_ok])
payload = asyncio.run(
wrapped(
case=EvalCase(inputs="a", name="a1"),
sneaky=EvalCase(inputs="b", name="b1"),
)
)
# The declared case wins; `sneaky` is not consulted.
assert payload.case_name == "a1"
assert payload.expected_output is None

def test_single_runtime_evalcase_kwarg_ok(self) -> None:
def task(case: EvalCase) -> str:
return str(case.inputs)

wrapped = make_eval_wrapper(task, [_ok])
payload = asyncio.run(wrapped(case=EvalCase(inputs="a", name="a1")))
assert payload.case_name == "a1"
Loading