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 <noreply@anthropic.com>

* 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 <noreply@anthropic.com>

* fix: update extension test tool count for mental model tools

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* 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 <noreply@anthropic.com>

* 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 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
DK09876 2026-02-10 14:40:43 -07:00 committed by GitHub
parent 90be7c6829
commit f641b30d83
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 1191 additions and 117 deletions

View file

@ -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,

View file

@ -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
)

View file

@ -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:

View file

@ -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)}

View file

@ -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"

View file

@ -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

View file

@ -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()

View file

@ -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