From bd06a01daaacde56313a5a817da9d3d60be6880b Mon Sep 17 00:00:00 2001 From: lifelmy Date: Tue, 1 Sep 2026 14:31:05 +0800 Subject: [PATCH] feat(providers): add session discovery + evaluate_sessions batch helper (#143) Extend TraceProvider with an optional, non-abstract list_sessions(SessionFilter) so backends that can enumerate sessions expose discovery, while the default raises NotImplementedError pointing at the known-session-id path. Add a SessionFilter pydantic model (start_time/end_time/limit + additional_fields). Add strands_evals.batch.evaluate_sessions(provider, evaluators, session_filter) that composes list_sessions -> Case-per-session -> Experiment.run, removing the discover/build/run boilerplate every provider user rewrites. It returns a single EvaluationReport to match the current Experiment API (flattened across evaluators), rather than the list[EvaluationReport] in the original sketch. Export SessionFilter and evaluate_sessions from the package root. --- src/strands_evals/__init__.py | 4 + src/strands_evals/batch.py | 64 +++++++++++++ src/strands_evals/providers/__init__.py | 2 + src/strands_evals/providers/trace_provider.py | 53 ++++++++++- .../providers/test_trace_provider.py | 57 ++++++++++++ tests/strands_evals/test_batch.py | 93 +++++++++++++++++++ 6 files changed, 272 insertions(+), 1 deletion(-) create mode 100644 src/strands_evals/batch.py create mode 100644 tests/strands_evals/test_batch.py diff --git a/src/strands_evals/__init__.py b/src/strands_evals/__init__.py index 5dbaee36..96531802 100644 --- a/src/strands_evals/__init__.py +++ b/src/strands_evals/__init__.py @@ -1,9 +1,11 @@ from . import chaos, detectors, evaluators, extractors, generators, providers, simulation, telemetry, types +from .batch import evaluate_sessions from .case import Case from .eval_task_handler import EvalTaskHandler, TracedHandler, eval_task from .evaluation_data_store import EvaluationDataStore from .experiment import Experiment from .local_file_task_result_store import LocalFileTaskResultStore +from .providers import SessionFilter from .simulation import ActorSimulator, UserSimulator from .telemetry import StrandsEvalsTelemetry, get_tracer from .types.detector import DiagnosisConfig @@ -19,6 +21,8 @@ "EvalTaskHandler", "TracedHandler", "eval_task", + "evaluate_sessions", + "SessionFilter", "chaos", "detectors", "evaluators", diff --git a/src/strands_evals/batch.py b/src/strands_evals/batch.py new file mode 100644 index 00000000..e7ab0740 --- /dev/null +++ b/src/strands_evals/batch.py @@ -0,0 +1,64 @@ +"""Batch evaluation over sessions discovered from a TraceProvider. + +Composes `TraceProvider.list_sessions` (session discovery) with +`TraceProvider.get_evaluation_data` (per-session retrieval) and an `Experiment` +so callers can evaluate every session matching a filter in a single call, +instead of writing the discover -> build cases -> run boilerplate by hand. +""" + +import asyncio +import logging +from collections.abc import Callable + +from .case import Case +from .evaluators.evaluator import Evaluator +from .experiment import Experiment +from .providers.trace_provider import SessionFilter, TraceProvider +from .types.evaluation import TaskOutput +from .types.evaluation_report import EvaluationReport + +logger = logging.getLogger(__name__) + + +def evaluate_sessions( + provider: TraceProvider, + evaluators: list[Evaluator], + session_filter: SessionFilter | None = None, + *, + max_workers: int = 1, +) -> EvaluationReport: + """Discover sessions from a provider and evaluate them all. + + Discovers session IDs via `provider.list_sessions(session_filter)`, wraps each + in a `Case` keyed on the session ID, and runs the given evaluators through an + `Experiment`. The per-case task fetches trace data via + `provider.get_evaluation_data(session_id)`. + + Args: + provider: A TraceProvider whose backend supports session discovery + (i.e. overrides `list_sessions`). + evaluators: Evaluators to run against each discovered session. + session_filter: Optional filter narrowing which sessions to evaluate. If + None, provider-specific defaults apply. + max_workers: Maximum number of parallel workers. Defaults to 1 (sequential), + matching `Experiment.run_evaluations`. Providers that share a single + network client or a `TracedHandler` exporter should keep this at 1. + + Returns: + A single `EvaluationReport` flattened across every (session, evaluator) pair, + with each row tagged by its evaluator via the `evaluator` field on `cases`. + + Raises: + NotImplementedError: If the provider does not support session discovery. + ProviderError: If the provider is unreachable or returns an error. + """ + session_ids = list(provider.list_sessions(session_filter)) + logger.debug("session_count=<%d> | discovered sessions for batch evaluation", len(session_ids)) + + cases: list[Case] = [Case(name=session_id, input=session_id, session_id=session_id) for session_id in session_ids] + + experiment = Experiment(cases=cases, evaluators=evaluators) + + task: Callable[[Case], TaskOutput] = provider.as_task() + + return asyncio.run(experiment.run_evaluations_async(task, max_workers=max_workers)) diff --git a/src/strands_evals/providers/__init__.py b/src/strands_evals/providers/__init__.py index 3cdb689d..cb110c65 100644 --- a/src/strands_evals/providers/__init__.py +++ b/src/strands_evals/providers/__init__.py @@ -6,6 +6,7 @@ TraceProviderError, ) from .trace_provider import ( + SessionFilter, TraceProvider, ) @@ -14,6 +15,7 @@ "LangfuseProvider", "OpenSearchProvider", "ProviderError", + "SessionFilter", "SessionNotFoundError", "TraceProvider", "TraceProviderError", diff --git a/src/strands_evals/providers/trace_provider.py b/src/strands_evals/providers/trace_provider.py index c3eff4f9..4b199e2b 100644 --- a/src/strands_evals/providers/trace_provider.py +++ b/src/strands_evals/providers/trace_provider.py @@ -1,12 +1,35 @@ """TraceProvider interface for retrieving agent trace data from observability backends.""" from abc import ABC, abstractmethod -from collections.abc import Callable +from collections.abc import Callable, Iterator +from datetime import datetime +from typing import Any + +from pydantic import BaseModel, Field from ..case import Case from ..types.evaluation import TaskOutput +class SessionFilter(BaseModel): + """Filter criteria for discovering sessions from a provider. + + Universal fields (`start_time`, `end_time`, `limit`) are defined here. + Provider-specific parameters that do not generalize go in `additional_fields`. + + Attributes: + start_time: Only include sessions at or after this time. None means no lower bound. + end_time: Only include sessions at or before this time. None means no upper bound. + limit: Maximum number of sessions to return. None means no limit. + additional_fields: Provider-specific filter parameters that have no universal field. + """ + + start_time: datetime | None = None + end_time: datetime | None = None + limit: int | None = None + additional_fields: dict[str, Any] = Field(default_factory=dict) + + class TraceProvider(ABC): """Retrieves agent trace data from observability backends for evaluation. @@ -43,3 +66,31 @@ def as_task(self) -> Callable[[Case], TaskOutput]: for that case's session. """ return lambda case: self.get_evaluation_data(case.session_id) + + def list_sessions( + self, + session_filter: SessionFilter | None = None, + ) -> Iterator[str]: + """Discover session IDs matching filter criteria. + + Returns session IDs that can be fed to `get_evaluation_data()`. This + method is intentionally not abstract: providers only override it when + their backend supports session discovery. The default raises + `NotImplementedError` with a message pointing at the known-session-id + access pattern. + + Args: + session_filter: Optional filter criteria. If None, provider-specific + defaults apply. + + Yields: + Session ID strings. + + Raises: + NotImplementedError: If the provider does not support session discovery. + ProviderError: If the provider is unreachable or returns an error. + """ + raise NotImplementedError( + "this provider does not support session discovery; " + "use get_evaluation_data() with a known session_id instead" + ) diff --git a/tests/strands_evals/providers/test_trace_provider.py b/tests/strands_evals/providers/test_trace_provider.py index 579c09be..efef6c51 100644 --- a/tests/strands_evals/providers/test_trace_provider.py +++ b/tests/strands_evals/providers/test_trace_provider.py @@ -1,5 +1,8 @@ """Tests for TraceProvider ABC and exception hierarchy.""" +from collections.abc import Iterator +from datetime import datetime + import pytest from strands_evals.providers.exceptions import ( @@ -8,6 +11,7 @@ TraceProviderError, ) from strands_evals.providers.trace_provider import ( + SessionFilter, TraceProvider, ) from strands_evals.types.evaluation import TaskOutput @@ -29,6 +33,23 @@ def get_evaluation_data(self, session_id: str) -> TaskOutput: ) +class DiscoverableProvider(ConcreteProvider): + """Provider that overrides list_sessions to support discovery.""" + + def __init__(self, session_ids: list[str], session: Session | None = None): + super().__init__(session=session) + self._session_ids = session_ids + self.last_filter: SessionFilter | None = None + + def list_sessions(self, session_filter: SessionFilter | None = None) -> Iterator[str]: + self.last_filter = session_filter + limit = session_filter.limit if session_filter else None + for i, session_id in enumerate(self._session_ids): + if limit is not None and i >= limit: + return + yield session_id + + class TestExceptionHierarchy: def test_trace_provider_error_is_exception(self): assert issubclass(TraceProviderError, Exception) @@ -89,3 +110,39 @@ class FakeCase: result = task(FakeCase()) assert result["output"] == "test response" assert result["trajectory"] == session + + +class TestSessionFilter: + def test_defaults_are_none_and_empty(self): + f = SessionFilter() + assert f.start_time is None + assert f.end_time is None + assert f.limit is None + assert f.additional_fields == {} + + def test_accepts_universal_and_additional_fields(self): + start = datetime(2026, 1, 1) + end = datetime(2026, 1, 2) + f = SessionFilter(start_time=start, end_time=end, limit=5, additional_fields={"env": "prod"}) + assert f.start_time == start + assert f.end_time == end + assert f.limit == 5 + assert f.additional_fields == {"env": "prod"} + + +class TestListSessions: + def test_default_list_sessions_raises_not_implemented(self): + provider = ConcreteProvider() + with pytest.raises(NotImplementedError, match="does not support session discovery"): + list(provider.list_sessions()) + + def test_overridden_list_sessions_yields_ids(self): + provider = DiscoverableProvider(["s1", "s2", "s3"]) + assert list(provider.list_sessions()) == ["s1", "s2", "s3"] + + def test_list_sessions_receives_filter(self): + provider = DiscoverableProvider(["s1", "s2", "s3"]) + f = SessionFilter(limit=2) + result = list(provider.list_sessions(f)) + assert result == ["s1", "s2"] + assert provider.last_filter is f diff --git a/tests/strands_evals/test_batch.py b/tests/strands_evals/test_batch.py new file mode 100644 index 00000000..c8bd27f4 --- /dev/null +++ b/tests/strands_evals/test_batch.py @@ -0,0 +1,93 @@ +"""Tests for evaluate_sessions batch evaluation.""" + +from collections.abc import Iterator + +from strands_evals.batch import evaluate_sessions +from strands_evals.evaluators.deterministic.output import Equals +from strands_evals.providers.exceptions import SessionNotFoundError +from strands_evals.providers.trace_provider import SessionFilter, TraceProvider +from strands_evals.types.evaluation import TaskOutput + + +class FakeProvider(TraceProvider): + """Provider with discovery over an in-memory session_id -> output map.""" + + def __init__(self, outputs: dict[str, str]): + self._outputs = outputs + self.fetched: list[str] = [] + self.last_filter: SessionFilter | None = None + + def list_sessions(self, session_filter: SessionFilter | None = None) -> Iterator[str]: + self.last_filter = session_filter + limit = session_filter.limit if session_filter else None + for i, session_id in enumerate(self._outputs): + if limit is not None and i >= limit: + return + yield session_id + + def get_evaluation_data(self, session_id: str) -> TaskOutput: + if session_id not in self._outputs: + raise SessionNotFoundError(session_id) + self.fetched.append(session_id) + return TaskOutput(output=self._outputs[session_id]) + + +class NoDiscoveryProvider(TraceProvider): + """Provider that does not override list_sessions.""" + + def get_evaluation_data(self, session_id: str) -> TaskOutput: + return TaskOutput(output="x") + + +def test_evaluate_sessions_evaluates_every_discovered_session(): + provider = FakeProvider({"s1": "hello", "s2": "world"}) + # Equals compares the task output against each case's expected_output. With no + # expected_output set (None), both cases fail, but the report still has one row + # per session and each session was fetched. + report = evaluate_sessions(provider, evaluators=[Equals()]) + + assert len(report.cases) == 2 + assert sorted(provider.fetched) == ["s1", "s2"] + assert {case["name"] for case in report.cases} == {"s1", "s2"} + + +def test_evaluate_sessions_passes_output_to_evaluator(): + provider = FakeProvider({"s1": "match", "s2": "other"}) + + # Equals("match") passes only for the session whose output is exactly "match". + report = evaluate_sessions(provider, evaluators=[Equals("match")]) + + passes = {case["name"]: passed for case, passed in zip(report.cases, report.test_passes, strict=True)} + assert passes["s1"] is True + assert passes["s2"] is False + + +def test_evaluate_sessions_forwards_filter(): + provider = FakeProvider({"s1": "a", "s2": "b", "s3": "c"}) + f = SessionFilter(limit=2) + + report = evaluate_sessions(provider, evaluators=[Equals()], session_filter=f) + + assert provider.last_filter is f + assert len(report.cases) == 2 + assert sorted(provider.fetched) == ["s1", "s2"] + + +def test_evaluate_sessions_empty_discovery_yields_empty_report(): + provider = FakeProvider({}) + + report = evaluate_sessions(provider, evaluators=[Equals()]) + + assert report.cases == [] + assert provider.fetched == [] + + +def test_evaluate_sessions_without_discovery_raises_not_implemented(): + provider = NoDiscoveryProvider() + + try: + evaluate_sessions(provider, evaluators=[Equals()]) + except NotImplementedError as err: + assert "does not support session discovery" in str(err) + else: + raise AssertionError("expected NotImplementedError")