diff --git a/agentfailbench/failures/__init__.py b/agentfailbench/failures/__init__.py index e69de29..1b9ae30 100644 --- a/agentfailbench/failures/__init__.py +++ b/agentfailbench/failures/__init__.py @@ -0,0 +1,5 @@ +"""Failure-injection suites for AgentFailBench.""" + +from agentfailbench.failures.base import FailureInjector, FailureInjectorRegistry + +__all__ = ["FailureInjector", "FailureInjectorRegistry"] diff --git a/agentfailbench/failures/base.py b/agentfailbench/failures/base.py new file mode 100644 index 0000000..b295b9c --- /dev/null +++ b/agentfailbench/failures/base.py @@ -0,0 +1,74 @@ +"""Shared failure-injector interface for AgentFailBench suites. + +Every failure-injection suite (tool drift, memory, planning, retrieval, data, +communication — see ``docs/roadmap.md`` Milestone 2) drives an episode through +the same three lifecycle hooks so a single dispatcher can coordinate +heterogeneous injectors without depending on their concrete types. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Protocol, runtime_checkable + +from runtime.schemas.episode import EnvObservation + + +@runtime_checkable +class FailureInjector(Protocol): + """Lifecycle contract every failure injector must implement. + + ``env`` is typed as ``Any`` on purpose: each suite injects failures into + a different environment (customer API today, memory/retrieval stores in + later milestones), so the interface stays environment-agnostic. + """ + + def before_step(self, env: Any, step: int) -> None: + """Arm/apply the failure ahead of the upcoming step, if triggered.""" + ... + + def after_step(self, env: Any, step: int, observation: EnvObservation) -> None: + """React to the outcome of a step that just executed.""" + ... + + def reset(self) -> None: + """Return the injector to its pre-trigger state for a fresh episode.""" + ... + + @property + def triggered(self) -> bool: + """Whether the injected failure has fired at least once.""" + ... + + +@dataclass +class FailureInjectorRegistry: + """Registers named injectors and dispatches lifecycle hooks to all of them.""" + + _injectors: dict[str, FailureInjector] = field(default_factory=dict) + + def register(self, name: str, injector: FailureInjector) -> None: + if name in self._injectors: + raise ValueError(f"Duplicate failure injector: {name}") + self._injectors[name] = injector + + def get(self, name: str) -> FailureInjector: + return self._injectors[name] + + def list_names(self) -> list[str]: + return sorted(self._injectors) + + def dispatch_before_step(self, env: Any, step: int) -> None: + for injector in self._injectors.values(): + injector.before_step(env, step) + + def dispatch_after_step(self, env: Any, step: int, observation: EnvObservation) -> None: + for injector in self._injectors.values(): + injector.after_step(env, step, observation) + + def dispatch_reset(self) -> None: + for injector in self._injectors.values(): + injector.reset() + + def any_triggered(self) -> bool: + return any(injector.triggered for injector in self._injectors.values()) diff --git a/agentfailbench/failures/tool_drift/__init__.py b/agentfailbench/failures/tool_drift/__init__.py index 18354af..1b12908 100644 --- a/agentfailbench/failures/tool_drift/__init__.py +++ b/agentfailbench/failures/tool_drift/__init__.py @@ -1,5 +1,8 @@ """ToolDrift suite package.""" -from agentfailbench.failures.tool_drift.injector import SemanticDriftInjector +from agentfailbench.failures.tool_drift.injector import ( + SemanticDriftInjector, + SemanticDriftInjectorAdapter, +) -__all__ = ["SemanticDriftInjector"] +__all__ = ["SemanticDriftInjector", "SemanticDriftInjectorAdapter"] diff --git a/agentfailbench/failures/tool_drift/injector.py b/agentfailbench/failures/tool_drift/injector.py index 425839d..b5d0baa 100644 --- a/agentfailbench/failures/tool_drift/injector.py +++ b/agentfailbench/failures/tool_drift/injector.py @@ -2,9 +2,10 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field from agentfailbench.environments.customer_api.env import CustomerApiEnv +from runtime.schemas.episode import EnvObservation @dataclass @@ -27,3 +28,30 @@ def triggered(self) -> bool: def reset(self) -> None: self._triggered = False + + +@dataclass +class SemanticDriftInjectorAdapter: + """Adapts :class:`SemanticDriftInjector` to the shared ``FailureInjector`` protocol. + + ``SemanticDriftInjector`` predates the shared interface (``agentfailbench.failures.base``) + and only needs a trigger-based ``before_step``/``reset``. This adapter fills in the + ``after_step`` hook so the injector can be registered and dispatched alongside future + suites without changing its existing behavior. + """ + + injector: SemanticDriftInjector = field(default_factory=SemanticDriftInjector) + + def before_step(self, env: CustomerApiEnv, step: int) -> None: + self.injector.before_step(env, step) + + def after_step(self, env: CustomerApiEnv, step: int, observation: EnvObservation) -> None: + """No-op: semantic drift is a one-shot contract flip with no post-step reaction.""" + del env, step, observation + + def reset(self) -> None: + self.injector.reset() + + @property + def triggered(self) -> bool: + return self.injector.triggered diff --git a/agentfailbench/runners/episode.py b/agentfailbench/runners/episode.py index f5c99a6..6f7c1ac 100644 --- a/agentfailbench/runners/episode.py +++ b/agentfailbench/runners/episode.py @@ -8,7 +8,11 @@ from agentfailbench.agents.scripted import ScriptedApiAgent from agentfailbench.environments.customer_api.env import CustomerApiEnv -from agentfailbench.failures.tool_drift.injector import SemanticDriftInjector +from agentfailbench.failures.base import FailureInjectorRegistry +from agentfailbench.failures.tool_drift.injector import ( + SemanticDriftInjector, + SemanticDriftInjectorAdapter, +) from agentfailbench.models import BenchmarkCaseModel from agentfailbench.registry import CaseRegistry from baselines.rules.detectors import BaseDetector, DetectionResult, all_detectors @@ -46,21 +50,27 @@ def _run_traced( case: BenchmarkCaseModel, *, inject_failure: bool, -) -> tuple[CustomerApiEnv, ScriptedApiAgent, TraceCollector, SemanticDriftInjector]: +) -> tuple[CustomerApiEnv, ScriptedApiAgent, TraceCollector, SemanticDriftInjectorAdapter]: task = build_task(case) env = CustomerApiEnv(task=task) agent = ScriptedApiAgent(env=env) - injector = SemanticDriftInjector(trigger_step=case.failure.trigger_step) + injector = SemanticDriftInjectorAdapter( + SemanticDriftInjector(trigger_step=case.failure.trigger_step) + ) + injectors = FailureInjectorRegistry() + injectors.register("tool_drift.semantic_drift", injector) collector = TraceCollector(task_id=case.case_id) for _ in range(case.task.expected_steps): upcoming_step = agent._step + 1 if inject_failure: - injector.before_step(env, upcoming_step) + injectors.dispatch_before_step(env, upcoming_step) action = agent.next_action() if action is None: break obs = env.step(action) + if inject_failure: + injectors.dispatch_after_step(env, upcoming_step, obs) progress = min(1.0, action.step / case.task.expected_steps) collector.record(action, obs, agent.expected_attributes(), progress) diff --git a/tests/unit/test_failure_injector.py b/tests/unit/test_failure_injector.py new file mode 100644 index 0000000..a1c29f2 --- /dev/null +++ b/tests/unit/test_failure_injector.py @@ -0,0 +1,116 @@ +"""Unit tests for the shared failure-injector interface and registry.""" + +from __future__ import annotations + +from dataclasses import dataclass, field + +import pytest + +from agentfailbench.environments.customer_api.env import CustomerApiEnv +from agentfailbench.failures.base import FailureInjector, FailureInjectorRegistry +from agentfailbench.failures.tool_drift.injector import ( + SemanticDriftInjector, + SemanticDriftInjectorAdapter, +) +from runtime.schemas.episode import EnvObservation, TaskSpec + + +def _make_env() -> CustomerApiEnv: + task = TaskSpec( + task_id="t-1", objective="update_customer_subscription", environment="customer_service_api" + ) + return CustomerApiEnv(task=task) + + +@dataclass +class _RecordingInjector: + """Minimal stand-in suite used to exercise registration/dispatch in isolation.""" + + calls: list[str] = field(default_factory=list) + _triggered: bool = False + + def before_step(self, env: object, step: int) -> None: + del env + self.calls.append(f"before:{step}") + self._triggered = True + + def after_step(self, env: object, step: int, observation: EnvObservation) -> None: + del env, observation + self.calls.append(f"after:{step}") + + def reset(self) -> None: + self.calls.append("reset") + self._triggered = False + + @property + def triggered(self) -> bool: + return self._triggered + + +def test_semantic_drift_adapter_satisfies_failure_injector_protocol() -> None: + adapter = SemanticDriftInjectorAdapter(SemanticDriftInjector(trigger_step=2)) + assert isinstance(adapter, FailureInjector) + + +def test_recording_injector_satisfies_failure_injector_protocol() -> None: + assert isinstance(_RecordingInjector(), FailureInjector) + + +def test_registry_rejects_duplicate_names() -> None: + registry = FailureInjectorRegistry() + registry.register("a", _RecordingInjector()) + with pytest.raises(ValueError): + registry.register("a", _RecordingInjector()) + + +def test_registry_get_and_list_names() -> None: + registry = FailureInjectorRegistry() + injector = _RecordingInjector() + registry.register("only", injector) + assert registry.list_names() == ["only"] + assert registry.get("only") is injector + + +def test_registry_dispatches_lifecycle_hooks_to_all_injectors() -> None: + registry = FailureInjectorRegistry() + first = _RecordingInjector() + second = _RecordingInjector() + registry.register("first", first) + registry.register("second", second) + env = _make_env() + obs = EnvObservation(success=True) + + registry.dispatch_before_step(env, 1) + registry.dispatch_after_step(env, 1, obs) + registry.dispatch_reset() + + assert first.calls == ["before:1", "after:1", "reset"] + assert second.calls == ["before:1", "after:1", "reset"] + + +def test_registry_any_triggered_reflects_injector_state() -> None: + registry = FailureInjectorRegistry() + registry.register("only", _RecordingInjector()) + assert registry.any_triggered() is False + registry.dispatch_before_step(_make_env(), 1) + assert registry.any_triggered() is True + + +def test_adapter_delegates_to_wrapped_semantic_drift_injector() -> None: + inner = SemanticDriftInjector(trigger_step=2) + adapter = SemanticDriftInjectorAdapter(inner) + env = _make_env() + + adapter.before_step(env, 1) + assert adapter.triggered is False + assert env.contract_version == "v1" + + adapter.before_step(env, 2) + assert adapter.triggered is True + assert env.contract_version == "v2" + + adapter.after_step(env, 2, EnvObservation(success=True)) # no-op, must not raise + + adapter.reset() + assert adapter.triggered is False + assert adapter.injector is inner