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
24 changes: 22 additions & 2 deletions agentfailbench/agents/scripted.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,22 +6,38 @@
from typing import Any

from agentfailbench.environments.customer_api.env import CustomerApiEnv, plan_id_meaning
from agentfailbench.failures.memory.injector import TARGET_PLAN_MEMORY_KEY
from agentfailbench.memory.store import AgentMemory
from runtime.schemas.episode import Action


@dataclass
class ScriptedApiAgent:
"""Eight-step scripted tool sequence for update_customer_subscription."""
"""Eight-step scripted tool sequence for update_customer_subscription.

When ``memory`` is provided, the upgrade target is read from
``TARGET_PLAN_MEMORY_KEY`` so MemoryGuard injectors can poison the plan.
"""

env: CustomerApiEnv
believed_meaning: str = "billing_plan_code"
believed_plan_id: str | None = None
memory: AgentMemory | None = None
history: list[dict[str, Any]] = field(default_factory=list)
_step: int = 0

def __post_init__(self) -> None:
if self.memory is not None and self.memory.get(TARGET_PLAN_MEMORY_KEY) is None:
self.memory.set(TARGET_PLAN_MEMORY_KEY, self.env.task.target_plan_code)
self.memory.checkpoint()

def belief_plan_id(self) -> str:
if self.believed_plan_id is not None:
return self.believed_plan_id
if self.memory is not None:
memorized = self.memory.get(TARGET_PLAN_MEMORY_KEY)
if memorized is not None:
return str(memorized)
return self.env.task.target_plan_code

def _sequence(self) -> list[Action]:
Expand Down Expand Up @@ -59,11 +75,15 @@ def next_action(self) -> Action | None:
return action

