Skip to content
Open
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
4 changes: 4 additions & 0 deletions src/strands_evals/__init__.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -19,6 +21,8 @@
"EvalTaskHandler",
"TracedHandler",
"eval_task",
"evaluate_sessions",
"SessionFilter",
"chaos",
"detectors",
"evaluators",
Expand Down
64 changes: 64 additions & 0 deletions src/strands_evals/batch.py
Original file line number Diff line number Diff line change
@@ -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))
2 changes: 2 additions & 0 deletions src/strands_evals/providers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
TraceProviderError,
)
from .trace_provider import (
SessionFilter,
TraceProvider,
)

Expand All @@ -14,6 +15,7 @@
"LangfuseProvider",
"OpenSearchProvider",
"ProviderError",
"SessionFilter",
"SessionNotFoundError",
"TraceProvider",
"TraceProviderError",
Expand Down
53 changes: 52 additions & 1 deletion src/strands_evals/providers/trace_provider.py
Original file line number Diff line number Diff line change
@@ -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.

Expand Down Expand Up @@ -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"
)
57 changes: 57 additions & 0 deletions tests/strands_evals/providers/test_trace_provider.py
Original file line number Diff line number Diff line change
@@ -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 (
Expand All @@ -8,6 +11,7 @@
TraceProviderError,
)
from strands_evals.providers.trace_provider import (
SessionFilter,
TraceProvider,
)
from strands_evals.types.evaluation import TaskOutput
Expand All @@ -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)
Expand Down Expand Up @@ -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
93 changes: 93 additions & 0 deletions tests/strands_evals/test_batch.py
Original file line number Diff line number Diff line change
@@ -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")
Loading