From f641b30d8342c9bdbb1db8d0a5a21882bb9126df Mon Sep 17 00:00:00 2001 From: DK09876 Date: Tue, 10 Feb 2026 14:40:43 -0700 Subject: [PATCH] feat: add mental model CRUD tools to MCP server (#337) * Add mental model CRUD tools to MCP server Expose mental models (pinned reflections) as 6 new MCP tools: - list_mental_models: List with optional tag filtering - get_mental_model: Get by ID - create_mental_model: Create with async content generation - update_mental_model: Update name/source_query/tags - delete_mental_model: Delete by ID - refresh_mental_model: Re-run source query to update content Both multi-bank (bank_id param) and single-bank modes supported, following the same patterns as existing retain/recall/reflect tools. Co-Authored-By: Claude Opus 4.6 * fix: include mental model tools in single-bank MCP mode and update tests The single-bank mode tool set was hardcoded to only retain/recall/reflect, excluding the new mental model tools. Updated all 3 test layers (unit, routing, HTTP integration) to assert mental model tool exposure. Co-Authored-By: Claude Opus 4.6 * fix: update extension test tool count for mental model tools Co-Authored-By: Claude Opus 4.6 * fix: move mental model usage metering into engine for MCP support Mental model validation hooks (validate_mental_model_get, validate_mental_model_refresh) were only called in REST HTTP handlers, not in the engine. MCP tools call engine methods directly, so usage metering was skipped entirely for MCP mental model operations. Moved pre-validation and post-completion hooks into memory_engine.py (matching the retain/recall/reflect pattern) and removed the duplicate code from http.py. Co-Authored-By: Claude Opus 4.6 * fix: remove double validation from create_mental_model and add internal checks - Remove pre-validation from create_mental_model since callers always call submit_async_refresh_mental_model next (which validates), preventing double credit checks - Add is_internal checks to mental model metering validators (matching the existing pattern for recall/reflect) so background worker tasks skip billing Co-Authored-By: Claude Opus 4.6 --------- Co-authored-by: Claude Opus 4.6 --- hindsight-api/hindsight_api/api/http.py | 70 --- hindsight-api/hindsight_api/api/mcp.py | 14 +- .../hindsight_api/engine/memory_engine.py | 46 +- hindsight-api/hindsight_api/mcp_tools.py | 560 +++++++++++++++++- .../tests/test_mcp_endpoint_routing.py | 12 +- hindsight-api/tests/test_mcp_extension.py | 4 +- hindsight-api/tests/test_mcp_routing.py | 115 ++-- hindsight-api/tests/test_mcp_tools.py | 487 ++++++++++++++- 8 files changed, 1191 insertions(+), 117 deletions(-) diff --git a/hindsight-api/hindsight_api/api/http.py b/hindsight-api/hindsight_api/api/http.py index bfb65769..beef29aa 100644 --- a/hindsight-api/hindsight_api/api/http.py +++ b/hindsight-api/hindsight_api/api/http.py @@ -2354,23 +2354,6 @@ def _register_routes(app: FastAPI): ): """Get a mental model by ID.""" try: - # Pre-operation validation hook - validator = app.state.memory._operation_validator - if validator: - from hindsight_api.extensions.operation_validator import MentalModelGetContext - - ctx = MentalModelGetContext( - bank_id=bank_id, - mental_model_id=mental_model_id, - request_context=request_context, - ) - validation = await validator.validate_mental_model_get(ctx) - if not validation.allowed: - raise OperationValidationError( - validation.reason or "Operation not allowed", - status_code=validation.status_code, - ) - mental_model = await app.state.memory.get_mental_model( bank_id=bank_id, mental_model_id=mental_model_id, @@ -2379,25 +2362,6 @@ def _register_routes(app: FastAPI): if mental_model is None: raise HTTPException(status_code=404, detail=f"Mental model '{mental_model_id}' not found") - # Post-operation hook - if validator: - from hindsight_api.extensions.operation_validator import MentalModelGetResult - - content = mental_model.get("content", "") - output_tokens = len(content) // 4 if content else 0 - - result_ctx = MentalModelGetResult( - bank_id=bank_id, - mental_model_id=mental_model_id, - request_context=request_context, - output_tokens=output_tokens, - success=True, - ) - try: - await validator.on_mental_model_get_complete(result_ctx) - except Exception as hook_err: - logger.warning(f"Post-mental-model-get hook error (non-fatal): {hook_err}") - return MentalModelResponse(**mental_model) except (AuthenticationError, HTTPException): raise @@ -2427,23 +2391,6 @@ def _register_routes(app: FastAPI): ): """Create a mental model (async - returns operation_id).""" try: - # Pre-operation validation hook - validator = app.state.memory._operation_validator - if validator: - from hindsight_api.extensions.operation_validator import MentalModelRefreshContext - - ctx = MentalModelRefreshContext( - bank_id=bank_id, - mental_model_id=None, # Not yet created - request_context=request_context, - ) - validation = await validator.validate_mental_model_refresh(ctx) - if not validation.allowed: - raise OperationValidationError( - validation.reason or "Operation not allowed", - status_code=validation.status_code, - ) - # 1. Create the mental model with placeholder content mental_model = await app.state.memory.create_mental_model( bank_id=bank_id, @@ -2491,23 +2438,6 @@ def _register_routes(app: FastAPI): ): """Refresh a mental model by re-running its source query (async).""" try: - # Pre-operation validation hook - validator = app.state.memory._operation_validator - if validator: - from hindsight_api.extensions.operation_validator import MentalModelRefreshContext - - ctx = MentalModelRefreshContext( - bank_id=bank_id, - mental_model_id=mental_model_id, - request_context=request_context, - ) - validation = await validator.validate_mental_model_refresh(ctx) - if not validation.allowed: - raise OperationValidationError( - validation.reason or "Operation not allowed", - status_code=validation.status_code, - ) - result = await app.state.memory.submit_async_refresh_mental_model( bank_id=bank_id, mental_model_id=mental_model_id, diff --git a/hindsight-api/hindsight_api/api/mcp.py b/hindsight-api/hindsight_api/api/mcp.py index 279fc2bb..640f0a80 100644 --- a/hindsight-api/hindsight_api/api/mcp.py +++ b/hindsight-api/hindsight_api/api/mcp.py @@ -90,7 +90,19 @@ def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP: tenant_id_resolver=get_current_tenant_id, # Propagate tenant_id for usage metering api_key_id_resolver=get_current_api_key_id, # Propagate api_key_id for usage metering include_bank_id_param=multi_bank, - tools=None if multi_bank else {"retain", "recall", "reflect"}, # Scoped tools for single-bank mode + tools=None + if multi_bank + else { + "retain", + "recall", + "reflect", + "list_mental_models", + "get_mental_model", + "create_mental_model", + "update_mental_model", + "delete_mental_model", + "refresh_mental_model", + }, # Scoped tools for single-bank mode (excludes bank management: list_banks, create_bank) retain_fire_and_forget=False, # HTTP MCP supports sync/async modes ) diff --git a/hindsight-api/hindsight_api/engine/memory_engine.py b/hindsight-api/hindsight_api/engine/memory_engine.py index c21285da..b2786e32 100644 --- a/hindsight-api/hindsight_api/engine/memory_engine.py +++ b/hindsight-api/hindsight_api/engine/memory_engine.py @@ -4693,6 +4693,18 @@ class MemoryEngine(MemoryEngineInterface): Pinned mental model dict or None if not found """ await self._authenticate_tenant(request_context) + + # Pre-operation validation (credit check / usage metering) + if self._operation_validator: + from hindsight_api.extensions.operation_validator import MentalModelGetContext + + ctx = MentalModelGetContext( + bank_id=bank_id, + mental_model_id=mental_model_id, + request_context=request_context, + ) + await self._validate_operation(self._operation_validator.validate_mental_model_get(ctx)) + pool = await self._get_pool() async with acquire_with_retry(pool) as conn: @@ -4708,7 +4720,28 @@ class MemoryEngine(MemoryEngineInterface): mental_model_id, ) - return self._row_to_mental_model(row) if row else None + result = self._row_to_mental_model(row) if row else None + + # Post-operation hook (usage recording) + if result and self._operation_validator: + from hindsight_api.extensions.operation_validator import MentalModelGetResult + + content = result.get("content", "") + output_tokens = len(content) // 4 if content else 0 + + result_ctx = MentalModelGetResult( + bank_id=bank_id, + mental_model_id=mental_model_id, + request_context=request_context, + output_tokens=output_tokens, + success=True, + ) + try: + await self._operation_validator.on_mental_model_get_complete(result_ctx) + except Exception as hook_err: + logger.warning(f"Post-mental-model-get hook error (non-fatal): {hook_err}") + + return result async def create_mental_model( self, @@ -5699,6 +5732,17 @@ class MemoryEngine(MemoryEngineInterface): """ await self._authenticate_tenant(request_context) + # Pre-operation validation (credit check) + if self._operation_validator: + from hindsight_api.extensions.operation_validator import MentalModelRefreshContext + + ctx = MentalModelRefreshContext( + bank_id=bank_id, + mental_model_id=mental_model_id, + request_context=request_context, + ) + await self._validate_operation(self._operation_validator.validate_mental_model_refresh(ctx)) + # Verify mental model exists mental_model = await self.get_mental_model(bank_id, mental_model_id, request_context=request_context) if not mental_model: diff --git a/hindsight-api/hindsight_api/mcp_tools.py b/hindsight-api/hindsight_api/mcp_tools.py index 3d652e56..3f4dced3 100644 --- a/hindsight-api/hindsight_api/mcp_tools.py +++ b/hindsight-api/hindsight_api/mcp_tools.py @@ -127,7 +127,19 @@ def register_mcp_tools( memory: MemoryEngine instance config: Tool configuration """ - tools_to_register = config.tools or {"retain", "recall", "reflect", "list_banks", "create_bank"} + tools_to_register = config.tools or { + "retain", + "recall", + "reflect", + "list_banks", + "create_bank", + "list_mental_models", + "get_mental_model", + "create_mental_model", + "update_mental_model", + "delete_mental_model", + "refresh_mental_model", + } if "retain" in tools_to_register: _register_retain(mcp, memory, config) @@ -144,6 +156,25 @@ def register_mcp_tools( if "create_bank" in tools_to_register: _register_create_bank(mcp, memory, config) + # Mental model tools + if "list_mental_models" in tools_to_register: + _register_list_mental_models(mcp, memory, config) + + if "get_mental_model" in tools_to_register: + _register_get_mental_model(mcp, memory, config) + + if "create_mental_model" in tools_to_register: + _register_create_mental_model(mcp, memory, config) + + if "update_mental_model" in tools_to_register: + _register_update_mental_model(mcp, memory, config) + + if "delete_mental_model" in tools_to_register: + _register_delete_mental_model(mcp, memory, config) + + if "refresh_mental_model" in tools_to_register: + _register_refresh_mental_model(mcp, memory, config) + def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None: """Register the retain tool.""" @@ -519,3 +550,530 @@ def _register_create_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo except Exception as e: logger.error(f"Error creating bank: {e}", exc_info=True) return f'{{"error": "{e}"}}' + + +# ========================================================================= +# MENTAL MODEL TOOLS +# ========================================================================= + + +def _register_list_mental_models(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None: + """Register the list_mental_models tool.""" + + if config.include_bank_id_param: + + @mcp.tool() + async def list_mental_models( + tags: list[str] | None = None, + bank_id: str | None = None, + ) -> str: + """ + List mental models (pinned reflections) for a memory bank. + + Mental models are living documents that stay current by periodically re-running + a source query through reflect. Use them to maintain up-to-date summaries, + preferences, or synthesized knowledge. + + Args: + tags: Optional tags to filter by (returns models matching any tag) + bank_id: Optional bank to list from (defaults to session bank). Use for cross-bank operations. + """ + try: + target_bank = bank_id or config.bank_id_resolver() + if target_bank is None: + return '{"error": "No bank_id configured", "items": []}' + + models = await memory.list_mental_models( + bank_id=target_bank, + tags=tags, + request_context=_get_request_context(config), + ) + return json.dumps({"items": models}, indent=2, default=str) + except Exception as e: + logger.error(f"Error listing mental models: {e}", exc_info=True) + return f'{{"error": "{e}", "items": []}}' + + else: + + @mcp.tool() + async def list_mental_models( + tags: list[str] | None = None, + ) -> dict: + """ + List mental models (pinned reflections) for this memory bank. + + Mental models are living documents that stay current by periodically re-running + a source query through reflect. Use them to maintain up-to-date summaries, + preferences, or synthesized knowledge. + + Args: + tags: Optional tags to filter by (returns models matching any tag) + """ + try: + target_bank = config.bank_id_resolver() + if target_bank is None: + return {"error": "No bank_id configured", "items": []} + + models = await memory.list_mental_models( + bank_id=target_bank, + tags=tags, + request_context=_get_request_context(config), + ) + return {"items": models} + except Exception as e: + logger.error(f"Error listing mental models: {e}", exc_info=True) + return {"error": str(e), "items": []} + + +def _register_get_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None: + """Register the get_mental_model tool.""" + + if config.include_bank_id_param: + + @mcp.tool() + async def get_mental_model( + mental_model_id: str, + bank_id: str | None = None, + ) -> str: + """ + Get a specific mental model by ID. + + Returns the full mental model including its generated content, source query, + and metadata. Use list_mental_models first to discover available model IDs. + + Args: + mental_model_id: The ID of the mental model to retrieve + bank_id: Optional bank (defaults to session bank). Use for cross-bank operations. + """ + try: + target_bank = bank_id or config.bank_id_resolver() + if target_bank is None: + return '{"error": "No bank_id configured"}' + + model = await memory.get_mental_model( + bank_id=target_bank, + mental_model_id=mental_model_id, + request_context=_get_request_context(config), + ) + if model is None: + return json.dumps({"error": f"Mental model '{mental_model_id}' not found"}) + return json.dumps(model, indent=2, default=str) + except Exception as e: + logger.error(f"Error getting mental model: {e}", exc_info=True) + return f'{{"error": "{e}"}}' + + else: + + @mcp.tool() + async def get_mental_model( + mental_model_id: str, + ) -> dict: + """ + Get a specific mental model by ID. + + Returns the full mental model including its generated content, source query, + and metadata. Use list_mental_models first to discover available model IDs. + + Args: + mental_model_id: The ID of the mental model to retrieve + """ + try: + target_bank = config.bank_id_resolver() + if target_bank is None: + return {"error": "No bank_id configured"} + + model = await memory.get_mental_model( + bank_id=target_bank, + mental_model_id=mental_model_id, + request_context=_get_request_context(config), + ) + if model is None: + return {"error": f"Mental model '{mental_model_id}' not found"} + return model + except Exception as e: + logger.error(f"Error getting mental model: {e}", exc_info=True) + return {"error": str(e)} + + +def _register_create_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None: + """Register the create_mental_model tool.""" + + if config.include_bank_id_param: + + @mcp.tool() + async def create_mental_model( + name: str, + source_query: str, + mental_model_id: str | None = None, + tags: list[str] | None = None, + max_tokens: int = 2048, + bank_id: str | None = None, + ) -> str: + """ + Create a new mental model (pinned reflection). + + A mental model is a living document generated by running the source_query through + reflect. The content is auto-generated asynchronously - use the returned operation_id + to track progress. + + EXAMPLES: + - name="Coding Preferences", source_query="What coding patterns and tools does the user prefer?" + - name="Project Goals", source_query="What are the user's current project goals and priorities?" + - name="Communication Style", source_query="How does the user prefer to communicate?" + + Args: + name: Human-readable name for the mental model + source_query: The query to run through reflect to generate content + mental_model_id: Optional custom ID (alphanumeric lowercase with hyphens). Auto-generated if not provided. + tags: Optional tags for scoped visibility filtering + max_tokens: Maximum tokens for generated content (256-8192, default: 2048) + bank_id: Optional bank (defaults to session bank). Use for cross-bank operations. + """ + try: + target_bank = bank_id or config.bank_id_resolver() + if target_bank is None: + return '{"error": "No bank_id configured"}' + + request_context = _get_request_context(config) + + # Create with placeholder content + model = await memory.create_mental_model( + bank_id=target_bank, + name=name, + source_query=source_query, + content="Generating content...", + mental_model_id=mental_model_id, + tags=tags, + max_tokens=max_tokens, + request_context=request_context, + ) + + # Schedule async refresh to generate actual content + result = await memory.submit_async_refresh_mental_model( + bank_id=target_bank, + mental_model_id=model["id"], + request_context=request_context, + ) + + return json.dumps( + { + "mental_model_id": model["id"], + "operation_id": result["operation_id"], + "status": "created", + "message": f"Mental model '{name}' created. Content is being generated asynchronously.", + } + ) + except ValueError as e: + return json.dumps({"error": str(e)}) + except Exception as e: + logger.error(f"Error creating mental model: {e}", exc_info=True) + return f'{{"error": "{e}"}}' + + else: + + @mcp.tool() + async def create_mental_model( + name: str, + source_query: str, + mental_model_id: str | None = None, + tags: list[str] | None = None, + max_tokens: int = 2048, + ) -> dict: + """ + Create a new mental model (pinned reflection). + + A mental model is a living document generated by running the source_query through + reflect. The content is auto-generated asynchronously - use the returned operation_id + to track progress. + + EXAMPLES: + - name="Coding Preferences", source_query="What coding patterns and tools does the user prefer?" + - name="Project Goals", source_query="What are the user's current project goals and priorities?" + - name="Communication Style", source_query="How does the user prefer to communicate?" + + Args: + name: Human-readable name for the mental model + source_query: The query to run through reflect to generate content + mental_model_id: Optional custom ID (alphanumeric lowercase with hyphens). Auto-generated if not provided. + tags: Optional tags for scoped visibility filtering + max_tokens: Maximum tokens for generated content (256-8192, default: 2048) + """ + try: + target_bank = config.bank_id_resolver() + if target_bank is None: + return {"error": "No bank_id configured"} + + request_context = _get_request_context(config) + + model = await memory.create_mental_model( + bank_id=target_bank, + name=name, + source_query=source_query, + content="Generating content...", + mental_model_id=mental_model_id, + tags=tags, + max_tokens=max_tokens, + request_context=request_context, + ) + + result = await memory.submit_async_refresh_mental_model( + bank_id=target_bank, + mental_model_id=model["id"], + request_context=request_context, + ) + + return { + "mental_model_id": model["id"], + "operation_id": result["operation_id"], + "status": "created", + "message": f"Mental model '{name}' created. Content is being generated asynchronously.", + } + except ValueError as e: + return {"error": str(e)} + except Exception as e: + logger.error(f"Error creating mental model: {e}", exc_info=True) + return {"error": str(e)} + + +def _register_update_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None: + """Register the update_mental_model tool.""" + + if config.include_bank_id_param: + + @mcp.tool() + async def update_mental_model( + mental_model_id: str, + name: str | None = None, + source_query: str | None = None, + max_tokens: int | None = None, + tags: list[str] | None = None, + bank_id: str | None = None, + ) -> str: + """ + Update a mental model's metadata. + + Changes the name, source query, or tags of an existing mental model. + To regenerate the content, use refresh_mental_model after updating the source query. + + Args: + mental_model_id: The ID of the mental model to update + name: New name (leave None to keep current) + source_query: New source query (leave None to keep current) + max_tokens: New max tokens for content generation (256-8192, leave None to keep current) + tags: New tags (leave None to keep current) + bank_id: Optional bank (defaults to session bank). Use for cross-bank operations. + """ + try: + target_bank = bank_id or config.bank_id_resolver() + if target_bank is None: + return '{"error": "No bank_id configured"}' + + model = await memory.update_mental_model( + bank_id=target_bank, + mental_model_id=mental_model_id, + name=name, + source_query=source_query, + max_tokens=max_tokens, + tags=tags, + request_context=_get_request_context(config), + ) + if model is None: + return json.dumps({"error": f"Mental model '{mental_model_id}' not found"}) + return json.dumps(model, indent=2, default=str) + except Exception as e: + logger.error(f"Error updating mental model: {e}", exc_info=True) + return f'{{"error": "{e}"}}' + + else: + + @mcp.tool() + async def update_mental_model( + mental_model_id: str, + name: str | None = None, + source_query: str | None = None, + max_tokens: int | None = None, + tags: list[str] | None = None, + ) -> dict: + """ + Update a mental model's metadata. + + Changes the name, source query, or tags of an existing mental model. + To regenerate the content, use refresh_mental_model after updating the source query. + + Args: + mental_model_id: The ID of the mental model to update + name: New name (leave None to keep current) + source_query: New source query (leave None to keep current) + max_tokens: New max tokens for content generation (256-8192, leave None to keep current) + tags: New tags (leave None to keep current) + """ + try: + target_bank = config.bank_id_resolver() + if target_bank is None: + return {"error": "No bank_id configured"} + + model = await memory.update_mental_model( + bank_id=target_bank, + mental_model_id=mental_model_id, + name=name, + source_query=source_query, + max_tokens=max_tokens, + tags=tags, + request_context=_get_request_context(config), + ) + if model is None: + return {"error": f"Mental model '{mental_model_id}' not found"} + return model + except Exception as e: + logger.error(f"Error updating mental model: {e}", exc_info=True) + return {"error": str(e)} + + +def _register_delete_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None: + """Register the delete_mental_model tool.""" + + if config.include_bank_id_param: + + @mcp.tool() + async def delete_mental_model( + mental_model_id: str, + bank_id: str | None = None, + ) -> str: + """ + Delete a mental model. + + Permanently removes a mental model and its generated content. + + Args: + mental_model_id: The ID of the mental model to delete + bank_id: Optional bank (defaults to session bank). Use for cross-bank operations. + """ + try: + target_bank = bank_id or config.bank_id_resolver() + if target_bank is None: + return '{"error": "No bank_id configured"}' + + deleted = await memory.delete_mental_model( + bank_id=target_bank, + mental_model_id=mental_model_id, + request_context=_get_request_context(config), + ) + if not deleted: + return json.dumps({"error": f"Mental model '{mental_model_id}' not found"}) + return json.dumps({"status": "deleted", "mental_model_id": mental_model_id}) + except Exception as e: + logger.error(f"Error deleting mental model: {e}", exc_info=True) + return f'{{"error": "{e}"}}' + + else: + + @mcp.tool() + async def delete_mental_model( + mental_model_id: str, + ) -> dict: + """ + Delete a mental model. + + Permanently removes a mental model and its generated content. + + Args: + mental_model_id: The ID of the mental model to delete + """ + try: + target_bank = config.bank_id_resolver() + if target_bank is None: + return {"error": "No bank_id configured"} + + deleted = await memory.delete_mental_model( + bank_id=target_bank, + mental_model_id=mental_model_id, + request_context=_get_request_context(config), + ) + if not deleted: + return {"error": f"Mental model '{mental_model_id}' not found"} + return {"status": "deleted", "mental_model_id": mental_model_id} + except Exception as e: + logger.error(f"Error deleting mental model: {e}", exc_info=True) + return {"error": str(e)} + + +def _register_refresh_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None: + """Register the refresh_mental_model tool.""" + + if config.include_bank_id_param: + + @mcp.tool() + async def refresh_mental_model( + mental_model_id: str, + bank_id: str | None = None, + ) -> str: + """ + Refresh a mental model by re-running its source query. + + Schedules an async task to re-run the source query through reflect and update the + mental model's content with fresh results. Use this after adding new memories or + when the mental model's content may be stale. + + Args: + mental_model_id: The ID of the mental model to refresh + bank_id: Optional bank (defaults to session bank). Use for cross-bank operations. + """ + try: + target_bank = bank_id or config.bank_id_resolver() + if target_bank is None: + return '{"error": "No bank_id configured"}' + + result = await memory.submit_async_refresh_mental_model( + bank_id=target_bank, + mental_model_id=mental_model_id, + request_context=_get_request_context(config), + ) + return json.dumps( + { + "operation_id": result["operation_id"], + "status": "queued", + "message": f"Refresh queued for mental model '{mental_model_id}'.", + } + ) + except ValueError as e: + return json.dumps({"error": str(e)}) + except Exception as e: + logger.error(f"Error refreshing mental model: {e}", exc_info=True) + return f'{{"error": "{e}"}}' + + else: + + @mcp.tool() + async def refresh_mental_model( + mental_model_id: str, + ) -> dict: + """ + Refresh a mental model by re-running its source query. + + Schedules an async task to re-run the source query through reflect and update the + mental model's content with fresh results. Use this after adding new memories or + when the mental model's content may be stale. + + Args: + mental_model_id: The ID of the mental model to refresh + """ + try: + target_bank = config.bank_id_resolver() + if target_bank is None: + return {"error": "No bank_id configured"} + + result = await memory.submit_async_refresh_mental_model( + bank_id=target_bank, + mental_model_id=mental_model_id, + request_context=_get_request_context(config), + ) + return { + "operation_id": result["operation_id"], + "status": "queued", + "message": f"Refresh queued for mental model '{mental_model_id}'.", + } + except ValueError as e: + return {"error": str(e)} + except Exception as e: + logger.error(f"Error refreshing mental model: {e}", exc_info=True) + return {"error": str(e)} diff --git a/hindsight-api/tests/test_mcp_endpoint_routing.py b/hindsight-api/tests/test_mcp_endpoint_routing.py index 6bbcc97e..38bcfcf3 100644 --- a/hindsight-api/tests/test_mcp_endpoint_routing.py +++ b/hindsight-api/tests/test_mcp_endpoint_routing.py @@ -39,12 +39,18 @@ async def test_mcp_endpoint_routing_integration(memory): multi_tools = {t.name for t in multi_result.tools} - # Multi-bank should have all tools including bank management + # Multi-bank should have all tools including bank management and mental models assert "retain" in multi_tools assert "recall" in multi_tools assert "reflect" in multi_tools assert "list_banks" in multi_tools, "Multi-bank should expose list_banks" assert "create_bank" in multi_tools, "Multi-bank should expose create_bank" + assert "list_mental_models" in multi_tools, "Multi-bank should expose list_mental_models" + assert "create_mental_model" in multi_tools, "Multi-bank should expose create_mental_model" + assert "get_mental_model" in multi_tools, "Multi-bank should expose get_mental_model" + assert "update_mental_model" in multi_tools, "Multi-bank should expose update_mental_model" + assert "delete_mental_model" in multi_tools, "Multi-bank should expose delete_mental_model" + assert "refresh_mental_model" in multi_tools, "Multi-bank should expose refresh_mental_model" # Multi-bank retain should have bank_id parameter retain_tool = next((t for t in multi_result.tools if t.name == "retain"), None) @@ -64,10 +70,12 @@ async def test_mcp_endpoint_routing_integration(memory): single_tools = {t.name for t in single_result.tools} - # Single-bank should only have scoped tools (no bank management) + # Single-bank should have scoped tools including mental models (no bank management) assert "retain" in single_tools assert "recall" in single_tools assert "reflect" in single_tools + assert "list_mental_models" in single_tools, "Single-bank should expose list_mental_models" + assert "create_mental_model" in single_tools, "Single-bank should expose create_mental_model" assert "list_banks" not in single_tools, "Single-bank should NOT expose list_banks" assert "create_bank" not in single_tools, "Single-bank should NOT expose create_bank" diff --git a/hindsight-api/tests/test_mcp_extension.py b/hindsight-api/tests/test_mcp_extension.py index f69013cc..39cd6133 100644 --- a/hindsight-api/tests/test_mcp_extension.py +++ b/hindsight-api/tests/test_mcp_extension.py @@ -165,5 +165,5 @@ class TestMCPExtensionIntegration: assert "create_bank" in tools # Extension tool also present assert "test_extension_tool" in tools - # Total: 5 core + 1 extension = 6 tools - assert len(tools) == 6 + # Total: 11 core + 1 extension = 12 tools + assert len(tools) == 12 diff --git a/hindsight-api/tests/test_mcp_routing.py b/hindsight-api/tests/test_mcp_routing.py index 02b5008d..d20d523e 100644 --- a/hindsight-api/tests/test_mcp_routing.py +++ b/hindsight-api/tests/test_mcp_routing.py @@ -1,8 +1,9 @@ """Test MCP server routing with dynamic bank_id.""" -import pytest from unittest.mock import AsyncMock, MagicMock +import pytest + @pytest.fixture def mock_memory(): @@ -17,7 +18,7 @@ def mock_memory(): @pytest.mark.asyncio async def test_mcp_context_variable(): """Test that context variable works correctly.""" - from hindsight_api.api.mcp import get_current_bank_id, _current_bank_id + from hindsight_api.api.mcp import _current_bank_id, get_current_bank_id # Initially None assert get_current_bank_id() is None @@ -36,7 +37,7 @@ async def test_mcp_context_variable(): @pytest.mark.asyncio async def test_mcp_tools_use_context_bank_id(mock_memory): """Test that MCP tools use bank_id from context.""" - from hindsight_api.api.mcp import create_mcp_server, _current_bank_id + from hindsight_api.api.mcp import _current_bank_id, create_mcp_server mcp_server = create_mcp_server(mock_memory) @@ -62,6 +63,7 @@ async def test_mcp_tools_use_context_bank_id(mock_memory): def test_path_parsing_logic(): """Test the path parsing logic for bank_id extraction.""" + def parse_path(path): """Simulate the path parsing logic from MCPMiddleware.""" if not path.startswith("/") or len(path) <= 1: @@ -102,7 +104,7 @@ def test_path_parsing_logic(): @pytest.mark.asyncio async def test_api_key_context_variable(): """Test that API key context variable works correctly.""" - from hindsight_api.api.mcp import get_current_api_key, _current_api_key + from hindsight_api.api.mcp import _current_api_key, get_current_api_key # Initially None assert get_current_api_key() is None @@ -121,7 +123,7 @@ async def test_api_key_context_variable(): @pytest.mark.asyncio async def test_mcp_tools_propagate_api_key(mock_memory): """Test that MCP tools propagate API key to RequestContext.""" - from hindsight_api.api.mcp import create_mcp_server, _current_bank_id, _current_api_key + from hindsight_api.api.mcp import _current_api_key, _current_bank_id, create_mcp_server mcp_server = create_mcp_server(mock_memory) tools = mcp_server._tool_manager._tools @@ -147,8 +149,10 @@ async def test_mcp_tools_propagate_api_key(mock_memory): async def test_tenant_id_context_variable(): """Test that tenant_id and api_key_id context variables work correctly.""" from hindsight_api.api.mcp import ( - get_current_tenant_id, _current_tenant_id, - get_current_api_key_id, _current_api_key_id, + _current_api_key_id, + _current_tenant_id, + get_current_api_key_id, + get_current_tenant_id, ) # Initially None @@ -179,9 +183,11 @@ async def test_mcp_tools_propagate_tenant_id_and_api_key_id(mock_memory): MCP operations get tenant_id="unknown" and billing is skipped entirely. """ from hindsight_api.api.mcp import ( + _current_api_key, + _current_api_key_id, + _current_bank_id, + _current_tenant_id, create_mcp_server, - _current_bank_id, _current_api_key, - _current_tenant_id, _current_api_key_id, ) mcp_server = create_mcp_server(mock_memory) @@ -210,20 +216,28 @@ async def test_mcp_tools_propagate_tenant_id_and_api_key_id(mock_memory): def test_multi_bank_mode_exposes_all_tools(mock_memory): - """Test that multi-bank mode exposes all tools including bank management.""" + """Test that multi-bank mode exposes all tools including bank management and mental models.""" from hindsight_api.api.mcp import create_mcp_server # Create server in multi-bank mode (default) mcp_server = create_mcp_server(mock_memory, multi_bank=True) tools = mcp_server._tool_manager._tools - # Should have all tools + # Core tools assert "retain" in tools assert "recall" in tools assert "reflect" in tools assert "list_banks" in tools assert "create_bank" in tools + # Mental model tools + assert "list_mental_models" in tools + assert "get_mental_model" in tools + assert "create_mental_model" in tools + assert "update_mental_model" in tools + assert "delete_mental_model" in tools + assert "refresh_mental_model" in tools + def test_single_bank_mode_excludes_bank_management_tools(mock_memory): """Test that single-bank mode only exposes bank-scoped tools.""" @@ -233,11 +247,19 @@ def test_single_bank_mode_excludes_bank_management_tools(mock_memory): mcp_server = create_mcp_server(mock_memory, multi_bank=False) tools = mcp_server._tool_manager._tools - # Should only have bank-scoped tools + # Should have bank-scoped tools assert "retain" in tools assert "recall" in tools assert "reflect" in tools + # Mental model tools should also be present (they're bank-scoped) + assert "list_mental_models" in tools + assert "get_mental_model" in tools + assert "create_mental_model" in tools + assert "update_mental_model" in tools + assert "delete_mental_model" in tools + assert "refresh_mental_model" in tools + # Should NOT have bank management tools assert "list_banks" not in tools assert "create_bank" not in tools @@ -245,46 +267,56 @@ def test_single_bank_mode_excludes_bank_management_tools(mock_memory): def test_multi_bank_mode_tools_have_bank_id_param(mock_memory): """Test that multi-bank mode tools include bank_id parameter.""" - from hindsight_api.api.mcp import create_mcp_server import inspect + from hindsight_api.api.mcp import create_mcp_server + mcp_server = create_mcp_server(mock_memory, multi_bank=True) tools = mcp_server._tool_manager._tools - # Check that tools have bank_id parameter - retain_tool = tools["retain"] - retain_sig = inspect.signature(retain_tool.fn) - assert "bank_id" in retain_sig.parameters - - recall_tool = tools["recall"] - recall_sig = inspect.signature(recall_tool.fn) - assert "bank_id" in recall_sig.parameters - - reflect_tool = tools["reflect"] - reflect_sig = inspect.signature(reflect_tool.fn) - assert "bank_id" in reflect_sig.parameters + # All bank-scoped tools should have bank_id parameter in multi-bank mode + bank_scoped_tools = [ + "retain", + "recall", + "reflect", + "list_mental_models", + "get_mental_model", + "create_mental_model", + "update_mental_model", + "delete_mental_model", + "refresh_mental_model", + ] + for tool_name in bank_scoped_tools: + tool = tools[tool_name] + sig = inspect.signature(tool.fn) + assert "bank_id" in sig.parameters, f"{tool_name} should have bank_id param in multi-bank mode" def test_single_bank_mode_tools_no_bank_id_param(mock_memory): """Test that single-bank mode tools do NOT include bank_id parameter.""" - from hindsight_api.api.mcp import create_mcp_server import inspect + from hindsight_api.api.mcp import create_mcp_server + mcp_server = create_mcp_server(mock_memory, multi_bank=False) tools = mcp_server._tool_manager._tools - # Check that tools do NOT have bank_id parameter - retain_tool = tools["retain"] - retain_sig = inspect.signature(retain_tool.fn) - assert "bank_id" not in retain_sig.parameters - - recall_tool = tools["recall"] - recall_sig = inspect.signature(recall_tool.fn) - assert "bank_id" not in recall_sig.parameters - - reflect_tool = tools["reflect"] - reflect_sig = inspect.signature(reflect_tool.fn) - assert "bank_id" not in reflect_sig.parameters + # No bank-scoped tool should have bank_id parameter in single-bank mode + bank_scoped_tools = [ + "retain", + "recall", + "reflect", + "list_mental_models", + "get_mental_model", + "create_mental_model", + "update_mental_model", + "delete_mental_model", + "refresh_mental_model", + ] + for tool_name in bank_scoped_tools: + tool = tools[tool_name] + sig = inspect.signature(tool.fn) + assert "bank_id" not in sig.parameters, f"{tool_name} should NOT have bank_id param in single-bank mode" @pytest.mark.asyncio @@ -308,10 +340,14 @@ async def test_middleware_handles_both_endpoints(mock_memory): assert "recall" in multi_bank_tools assert "list_banks" in multi_bank_tools assert "create_bank" in multi_bank_tools + assert "list_mental_models" in multi_bank_tools + assert "create_mental_model" in multi_bank_tools # Single-bank should only have scoped tools assert "retain" in single_bank_tools assert "recall" in single_bank_tools + assert "list_mental_models" in single_bank_tools + assert "create_mental_model" in single_bank_tools assert "list_banks" not in single_bank_tools assert "create_bank" not in single_bank_tools @@ -319,9 +355,10 @@ async def test_middleware_handles_both_endpoints(mock_memory): @pytest.mark.asyncio async def test_routing_logic_from_url_path(): """Test that routing correctly selects server based on URL structure.""" - from hindsight_api.api.mcp import MCPMiddleware from unittest.mock import AsyncMock + from hindsight_api.api.mcp import MCPMiddleware + # Mock memory mock_memory = MagicMock() diff --git a/hindsight-api/tests/test_mcp_tools.py b/hindsight-api/tests/test_mcp_tools.py index 12c15b1f..b3e65c9a 100644 --- a/hindsight-api/tests/test_mcp_tools.py +++ b/hindsight-api/tests/test_mcp_tools.py @@ -1,10 +1,11 @@ """Tests for the shared MCP tools module.""" from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock import pytest -from hindsight_api.mcp_tools import build_content_dict, parse_timestamp +from hindsight_api.mcp_tools import MCPToolsConfig, build_content_dict, parse_timestamp, register_mcp_tools class TestParseTimestamp: @@ -61,3 +62,487 @@ class TestBuildContentDict: result, error = build_content_dict("test content", "test_context", None) assert error is None assert "event_date" not in result + + +# ========================================================================= +# Mental Model MCP Tool Tests +# ========================================================================= + + +@pytest.fixture +def mock_memory(): + """Create a mock MemoryEngine with mental model methods.""" + memory = MagicMock() + memory.list_mental_models = AsyncMock( + return_value=[ + {"id": "mm-1", "name": "Coding Prefs", "source_query": "coding preferences?", "content": "Prefers Python"}, + {"id": "mm-2", "name": "Goals", "source_query": "current goals?", "content": "Ship v2"}, + ] + ) + memory.get_mental_model = AsyncMock( + return_value={ + "id": "mm-1", + "name": "Coding Prefs", + "source_query": "coding preferences?", + "content": "Prefers Python", + } + ) + memory.create_mental_model = AsyncMock(return_value={"id": "mm-new"}) + memory.submit_async_refresh_mental_model = AsyncMock(return_value={"operation_id": "op-123"}) + memory.update_mental_model = AsyncMock( + return_value={ + "id": "mm-1", + "name": "Updated Name", + "source_query": "new query?", + "content": "Updated", + } + ) + memory.delete_mental_model = AsyncMock(return_value=True) + return memory + + +@pytest.fixture +def mcp_server_with_mental_models(mock_memory): + """Create a FastMCP server with mental model tools registered (multi-bank mode).""" + from fastmcp import FastMCP + + mcp = FastMCP("test", stateless_http=True) + config = MCPToolsConfig( + bank_id_resolver=lambda: "test-bank", + include_bank_id_param=True, + tools={ + "list_mental_models", + "get_mental_model", + "create_mental_model", + "update_mental_model", + "delete_mental_model", + "refresh_mental_model", + }, + ) + register_mcp_tools(mcp, mock_memory, config) + return mcp + + +@pytest.fixture +def mcp_server_single_bank(mock_memory): + """Create a FastMCP server with mental model tools registered (single-bank mode).""" + from fastmcp import FastMCP + + mcp = FastMCP("test") + config = MCPToolsConfig( + bank_id_resolver=lambda: "fixed-bank", + include_bank_id_param=False, + tools={ + "list_mental_models", + "get_mental_model", + "create_mental_model", + "update_mental_model", + "delete_mental_model", + "refresh_mental_model", + }, + ) + register_mcp_tools(mcp, mock_memory, config) + return mcp + + +class TestMentalModelToolRegistration: + """Test that mental model tools are registered correctly.""" + + def test_tools_registered_multi_bank(self, mcp_server_with_mental_models): + tools = mcp_server_with_mental_models._tool_manager._tools + expected = { + "list_mental_models", + "get_mental_model", + "create_mental_model", + "update_mental_model", + "delete_mental_model", + "refresh_mental_model", + } + assert expected == set(tools.keys()) + + def test_tools_registered_single_bank(self, mcp_server_single_bank): + tools = mcp_server_single_bank._tool_manager._tools + expected = { + "list_mental_models", + "get_mental_model", + "create_mental_model", + "update_mental_model", + "delete_mental_model", + "refresh_mental_model", + } + assert expected == set(tools.keys()) + + @pytest.mark.asyncio + async def test_list_mental_models_propagates_request_context(self, mock_memory): + from fastmcp import FastMCP + + mcp = FastMCP("test", stateless_http=True) + config = MCPToolsConfig( + bank_id_resolver=lambda: "test-bank", + api_key_resolver=lambda: "test-api-key", + include_bank_id_param=True, + tools={"list_mental_models"}, + ) + register_mcp_tools(mcp, mock_memory, config) + await _tools(mcp)["list_mental_models"].fn() + request_context = mock_memory.list_mental_models.call_args.kwargs["request_context"] + assert request_context.api_key == "test-api-key" + + @pytest.mark.asyncio + async def test_create_mental_model_propagates_request_context(self, mock_memory): + from fastmcp import FastMCP + + mcp = FastMCP("test", stateless_http=True) + config = MCPToolsConfig( + bank_id_resolver=lambda: "test-bank", + api_key_resolver=lambda: "test-api-key", + include_bank_id_param=True, + tools={"create_mental_model"}, + ) + register_mcp_tools(mcp, mock_memory, config) + await _tools(mcp)["create_mental_model"].fn(name="Test", source_query="query") + request_context = mock_memory.create_mental_model.call_args.kwargs["request_context"] + assert request_context.api_key == "test-api-key" + + def test_mental_model_tools_in_default_set(self): + """Mental model tools should be in the default tools set when config.tools is None.""" + from fastmcp import FastMCP + + memory = MagicMock() + # Mock all engine methods that tools reference + memory.retain_batch_async = AsyncMock() + memory.submit_async_retain = AsyncMock(return_value={"operation_id": "op"}) + memory.recall_async = AsyncMock(return_value=MagicMock(results=[])) + memory.reflect_async = AsyncMock() + memory.list_banks = AsyncMock(return_value=[]) + memory.get_bank_profile = AsyncMock(return_value={}) + memory.update_bank = AsyncMock() + memory.list_mental_models = AsyncMock(return_value=[]) + memory.get_mental_model = AsyncMock() + memory.create_mental_model = AsyncMock() + memory.submit_async_refresh_mental_model = AsyncMock() + memory.update_mental_model = AsyncMock() + memory.delete_mental_model = AsyncMock() + + mcp = FastMCP("test", stateless_http=True) + config = MCPToolsConfig( + bank_id_resolver=lambda: "bank", + include_bank_id_param=True, + tools=None, # Default - all tools + ) + register_mcp_tools(mcp, memory, config) + tools = mcp._tool_manager._tools + assert "list_mental_models" in tools + assert "create_mental_model" in tools + assert "refresh_mental_model" in tools + + +@pytest.fixture +def no_bank_mcp_server(mock_memory): + """Create a multi-bank MCP server where bank_id_resolver returns None.""" + from fastmcp import FastMCP + + mcp = FastMCP("test", stateless_http=True) + config = MCPToolsConfig( + bank_id_resolver=lambda: None, + include_bank_id_param=True, + tools={ + "list_mental_models", + "get_mental_model", + "create_mental_model", + "update_mental_model", + "delete_mental_model", + "refresh_mental_model", + }, + ) + register_mcp_tools(mcp, mock_memory, config) + return mcp + + +def _tools(mcp_server): + """Helper to get tools dict from MCP server.""" + return mcp_server._tool_manager._tools + + +@pytest.mark.asyncio +class TestListMentalModels: + async def test_list_multi_bank(self, mcp_server_with_mental_models, mock_memory): + result = await _tools(mcp_server_with_mental_models)["list_mental_models"].fn() + assert '"mm-1"' in result + assert '"mm-2"' in result + mock_memory.list_mental_models.assert_called_once() + assert mock_memory.list_mental_models.call_args.kwargs["bank_id"] == "test-bank" + + async def test_list_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory): + """Explicit bank_id should override the resolver.""" + await _tools(mcp_server_with_mental_models)["list_mental_models"].fn(bank_id="other-bank") + assert mock_memory.list_mental_models.call_args.kwargs["bank_id"] == "other-bank" + + async def test_list_with_tags(self, mcp_server_with_mental_models, mock_memory): + await _tools(mcp_server_with_mental_models)["list_mental_models"].fn(tags=["work"]) + assert mock_memory.list_mental_models.call_args.kwargs["tags"] == ["work"] + + async def test_list_single_bank(self, mcp_server_single_bank, mock_memory): + result = await _tools(mcp_server_single_bank)["list_mental_models"].fn() + assert isinstance(result, dict) + assert len(result["items"]) == 2 + assert mock_memory.list_mental_models.call_args.kwargs["bank_id"] == "fixed-bank" + + async def test_list_no_bank_returns_error(self, no_bank_mcp_server): + result = await _tools(no_bank_mcp_server)["list_mental_models"].fn() + assert "error" in result + + async def test_list_engine_error_multi_bank(self, mcp_server_with_mental_models, mock_memory): + mock_memory.list_mental_models.side_effect = RuntimeError("DB connection lost") + result = await _tools(mcp_server_with_mental_models)["list_mental_models"].fn() + assert "error" in result + assert "DB connection lost" in result + + async def test_list_engine_error_single_bank(self, mcp_server_single_bank, mock_memory): + mock_memory.list_mental_models.side_effect = RuntimeError("DB connection lost") + result = await _tools(mcp_server_single_bank)["list_mental_models"].fn() + assert isinstance(result, dict) + assert "error" in result + + +@pytest.mark.asyncio +class TestGetMentalModel: + async def test_get_multi_bank(self, mcp_server_with_mental_models, mock_memory): + result = await _tools(mcp_server_with_mental_models)["get_mental_model"].fn(mental_model_id="mm-1") + assert '"mm-1"' in result + assert mock_memory.get_mental_model.call_args.kwargs["mental_model_id"] == "mm-1" + + async def test_get_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory): + await _tools(mcp_server_with_mental_models)["get_mental_model"].fn(mental_model_id="mm-1", bank_id="other-bank") + assert mock_memory.get_mental_model.call_args.kwargs["bank_id"] == "other-bank" + + async def test_get_not_found_multi_bank(self, mcp_server_with_mental_models, mock_memory): + mock_memory.get_mental_model.return_value = None + result = await _tools(mcp_server_with_mental_models)["get_mental_model"].fn(mental_model_id="missing") + assert "not found" in result + + async def test_get_not_found_single_bank(self, mcp_server_single_bank, mock_memory): + mock_memory.get_mental_model.return_value = None + result = await _tools(mcp_server_single_bank)["get_mental_model"].fn(mental_model_id="missing") + assert isinstance(result, dict) + assert "not found" in result["error"] + + async def test_get_single_bank(self, mcp_server_single_bank, mock_memory): + result = await _tools(mcp_server_single_bank)["get_mental_model"].fn(mental_model_id="mm-1") + assert isinstance(result, dict) + assert result["id"] == "mm-1" + + async def test_get_no_bank_returns_error(self, no_bank_mcp_server): + result = await _tools(no_bank_mcp_server)["get_mental_model"].fn(mental_model_id="mm-1") + assert "error" in result + + async def test_get_engine_error(self, mcp_server_with_mental_models, mock_memory): + mock_memory.get_mental_model.side_effect = RuntimeError("DB error") + result = await _tools(mcp_server_with_mental_models)["get_mental_model"].fn(mental_model_id="mm-1") + assert "error" in result + + +@pytest.mark.asyncio +class TestCreateMentalModel: + async def test_create_multi_bank(self, mcp_server_with_mental_models, mock_memory): + result = await _tools(mcp_server_with_mental_models)["create_mental_model"].fn( + name="Test Model", + source_query="What are the user's preferences?", + ) + assert '"mm-new"' in result + assert '"op-123"' in result + mock_memory.create_mental_model.assert_called_once() + call_kwargs = mock_memory.create_mental_model.call_args.kwargs + assert call_kwargs["name"] == "Test Model" + assert call_kwargs["source_query"] == "What are the user's preferences?" + assert call_kwargs["content"] == "Generating content..." + # Verify async refresh was scheduled + mock_memory.submit_async_refresh_mental_model.assert_called_once() + assert mock_memory.submit_async_refresh_mental_model.call_args.kwargs["mental_model_id"] == "mm-new" + + async def test_create_with_custom_id(self, mcp_server_with_mental_models, mock_memory): + await _tools(mcp_server_with_mental_models)["create_mental_model"].fn( + name="Test", source_query="query", mental_model_id="custom-id" + ) + assert mock_memory.create_mental_model.call_args.kwargs["mental_model_id"] == "custom-id" + + async def test_create_with_tags_and_max_tokens(self, mcp_server_with_mental_models, mock_memory): + await _tools(mcp_server_with_mental_models)["create_mental_model"].fn( + name="Test", source_query="query", tags=["work", "coding"], max_tokens=4096 + ) + call_kwargs = mock_memory.create_mental_model.call_args.kwargs + assert call_kwargs["tags"] == ["work", "coding"] + assert call_kwargs["max_tokens"] == 4096 + + async def test_create_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory): + await _tools(mcp_server_with_mental_models)["create_mental_model"].fn( + name="Test", source_query="query", bank_id="other-bank" + ) + assert mock_memory.create_mental_model.call_args.kwargs["bank_id"] == "other-bank" + assert mock_memory.submit_async_refresh_mental_model.call_args.kwargs["bank_id"] == "other-bank" + + async def test_create_single_bank(self, mcp_server_single_bank, mock_memory): + result = await _tools(mcp_server_single_bank)["create_mental_model"].fn(name="Test", source_query="query") + assert isinstance(result, dict) + assert result["mental_model_id"] == "mm-new" + assert result["operation_id"] == "op-123" + + async def test_create_no_bank_returns_error(self, no_bank_mcp_server): + result = await _tools(no_bank_mcp_server)["create_mental_model"].fn(name="Test", source_query="query") + assert "error" in result + + async def test_create_value_error_multi_bank(self, mcp_server_with_mental_models, mock_memory): + """ValueError from engine (e.g. invalid ID format) should return error, not crash.""" + mock_memory.create_mental_model.side_effect = ValueError("ID must be alphanumeric lowercase") + result = await _tools(mcp_server_with_mental_models)["create_mental_model"].fn( + name="Test", source_query="query", mental_model_id="INVALID!!" + ) + assert "alphanumeric" in result + + async def test_create_value_error_single_bank(self, mcp_server_single_bank, mock_memory): + mock_memory.create_mental_model.side_effect = ValueError("ID must be alphanumeric lowercase") + result = await _tools(mcp_server_single_bank)["create_mental_model"].fn( + name="Test", source_query="query", mental_model_id="INVALID!!" + ) + assert isinstance(result, dict) + assert "alphanumeric" in result["error"] + + async def test_create_engine_error(self, mcp_server_with_mental_models, mock_memory): + mock_memory.create_mental_model.side_effect = RuntimeError("DB error") + result = await _tools(mcp_server_with_mental_models)["create_mental_model"].fn( + name="Test", source_query="query" + ) + assert "error" in result + + +@pytest.mark.asyncio +class TestUpdateMentalModel: + async def test_update_multi_bank(self, mcp_server_with_mental_models, mock_memory): + result = await _tools(mcp_server_with_mental_models)["update_mental_model"].fn( + mental_model_id="mm-1", name="Updated Name" + ) + assert '"Updated Name"' in result + call_kwargs = mock_memory.update_mental_model.call_args.kwargs + assert call_kwargs["name"] == "Updated Name" + assert call_kwargs["source_query"] is None # Not updated + + async def test_update_multiple_fields(self, mcp_server_with_mental_models, mock_memory): + await _tools(mcp_server_with_mental_models)["update_mental_model"].fn( + mental_model_id="mm-1", name="New Name", source_query="new query?", tags=["updated"], max_tokens=4096 + ) + call_kwargs = mock_memory.update_mental_model.call_args.kwargs + assert call_kwargs["name"] == "New Name" + assert call_kwargs["source_query"] == "new query?" + assert call_kwargs["tags"] == ["updated"] + assert call_kwargs["max_tokens"] == 4096 + + async def test_update_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory): + await _tools(mcp_server_with_mental_models)["update_mental_model"].fn( + mental_model_id="mm-1", name="X", bank_id="other-bank" + ) + assert mock_memory.update_mental_model.call_args.kwargs["bank_id"] == "other-bank" + + async def test_update_not_found_multi_bank(self, mcp_server_with_mental_models, mock_memory): + mock_memory.update_mental_model.return_value = None + result = await _tools(mcp_server_with_mental_models)["update_mental_model"].fn( + mental_model_id="missing", name="X" + ) + assert "not found" in result + + async def test_update_single_bank(self, mcp_server_single_bank, mock_memory): + result = await _tools(mcp_server_single_bank)["update_mental_model"].fn(mental_model_id="mm-1", name="Updated") + assert isinstance(result, dict) + assert mock_memory.update_mental_model.call_args.kwargs["bank_id"] == "fixed-bank" + + async def test_update_not_found_single_bank(self, mcp_server_single_bank, mock_memory): + mock_memory.update_mental_model.return_value = None + result = await _tools(mcp_server_single_bank)["update_mental_model"].fn(mental_model_id="missing", name="X") + assert isinstance(result, dict) + assert "not found" in result["error"] + + async def test_update_no_bank_returns_error(self, no_bank_mcp_server): + result = await _tools(no_bank_mcp_server)["update_mental_model"].fn(mental_model_id="mm-1", name="X") + assert "error" in result + + async def test_update_engine_error(self, mcp_server_with_mental_models, mock_memory): + mock_memory.update_mental_model.side_effect = RuntimeError("DB error") + result = await _tools(mcp_server_with_mental_models)["update_mental_model"].fn(mental_model_id="mm-1", name="X") + assert "error" in result + + +@pytest.mark.asyncio +class TestDeleteMentalModel: + async def test_delete_multi_bank(self, mcp_server_with_mental_models, mock_memory): + result = await _tools(mcp_server_with_mental_models)["delete_mental_model"].fn(mental_model_id="mm-1") + assert '"deleted"' in result + assert mock_memory.delete_mental_model.call_args.kwargs["mental_model_id"] == "mm-1" + + async def test_delete_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory): + await _tools(mcp_server_with_mental_models)["delete_mental_model"].fn( + mental_model_id="mm-1", bank_id="other-bank" + ) + assert mock_memory.delete_mental_model.call_args.kwargs["bank_id"] == "other-bank" + + async def test_delete_not_found_multi_bank(self, mcp_server_with_mental_models, mock_memory): + mock_memory.delete_mental_model.return_value = False + result = await _tools(mcp_server_with_mental_models)["delete_mental_model"].fn(mental_model_id="missing") + assert "not found" in result + + async def test_delete_not_found_single_bank(self, mcp_server_single_bank, mock_memory): + mock_memory.delete_mental_model.return_value = False + result = await _tools(mcp_server_single_bank)["delete_mental_model"].fn(mental_model_id="missing") + assert isinstance(result, dict) + assert "not found" in result["error"] + + async def test_delete_single_bank(self, mcp_server_single_bank, mock_memory): + result = await _tools(mcp_server_single_bank)["delete_mental_model"].fn(mental_model_id="mm-1") + assert isinstance(result, dict) + assert result["status"] == "deleted" + + async def test_delete_no_bank_returns_error(self, no_bank_mcp_server): + result = await _tools(no_bank_mcp_server)["delete_mental_model"].fn(mental_model_id="mm-1") + assert "error" in result + + async def test_delete_engine_error(self, mcp_server_with_mental_models, mock_memory): + mock_memory.delete_mental_model.side_effect = RuntimeError("DB error") + result = await _tools(mcp_server_with_mental_models)["delete_mental_model"].fn(mental_model_id="mm-1") + assert "error" in result + + +@pytest.mark.asyncio +class TestRefreshMentalModel: + async def test_refresh_multi_bank(self, mcp_server_with_mental_models, mock_memory): + result = await _tools(mcp_server_with_mental_models)["refresh_mental_model"].fn(mental_model_id="mm-1") + assert '"op-123"' in result + assert '"queued"' in result + + async def test_refresh_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory): + await _tools(mcp_server_with_mental_models)["refresh_mental_model"].fn( + mental_model_id="mm-1", bank_id="other-bank" + ) + assert mock_memory.submit_async_refresh_mental_model.call_args.kwargs["bank_id"] == "other-bank" + + async def test_refresh_not_found_multi_bank(self, mcp_server_with_mental_models, mock_memory): + mock_memory.submit_async_refresh_mental_model.side_effect = ValueError("Mental model 'missing' not found") + result = await _tools(mcp_server_with_mental_models)["refresh_mental_model"].fn(mental_model_id="missing") + assert "not found" in result + + async def test_refresh_not_found_single_bank(self, mcp_server_single_bank, mock_memory): + mock_memory.submit_async_refresh_mental_model.side_effect = ValueError("not found") + result = await _tools(mcp_server_single_bank)["refresh_mental_model"].fn(mental_model_id="missing") + assert isinstance(result, dict) + assert "not found" in result["error"] + + async def test_refresh_single_bank(self, mcp_server_single_bank, mock_memory): + result = await _tools(mcp_server_single_bank)["refresh_mental_model"].fn(mental_model_id="mm-1") + assert isinstance(result, dict) + assert result["operation_id"] == "op-123" + + async def test_refresh_no_bank_returns_error(self, no_bank_mcp_server): + result = await _tools(no_bank_mcp_server)["refresh_mental_model"].fn(mental_model_id="mm-1") + assert "error" in result + + async def test_refresh_engine_error(self, mcp_server_with_mental_models, mock_memory): + mock_memory.submit_async_refresh_mental_model.side_effect = RuntimeError("DB error") + result = await _tools(mcp_server_with_mental_models)["refresh_mental_model"].fn(mental_model_id="mm-1") + assert "error" in result