"""Tool handlers for the oGMemory MCP adapter."""
from __future__ import annotations
import json
from typing import Any
from core.models import RetrievedBlock
from server.mcp_resources import (
MEMORY_TYPES,
external_id_from_uri,
external_to_internal_uri_candidates,
internal_to_external_uri,
parse_ogmemory_uri,
public_type_from_internal,
render_resource,
)
from server.mcp_security import (
UNTRUSTED_DATA_NOTICE,
build_mcp_context,
is_not_found_or_not_visible_error,
raise_not_found_or_not_visible,
)
_FORBIDDEN_WRITE_FIELDS = {
"type",
"memory_type",
"memory_id",
"uri",
"abstract",
"overview",
}
def query_memory(
service: Any,
*,
query: str,
types: list[str] | None = None,
top_k: int = 5,
score_threshold: float | None = None,
session_id: str | None = None,
agent_id: str | None = None,
transport_headers: dict[str, Any] | None = None,
allow_stdio_env_auth: bool = True,
resolved_identity: Any | None = None,
) -> dict[str, Any]:
if not (query or "").strip():
raise ValueError("query must not be empty")
if top_k < 1 or top_k > 20:
raise ValueError("top_k must be between 1 and 20")
requested_types = types or []
unknown = sorted(t for t in requested_types if t not in MEMORY_TYPES)
if unknown:
raise ValueError(f"Unknown memory types: {', '.join(unknown)}")
_validate_request_size(
service,
{
"query": query,
"types": requested_types,
"top_k": top_k,
"score_threshold": score_threshold,
"session_id": session_id,
"agent_id": agent_id,
},
)
arguments = {"session_id": session_id, "agent_id": agent_id}
ctx = build_mcp_context(
service,
arguments,
transport_headers=transport_headers,
allow_stdio_env_auth=allow_stdio_env_auth,
resolved_identity=resolved_identity,
)
categories = [MEMORY_TYPES[t].internal_category for t in requested_types] or None
try:
result = service.get_read_api().search_memory(
query.strip(),
ctx,
top_k=top_k,
categories=categories,
score_threshold=score_threshold,
include_debug=False,
fill_content_for_top_k=0,
)
except Exception as exc:
raise_not_found_or_not_visible(exc)
raise
hits = []
for hit in result.hits:
try:
hits.append(_hit_to_public(hit))
except ValueError:
continue
payload = {
"request_id": result.request_id,
"query": result.query or query,
"hits": hits,
"untrusted_data_notice": UNTRUSTED_DATA_NOTICE,
}
return _truncate_query_payload(payload, _configured_max_response_chars(service))
def read_memory(
service: Any,
*,
uri: str,
level: str = "all",
max_chars: int = 12000,
session_id: str | None = None,
agent_id: str | None = None,
transport_headers: dict[str, Any] | None = None,
allow_stdio_env_auth: bool = True,
resolved_identity: Any | None = None,
) -> dict[str, Any]:
if level not in {"abstract", "overview", "content", "all"}:
raise ValueError("level must be one of: abstract, overview, content, all")
if max_chars < 1 or max_chars > 50000:
raise ValueError("max_chars must be between 1 and 50000")
_validate_request_size(
service,
{
"uri": uri,
"level": level,
"max_chars": max_chars,
"session_id": session_id,
"agent_id": agent_id,
},
)
max_chars = min(max_chars, _configured_max_response_chars(service))
arguments = {"session_id": session_id, "agent_id": agent_id}
ctx = build_mcp_context(
service,
arguments,
transport_headers=transport_headers,
allow_stdio_env_auth=allow_stdio_env_auth,
resolved_identity=resolved_identity,
)
parsed = parse_ogmemory_uri(uri)
if parsed.kind != "item":
payload = render_resource(service, uri, ctx)
return _truncate_payload(payload, max_chars)
try:
context_fs = service._get_context_fs()
last_exc: Exception | None = None
node = None
for internal_uri in external_to_internal_uri_candidates(uri, ctx):
try:
if hasattr(context_fs, "read_node_fields"):
node = context_fs.read_node_fields(internal_uri, ctx, level)
else:
node = context_fs.read_node(internal_uri, ctx)
break
except Exception as exc:
if not is_not_found_or_not_visible_error(exc):
raise
last_exc = exc
if node is None:
if last_exc is not None:
raise last_exc
raise ValueError(f"URI does not reference a single memory: {uri}")
canonical_uri = internal_to_external_uri(node.uri, category=node.category)
canonical_parsed = parse_ogmemory_uri(canonical_uri)
payload = {
"uri": canonical_uri,
"type": canonical_parsed.memory_type,
"id": canonical_parsed.memory_id,
"untrusted_data_notice": UNTRUSTED_DATA_NOTICE,
}
if level in {"abstract", "all"}:
payload["abstract"] = node.abstract
if level in {"overview", "all"}:
payload["overview"] = node.overview
if level in {"content", "all"}:
payload["content"] = node.content
return _truncate_payload(payload, max_chars)
except Exception as exc:
raise_not_found_or_not_visible(exc)
raise
def delete_memory(
service: Any,
*,
uri: str,
reason: str | None = None,
session_id: str | None = None,
agent_id: str | None = None,
client_request_id: str | None = None,
transport_headers: dict[str, Any] | None = None,
allow_stdio_env_auth: bool = True,
resolved_identity: Any | None = None,
) -> dict[str, Any]:
parsed = parse_ogmemory_uri(uri)
if parsed.kind != "item":
raise ValueError("delete_memory requires a single ogmemory://memories/{type}/{id} URI")
_validate_request_size(
service,
{
"uri": uri,
"reason": reason,
"session_id": session_id,
"agent_id": agent_id,
"client_request_id": client_request_id,
},
)
arguments = {"session_id": session_id, "agent_id": agent_id}
ctx = build_mcp_context(
service,
arguments,
transport_headers=transport_headers,
allow_stdio_env_auth=allow_stdio_env_auth,
resolved_identity=resolved_identity,
)
try:
last_exc: Exception | None = None
result = None
for internal_uri in external_to_internal_uri_candidates(uri, ctx):
try:
result = service.forget_memory(
internal_uri,
ctx,
external_uri=uri,
reason=reason,
client_request_id=client_request_id,
)
break
except Exception as exc:
if not is_not_found_or_not_visible_error(exc):
raise
last_exc = exc
if result is None:
if last_exc is not None:
raise last_exc
raise ValueError(f"URI does not reference a single memory: {uri}")
except Exception as exc:
raise_not_found_or_not_visible(exc)
raise
return {
"ok": bool(result.get("ok", True)),
"deleted": bool(result.get("deleted", True)),
"mode": result.get("mode", "archive"),
"uri": uri,
"type": parsed.memory_type,
"id": parsed.memory_id,
"client_request_id": client_request_id,
}
def _hit_to_public(hit: RetrievedBlock) -> dict[str, Any]:
external_uri = internal_to_external_uri(hit.uri, category=hit.category)
parsed = parse_ogmemory_uri(external_uri)
return {
"uri": external_uri,
"type": parsed.memory_type or public_type_from_internal(hit.category, hit.uri),
"id": external_id_from_uri(external_uri),
"score": hit.score,
"abstract": hit.abstract,
"has_content": hit.has_content or bool(hit.content_excerpt),
}
def _truncate_payload(payload: dict[str, Any], max_chars: int) -> dict[str, Any]:
remaining = max_chars
truncated = bool(payload.get("truncated", False))
result: dict[str, Any] = {}
for key, value in payload.items():
if isinstance(value, str) and key in {"abstract", "overview", "content"}:
if len(value) > remaining:
result[key] = value[: max(0, remaining)]
truncated = True
remaining = 0
else:
result[key] = value
remaining -= len(value)
else:
result[key] = value
result["truncated"] = truncated
return result
def _truncate_query_payload(payload: dict[str, Any], max_chars: int) -> dict[str, Any]:
remaining = max_chars
truncated = False
result = dict(payload)
truncated_hits: list[dict[str, Any]] = []
for hit in payload.get("hits", []):
if remaining <= 0:
truncated = True
break
row = dict(hit)
abstract = row.get("abstract")
if isinstance(abstract, str):
if len(abstract) > remaining:
row["abstract"] = abstract[: max(0, remaining)]
truncated = True
remaining = 0
else:
remaining -= len(abstract)
truncated_hits.append(row)
result["hits"] = truncated_hits
result["truncated"] = truncated
return result
def _configured_max_request_bytes(service: Any) -> int:
value = int(getattr(service._cfg, "mcp_max_request_bytes", 1048576) or 1048576)
return max(1, value)
def _configured_max_response_chars(service: Any) -> int:
value = int(getattr(service._cfg, "mcp_max_response_chars", 50000) or 50000)
return max(1, min(value, 50000))
def _validate_request_size(service: Any, payload: dict[str, Any]) -> None:
max_bytes = _configured_max_request_bytes(service)
encoded = json.dumps(payload, ensure_ascii=False, default=str, sort_keys=True).encode("utf-8")
if len(encoded) > max_bytes:
raise ValueError(f"MCP request exceeds mcp_max_request_bytes ({max_bytes})")
def write_memory(
service: Any,
*,
text: str | None = None,
messages: list[dict[str, Any]] | None = None,
session_id: str | None = None,
agent_id: str | None = None,
client_request_id: str | None = None,
wait: bool = False,
transport_headers: dict[str, Any] | None = None,
allow_stdio_env_auth: bool = True,
resolved_identity: Any | None = None,
) -> dict[str, Any]:
if session_id is None or not session_id.strip():
raise ValueError("session_id is required for write_memory")
normalized_messages = _normalize_write_messages(text=text, messages=messages)
_validate_request_size(
service,
{
"messages": normalized_messages,
"session_id": session_id,
"agent_id": agent_id,
"client_request_id": client_request_id,
"wait": wait,
},
)
arguments = {"session_id": session_id, "agent_id": agent_id}
ctx = build_mcp_context(
service,
arguments,
transport_headers=transport_headers,
allow_stdio_env_auth=allow_stdio_env_auth,
resolved_identity=resolved_identity,
)
params: dict[str, Any] = {
"sessionId": session_id,
"messages": normalized_messages,
"_ctx": ctx,
}
if agent_id:
params["agentId"] = agent_id
if client_request_id:
params["clientRequestId"] = client_request_id
params["client_request_id"] = client_request_id
if wait:
params["wait"] = True
try:
result = service.after_turn(params)
except Exception as exc:
_record_mcp_write_audit(
service,
ctx,
session_id=session_id,
agent_id=agent_id,
client_request_id=client_request_id,
result="failed",
details={"error": f"{type(exc).__name__}: {exc}"},
)
raise
_record_mcp_write_audit(
service,
ctx,
session_id=session_id,
agent_id=agent_id,
client_request_id=client_request_id,
result="success" if result.get("ok", True) else "failed",
details={
"status": result.get("status", "accepted"),
"accepted_messages": len(normalized_messages),
"wait_requested": wait,
},
)
return {
"ok": bool(result.get("ok", True)),
"status": result.get("status", "accepted"),
"session_id": session_id,
"agent_id": agent_id or ctx.agent_id,
"client_request_id": client_request_id,
"accepted_messages": len(normalized_messages),
"stats": _public_write_stats(result),
}
def _normalize_write_messages(
*,
text: str | None,
messages: list[dict[str, Any]] | None,
) -> list[dict[str, Any]]:
if text and messages:
raise ValueError("Provide either text or messages, not both")
if text is not None:
if not text.strip():
raise ValueError("text must not be empty")
return [{"role": "user", "content": text}]
if not messages:
raise ValueError("messages or text is required")
normalized: list[dict[str, Any]] = []
for index, message in enumerate(messages):
if not isinstance(message, dict):
raise ValueError(f"messages[{index}] must be an object")
forbidden = sorted(k for k in message if k in _FORBIDDEN_WRITE_FIELDS)
if forbidden:
raise ValueError(
f"messages[{index}] cannot set atomic memory fields: {', '.join(forbidden)}"
)
role = str(message.get("role") or "").strip()
content = message.get("content")
if role not in {"system", "user", "assistant", "tool"}:
raise ValueError(f"messages[{index}].role must be system, user, assistant, or tool")
if not isinstance(content, str) or not content.strip():
raise ValueError(f"messages[{index}].content must be a non-empty string")
normalized_message: dict[str, Any] = {"role": role, "content": content}
if message.get("id"):
normalized_message["id"] = str(message["id"])
if message.get("created_at"):
normalized_message["created_at"] = str(message["created_at"])
normalized.append(normalized_message)
return normalized
def _record_mcp_write_audit(
service: Any,
ctx: Any,
*,
session_id: str,
agent_id: str | None,
client_request_id: str | None,
result: str,
details: dict[str, Any],
) -> None:
get_audit_service = getattr(service, "get_audit_service", None)
if not callable(get_audit_service):
return
try:
get_audit_service().record(
ctx.account_id,
actor=ctx.user_id,
target=session_id,
action="mcp.write_memory",
result=result,
trace_id=ctx.trace_id,
details={
"role": str(ctx.role),
"agent_id": agent_id or ctx.agent_id,
"session_id": session_id,
"client_request_id": client_request_id,
**details,
},
)
except Exception:
return
def _public_write_stats(result: dict[str, Any]) -> dict[str, Any]:
allowed_fields = {
"pending_tokens",
"buffered_tokens",
"candidates_extracted",
"writes_completed",
"writes_failed",
"task_id",
"reason",
}
stats: dict[str, Any] = {}
for field in allowed_fields:
value = result.get(field)
if isinstance(value, (str, int, float, bool)) or value is None:
stats[field] = value
return stats