From 0f1e8c1c00fc08a4c1c7e3d3c05bb10b4e155e64 Mon Sep 17 00:00:00 2001 From: Abhinaysai Kamineni <66816045+askmy-stack@users.noreply.github.com> Date: Fri, 24 Jul 2026 11:06:22 -0400 Subject: [PATCH] Add MemoryGuard stale-memory injector and conflict signal. Introduce AgentMemory with checkpoints, StaleMemoryInjector at trigger step, stale-memory-001 case, and a memory-vs-live-target feature for early detection without exceptions. Co-authored-by: Cursor --- agentfailbench/agents/scripted.py | 24 ++- agentfailbench/failures/memory/__init__.py | 8 + agentfailbench/failures/memory/injector.py | 51 +++++++ .../failures/memory/stale-memory-001.yaml | 25 ++++ agentfailbench/memory/__init__.py | 5 + agentfailbench/memory/store.py | 36 +++++ runtime/features/__init__.py | 8 +- runtime/features/memory_conflict.py | 30 ++++ runtime/schemas/root_cause.py | 2 + tests/unit/test_root_cause.py | 5 + tests/unit/test_stale_memory.py | 137 ++++++++++++++++++ 11 files changed, 328 insertions(+), 3 deletions(-) create mode 100644 agentfailbench/failures/memory/injector.py create mode 100644 agentfailbench/failures/memory/stale-memory-001.yaml create mode 100644 agentfailbench/memory/__init__.py create mode 100644 agentfailbench/memory/store.py create mode 100644 runtime/features/memory_conflict.py create mode 100644 tests/unit/test_stale_memory.py diff --git a/agentfailbench/agents/scripted.py b/agentfailbench/agents/scripted.py index e027ac8..95de9ad 100644 --- a/agentfailbench/agents/scripted.py +++ b/agentfailbench/agents/scripted.py @@ -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]: @@ -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 diff --git a/agentfailbench/failures/memory/__init__.py b/agentfailbench/failures/memory/__init__.py index e69de29..50c586e 100644 --- a/agentfailbench/failures/memory/__init__.py +++ b/agentfailbench/failures/memory/__init__.py @@ -0,0 +1,8 @@ +"""MemoryGuard failure suite.""" + +from agentfailbench.failures.memory.injector import ( + TARGET_PLAN_MEMORY_KEY, + StaleMemoryInjector, +) + +__all__ = ["TARGET_PLAN_MEMORY_KEY", "StaleMemoryInjector"] diff --git a/agentfailbench/failures/memory/injector.py b/agentfailbench/failures/memory/injector.py new file mode 100644 index 0000000..dcc4df4 --- /dev/null +++ b/agentfailbench/failures/memory/injector.py @@ -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 diff --git a/agentfailbench/failures/memory/stale-memory-001.yaml b/agentfailbench/failures/memory/stale-memory-001.yaml new file mode 100644 index 0000000..5205da6 --- /dev/null +++ b/agentfailbench/failures/memory/stale-memory-001.yaml @@ -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 diff --git a/agentfailbench/memory/__init__.py b/agentfailbench/memory/__init__.py new file mode 100644 index 0000000..5768a87 --- /dev/null +++ b/agentfailbench/memory/__init__.py @@ -0,0 +1,5 @@ +"""Agent memory utilities.""" + +from agentfailbench.memory.store import AgentMemory + +__all__ = ["AgentMemory"] diff --git a/agentfailbench/memory/store.py b/agentfailbench/memory/store.py new file mode 100644 index 0000000..af7bdcf --- /dev/null +++ b/agentfailbench/memory/store.py @@ -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 diff --git a/runtime/features/__init__.py b/runtime/features/__init__.py index 761f5a2..6883d5c 100644 --- a/runtime/features/__init__.py +++ b/runtime/features/__init__.py @@ -1,5 +1,11 @@ """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 -__all__ = ["first_mismatch_step", "semantic_mismatch"] +__all__ = [ + "first_memory_conflict_step", + "first_mismatch_step", + "memory_live_conflict", + "semantic_mismatch", +] diff --git a/runtime/features/memory_conflict.py b/runtime/features/memory_conflict.py new file mode 100644 index 0000000..dd59638 --- /dev/null +++ b/runtime/features/memory_conflict.py @@ -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 diff --git a/runtime/schemas/root_cause.py b/runtime/schemas/root_cause.py index a96653f..e8b646e 100644 --- a/runtime/schemas/root_cause.py +++ b/runtime/schemas/root_cause.py @@ -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, } diff --git a/tests/unit/test_root_cause.py b/tests/unit/test_root_cause.py index 0b2a1b8..ec9f629 100644 --- a/tests/unit/test_root_cause.py +++ b/tests/unit/test_root_cause.py @@ -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 diff --git a/tests/unit/test_stale_memory.py b/tests/unit/test_stale_memory.py new file mode 100644 index 0000000..3464349 --- /dev/null +++ b/tests/unit/test_stale_memory.py @@ -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"