From 5a7a974680d4bb8cbf0bf662b34768468e6a8e42 Mon Sep 17 00:00:00 2001 From: Atharva-Kanherkar Date: Mon, 30 Mar 2026 15:23:31 +0530 Subject: [PATCH 1/2] Add FastMCP server and startup path --- README.md | 23 ++ mcp_server/__init__.py | 10 + mcp_server/__main__.py | 5 + mcp_server/server.py | 797 +++++++++++++++++++++++++++++++++++++++ requirements.txt | 1 + tests/test_mcp_server.py | 113 ++++++ 6 files changed, 949 insertions(+) create mode 100644 mcp_server/__main__.py create mode 100644 mcp_server/server.py create mode 100644 tests/test_mcp_server.py diff --git a/README.md b/README.md index ba302f4..ef4a585 100644 --- a/README.md +++ b/README.md @@ -61,6 +61,29 @@ The playground is now running at `http://localhost:3000`. --- +## MCP Server + +The repo also ships a thin MCP server under `mcp_server/` that reuses the shared runtime. + +### Local stdio transport + +```bash +MEMORY_MCP_TRANSPORT=stdio .venv/bin/python -m mcp_server +``` + +### Streamable HTTP transport + +```bash +MEMORY_MCP_TRANSPORT=streamable-http \ +MEMORY_MCP_HOST=127.0.0.1 \ +MEMORY_MCP_PORT=8000 \ +.venv/bin/python -m mcp_server +``` + +When running over HTTP, the MCP endpoint is mounted at `/mcp` and health is exposed at `/health`. + +--- + ## CLI ```bash diff --git a/mcp_server/__init__.py b/mcp_server/__init__.py index 1f2f97c..0c64c4d 100644 --- a/mcp_server/__init__.py +++ b/mcp_server/__init__.py @@ -1,5 +1,11 @@ """Shared MCP contracts and helpers for the memory runtime.""" +from .server import ( + MemoryMCPServer, + ServerConfig, + create_http_app, + create_mcp_server, +) from .schemas import ( ContradictionCandidateModel, ContradictionCheckModel, @@ -24,6 +30,10 @@ ) __all__ = [ + "MemoryMCPServer", + "ServerConfig", + "create_http_app", + "create_mcp_server", "ContradictionCandidateModel", "ContradictionCheckModel", "ForgettingCycleRequest", diff --git a/mcp_server/__main__.py b/mcp_server/__main__.py new file mode 100644 index 0000000..5b9eec7 --- /dev/null +++ b/mcp_server/__main__.py @@ -0,0 +1,5 @@ +from .server import main + + +if __name__ == "__main__": + main() diff --git a/mcp_server/server.py b/mcp_server/server.py new file mode 100644 index 0000000..a3bfac5 --- /dev/null +++ b/mcp_server/server.py @@ -0,0 +1,797 @@ +from __future__ import annotations + +import asyncio +import base64 +import binascii +import logging +import mimetypes +import os +import tempfile +from contextlib import asynccontextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import config +from mcp_server.schemas import ( + DEFAULT_MAX_CONTRADICTIONS, + ForgettingCycleRequest, + ForgettingReportModel, + GetMemoryRequest, + GetMemoryResponse, + MCPToolError, + MediaInputModel, + MemoryOverviewModel, + RecallEpisodesRequestAdapter, + RecallEpisodesResponse, + RecallMemoriesRequest, + RecallMemoriesResponse, + RecallProceduresRequest, + RecallProceduresResponse, + RecordProcedureOutcomeRequest, + RecordProcedureOutcomeResponse, + RememberEpisodeRequest, + RememberEpisodeResponse, + RememberFactRequest, + RememberFactResponse, + RememberProcedureRequest, + RememberProcedureResponse, + ResolveContradictionRequest, + ResolveContradictionResponse, + ToolErrorModel, + get_memory_by_type, + serialise_forgetting_report, + serialise_memory_overview, + serialise_procedural_match, + serialise_ranked_result, + serialise_record, + serialise_remember_fact_response, +) +from models.base import normalize_modality +from models.episodic import EpisodicMemory +from models.procedural import ProceduralMemory +from models.semantic import SemanticMemory +from runtime import MemoryRuntime, build_runtime +from stores.episodic_store import EpisodicStoreError, MediaTooLargeError +from stores.media_store import MediaStore +from utils.embeddings import EmbeddingProviderError, TextEmbedder + +logger = logging.getLogger(__name__) + +_TRANSPORTS = {"stdio", "streamable-http"} +_MCP_SERVER_NAME = "agentic-memory" + + +@dataclass(frozen=True, slots=True) +class ServerConfig: + transport: str = "stdio" + host: str = "127.0.0.1" + port: int = 8000 + chroma_path: str | None = None + media_root: str | None = None + embedding_dimensions: int | None = None + streamable_http_path: str = "/mcp" + + @classmethod + def from_env(cls) -> "ServerConfig": + embedding_dimensions = os.getenv("EMBEDDING_DIMENSIONS") + port = os.getenv("MEMORY_MCP_PORT") + return cls( + transport=os.getenv("MEMORY_MCP_TRANSPORT", "stdio").strip().lower(), + host=os.getenv("MEMORY_MCP_HOST", "127.0.0.1").strip(), + port=int(port) if port else 8000, + chroma_path=os.getenv("MEMORY_CHROMA_PATH") or None, + media_root=os.getenv("MEMORY_MEDIA_DIR") or None, + embedding_dimensions=int(embedding_dimensions) if embedding_dimensions else None, + streamable_http_path=os.getenv("MEMORY_MCP_PATH", "/mcp").strip() or "/mcp", + ) + + def validate(self) -> "ServerConfig": + if self.transport not in _TRANSPORTS: + raise ValueError( + "MEMORY_MCP_TRANSPORT must be one of: stdio, streamable-http" + ) + if self.port <= 0: + raise ValueError("MEMORY_MCP_PORT must be a positive integer") + if not self.streamable_http_path.startswith("/"): + raise ValueError("MEMORY_MCP_PATH must start with '/'") + if self.streamable_http_path != "/" and self.streamable_http_path.endswith("/"): + raise ValueError("MEMORY_MCP_PATH must not end with '/'") + + if self.transport == "streamable-http": + for label, raw_path in ( + ("MEMORY_CHROMA_PATH", self.chroma_path or config.CHROMA_DB_PATH), + ("MEMORY_MEDIA_DIR", self.media_root or config.MEDIA_STORAGE_PATH), + ): + if not Path(raw_path).is_absolute(): + logger.warning( + "%s is relative (%s). This is safe for local dev but risky for remote " + "streamable-http deployments due to working-directory drift.", + label, + raw_path, + ) + return self + + +@dataclass(frozen=True, slots=True) +class PreparedMedia: + path: Path + media_type: str + cleanup: bool = False + + +def _map_exception(exc: Exception) -> MCPToolError: + if isinstance(exc, MCPToolError): + return exc + if isinstance(exc, MediaTooLargeError): + return MCPToolError( + code="media_too_large", + message=str(exc), + details={"exception_type": exc.__class__.__name__}, + ) + if isinstance(exc, (FileNotFoundError, ValueError, EpisodicStoreError, binascii.Error)): + return MCPToolError( + code="invalid_request", + message=str(exc), + details={"exception_type": exc.__class__.__name__}, + ) + if isinstance(exc, EmbeddingProviderError): + return MCPToolError( + code="dependency_unavailable", + message="Gemini embedding provider failed after retries", + retryable=True, + details={"exception_type": exc.__class__.__name__}, + ) + logger.exception("Unhandled MCP server error", exc_info=exc) + return MCPToolError( + code="internal_error", + message=str(exc) or exc.__class__.__name__, + details={"exception_type": exc.__class__.__name__}, + ) + + +def _safe_contradiction_lookup( + runtime: MemoryRuntime, + record: SemanticMemory, + *, + max_contradictions: int = DEFAULT_MAX_CONTRADICTIONS, +) -> RememberFactResponse: + try: + candidates = runtime.contradiction_detector.find_potential_contradictions( + record, + top_k=max_contradictions, + ) + return serialise_remember_fact_response( + record=record, + contradiction_status="completed", + contradiction_candidates=candidates, + max_contradictions=max_contradictions, + ) + except ValueError as exc: + return serialise_remember_fact_response( + record=record, + contradiction_status="skipped", + contradiction_candidates=[], + contradiction_detail=str(exc), + max_contradictions=max_contradictions, + ) + except Exception as exc: # pragma: no cover - defensive path + logger.exception("Contradiction lookup failed for semantic record %s", record.id) + return serialise_remember_fact_response( + record=record, + contradiction_status="error", + contradiction_candidates=[], + contradiction_detail=str(exc), + max_contradictions=max_contradictions, + ) + + +class MemoryMCPServer: + def __init__( + self, + *, + config: ServerConfig, + embedder: TextEmbedder | None = None, + ) -> None: + self._config = config.validate() + self._embedder = embedder + self._runtime: MemoryRuntime | None = None + self._forgetting_lock = asyncio.Lock() + + @property + def config(self) -> ServerConfig: + return self._config + + def runtime(self) -> MemoryRuntime: + if self._runtime is None: + self._runtime = build_runtime( + chroma_path=self._config.chroma_path, + media_root=self._config.media_root, + embedder=self._embedder, + embedding_dimensions=self._config.embedding_dimensions, + ) + return self._runtime + + def validate_startup(self) -> None: + self.runtime() + + def overview(self) -> MemoryOverviewModel: + runtime = self.runtime() + payload = { + "semantic_count": runtime.semantic_store._collection.count(), + "episodic_count": runtime.episodic_store._collection.count(), + "procedural_count": runtime.procedural_store._collection.count(), + "recent_sessions": sorted( + {record.session_id for record in runtime.episodic_store.get_recent(5)} + ), + "latest_events": runtime.events.snapshot(10), + } + return serialise_memory_overview(payload) + + def remember_fact(self, request: RememberFactRequest) -> RememberFactResponse: + runtime = self.runtime() + prepared_media = self._prepare_media(request.media, request.media_type) + try: + modality, media_type = self._resolve_media_contract( + modality=request.modality, + prepared_media=prepared_media, + requested_media_type=request.media_type, + ) + record = SemanticMemory( + content=request.content, + importance=request.importance, + category=request.category, + domain=request.domain, + confidence=request.confidence, + supersedes=request.supersedes, + related_ids=list(request.related_ids), + has_visual=request.has_visual, + modality=modality, + media_ref=str(prepared_media.path) if prepared_media else None, + media_type=media_type, + text_description=request.text_description, + ) + runtime.semantic_store.store(record) + return _safe_contradiction_lookup( + runtime, + record, + max_contradictions=request.max_contradictions, + ) + finally: + self._cleanup_prepared_media(prepared_media) + + def remember_episode(self, request: RememberEpisodeRequest) -> RememberEpisodeResponse: + runtime = self.runtime() + prepared_media = self._prepare_media(request.media, request.media_type) + try: + modality, media_type = self._resolve_media_contract( + modality=request.modality, + prepared_media=prepared_media, + requested_media_type=request.media_type, + ) + record = EpisodicMemory( + content=request.text, + session_id=request.session_id, + turn_number=request.turn_number, + participants=list(request.participants), + summary=request.summary, + emotional_profile=dict(request.emotional_profile), + emotional_valence=request.emotional_valence, + importance=request.importance, + modality=modality, + media_ref=str(prepared_media.path) if prepared_media else None, + media_type=media_type, + source_mime_type=mimetypes.guess_type(prepared_media.path.name)[0] + if prepared_media + else request.source_mime_type, + ) + runtime.episodic_store.store(record) + return RememberEpisodeResponse(record=serialise_record(record)) + finally: + self._cleanup_prepared_media(prepared_media) + + def remember_procedure(self, request: RememberProcedureRequest) -> RememberProcedureResponse: + runtime = self.runtime() + prepared_media = self._prepare_media(request.media, request.media_type) + try: + modality, media_type = self._resolve_media_contract( + modality=request.modality, + prepared_media=prepared_media, + requested_media_type=request.media_type, + ) + record = ProceduralMemory( + content=request.content, + steps=list(request.steps), + preconditions=list(request.preconditions), + importance=request.importance, + modality=modality, + media_ref=str(prepared_media.path) if prepared_media else None, + media_type=media_type, + text_description=request.text_description, + ) + runtime.procedural_store.store(record) + return RememberProcedureResponse(record=serialise_record(record)) + finally: + self._cleanup_prepared_media(prepared_media) + + def resolve_contradiction( + self, + request: ResolveContradictionRequest, + ) -> ResolveContradictionResponse: + self.runtime().contradiction_detector.resolve_supersession( + superseded_id=request.supersede_id, + kept_id=request.keep_id, + ) + return ResolveContradictionResponse( + kept_id=request.keep_id, + superseded_id=request.supersede_id, + ) + + def record_procedure_outcome( + self, + request: RecordProcedureOutcomeRequest, + ) -> RecordProcedureOutcomeResponse: + runtime = self.runtime() + record = runtime.procedural_store.get_by_id(request.record_id) + if record is None: + raise ValueError(f"Procedural memory '{request.record_id}' not found") + runtime.procedural_store.record_outcome(request.record_id, request.success) + updated = runtime.procedural_store.get_by_id(request.record_id) + if updated is None: + raise ValueError(f"Procedural memory '{request.record_id}' not found after update") + return RecordProcedureOutcomeResponse(record=serialise_record(updated)) + + def recall_memories(self, request: RecallMemoriesRequest) -> RecallMemoriesResponse: + results = self.runtime().retriever.query( + request.query, + top_k=request.top_k, + memory_types=list(request.memory_types) if request.memory_types else None, + ) + return RecallMemoriesResponse( + results=[serialise_ranked_result(result) for result in results] + ) + + def recall_procedures(self, request: RecallProceduresRequest) -> RecallProceduresResponse: + matches = self.runtime().procedural_store.get_best_procedure_matches( + request.task, + top_k=request.top_k, + ) + return RecallProceduresResponse( + results=[serialise_procedural_match(match) for match in matches] + ) + + def recall_episodes(self, payload: dict[str, Any]) -> RecallEpisodesResponse: + runtime = self.runtime() + request = RecallEpisodesRequestAdapter.validate_python(payload) + if request.mode == "recent": + records = runtime.retriever.query_recent(request.limit) + elif request.mode == "session": + records = runtime.episodic_store.get_by_session(request.session_id) + else: + records = runtime.retriever.query_time_range(request.start, request.end) + + return RecallEpisodesResponse( + mode=request.mode, + records=[serialise_record(record) for record in records], + ) + + def get_memory(self, request: GetMemoryRequest) -> GetMemoryResponse: + record = get_memory_by_type(self.runtime(), request) + return GetMemoryResponse( + record=serialise_record(record) if record is not None else None + ) + + async def preview_forgetting_cycle( + self, + request: ForgettingCycleRequest, + ) -> ForgettingReportModel: + return await self._run_forgetting_cycle(request=request, dry_run=True) + + async def run_forgetting_cycle( + self, + request: ForgettingCycleRequest, + ) -> ForgettingReportModel: + return await self._run_forgetting_cycle(request=request, dry_run=False) + + async def _run_forgetting_cycle( + self, + *, + request: ForgettingCycleRequest, + dry_run: bool, + ) -> ForgettingReportModel: + if self._forgetting_lock.locked(): + return ForgettingReportModel(status="already_running", dry_run=dry_run) + + async with self._forgetting_lock: + report = self.runtime().forgetting_service.run_cycle(dry_run=dry_run) + return serialise_forgetting_report(report, max_decisions=request.max_decisions) + + def _prepare_media( + self, + media: MediaInputModel | None, + requested_media_type: str | None, + ) -> PreparedMedia | None: + if media is None: + return None + if media.file_path: + path = Path(media.file_path).expanduser() + if not path.exists(): + raise FileNotFoundError(f"Media file not found: {path}") + if not path.is_file(): + raise ValueError(f"Media path is not a file: {path}") + if not os.access(path, os.R_OK): + raise ValueError(f"Media file is not readable: {path}") + media_type = MediaStore.resolve_media_type(path, requested_media_type) + return PreparedMedia(path=path, media_type=media_type) + + inline = media.inline_content + if inline is None: + raise ValueError("Media input requires file_path or inline_content") + + try: + payload = base64.b64decode(inline.inline_base64, validate=True) + except binascii.Error as exc: + raise ValueError("inline_base64 must be valid base64") from exc + if not payload: + raise ValueError("Cannot store empty media payload") + + suffix = Path(inline.filename or "").suffix + if not suffix: + suffix = mimetypes.guess_extension(inline.mime_type) or "" + handle = tempfile.NamedTemporaryFile( + prefix="agentic_memory_mcp_", + suffix=suffix, + delete=False, + ) + try: + handle.write(payload) + handle.flush() + finally: + handle.close() + + path = Path(handle.name) + media_type = MediaStore.resolve_media_type(path, requested_media_type) + return PreparedMedia(path=path, media_type=media_type, cleanup=True) + + @staticmethod + def _cleanup_prepared_media(prepared_media: PreparedMedia | None) -> None: + if prepared_media is None or not prepared_media.cleanup: + return + try: + prepared_media.path.unlink(missing_ok=True) + except OSError: + logger.warning("Failed to remove temporary media file %s", prepared_media.path) + + @staticmethod + def _resolve_media_contract( + *, + modality: str, + prepared_media: PreparedMedia | None, + requested_media_type: str | None, + ) -> tuple[str, str | None]: + resolved_modality = normalize_modality(modality) + if prepared_media is None: + if resolved_modality != "text": + raise ValueError("media is required when modality is not text") + return resolved_modality, requested_media_type + + resolved_media_type = prepared_media.media_type + if requested_media_type and requested_media_type != resolved_media_type: + raise ValueError( + f"media_type '{requested_media_type}' does not match detected media type " + f"'{resolved_media_type}'" + ) + + inferred_modality = "multimodal" if resolved_media_type == "pdf" else resolved_media_type + if resolved_modality == "text": + resolved_modality = inferred_modality + elif resolved_modality == "multimodal" and resolved_media_type != "pdf": + pass + elif resolved_modality != inferred_modality: + raise ValueError( + f"modality '{resolved_modality}' does not match media type '{resolved_media_type}'" + ) + + return resolved_modality, resolved_media_type + + +def _load_fastmcp(): + try: + from mcp.server.fastmcp import FastMCP + except ImportError as exc: # pragma: no cover - import depends on optional dependency + raise RuntimeError( + "The 'mcp' package is required to run the MCP server. " + "Install it with: pip install \"mcp[cli]\"" + ) from exc + return FastMCP + + +def create_mcp_server( + server: MemoryMCPServer, +): + FastMCP = _load_fastmcp() + mcp = FastMCP( + _MCP_SERVER_NAME, + json_response=True, + streamable_http_path="/", + ) + + @mcp.tool() + def remember_fact( + content: str, + importance: float = 0.5, + category: str = "general", + domain: str | None = None, + confidence: float = 1.0, + supersedes: str | None = None, + related_ids: list[str] | None = None, + has_visual: bool = False, + modality: str = "text", + media: dict[str, Any] | None = None, + media_type: str | None = None, + text_description: str | None = None, + max_contradictions: int = DEFAULT_MAX_CONTRADICTIONS, + ) -> RememberFactResponse | ToolErrorModel: + """Store a persistent fact, preference, or durable piece of knowledge.""" + try: + request = RememberFactRequest( + content=content, + importance=importance, + category=category, + domain=domain, + confidence=confidence, + supersedes=supersedes, + related_ids=related_ids or [], + has_visual=has_visual, + modality=modality, + media=MediaInputModel.model_validate(media) if media else None, + media_type=media_type, + text_description=text_description, + max_contradictions=max_contradictions, + ) + return server.remember_fact(request) + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + def remember_episode( + session_id: str, + text: str, + turn_number: int | None = None, + participants: list[str] | None = None, + summary: str | None = None, + emotional_profile: dict[str, float] | None = None, + emotional_valence: float | None = None, + importance: float = 0.5, + modality: str = "text", + media: dict[str, Any] | None = None, + media_type: str | None = None, + source_mime_type: str | None = None, + ) -> RememberEpisodeResponse | ToolErrorModel: + """Store a concrete event, interaction, or session-specific memory.""" + try: + request = RememberEpisodeRequest( + session_id=session_id, + text=text, + turn_number=turn_number, + participants=participants or ["user", "agent"], + summary=summary, + emotional_profile=emotional_profile or {}, + emotional_valence=emotional_valence, + importance=importance, + modality=modality, + media=MediaInputModel.model_validate(media) if media else None, + media_type=media_type, + source_mime_type=source_mime_type, + ) + return server.remember_episode(request) + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + def remember_procedure( + content: str, + steps: list[str], + preconditions: list[str] | None = None, + importance: float = 0.5, + modality: str = "text", + media: dict[str, Any] | None = None, + media_type: str | None = None, + text_description: str | None = None, + ) -> RememberProcedureResponse | ToolErrorModel: + """Store a reusable workflow, skill, or step-by-step procedure.""" + try: + request = RememberProcedureRequest( + content=content, + steps=steps, + preconditions=preconditions or [], + importance=importance, + modality=modality, + media=MediaInputModel.model_validate(media) if media else None, + media_type=media_type, + text_description=text_description, + ) + return server.remember_procedure(request) + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + def resolve_contradiction( + keep_id: str, + supersede_id: str, + ) -> ResolveContradictionResponse | ToolErrorModel: + """Mark an older fact as superseded by the kept fact.""" + try: + return server.resolve_contradiction( + ResolveContradictionRequest(keep_id=keep_id, supersede_id=supersede_id) + ) + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + def record_procedure_outcome( + record_id: str, + success: bool, + ) -> RecordProcedureOutcomeResponse | ToolErrorModel: + """Record whether a stored procedure succeeded or failed when executed.""" + try: + return server.record_procedure_outcome( + RecordProcedureOutcomeRequest(record_id=record_id, success=success) + ) + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + def recall_memories( + query: str, + top_k: int = 5, + memory_types: list[str] | None = None, + ) -> RecallMemoriesResponse | ToolErrorModel: + """Recall the most relevant memories for a natural-language query.""" + try: + return server.recall_memories( + RecallMemoriesRequest(query=query, top_k=top_k, memory_types=memory_types) + ) + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + def recall_procedures( + task: str, + top_k: int = 3, + ) -> RecallProceduresResponse | ToolErrorModel: + """Recall the best procedures for a task using similarity plus outcome ranking.""" + try: + return server.recall_procedures(RecallProceduresRequest(task=task, top_k=top_k)) + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + def recall_episodes( + mode: str, + limit: int = 5, + session_id: str | None = None, + start: str | None = None, + end: str | None = None, + ) -> RecallEpisodesResponse | ToolErrorModel: + """Recall episodic memories by recent activity, session, or time range.""" + try: + payload: dict[str, Any] = {"mode": mode} + if mode == "recent": + payload["limit"] = limit + elif mode == "session": + payload["session_id"] = session_id + elif mode == "time_range": + payload["start"] = start + payload["end"] = end + return server.recall_episodes(payload) + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + def get_memory( + memory_type: str, + record_id: str, + ) -> GetMemoryResponse | ToolErrorModel: + """Fetch one memory by explicit type and record id.""" + try: + return server.get_memory( + GetMemoryRequest(memory_type=memory_type, record_id=record_id) + ) + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + def get_memory_overview() -> MemoryOverviewModel | ToolErrorModel: + """Get memory counts, recent sessions, and the latest lifecycle events.""" + try: + return server.overview() + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + async def preview_forgetting_cycle( + max_decisions: int = 50, + ) -> ForgettingReportModel | ToolErrorModel: + """Preview forgetting decisions without mutating stored memories.""" + try: + return await server.preview_forgetting_cycle( + ForgettingCycleRequest(max_decisions=max_decisions) + ) + except Exception as exc: + return _map_exception(exc).payload + + @mcp.tool() + async def run_forgetting_cycle( + max_decisions: int = 50, + ) -> ForgettingReportModel | ToolErrorModel: + """Run one forgetting cycle and return a bounded report of what happened.""" + try: + return await server.run_forgetting_cycle( + ForgettingCycleRequest(max_decisions=max_decisions) + ) + except Exception as exc: + return _map_exception(exc).payload + + return mcp + + +def create_http_app(server: MemoryMCPServer): + try: + from starlette.applications import Starlette + from starlette.responses import JSONResponse + from starlette.routing import Mount, Route + except ImportError as exc: # pragma: no cover - import depends on optional dependency + raise RuntimeError( + "Starlette is required for streamable-http transport. " + "Install the API or MCP dependencies before starting the HTTP server." + ) from exc + + mcp = create_mcp_server(server) + + async def health(_request) -> JSONResponse: + server.validate_startup() + return JSONResponse( + { + "status": "ok", + "transport": server.config.transport, + "mcp_path": server.config.streamable_http_path, + } + ) + + @asynccontextmanager + async def lifespan(_app): + server.validate_startup() + async with mcp.session_manager.run(): + yield + + return Starlette( + routes=[ + Route("/health", endpoint=health), + Mount(server.config.streamable_http_path, app=mcp.streamable_http_app()), + ], + lifespan=lifespan, + ) + + +def main() -> None: + server = MemoryMCPServer(config=ServerConfig.from_env()) + server.validate_startup() + + if server.config.transport == "stdio": + create_mcp_server(server).run(transport="stdio") + return + + try: + import uvicorn + except ImportError as exc: # pragma: no cover - import depends on optional dependency + raise RuntimeError( + "uvicorn is required for streamable-http transport." + ) from exc + + app = create_http_app(server) + uvicorn.run(app, host=server.config.host, port=server.config.port) + + +if __name__ == "__main__": + main() diff --git a/requirements.txt b/requirements.txt index 261c891..2fbd39a 100644 --- a/requirements.txt +++ b/requirements.txt @@ -2,6 +2,7 @@ chromadb==1.5.5 fastapi==0.116.1 google-genai==1.68.0 httpx==0.28.1 +mcp[cli]==1.26.0 numpy==2.4.3 python-multipart==0.0.20 python-dotenv==1.2.2 diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py new file mode 100644 index 0000000..55d724e --- /dev/null +++ b/tests/test_mcp_server.py @@ -0,0 +1,113 @@ +from __future__ import annotations + +import base64 +from pathlib import Path + +import pytest + +from mcp_server.server import MemoryMCPServer, ServerConfig +from mcp_server.schemas import ( + ForgettingCycleRequest, + GetMemoryRequest, + RememberEpisodeRequest, + RememberFactRequest, +) +from tests.helpers import HashingEmbedder, cleanup_dir, make_temp_chroma_dir + + +def make_server( + *, + transport: str = "stdio", + chroma_path: str | None = None, + media_root: str | None = None, +) -> MemoryMCPServer: + resolved_chroma = chroma_path or make_temp_chroma_dir("mcp_server_") + resolved_media_root = media_root or make_temp_chroma_dir("mcp_media_") + return MemoryMCPServer( + config=ServerConfig( + transport=transport, + chroma_path=resolved_chroma, + media_root=resolved_media_root, + embedding_dimensions=32, + ), + embedder=HashingEmbedder(dimensions=32), + ) + + +def test_server_remember_fact_round_trip(): + server = make_server() + try: + response = server.remember_fact( + RememberFactRequest( + content="The MCP server should reuse one runtime per process", + category="architecture", + ) + ) + + assert response.record.content == "The MCP server should reuse one runtime per process" + assert response.contradiction_check.status == "completed" + + stored = server.get_memory( + GetMemoryRequest(memory_type="semantic", record_id=response.record.id) + ) + assert stored.record is not None + assert stored.record.id == response.record.id + finally: + cleanup_dir(server.config.chroma_path or "") + cleanup_dir(server.config.media_root or "") + + +def test_server_materializes_inline_media_for_episode_storage(): + server = make_server() + try: + payload = base64.b64encode(b"diagram contents for multimodal recall").decode("ascii") + response = server.remember_episode( + RememberEpisodeRequest( + session_id="session-1", + text="Architecture diagram from planning session", + media={ + "inline_content": { + "inline_base64": payload, + "mime_type": "image/png", + "filename": "diagram.png", + } + }, + ) + ) + + assert response.record.media_ref is not None + assert Path(response.record.media_ref).exists() + assert response.record.modality == "image" + finally: + cleanup_dir(server.config.chroma_path or "") + cleanup_dir(server.config.media_root or "") + + +@pytest.mark.anyio +async def test_server_forgetting_cycle_is_single_flight(): + server = make_server() + try: + await server._forgetting_lock.acquire() # type: ignore[attr-defined] + try: + second = await server.preview_forgetting_cycle(ForgettingCycleRequest(max_decisions=5)) + finally: + server._forgetting_lock.release() # type: ignore[attr-defined] + assert second.status == "already_running" + finally: + cleanup_dir(server.config.chroma_path or "") + cleanup_dir(server.config.media_root or "") + + +def test_streamable_http_config_allows_relative_paths_but_keeps_transport_valid(): + config = ServerConfig( + transport="streamable-http", + chroma_path="./chroma_db", + media_root="./data/media", + ).validate() + + assert config.transport == "streamable-http" + + +def test_invalid_transport_is_rejected(): + with pytest.raises(ValueError): + ServerConfig(transport="http2").validate() From f01f9f17a924673b125a5efe5406c600d396ae6d Mon Sep 17 00:00:00 2001 From: Atharva-Kanherkar Date: Mon, 30 Mar 2026 15:46:59 +0530 Subject: [PATCH 2/2] Fix MCP server cleanup and forgetting execution --- mcp_server/server.py | 21 +++++++++++++-- tests/test_mcp_server.py | 58 ++++++++++++++++++++++++++++++++++++++-- 2 files changed, 75 insertions(+), 4 deletions(-) diff --git a/mcp_server/server.py b/mcp_server/server.py index a3bfac5..33c77f9 100644 --- a/mcp_server/server.py +++ b/mcp_server/server.py @@ -3,6 +3,7 @@ import asyncio import base64 import binascii +import functools import logging import mimetypes import os @@ -403,7 +404,13 @@ async def _run_forgetting_cycle( return ForgettingReportModel(status="already_running", dry_run=dry_run) async with self._forgetting_lock: - report = self.runtime().forgetting_service.run_cycle(dry_run=dry_run) + report = await asyncio.get_running_loop().run_in_executor( + None, + functools.partial( + self.runtime().forgetting_service.run_cycle, + dry_run=dry_run, + ), + ) return serialise_forgetting_report(report, max_decisions=request.max_decisions) def _prepare_media( @@ -450,7 +457,11 @@ def _prepare_media( handle.close() path = Path(handle.name) - media_type = MediaStore.resolve_media_type(path, requested_media_type) + try: + media_type = MediaStore.resolve_media_type(path, requested_media_type) + except Exception: + path.unlink(missing_ok=True) + raise return PreparedMedia(path=path, media_type=media_type, cleanup=True) @staticmethod @@ -676,12 +687,18 @@ def recall_episodes( ) -> RecallEpisodesResponse | ToolErrorModel: """Recall episodic memories by recent activity, session, or time range.""" try: + if mode not in {"recent", "session", "time_range"}: + raise ValueError("mode must be one of: recent, session, time_range") payload: dict[str, Any] = {"mode": mode} if mode == "recent": payload["limit"] = limit elif mode == "session": + if not session_id: + raise ValueError("session_id is required when mode='session'") payload["session_id"] = session_id elif mode == "time_range": + if not start or not end: + raise ValueError("start and end are required when mode='time_range'") payload["start"] = start payload["end"] = end return server.recall_episodes(payload) diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py index 55d724e..cbafe00 100644 --- a/tests/test_mcp_server.py +++ b/tests/test_mcp_server.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import base64 from pathlib import Path @@ -83,15 +84,63 @@ def test_server_materializes_inline_media_for_episode_storage(): cleanup_dir(server.config.media_root or "") +def test_prepare_media_cleans_up_temp_file_when_media_type_resolution_fails(): + server = make_server() + before = {path for path in Path("/tmp").glob("agentic_memory_mcp_*")} + try: + payload = base64.b64encode(b"not a supported media type").decode("ascii") + with pytest.raises(ValueError): + server._prepare_media( # type: ignore[attr-defined] + RememberEpisodeRequest( + session_id="session-1", + text="Unsupported inline media", + media={ + "inline_content": { + "inline_base64": payload, + "mime_type": "application/octet-stream", + "filename": "blob.bin", + } + }, + ).media, + None, + ) + finally: + after = {path for path in Path("/tmp").glob("agentic_memory_mcp_*")} + cleanup_dir(server.config.chroma_path or "") + cleanup_dir(server.config.media_root or "") + + assert after == before + + @pytest.mark.anyio async def test_server_forgetting_cycle_is_single_flight(): server = make_server() try: - await server._forgetting_lock.acquire() # type: ignore[attr-defined] + loop = asyncio.get_running_loop() + started = asyncio.Event() + release = asyncio.Event() + + original = server.runtime().forgetting_service.run_cycle + + def slow_run_cycle(*, dry_run: bool): + started.set() + future = asyncio.run_coroutine_threadsafe(release.wait(), loop) + future.result(timeout=2) + return original(dry_run=dry_run) + + server.runtime().forgetting_service.run_cycle = slow_run_cycle # type: ignore[method-assign] try: + first_task = asyncio.create_task( + server.preview_forgetting_cycle(ForgettingCycleRequest(max_decisions=5)) + ) + await asyncio.wait_for(started.wait(), timeout=1) second = await server.preview_forgetting_cycle(ForgettingCycleRequest(max_decisions=5)) + release.set() + first = await asyncio.wait_for(first_task, timeout=1) finally: - server._forgetting_lock.release() # type: ignore[attr-defined] + server.runtime().forgetting_service.run_cycle = original # type: ignore[method-assign] + + assert first.status == "completed" assert second.status == "already_running" finally: cleanup_dir(server.config.chroma_path or "") @@ -111,3 +160,8 @@ def test_streamable_http_config_allows_relative_paths_but_keeps_transport_valid( def test_invalid_transport_is_rejected(): with pytest.raises(ValueError): ServerConfig(transport="http2").validate() + + +def test_invalid_port_is_rejected(): + with pytest.raises(ValueError): + ServerConfig(port=0).validate()