diff --git a/services/api/app/main.py b/services/api/app/main.py index f833312..2de6f53 100644 --- a/services/api/app/main.py +++ b/services/api/app/main.py @@ -27,6 +27,7 @@ from app.routes.skills import context_router as skills_context_router from app.routes.skills import router as skills_router from app.models import memory, audit, agent_config +from app.models import memory_forget # noqa: F401 - registers the table for create_all from app.models import auth_setting from app.models import evidence_source, evidence_object, analysis_object from app.models import episode_object, object_link diff --git a/services/api/app/models/memory_forget.py b/services/api/app/models/memory_forget.py new file mode 100644 index 0000000..e429732 --- /dev/null +++ b/services/api/app/models/memory_forget.py @@ -0,0 +1,18 @@ +"""Durable owner-forget receipts: identity and outcome only, never the forgotten text.""" + +from app.core.db import Base +from sqlalchemy import DateTime, Integer, String, func +from sqlalchemy.orm import Mapped, mapped_column + + +class MemoryForget(Base): + __tablename__ = "memory_forgets" + agent_id: Mapped[str] = mapped_column(String(128), primary_key=True) + request_id: Mapped[str] = mapped_column(String(128), primary_key=True) + memory_id: Mapped[str] = mapped_column(String(200), nullable=False, index=True) + payload_hash: Mapped[str] = mapped_column(String(64), nullable=False) + revision: Mapped[int] = mapped_column(Integer, nullable=False) + index_removal: Mapped[str] = mapped_column(String(24), nullable=False, default="pending") + created_at: Mapped[DateTime] = mapped_column( + DateTime(timezone=True), server_default=func.now() + ) diff --git a/services/api/app/routes/corrections.py b/services/api/app/routes/corrections.py index abf506a..6142038 100644 --- a/services/api/app/routes/corrections.py +++ b/services/api/app/routes/corrections.py @@ -5,7 +5,7 @@ import secrets from typing import Annotated -from app.services import memory_corrections +from app.services import memory_corrections, memory_forgetting from fastapi import APIRouter, Depends, Header, HTTPException, Path from fastapi.exceptions import RequestValidationError from fastapi.responses import JSONResponse @@ -78,6 +78,12 @@ def nonblank(cls, value): return value +class Forget(BaseModel): + model_config = ConfigDict(extra="forbid", strict=True) + memory_id: str = Field(pattern=MEMORY_ID) + expected_revision: int = Field(ge=1, le=2147483646) + + router = APIRouter( prefix="/runtime/corrections", tags=["reviewed memory correction"], @@ -93,6 +99,24 @@ def memory( return memory_corrections.read_memory(agent_id, memory_id) +@router.put("/forget/{request_id}") +def forget( + request_id: Annotated[str, Path(pattern=REQUEST_ID)], + payload: Forget, + agent_id: str = Depends(require_correction_key), +): + """Owner-reviewed forget: removes the memory and every stored copy of its text.""" + return memory_forgetting.forget(agent_id, request_id, payload.model_dump()) + + +@router.get("/forget/{request_id}") +def forget_receipt( + request_id: Annotated[str, Path(pattern=REQUEST_ID)], + agent_id: str = Depends(require_correction_key), +): + return memory_forgetting.read_receipt(agent_id, request_id) + + @router.put("/{request_id}") def apply( request_id: Annotated[str, Path(pattern=REQUEST_ID)], diff --git a/services/api/app/routes/memory.py b/services/api/app/routes/memory.py index a4e490e..4edbd13 100644 --- a/services/api/app/routes/memory.py +++ b/services/api/app/routes/memory.py @@ -498,10 +498,9 @@ def delete_memory(memory_id: str, agent_id: str = Depends(get_agent_id), raise HTTPException(404, "Memory not found") _check_revision(row, expected_revision) - from app.models.deletion_receipt import record - record(db, agent_id, "memory", row.id) - db.add(MemoryAudit(action="delete", memory_id=row.id, payload_json=json.dumps({"text": row.text}))) - db.delete(row) + from app.services.memory_forgetting import erase + # A deleted memory's text must not survive in revisions, conflicts or audit. + erase(db, agent_id, row, "delete", {"revision": row.revision}) db.commit() try: diff --git a/services/api/app/services/memory_forgetting.py b/services/api/app/services/memory_forgetting.py new file mode 100644 index 0000000..bcf3cac --- /dev/null +++ b/services/api/app/services/memory_forgetting.py @@ -0,0 +1,138 @@ +"""Forgetting a memory removes every stored copy of its text. + +The memory row, its revision snapshots, conflict records naming it, and the text in +its audit entries are removed or redacted in one transaction. A content-free audit +entry and a deletion receipt (for recovery replay) record that it happened. The +vector point is removed after commit; a failed removal leaves a visible receipt. + +The source conversation is not touched: forgetting a memory is not forgetting the +chat it came from. +""" + +import hashlib +import json + +from app.core.db import SessionLocal +from app.models.audit import MemoryAudit +from app.models import deletion_receipt +from app.models.memory import Memory +from app.models.memory_conflict import MemoryConflict +from app.models.memory_forget import MemoryForget +from app.models.memory_revision import MemoryRevision +from app.services.qdrant_store import delete_memory_embedding, index_after_commit +from fastapi import HTTPException +from sqlalchemy import delete, or_, select, text, update +from sqlalchemy.exc import IntegrityError, SQLAlchemyError + + +def erase(db, agent_id: str, memory: Memory, action: str, detail: dict) -> None: + """Remove the memory and every text copy inside the caller's transaction.""" + db.execute(delete(MemoryRevision).where(MemoryRevision.memory_id == memory.id)) + db.execute( + delete(MemoryConflict).where( + or_( + MemoryConflict.memory_id == memory.id, + MemoryConflict.conflicting_memory_id == memory.id, + ) + ) + ) + db.execute( + update(MemoryAudit) + .where(MemoryAudit.memory_id == memory.id) + .values(payload_json=json.dumps({"redacted": action})) + ) + deletion_receipt.record(db, agent_id, "memory", memory.id) + db.add(MemoryAudit(action=action, memory_id=memory.id, payload_json=json.dumps(detail))) + db.delete(memory) + + +def _view(row): + return { + "request_id": row.request_id, + "memory_id": row.memory_id, + "agent_id": row.agent_id, + "revision": row.revision, + "status": "forgotten", + "index_removal": row.index_removal, + } + + +def _prior(db, agent_id, request_id, digest=None): + row = db.get(MemoryForget, (agent_id, request_id)) + if row is not None and digest is not None and row.payload_hash != digest: + raise HTTPException(409, "Forget request identity already used") + return row + + +def read_receipt(agent_id, request_id): + with SessionLocal() as db: + row = _prior(db, agent_id, request_id) + if row is None: + raise HTTPException(404, "Forget receipt not found") + return _view(row) + + +def forget(agent_id, request_id, payload): + digest = hashlib.sha256( + json.dumps( + {"agent_id": agent_id, **payload}, sort_keys=True, separators=(",", ":") + ).encode() + ).hexdigest() + try: + with SessionLocal() as db: + if db.get_bind().dialect.name == "sqlite": + db.execute(text("BEGIN IMMEDIATE")) + prior = _prior(db, agent_id, request_id, digest) + if prior is not None: + return _view(prior) + memory = db.execute( + select(Memory) + .where(Memory.id == payload["memory_id"], Memory.agent_id == agent_id) + .with_for_update() + ).scalar_one_or_none() + if memory is None: + raise HTTPException(404, "Memory not found") + prior = _prior(db, agent_id, request_id, digest) + if prior is not None: + return _view(prior) + if memory.revision != payload["expected_revision"]: + raise HTTPException(409, "Memory revision changed; review the current record") + receipt = MemoryForget( + agent_id=agent_id, + request_id=request_id, + memory_id=memory.id, + payload_hash=digest, + revision=memory.revision, + index_removal="pending", + ) + erase( + db, + agent_id, + memory, + "owner_forget", + {"request_id": request_id, "revision": memory.revision}, + ) + db.add(receipt) + db.commit() + result = _view(receipt) + except IntegrityError: + with SessionLocal() as db: + prior = _prior(db, agent_id, request_id, digest) + if prior is not None: + return _view(prior) + raise HTTPException(409, "Memory or forget request changed") from None + + removal = index_after_commit(delete_memory_embedding, payload["memory_id"]) + state = "removed" if removal.get("status") == "ok" else "degraded" + try: + with SessionLocal() as db: + db.execute( + update(MemoryForget) + .where(MemoryForget.agent_id == agent_id, MemoryForget.request_id == request_id) + .values(index_removal=state) + ) + db.commit() + result["index_removal"] = state + except SQLAlchemyError: + pass # The forget is durable; the receipt stays visibly pending. + return result diff --git a/services/api/tests/test_memory_forgetting.py b/services/api/tests/test_memory_forgetting.py new file mode 100644 index 0000000..fe3525a --- /dev/null +++ b/services/api/tests/test_memory_forgetting.py @@ -0,0 +1,143 @@ +"""Forgetting removes every stored copy of a memory's text. Temporary SQLite; no network.""" + +import json + +import pytest +from app.core.db import Base +from app.models.audit import MemoryAudit +from app.models.deletion_receipt import DeletionReceipt +from app.models.memory import Memory +from app.models.memory_conflict import MemoryConflict +from app.models.memory_forget import MemoryForget +from app.models.memory_revision import MemoryRevision +from app.routes import corrections +from app.routes import memory as memory_routes +from app.services import memory_corrections, memory_forgetting +from fastapi import FastAPI +from fastapi.testclient import TestClient +from sqlalchemy import create_engine, inspect, select, text +from sqlalchemy.orm import sessionmaker + +KEY = "synthetic-correction-key-for-tests-12345" +HEADERS = {"X-MemoryGate-Correction-Key": KEY} +SECRET = "Owner lives at 12 Secret Street" + + +@pytest.fixture +def setup(tmp_path, monkeypatch): + engine = create_engine(f"sqlite:///{tmp_path / 'forget.db'}", connect_args={"check_same_thread": False}) + Base.metadata.create_all(engine) + sessions = sessionmaker(engine, autoflush=False) + monkeypatch.setenv("MEMORYGATE_CORRECTION_KEY", KEY) + monkeypatch.setenv("MEMORYGATE_CORRECTION_AGENT_ID", "owner") + for module in (memory_corrections, memory_forgetting, memory_routes): + monkeypatch.setattr(module, "SessionLocal", sessions) + removed = [] + monkeypatch.setattr(memory_corrections, "index_after_commit", lambda *a, **k: {"status": "ok"}) + monkeypatch.setattr(memory_forgetting, "delete_memory_embedding", removed.append) + monkeypatch.setattr(memory_routes, "delete_memory_embedding", removed.append) + with sessions() as db: + for identity, agent, value in (("one", "owner", SECRET), ("two", "owner", "Likes tea"), ("foreign", "other", SECRET)): + db.add(Memory(id=identity, agent_id=agent, text=value, summary=value, source_type="stated", confidence="high")) + db.add(MemoryAudit(action="write", memory_id="one", payload_json=json.dumps({"text": SECRET}))) + db.add(MemoryConflict(agent_id="owner", memory_id="two", conflicting_memory_id="one", reason=f"Conflicts with {SECRET}")) + db.commit() + app = FastAPI() + app.include_router(corrections.router) + with TestClient(app, raise_server_exceptions=False) as client: + yield client, sessions, engine, removed + engine.dispose() + + +def forget(http, request_id="forget_request_0001", **changes): + return http.put( + "/runtime/corrections/forget/" + request_id, + headers=HEADERS, + json={"memory_id": "one", "expected_revision": 2, **changes}, + ) + + +def texts_everywhere(engine): + """Every text value in every table, to prove the secret survives nowhere.""" + found = [] + with engine.connect() as db: + for table in inspect(engine).get_table_names(): + for row in db.execute(text(f'SELECT * FROM "{table}"')): + found.extend(str(value) for value in row) + return " ".join(found) + + +def test_forget_removes_the_memory_and_every_copy_of_its_text(setup): + http, sessions, engine, removed = setup + # A correction stores the old text in a revision snapshot; forgetting must remove it. + corrected = http.put( + "/runtime/corrections/correction_request_01", + headers=HEADERS, + json={"memory_id": "one", "expected_revision": 1, "text": SECRET + " (corrected)"}, + ) + assert corrected.status_code == 200 + assert SECRET in texts_everywhere(engine) + + response = forget(http) + assert response.status_code == 200 + expected = {"request_id": "forget_request_0001", "memory_id": "one", "agent_id": "owner", "revision": 2, "status": "forgotten", "index_removal": "removed"} + assert response.json() == expected + assert removed == ["one"] + # Only the other agent's unrelated memory still holds that sentence. + with sessions() as db: + db.delete(db.get(Memory, "foreign")) + db.commit() + assert SECRET not in texts_everywhere(engine) + with sessions() as db: + assert db.get(Memory, "one") is None + assert db.get(Memory, "two") is not None + assert not db.scalars(select(MemoryRevision).where(MemoryRevision.memory_id == "one")).all() + assert not db.scalars(select(MemoryConflict)).all() + assert db.get(DeletionReceipt, ("owner", "memory", "one")) is not None + actions = [(row.action, json.loads(row.payload_json)) for row in db.scalars(select(MemoryAudit).where(MemoryAudit.memory_id == "one"))] + assert ("owner_forget", {"request_id": "forget_request_0001", "revision": 2}) in actions + assert all("text" not in payload for _, payload in actions) + + # Replays return the saved receipt; a different request for the gone memory is 404. + assert forget(http).json() == expected + assert http.get("/runtime/corrections/forget/forget_request_0001", headers=HEADERS).json() == expected + assert forget(http, request_id="forget_request_0002").status_code == 404 + assert forget(http, expected_revision=3).status_code == 409 + assert removed == ["one"] + + +def test_forget_is_bounded_by_revision_namespace_and_capability(setup): + http, sessions, _, removed = setup + assert forget(http, expected_revision=1, memory_id="two").status_code == 200 + assert forget(http, request_id="forget_request_0003", memory_id="one", expected_revision=5).status_code == 409 + assert forget(http, request_id="forget_request_0004", memory_id="foreign", expected_revision=1).status_code == 404 + assert http.put("/runtime/corrections/forget/forget_request_0005", json={"memory_id": "one", "expected_revision": 1}).status_code == 401 + wrong = http.put("/runtime/corrections/forget/forget_request_0006", headers={**HEADERS, "X-Agent-Id": "other"}, json={"memory_id": "one", "expected_revision": 1}) + assert wrong.status_code == 403 + assert http.put("/runtime/corrections/forget/forget_request_0007", headers=HEADERS, json={"memory_id": "one", "expected_revision": 1, "text": "extra"}).status_code == 422 + with sessions() as db: + assert db.get(Memory, "one") is not None and db.get(Memory, "foreign") is not None + assert db.get(MemoryForget, ("owner", "forget_request_0003")) is None + assert removed == ["two"] + + +def test_index_outage_keeps_the_forget_and_says_so(setup, monkeypatch): + http, sessions, _, _ = setup + + def unreachable(memory_id): + raise ConnectionError("qdrant down") + + monkeypatch.setattr(memory_forgetting, "delete_memory_embedding", unreachable) + response = http.put("/runtime/corrections/forget/forget_request_0008", headers=HEADERS, json={"memory_id": "two", "expected_revision": 1}) + assert response.status_code == 200 and response.json()["index_removal"] == "degraded" + with sessions() as db: + assert db.get(Memory, "two") is None + + +def test_admin_delete_no_longer_copies_text_into_the_audit(setup): + _, sessions, engine, _ = setup + memory_routes.delete_memory("one", agent_id="owner", expected_revision=None) + with sessions() as db: + db.delete(db.get(Memory, "foreign")) + db.commit() + assert SECRET not in texts_everywhere(engine)