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
5 changes: 5 additions & 0 deletions agentfailbench/failures/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
"""Failure-injection suites for AgentFailBench."""

from agentfailbench.failures.base import FailureInjector, FailureInjectorRegistry

__all__ = ["FailureInjector", "FailureInjectorRegistry"]
74 changes: 74 additions & 0 deletions agentfailbench/failures/base.py
Original file line number Diff line number Diff line change
@@ -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())
7 changes: 5 additions & 2 deletions agentfailbench/failures/tool_drift/__init__.py
Original file line number Diff line number Diff line change
@@ -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"]
30 changes: 29 additions & 1 deletion agentfailbench/failures/tool_drift/injector.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
18 changes: 14 additions & 4 deletions agentfailbench/runners/episode.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down
116 changes: 116 additions & 0 deletions tests/unit/test_failure_injector.py
Original file line number Diff line number Diff line change
@@ -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
Loading