def expected_attributes(self) -> dict[str, Any]:
return {
attrs: dict[str, Any] = {
"plan_id_meaning": self.believed_meaning,
"plan_id": self.belief_plan_id(),
"contract_belief": "v1" if self.believed_meaning == "billing_plan_code" else "v2",
"live_target_plan_id": self.env.task.target_plan_code,
}
if self.memory is not None:
attrs["memory_plan_id"] = self.memory.get(TARGET_PLAN_MEMORY_KEY)
return attrs

def apply_contract_refresh(self, version: str = "v2") -> None:
from typing import cast
Expand Down
8 changes: 8 additions & 0 deletions agentfailbench/failures/memory/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
"""MemoryGuard failure suite."""

from agentfailbench.failures.memory.injector import (
TARGET_PLAN_MEMORY_KEY,
StaleMemoryInjector,
)

__all__ = ["TARGET_PLAN_MEMORY_KEY", "StaleMemoryInjector"]
51 changes: 51 additions & 0 deletions agentfailbench/failures/memory/injector.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
"""MemoryGuard: stale memory failure injector."""

from __future__ import annotations

from dataclasses import dataclass, field
from typing import Any

from agentfailbench.memory.store import AgentMemory
from runtime.schemas.episode import EnvObservation

# Memory key used by the scripted subscription agent for the upgrade target.
TARGET_PLAN_MEMORY_KEY = "target_plan_id"


@dataclass
class StaleMemoryInjector:
"""Overwrite a memory fact with a stale value at ``trigger_step`` (no exceptions).

Implements the shared :class:`~agentfailbench.failures.base.FailureInjector`
protocol directly so it can be registered alongside tool-drift injectors.
"""

memory: AgentMemory
key: str = TARGET_PLAN_MEMORY_KEY
stale_value: str = "BASIC_MONTHLY"
trigger_step: int = 4
_armed: bool = True
_triggered: bool = False
_pre_inject_value: Any = field(default=None, repr=False)

def before_step(self, env: object, step: int) -> None:
del env
if self._armed and not self._triggered and step >= self.trigger_step:
self._pre_inject_value = self.memory.get(self.key)
# Ensure a checkpoint exists so recovery can roll back.
if self.memory._checkpoint is None:
self.memory.checkpoint()
self.memory.set(self.key, self.stale_value)
self._triggered = True

def after_step(self, env: object, step: int, observation: EnvObservation) -> None:
"""No-op: stale memory is a one-shot overwrite with no post-step reaction."""
del env, step, observation

def reset(self) -> None:
self._triggered = False
self._pre_inject_value = None

@property
def triggered(self) -> bool:
return self._triggered
25 changes: 25 additions & 0 deletions agentfailbench/failures/memory/stale-memory-001.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
case_id: stale-memory-001

task:
objective: update_customer_subscription
environment: customer_service_api
expected_steps: 8

failure:
category: memory_stale
trigger_step: 4
visible_exception: false
reversible: true

ground_truth:
root_cause: stale_memory
first_detectable_step: 4
final_failure_step: 6
expected_recovery:
- rollback_memory
- refresh_context
- replan

risk:
severity: medium
impact: incorrect_account_update_from_stale_plan_target
5 changes: 5 additions & 0 deletions agentfailbench/memory/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
"""Agent memory utilities."""

from agentfailbench.memory.store import AgentMemory

__all__ = ["AgentMemory"]
36 changes: 36 additions & 0 deletions agentfailbench/memory/store.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
"""Persistent agent memory with checkpoint / rollback support."""

from __future__ import annotations

from copy import deepcopy
from dataclasses import dataclass, field
from typing import Any


@dataclass
class AgentMemory:
"""Simple key-value memory used by scripted agents and MemoryGuard injectors."""

store: dict[str, Any] = field(default_factory=dict)
_checkpoint: dict[str, Any] | None = None

def get(self, key: str, default: Any = None) -> Any:
return self.store.get(key, default)

def set(self, key: str, value: Any) -> None:
self.store[key] = value

def checkpoint(self) -> None:
"""Save a rollback snapshot of the current store."""
self._checkpoint = deepcopy(self.store)

def rollback(self) -> bool:
"""Restore the last checkpoint. Returns False if none exists."""
if self._checkpoint is None:
return False
self.store = deepcopy(self._checkpoint)
return True

def clear(self) -> None:
self.store.clear()
self._checkpoint = None
5 changes: 4 additions & 1 deletion runtime/features/__init__.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,14 @@
"""Feature extractors."""

from runtime.features.memory_conflict import first_memory_conflict_step, memory_live_conflict
from runtime.features.mismatch import first_mismatch_step, semantic_mismatch
from runtime.features.progress import flat_progress_step, repeated_action_step

__all__ = [
"first_memory_conflict_step",
"first_mismatch_step",
"semantic_mismatch",
"flat_progress_step",
"memory_live_conflict",
"repeated_action_step",
"semantic_mismatch",
]
30 changes: 30 additions & 0 deletions runtime/features/memory_conflict.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
"""Memory vs live-tool conflict features (MemoryGuard)."""

from __future__ import annotations

from runtime.schemas.trace import TraceEvent


def memory_live_conflict(event: TraceEvent) -> bool:
"""True when memorized plan target disagrees with the live task target.

Live tool metadata is represented by ``live_target_plan_id`` (from the task
objective). A stale memory write changes ``memory_plan_id`` / ``plan_id``
without raising an exception — the conflict signal for MemoryGuard.
"""
if event.expectation is None:
return False
attrs = event.expectation.attributes
memory_plan = attrs.get("memory_plan_id", attrs.get("plan_id"))
live_target = attrs.get("live_target_plan_id")
if memory_plan is None or live_target is None:
return False
return str(memory_plan) != str(live_target)


def first_memory_conflict_step(events: list[TraceEvent]) -> int | None:
for event in events:
if memory_live_conflict(event):
step = event.attributes.get("step")
return int(step) if step is not None else None
return None
2 changes: 2 additions & 0 deletions runtime/schemas/root_cause.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,11 +23,13 @@ class RootCauseCode(StrEnum):

PLAN_IDENTIFIER_SEMANTICS_CHANGED = "plan_identifier_semantics_changed"
TOOL_TRANSPORT_ERROR = "tool_transport_error"
STALE_MEMORY = "stale_memory"


ROOT_CAUSE_CATEGORY: dict[RootCauseCode, FailureCategory] = {
RootCauseCode.PLAN_IDENTIFIER_SEMANTICS_CHANGED: FailureCategory.TOOL,
RootCauseCode.TOOL_TRANSPORT_ERROR: FailureCategory.TOOL,
RootCauseCode.STALE_MEMORY: FailureCategory.MEMORY,
}


Expand Down
5 changes: 5 additions & 0 deletions tests/unit/test_root_cause.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,11 @@ def test_root_cause_label_infers_category() -> None:
assert label.schema_version == ROOT_CAUSE_LABEL_SCHEMA_VERSION


def test_stale_memory_root_cause_maps_to_memory_category() -> None:
label = root_cause_label(RootCauseCode.STALE_MEMORY)
assert label.category == FailureCategory.MEMORY


def test_root_cause_label_accepts_raw_string() -> None:
label = root_cause_label("tool_transport_error")
assert label.code is RootCauseCode.TOOL_TRANSPORT_ERROR
Expand Down
137 changes: 137 additions & 0 deletions tests/unit/test_stale_memory.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,137 @@
"""Unit tests for MemoryGuard stale-memory injector and conflict feature."""

from __future__ import annotations

from pathlib import Path

from agentfailbench.agents.scripted import ScriptedApiAgent
from agentfailbench.environments.customer_api.env import CustomerApiEnv
from agentfailbench.failures.base import FailureInjector, FailureInjectorRegistry
from agentfailbench.failures.memory.injector import (
TARGET_PLAN_MEMORY_KEY,
StaleMemoryInjector,
)
from agentfailbench.memory.store import AgentMemory
from agentfailbench.registry import CaseRegistry
from runtime.features.memory_conflict import first_memory_conflict_step, memory_live_conflict
from runtime.schemas.episode import TaskSpec
from runtime.schemas.root_cause import RootCauseCode
from runtime.schemas.taxonomy import FailureCategory
from runtime.schemas.trace import Expectation, TraceEvent
from runtime.tracing.collector import TraceCollector


def _task() -> TaskSpec:
return TaskSpec(
task_id="stale-mem-test",
objective="update_customer_subscription",
environment="customer_service_api",
)


def test_stale_memory_injector_satisfies_protocol() -> None:
memory = AgentMemory()
injector = StaleMemoryInjector(memory=memory, trigger_step=2)
assert isinstance(injector, FailureInjector)


def test_stale_memory_injector_overwrites_at_trigger_step() -> None:
memory = AgentMemory()
memory.set(TARGET_PLAN_MEMORY_KEY, "GOLD_ANNUAL")
memory.checkpoint()
injector = StaleMemoryInjector(memory=memory, trigger_step=4, stale_value="BASIC_MONTHLY")
env = CustomerApiEnv(task=_task())

injector.before_step(env, 3)
assert not injector.triggered
assert memory.get(TARGET_PLAN_MEMORY_KEY) == "GOLD_ANNUAL"

injector.before_step(env, 4)
assert injector.triggered
assert memory.get(TARGET_PLAN_MEMORY_KEY) == "BASIC_MONTHLY"


def test_stale_memory_causes_failed_subscription_update() -> None:
env = CustomerApiEnv(task=_task())
memory = AgentMemory()
agent = ScriptedApiAgent(env=env, memory=memory)
injector = StaleMemoryInjector(memory=memory, trigger_step=4)
registry = FailureInjectorRegistry()
registry.register("memory.stale", injector)
collector = TraceCollector(task_id="stale-mem-test")

for _ in range(8):
upcoming = agent._step + 1
registry.dispatch_before_step(env, upcoming)
action = agent.next_action()
if action is None:
break
obs = env.step(action)
registry.dispatch_after_step(env, upcoming, obs)
collector.record(action, obs, agent.expected_attributes(), upcoming / 8)

assert injector.triggered
assert not env.validate_success()
assert first_memory_conflict_step(collector.events) == 4
assert any(memory_live_conflict(e) for e in collector.events)


def test_memory_live_conflict_feature() -> None:
from datetime import UTC, datetime

conflict = TraceEvent(
event_id="e1",
timestamp=datetime.now(UTC),
task_id="t",
scaffold="s",
entity="tool",
name="get_plan",
expectation=Expectation(
attributes={
"memory_plan_id": "BASIC_MONTHLY",
"live_target_plan_id": "GOLD_ANNUAL",
"plan_id": "BASIC_MONTHLY",
}
),
attributes={"step": 4},
)
aligned = TraceEvent(
event_id="e2",
timestamp=datetime.now(UTC),
task_id="t",
scaffold="s",
entity="tool",
name="get_plan",
expectation=Expectation(
attributes={
"memory_plan_id": "GOLD_ANNUAL",
"live_target_plan_id": "GOLD_ANNUAL",
}
),
attributes={"step": 2},
)
assert memory_live_conflict(conflict) is True
assert memory_live_conflict(aligned) is False


def test_stale_memory_case_loads_with_root_cause() -> None:
path = (
Path(__file__).resolve().parents[2]
/ "agentfailbench"
/ "failures"
/ "memory"
/ "stale-memory-001.yaml"
)
registry = CaseRegistry.from_yaml_file(path)
case = registry.get("stale-memory-001")
assert case.model.ground_truth.root_cause == RootCauseCode.STALE_MEMORY
assert case.root_cause_label.category == FailureCategory.MEMORY


def test_memory_rollback_restores_checkpoint() -> None:
memory = AgentMemory()
memory.set(TARGET_PLAN_MEMORY_KEY, "GOLD_ANNUAL")
memory.checkpoint()
memory.set(TARGET_PLAN_MEMORY_KEY, "BASIC_MONTHLY")
assert memory.rollback() is True
assert memory.get(TARGET_PLAN_MEMORY_KEY) == "GOLD_ANNUAL"
Loading