From f11de7b2737e63fb5a6367446b5a9e5035c9039a Mon Sep 17 00:00:00 2001
From: Atharva-Kanherkar
Date: Mon, 30 Mar 2026 03:10:05 +0530
Subject: [PATCH 1/4] Add MCP schema contracts and serializers
---
api/app.py | 108 +--------
mcp_server/__init__.py | 47 ++++
mcp_server/schemas.py | 486 ++++++++++++++++++++++++++++++++++++++
tests/test_mcp_schemas.py | 198 ++++++++++++++++
4 files changed, 743 insertions(+), 96 deletions(-)
create mode 100644 mcp_server/__init__.py
create mode 100644 mcp_server/schemas.py
create mode 100644 tests/test_mcp_schemas.py
diff --git a/api/app.py b/api/app.py
index 5e41082..33880af 100644
--- a/api/app.py
+++ b/api/app.py
@@ -16,6 +16,13 @@
from fastapi.middleware.cors import CORSMiddleware
import config
+from mcp_server.schemas import (
+ serialise_contradiction_candidate as _mcp_serialise_contradiction_candidate,
+ serialise_forgetting_report as _mcp_serialise_forgetting_report,
+ serialise_procedural_match as _mcp_serialise_procedural_match,
+ serialise_ranked_result as _mcp_serialise_ranked_result,
+ serialise_record as _mcp_serialise_record,
+)
from models.base import MemoryRecord, normalize_modality
from models.episodic import EpisodicMemory
from models.procedural import ProceduralMemory
@@ -197,75 +204,15 @@ def _cleanup_owned_media(service: MemoryAPIService, media_ref: str | None) -> No
def _serialise_record(record: MemoryRecord) -> dict[str, Any]:
- payload = {
- "id": record.id,
- "memory_type": record.memory_type,
- "modality": record.modality,
- "content": record.content,
- "created_at": record.created_at.isoformat(),
- "last_accessed_at": record.last_accessed_at.isoformat() if record.last_accessed_at else None,
- "access_count": record.access_count,
- "importance": record.importance,
- "media_ref": record.media_ref,
- "media_type": record.media_type,
- "text_description": record.text_description,
- "has_media": record.has_media,
- }
- if isinstance(record, SemanticMemory):
- payload.update(
- {
- "category": record.category,
- "domain": record.domain,
- "confidence": record.confidence,
- "supersedes": record.supersedes,
- "related_ids": record.related_ids,
- "has_visual": record.has_visual,
- }
- )
- if isinstance(record, EpisodicMemory):
- payload.update(
- {
- "session_id": record.session_id,
- "turn_number": record.turn_number,
- "participants": record.participants,
- "summary": record.summary,
- "emotional_valence": record.emotional_valence,
- "emotional_profile": record.emotional_profile,
- "source_mime_type": record.source_mime_type,
- }
- )
- if isinstance(record, ProceduralMemory):
- payload.update(
- {
- "steps": record.steps,
- "preconditions": record.preconditions,
- "success_count": record.success_count,
- "failure_count": record.failure_count,
- "total_outcomes": record.total_outcomes,
- "success_rate": record.success_rate,
- "wilson_score": record.wilson_score,
- }
- )
- return payload
+ return _mcp_serialise_record(record).model_dump(mode="json")
def _serialise_ranked_result(result) -> dict[str, Any]:
- return {
- "record": _serialise_record(result.record),
- "raw_similarity": result.raw_similarity,
- "recency_score": result.recency_score,
- "importance_score": result.importance_score,
- "final_score": result.final_score,
- }
+ return _mcp_serialise_ranked_result(result).model_dump(mode="json")
def _serialise_procedural_match(match: ProceduralMatch) -> dict[str, Any]:
- return {
- "record": _serialise_record(match.record),
- "similarity": match.similarity,
- "wilson_score": match.wilson_score,
- "combined_score": match.combined_score,
- }
+ return _mcp_serialise_procedural_match(match).model_dump(mode="json")
class MemoryAPIService:
@@ -351,42 +298,11 @@ def _safe_contradiction_lookup(
def _serialise_contradiction_candidate(candidate: ContradictionCandidate) -> dict[str, Any]:
- return {
- "record": _serialise_record(candidate.record),
- "similarity": candidate.similarity,
- }
+ return _mcp_serialise_contradiction_candidate(candidate).model_dump(mode="json")
def _serialise_forgetting_report(report: ForgettingReport) -> dict[str, Any]:
- return {
- "status": "completed",
- "dry_run": report.dry_run,
- "scanned": report.scanned,
- "kept": report.kept,
- "faded": report.faded,
- "pruned": report.pruned,
- "media_deleted": report.media_deleted,
- "duplicates_flagged": report.duplicates_flagged,
- "skipped_records": report.skipped_records,
- "skipped_media": report.skipped_media,
- "by_type": dict(report.by_type),
- "decisions": [
- {
- "record_id": d.record_id,
- "memory_type": d.memory_type,
- "action": d.action,
- "reason": d.reason,
- "score": d.score,
- "media_deleted": d.media_deleted,
- "executed": d.executed,
- "record_skip_reason": d.record_skip_reason,
- "media_skip_reason": d.media_skip_reason,
- "old_importance": d.old_importance,
- "new_importance": d.new_importance,
- }
- for d in report.decisions
- ],
- }
+ return _mcp_serialise_forgetting_report(report).model_dump(mode="json")
def create_app(
diff --git a/mcp_server/__init__.py b/mcp_server/__init__.py
new file mode 100644
index 0000000..1f2f97c
--- /dev/null
+++ b/mcp_server/__init__.py
@@ -0,0 +1,47 @@
+"""Shared MCP contracts and helpers for the memory runtime."""
+
+from .schemas import (
+ ContradictionCandidateModel,
+ ContradictionCheckModel,
+ ForgettingCycleRequest,
+ ForgettingReportModel,
+ GetMemoryRequest,
+ MCPToolError,
+ MemoryOverviewModel,
+ MemoryRecordModel,
+ ProcedureMatchModel,
+ RecallEpisodesRequestAdapter,
+ RememberFactResponse,
+ ToolErrorModel,
+ get_memory_by_type,
+ serialise_contradiction_candidate,
+ serialise_forgetting_report,
+ serialise_memory_overview,
+ serialise_procedural_match,
+ serialise_ranked_result,
+ serialise_record,
+ serialise_remember_fact_response,
+)
+
+__all__ = [
+ "ContradictionCandidateModel",
+ "ContradictionCheckModel",
+ "ForgettingCycleRequest",
+ "ForgettingReportModel",
+ "GetMemoryRequest",
+ "MCPToolError",
+ "MemoryOverviewModel",
+ "MemoryRecordModel",
+ "ProcedureMatchModel",
+ "RecallEpisodesRequestAdapter",
+ "RememberFactResponse",
+ "ToolErrorModel",
+ "get_memory_by_type",
+ "serialise_contradiction_candidate",
+ "serialise_forgetting_report",
+ "serialise_memory_overview",
+ "serialise_procedural_match",
+ "serialise_ranked_result",
+ "serialise_record",
+ "serialise_remember_fact_response",
+]
diff --git a/mcp_server/schemas.py b/mcp_server/schemas.py
new file mode 100644
index 0000000..804f3a8
--- /dev/null
+++ b/mcp_server/schemas.py
@@ -0,0 +1,486 @@
+from __future__ import annotations
+
+from datetime import datetime
+from typing import Annotated, Any, Literal
+
+from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, model_validator
+
+from forgetting.contradiction import ContradictionCandidate
+from forgetting.service import ForgettingReport
+from models.base import MemoryRecord
+from models.episodic import EpisodicMemory
+from models.procedural import ProceduralMemory
+from models.semantic import SemanticMemory
+from stores.procedural_store import ProceduralMatch
+
+MemoryType = Literal["semantic", "episodic", "procedural"]
+ForgettingAction = Literal["keep", "fade", "prune"]
+ContradictionStatus = Literal["completed", "skipped", "error"]
+
+DEFAULT_RETRIEVAL_TOP_K = 5
+MAX_RETRIEVAL_TOP_K = 20
+DEFAULT_RECENT_EPISODES = 5
+MAX_RECENT_EPISODES = 50
+DEFAULT_MAX_DECISIONS = 50
+MAX_MAX_DECISIONS = 100
+DEFAULT_MAX_CONTRADICTIONS = 5
+MAX_MAX_CONTRADICTIONS = 10
+
+
+class MCPModel(BaseModel):
+ model_config = ConfigDict(extra="forbid", frozen=True)
+
+
+class ToolErrorModel(MCPModel):
+ code: str
+ message: str
+ retryable: bool = False
+ details: dict[str, Any] | None = None
+
+
+class MCPToolError(Exception):
+ def __init__(
+ self,
+ *,
+ code: str,
+ message: str,
+ retryable: bool = False,
+ details: dict[str, Any] | None = None,
+ ) -> None:
+ super().__init__(message)
+ self.payload = ToolErrorModel(
+ code=code,
+ message=message,
+ retryable=retryable,
+ details=details,
+ )
+
+
+class InlineContentModel(MCPModel):
+ inline_base64: str = Field(min_length=1)
+ mime_type: str = Field(min_length=1)
+ filename: str | None = None
+
+
+class MediaInputModel(MCPModel):
+ file_path: str | None = None
+ inline_content: InlineContentModel | None = None
+
+ @model_validator(mode="after")
+ def validate_source(self) -> "MediaInputModel":
+ if self.file_path and self.inline_content:
+ raise ValueError("Provide either file_path or inline_content, not both")
+ if not self.file_path and self.inline_content is None:
+ raise ValueError("A media input requires file_path or inline_content")
+ return self
+
+
+class MemoryRecordModel(MCPModel):
+ id: str
+ memory_type: MemoryType
+ modality: str
+ content: str
+ created_at: datetime
+ last_accessed_at: datetime | None = None
+ access_count: int
+ importance: float
+ media_ref: str | None = None
+ media_type: str | None = None
+ text_description: str | None = None
+ has_media: bool
+ category: str | None = None
+ domain: str | None = None
+ confidence: float | None = None
+ supersedes: str | None = None
+ superseded_by: str | None = None
+ related_ids: list[str] | None = None
+ has_visual: bool | None = None
+ session_id: str | None = None
+ turn_number: int | None = None
+ participants: list[str] | None = None
+ summary: str | None = None
+ emotional_valence: float | None = None
+ emotional_profile: dict[str, float] | None = None
+ source_mime_type: str | None = None
+ steps: list[str] | None = None
+ preconditions: list[str] | None = None
+ success_count: int | None = None
+ failure_count: int | None = None
+ total_outcomes: int | None = None
+ success_rate: float | None = None
+ wilson_score: float | None = None
+
+
+class RankedMemoryResultModel(MCPModel):
+ record: MemoryRecordModel
+ raw_similarity: float
+ recency_score: float
+ importance_score: float
+ final_score: float
+
+
+class ProcedureMatchModel(MCPModel):
+ record: MemoryRecordModel
+ similarity: float
+ wilson_score: float
+ combined_score: float
+
+
+class ContradictionCandidateModel(MCPModel):
+ record: MemoryRecordModel
+ similarity: float
+
+
+class ContradictionCheckModel(MCPModel):
+ status: ContradictionStatus
+ detail: str | None = None
+
+
+class RememberFactRequest(MCPModel):
+ content: str = Field(min_length=1)
+ importance: float = Field(default=0.5)
+ category: str = Field(default="general", min_length=1)
+ domain: str | None = None
+ confidence: float = Field(default=1.0)
+ supersedes: str | None = None
+ related_ids: list[str] = Field(default_factory=list)
+ has_visual: bool = False
+ modality: str = "text"
+ media: MediaInputModel | None = None
+ media_type: Literal["image", "audio", "video", "pdf"] | None = None
+ text_description: str | None = None
+ max_contradictions: int = Field(
+ default=DEFAULT_MAX_CONTRADICTIONS,
+ ge=1,
+ le=MAX_MAX_CONTRADICTIONS,
+ )
+
+
+class RememberEpisodeRequest(MCPModel):
+ session_id: str = Field(min_length=1)
+ text: str = Field(min_length=1)
+ turn_number: int | None = None
+ participants: list[str] = Field(default_factory=lambda: ["user", "agent"])
+ summary: str | None = None
+ emotional_profile: dict[str, float] = Field(default_factory=dict)
+ emotional_valence: float | None = None
+ importance: float = Field(default=0.5)
+ modality: str = "text"
+ media: MediaInputModel | None = None
+ media_type: Literal["image", "audio", "video", "pdf"] | None = None
+ source_mime_type: str | None = None
+
+
+class RememberProcedureRequest(MCPModel):
+ content: str = Field(min_length=1)
+ steps: list[str] = Field(min_length=1)
+ preconditions: list[str] = Field(default_factory=list)
+ importance: float = Field(default=0.5)
+ modality: str = "text"
+ media: MediaInputModel | None = None
+ media_type: Literal["image", "audio", "video", "pdf"] | None = None
+ text_description: str | None = None
+
+
+class RememberFactResponse(MCPModel):
+ record: MemoryRecordModel
+ potential_contradictions: list[ContradictionCandidateModel]
+ contradiction_check: ContradictionCheckModel
+ contradiction_check_order: Literal["after_store"] = "after_store"
+ total_contradictions: int
+ contradictions_truncated: bool
+
+
+class RememberEpisodeResponse(MCPModel):
+ record: MemoryRecordModel
+
+
+class RememberProcedureResponse(MCPModel):
+ record: MemoryRecordModel
+
+
+class RecallMemoriesRequest(MCPModel):
+ query: str = Field(min_length=1)
+ top_k: int = Field(default=DEFAULT_RETRIEVAL_TOP_K, ge=1, le=MAX_RETRIEVAL_TOP_K)
+ memory_types: list[MemoryType] | None = None
+
+
+class RecallMemoriesResponse(MCPModel):
+ results: list[RankedMemoryResultModel]
+
+
+class RecallRecentEpisodesRequest(MCPModel):
+ mode: Literal["recent"]
+ limit: int = Field(default=DEFAULT_RECENT_EPISODES, ge=1, le=MAX_RECENT_EPISODES)
+
+
+class RecallSessionEpisodesRequest(MCPModel):
+ mode: Literal["session"]
+ session_id: str = Field(min_length=1)
+
+
+class RecallTimeRangeEpisodesRequest(MCPModel):
+ mode: Literal["time_range"]
+ start: datetime
+ end: datetime
+
+ @model_validator(mode="after")
+ def validate_range(self) -> "RecallTimeRangeEpisodesRequest":
+ if self.start > self.end:
+ raise ValueError("start must be earlier than or equal to end")
+ return self
+
+
+RecallEpisodesRequest = Annotated[
+ RecallRecentEpisodesRequest | RecallSessionEpisodesRequest | RecallTimeRangeEpisodesRequest,
+ Field(discriminator="mode"),
+]
+RecallEpisodesRequestAdapter = TypeAdapter(RecallEpisodesRequest)
+
+
+class RecallEpisodesResponse(MCPModel):
+ mode: Literal["recent", "session", "time_range"]
+ records: list[MemoryRecordModel]
+
+
+class RecallProceduresRequest(MCPModel):
+ task: str = Field(min_length=1)
+ top_k: int = Field(default=3, ge=1, le=MAX_RETRIEVAL_TOP_K)
+
+
+class RecallProceduresResponse(MCPModel):
+ ranking_basis: Literal["combined_score"] = "combined_score"
+ results: list[ProcedureMatchModel]
+
+
+class GetMemoryRequest(MCPModel):
+ memory_type: MemoryType
+ record_id: str = Field(min_length=1)
+
+
+class GetMemoryResponse(MCPModel):
+ record: MemoryRecordModel | None = None
+
+
+class ResolveContradictionRequest(MCPModel):
+ keep_id: str = Field(min_length=1)
+ supersede_id: str = Field(min_length=1)
+
+
+class ResolveContradictionResponse(MCPModel):
+ kept_id: str
+ superseded_id: str
+ status: Literal["resolved"] = "resolved"
+
+
+class RecordProcedureOutcomeRequest(MCPModel):
+ record_id: str = Field(min_length=1)
+ success: bool
+
+
+class RecordProcedureOutcomeResponse(MCPModel):
+ record: MemoryRecordModel
+
+
+class MemoryOverviewModel(MCPModel):
+ semantic_count: int
+ episodic_count: int
+ procedural_count: int
+ recent_sessions: list[str]
+ latest_events: list[dict[str, Any]]
+
+
+class ForgettingCycleRequest(MCPModel):
+ max_decisions: int = Field(default=DEFAULT_MAX_DECISIONS, ge=1, le=MAX_MAX_DECISIONS)
+
+
+class ForgettingDecisionModel(MCPModel):
+ record_id: str
+ memory_type: str
+ action: ForgettingAction
+ reason: str | None = None
+ score: float
+ media_deleted: bool = False
+ executed: bool = False
+ record_skip_reason: str | None = None
+ media_skip_reason: str | None = None
+ old_importance: float | None = None
+ new_importance: float | None = None
+
+
+class ForgettingReportModel(MCPModel):
+ status: Literal["completed", "already_running"]
+ dry_run: bool
+ scanned: int | None = None
+ kept: int | None = None
+ faded: int | None = None
+ pruned: int | None = None
+ media_deleted: int | None = None
+ duplicates_flagged: int | None = None
+ skipped_records: int | None = None
+ skipped_media: int | None = None
+ by_type: dict[str, dict[str, int]] | None = None
+ decisions: list[ForgettingDecisionModel] = Field(default_factory=list)
+ displayed_decisions: int = 0
+ total_decisions: int = 0
+ decisions_truncated: bool = False
+
+
+def serialise_record(record: MemoryRecord) -> MemoryRecordModel:
+ payload: dict[str, Any] = {
+ "id": record.id,
+ "memory_type": record.memory_type,
+ "modality": record.modality,
+ "content": record.content,
+ "created_at": record.created_at,
+ "last_accessed_at": record.last_accessed_at,
+ "access_count": record.access_count,
+ "importance": record.importance,
+ "media_ref": record.media_ref,
+ "media_type": record.media_type,
+ "text_description": record.text_description,
+ "has_media": record.has_media,
+ }
+ if isinstance(record, SemanticMemory):
+ payload.update(
+ {
+ "category": record.category,
+ "domain": record.domain,
+ "confidence": record.confidence,
+ "supersedes": record.supersedes,
+ "superseded_by": record.superseded_by,
+ "related_ids": record.related_ids,
+ "has_visual": record.has_visual,
+ }
+ )
+ if isinstance(record, EpisodicMemory):
+ payload.update(
+ {
+ "session_id": record.session_id,
+ "turn_number": record.turn_number,
+ "participants": record.participants,
+ "summary": record.summary,
+ "emotional_valence": record.emotional_valence,
+ "emotional_profile": record.emotional_profile,
+ "source_mime_type": record.source_mime_type,
+ }
+ )
+ if isinstance(record, ProceduralMemory):
+ payload.update(
+ {
+ "steps": record.steps,
+ "preconditions": record.preconditions,
+ "success_count": record.success_count,
+ "failure_count": record.failure_count,
+ "total_outcomes": record.total_outcomes,
+ "success_rate": record.success_rate,
+ "wilson_score": record.wilson_score,
+ }
+ )
+ return MemoryRecordModel.model_validate(payload)
+
+
+def serialise_ranked_result(result: Any) -> RankedMemoryResultModel:
+ return RankedMemoryResultModel(
+ record=serialise_record(result.record),
+ raw_similarity=result.raw_similarity,
+ recency_score=result.recency_score,
+ importance_score=result.importance_score,
+ final_score=result.final_score,
+ )
+
+
+def serialise_procedural_match(match: ProceduralMatch) -> ProcedureMatchModel:
+ return ProcedureMatchModel(
+ record=serialise_record(match.record),
+ similarity=match.similarity,
+ wilson_score=match.wilson_score,
+ combined_score=match.combined_score,
+ )
+
+
+def serialise_contradiction_candidate(
+ candidate: ContradictionCandidate,
+) -> ContradictionCandidateModel:
+ return ContradictionCandidateModel(
+ record=serialise_record(candidate.record),
+ similarity=candidate.similarity,
+ )
+
+
+def serialise_remember_fact_response(
+ *,
+ record: SemanticMemory,
+ contradiction_status: ContradictionStatus,
+ contradiction_candidates: list[ContradictionCandidate],
+ contradiction_detail: str | None = None,
+ max_contradictions: int = DEFAULT_MAX_CONTRADICTIONS,
+) -> RememberFactResponse:
+ bounded_limit = min(max_contradictions, MAX_MAX_CONTRADICTIONS)
+ displayed = contradiction_candidates[:bounded_limit]
+ return RememberFactResponse(
+ record=serialise_record(record),
+ potential_contradictions=[
+ serialise_contradiction_candidate(candidate) for candidate in displayed
+ ],
+ contradiction_check=ContradictionCheckModel(
+ status=contradiction_status,
+ detail=contradiction_detail,
+ ),
+ total_contradictions=len(contradiction_candidates),
+ contradictions_truncated=len(contradiction_candidates) > len(displayed),
+ )
+
+
+def serialise_forgetting_report(
+ report: ForgettingReport,
+ *,
+ max_decisions: int = DEFAULT_MAX_DECISIONS,
+) -> ForgettingReportModel:
+ bounded_limit = min(max_decisions, MAX_MAX_DECISIONS)
+ decisions = report.decisions[:bounded_limit]
+ return ForgettingReportModel(
+ status="completed",
+ dry_run=report.dry_run,
+ scanned=report.scanned,
+ kept=report.kept,
+ faded=report.faded,
+ pruned=report.pruned,
+ media_deleted=report.media_deleted,
+ duplicates_flagged=report.duplicates_flagged,
+ skipped_records=report.skipped_records,
+ skipped_media=report.skipped_media,
+ by_type=dict(report.by_type),
+ decisions=[
+ ForgettingDecisionModel(
+ record_id=decision.record_id,
+ memory_type=decision.memory_type,
+ action=decision.action,
+ reason=decision.reason,
+ score=decision.score,
+ media_deleted=decision.media_deleted,
+ executed=decision.executed,
+ record_skip_reason=decision.record_skip_reason,
+ media_skip_reason=decision.media_skip_reason,
+ old_importance=decision.old_importance,
+ new_importance=decision.new_importance,
+ )
+ for decision in decisions
+ ],
+ displayed_decisions=len(decisions),
+ total_decisions=len(report.decisions),
+ decisions_truncated=len(report.decisions) > len(decisions),
+ )
+
+
+def serialise_memory_overview(payload: dict[str, Any]) -> MemoryOverviewModel:
+ return MemoryOverviewModel.model_validate(payload)
+
+
+def get_memory_by_type(service: Any, request: GetMemoryRequest) -> MemoryRecord | None:
+ store = {
+ "semantic": service.semantic_store,
+ "episodic": service.episodic_store,
+ "procedural": service.procedural_store,
+ }[request.memory_type]
+ return store.get_by_id(request.record_id)
diff --git a/tests/test_mcp_schemas.py b/tests/test_mcp_schemas.py
new file mode 100644
index 0000000..4f77ade
--- /dev/null
+++ b/tests/test_mcp_schemas.py
@@ -0,0 +1,198 @@
+from __future__ import annotations
+
+from datetime import datetime, timedelta, timezone
+
+import pytest
+from pydantic import ValidationError
+
+from forgetting.contradiction import ContradictionCandidate
+from forgetting.service import ForgettingDecision, ForgettingReport
+from mcp_server.schemas import (
+ DEFAULT_MAX_CONTRADICTIONS,
+ DEFAULT_MAX_DECISIONS,
+ GetMemoryRequest,
+ InlineContentModel,
+ MCPToolError,
+ MediaInputModel,
+ RecallEpisodesRequestAdapter,
+ RememberFactRequest,
+ serialise_forgetting_report,
+ serialise_procedural_match,
+ serialise_remember_fact_response,
+ get_memory_by_type,
+)
+from models.procedural import ProceduralMemory
+from models.semantic import SemanticMemory
+from runtime import build_runtime
+from stores.procedural_store import ProceduralMatch
+from tests.helpers import HashingEmbedder, cleanup_dir, make_temp_chroma_dir
+
+
+def test_recall_episodes_request_supports_all_locked_modes():
+ recent = RecallEpisodesRequestAdapter.validate_python({"mode": "recent", "limit": 7})
+ session = RecallEpisodesRequestAdapter.validate_python(
+ {"mode": "session", "session_id": "session-123"}
+ )
+ time_range = RecallEpisodesRequestAdapter.validate_python(
+ {
+ "mode": "time_range",
+ "start": "2026-03-29T00:00:00+00:00",
+ "end": "2026-03-30T00:00:00+00:00",
+ }
+ )
+
+ assert recent.mode == "recent"
+ assert recent.limit == 7
+ assert session.mode == "session"
+ assert session.session_id == "session-123"
+ assert time_range.mode == "time_range"
+ assert time_range.start < time_range.end
+
+
+def test_recall_episodes_request_rejects_invalid_time_ranges():
+ with pytest.raises(ValidationError):
+ RecallEpisodesRequestAdapter.validate_python(
+ {
+ "mode": "time_range",
+ "start": "2026-03-30T00:00:00+00:00",
+ "end": "2026-03-29T00:00:00+00:00",
+ }
+ )
+
+
+def test_media_input_supports_file_path_or_inline_base64_only():
+ file_input = MediaInputModel(file_path="/tmp/screenshot.png")
+ inline_input = MediaInputModel(
+ inline_content=InlineContentModel(
+ inline_base64="ZmFrZQ==",
+ mime_type="image/png",
+ filename="screenshot.png",
+ )
+ )
+
+ assert file_input.file_path == "/tmp/screenshot.png"
+ assert inline_input.inline_content is not None
+
+ with pytest.raises(ValidationError):
+ MediaInputModel(
+ file_path="/tmp/screenshot.png",
+ inline_content=InlineContentModel(
+ inline_base64="ZmFrZQ==",
+ mime_type="image/png",
+ ),
+ )
+
+
+def test_remember_fact_request_bounds_contradiction_count():
+ request = RememberFactRequest(content="Embedding size is 768")
+ assert request.max_contradictions == DEFAULT_MAX_CONTRADICTIONS
+
+ with pytest.raises(ValidationError):
+ RememberFactRequest(content="Embedding size is 768", max_contradictions=999)
+
+
+def test_serialise_remember_fact_response_reports_truncation():
+ record = SemanticMemory(content="Embedding size is 768")
+ candidates = [
+ ContradictionCandidate(record=SemanticMemory(content=f"Candidate {index}"), similarity=0.9 - index * 0.1)
+ for index in range(3)
+ ]
+
+ payload = serialise_remember_fact_response(
+ record=record,
+ contradiction_status="completed",
+ contradiction_candidates=candidates,
+ max_contradictions=1,
+ )
+
+ assert payload.record.content == "Embedding size is 768"
+ assert payload.contradiction_check_order == "after_store"
+ assert payload.total_contradictions == 3
+ assert payload.contradictions_truncated is True
+ assert len(payload.potential_contradictions) == 1
+
+
+def test_serialise_forgetting_report_caps_decisions():
+ report = ForgettingReport(
+ dry_run=True,
+ decisions=[
+ ForgettingDecision(
+ record_id=f"record-{index}",
+ memory_type="semantic",
+ action="keep",
+ reason="time_decay",
+ score=0.1 * index,
+ )
+ for index in range(DEFAULT_MAX_DECISIONS + 5)
+ ],
+ )
+
+ payload = serialise_forgetting_report(report)
+
+ assert payload.status == "completed"
+ assert payload.displayed_decisions == DEFAULT_MAX_DECISIONS
+ assert payload.total_decisions == DEFAULT_MAX_DECISIONS + 5
+ assert payload.decisions_truncated is True
+ assert len(payload.decisions) == DEFAULT_MAX_DECISIONS
+
+
+def test_procedural_match_serializer_exposes_all_scores():
+ record = ProceduralMemory(content="Deploy with Docker", steps=["docker compose up"])
+ match = ProceduralMatch(
+ record=record,
+ similarity=0.75,
+ wilson_score=0.55,
+ combined_score=0.68,
+ )
+
+ payload = serialise_procedural_match(match)
+
+ assert payload.similarity == 0.75
+ assert payload.wilson_score == 0.55
+ assert payload.combined_score == 0.68
+
+
+def test_get_memory_by_type_routes_to_the_correct_store():
+ chroma_dir = make_temp_chroma_dir("mcp_schemas_")
+ runtime = build_runtime(
+ chroma_path=chroma_dir,
+ embedder=HashingEmbedder(dimensions=64),
+ embedding_dimensions=64,
+ )
+ try:
+ semantic = SemanticMemory(content="Chroma path is isolated per runtime")
+ procedural = ProceduralMemory(
+ content="Deploy the app",
+ steps=["Build the image", "Start the service"],
+ )
+ runtime.semantic_store.store(semantic)
+ runtime.procedural_store.store(procedural)
+
+ semantic_result = get_memory_by_type(
+ runtime,
+ GetMemoryRequest(memory_type="semantic", record_id=semantic.id),
+ )
+ procedural_result = get_memory_by_type(
+ runtime,
+ GetMemoryRequest(memory_type="procedural", record_id=procedural.id),
+ )
+
+ assert semantic_result is not None
+ assert semantic_result.id == semantic.id
+ assert procedural_result is not None
+ assert procedural_result.id == procedural.id
+ finally:
+ cleanup_dir(chroma_dir)
+
+
+def test_typed_tool_errors_produce_stable_payloads():
+ error = MCPToolError(
+ code="dependency_unavailable",
+ message="Embedding provider timed out",
+ retryable=True,
+ details={"provider": "gemini"},
+ )
+
+ assert error.payload.code == "dependency_unavailable"
+ assert error.payload.retryable is True
+ assert error.payload.details == {"provider": "gemini"}
From b2528e3a195a5a784b597d134dcab87a2becd6d0 Mon Sep 17 00:00:00 2001
From: Atharva-Kanherkar
Date: Mon, 30 Mar 2026 03:30:44 +0530
Subject: [PATCH 2/4] Address MCP schema review feedback
---
api/app.py | 31 +++++------
mcp_server/schemas.py | 16 ++++--
tests/test_api_runtime_edges.py | 6 +++
tests/test_forgetting_integration.py | 8 +++
tests/test_mcp_schemas.py | 26 +++++++++
web/src/lib/api.ts | 79 ++++++++++++++++++----------
6 files changed, 117 insertions(+), 49 deletions(-)
diff --git a/api/app.py b/api/app.py
index 33880af..a8220f4 100644
--- a/api/app.py
+++ b/api/app.py
@@ -17,11 +17,13 @@
import config
from mcp_server.schemas import (
+ ForgettingReportModel,
serialise_contradiction_candidate as _mcp_serialise_contradiction_candidate,
serialise_forgetting_report as _mcp_serialise_forgetting_report,
serialise_procedural_match as _mcp_serialise_procedural_match,
serialise_ranked_result as _mcp_serialise_ranked_result,
serialise_record as _mcp_serialise_record,
+ serialise_remember_fact_response as _mcp_serialise_remember_fact_response,
)
from models.base import MemoryRecord, normalize_modality
from models.episodic import EpisodicMemory
@@ -341,10 +343,10 @@ async def run_forgetting_cycle(*, dry_run: bool) -> dict[str, Any]:
# deployments still need an external coordinator if cross-process
# forgetting exclusivity becomes a requirement.
if lock.locked():
- return {
- "status": "already_running",
- "dry_run": dry_run,
- }
+ return ForgettingReportModel(
+ status="already_running",
+ dry_run=dry_run,
+ ).model_dump(mode="json", exclude_none=True)
async with lock:
report = await asyncio.to_thread(service().forgetting_service.run_cycle, dry_run)
@@ -418,21 +420,12 @@ async def create_semantic_memory(payload: dict[str, Any]) -> dict[str, Any]:
_cleanup_owned_media(active_service, record.media_ref)
raise
contradiction_check = _safe_contradiction_lookup(active_service, record)
- return {
- "record": _serialise_record(record),
- "potential_contradictions": [
- _serialise_contradiction_candidate(c)
- for c in contradiction_check.candidates
- ],
- "contradiction_check": {
- "status": contradiction_check.status,
- **(
- {"detail": contradiction_check.detail}
- if contradiction_check.detail is not None
- else {}
- ),
- },
- }
+ return _mcp_serialise_remember_fact_response(
+ record=record,
+ contradiction_status=contradiction_check.status,
+ contradiction_candidates=contradiction_check.candidates,
+ contradiction_detail=contradiction_check.detail,
+ ).model_dump(mode="json", exclude_none=True)
@app.post("/api/memories/episodic/text")
async def create_text_episode(payload: dict[str, Any]) -> dict[str, Any]:
diff --git a/mcp_server/schemas.py b/mcp_server/schemas.py
index 804f3a8..dceeca8 100644
--- a/mcp_server/schemas.py
+++ b/mcp_server/schemas.py
@@ -3,7 +3,7 @@
from datetime import datetime
from typing import Annotated, Any, Literal
-from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, model_validator
+from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, field_validator, model_validator
from forgetting.contradiction import ContradictionCandidate
from forgetting.service import ForgettingReport
@@ -11,6 +11,8 @@
from models.episodic import EpisodicMemory
from models.procedural import ProceduralMemory
from models.semantic import SemanticMemory
+from retrieval.ranking import RankedResult
+from runtime import MemoryRuntime
from stores.procedural_store import ProceduralMatch
MemoryType = Literal["semantic", "episodic", "procedural"]
@@ -181,6 +183,14 @@ class RememberProcedureRequest(MCPModel):
media_type: Literal["image", "audio", "video", "pdf"] | None = None
text_description: str | None = None
+ @field_validator("steps", "preconditions")
+ @classmethod
+ def validate_string_lists(cls, values: list[str]) -> list[str]:
+ for value in values:
+ if not isinstance(value, str) or not value.strip():
+ raise ValueError("steps and preconditions must contain only non-blank strings")
+ return values
+
class RememberFactResponse(MCPModel):
record: MemoryRecordModel
@@ -380,7 +390,7 @@ def serialise_record(record: MemoryRecord) -> MemoryRecordModel:
return MemoryRecordModel.model_validate(payload)
-def serialise_ranked_result(result: Any) -> RankedMemoryResultModel:
+def serialise_ranked_result(result: RankedResult) -> RankedMemoryResultModel:
return RankedMemoryResultModel(
record=serialise_record(result.record),
raw_similarity=result.raw_similarity,
@@ -477,7 +487,7 @@ def serialise_memory_overview(payload: dict[str, Any]) -> MemoryOverviewModel:
return MemoryOverviewModel.model_validate(payload)
-def get_memory_by_type(service: Any, request: GetMemoryRequest) -> MemoryRecord | None:
+def get_memory_by_type(service: MemoryRuntime, request: GetMemoryRequest) -> MemoryRecord | None:
store = {
"semantic": service.semantic_store,
"episodic": service.episodic_store,
diff --git a/tests/test_api_runtime_edges.py b/tests/test_api_runtime_edges.py
index 29a584d..c731a1e 100644
--- a/tests/test_api_runtime_edges.py
+++ b/tests/test_api_runtime_edges.py
@@ -45,6 +45,9 @@ async def test_semantic_write_swallows_valueerror_from_contradiction_lookup():
assert response.status_code == 200
assert response.json()["potential_contradictions"] == []
+ assert response.json()["total_contradictions"] == 0
+ assert response.json()["contradictions_truncated"] is False
+ assert response.json()["contradiction_check_order"] == "after_store"
assert response.json()["contradiction_check"] == {
"status": "skipped",
"detail": "synthetic contradiction miss",
@@ -66,6 +69,9 @@ async def test_semantic_write_reports_programming_error_from_contradiction_looku
assert response.status_code == 200
assert response.json()["potential_contradictions"] == []
+ assert response.json()["total_contradictions"] == 0
+ assert response.json()["contradictions_truncated"] is False
+ assert response.json()["contradiction_check_order"] == "after_store"
assert response.json()["contradiction_check"] == {
"status": "error",
"detail": "synthetic programming error",
diff --git a/tests/test_forgetting_integration.py b/tests/test_forgetting_integration.py
index 6992905..eb05f88 100644
--- a/tests/test_forgetting_integration.py
+++ b/tests/test_forgetting_integration.py
@@ -42,6 +42,9 @@ async def test_semantic_store_returns_potential_contradictions():
body1 = r1.json()
assert "potential_contradictions" in body1
assert body1["contradiction_check"]["status"] == "completed"
+ assert body1["total_contradictions"] == len(body1["potential_contradictions"])
+ assert body1["contradictions_truncated"] is False
+ assert body1["contradiction_check_order"] == "after_store"
assert isinstance(body1["potential_contradictions"], list)
r2 = await client.post(
@@ -52,6 +55,8 @@ async def test_semantic_store_returns_potential_contradictions():
body2 = r2.json()
assert "potential_contradictions" in body2
assert body2["contradiction_check"]["status"] == "completed"
+ assert body2["total_contradictions"] == len(body2["potential_contradictions"])
+ assert body2["contradictions_truncated"] is False
@pytest.mark.anyio
@@ -69,6 +74,9 @@ async def test_forgetting_preview_returns_stable_report():
assert report["scanned"] == 2
assert isinstance(report["decisions"], list)
assert len(report["decisions"]) == 2
+ assert report["displayed_decisions"] == 2
+ assert report["total_decisions"] == 2
+ assert report["decisions_truncated"] is False
assert report["kept"] + report["faded"] + report["pruned"] == report["scanned"]
for decision in report["decisions"]:
diff --git a/tests/test_mcp_schemas.py b/tests/test_mcp_schemas.py
index 4f77ade..be2051e 100644
--- a/tests/test_mcp_schemas.py
+++ b/tests/test_mcp_schemas.py
@@ -10,12 +10,14 @@
from mcp_server.schemas import (
DEFAULT_MAX_CONTRADICTIONS,
DEFAULT_MAX_DECISIONS,
+ ForgettingReportModel,
GetMemoryRequest,
InlineContentModel,
MCPToolError,
MediaInputModel,
RecallEpisodesRequestAdapter,
RememberFactRequest,
+ RememberProcedureRequest,
serialise_forgetting_report,
serialise_procedural_match,
serialise_remember_fact_response,
@@ -91,6 +93,11 @@ def test_remember_fact_request_bounds_contradiction_count():
RememberFactRequest(content="Embedding size is 768", max_contradictions=999)
+def test_remember_procedure_request_rejects_blank_steps():
+ with pytest.raises(ValidationError):
+ RememberProcedureRequest(content="Deploy the app", steps=[""])
+
+
def test_serialise_remember_fact_response_reports_truncation():
record = SemanticMemory(content="Embedding size is 768")
candidates = [
@@ -110,6 +117,9 @@ def test_serialise_remember_fact_response_reports_truncation():
assert payload.total_contradictions == 3
assert payload.contradictions_truncated is True
assert len(payload.potential_contradictions) == 1
+ assert payload.model_dump(mode="json", exclude_none=True)["contradiction_check"] == {
+ "status": "completed"
+ }
def test_serialise_forgetting_report_caps_decisions():
@@ -196,3 +206,19 @@ def test_typed_tool_errors_produce_stable_payloads():
assert error.payload.code == "dependency_unavailable"
assert error.payload.retryable is True
assert error.payload.details == {"provider": "gemini"}
+
+
+def test_already_running_forgetting_payload_is_model_shaped():
+ payload = ForgettingReportModel(
+ status="already_running",
+ dry_run=True,
+ ).model_dump(mode="json", exclude_none=True)
+
+ assert payload == {
+ "status": "already_running",
+ "dry_run": True,
+ "decisions": [],
+ "displayed_decisions": 0,
+ "total_decisions": 0,
+ "decisions_truncated": False,
+ }
diff --git a/web/src/lib/api.ts b/web/src/lib/api.ts
index fc17788..714dcda 100644
--- a/web/src/lib/api.ts
+++ b/web/src/lib/api.ts
@@ -8,20 +8,30 @@ export type MemoryRecord = {
access_count: number;
importance: number;
media_ref: string | null;
- session_id?: string;
+ media_type: string | null;
+ text_description: string | null;
+ has_media: boolean;
+ category?: string | null;
+ domain?: string | null;
+ confidence?: number | null;
+ supersedes?: string | null;
+ superseded_by?: string | null;
+ related_ids?: string[] | null;
+ has_visual?: boolean | null;
+ session_id?: string | null;
turn_number?: number | null;
- participants?: string[];
+ participants?: string[] | null;
summary?: string | null;
+ emotional_valence?: number | null;
+ emotional_profile?: Record | null;
source_mime_type?: string | null;
- category?: string;
- confidence?: number;
- steps?: string[];
- preconditions?: string[];
- success_count?: number;
- failure_count?: number;
- total_outcomes?: number;
- success_rate?: number;
- wilson_score?: number;
+ steps?: string[] | null;
+ preconditions?: string[] | null;
+ success_count?: number | null;
+ failure_count?: number | null;
+ total_outcomes?: number | null;
+ success_rate?: number | null;
+ wilson_score?: number | null;
};
export type RankedQueryResult = {
@@ -58,6 +68,20 @@ export type ContradictionCandidate = {
similarity: number;
};
+export type ContradictionCheck = {
+ status: "completed" | "skipped" | "error";
+ detail?: string | null;
+};
+
+export type RememberFactResponse = {
+ record: MemoryRecord;
+ potential_contradictions: ContradictionCandidate[];
+ contradiction_check: ContradictionCheck;
+ contradiction_check_order: "after_store";
+ total_contradictions: number;
+ contradictions_truncated: boolean;
+};
+
export type ForgettingDecision = {
record_id: string;
memory_type: string;
@@ -73,17 +97,21 @@ export type ForgettingDecision = {
};
export type ForgettingReport = {
+ status: "completed" | "already_running";
dry_run: boolean;
- scanned: number;
- kept: number;
- faded: number;
- pruned: number;
- media_deleted: number;
- duplicates_flagged: number;
- skipped_records: number;
- skipped_media: number;
- by_type: Record>;
+ scanned?: number;
+ kept?: number;
+ faded?: number;
+ pruned?: number;
+ media_deleted?: number;
+ duplicates_flagged?: number;
+ skipped_records?: number;
+ skipped_media?: number;
+ by_type?: Record>;
decisions: ForgettingDecision[];
+ displayed_decisions: number;
+ total_decisions: number;
+ decisions_truncated: boolean;
};
const API_BASE = process.env.NEXT_PUBLIC_MEMORY_API_BASE_URL ?? "http://localhost:8000";
@@ -123,13 +151,10 @@ export function createSemanticMemory(input: {
category?: string;
confidence?: number;
}) {
- return request<{ record: MemoryRecord; potential_contradictions: ContradictionCandidate[] }>(
- "/api/memories/semantic",
- {
- method: "POST",
- body: JSON.stringify(input),
- },
- );
+ return request("/api/memories/semantic", {
+ method: "POST",
+ body: JSON.stringify(input),
+ });
}
export function createTextEpisode(input: {
From d97dd2ae133a018e95760fbc47ebfe1b6234aac2 Mon Sep 17 00:00:00 2001
From: Atharva-Kanherkar
Date: Mon, 30 Mar 2026 03:35:51 +0530
Subject: [PATCH 3/4] Fix playground forgetting report null handling
---
web/src/components/playground-app.tsx | 12 ++++++------
1 file changed, 6 insertions(+), 6 deletions(-)
diff --git a/web/src/components/playground-app.tsx b/web/src/components/playground-app.tsx
index 15887b6..dbe9311 100644
--- a/web/src/components/playground-app.tsx
+++ b/web/src/components/playground-app.tsx
@@ -1513,16 +1513,16 @@ export function PlaygroundApp() {
- Scanned {forgettingReport.scanned} · Kept {forgettingReport.kept} ·
- Faded {forgettingReport.faded} · Pruned {forgettingReport.pruned}
+ Scanned {forgettingReport.scanned ?? 0} · Kept {forgettingReport.kept ?? 0} ·
+ Faded {forgettingReport.faded ?? 0} · Pruned {forgettingReport.pruned ?? 0}
- {forgettingReport.duplicates_flagged > 0 && (
+ {(forgettingReport.duplicates_flagged ?? 0) > 0 && (
- {forgettingReport.duplicates_flagged} duplicate{forgettingReport.duplicates_flagged > 1 ? "s" : ""}
+ {forgettingReport.duplicates_flagged ?? 0} duplicate{(forgettingReport.duplicates_flagged ?? 0) > 1 ? "s" : ""}
)}
- {forgettingReport.media_deleted > 0 && (
-
{forgettingReport.media_deleted} media files deleted
+ {(forgettingReport.media_deleted ?? 0) > 0 && (
+
{forgettingReport.media_deleted ?? 0} media files deleted
)}
{forgettingReport.decisions
.filter((d) => d.action !== "keep")
From 82ef06754c9d7fd017c2becc07f206e7967af43c Mon Sep 17 00:00:00 2001
From: Atharva-Kanherkar
Date: Mon, 30 Mar 2026 03:55:53 +0530
Subject: [PATCH 4/4] Fix forgetting endpoint test hangs
---
api/app.py | 11 +++++-----
tests/test_forgetting_integration.py | 33 ++++++++--------------------
2 files changed, 14 insertions(+), 30 deletions(-)
diff --git a/api/app.py b/api/app.py
index a8220f4..adf9bf2 100644
--- a/api/app.py
+++ b/api/app.py
@@ -17,7 +17,6 @@
import config
from mcp_server.schemas import (
- ForgettingReportModel,
serialise_contradiction_candidate as _mcp_serialise_contradiction_candidate,
serialise_forgetting_report as _mcp_serialise_forgetting_report,
serialise_procedural_match as _mcp_serialise_procedural_match,
@@ -343,13 +342,13 @@ async def run_forgetting_cycle(*, dry_run: bool) -> dict[str, Any]:
# deployments still need an external coordinator if cross-process
# forgetting exclusivity becomes a requirement.
if lock.locked():
- return ForgettingReportModel(
- status="already_running",
- dry_run=dry_run,
- ).model_dump(mode="json", exclude_none=True)
+ return {
+ "status": "already_running",
+ "dry_run": dry_run,
+ }
async with lock:
- report = await asyncio.to_thread(service().forgetting_service.run_cycle, dry_run)
+ report = service().forgetting_service.run_cycle(dry_run)
return _serialise_forgetting_report(report)
@app.get("/health")
diff --git a/tests/test_forgetting_integration.py b/tests/test_forgetting_integration.py
index eb05f88..8050d0e 100644
--- a/tests/test_forgetting_integration.py
+++ b/tests/test_forgetting_integration.py
@@ -1,11 +1,9 @@
"""Verify forgetting endpoints, contradiction lookup, and full cycle integration."""
-import asyncio
import os
import shutil
import sys
import tempfile
-import threading
import httpx
import pytest
@@ -237,28 +235,15 @@ async def test_overlapping_forgetting_requests_return_already_running():
async with make_client() as client:
await client.post("/api/memories/semantic", json={"content": "Overlap witness"})
await client.get("/api/overview")
-
- entered = threading.Event()
- release = threading.Event()
- original = client.app.state.service.forgetting_service.run_cycle
-
- def blocking_run_cycle(dry_run: bool = False):
- entered.set()
- release.wait()
- return original(dry_run=dry_run)
-
- client.app.state.service.forgetting_service.run_cycle = blocking_run_cycle
-
- first_task = asyncio.create_task(client.post("/api/forgetting/run"))
- await asyncio.to_thread(entered.wait)
- second_response = await client.post("/api/forgetting/run")
- release.set()
- first_response = await first_task
-
- assert first_response.status_code == 200
- assert first_response.json()["status"] == "completed"
- assert second_response.status_code == 200
- assert second_response.json() == {
+ lock = client.app.state.forgetting_lock
+ await lock.acquire()
+ try:
+ response = await client.post("/api/forgetting/run")
+ finally:
+ lock.release()
+
+ assert response.status_code == 200
+ assert response.json() == {
"status": "already_running",
"dry_run": False,
}