diff --git a/api/app.py b/api/app.py index 5e41082..adf9bf2 100644 --- a/api/app.py +++ b/api/app.py @@ -16,6 +16,14 @@ 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, + serialise_remember_fact_response as _mcp_serialise_remember_fact_response, +) from models.base import MemoryRecord, normalize_modality from models.episodic import EpisodicMemory from models.procedural import ProceduralMemory @@ -197,75 +205,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 +299,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( @@ -431,7 +348,7 @@ async def run_forgetting_cycle(*, dry_run: bool) -> dict[str, Any]: } 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") @@ -502,21 +419,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/__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..dceeca8 --- /dev/null +++ b/mcp_server/schemas.py @@ -0,0 +1,496 @@ +from __future__ import annotations + +from datetime import datetime +from typing import Annotated, Any, Literal + +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, field_validator, 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 retrieval.ranking import RankedResult +from runtime import MemoryRuntime +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 + + @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 + 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: RankedResult) -> 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: MemoryRuntime, 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_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..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 @@ -42,6 +40,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 +53,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 +72,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"]: @@ -229,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, } diff --git a/tests/test_mcp_schemas.py b/tests/test_mcp_schemas.py new file mode 100644 index 0000000..be2051e --- /dev/null +++ b/tests/test_mcp_schemas.py @@ -0,0 +1,224 @@ +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, + ForgettingReportModel, + GetMemoryRequest, + InlineContentModel, + MCPToolError, + MediaInputModel, + RecallEpisodesRequestAdapter, + RememberFactRequest, + RememberProcedureRequest, + 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_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 = [ + 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 + assert payload.model_dump(mode="json", exclude_none=True)["contradiction_check"] == { + "status": "completed" + } + + +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"} + + +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/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") 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