From 1b37c9502917bc54a0fa88715fb3705c7b268f3d Mon Sep 17 00:00:00 2001 From: Abhinaysai Kamineni <66816045+askmy-stack@users.noreply.github.com> Date: Tue, 14 Jul 2026 09:25:43 -0400 Subject: [PATCH] Define shared FailureInjector interface and registry (#12) SemanticDriftInjector previously wired directly into runners/episode.py with no common contract for future suites (memory, planning, retrieval, data, communication). Add a FailureInjector protocol (before_step/ after_step/reset + triggered) and a FailureInjectorRegistry that registers named injectors and dispatches lifecycle hooks to all of them, then adapt SemanticDriftInjector to the protocol and wire the episode runner through the registry instead of the concrete class. Fixes #12 Co-authored-by: Cursor --- agentfailbench/failures/__init__.py | 5 + agentfailbench/failures/base.py | 74 +++++++++++ .../failures/tool_drift/__init__.py | 7 +- .../failures/tool_drift/injector.py | 30 ++++- agentfailbench/runners/episode.py | 18 ++- tests/unit/test_failure_injector.py | 116 ++++++++++++++++++ 6 files changed, 243 insertions(+), 7 deletions(-) create mode 100644 agentfailbench/failures/base.py create mode 100644 tests/unit/test_failure_injector.py 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