From 6e455ef65e066b39e349f65baadca88ccd3e6758 Mon Sep 17 00:00:00 2001 From: Caleb Evans Date: Tue, 11 Aug 2026 00:38:04 -0600 Subject: [PATCH 1/8] fix: measure content limits in bytes and say so in the error Every length check used str::len() (UTF-8 bytes) while the error message and the advertised MCP JSON schema both said "characters". A summary of 1900 characters written with em dashes is 5696 bytes and was rejected as "Summary exceeds maximum length of 2000 characters", while 1900 ASCII characters was accepted. Verified against a running server. Byte semantics are intentional and kept: the on-disk record encodes the summary length as a u16. What was wrong was the reporting. - Add src/model/validation.rs as the single source of truth. The HTTP create path, the HTTP batch path and both MCP store handlers each carried a hand-rolled copy of these checks, and the copies had drifted. - Errors now name the field, the measured byte count, the character count and the limit, and explain the discrepancy only when one exists: "summary is 5700 bytes (1900 characters), which exceeds the 2000-byte limit. Limits are measured in UTF-8 bytes, not characters: ..." - Memory::validate() already produced correct messages but had zero call sites. It now delegates here instead of being a second source of truth. - Validate the HTTP batch endpoint, which previously checked only that the summary was non-empty. A summary of 65536 bytes or more reached the wrapping `as u16` cast in record.rs and corrupted the stored record. - Fix a latent panic in the CLI: &text[..200] aborts when byte 200 is not a char boundary. Replaced with truncate_on_char_boundary(). - Replace magic literals with named constants and advertise the real limits in the tool schemas (fullText maxLength, array maxItems). - Correct docs/mcp.md and docs/guide.md, which stated the wrong unit. Behavioral changes: entities/topics/emotions counts are now enforced on the HTTP path; an empty-string summary is now rejected over MCP; namespace names are ASCII-only on both paths (previously HTTP accepted Unicode). Removing three public ValidationError variants is semver-breaking, which is acceptable at 0.1.x. Error message text changed on every path; the machine-readable error code and field are unchanged. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01JNhUehmChiQKJUjv3QPnkH --- docs/guide.md | 4 +- docs/mcp.md | 16 +- src/api/errors.rs | 72 +++++++++ src/api/handlers.rs | 84 ++++------- src/cli/client.rs | 14 +- src/mcp/tools.rs | 222 +++++++++++++++++----------- src/model/constants.rs | 41 ++++- src/model/error.rs | 85 +++++++++-- src/model/memory.rs | 32 +--- src/model/mod.rs | 2 + src/model/validation.rs | 320 ++++++++++++++++++++++++++++++++++++++++ 11 files changed, 700 insertions(+), 192 deletions(-) create mode 100644 src/model/validation.rs diff --git a/docs/guide.md b/docs/guide.md index e821218..d1b0335 100644 --- a/docs/guide.md +++ b/docs/guide.md @@ -754,8 +754,8 @@ Store a new memory. | Parameter | Type | Required | Description | |---|---|---|---| -| `summary` | string | Yes | Short description (max 2000 chars). | -| `fullText` | string | No | Detailed content (max 1 MB). Dropped as memory decays to ghost phase. | +| `summary` | string | Yes | Short description (max 2000 bytes of UTF-8; ~2000 ASCII characters, fewer with em dashes/curly quotes/emoji). | +| `fullText` | string | No | Detailed content (max 1 MB = 1 048 576 bytes of UTF-8, counted on the raw text before JSON escaping). Dropped as memory decays to ghost phase. | | `tags` | string[] | No | Categorization tags, e.g. `["topic/rust", "type/observation"]`. Max 64. | | `entities` | string[] | No | Named entities (people, places, orgs). Used for search indexing and graph linking. Max 32. | | `topics` | string[] | No | Topic keywords, e.g. `["rust", "cooking"]`. Max 32. | diff --git a/docs/mcp.md b/docs/mcp.md index f2dd210..bfa1d7b 100644 --- a/docs/mcp.md +++ b/docs/mcp.md @@ -93,6 +93,14 @@ Configure per-project namespace defaults in a `.recalld.toml` file (see above), ## Available tools +### A note on limits + +Every length limit below is measured in **UTF-8 bytes, not characters**. For plain ASCII the two are the same, so a 2000-byte `summary` holds 2000 characters. Non-ASCII characters cost more: accented letters and curly quotes are 2-3 bytes, em dashes (`—`) are 3 bytes, and emoji are 4 bytes. A summary of 1900 em dashes is 5700 bytes and will be rejected. Rejection messages report both numbers, e.g. `summary is 5700 bytes (1900 characters), which exceeds the 2000-byte limit.` + +Array limits (`tags`, `entities`, `topics`, `emotions`) are plain item counts. + +> **Known limitation:** `entities`, `topics`, and `emotions` are converted into `entity/…`, `topic/…`, and `emotion/…` tags *after* validation runs, and the `tags` limit of 64 is checked against the tags you supplied. A memory that supplies 64 tags plus 32 entities, 32 topics, and 32 emotions can therefore end up with up to 160 tags stored. This is accepted today and is not treated as an error. + ### store_memory Store a new observation, fact, or piece of context. The system automatically generates an embedding for semantic search. Memories decay over time unless reinforced. @@ -101,8 +109,8 @@ Store a new observation, fact, or piece of context. The system automatically gen | Parameter | Type | Required | Default | Description | |-----------|------|----------|---------|-------------| -| `summary` | string | yes | -- | Short description (max 2000 chars) | -| `fullText` | string | no | -- | Detailed content. Dropped when memory decays to ghost phase. Max 1 MB. | +| `summary` | string | yes | -- | Short description (max 2000 bytes of UTF-8; ~2000 ASCII characters, fewer with em dashes/curly quotes/emoji) | +| `fullText` | string | no | -- | Detailed content. Dropped when memory decays to ghost phase. Max 1 MB = 1 048 576 bytes of UTF-8, counted on the raw text before JSON escaping. | | `tags` | string[] | no | `[]` | Categorization tags, e.g. `["topic/rust", "type/observation"]`. Max 64. | | `entities` | string[] | no | `[]` | Named entities (people, places, orgs). Used for search indexing and graph linking. Max 32. | | `topics` | string[] | no | `[]` | Topic keywords, e.g. `["rust", "cooking"]`. Max 32. | @@ -152,8 +160,8 @@ Each object in the `memories` array accepts: | Field | Type | Required | Default | Description | |-------|------|----------|---------|-------------| -| `summary` | string | yes | -- | Short description (max 2000 chars) | -| `fullText` | string | no | -- | Detailed content. Max 1 MB. | +| `summary` | string | yes | -- | Short description (max 2000 bytes of UTF-8; ~2000 ASCII characters, fewer with em dashes/curly quotes/emoji) | +| `fullText` | string | no | -- | Detailed content. Max 1 MB = 1 048 576 bytes of UTF-8, counted on the raw text before JSON escaping. | | `tags` | string[] | no | `[]` | Categorization tags. Max 64. | | `entities` | string[] | no | `[]` | Named entities. Max 32. | | `topics` | string[] | no | `[]` | Topic keywords. Max 32. | diff --git a/src/api/errors.rs b/src/api/errors.rs index a599334..5999bf8 100644 --- a/src/api/errors.rs +++ b/src/api/errors.rs @@ -11,6 +11,7 @@ use axum::{ response::{IntoResponse, Response}, }; +use crate::model::error::ValidationError; use crate::serialization::ApiError; // ═══════════════════════════════════════════════════════════════════════ @@ -229,6 +230,29 @@ impl From for AppError { } } +impl From for AppError { + fn from(err: ValidationError) -> Self { + match &err { + // Embedding shape problems are semantically invalid rather + // than malformed — 422, matching the inline check in + // `create_memory`. + ValidationError::DimensionMismatch { .. } + | ValidationError::NonFiniteEmbedding { .. } => AppError::UnprocessableEntity { + message: err.to_string(), + field: err.field().map(Into::into), + }, + ValidationError::NamespaceNotFound { name } => AppError::NotFound { + resource: "namespace", + id: name.clone(), + }, + _ => AppError::BadRequest { + message: err.to_string(), + field: err.field().map(Into::into), + }, + } + } +} + impl From for AppError { fn from(err: GraphError) -> Self { match &err { @@ -245,3 +269,51 @@ impl From for AppError { } } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn validation_error_becomes_bad_request_with_field() { + let source = ValidationError::FieldTooLong { + field: "summary", + bytes: 5700, + chars: 1900, + max_bytes: 2000, + }; + let expected = source.to_string(); + match AppError::from(source) { + AppError::BadRequest { message, field } => { + assert_eq!(field.as_deref(), Some("summary")); + assert_eq!(message, expected); + assert!(message.contains("5700 bytes (1900 characters)")); + } + other => panic!("expected BadRequest, got {other:?}"), + } + } + + #[test] + fn namespace_not_found_becomes_not_found() { + let err = AppError::from(ValidationError::NamespaceNotFound { + name: "missing".to_string(), + }); + match err { + AppError::NotFound { resource, id } => { + assert_eq!(resource, "namespace"); + assert_eq!(id, "missing"); + } + other => panic!("expected NotFound, got {other:?}"), + } + } + + #[test] + fn dimension_mismatch_becomes_unprocessable_entity() { + let err = AppError::from(ValidationError::DimensionMismatch { + expected: 1536, + actual: 768, + namespace: "default".to_string(), + }); + assert!(matches!(err, AppError::UnprocessableEntity { .. })); + } +} diff --git a/src/api/handlers.rs b/src/api/handlers.rs index 1dc0d3a..8862c59 100644 --- a/src/api/handlers.rs +++ b/src/api/handlers.rs @@ -19,9 +19,10 @@ use super::errors::AppError; use super::models::*; use super::state::{AppState, QueryInput, SearchQuery}; use crate::health::report as health_report_compute; -use crate::model::constants::NAMESPACE_NAME_MAX_BYTES; +use crate::model::constants::MAX_BATCH_MEMORIES; use crate::model::id::{MemoryId, NamespaceId}; use crate::model::memory::AccessKind; +use crate::model::validation::{MemoryInputRef, validate_memory_input, validate_namespace_name}; use crate::serialization::{ ApiResponse, MemoryResponse, NamespaceRequest, NamespaceResponse, SearchHit, SearchRequest, SearchResponse, @@ -51,7 +52,8 @@ fn health_report_cache() /// POST /memories -- create a new memory. /// /// Steps: -/// 1. Validate request: summary non-empty, tags <= 64, namespace exists. +/// 1. Validate request content limits (see [`validate_memory_input`]); +/// all length limits are UTF-8 bytes, not characters. /// 2. Resolve namespace by name -> NamespaceId. /// 3. Generate embedding if not provided (calls embedding provider). /// 4. Validate embedding dimensionality against namespace config. @@ -66,32 +68,14 @@ pub async fn create_memory( let start = Instant::now(); // --- Validation --- - if req.summary.is_empty() { - return Err(AppError::BadRequest { - message: "summary must not be empty".into(), - field: Some("summary".into()), - }); - } - if req.summary.len() > 2000 { - return Err(AppError::BadRequest { - message: "summary exceeds 2,000 byte limit".into(), - field: Some("summary".into()), - }); - } - if let Some(ref text) = req.full_text { - if text.len() > 1_048_576 { - return Err(AppError::BadRequest { - message: "full_text exceeds 1 MB limit".into(), - field: Some("fullText".into()), - }); - } - } - if req.tags.len() > 64 { - return Err(AppError::BadRequest { - message: "too many tags (max 64)".into(), - field: Some("tags".into()), - }); - } + validate_memory_input(MemoryInputRef { + summary: &req.summary, + full_text: req.full_text.as_deref(), + tags: &req.tags, + entities: &req.entities, + topics: &req.topics, + emotions: &req.emotions, + })?; // --- Resolve namespace --- let ns = state @@ -914,7 +898,8 @@ pub async fn list_namespaces( /// POST /namespaces -- create a new namespace. /// /// Steps: -/// 1. Validate name format (1-64 chars, alphanumeric + hyphens + underscores). +/// 1. Validate name format (1-64 UTF-8 bytes, ASCII alphanumeric + +/// hyphens + underscores). /// 2. Check for duplicate name. /// 3. Register namespace with fixed embedding dimensionality. /// 4. Return 201 with namespace details. @@ -925,25 +910,7 @@ pub async fn create_namespace( let start = Instant::now(); // Validate name format - if req.name.is_empty() || req.name.len() > NAMESPACE_NAME_MAX_BYTES { - return Err(AppError::BadRequest { - message: format!("namespace name must be 1-{NAMESPACE_NAME_MAX_BYTES} characters") - .into(), - field: Some("name".into()), - }); - } - if !req - .name - .chars() - .all(|c| c.is_alphanumeric() || c == '-' || c == '_') - { - return Err(AppError::BadRequest { - message: - "namespace name may only contain alphanumeric characters, hyphens, and underscores" - .into(), - field: Some("name".into()), - }); - } + validate_namespace_name(&req.name)?; // Validate embedding dimensions if provided if let Some(dim) = req.embedding_dim { @@ -1252,9 +1219,9 @@ pub async fn batch_store( }); } - if req.memories.len() > 100 { + if req.memories.len() > MAX_BATCH_MEMORIES { return Err(AppError::BadRequest { - message: "batch size exceeds 100".into(), + message: format!("batch size exceeds the maximum of {MAX_BATCH_MEMORIES}"), field: Some("memories".into()), }); } @@ -1262,8 +1229,21 @@ pub async fn batch_store( let mut created = Vec::with_capacity(req.memories.len()); for mem_req in req.memories { - // Validate - if mem_req.summary.is_empty() { + // Validate. This endpoint silently skips invalid items rather + // than failing the whole batch, so an error here is a `continue` + // — but the limits themselves are the shared ones, which keeps + // an oversized summary from reaching the u16-prefixed on-disk + // record encoder. + if validate_memory_input(MemoryInputRef { + summary: &mem_req.summary, + full_text: mem_req.full_text.as_deref(), + tags: &mem_req.tags, + entities: &mem_req.entities, + topics: &mem_req.topics, + emotions: &mem_req.emotions, + }) + .is_err() + { continue; } diff --git a/src/cli/client.rs b/src/cli/client.rs index 5ce7191..52ed0ff 100644 --- a/src/cli/client.rs +++ b/src/cli/client.rs @@ -11,6 +11,8 @@ use crate::cli::output::{ ForgetResult, HealthReportView, InspectView, ListResult, MemoryView, NamespaceStatsView, NamespaceView, ReinforceResult, SearchResult, StatusView, StoreResult, SweepResult, }; +use crate::model::constants::SUMMARY_MAX_BYTES; +use crate::model::validation::truncate_on_char_boundary; /// HTTP client for the Recalld API server. /// @@ -116,10 +118,14 @@ impl RecalldClient { parent_id: Option<&'a uuid::Uuid>, } - // If text exceeds 2,000 bytes, treat it as full_text and let - // the server generate a summary. Otherwise, use it as the summary. - let (summary, full_text) = if text.len() > 2000 { - (&text[..200], Some(text)) // Use first 200 chars as provisional summary + // If text exceeds the summary byte limit, treat it as full_text + // and let the server generate a summary. Otherwise, use it as + // the summary. + let (summary, full_text) = if text.len() > SUMMARY_MAX_BYTES { + // First 200 UTF-8 *bytes*, cut on a character boundary — + // slicing `&text[..200]` panics when byte 200 lands inside a + // multi-byte character. + (truncate_on_char_boundary(text, 200), Some(text)) } else { (text, None) }; diff --git a/src/mcp/tools.rs b/src/mcp/tools.rs index b75f428..995ee39 100644 --- a/src/mcp/tools.rs +++ b/src/mcp/tools.rs @@ -8,6 +8,24 @@ use serde_json::json; use crate::mcp::bridge::McpBridge; use crate::mcp::protocol::{ToolAnnotations, ToolCallResult, ToolInfo}; +use crate::model::constants::{ + FULL_TEXT_MAX_BYTES, MAX_BATCH_MEMORIES, MAX_EMOTIONS, MAX_ENTITIES, MAX_TAGS, MAX_TOPICS, + SUMMARY_MAX_BYTES, +}; +use crate::model::validation::{MemoryInputRef, validate_memory_input, validate_namespace_name}; + +/// Schema description for the `summary` field. +/// +/// The limit is enforced in UTF-8 bytes; the previous wording said +/// "chars", which made a 1,900-character em-dash summary look like it +/// should fit when it is 5,700 bytes. +const SUMMARY_DESC: &str = "Short description of the memory. Limit is 2000 UTF-8 bytes, \ + which is 2000 ASCII characters but fewer for text containing em dashes, curly quotes, \ + accents, or emoji (2-4 bytes each)."; + +/// Schema description for the `fullText` field. +const FULL_TEXT_DESC: &str = "Detailed content. Removed as memory decays to ghost phase. \ + Limit is 1048576 UTF-8 bytes (1 MiB)."; // ═══════════════════════════════════════════════════════════════════════ // Registry and dispatch @@ -68,31 +86,40 @@ fn store_memory_def() -> ToolInfo { "properties": { "summary": { "type": "string", - "description": "Short description of the memory (max 2000 chars)", - "maxLength": 2000 + "description": SUMMARY_DESC, + // Counts code points, and chars <= bytes always, so + // this is strictly looser than the server's byte + // limit — a useful client-side guard that can never + // reject input the server would accept. + "maxLength": SUMMARY_MAX_BYTES }, "fullText": { "type": "string", - "description": "Detailed content. Removed as memory decays to ghost phase." + "description": FULL_TEXT_DESC, + "maxLength": FULL_TEXT_MAX_BYTES }, "tags": { "type": "array", "items": { "type": "string" }, + "maxItems": MAX_TAGS, "description": "Categorization tags, e.g. [\"topic/rust\", \"type/observation\"]" }, "entities": { "type": "array", "items": { "type": "string" }, + "maxItems": MAX_ENTITIES, "description": "Named entities (people, places, orgs, titles) mentioned in this memory. Used for search indexing and graph linking." }, "topics": { "type": "array", "items": { "type": "string" }, + "maxItems": MAX_TOPICS, "description": "Topic keywords describing what the memory is about, e.g. [\"rust\", \"cooking\", \"career\"]" }, "emotions": { "type": "array", "items": { "type": "string" }, + "maxItems": MAX_EMOTIONS, "description": "Emotional tone if relevant, e.g. [\"happy\", \"anxious\", \"grateful\"]" }, "namespace": { @@ -126,24 +153,11 @@ async fn handle_store_memory(bridge: &McpBridge, arguments: serde_json::Value) - None => return ToolCallResult::error("Missing required parameter: summary"), }; - // Issue 2: Enforce summary length limit - if summary.len() > 2000 { - return ToolCallResult::error("Summary exceeds maximum length of 2000 characters"); - } - let full_text = arguments .get("fullText") .and_then(|v| v.as_str()) .map(String::from); - // Issue 2: Enforce full_text length limit (1 MB) - const MAX_FULL_TEXT_BYTES: usize = 1_048_576; - if let Some(ref ft) = full_text { - if ft.len() > MAX_FULL_TEXT_BYTES { - return ToolCallResult::error("fullText exceeds maximum length of 1 MB"); - } - } - let tags: Vec = arguments .get("tags") .and_then(|v| serde_json::from_value(v.clone()).ok()) @@ -161,19 +175,19 @@ async fn handle_store_memory(bridge: &McpBridge, arguments: serde_json::Value) - .and_then(|v| serde_json::from_value(v.clone()).ok()) .unwrap_or_default(); - // Issue 3: Enforce array size limits - if tags.len() > 64 { - return ToolCallResult::error("Too many tags (maximum 64)"); - } - if entities.len() > 32 { - return ToolCallResult::error("Too many entities (maximum 32)"); - } - if topics.len() > 32 { - return ToolCallResult::error("Too many topics (maximum 32)"); - } - if emotions.len() > 32 { - return ToolCallResult::error("Too many emotions (maximum 32)"); + // Content limits — shared with the HTTP API so the two paths cannot + // disagree about what fits. + if let Err(e) = validate_memory_input(MemoryInputRef { + summary: &summary, + full_text: full_text.as_deref(), + tags: &tags, + entities: &entities, + topics: &topics, + emotions: &emotions, + }) { + return ToolCallResult::error(e.to_string()); } + let namespace = arguments .get("namespace") .and_then(|v| v.as_str()) @@ -231,37 +245,42 @@ fn store_memories_def() -> ToolInfo { "memories": { "type": "array", "description": "Array of memories to store (max 100 per call)", - "maxItems": 100, + "maxItems": MAX_BATCH_MEMORIES, "items": { "type": "object", "properties": { "summary": { "type": "string", - "description": "Short description of the memory (max 2000 chars)", - "maxLength": 2000 + "description": SUMMARY_DESC, + "maxLength": SUMMARY_MAX_BYTES }, "fullText": { "type": "string", - "description": "Detailed content. Removed as memory decays to ghost phase." + "description": FULL_TEXT_DESC, + "maxLength": FULL_TEXT_MAX_BYTES }, "tags": { "type": "array", "items": { "type": "string" }, + "maxItems": MAX_TAGS, "description": "Categorization tags, e.g. [\"topic/rust\", \"type/observation\"]" }, "entities": { "type": "array", "items": { "type": "string" }, + "maxItems": MAX_ENTITIES, "description": "Named entities (people, places, orgs, titles) mentioned in this memory." }, "topics": { "type": "array", "items": { "type": "string" }, + "maxItems": MAX_TOPICS, "description": "Topic keywords describing what the memory is about" }, "emotions": { "type": "array", "items": { "type": "string" }, + "maxItems": MAX_EMOTIONS, "description": "Emotional tone if relevant" }, "namespace": { @@ -308,8 +327,10 @@ async fn handle_store_memories(bridge: &McpBridge, arguments: serde_json::Value) return ToolCallResult::error("Parameter 'memories' must not be empty"); } - if memories_arr.len() > 100 { - return ToolCallResult::error("Too many memories (maximum 100 per call)"); + if memories_arr.len() > MAX_BATCH_MEMORIES { + return ToolCallResult::error(format!( + "Too many memories (maximum {MAX_BATCH_MEMORIES} per call)" + )); } let mut results: Vec = Vec::with_capacity(memories_arr.len()); @@ -326,30 +347,11 @@ async fn handle_store_memories(bridge: &McpBridge, arguments: serde_json::Value) } }; - if summary.len() > 2000 { - results.push(json!({ - "index": index, - "error": "Summary exceeds maximum length of 2000 characters" - })); - continue; - } - let full_text = item .get("fullText") .and_then(|v| v.as_str()) .map(String::from); - const MAX_FULL_TEXT_BYTES: usize = 1_048_576; - if let Some(ref ft) = full_text { - if ft.len() > MAX_FULL_TEXT_BYTES { - results.push(json!({ - "index": index, - "error": "fullText exceeds maximum length of 1 MB" - })); - continue; - } - } - let tags: Vec = item .get("tags") .and_then(|v| serde_json::from_value(v.clone()).ok()) @@ -367,31 +369,17 @@ async fn handle_store_memories(bridge: &McpBridge, arguments: serde_json::Value) .and_then(|v| serde_json::from_value(v.clone()).ok()) .unwrap_or_default(); - if tags.len() > 64 { + if let Err(e) = validate_memory_input(MemoryInputRef { + summary: &summary, + full_text: full_text.as_deref(), + tags: &tags, + entities: &entities, + topics: &topics, + emotions: &emotions, + }) { results.push(json!({ "index": index, - "error": "Too many tags (maximum 64)" - })); - continue; - } - if entities.len() > 32 { - results.push(json!({ - "index": index, - "error": "Too many entities (maximum 32)" - })); - continue; - } - if topics.len() > 32 { - results.push(json!({ - "index": index, - "error": "Too many topics (maximum 32)" - })); - continue; - } - if emotions.len() > 32 { - results.push(json!({ - "index": index, - "error": "Too many emotions (maximum 32)" + "error": e.to_string() })); continue; } @@ -1118,16 +1106,9 @@ async fn handle_create_namespace( None => return ToolCallResult::error("Missing required parameter: name"), }; - // Issue 1: Server-side validation of namespace name - if name.is_empty() - || name.len() > 64 - || !name - .chars() - .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-') - { - return ToolCallResult::error( - "Invalid namespace name: must be 1-64 characters, alphanumeric, hyphens, or underscores only", - ); + // Server-side validation of namespace name, shared with the HTTP API. + if let Err(e) = validate_namespace_name(&name) { + return ToolCallResult::error(e.to_string()); } let embedding_dim = arguments @@ -1348,3 +1329,72 @@ async fn handle_list_memories(bridge: &McpBridge, arguments: serde_json::Value) Err(e) => ToolCallResult::error(format!("List memories failed: {e}")), } } + +// ═══════════════════════════════════════════════════════════════════════ +// Tests +// ═══════════════════════════════════════════════════════════════════════ + +#[cfg(test)] +mod tests { + use super::*; + + /// The `summary` sub-schema of each store tool, keyed by tool name. + fn summary_schemas() -> Vec<(&'static str, serde_json::Value)> { + let single = store_memory_def().input_schema["properties"]["summary"].clone(); + let batch = store_memories_def().input_schema["properties"]["memories"]["items"] + ["properties"]["summary"] + .clone(); + vec![("store_memory", single), ("store_memories", batch)] + } + + #[test] + fn summary_schema_describes_the_limit_in_bytes_not_chars() { + for (tool, schema) in summary_schemas() { + let desc = schema["description"] + .as_str() + .unwrap_or_else(|| panic!("{tool}: summary has no description")); + assert!(desc.contains("bytes"), "{tool}: {desc}"); + assert!(!desc.contains("chars"), "{tool}: {desc}"); + } + } + + #[test] + fn summary_schema_keeps_max_length_guard() { + // maxLength counts code points and chars <= bytes always, so + // 2000 is strictly looser than the server's byte limit; it can + // never reject something the server would accept. + for (tool, schema) in summary_schemas() { + assert_eq!(schema["maxLength"].as_u64(), Some(2000), "{tool}"); + } + } + + #[test] + fn array_fields_declare_max_items() { + let single = store_memory_def().input_schema; + let batch = store_memories_def().input_schema; + let batch_item = &batch["properties"]["memories"]["items"]; + for props in [&single["properties"], &batch_item["properties"]] { + assert_eq!(props["tags"]["maxItems"].as_u64(), Some(MAX_TAGS as u64)); + assert_eq!( + props["entities"]["maxItems"].as_u64(), + Some(MAX_ENTITIES as u64) + ); + assert_eq!( + props["topics"]["maxItems"].as_u64(), + Some(MAX_TOPICS as u64) + ); + assert_eq!( + props["emotions"]["maxItems"].as_u64(), + Some(MAX_EMOTIONS as u64) + ); + assert_eq!( + props["fullText"]["maxLength"].as_u64(), + Some(FULL_TEXT_MAX_BYTES as u64) + ); + } + assert_eq!( + batch["properties"]["memories"]["maxItems"].as_u64(), + Some(MAX_BATCH_MEMORIES as u64) + ); + } +} diff --git a/src/model/constants.rs b/src/model/constants.rs index 765a1d6..69c27d7 100644 --- a/src/model/constants.rs +++ b/src/model/constants.rs @@ -46,16 +46,47 @@ pub const DEFAULT_SUMMARY_THRESHOLD: f32 = 0.3; pub const DEFAULT_GHOST_THRESHOLD: f32 = 0.05; // ── Content Limits ─────────────────────────────────────────────────── -/// Maximum byte length of `summary` (UTF-8). +/// Maximum length of `summary`, measured in UTF-8 **bytes**, not +/// characters. ASCII text gets 2,000 characters; text containing em +/// dashes, curly quotes, accents, or emoji gets fewer, because those +/// characters occupy 2-4 bytes each. +/// +/// The limit exists because the on-disk record encodes the summary with +/// a `u16` length prefix (see `DiskRecord::to_bytes` in +/// [`crate::model::record`]); it must stay well below `u16::MAX`. pub const SUMMARY_MAX_BYTES: usize = 2_000; -/// Maximum byte length of `full_text` (UTF-8). 1 MiB. +/// Maximum length of `full_text`, measured in UTF-8 **bytes**, not +/// characters. 1 MiB = 1,048,576 bytes. +/// +/// This bounds the *raw* text, measured before JSON escaping. A request +/// carrying the maximum `full_text` can serialize to considerably more +/// than 1 MiB on the wire (`\n` doubles, control bytes expand up to 6x +/// as `\u00XX`), which is why the daemon's frame limit sits well above +/// it — see [`crate::daemon::protocol::MAX_MESSAGE_SIZE`]. pub const FULL_TEXT_MAX_BYTES: usize = 1_048_576; // ── Tag Limits ─────────────────────────────────────────────────────── /// Maximum number of tags per memory. pub const MAX_TAGS: usize = 64; -/// Maximum byte length of a single tag (UTF-8). +/// Maximum length of a single tag, measured in UTF-8 **bytes**, not +/// characters. pub const TAG_MAX_BYTES: usize = 128; +/// Maximum byte length of a single entity, topic, or emotion label +/// (UTF-8). Bounds the worst-case serialized request size; see +/// [`crate::daemon::protocol`]. +pub const LABEL_MAX_BYTES: usize = 128; + +// ── Structured Metadata Limits ─────────────────────────────────────── +/// Maximum number of `entities` accepted on a single store request. +pub const MAX_ENTITIES: usize = 32; +/// Maximum number of `topics` accepted on a single store request. +pub const MAX_TOPICS: usize = 32; +/// Maximum number of `emotions` accepted on a single store request. +pub const MAX_EMOTIONS: usize = 32; + +// ── Batch Limits ───────────────────────────────────────────────────── +/// Maximum number of memories accepted in a single batch store call. +pub const MAX_BATCH_MEMORIES: usize = 100; // ── Access History ─────────────────────────────────────────────────── /// Maximum number of `AccessEvent` entries retained per memory. @@ -66,7 +97,9 @@ pub const ACCESS_HISTORY_MAX: usize = 32; pub const DEFAULT_DESIRED_RETENTION: f32 = 0.9; // ── Namespace ──────────────────────────────────────────────────────── -/// Maximum byte length of a namespace name. +/// Maximum length of a namespace name, measured in UTF-8 **bytes**, not +/// characters. Namespace names are restricted to ASCII (they are +/// filesystem-adjacent), so for valid names bytes and characters agree. pub const NAMESPACE_NAME_MAX_BYTES: usize = 64; /// Default embedding dimensionality (OpenAI text-embedding-3-small). diff --git a/src/model/error.rs b/src/model/error.rs index dd30092..f5d1fa3 100644 --- a/src/model/error.rs +++ b/src/model/error.rs @@ -30,28 +30,59 @@ pub enum TagError { // ValidationError // ═══════════════════════════════════════════════════════════════════════ +/// Returns a clarifying sentence when a value's byte length differs +/// from its character count, and an empty string otherwise. +/// +/// Length limits in Recalld are measured in UTF-8 bytes. For pure ASCII +/// input that is indistinguishable from a character count, so the hint +/// would be noise; for text containing multi-byte characters it is the +/// whole explanation. +fn utf8_hint(bytes: usize, chars: usize) -> &'static str { + if bytes == chars { + "" + } else { + " Limits are measured in UTF-8 bytes, not characters: \ + non-ASCII characters such as em dashes and curly quotes \ + count as 2-4 bytes each." + } +} + /// Errors from `Memory::validate()` and memory creation validation. /// -/// Each variant carries a machine-readable `code()` suitable for the -/// JSON error response `"error"` field, plus a human-readable message -/// via `Display`. +/// Each variant renders a human-readable message via `Display`. Variants +/// that pertain to a specific request field also expose that field's +/// camelCase wire name via [`ValidationError::field`], suitable for the +/// JSON error response `"field"` member. #[derive(Debug, Error)] pub enum ValidationError { /// Summary field is empty. #[error("summary must not be empty")] SummaryEmpty, - /// Summary exceeds the maximum byte length. - #[error("summary is {len} bytes, max is {max}")] - SummaryTooLong { len: usize, max: usize }, - - /// Full text exceeds the maximum byte length. - #[error("full_text is {len} bytes, max is {max}")] - FullTextTooLong { len: usize, max: usize }, + /// A text field exceeds its maximum byte length. + /// + /// Carries both units so the message can explain the difference: + /// the limit is on UTF-8 bytes, but callers usually think in + /// characters. + #[error( + "{field} is {bytes} bytes ({chars} characters), which exceeds \ + the {max_bytes}-byte limit.{}", + utf8_hint(*bytes, *chars) + )] + FieldTooLong { + field: &'static str, + bytes: usize, + chars: usize, + max_bytes: usize, + }, - /// Too many tags on a single memory. - #[error("too many tags: {count}, max is {max}")] - TooManyTags { count: usize, max: usize }, + /// An array field has more items than allowed. + #[error("{field} has {count} items, which exceeds the maximum of {max}")] + TooManyItems { + field: &'static str, + count: usize, + max: usize, + }, /// A tag failed validation. #[error("invalid tag: {source}")] @@ -105,8 +136,8 @@ pub enum ValidationError { /// Namespace name is invalid (empty, too long, or bad characters). #[error( - "namespace name must be 1-{max} characters, \ - alphanumeric/hyphens/underscores" + "namespace name must be 1-{max} bytes (UTF-8) of ASCII \ + alphanumerics, hyphens, or underscores" )] InvalidNamespaceName { max: usize }, @@ -115,6 +146,30 @@ pub enum ValidationError { InvalidDecayMultiplier { value: f32 }, } +impl ValidationError { + /// The camelCase wire name of the request field this error concerns, + /// or `None` for errors not attributable to a single input field. + pub fn field(&self) -> Option<&'static str> { + match self { + ValidationError::SummaryEmpty => Some("summary"), + ValidationError::FieldTooLong { field, .. } + | ValidationError::TooManyItems { field, .. } => Some(field), + ValidationError::InvalidTag { .. } => Some("tags"), + ValidationError::DimensionMismatch { .. } + | ValidationError::NonFiniteEmbedding { .. } => Some("embedding"), + ValidationError::NamespaceNotFound { .. } => Some("namespace"), + ValidationError::InvalidNamespaceName { .. } => Some("name"), + ValidationError::InvalidStability { .. } => Some("initialStability"), + ValidationError::InvalidDecayMultiplier { .. } => Some("decayRateMultiplier"), + ValidationError::StrengthOutOfRange(_) + | ValidationError::DecayStrengthOutOfRange(_) + | ValidationError::StabilityOutOfRange { .. } + | ValidationError::DifficultyOutOfRange { .. } + | ValidationError::TimestampOrdering { .. } => None, + } + } +} + // ═══════════════════════════════════════════════════════════════════════ // DecodeError // ═══════════════════════════════════════════════════════════════════════ diff --git a/src/model/memory.rs b/src/model/memory.rs index 21bf2c0..563e47a 100644 --- a/src/model/memory.rs +++ b/src/model/memory.rs @@ -8,6 +8,7 @@ use crate::model::decay::DecayPhase; use crate::model::error::ValidationError; use crate::model::id::MemoryId; use crate::model::tag::Tag; +use crate::model::validation::{validate_count, validate_text_len}; // ═══════════════════════════════════════════════════════════════════════ // AccessKind @@ -118,36 +119,17 @@ impl Memory { /// consistent. This does NOT validate namespace existence or /// embedding dimensions (those require external context). pub fn validate(&self) -> Result<(), ValidationError> { - // Summary non-empty + // Content limits — delegated to the shared validators so this + // does not become a second source of truth alongside the + // request-time checks in `model::validation`. if self.summary.is_empty() { return Err(ValidationError::SummaryEmpty); } - - // Summary length - if self.summary.len() > SUMMARY_MAX_BYTES { - return Err(ValidationError::SummaryTooLong { - len: self.summary.len(), - max: SUMMARY_MAX_BYTES, - }); - } - - // Full text length + validate_text_len("summary", &self.summary, SUMMARY_MAX_BYTES)?; if let Some(ref ft) = self.full_text { - if ft.len() > FULL_TEXT_MAX_BYTES { - return Err(ValidationError::FullTextTooLong { - len: ft.len(), - max: FULL_TEXT_MAX_BYTES, - }); - } - } - - // Tag count - if self.tags.len() > MAX_TAGS { - return Err(ValidationError::TooManyTags { - count: self.tags.len(), - max: MAX_TAGS, - }); + validate_text_len("fullText", ft, FULL_TEXT_MAX_BYTES)?; } + validate_count("tags", self.tags.len(), MAX_TAGS)?; // Strength range if !(0.0..=1.0).contains(&self.strength) { diff --git a/src/model/mod.rs b/src/model/mod.rs index e03e9da..83d882e 100644 --- a/src/model/mod.rs +++ b/src/model/mod.rs @@ -25,6 +25,7 @@ pub mod memory; pub mod namespace; pub mod record; pub mod tag; +pub mod validation; // Re-export primary types at module level for convenience. pub use self::decay::DecayPhase; @@ -35,3 +36,4 @@ pub use self::memory::{AccessEvent, AccessKind, Memory}; pub use self::namespace::{NamespaceConfig, PhaseThresholds}; pub use self::record::{CachedRecord, DiskRecord}; pub use self::tag::{StructuredMetadata, Tag, entity_overlap, parse_structured_tags}; +pub use self::validation::{MemoryInputRef, validate_memory_input}; diff --git a/src/model/validation.rs b/src/model/validation.rs new file mode 100644 index 0000000..86713d1 --- /dev/null +++ b/src/model/validation.rs @@ -0,0 +1,320 @@ +//! Shared input validation for memory-creating entry points. +//! +//! The HTTP API, the batch HTTP endpoint, and the MCP tool handlers all +//! accept the same logical payload but each used to carry its own hand- +//! rolled copy of the limit checks. Those copies drifted: the batch +//! endpoint checked nothing but emptiness, and every message claimed the +//! limits were measured in *characters* when the code measured UTF-8 +//! *bytes*. This module is the single source of truth. +//! +//! The functions here operate on borrowed request data, before any +//! [`Memory`](crate::model::Memory) exists, which is why they are free +//! functions rather than methods on `Memory`. `Memory::validate()` +//! delegates its content-limit checks here so the two cannot diverge. + +use super::constants::{ + FULL_TEXT_MAX_BYTES, MAX_EMOTIONS, MAX_ENTITIES, MAX_TAGS, MAX_TOPICS, + NAMESPACE_NAME_MAX_BYTES, SUMMARY_MAX_BYTES, +}; +use super::error::ValidationError; + +/// A borrowed view of the user-supplied fields of a store request. +/// +/// `entities`, `topics`, and `emotions` are input-only: callers merge +/// them into `tags` after validation, so they must be checked here while +/// they are still distinguishable. +pub struct MemoryInputRef<'a> { + pub summary: &'a str, + pub full_text: Option<&'a str>, + pub tags: &'a [String], + pub entities: &'a [String], + pub topics: &'a [String], + pub emotions: &'a [String], +} + +/// Validate every content limit on a store request. +/// +/// Checks run in a fixed order — summary emptiness, summary length, +/// full text length, then the array counts — so a request violating +/// several limits always reports the same one. +pub fn validate_memory_input(input: MemoryInputRef<'_>) -> Result<(), ValidationError> { + if input.summary.is_empty() { + return Err(ValidationError::SummaryEmpty); + } + validate_text_len("summary", input.summary, SUMMARY_MAX_BYTES)?; + if let Some(ft) = input.full_text { + validate_text_len("fullText", ft, FULL_TEXT_MAX_BYTES)?; + } + validate_count("tags", input.tags.len(), MAX_TAGS)?; + validate_count("entities", input.entities.len(), MAX_ENTITIES)?; + validate_count("topics", input.topics.len(), MAX_TOPICS)?; + validate_count("emotions", input.emotions.len(), MAX_EMOTIONS)?; + Ok(()) +} + +/// Check that `value` fits within `max_bytes` UTF-8 bytes. +/// +/// The character count is computed only when the check fails, so the +/// happy path costs one `len()`. +pub fn validate_text_len( + field: &'static str, + value: &str, + max_bytes: usize, +) -> Result<(), ValidationError> { + let bytes = value.len(); + if bytes > max_bytes { + return Err(ValidationError::FieldTooLong { + field, + bytes, + chars: value.chars().count(), + max_bytes, + }); + } + Ok(()) +} + +/// Check that an array field holds at most `max` items. +pub fn validate_count( + field: &'static str, + count: usize, + max: usize, +) -> Result<(), ValidationError> { + if count > max { + return Err(ValidationError::TooManyItems { field, count, max }); + } + Ok(()) +} + +/// Validate a namespace name: 1-64 bytes of ASCII alphanumerics, +/// hyphens, or underscores. +/// +/// Namespace names become directory names on disk, so the character set +/// is deliberately restricted to ASCII rather than Unicode +/// alphanumerics. +pub fn validate_namespace_name(name: &str) -> Result<(), ValidationError> { + let valid = !name.is_empty() + && name.len() <= NAMESPACE_NAME_MAX_BYTES + && name + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'); + if valid { + Ok(()) + } else { + Err(ValidationError::InvalidNamespaceName { + max: NAMESPACE_NAME_MAX_BYTES, + }) + } +} + +/// Return the longest prefix of `text` that is at most `max_bytes` long +/// and ends on a character boundary. +/// +/// Slicing a `&str` at an arbitrary byte offset panics when the offset +/// splits a multi-byte character; this is the non-panicking form. +/// (`str::floor_char_boundary` would do the same job but is still +/// unstable.) +pub fn truncate_on_char_boundary(text: &str, max_bytes: usize) -> &str { + if text.len() <= max_bytes { + return text; + } + let end = text + .char_indices() + .map(|(i, _)| i) + .take_while(|&i| i <= max_bytes) + .last() + .unwrap_or(0); + &text[..end] +} + +#[cfg(test)] +mod tests { + use super::*; + + fn input(summary: &str) -> MemoryInputRef<'_> { + MemoryInputRef { + summary, + full_text: None, + tags: &[], + entities: &[], + topics: &[], + emotions: &[], + } + } + + fn strings(n: usize) -> Vec { + (0..n).map(|i| format!("t{i}")).collect() + } + + const HINT: &str = "Limits are measured in UTF-8 bytes, not characters"; + + // ── Summary length: bytes vs characters ────────────────────────── + + #[test] + fn em_dash_summary_reports_both_units_and_hint() { + // 1900 em dashes = 1900 characters but 5700 bytes. This is the + // case that produced the misleading "exceeds maximum length of + // 2000 characters" message. + let summary = "—".repeat(1900); + assert_eq!(summary.len(), 5700); + let err = validate_memory_input(input(&summary)).unwrap_err(); + let msg = err.to_string(); + assert!(msg.contains("5700 bytes (1900 characters)"), "{msg}"); + assert!(msg.contains("2000-byte limit"), "{msg}"); + assert!(msg.contains(HINT), "{msg}"); + assert_eq!(err.field(), Some("summary")); + } + + #[test] + fn ascii_1900_is_accepted() { + assert!(validate_memory_input(input(&"a".repeat(1900))).is_ok()); + } + + #[test] + fn ascii_2100_reports_equal_units_without_hint() { + let err = validate_memory_input(input(&"a".repeat(2100))).unwrap_err(); + let msg = err.to_string(); + assert!(msg.contains("2100 bytes (2100 characters)"), "{msg}"); + assert!(!msg.contains(HINT), "{msg}"); + } + + #[test] + fn summary_byte_boundary_is_inclusive() { + assert!(validate_memory_input(input(&"a".repeat(2000))).is_ok()); + assert!(validate_memory_input(input(&"a".repeat(2001))).is_err()); + } + + #[test] + fn em_dash_boundary_counts_bytes_not_chars() { + let under = "—".repeat(666); // 1998 bytes + assert_eq!(under.len(), 1998); + assert!(validate_memory_input(input(&under)).is_ok()); + + let over = "—".repeat(667); // 2001 bytes, 667 characters + assert_eq!(over.len(), 2001); + let msg = validate_memory_input(input(&over)).unwrap_err().to_string(); + assert!(msg.contains("2001 bytes (667 characters)"), "{msg}"); + } + + #[test] + fn four_byte_chars_at_exact_byte_limit_are_accepted() { + let summary = "🧠".repeat(500); // 2000 bytes, 500 characters + assert_eq!(summary.len(), 2000); + assert!(validate_memory_input(input(&summary)).is_ok()); + } + + #[test] + fn empty_summary_is_rejected() { + let err = validate_memory_input(input("")).unwrap_err(); + assert!(matches!(err, ValidationError::SummaryEmpty)); + assert_eq!(err.field(), Some("summary")); + } + + // ── full_text ──────────────────────────────────────────────────── + + #[test] + fn full_text_boundary() { + let at_limit = "a".repeat(FULL_TEXT_MAX_BYTES); + let mut inp = input("ok"); + inp.full_text = Some(&at_limit); + assert!(validate_memory_input(inp).is_ok()); + + let over = "a".repeat(FULL_TEXT_MAX_BYTES + 1); + let mut inp = input("ok"); + inp.full_text = Some(&over); + let err = validate_memory_input(inp).unwrap_err(); + assert_eq!(err.field(), Some("fullText")); + assert!(err.to_string().contains("1048576-byte limit")); + } + + // ── Array counts ───────────────────────────────────────────────── + + #[test] + fn array_count_limits() { + let cases: [(&str, usize); 4] = [ + ("tags", MAX_TAGS), + ("entities", MAX_ENTITIES), + ("topics", MAX_TOPICS), + ("emotions", MAX_EMOTIONS), + ]; + for (field, max) in cases { + let at = strings(max); + let over = strings(max + 1); + for (items, should_err) in [(&at, false), (&over, true)] { + let mut inp = input("ok"); + match field { + "tags" => inp.tags = items.as_slice(), + "entities" => inp.entities = items.as_slice(), + "topics" => inp.topics = items.as_slice(), + _ => inp.emotions = items.as_slice(), + } + let result = validate_memory_input(inp); + if should_err { + let err = result.unwrap_err(); + assert_eq!(err.field(), Some(field)); + match err { + ValidationError::TooManyItems { count, max: m, .. } => { + assert_eq!(count, max + 1); + assert_eq!(m, max); + } + other => panic!("unexpected error for {field}: {other}"), + } + } else { + assert!(result.is_ok(), "{field} at {max} should be accepted"); + } + } + } + } + + #[test] + fn summary_is_reported_before_tag_count() { + let summary = "a".repeat(3000); + let tags = strings(100); + let mut inp = input(&summary); + inp.tags = &tags; + let err = validate_memory_input(inp).unwrap_err(); + assert_eq!(err.field(), Some("summary")); + } + + // ── truncate_on_char_boundary ──────────────────────────────────── + + #[test] + fn truncate_never_splits_a_character() { + let text = "—".repeat(300); // 900 bytes, 3 bytes per char + let cut = truncate_on_char_boundary(&text, 200); + assert!(cut.len() <= 200); + assert_eq!(cut.len(), 198); // largest multiple of 3 <= 200 + assert!(text.starts_with(cut)); + } + + #[test] + fn truncate_returns_whole_string_when_short_enough() { + assert_eq!(truncate_on_char_boundary("hello", 200), "hello"); + assert_eq!(truncate_on_char_boundary("hello", 5), "hello"); + } + + #[test] + fn truncate_to_zero_or_below_first_char_yields_empty() { + assert_eq!(truncate_on_char_boundary("—abc", 0), ""); + assert_eq!(truncate_on_char_boundary("—abc", 2), ""); + } + + // ── Namespace names ────────────────────────────────────────────── + + #[test] + fn namespace_names() { + assert!(validate_namespace_name("a-b_c1").is_ok()); + assert!(validate_namespace_name(&"a".repeat(64)).is_ok()); + + assert!(validate_namespace_name("").is_err()); + assert!(validate_namespace_name(&"a".repeat(65)).is_err()); + assert!(validate_namespace_name("a b").is_err()); + assert!(validate_namespace_name("ünïcode").is_err()); + } + + #[test] + fn namespace_error_message_says_bytes() { + let err = validate_namespace_name("").unwrap_err(); + assert!(err.to_string().contains("1-64 bytes"), "{err}"); + assert_eq!(err.field(), Some("name")); + } +} From 42621ce7321caa43898e7a1bffbfa5211299f623 Mon Sep 17 00:00:00 2001 From: Caleb Evans Date: Tue, 11 Aug 2026 00:38:57 -0600 Subject: [PATCH 2/8] fix: stop an oversized daemon frame from killing the connection fullText's 1 MiB limit was exactly the daemon's 1 MiB frame limit, so a request that passed validation could exceed the frame once wrapped in its JSON envelope. Reproduced against a running server: four sequential store_memory calls on one connection gave OK, then a 1,048,570-byte write failed with "Broken pipe", and then every later call failed the same way forever. Only restarting the MCP server recovered. The response direction was worse and silent. A single get_memory on a large memory, or a recall_memories returning several, made the daemon write an over-limit frame. The client rejected it after consuming the 4-byte length prefix but before draining the payload, leaving unread JSON in the socket -- no broken pipe, just every subsequent read_u32 parsing from the middle of a message. - Raise the frame limit to 8 MiB and keep fullText at 1 MiB. What overflows is JSON escaping of the text (up to 6x for control bytes), which is a transport artifact, so the headroom belongs in the transport. The worst-case arithmetic is documented on the constant. - Check the size before any byte reaches the socket, so an oversized request fails cleanly and leaves the connection usable. - Drain an oversized but well-formed incoming frame into a sink under a timeout. Length-prefixed framing makes the boundary known, so this resynchronizes exactly rather than best-effort. A length beyond the drain limit is treated as unrecoverable desync and closes instead. - Stop dropping the whole connection for one bad frame; recoverable errors now reply and keep serving. - Add lazy reconnect to DaemonClient with bounded attempts. Retry only read-only methods: store_memory mints a server-side id, so a transparent retry could double-write, and neither a failed write nor a failed read proves the operation did not execute. Mutating calls return an explicit "may or may not have been applied" error instead. - Correlate response ids. Stale replies are skipped up to a bound, and an impossible id resets the connection rather than looping forever. - Negotiate the frame limit via ping, defaulting to 1 MiB for an unannounced peer, so a v0.1.10 daemon or client on either side of the socket is never sent a frame it will reject. - Dispatch "list_memories", which the client called but the daemon had no arm for, leaving that tool broken in daemon mode. The primary fix does not depend on retry: poisoning plus lazy reconnect alone ends the permanent-failure behavior. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01JNhUehmChiQKJUjv3QPnkH --- docs/architecture.md | 48 +++ src/daemon/client.rs | 816 +++++++++++++++++++++++++++++++++++++++-- src/daemon/protocol.rs | 548 ++++++++++++++++++++++++--- src/daemon/server.rs | 349 +++++++++++++++++- src/mcp/bridge.rs | 18 + 5 files changed, 1685 insertions(+), 94 deletions(-) diff --git a/docs/architecture.md b/docs/architecture.md index 5882318..767c507 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -693,6 +693,54 @@ Two transports expose the same 9 MCP tools (`store_memory`, `store_memories`, `r - Stale socket cleanup on startup (checks if PID is alive). - MCP clients connect to the daemon via JSON-RPC 2.0. +#### Framing + +Messages are length-prefixed JSON: a 4-byte big-endian byte count followed by +the JSON payload. + +| Limit | Value | Meaning | +|---|---|---| +| `MAX_MESSAGE_SIZE` | 8 MiB | Largest frame this build will send or accept. | +| `LEGACY_MAX_MESSAGE_SIZE` | 1 MiB | What a protocol v1 peer (recalld <= 0.1.10) enforces. | +| `MAX_DRAIN_SIZE` | 64 MiB | Largest declared length still worth skipping past. | +| `DRAIN_TIMEOUT` | 30 s | How long a peer may take to deliver a frame being skipped. | + +The frame limit is deliberately much larger than the 1 MiB `fullText` content +limit. The content limit measures raw UTF-8 bytes; the frame carries those +bytes *after* JSON escaping, which expands `\n` 2x and control bytes up to 6x. +A worst-case `store_memory` request with a maximum `fullText` serializes to +roughly 6.1 MiB, so the headroom belongs in the transport rather than in the +content limit. + +**Drain and resync.** Once the length prefix is read, the frame boundary is +known. An oversized but believable frame (over the peer limit, under +`MAX_DRAIN_SIZE`) is skipped exactly — copied to a sink, so no allocation — +which leaves the stream sitting on the next frame boundary. Both sides then +report the error and keep the connection. Only three conditions close a +connection: a clean EOF, a declared length beyond `MAX_DRAIN_SIZE` (the +boundary is fiction), and a drain that times out. + +Oversized messages are rejected at serialization time, **before** any byte +reaches the socket, so an outsized request cannot desynchronize a connection +for the calls that follow it. + +**Protocol negotiation.** `ping` doubles as a handshake: the client sends +`{"protocolVersion": 2, "maxMessageSize": 8388608}` and the daemon answers in +kind. Each side writes frames sized for what the other announced, defaulting +to 1 MiB for a peer that announces nothing. This matters because a protocol v1 +peer rejects an oversized length prefix *without* draining the payload, which +corrupts its stream permanently — so a v2 peer never sends one an oversized +frame. Negotiation is repeated on every reconnect. + +**Client recovery.** A transport failure poisons the client connection; the +next call re-establishes it (3 attempts, 50/200/500 ms backoff). Read-only +methods are replayed transparently on the fresh connection. Mutating methods +are never replayed — `store_memory` mints a server-side id, `reinforce_memory` +advances the FSRS schedule — and instead return an error stating that the +operation may or may not have been applied. Request ids are monotonic across +reconnects, so a response left over from an abandoned call can be recognized +and discarded rather than mistaken for the current one. + ### CLI client (recalld-cli) - Separate binary that communicates with the HTTP API server (default `http://localhost:7680`). diff --git a/src/daemon/client.rs b/src/daemon/client.rs index b0c46ca..91e0d3c 100644 --- a/src/daemon/client.rs +++ b/src/daemon/client.rs @@ -1,74 +1,824 @@ -use std::path::Path; +use std::cmp::Ordering; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering}; +use std::time::Duration; use tokio::io::BufReader; +use tokio::net::UnixStream; +use tokio::net::unix::{OwnedReadHalf, OwnedWriteHalf}; use tokio::sync::Mutex; -use super::protocol::{DaemonRequest, read_framed_response, write_framed_request}; +use super::protocol::{ + DaemonRequest, FrameError, LEGACY_MAX_MESSAGE_SIZE, MAX_MESSAGE_SIZE, PROTOCOL_VERSION, + read_framed_response, write_framed_request, +}; use crate::mcp::bridge::BridgeError; -/// Client for communicating with the Recalld daemon over a Unix socket. -pub struct DaemonClient { - inner: Mutex, +/// Connect attempts made before a reconnect is declared hopeless. +const RECONNECT_ATTEMPTS: usize = 3; +/// Backoff between reconnect attempts. The final entry is never slept on. +const RECONNECT_BACKOFF_MS: [u64; RECONNECT_ATTEMPTS] = [50, 200, 500]; +/// How many stale (already-abandoned) responses may be skipped before the +/// stream is treated as corrupt. +const MAX_STALE_SKIPS: usize = 8; + +/// Whether a method may be replayed on a fresh connection after the original +/// attempt failed with an unknown outcome. +/// +/// Only read-only methods qualify. Every mutating method here either mints +/// new server-side state (`store_memory` allocates a fresh id, so a replay +/// silently creates a duplicate the caller can never reconcile), advances +/// state on each call (`reinforce_memory` moves the FSRS schedule), or has a +/// return value that is not idempotent even where the effect is +/// (`delete_memory` returns whether the memory existed). +fn is_retry_safe(method: &str) -> bool { + matches!( + method, + "ping" + | "check_health" + | "search" + | "find_similar" + | "scan_duplicates" + | "get_memory" + | "list_namespaces" + | "namespace_stats" + | "list_memories" + ) +} + +/// An established socket connection to the daemon. +struct Connection { + reader: BufReader, + writer: OwnedWriteHalf, +} + +impl Connection { + fn new(stream: UnixStream) -> Self { + let (reader, writer) = stream.into_split(); + Self { + reader: BufReader::new(reader), + writer, + } + } } -struct DaemonClientInner { - reader: BufReader, - writer: tokio::net::unix::OwnedWriteHalf, +/// Mutable client state, serialized by the client's mutex. +struct ClientState { + /// `None` once the connection has been poisoned; re-established lazily. + conn: Option, + /// Monotonic *across reconnects*, so a frame left over from a dropped + /// connection can never alias the id of a live call. next_id: u64, + /// Largest frame the daemon is known to accept. Held at the protocol v1 + /// limit until a `ping` negotiates something larger. + frame_limit: usize, +} + +/// A failed exchange, plus whether the connection survived it. +struct ExchangeError { + error: BridgeError, + /// The connection was poisoned and must be re-established. When set, the + /// request may or may not have been applied by the daemon: neither a + /// failed write nor a failed read proves the daemon did not act. + lost: bool, +} + +impl ExchangeError { + /// The call failed but the connection is still usable and the outcome is + /// known (nothing was sent, or the daemon answered definitively). + fn settled(error: BridgeError) -> Self { + Self { error, lost: false } + } + + /// The connection was poisoned; the outcome of the call is unknown. + fn lost(error: BridgeError) -> Self { + Self { error, lost: true } + } +} + +/// Client for communicating with the Recalld daemon over a Unix socket. +/// +/// The connection is self-healing: a transport failure poisons it, and the +/// next call transparently re-establishes it. An oversized request is +/// rejected before a single byte reaches the socket, so it cannot corrupt +/// the stream for subsequent calls. +pub struct DaemonClient { + socket_path: PathBuf, + state: Mutex, + /// Mirrors `state.conn.is_some()` so [`DaemonClient::is_connected`] can + /// stay synchronous. + connected: AtomicBool, } impl DaemonClient { /// Connects to the daemon at the given Unix socket path. + /// + /// The frame limit starts at the conservative protocol v1 value; call + /// [`DaemonClient::ping`] to negotiate a larger one. pub async fn connect(socket_path: &Path) -> std::io::Result { - let stream = tokio::net::UnixStream::connect(socket_path).await?; - let (reader, writer) = stream.into_split(); + let stream = UnixStream::connect(socket_path).await?; Ok(Self { - inner: Mutex::new(DaemonClientInner { - reader: BufReader::new(reader), - writer, + socket_path: socket_path.to_path_buf(), + state: Mutex::new(ClientState { + conn: Some(Connection::new(stream)), next_id: 1, + frame_limit: LEGACY_MAX_MESSAGE_SIZE, }), + connected: AtomicBool::new(true), }) } - /// The Mutex serializes requests on this connection. + /// Whether the client currently holds a live connection. + /// + /// `false` only means the next call will reconnect first, not that the + /// daemon is gone. + pub fn is_connected(&self) -> bool { + self.connected.load(AtomicOrdering::Relaxed) + } + + /// Sends a request and awaits its response. + /// + /// The mutex is held for the whole exchange, so calls are serialized on + /// the connection. + /// + /// **Not cancel-safe.** Dropping the returned future between the write + /// and the read abandons a response that a later call must then skip; + /// the client tolerates a bounded number of such strays but does not + /// treat cancellation as a supported pattern. pub async fn call( &self, method: &str, params: serde_json::Value, ) -> Result { - let mut inner = self.inner.lock().await; - let id = inner.next_id; - inner.next_id += 1; + let mut state = self.state.lock().await; + let outcome = self.call_locked(&mut state, method, params).await; + self.connected + .store(state.conn.is_some(), AtomicOrdering::Relaxed); + outcome + } - let request = DaemonRequest { + async fn call_locked( + &self, + state: &mut ClientState, + method: &str, + params: serde_json::Value, + ) -> Result { + if state.conn.is_none() { + self.reconnect(state).await?; + } + + let mut request = DaemonRequest { jsonrpc: "2.0".into(), - id, + id: 0, method: method.into(), params, }; - write_framed_request(&mut inner.writer, &request) - .await - .map_err(|e| BridgeError::Internal(format!("daemon write error: {e}")))?; + let failure = match state.exchange(&mut request).await { + Ok(result) => return Ok(result), + Err(e) => e, + }; - let response = read_framed_response(&mut inner.reader) - .await - .map_err(|e| BridgeError::Internal(format!("daemon read error: {e}")))? - .ok_or_else(|| BridgeError::Internal("daemon connection closed".into()))?; + // The connection is intact: the answer, good or bad, is final. + if !failure.lost { + return Err(failure.error); + } - if let Some(err) = response.error { - return Err(err.into_bridge_error()); + let reconnected = self.reconnect(state).await; + + if !is_retry_safe(method) { + let mut message = format!( + "daemon connection lost while awaiting the response to `{method}` \ + ({}); the operation may or may not have been applied.", + failure.error + ); + if reconnected.is_ok() { + message.push_str( + " The connection has been re-established - verify with recall_memories \ + before retrying.", + ); + } else { + message.push_str( + " The daemon is currently unreachable - verify with recall_memories \ + once it is back before retrying.", + ); + } + return Err(BridgeError::Internal(message)); } - response - .result - .ok_or_else(|| BridgeError::Internal("empty daemon response".into())) + reconnected?; + state.exchange(&mut request).await.map_err(|e| e.error) } - /// Sends a ping to verify the daemon is responsive. + /// Sends a ping to verify the daemon is responsive and negotiate the + /// frame limit for this connection. pub async fn ping(&self) -> Result<(), BridgeError> { - self.call("ping", serde_json::json!({})).await?; + let result = self.call("ping", handshake_params()).await?; + let limit = negotiated_limit(&result); + self.state.lock().await.frame_limit = limit; Ok(()) } + + /// Re-establishes the connection, negotiating the frame limit on the way. + /// + /// Bounded: at most [`RECONNECT_ATTEMPTS`] attempts with a fixed backoff, + /// so this can never spin. + async fn reconnect(&self, state: &mut ClientState) -> Result<(), BridgeError> { + state.conn = None; + state.frame_limit = LEGACY_MAX_MESSAGE_SIZE; + let mut last: Option = None; + + for (attempt, backoff_ms) in RECONNECT_BACKOFF_MS.iter().enumerate() { + match UnixStream::connect(&self.socket_path).await { + Ok(stream) => { + state.conn = Some(Connection::new(stream)); + state.frame_limit = LEGACY_MAX_MESSAGE_SIZE; + match state.negotiate().await { + Ok(()) => return Ok(()), + Err(e) => { + state.conn = None; + last = Some(e.to_string()); + } + } + } + Err(e) => last = Some(e.to_string()), + } + + // No point sleeping after the final attempt. + if attempt + 1 < RECONNECT_ATTEMPTS { + tokio::time::sleep(Duration::from_millis(*backoff_ms)).await; + } + } + + Err(BridgeError::Internal(format!( + "daemon unavailable at {} after {RECONNECT_ATTEMPTS} attempts: {}", + self.socket_path.display(), + last.unwrap_or_else(|| "unknown error".into()) + ))) + } + + /// Test hook: the frame limit currently in force. + #[cfg(test)] + async fn frame_limit(&self) -> usize { + self.state.lock().await.frame_limit + } +} + +impl ClientState { + /// Performs the handshake on a freshly established connection. + /// + /// Never recurses into reconnect: it drives [`ClientState::exchange`] + /// directly. + async fn negotiate(&mut self) -> Result<(), BridgeError> { + let mut request = DaemonRequest { + jsonrpc: "2.0".into(), + id: 0, + method: "ping".into(), + params: handshake_params(), + }; + let result = self.exchange(&mut request).await.map_err(|e| e.error)?; + self.frame_limit = negotiated_limit(&result); + Ok(()) + } + + /// One request/response round trip on the current connection. + /// + /// Assigns `request.id` so a retry gets a fresh id while reusing the + /// (potentially large) params without cloning them. + async fn exchange( + &mut self, + request: &mut DaemonRequest, + ) -> Result { + let method = request.method.clone(); + let limit = self.frame_limit; + request.id = self.next_id; + self.next_id += 1; + let id = request.id; + + let written = match self.conn.as_mut() { + Some(conn) => write_framed_request(&mut conn.writer, request, limit).await, + None => { + return Err(ExchangeError::lost(BridgeError::Internal( + "daemon connection is not established".into(), + ))); + } + }; + + if let Err(e) = written { + return Err(match e { + // Rejected before anything reached the socket: the connection + // is untouched and stays usable. This is what keeps one + // oversized request from breaking every call after it. + FrameError::TooLarge { size, limit } => { + ExchangeError::settled(BridgeError::too_large(size, limit)) + } + FrameError::Malformed(detail) => ExchangeError::settled(BridgeError::InvalidInput( + format!("could not serialize `{method}` request: {detail}"), + )), + other => { + self.conn = None; + ExchangeError::lost(BridgeError::Internal(format!( + "daemon write error on `{method}`: {other}" + ))) + } + }); + } + + for _ in 0..=MAX_STALE_SKIPS { + let read = match self.conn.as_mut() { + Some(conn) => read_framed_response(&mut conn.reader, limit).await, + None => { + return Err(ExchangeError::lost(BridgeError::Internal( + "daemon connection is not established".into(), + ))); + } + }; + + let response = match read { + Ok(response) => response, + // The frame was skipped or unparseable but the stream is back + // on a boundary: only this call is lost. + Err(e) if e.is_recoverable() => { + return Err(ExchangeError::settled(BridgeError::TooLarge(format!( + "the daemon's response to `{method}` could not be delivered: {e}" + )))); + } + Err(e) => { + self.conn = None; + return Err(ExchangeError::lost(BridgeError::Internal(format!( + "daemon connection lost while awaiting the response to `{method}`: {e}" + )))); + } + }; + + // The daemon uses id 0 when a framing error hid the request id. + if response.id == 0 { + let error = response + .error + .map(|e| e.into_bridge_error()) + .unwrap_or_else(|| { + BridgeError::Internal( + "daemon reported an unattributable protocol error".into(), + ) + }); + return Err(ExchangeError::settled(error)); + } + + match response.id.cmp(&id) { + Ordering::Equal => { + if let Some(err) = response.error { + return Err(ExchangeError::settled(err.into_bridge_error())); + } + return match response.result { + Some(result) => Ok(result), + None => Err(ExchangeError::settled(BridgeError::Internal( + "empty daemon response".into(), + ))), + }; + } + // Ids are monotonic, so a lower id belongs to a call that was + // already abandoned. Provably not ours; discard it. + Ordering::Less => { + tracing::debug!( + got = response.id, + expected = id, + "discarding stale daemon response" + ); + } + // Impossible under monotonic ids: the stream is corrupt. + Ordering::Greater => { + self.conn = None; + return Err(ExchangeError::lost(BridgeError::Internal(format!( + "daemon response id mismatch: expected {id}, got {}; \ + connection reset to resynchronize", + response.id + )))); + } + } + } + + self.conn = None; + Err(ExchangeError::lost(BridgeError::Internal(format!( + "daemon sent more than {MAX_STALE_SKIPS} stale responses while awaiting \ + `{method}`; connection reset to resynchronize" + )))) + } +} + +/// Params announcing this build's protocol capabilities on `ping`. +fn handshake_params() -> serde_json::Value { + serde_json::json!({ + "protocolVersion": PROTOCOL_VERSION, + "maxMessageSize": MAX_MESSAGE_SIZE, + }) +} + +/// Frame limit implied by a `ping` result. +/// +/// A daemon that says nothing is protocol v1 and is held at 1 MiB, which +/// keeps us from ever handing it a frame it would reject without draining. +fn negotiated_limit(result: &serde_json::Value) -> usize { + let version = result + .get("protocolVersion") + .and_then(serde_json::Value::as_u64) + .unwrap_or(1); + if version < PROTOCOL_VERSION as u64 { + return LEGACY_MAX_MESSAGE_SIZE; + } + let announced = result + .get("maxMessageSize") + .and_then(serde_json::Value::as_u64) + .unwrap_or(LEGACY_MAX_MESSAGE_SIZE as u64); + (announced as usize).clamp(LEGACY_MAX_MESSAGE_SIZE, MAX_MESSAGE_SIZE) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + use std::sync::Mutex as StdMutex; + + use serde_json::json; + use tokio::net::UnixListener; + + use super::super::protocol::{DaemonResponse, read_framed_message, write_framed_message}; + use super::*; + + /// Scripted daemon behaviours used to drive the client through each + /// failure mode. All counters below are global across connections, so a + /// behaviour that fires "on the first call" does not fire again on the + /// connection the client reconnects with. + #[derive(Clone, Copy, PartialEq, Eq, Debug)] + enum Behavior { + /// Answer every call successfully; negotiate protocol v2 on ping. + Healthy, + /// Answer ping with `{}`, as recalld <= 0.1.10 does. + LegacyPing, + /// Answer the first call with a response larger than the client's limit. + HugeResponse, + /// Read the first call, then hang up without replying. + CloseAfterFirstCall, + /// Prefix the second call's response with a stale one. + StaleThenCorrect, + /// Answer the first call with an id from the future. + FutureId, + } + + struct Stub { + path: PathBuf, + received: Arc>>, + _dir: tempfile::TempDir, + handle: tokio::task::JoinHandle<()>, + } + + impl Stub { + fn start(behavior: Behavior) -> Self { + let dir = tempfile::tempdir().expect("tempdir"); + let path = dir.path().join("t.sock"); + let listener = UnixListener::bind(&path).expect("bind"); + let received = Arc::new(StdMutex::new(Vec::new())); + let handle = tokio::spawn(serve(listener, behavior, received.clone())); + Self { + path, + received, + _dir: dir, + handle, + } + } + + fn methods(&self) -> Vec { + self.received + .lock() + .unwrap() + .iter() + .map(|r| r.method.clone()) + .collect() + } + + fn count(&self, method: &str) -> usize { + self.methods().iter().filter(|m| *m == method).count() + } + + /// Stops listening and removes the socket so reconnects fail. + async fn shutdown(&self) { + self.handle.abort(); + let _ = std::fs::remove_file(&self.path); + // Give the abort a chance to land so the live connection drops. + tokio::time::sleep(Duration::from_millis(50)).await; + } + } + + async fn serve( + listener: UnixListener, + behavior: Behavior, + received: Arc>>, + ) { + let mut calls = 0usize; + loop { + let Ok((stream, _)) = listener.accept().await else { + return; + }; + let (reader, writer) = stream.into_split(); + let mut reader = BufReader::new(reader); + let mut writer = writer; + + loop { + let Ok(request) = read_framed_message(&mut reader, MAX_MESSAGE_SIZE).await else { + break; + }; + received.lock().unwrap().push(request.clone()); + + if request.method == "ping" { + let result = if behavior == Behavior::LegacyPing { + json!({}) + } else { + json!({ + "protocolVersion": PROTOCOL_VERSION, + "maxMessageSize": MAX_MESSAGE_SIZE, + }) + }; + let response = DaemonResponse::success(request.id, result); + if write_framed_message(&mut writer, &response, MAX_MESSAGE_SIZE) + .await + .is_err() + { + break; + } + continue; + } + + calls += 1; + let replies: Vec = match behavior { + Behavior::CloseAfterFirstCall if calls == 1 => break, + Behavior::HugeResponse if calls == 1 => vec![DaemonResponse::success( + request.id, + json!({ "blob": "x".repeat(2 * 1024 * 1024) }), + )], + Behavior::StaleThenCorrect if calls == 2 => vec![ + DaemonResponse::success(request.id - 1, json!({ "stale": true })), + DaemonResponse::success(request.id, json!({ "ok": true })), + ], + Behavior::FutureId if calls == 1 => { + vec![DaemonResponse::success(request.id + 5, json!({}))] + } + _ => vec![DaemonResponse::success(request.id, json!({ "ok": true }))], + }; + + let mut failed = false; + for reply in &replies { + if write_framed_message(&mut writer, reply, MAX_MESSAGE_SIZE) + .await + .is_err() + { + failed = true; + break; + } + } + if failed { + break; + } + } + } + } + + /// A params blob whose serialized form comfortably exceeds `bytes`. + fn params_of_bytes(bytes: usize) -> serde_json::Value { + json!({ "fullText": "y".repeat(bytes) }) + } + + // 14 — direct regression for the reported bug. + #[tokio::test] + async fn oversize_request_fails_fast_and_connection_survives() { + let stub = Stub::start(Behavior::Healthy); + let client = DaemonClient::connect(&stub.path).await.unwrap(); + + client + .call("get_memory", json!({ "id": "a" })) + .await + .unwrap(); + + let err = client + .call("store_memory", params_of_bytes(2 * 1024 * 1024)) + .await + .unwrap_err(); + let text = err.to_string(); + assert!(matches!(err, BridgeError::TooLarge(_)), "got {text}"); + assert!( + text.contains(&LEGACY_MAX_MESSAGE_SIZE.to_string()), + "{text}" + ); + assert!(text.contains("bytes"), "{text}"); + + // The socket was never touched, so the next call still works. + assert!(client.is_connected()); + client + .call("get_memory", json!({ "id": "b" })) + .await + .unwrap(); + + // Exactly the two small calls reached the wire. + assert_eq!(stub.methods().len(), 2); + assert_eq!(stub.count("get_memory"), 2); + assert_eq!(stub.count("store_memory"), 0); + } + + // 15 + #[tokio::test] + async fn reconnects_after_server_closes_connection() { + let stub = Stub::start(Behavior::CloseAfterFirstCall); + let client = DaemonClient::connect(&stub.path).await.unwrap(); + + // The first call is swallowed; the client reconnects and retries. + client + .call("get_memory", json!({ "id": "a" })) + .await + .unwrap(); + assert!(client.is_connected()); + + client + .call("get_memory", json!({ "id": "b" })) + .await + .unwrap(); + assert_eq!(stub.count("ping"), 1, "reconnect renegotiates exactly once"); + } + + // 16 + #[tokio::test] + async fn read_only_method_retries_transparently_after_mid_call_close() { + let stub = Stub::start(Behavior::CloseAfterFirstCall); + let client = DaemonClient::connect(&stub.path).await.unwrap(); + + let result = client.call("search", json!({})).await.unwrap(); + assert_eq!(result, json!({ "ok": true })); + assert_eq!(stub.count("search"), 2, "the read-only call was replayed"); + } + + // 17 + #[tokio::test] + async fn mutating_method_is_not_retried() { + let stub = Stub::start(Behavior::CloseAfterFirstCall); + let client = DaemonClient::connect(&stub.path).await.unwrap(); + + let err = client + .call("store_memory", json!({ "summary": "s" })) + .await + .unwrap_err(); + let text = err.to_string(); + assert!(text.contains("may or may not have been applied"), "{text}"); + assert!(text.contains("re-established"), "{text}"); + + assert_eq!( + stub.count("store_memory"), + 1, + "the mutating call must reach the daemon exactly once" + ); + // The client is healthy again for the next call. + assert!(client.is_connected()); + client + .call("get_memory", json!({ "id": "a" })) + .await + .unwrap(); + } + + // 18 + #[tokio::test] + async fn reconnect_gives_up_after_bounded_attempts() { + let stub = Stub::start(Behavior::Healthy); + let client = DaemonClient::connect(&stub.path).await.unwrap(); + client + .call("get_memory", json!({ "id": "a" })) + .await + .unwrap(); + + stub.shutdown().await; + + let err = tokio::time::timeout(Duration::from_secs(5), async { + // The first call discovers the dead peer; the second starts with + // no connection at all. Both must terminate. + let _ = client.call("get_memory", json!({ "id": "b" })).await; + client.call("get_memory", json!({ "id": "c" })).await + }) + .await + .expect("reconnect must not loop forever") + .unwrap_err(); + + let text = err.to_string(); + assert!(text.contains(&stub.path.display().to_string()), "{text}"); + assert!(text.contains("3 attempts"), "{text}"); + assert!(!client.is_connected()); + } + + // 19 + #[tokio::test] + async fn stale_response_id_is_skipped() { + let stub = Stub::start(Behavior::StaleThenCorrect); + let client = DaemonClient::connect(&stub.path).await.unwrap(); + + client + .call("get_memory", json!({ "id": "a" })) + .await + .unwrap(); + let result = client + .call("get_memory", json!({ "id": "b" })) + .await + .unwrap(); + + assert_eq!(result, json!({ "ok": true }), "the stale reply was skipped"); + assert!(client.is_connected()); + assert_eq!(stub.count("ping"), 0, "no reconnect was needed"); + } + + // 20 + #[tokio::test] + async fn future_response_id_poisons_and_next_call_recovers() { + let stub = Stub::start(Behavior::FutureId); + let client = DaemonClient::connect(&stub.path).await.unwrap(); + + // A mutating method, so the mismatch surfaces instead of being + // silently replayed. + let err = client.call("store_memory", json!({})).await.unwrap_err(); + let text = err.to_string(); + assert!(text.contains("id mismatch"), "{text}"); + + let result = client + .call("get_memory", json!({ "id": "a" })) + .await + .unwrap(); + assert_eq!(result, json!({ "ok": true })); + } + + // 21 — regression for the silently-desyncing response path. + #[tokio::test] + async fn oversized_response_is_drained_and_connection_survives() { + let stub = Stub::start(Behavior::HugeResponse); + let client = DaemonClient::connect(&stub.path).await.unwrap(); + + let err = client + .call("get_memory", json!({ "id": "a" })) + .await + .unwrap_err(); + assert!(matches!(err, BridgeError::TooLarge(_)), "{err}"); + + // The oversized frame was drained, so the stream is still aligned. + assert!(client.is_connected()); + let result = client + .call("get_memory", json!({ "id": "b" })) + .await + .unwrap(); + assert_eq!(result, json!({ "ok": true })); + assert_eq!(stub.count("ping"), 0, "no reconnect was needed"); + } + + // 22 + #[tokio::test] + async fn ping_negotiates_frame_limit() { + let stub = Stub::start(Behavior::Healthy); + let client = DaemonClient::connect(&stub.path).await.unwrap(); + assert_eq!(client.frame_limit().await, LEGACY_MAX_MESSAGE_SIZE); + + client.ping().await.unwrap(); + assert_eq!(client.frame_limit().await, MAX_MESSAGE_SIZE); + } + + // 22b + #[tokio::test] + async fn legacy_daemon_returning_empty_ping_clamps_to_1mib() { + let stub = Stub::start(Behavior::LegacyPing); + let client = DaemonClient::connect(&stub.path).await.unwrap(); + + client.ping().await.unwrap(); + assert_eq!(client.frame_limit().await, LEGACY_MAX_MESSAGE_SIZE); + } + + #[test] + fn mutating_methods_are_never_retry_safe() { + for method in [ + "store_memory", + "create_namespace", + "reinforce_memory", + "delete_memory", + "shutdown", + ] { + assert!(!is_retry_safe(method), "{method} must not be replayed"); + } + for method in ["ping", "search", "get_memory", "list_memories"] { + assert!(is_retry_safe(method), "{method} should be replayable"); + } + } + + #[test] + fn negotiated_limit_ignores_unknown_and_absurd_announcements() { + assert_eq!(negotiated_limit(&json!({})), LEGACY_MAX_MESSAGE_SIZE); + assert_eq!( + negotiated_limit(&json!({ "protocolVersion": 1, "maxMessageSize": 99_999_999 })), + LEGACY_MAX_MESSAGE_SIZE + ); + assert_eq!( + negotiated_limit(&json!({ "protocolVersion": 2, "maxMessageSize": 99_999_999 })), + MAX_MESSAGE_SIZE + ); + assert_eq!( + negotiated_limit(&json!({ "protocolVersion": 2, "maxMessageSize": 1024 })), + LEGACY_MAX_MESSAGE_SIZE + ); + } } diff --git a/src/daemon/protocol.rs b/src/daemon/protocol.rs index 3611eb3..b13a524 100644 --- a/src/daemon/protocol.rs +++ b/src/daemon/protocol.rs @@ -1,5 +1,7 @@ +use std::time::Duration; + use serde::{Deserialize, Serialize}; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; use crate::mcp::bridge::BridgeError; @@ -112,6 +114,8 @@ pub const ERR_STORAGE: i32 = -32003; pub const ERR_SEARCH: i32 = -32004; /// RPC error code: internal server error. pub const ERR_INTERNAL: i32 = -32603; +/// RPC error code: the request or response did not fit in a protocol frame. +pub const ERR_TOO_LARGE: i32 = -32005; // ── BridgeError conversion ─────────────────────────────────────────── @@ -138,6 +142,10 @@ impl From<&BridgeError> for DaemonRpcError { code: ERR_INTERNAL, message: msg.clone(), }, + BridgeError::TooLarge(msg) => DaemonRpcError { + code: ERR_TOO_LARGE, + message: msg.clone(), + }, } } } @@ -150,6 +158,7 @@ impl DaemonRpcError { ERR_INVALID_INPUT => BridgeError::InvalidInput(self.message), ERR_STORAGE => BridgeError::Storage(self.message), ERR_SEARCH => BridgeError::Search(self.message), + ERR_TOO_LARGE => BridgeError::TooLarge(self.message), _ => BridgeError::Internal(self.message), } } @@ -183,76 +192,517 @@ impl DaemonResponse { // // Length-prefixed JSON: [4-byte big-endian length][JSON payload]. -const MAX_MESSAGE_SIZE: usize = 1024 * 1024; // 1 MB +/// Version of the daemon framing protocol implemented by this build. +/// +/// Version 1 (recalld <= 0.1.10) had no negotiation and a hard 1 MiB frame +/// limit on both sides. Version 2 raises the limit and adds recovery from +/// oversized frames. A v2 peer announces itself in the `ping` request / +/// response; a peer that says nothing is assumed to be v1 and is only ever +/// sent frames within [`LEGACY_MAX_MESSAGE_SIZE`]. +pub const PROTOCOL_VERSION: u32 = 2; + +/// Maximum size in bytes of a single protocol frame. +/// +/// This bounds the *serialized* JSON envelope, not the content limits the +/// MCP layer enforces. The gap between the two is JSON escaping, which is +/// a pure transport artifact, so the headroom belongs here rather than in +/// [`crate::model::constants::FULL_TEXT_MAX_BYTES`]. +/// +/// `serde_json` expands a byte by at most 6x (a control byte becomes +/// `\u00XX`), which puts the worst-case `store_memory` request at: +/// +/// ```text +/// fullText 1,048,576 x 6 = 6,291,456 +/// summary 2,000 x 6 = 12,000 +/// tags 64 x (128 x 6 + 3) = 49,344 +/// ent+top+emo 96 x (128 x 6 + 3) = 74,016 +/// namespace 64 x 6 = 384 +/// scaffold < 1,024 +/// total ~ 6,428,224 (6.13 MiB) +/// ``` +/// +/// 8 MiB is the smallest power of two that clears that with headroom: 2 MiB +/// is not enough for a newline-heavy 1 MiB `fullText` (`\n` alone doubles), +/// and 4 MiB sits under the worst case. +pub const MAX_MESSAGE_SIZE: usize = 8 * 1024 * 1024; + +/// The frame limit a protocol v1 peer (recalld <= 0.1.10) enforces. +/// +/// Used as the assumed limit for a peer that has not announced itself, so +/// we never hand an old peer a frame it will reject — an old peer rejects +/// the length prefix *without* draining the payload, which corrupts its +/// stream permanently. +pub const LEGACY_MAX_MESSAGE_SIZE: usize = 1024 * 1024; + +/// Largest declared frame length that is still worth draining. +/// +/// Draining is allocation-free (it copies into [`tokio::io::sink`]), so this +/// bounds only how much garbage we are willing to *skip*, not how much we +/// allocate — hence it is much larger than [`MAX_MESSAGE_SIZE`]. Beyond it, +/// the declared length is more likely fiction than truth and the frame +/// boundary cannot be trusted, so the connection is closed instead. +pub const MAX_DRAIN_SIZE: usize = 64 * 1024 * 1024; + +/// How long a peer may take to deliver an oversized frame we are draining. +/// +/// Without this, a peer could declare [`MAX_DRAIN_SIZE`] and then dribble +/// bytes forever, pinning one of [`MAX_CONNECTIONS`] connection tasks. +pub const DRAIN_TIMEOUT: Duration = Duration::from_secs(30); /// Maximum number of concurrent client connections the daemon will accept. /// Additional connections are dropped immediately until an existing one closes. pub const MAX_CONNECTIONS: u32 = 32; -/// Reads a length-prefixed JSON-RPC request from the stream. -pub async fn read_framed_message( - reader: &mut R, -) -> std::io::Result> { - let len = match reader.read_u32().await { - Ok(n) => n as usize, - Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None), - Err(e) => return Err(e), - }; - if len > MAX_MESSAGE_SIZE { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - format!("message too large: {len} bytes (max {MAX_MESSAGE_SIZE})"), - )); +/// An error encountered while reading or writing a protocol frame. +#[derive(Debug, thiserror::Error)] +pub enum FrameError { + /// The peer closed the connection cleanly between frames. + #[error("connection closed by peer")] + Closed, + + /// The outgoing message does not fit in a frame. Nothing was written. + #[error("frame too large: {size} bytes exceeds the {limit} byte protocol limit")] + TooLarge { + /// Serialized size of the message. + size: usize, + /// Frame limit in force for this peer. + limit: usize, + }, + + /// An oversized incoming frame was skipped. The stream is resynchronized. + #[error( + "discarded oversized incoming frame of {size} bytes (limit {limit}); stream resynchronized" + )] + OversizedDiscarded { + /// Declared length of the discarded frame. + size: usize, + /// Frame limit in force for this peer. + limit: usize, + }, + + /// A frame was read whole but its payload is not a valid message. + /// The stream is resynchronized. + #[error("malformed frame payload: {0}")] + Malformed(String), + + /// The declared frame length is too large to be believed. The frame + /// boundary is unknown, so the stream cannot be resynchronized. + #[error("unrecoverable framing desync: declared length {size} exceeds drain limit {limit}")] + Desync { + /// Declared length of the frame. + size: usize, + /// [`MAX_DRAIN_SIZE`]. + limit: usize, + }, + + /// The peer stopped sending mid-drain. The frame boundary was never + /// reached, so the stream cannot be resynchronized. + #[error( + "timed out draining an oversized frame of {size} bytes; stream cannot be resynchronized" + )] + DrainTimeout { + /// Declared length of the frame being drained. + size: usize, + }, + + /// Transport failure. + #[error("io error: {0}")] + Io(#[from] std::io::Error), +} + +impl FrameError { + /// True when the stream is known to sit on a frame boundary, so the + /// connection can keep serving after this error. + pub fn is_recoverable(&self) -> bool { + matches!(self, Self::OversizedDiscarded { .. } | Self::Malformed(_)) } - let mut buf = vec![0u8; len]; - reader.read_exact(&mut buf).await?; - let request = serde_json::from_slice(&buf) - .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?; - Ok(Some(request)) } -/// Writes a length-prefixed JSON-RPC response to the stream. -pub async fn write_framed_message( +/// Serializes `value` and checks it against `limit` **before** any byte +/// reaches the socket. +/// +/// This is what keeps an oversized message from poisoning a connection: the +/// caller gets [`FrameError::TooLarge`] with the stream untouched. +fn encode_frame(value: &T, limit: usize) -> Result, FrameError> { + let payload = serde_json::to_vec(value).map_err(|e| FrameError::Malformed(e.to_string()))?; + if payload.len() > limit { + return Err(FrameError::TooLarge { + size: payload.len(), + limit, + }); + } + Ok(payload) +} + +/// Writes an already-encoded payload as a length-prefixed frame. +async fn write_frame( writer: &mut W, - response: &DaemonResponse, -) -> std::io::Result<()> { - let payload = serde_json::to_vec(response) - .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?; + payload: &[u8], +) -> Result<(), FrameError> { writer.write_u32(payload.len() as u32).await?; - writer.write_all(&payload).await?; - writer.flush().await + writer.write_all(payload).await?; + writer.flush().await?; + Ok(()) } -/// Reads a length-prefixed JSON-RPC response from the stream. -pub async fn read_framed_response( +/// Reads one length-prefixed frame, recovering where recovery is provable. +/// +/// Once the length prefix has been read the frame boundary is known, so an +/// oversized-but-believable frame can be skipped exactly and the stream left +/// sitting on the next frame. See [`FrameError::is_recoverable`]. +async fn read_frame( reader: &mut R, -) -> std::io::Result> { + limit: usize, +) -> Result, FrameError> { let len = match reader.read_u32().await { Ok(n) => n as usize, - Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None), - Err(e) => return Err(e), + Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Err(FrameError::Closed), + Err(e) => return Err(FrameError::Io(e)), }; - if len > MAX_MESSAGE_SIZE { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - format!("message too large: {len} bytes (max {MAX_MESSAGE_SIZE})"), - )); + + if len > MAX_DRAIN_SIZE { + return Err(FrameError::Desync { + size: len, + limit: MAX_DRAIN_SIZE, + }); } + + if len > limit { + let mut limited = (&mut *reader).take(len as u64); + let mut sink = tokio::io::sink(); + let drain = tokio::io::copy(&mut limited, &mut sink); + return match tokio::time::timeout(DRAIN_TIMEOUT, drain).await { + Ok(Ok(n)) if n == len as u64 => { + Err(FrameError::OversizedDiscarded { size: len, limit }) + } + // EOF before the frame boundary: nothing left to resynchronize to. + Ok(Ok(_)) => Err(FrameError::Closed), + Ok(Err(e)) => Err(FrameError::Io(e)), + Err(_) => Err(FrameError::DrainTimeout { size: len }), + }; + } + let mut buf = vec![0u8; len]; reader.read_exact(&mut buf).await?; - let response = serde_json::from_slice(&buf) - .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?; - Ok(Some(response)) + Ok(buf) +} + +/// Reads a length-prefixed JSON-RPC request from the stream. +pub async fn read_framed_message( + reader: &mut R, + limit: usize, +) -> Result { + let buf = read_frame(reader, limit).await?; + serde_json::from_slice(&buf).map_err(|e| FrameError::Malformed(e.to_string())) +} + +/// Reads a length-prefixed JSON-RPC response from the stream. +pub async fn read_framed_response( + reader: &mut R, + limit: usize, +) -> Result { + let buf = read_frame(reader, limit).await?; + serde_json::from_slice(&buf).map_err(|e| FrameError::Malformed(e.to_string())) +} + +/// Writes a length-prefixed JSON-RPC response to the stream. +/// +/// Returns [`FrameError::TooLarge`] without writing anything if the response +/// exceeds `limit`. +pub async fn write_framed_message( + writer: &mut W, + response: &DaemonResponse, + limit: usize, +) -> Result<(), FrameError> { + let payload = encode_frame(response, limit)?; + write_frame(writer, &payload).await } /// Writes a length-prefixed JSON-RPC request to the stream. -pub async fn write_framed_request( +/// +/// Returns [`FrameError::TooLarge`] without writing anything if the request +/// exceeds `limit`. +pub async fn write_framed_request( writer: &mut W, request: &DaemonRequest, -) -> std::io::Result<()> { - let payload = serde_json::to_vec(request) - .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?; - writer.write_u32(payload.len() as u32).await?; - writer.write_all(&payload).await?; - writer.flush().await + limit: usize, +) -> Result<(), FrameError> { + let payload = encode_frame(request, limit)?; + write_frame(writer, &payload).await +} + +#[cfg(test)] +mod tests { + use super::*; + + /// Frames used in these tests are tiny, so a tiny limit keeps every + /// assertion sub-millisecond. + const TINY: usize = 64; + const ROOMY: usize = 64 * 1024; + + fn request(id: u64) -> DaemonRequest { + DaemonRequest { + jsonrpc: "2.0".into(), + id, + method: "ping".into(), + params: serde_json::json!({}), + } + } + + /// A request whose serialized form is far larger than [`TINY`]. + fn fat_request() -> DaemonRequest { + DaemonRequest { + jsonrpc: "2.0".into(), + id: 1, + method: "store_memory".into(), + params: serde_json::json!({ "fullText": "z".repeat(500) }), + } + } + + // 1 + #[tokio::test] + async fn round_trip_request() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + let sent = DaemonRequest { + jsonrpc: "2.0".into(), + id: 7, + method: "get_memory".into(), + params: serde_json::json!({ "id": "abc" }), + }; + + write_framed_request(&mut a, &sent, ROOMY).await.unwrap(); + let got = read_framed_message(&mut b, ROOMY).await.unwrap(); + + assert_eq!(got.id, 7); + assert_eq!(got.method, "get_memory"); + assert_eq!(got.params, serde_json::json!({ "id": "abc" })); + } + + // 2 + #[tokio::test] + async fn round_trip_response_success() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + let sent = DaemonResponse::success(3, serde_json::json!({ "ok": true })); + + write_framed_message(&mut a, &sent, ROOMY).await.unwrap(); + let got = read_framed_response(&mut b, ROOMY).await.unwrap(); + + assert_eq!(got.id, 3); + assert_eq!(got.result, Some(serde_json::json!({ "ok": true }))); + assert!(got.error.is_none()); + } + + // 2b + #[tokio::test] + async fn round_trip_response_error() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + let sent = DaemonResponse::error( + 9, + DaemonRpcError { + code: ERR_NOT_FOUND, + message: "no such memory".into(), + }, + ); + + write_framed_message(&mut a, &sent, ROOMY).await.unwrap(); + let got = read_framed_response(&mut b, ROOMY).await.unwrap(); + + let err = got.error.expect("error payload"); + assert_eq!(got.id, 9); + assert_eq!(err.code, ERR_NOT_FOUND); + assert_eq!(err.message, "no such memory"); + assert!(got.result.is_none()); + } + + // 3 + #[tokio::test] + async fn three_frames_read_back_in_order() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + for id in 1..=3 { + write_framed_request(&mut a, &request(id), ROOMY) + .await + .unwrap(); + } + + for id in 1..=3 { + assert_eq!(read_framed_message(&mut b, ROOMY).await.unwrap().id, id); + } + } + + // 4 + #[tokio::test] + async fn read_returns_closed_on_clean_eof() { + let (a, mut b) = tokio::io::duplex(ROOMY); + drop(a); + + let err = read_framed_message(&mut b, ROOMY).await.unwrap_err(); + assert!(matches!(err, FrameError::Closed), "got {err}"); + assert!(!err.is_recoverable()); + } + + // 5 + #[tokio::test] + async fn write_rejects_oversized_request_without_writing() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + + match write_framed_request(&mut a, &fat_request(), TINY) + .await + .unwrap_err() + { + FrameError::TooLarge { size, limit } => { + assert!(size > TINY); + assert_eq!(limit, TINY); + } + other => panic!("expected TooLarge, got {other}"), + } + + // Nothing at all reached the peer. + let peeked = tokio::time::timeout(Duration::from_millis(10), b.read_u8()).await; + assert!(peeked.is_err(), "the socket must be untouched"); + } + + // 6 + #[tokio::test] + async fn write_rejects_oversized_response_without_writing() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + let fat = DaemonResponse::success(1, serde_json::json!({ "blob": "z".repeat(500) })); + + let err = write_framed_message(&mut a, &fat, TINY).await.unwrap_err(); + assert!( + matches!(err, FrameError::TooLarge { limit: TINY, .. }), + "got {err}" + ); + + let peeked = tokio::time::timeout(Duration::from_millis(10), b.read_u8()).await; + assert!(peeked.is_err(), "the socket must be untouched"); + } + + // 7 — the flagship: an oversized frame must not cost us the connection. + #[tokio::test] + async fn oversized_incoming_frame_is_drained_and_stream_resyncs() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + + let junk = vec![b'x'; 200]; + a.write_u32(junk.len() as u32).await.unwrap(); + a.write_all(&junk).await.unwrap(); + write_framed_request(&mut a, &request(42), ROOMY) + .await + .unwrap(); + + let err = read_framed_message(&mut b, TINY).await.unwrap_err(); + assert!( + matches!( + err, + FrameError::OversizedDiscarded { + size: 200, + limit: TINY + } + ), + "got {err}" + ); + assert!(err.is_recoverable()); + + // The stream landed exactly on the next frame boundary. + assert_eq!(read_framed_message(&mut b, ROOMY).await.unwrap().id, 42); + } + + // 8 + #[tokio::test] + async fn desync_length_beyond_drain_limit_is_fatal() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + a.write_u32(MAX_DRAIN_SIZE as u32 + 1).await.unwrap(); + + let err = read_framed_message(&mut b, TINY).await.unwrap_err(); + assert!( + matches!(err, FrameError::Desync { limit, .. } if limit == MAX_DRAIN_SIZE), + "got {err}" + ); + assert!(!err.is_recoverable()); + } + + // 9 + #[tokio::test(start_paused = true)] + async fn drain_times_out_when_payload_never_arrives() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + a.write_u32(200).await.unwrap(); + a.write_all(b"only ten b").await.unwrap(); + + // `a` stays alive, so the read blocks rather than seeing EOF. + let err = read_framed_message(&mut b, TINY).await.unwrap_err(); + assert!( + matches!(err, FrameError::DrainTimeout { size: 200 }), + "got {err}" + ); + assert!(!err.is_recoverable()); + drop(a); + } + + // 10 + #[tokio::test] + async fn malformed_json_is_recoverable_and_stream_resyncs() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + let junk = b"not json at all"; + a.write_u32(junk.len() as u32).await.unwrap(); + a.write_all(junk).await.unwrap(); + write_framed_request(&mut a, &request(5), ROOMY) + .await + .unwrap(); + + let err = read_framed_message(&mut b, ROOMY).await.unwrap_err(); + assert!(matches!(err, FrameError::Malformed(_)), "got {err}"); + assert!(err.is_recoverable()); + + assert_eq!(read_framed_message(&mut b, ROOMY).await.unwrap().id, 5); + } + + // 11 + #[tokio::test] + async fn zero_length_frame_is_malformed_not_fatal() { + let (mut a, mut b) = tokio::io::duplex(ROOMY); + a.write_u32(0).await.unwrap(); + write_framed_request(&mut a, &request(6), ROOMY) + .await + .unwrap(); + + let err = read_framed_message(&mut b, ROOMY).await.unwrap_err(); + assert!(matches!(err, FrameError::Malformed(_)), "got {err}"); + assert!(err.is_recoverable()); + + assert_eq!(read_framed_message(&mut b, ROOMY).await.unwrap().id, 6); + } + + // 12 + #[test] + fn too_large_error_message_names_size_and_limit() { + let err = FrameError::TooLarge { + size: 2_000_000, + limit: 1_048_576, + }; + let text = err.to_string(); + assert!(text.contains("2000000"), "{text}"); + assert!(text.contains("1048576"), "{text}"); + } + + // 13 + #[test] + fn err_too_large_round_trips_through_daemon_rpc_error() { + let original = BridgeError::too_large(9_000_000, MAX_MESSAGE_SIZE); + let wire = DaemonRpcError::from(&original); + assert_eq!(wire.code, ERR_TOO_LARGE); + + let back = wire.into_bridge_error(); + assert!(matches!(back, BridgeError::TooLarge(_)), "got {back}"); + assert_eq!(back.to_string(), original.to_string()); + assert!(back.to_string().contains("9000000")); + } + + #[test] + fn frame_limits_leave_headroom_over_the_content_limit() { + // The transport must be able to carry a maximum-size `fullText` + // after worst-case JSON escaping (one byte -> `\u00XX`). + const _: () = assert!(MAX_MESSAGE_SIZE >= crate::model::constants::FULL_TEXT_MAX_BYTES * 6); + // Skipping garbage costs no memory, so the drain bound is looser. + const _: () = assert!(MAX_DRAIN_SIZE > MAX_MESSAGE_SIZE); + const _: () = assert!(LEGACY_MAX_MESSAGE_SIZE < MAX_MESSAGE_SIZE); + } } diff --git a/src/daemon/server.rs b/src/daemon/server.rs index d580178..63cd47e 100644 --- a/src/daemon/server.rs +++ b/src/daemon/server.rs @@ -12,8 +12,8 @@ use super::lifecycle; use super::protocol::*; use crate::Recalld; use crate::mcp::bridge::{ - BridgeError, CreateNamespaceInput, HealthChecker, NamespaceRegistry, SearchInput, - SearchPipeline, StorageEngine as BridgeStorageEngine, StoreInput, + BridgeError, CreateNamespaceInput, HealthChecker, ListMemoriesInput, NamespaceRegistry, + SearchInput, SearchPipeline, StorageEngine as BridgeStorageEngine, StoreInput, }; use crate::mcp::bridge_adapters::*; use crate::model::MemoryId; @@ -174,14 +174,22 @@ async fn handle_connection( let mut reader = BufReader::new(reader); let mut writer = writer; + // Largest frame this peer is known to accept. Stays at the protocol v1 + // limit until the client announces v2 in its `ping`, so an old client + // never receives a frame it would reject without draining. + let mut peer_limit = LEGACY_MAX_MESSAGE_SIZE; + loop { tokio::select! { - msg = read_framed_message(&mut reader) => { + msg = read_framed_message(&mut reader, MAX_MESSAGE_SIZE) => { match msg { - Ok(None) => break, - Ok(Some(request)) => { + Ok(request) => { let method = request.method.clone(); + if method == "ping" { + peer_limit = peer_frame_limit(&request.params); + } + let response = dispatch( &*search, &*storage, &*namespaces, &*health, &request.method, request.params, @@ -192,8 +200,38 @@ async fn handle_connection( Err(e) => DaemonResponse::error(request.id, DaemonRpcError::from(&e)), }; - if write_framed_message(&mut writer, &resp).await.is_err() { - break; + match write_framed_message(&mut writer, &resp, peer_limit).await { + Ok(()) => {} + // The result does not fit in a frame. Sending it anyway + // would desynchronize the client's stream, so send an + // error that does fit instead. + Err(FrameError::TooLarge { size, limit }) => { + tracing::warn!( + method = %method, size, limit, + "response exceeds the peer frame limit, replying with an error" + ); + let fallback = DaemonResponse::error( + request.id, + DaemonRpcError { + code: ERR_TOO_LARGE, + message: format!( + "response is {size} bytes, exceeding the {limit} byte \ + message limit. Reduce `limit`, or fetch large memories \ + individually with get_memory." + ), + }, + ); + if write_framed_message(&mut writer, &fallback, peer_limit) + .await + .is_err() + { + break; + } + } + Err(e) => { + tracing::warn!(%e, "client write error, closing connection"); + break; + } } if !matches!(method.as_str(), "ping" | "check_health" | "shutdown") { @@ -205,10 +243,21 @@ async fn handle_connection( break; } } - Err(e) => { - tracing::warn!(%e, "client read error"); - break; - } + Err(e) => match action_for_frame_error(&e) { + // The stream is provably still on a frame boundary, so + // the connection survives; only this request is lost. + ConnectionAction::Reply(err) => { + tracing::warn!(%e, "recoverable client frame error, replying and continuing"); + let resp = DaemonResponse::error(UNATTRIBUTED_ID, err); + if write_framed_message(&mut writer, &resp, peer_limit).await.is_err() { + break; + } + } + ConnectionAction::Close => { + tracing::warn!(%e, "client read error, closing connection"); + break; + } + }, } } _ = shutdown_rx.changed() => break, @@ -218,6 +267,67 @@ async fn handle_connection( connection_count.fetch_sub(1, Ordering::Relaxed); } +/// Response id used when a framing error leaves the request id unknown. +/// Client request ids start at 1, so 0 can never collide with a live call. +const UNATTRIBUTED_ID: u64 = 0; + +/// What a connection should do after a [`FrameError`]. +#[derive(Debug)] +enum ConnectionAction { + /// Report the error to the client and keep serving. + Reply(DaemonRpcError), + /// Close the connection; the stream cannot be trusted. + Close, +} + +/// Classifies a framing error into a connection-level action. +/// +/// Extracted as a free function so the policy is unit-testable without a +/// live `Recalld` instance. +fn action_for_frame_error(e: &FrameError) -> ConnectionAction { + match e { + FrameError::OversizedDiscarded { size, limit } => ConnectionAction::Reply(DaemonRpcError { + code: ERR_TOO_LARGE, + message: format!( + "request is {size} bytes, exceeding the {limit} byte daemon message limit; \ + the frame was discarded. Reduce fullText (max 1 MiB), shorten \ + tags/entities/topics/emotions, or split this into multiple memories." + ), + }), + FrameError::Malformed(detail) => ConnectionAction::Reply(DaemonRpcError { + code: ERR_INVALID_INPUT, + message: format!("malformed request frame: {detail}"), + }), + // `TooLarge` is a write-side error and never reaches this path, but + // closing is the safe default if it ever does. + FrameError::Closed + | FrameError::Desync { .. } + | FrameError::DrainTimeout { .. } + | FrameError::Io(_) + | FrameError::TooLarge { .. } => ConnectionAction::Close, + } +} + +/// Frame limit to use when writing to the peer that sent these `ping` params. +/// +/// A protocol v1 client sends no version field; it is held at the legacy +/// 1 MiB limit because it rejects an oversized length prefix without +/// draining the payload, which would corrupt its stream permanently. +fn peer_frame_limit(params: &serde_json::Value) -> usize { + let version = params + .get("protocolVersion") + .and_then(serde_json::Value::as_u64) + .unwrap_or(1); + if version < PROTOCOL_VERSION as u64 { + return LEGACY_MAX_MESSAGE_SIZE; + } + let announced = params + .get("maxMessageSize") + .and_then(serde_json::Value::as_u64) + .unwrap_or(LEGACY_MAX_MESSAGE_SIZE as u64); + (announced as usize).clamp(LEGACY_MAX_MESSAGE_SIZE, MAX_MESSAGE_SIZE) +} + fn parse_memory_id(s: &str) -> Result { let uuid = Uuid::parse_str(s) .map_err(|e| BridgeError::InvalidInput(format!("invalid memory ID: {e}")))?; @@ -233,7 +343,12 @@ async fn dispatch( params: serde_json::Value, ) -> Result { match method { - "ping" => Ok(serde_json::json!({})), + // Doubles as the protocol handshake. A v1 client ignores the extra + // fields; a v2 client uses them to size its outgoing frames. + "ping" => Ok(serde_json::json!({ + "protocolVersion": PROTOCOL_VERSION, + "maxMessageSize": MAX_MESSAGE_SIZE, + })), "shutdown" => Ok(serde_json::json!({})), @@ -294,6 +409,13 @@ async fn dispatch( serde_json::to_value(result).map_err(|e| BridgeError::Internal(e.to_string())) } + "list_memories" => { + let input: ListMemoriesInput = serde_json::from_value(params) + .map_err(|e| BridgeError::InvalidInput(e.to_string()))?; + let result = storage.list_memories(input).await?; + serde_json::to_value(result).map_err(|e| BridgeError::Internal(e.to_string())) + } + "list_namespaces" => { let result = namespaces.list_namespaces().await?; serde_json::to_value(result).map_err(|e| BridgeError::Internal(e.to_string())) @@ -323,3 +445,206 @@ async fn dispatch( ))), } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::mcp::bridge::{ + DuplicateCluster, HealthStatus, ListMemoriesResponse, MemoryRecord, NamespaceInfo, + NamespaceStats, ReinforceResult, SearchHit, SearchResponse, StoredMemory, + }; + + /// Stub subsystems for the dispatch tests. Only the unknown-method path + /// is exercised, which never reaches any of these. + struct Unused; + + #[async_trait::async_trait] + impl SearchPipeline for Unused { + async fn search(&self, _: SearchInput) -> Result { + unreachable!() + } + async fn find_similar( + &self, + _: MemoryId, + _: usize, + _: Option, + _: bool, + ) -> Result, BridgeError> { + unreachable!() + } + async fn scan_duplicates( + &self, + _: &str, + _: f32, + _: usize, + ) -> Result, BridgeError> { + unreachable!() + } + } + + #[async_trait::async_trait] + impl BridgeStorageEngine for Unused { + async fn store_memory(&self, _: StoreInput) -> Result { + unreachable!() + } + async fn get_memory(&self, _: MemoryId) -> Result, BridgeError> { + unreachable!() + } + async fn delete_memory(&self, _: MemoryId) -> Result { + unreachable!() + } + async fn reinforce_memory( + &self, + _: MemoryId, + _: u8, + ) -> Result { + unreachable!() + } + async fn list_memories( + &self, + _: ListMemoriesInput, + ) -> Result { + unreachable!() + } + } + + #[async_trait::async_trait] + impl NamespaceRegistry for Unused { + async fn list_namespaces(&self) -> Result, BridgeError> { + unreachable!() + } + async fn create_namespace( + &self, + _: CreateNamespaceInput, + ) -> Result { + unreachable!() + } + async fn namespace_stats(&self, _: &str) -> Result { + unreachable!() + } + } + + #[async_trait::async_trait] + impl HealthChecker for Unused { + async fn check_health(&self) -> HealthStatus { + unreachable!() + } + } + + async fn dispatch_stub( + method: &str, + params: serde_json::Value, + ) -> Result { + dispatch(&Unused, &Unused, &Unused, &Unused, method, params).await + } + + // 23 + #[test] + fn action_for_frame_error_classification() { + let recoverable = [ + ( + FrameError::OversizedDiscarded { + size: 9_000_000, + limit: 1_048_576, + }, + ERR_TOO_LARGE, + ), + ( + FrameError::Malformed("expected value".into()), + ERR_INVALID_INPUT, + ), + ]; + for (error, code) in recoverable { + match action_for_frame_error(&error) { + ConnectionAction::Reply(rpc) => { + assert_eq!(rpc.code, code, "for {error}"); + assert!(!rpc.message.is_empty()); + } + ConnectionAction::Close => panic!("{error} must not close the connection"), + } + assert!(error.is_recoverable(), "{error}"); + } + + let fatal = [ + FrameError::Closed, + FrameError::Desync { + size: u32::MAX as usize, + limit: MAX_DRAIN_SIZE, + }, + FrameError::DrainTimeout { size: 9_000_000 }, + FrameError::Io(std::io::Error::other("boom")), + FrameError::TooLarge { + size: 9_000_000, + limit: 1_048_576, + }, + ]; + for error in fatal { + assert!( + matches!(action_for_frame_error(&error), ConnectionAction::Close), + "{error} must close the connection" + ); + assert!(!error.is_recoverable(), "{error}"); + } + } + + // 23b — the oversized-request reply must name the limit so the caller can act. + #[test] + fn oversized_request_reply_explains_the_limit() { + let action = action_for_frame_error(&FrameError::OversizedDiscarded { + size: 9_000_000, + limit: 8_388_608, + }); + let ConnectionAction::Reply(rpc) = action else { + panic!("expected a reply"); + }; + assert!(rpc.message.contains("9000000"), "{}", rpc.message); + assert!(rpc.message.contains("8388608"), "{}", rpc.message); + } + + // 24 + #[tokio::test] + async fn dispatch_unknown_method_returns_invalid_input() { + let err = dispatch_stub("no_such_method", serde_json::json!({})) + .await + .unwrap_err(); + assert!(matches!(err, BridgeError::InvalidInput(_)), "{err}"); + assert!(err.to_string().contains("no_such_method"), "{err}"); + } + + #[tokio::test] + async fn dispatch_ping_announces_the_protocol_version() { + let result = dispatch_stub("ping", serde_json::json!({})).await.unwrap(); + assert_eq!(result["protocolVersion"], PROTOCOL_VERSION); + assert_eq!(result["maxMessageSize"], MAX_MESSAGE_SIZE); + } + + #[test] + fn peer_frame_limit_holds_unannounced_peers_at_the_legacy_limit() { + // A protocol v1 client sends no version field. + assert_eq!( + peer_frame_limit(&serde_json::json!({})), + LEGACY_MAX_MESSAGE_SIZE + ); + assert_eq!( + peer_frame_limit(&serde_json::Value::Null), + LEGACY_MAX_MESSAGE_SIZE + ); + assert_eq!( + peer_frame_limit(&serde_json::json!({ "protocolVersion": 2 })), + LEGACY_MAX_MESSAGE_SIZE + ); + assert_eq!( + peer_frame_limit( + &serde_json::json!({ "protocolVersion": 2, "maxMessageSize": MAX_MESSAGE_SIZE }) + ), + MAX_MESSAGE_SIZE + ); + // An absurd announcement is clamped rather than trusted. + assert_eq!( + peer_frame_limit( + &serde_json::json!({ "protocolVersion": 2, "maxMessageSize": 1_u64 << 40 }) + ), + MAX_MESSAGE_SIZE + ); + } +} diff --git a/src/mcp/bridge.rs b/src/mcp/bridge.rs index ecdd17b..76a29e9 100644 --- a/src/mcp/bridge.rs +++ b/src/mcp/bridge.rs @@ -495,6 +495,24 @@ pub enum BridgeError { /// An unexpected internal error occurred. #[error("Internal error: {0}")] Internal(String), + + /// The request or response exceeded the daemon's frame size limit. + #[error("Too large: {0}")] + TooLarge(String), +} + +impl BridgeError { + /// Builds a [`BridgeError::TooLarge`] describing an oversized message. + /// + /// The payload is a plain string rather than structured fields so the + /// message survives the `DaemonRpcError` wire round trip intact. + pub fn too_large(size: usize, limit: usize) -> Self { + Self::TooLarge(format!( + "request is {size} bytes, which exceeds the {limit} byte daemon message limit. \ + Reduce fullText (max 1 MiB), shorten tags/entities/topics/emotions, \ + or split this into multiple memories." + )) + } } // ═══════════════════════════════════════════════════════════════════════ From cbb3f950645bf09106c11986f3e598c9b3a6e72c Mon Sep 17 00:00:00 2001 From: Caleb Evans Date: Tue, 11 Aug 2026 01:22:10 -0600 Subject: [PATCH 3/8] fix: reject malformed MCP tool arguments instead of guessing A bug report described store_memory silently corrupting records: trailing parameters absorbed into the preceding field as literal text, tags empty, namespace wrong, supersedes ignored -- all returned as success. The corruption itself originates client-side, in tool-call serialization; the server does structured serde_json parsing over newline-delimited stdio with no text splicing, no shared buffers and no cross-request state, and the report's own reproduction stores cleanly here. That part is not ours to fix. What is ours is that the server accepted it cheerfully. Arguments were read with .get().and_then(..).ok().unwrap_or_default(), so a tags array sent as a string became an empty vec, a wrong-typed namespace fell back to the default partition, and an unparseable supersedes became None -- each silently, each reported as success. That is why roughly ten bad writes accumulated before anyone noticed. - Add src/mcp/args.rs: typed extractors that fail loudly and name the field, shared by store_memory and store_memories so the two cannot drift again. Length and count limits delegate to model::validation, so there is still one source of truth for them. - Detect literal tool-call markup in summary and fullText and reject the write. The report suggested keying on any parameter-named tag; that would be unusable, since is standard HTML5, the backbone of C# doc comments, and an Atom element. The detector instead anchors on tokens that are not English words in tag position -- the client's reserved vendor prefix, function-call literals, the exact camelCase parameter tags -- plus closing-tag adjacency, which needs two independent signals before a generic name counts. Errors carry the matched token and byte offset so a false positive is diagnosable in seconds. - Reject rather than sanitise. Stripping the markup would leave a record whose tags, namespace and supersedes are still silently wrong, and whose embedding and graph edges are computed from garbage-adjacent text. A rejected write costs one round trip and lands in the agent's transcript, where it self-corrects. - Log rejections at warn with the field and rule, never the text body. Explicit JSON null keeps meaning "absent" for every optional field: many SDKs serialize optionals as null, and treating that as a type error would break them on every call. Behavioral change: a wrong-typed tags/namespace, or an unparseable parentId/supersedes, now fails the whole write where it previously stored a quietly incomplete memory. Batch items that used to store minus their tags now report an error instead, so stored counts drop and errors rise for affected clients. The per-item error object gains an additive "field". Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01JNhUehmChiQKJUjv3QPnkH --- docs/mcp.md | 38 ++ src/mcp/args.rs | 1461 ++++++++++++++++++++++++++++++++++++++++++++++ src/mcp/mod.rs | 1 + src/mcp/tools.rs | 174 ++---- 4 files changed, 1534 insertions(+), 140 deletions(-) create mode 100644 src/mcp/args.rs diff --git a/docs/mcp.md b/docs/mcp.md index bfa1d7b..da78e68 100644 --- a/docs/mcp.md +++ b/docs/mcp.md @@ -101,6 +101,44 @@ Array limits (`tags`, `entities`, `topics`, `emotions`) are plain item counts. > **Known limitation:** `entities`, `topics`, and `emotions` are converted into `entity/…`, `topic/…`, and `emotion/…` tags *after* validation runs, and the `tags` limit of 64 is checked against the tags you supplied. A memory that supplies 64 tags plus 32 entities, 32 topics, and 32 emotions can therefore end up with up to 160 tags stored. This is accepted today and is not treated as an error. +### Argument validation + +`store_memory` and `store_memories` validate their arguments strictly. A wrong-typed or unparseable argument is an error naming the field; it is never coerced, defaulted, or dropped. Earlier versions guessed silently — a `tags` array sent as a comma-joined string stored zero tags, a wrong-typed `namespace` fell back to the default partition, and an unparseable `supersedes` stored the memory with no link — so a client with a broken serializer could write a long run of quietly wrong memories without a single error reaching the agent. + +**Explicit `null` still means "absent."** `{"fullText": null}` is treated exactly like omitting `fullText`, for every optional field. SDKs that serialize optional fields as `null` need no changes. + +Example messages: + +| Situation | Error | +|---|---| +| `"summary"` missing or `null` | `Missing required parameter: summary` | +| `"fullText": ["a"]` | `Parameter 'fullText' must be a string (got array)` | +| `"tags": "a,b"` | `Parameter 'tags' must be an array of strings (got string)` | +| `"tags": ["a", 3]` | `Parameter 'tags' must be an array of strings (element at index 2 is a number)` | +| `"supersedes": 123` | `Parameter 'supersedes' must be a string containing a UUID (got number)` | +| `"supersedes": "mem-1234"` | `Invalid UUID in parameter 'supersedes': "mem-1234"` | +| a `memories` entry that is not an object | `Item at index 3 must be an object (got string)` | + +`parentId` and `supersedes` accept any spelling `Uuid::parse_str` accepts: canonical, hyphenless, braced, and `urn:uuid:` forms. + +In `store_memories`, a rejected item does not abort the batch: its entry in `results` carries `error` plus a `field` naming the offending parameter (`""` when the item as a whole is at fault), and the remaining items are still stored. The `field` member is additive; the rest of the response shape is unchanged. + +#### Tool-call markup in text fields + +`summary` and `fullText` are additionally scanned for literal tool-call markup — the XML-ish tags a client emits when serializing a tool call, spliced into the argument text instead of being parsed. This happens when the client's serializer breaks; the symptom is a memory whose `fullText` ends in parameter closing tags and whose `tags`, `entities`, and `namespace` are silently wrong. Such a call is rejected outright rather than cleaned up, because stripping the tags cannot recover the fields that never arrived, and a sanitized write becomes a permanent record whose embedding, full-text index, and graph edges were all computed from corrupt text. + +The rejection names the matched token and its byte offset, and states plainly that nothing was stored: + +``` +Parameter 'fullText' contains literal tool-call markup ("" at byte offset 2731). +Arguments appear to have been serialized incorrectly; nothing was stored. Re-send with +structured JSON arguments and no XML tool-call tags in text fields. +``` + +The detector is deliberately narrow, because a memory store is used to record notes *about* markup. It anchors on the client's reserved vendor-namespace tag prefix, on invoke/function-call literals and the attribute-first `` and ``, on a parameter closing tag followed across whitespace only by another parameter tag or a call terminator, and on text that ends with a call terminator. Ordinary technical prose is unaffected: C# `/// ` doc comments, HTML `
`, Atom entries, DocBook ``, WS-BPEL ``, JSX, Maven POM fragments, Rust generics such as `Vec`, shell redirection, and prose quoting recalld's own parameter names in JSON all store normally. + +The one known false positive is a note whose **last line** is a bare `` quoted from a BPEL snippet. That rule exists to catch truncated corruption where the parameter closing tag was consumed but the terminator survived; if it ever bites you in practice, end the note with a sentence after the quoted snippet (trailing whitespace alone is ignored) and report it. + ### store_memory Store a new observation, fact, or piece of context. The system automatically generates an embedding for semantic search. Memories decay over time unless reinforced. diff --git a/src/mcp/args.rs b/src/mcp/args.rs new file mode 100644 index 0000000..d35aa05 --- /dev/null +++ b/src/mcp/args.rs @@ -0,0 +1,1461 @@ +//! Strict argument parsing for the MCP store tools. +//! +//! The store handlers used to pull their fields out of the raw +//! `serde_json::Value` with `.get(k).and_then(|v| v.as_str())` and +//! `serde_json::from_value(..).ok()`, so every wrong-typed or +//! unparseable argument became a silent default: a `tags` array sent as +//! a comma-joined string stored zero tags, a `namespace` sent as an +//! array fell back to the default namespace, and an unparseable +//! `supersedes` UUID dropped the link. A client whose tool-call +//! serializer mangled its arguments could therefore write a long run of +//! quietly wrong memories without a single error reaching the agent's +//! transcript. +//! +//! This module is the boundary that makes those cases loud. It does +//! three things and delegates the fourth: +//! +//! 1. **Extraction** with explicit types ([`require_str`], [`opt_str`], +//! [`opt_string_array`], [`opt_uuid`], [`require_uuid`]). +//! 2. **Type strictness** — a wrong-typed argument is an error naming +//! the field and the type that arrived, never a default. +//! 3. **Markup detection** ([`scan_client_markup`]) — text fields that +//! contain literal tool-call markup are rejected rather than stored. +//! 4. Everything about *how long* or *how many* is delegated to +//! [`crate::model::validation::validate_memory_input`], which the +//! HTTP API shares. +//! +//! Nothing here needs a bridge, a runtime, or storage: the composite +//! parser takes the default namespace as a `&str`, which is what makes +//! the whole boundary unit-testable with no mocks. +//! +//! # Explicit JSON `null` means "absent" +//! +//! [`field_opt`] deliberately maps both a missing key and an explicit +//! `null` to `None`. Many client SDKs serialize optional fields as +//! `null` rather than omitting them, and the previous +//! `.and_then(|v| v.as_str())` chain treated the two identically. Any +//! change to that would break those clients on every call. + +use serde_json::Value; + +use crate::mcp::bridge::StoreInput; +use crate::model::MemoryId; +use crate::model::validation::{MemoryInputRef, truncate_on_char_boundary, validate_memory_input}; + +// ═══════════════════════════════════════════════════════════════════════ +// Error type +// ═══════════════════════════════════════════════════════════════════════ + +/// A rejected tool argument. +/// +/// `Display` renders exactly `message`, so a handler can pass +/// `e.to_string()` straight to `ToolCallResult::error`, while `field` +/// feeds the per-item `field` member of the batch response. `field` is +/// the empty string for errors that are not attributable to one field +/// (such as "the item is not an object at all"). +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +#[error("{message}")] +pub struct ArgError { + /// camelCase wire name of the offending parameter, or `""`. + pub field: String, + /// Human-readable, agent-facing explanation. + pub message: String, +} + +impl ArgError { + fn new(field: impl Into, message: impl Into) -> Self { + Self { + field: field.into(), + message: message.into(), + } + } +} + +// ═══════════════════════════════════════════════════════════════════════ +// Helpers +// ═══════════════════════════════════════════════════════════════════════ + +/// The JSON type name of `v`, as used in error messages. +fn json_type_name(v: &Value) -> &'static str { + match v { + Value::Null => "null", + Value::Bool(_) => "boolean", + Value::Number(_) => "number", + Value::String(_) => "string", + Value::Array(_) => "array", + Value::Object(_) => "object", + } +} + +/// Truncate `s` to at most `max` UTF-8 bytes on a character boundary, +/// appending an ellipsis when anything was removed. +/// +/// Used to echo a bad value back to the caller without letting a 1 MiB +/// argument turn into a 1 MiB error message. +fn elide(s: &str, max: usize) -> String { + let kept = truncate_on_char_boundary(s, max); + if kept.len() == s.len() { + s.to_string() + } else { + format!("{kept}...") + } +} + +/// Look up `key`, treating an explicit JSON `null` exactly like an +/// absent key. +/// +/// See the module docs: this equivalence is load-bearing for client +/// SDKs that serialize `Option::None` as `null`. +fn field_opt<'a>(obj: &'a Value, key: &str) -> Option<&'a Value> { + match obj.get(key) { + None | Some(Value::Null) => None, + Some(v) => Some(v), + } +} + +/// The number of bytes of a bad value echoed back in an error message. +const ECHO_MAX_BYTES: usize = 64; + +// ═══════════════════════════════════════════════════════════════════════ +// Typed extractors +// ═══════════════════════════════════════════════════════════════════════ + +/// Extract a required string parameter. +/// +/// A missing key and an explicit `null` produce the same "missing" +/// error; a present value of the wrong type names the type that +/// arrived. +pub fn require_str(obj: &Value, key: &str) -> Result { + match field_opt(obj, key) { + None => Err(ArgError::new( + key, + format!("Missing required parameter: {key}"), + )), + Some(v) => match v.as_str() { + Some(s) => Ok(s.to_string()), + None => Err(ArgError::new( + key, + format!( + "Parameter '{key}' must be a string (got {})", + json_type_name(v) + ), + )), + }, + } +} + +/// Extract an optional string parameter. Absent and `null` both yield +/// `None`; an empty string is a legitimate `Some("")`. +pub fn opt_str(obj: &Value, key: &str) -> Result, ArgError> { + match field_opt(obj, key) { + None => Ok(None), + Some(v) => match v.as_str() { + Some(s) => Ok(Some(s.to_string())), + None => Err(ArgError::new( + key, + format!( + "Parameter '{key}' must be a string (got {})", + json_type_name(v) + ), + )), + }, + } +} + +/// Extract an optional array-of-strings parameter, defaulting to an +/// empty vector. +/// +/// A bare string is *not* coerced into a one-element array: a client +/// sending `"tags": "a,b"` has been losing its tags silently, and +/// guessing here would re-establish exactly the pattern this module +/// exists to remove. +pub fn opt_string_array(obj: &Value, key: &str) -> Result, ArgError> { + let Some(v) = field_opt(obj, key) else { + return Ok(Vec::new()); + }; + let Some(arr) = v.as_array() else { + return Err(ArgError::new( + key, + format!( + "Parameter '{key}' must be an array of strings (got {})", + json_type_name(v) + ), + )); + }; + let mut out = Vec::with_capacity(arr.len()); + for (index, item) in arr.iter().enumerate() { + match item.as_str() { + Some(s) => out.push(s.to_string()), + None => { + return Err(ArgError::new( + key, + format!( + "Parameter '{key}' must be an array of strings \ + (element at index {index} is a {})", + json_type_name(item) + ), + )); + } + } + } + Ok(out) +} + +/// Extract an optional UUID parameter. +/// +/// Parsing is `Uuid::parse_str`, which also accepts the hyphenless, +/// braced, and `urn:uuid:` forms. A present but unparseable value is an +/// error rather than a silently dropped link. +pub fn opt_uuid(obj: &Value, key: &str) -> Result, ArgError> { + let Some(v) = field_opt(obj, key) else { + return Ok(None); + }; + let Some(s) = v.as_str() else { + return Err(ArgError::new( + key, + format!( + "Parameter '{key}' must be a string containing a UUID (got {})", + json_type_name(v) + ), + )); + }; + match uuid::Uuid::parse_str(s) { + Ok(u) => Ok(Some(MemoryId::from_uuid(u))), + Err(_) => Err(ArgError::new( + key, + format!( + "Invalid UUID in parameter '{key}': \"{}\"", + elide(s, ECHO_MAX_BYTES) + ), + )), + } +} + +/// Extract a required UUID parameter. +pub fn require_uuid(obj: &Value, key: &str) -> Result { + match opt_uuid(obj, key)? { + Some(id) => Ok(id), + None => Err(ArgError::new( + key, + format!("Missing required parameter: {key}"), + )), + } +} + +// ═══════════════════════════════════════════════════════════════════════ +// Client tool-call markup detection +// ═══════════════════════════════════════════════════════════════════════ + +// The literals below are assembled with `concat!` rather than written +// out in one piece. Emitting the reserved vendor prefix verbatim inside +// a tool call is exactly the corruption this detector defends against, +// and it damages files written through such a call. + +/// Reserved client vendor-namespace prefix in an opening tag. +const VENDOR_OPEN: &str = concat!("<", "ant", "ml", ":"); +/// Reserved client vendor-namespace prefix in a closing tag. +const VENDOR_CLOSE: &str = concat!("` leads with `partnerLink=`/`operation=`, and DocBook's +/// `` takes no attributes at all. +const INVOKE_LITERALS: [&str; 8] = [ + "", + "", + ""), + concat!(""), + concat!("<", "ant", "ml", ":", "invoke name="), + concat!("<", "ant", "ml", ":", "parameter name="), +]; + +/// R3 anchors: exact open/close tags of the two camelCase parameter +/// names. +/// +/// Real XML element names are kebab-case, snake_case, or PascalCase; an +/// element spelled exactly `fullText` or `parentId` in lowerCamel, +/// matching one of our own parameters character for character, is +/// essentially only producible by a serializer that turned our schema +/// into tags. The lowercase English names (`summary`, `tags`, +/// `namespace`, `supersedes`) are deliberately excluded — they get no +/// single-signal power. +const DISTINCTIVE_PARAM_TOKENS: [&str; 4] = + ["", "", "", ""]; + +/// Parameter names that may participate in an R4 adjacency. +const PARAM_TAGS: [&str; 10] = [ + "emotions", + "entities", + "fullText", + "memories", + "namespace", + "parentId", + "summary", + "supersedes", + "tags", + "topics", +]; + +/// Terminators that may *follow* a parameter close tag (R4). +const FOLLOWER_TERMINATORS: [&str; 3] = ["", "", ""]; + +/// Terminators that may *end* a text field (R5). +/// +/// `` is excluded here: DocBook makes a note ending in one +/// plausible. It participates only as an R4 follower, where a second +/// signal backs it up. +const TRAILING_TERMINATORS: [&str; 2] = ["", ""]; + +/// Whitespace permitted between an R4 close tag and its follower. +const WS: [char; 4] = [' ', '\t', '\n', '\r']; + +/// Longest window searched for a tag name after a `` anywhere. +const TAG_NAME_SCAN_LIMIT: usize = 32; + +/// Which detection rule fired. +/// +/// The variants are numbered R1-R5 in the design note in that +/// declaration order. Evaluation order is *not* the same: R2 is tested +/// first because every vendor-prefixed invoke literal also contains a +/// vendor-namespace literal, and reporting the longer, more specific +/// token is strictly more useful to whoever has to debug the client. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum MarkupRule { + /// R1 — the reserved client vendor-namespace prefix in tag position. + VendorNamespace, + /// R2 — an invoke/function-call literal, or an attribute-first + /// ` Option<(usize, &'static str)> { + literals + .iter() + .filter_map(|lit| text.find(lit).map(|at| (at, *lit))) + .min_by(|a, b| a.0.cmp(&b.0).then(b.1.len().cmp(&a.1.len()))) +} + +/// Run one literal-table rule. +fn scan_literals(text: &str, literals: &[&'static str], rule: MarkupRule) -> Option { + first_literal_match(text, literals).map(|(offset, token)| MarkupFinding { + rule, + token: token.to_string(), + offset, + }) +} + +/// True when `s` begins with an opening tag whose name is a parameter +/// name, i.e. `` or ` bool { + let Some(rest) = s.strip_prefix('<') else { + return false; + }; + PARAM_TAGS.iter().any(|tag| { + rest.strip_prefix(tag) + .is_some_and(|after| after.starts_with('>') || after.starts_with(' ')) + }) +} + +/// R4 — a `` close tag followed, across whitespace only, by +/// another parameter tag or by a call terminator. +/// +/// Two independent signals, which is what lets generic names such as +/// `summary` participate without firing on prose. Real documentation +/// quoting `
…` continues with `

`, +/// `

`, or text — never with a second recalld parameter name +/// separated by nothing but whitespace. +fn scan_tag_adjacency(text: &str) -> Option { + for (start, _) in text.match_indices("') else { + continue; + }; + let name = &window[..gt]; + if !PARAM_TAGS.contains(&name) { + continue; + } + let follower = rest[gt + 1..].trim_start_matches(WS); + let adjacent = FOLLOWER_TERMINATORS.iter().any(|t| follower.starts_with(t)) + || starts_with_param_open(follower); + if adjacent { + return Some(MarkupFinding { + rule: MarkupRule::TagAdjacency, + token: format!(""), + offset: start, + }); + } + } + None +} + +/// R5 — the text ends with a call terminator. +/// +/// This catches truncated corruption where the parameter close tag was +/// consumed but the invoke terminator survived. It carries the highest +/// false-positive risk of the five rules (a note ending in a quoted +/// BPEL snippet fires), which is accepted deliberately: the message +/// names the token and offset, and this is the first rule to drop if a +/// real user is ever bitten. +fn scan_trailing_terminator(text: &str) -> Option { + let trimmed = text.trim_end(); + TRAILING_TERMINATORS + .iter() + .find(|t| trimmed.ends_with(**t)) + .map(|t| MarkupFinding { + rule: MarkupRule::TrailingTerminator, + token: (*t).to_string(), + offset: trimmed.len() - t.len(), + }) +} + +/// Scan `text` for literal client tool-call markup. +/// +/// Returns the first rule that fires, evaluated most-specific-first +/// (see [`MarkupRule`]). Every rule is a linear pass over fixed +/// literals, so the whole scan is linear in the length of `text` and is +/// dwarfed by the embedding generation that follows a successful store. +pub fn scan_client_markup(text: &str) -> Option { + scan_literals(text, &INVOKE_LITERALS, MarkupRule::InvokeAttribute) + .or_else(|| scan_literals(text, &VENDOR_LITERALS, MarkupRule::VendorNamespace)) + .or_else(|| { + scan_literals( + text, + &DISTINCTIVE_PARAM_TOKENS, + MarkupRule::DistinctiveParamTag, + ) + }) + .or_else(|| scan_tag_adjacency(text)) + .or_else(|| scan_trailing_terminator(text)) +} + +/// Reject `text` if it contains literal client tool-call markup. +/// +/// The whole call is refused rather than sanitised: stripping the tags +/// out cannot recover the fields that never arrived (in the reported +/// incident `tags` and `entities` were absent from the call entirely), +/// and a sanitised write becomes a permanent record whose embedding, +/// full-text index, and autolink edges were all computed from +/// garbage-adjacent text. A rejection costs one round trip and lands in +/// the agent's transcript, where a well-behaved client retries. +pub fn reject_client_markup(field: &str, text: &str) -> Result<(), ArgError> { + match scan_client_markup(text) { + None => Ok(()), + Some(f) => Err(ArgError::new( + field, + format!( + "Parameter '{field}' contains literal tool-call markup \ + (\"{}\" at byte offset {}). {MARKUP_ADVICE}", + f.token, f.offset + ), + )), + } +} + +// ═══════════════════════════════════════════════════════════════════════ +// Composite parser +// ═══════════════════════════════════════════════════════════════════════ + +/// Parse the arguments of a `store_memory` call (or one item of a +/// `store_memories` batch) into a [`StoreInput`]. +/// +/// `default_namespace` is used only when the caller omits `namespace`; +/// a *wrong-typed* `namespace` is an error, never a fallback. +/// +/// Checks run in a fixed order so that an argument object broken in +/// several ways always reports the same failure: object-ness, then +/// `summary` (presence, type, markup), `fullText` (type, markup), the +/// four array fields, `namespace`, `parentId`, `supersedes`, and +/// finally the shared content limits from +/// [`validate_memory_input`](crate::model::validation::validate_memory_input). +pub fn parse_store_input(item: &Value, default_namespace: &str) -> Result { + if !item.is_object() { + return Err(ArgError::new( + "", + format!("Item must be an object (got {})", json_type_name(item)), + )); + } + + let summary = require_str(item, "summary")?; + reject_client_markup("summary", &summary)?; + + let full_text = opt_str(item, "fullText")?; + if let Some(ft) = full_text.as_deref() { + reject_client_markup("fullText", ft)?; + } + + let tags = opt_string_array(item, "tags")?; + let entities = opt_string_array(item, "entities")?; + let topics = opt_string_array(item, "topics")?; + let emotions = opt_string_array(item, "emotions")?; + + let namespace = opt_str(item, "namespace")?.unwrap_or_else(|| default_namespace.to_string()); + + let parent_id = opt_uuid(item, "parentId")?; + let supersedes = opt_uuid(item, "supersedes")?; + + validate_memory_input(MemoryInputRef { + summary: &summary, + full_text: full_text.as_deref(), + tags: &tags, + entities: &entities, + topics: &topics, + emotions: &emotions, + }) + .map_err(|e| ArgError::new(e.field().unwrap_or(""), e.to_string()))?; + + Ok(StoreInput { + summary, + full_text, + tags, + entities, + topics, + emotions, + namespace, + embedding: None, + initial_stability: None, + parent_id, + supersedes, + }) +} + +/// Parse one item of a `store_memories` batch. +/// +/// Identical to [`parse_store_input`] except that an item which is not +/// an object names its position, which is the only context the caller +/// cannot recover from the message alone. +pub fn parse_store_item( + item: &Value, + index: usize, + default_namespace: &str, +) -> Result { + if !item.is_object() { + return Err(ArgError::new( + "", + format!( + "Item at index {index} must be an object (got {})", + json_type_name(item) + ), + )); + } + parse_store_input(item, default_namespace) +} + +// ═══════════════════════════════════════════════════════════════════════ +// Tests +// ═══════════════════════════════════════════════════════════════════════ + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + // ── Helpers ────────────────────────────────────────────────────── + + /// Assert that `text` is rejected by exactly `rule`, with exactly + /// `token` at exactly `offset`. + fn assert_rejected(text: &str, rule: MarkupRule, token: &str, offset: usize) { + let found = scan_client_markup(text) + .unwrap_or_else(|| panic!("expected a rejection, got none for {text:?}")); + assert_eq!(found.rule, rule, "rule for {text:?}"); + assert_eq!(found.token, token, "token for {text:?}"); + assert_eq!(found.offset, offset, "offset for {text:?}"); + assert_eq!( + &text[found.offset..found.offset + found.token.len()], + found.token, + "offset does not point at the token for {text:?}" + ); + } + + /// Assert that `text` is legitimate technical prose. + fn assert_accepted(text: &str) { + assert_eq!( + scan_client_markup(text), + None, + "false positive on legitimate text: {text:?}" + ); + } + + /// The corrupted `fullText` tail from the incident report. + fn report_artifact() -> String { + "Reviewed the storage layer and wrote up the design notes\n\ + [\"recalld\"]\n\ + [\"topic/rust\"]\n\ +
\n" + .to_string() + } + + // ═══════════════════════════════════════════════════════════════ + // A. Markup detector MUST REJECT + // ═══════════════════════════════════════════════════════════════ + + #[test] + fn a1_report_artifact_tail_is_rejected() { + let text = report_artifact(); + // R3 wins over the R4 adjacency and the R5 terminator that the + // same string also contains. + assert_rejected( + &text, + MarkupRule::DistinctiveParamTag, + "", + text.find("").unwrap(), + ); + } + + #[test] + fn a2_vendor_prefix_inside_a_sentence_is_rejected() { + let text = format!("The client spliced {VENDOR_OPEN}invoke tags into my note."); + assert_rejected( + &text, + MarkupRule::VendorNamespace, + VENDOR_OPEN, + text.find(VENDOR_OPEN).unwrap(), + ); + } + + #[test] + fn a3_vendor_prefixed_parameter_open_tag_is_rejected() { + let text = format!("{VENDOR_OPEN}parameter>value"); + assert_rejected(&text, MarkupRule::VendorNamespace, VENDOR_OPEN, 0); + } + + #[test] + fn a4_vendor_prefixed_invoke_attribute_is_rejected() { + let literal = concat!("<", "ant", "ml", ":", "invoke name="); + let text = format!("{literal}\"store_memory\">"); + assert_rejected(&text, MarkupRule::InvokeAttribute, literal, 0); + } + + #[test] + fn a5_vendor_prefixed_parameter_attribute_is_rejected() { + let literal = concat!("<", "ant", "ml", ":", "parameter name="); + let text = format!("{literal}\"fullText\">"); + assert_rejected(&text, MarkupRule::InvokeAttribute, literal, 0); + } + + #[test] + fn a6_vendor_prefixed_function_calls_open_and_close_are_rejected() { + let open = concat!("<", "ant", "ml", ":", "function_calls>"); + let close = concat!(""); + assert_rejected(open, MarkupRule::InvokeAttribute, open, 0); + let text = format!("note body{close}"); + assert_rejected( + &text, + MarkupRule::InvokeAttribute, + close, + text.find(close).unwrap(), + ); + } + + #[test] + fn a6b_bare_function_calls_tags_are_rejected() { + // Underscore-joined and plural: not an element in any real + // schema, so the unprefixed spelling is anchored too. + assert_rejected( + "", + MarkupRule::InvokeAttribute, + "", + 0, + ); + assert_rejected( + "body", + MarkupRule::InvokeAttribute, + "", + 4, + ); + } + + #[test] + fn a6c_bare_attribute_first_invoke_and_parameter_are_rejected() { + assert_rejected( + "", + MarkupRule::InvokeAttribute, + "", + MarkupRule::InvokeAttribute, + "", + MarkupRule::DistinctiveParamTag, + "", + 4, + ); + } + + #[test] + fn a8_bare_parent_id_close_tag_is_rejected() { + assert_rejected( + "xy", + MarkupRule::DistinctiveParamTag, + "", + 1, + ); + } + + #[test] + fn a8b_bare_camel_case_open_tags_are_rejected() { + assert_rejected( + "ab", + MarkupRule::DistinctiveParamTag, + "", + 1, + ); + assert_rejected( + "ab", + MarkupRule::DistinctiveParamTag, + "", + 1, + ); + } + + #[test] + fn a9_two_generic_param_tags_in_adjacency_are_rejected() { + assert_rejected( + "
\n[]", + MarkupRule::TagAdjacency, + "
", + 0, + ); + } + + #[test] + fn a10_param_close_then_whitespace_run_then_terminator_is_rejected() { + assert_rejected( + "\n\n
", + MarkupRule::TagAdjacency, + "", + 0, + ); + } + + #[test] + fn a11_trailing_terminator_with_trailing_whitespace_is_rejected() { + let text = "Notes on the storage layer rewrite.\n\n\n"; + assert_rejected( + text, + MarkupRule::TrailingTerminator, + "", + text.find("").unwrap(), + ); + } + + #[test] + fn a12_generic_to_generic_adjacency_with_single_space_is_rejected() { + assert_rejected( + " ", + MarkupRule::TagAdjacency, + "", + 0, + ); + } + + #[test] + fn a13_adjacency_follower_may_carry_attributes() { + // `` as a follower still counts: the + // second signal is the parameter name, not the tag's shape. + assert_rejected( + "\nx", + MarkupRule::TagAdjacency, + "", + 0, + ); + } + + // ═══════════════════════════════════════════════════════════════ + // B. Markup detector MUST ACCEPT (false-positive proof) + // ═══════════════════════════════════════════════════════════════ + + #[test] + fn b1_csharp_doc_comment_is_accepted() { + assert_accepted( + "/// Returns the widget count.\n/// int", + ); + } + + #[test] + fn b2_html_details_summary_is_accepted() { + assert_accepted("
Stack trace
panic at...
"); + } + + #[test] + fn b3_atom_entry_is_accepted() { + assert_accepted("PostBlurb"); + } + + #[test] + fn b4_docbook_parameter_element_is_accepted() { + assert_accepted("The timeout element controls retry backoff."); + } + + #[test] + fn b5_bpel_invoke_mid_prose_is_accepted() { + assert_accepted( + "Legacy BPEL wrapped calls in \ + ... inside a scope, which we replaced with gRPC.", + ); + } + + #[test] + fn b6_jsx_fragment_is_accepted() { + assert_accepted("return ;"); + } + + #[test] + fn b7_self_referential_prose_about_these_very_parameters_is_accepted() { + // The normal way to discuss the store schema. If this ever + // fires, the detector is unusable for a memory store whose + // users are engineers working on the memory store. + assert_accepted( + "The MCP store_memory tool takes {\"summary\": \"...\", \"fullText\": \"...\", \ + \"tags\": []}. A client bug spliced parameter markup into fullText, so tags \ + arrived empty.", + ); + } + + #[test] + fn b8_xml_config_with_a_namespace_element_is_accepted() { + // `` followed by `` — adjacent tags, but the + // follower is not one of our parameter names. + assert_accepted("prod\ningest-svc"); + } + + #[test] + fn b9_rust_generics_and_comparisons_are_accepted() { + assert_accepted("Vec, HashMap, impl Iterator"); + assert_accepted("the guard is a < b && c > d, which held under load"); + } + + #[test] + fn b10_maven_dependency_block_is_accepted() { + assert_accepted("xy"); + } + + #[test] + fn b11_shell_redirection_is_accepted() { + assert_accepted("cmd output.txt 2>&1"); + } + + #[test] + fn b12_adjacency_whose_follower_is_not_a_param_tag_is_accepted() { + assert_accepted("
"); + } + + #[test] + fn b13_plain_word_usage_is_accepted() { + assert_accepted( + "We invoke the parameter validator; the new note supersedes the old summary.", + ); + } + + #[test] + fn b14_generic_close_newline_non_param_follower_is_accepted() { + assert_accepted("Overview\n

Body text

"); + } + + #[test] + fn b15_empty_and_one_mib_inputs_are_accepted_quickly() { + assert_accepted(""); + // A megabyte of markup-shaped prose. The rules are linear, so + // this returns immediately; a quadratic scan would hang the + // test suite, which is the point of the case. + let para = "

Lorem ipsum dolor sit amet, consectetur adipiscing elit, \ + sed do eiusmod tempor incididunt ut labore.

\n"; + let mut big = String::with_capacity(1_100_000); + while big.len() < 1_048_576 { + big.push_str(para); + } + assert!(big.len() >= 1_048_576); + // Not `assert_accepted`: a failure there would dump a megabyte. + assert!(scan_client_markup(&big).is_none()); + } + + #[test] + fn b16_param_close_far_from_its_follower_is_accepted() { + // Whitespace only is the rule; intervening prose breaks it. + assert_accepted("
and then the element follows much later"); + } + + // ═══════════════════════════════════════════════════════════════ + // B'. Documented, deliberately accepted false positives + // ═══════════════════════════════════════════════════════════════ + + #[test] + fn b_prime_1_note_ending_in_a_quoted_bpel_terminator_is_rejected() { + // THE ONE KNOWINGLY-ACCEPTED FALSE POSITIVE. A note whose final + // line happens to be a quoted `` fires R5. R5 exists + // to catch truncated corruption where the parameter close tag + // was consumed but the terminator survived; that is worth more + // than this rare shape, and R5 is the first rule to drop if a + // real user is ever bitten. The message names the token and its + // offset so the caller can see exactly why. + let text = "The BPEL scope ended with:\n"; + assert_rejected( + text, + MarkupRule::TrailingTerminator, + "", + text.find("").unwrap(), + ); + } + + #[test] + fn b_prime_2_generic_tags_never_fire_alone() { + // Locks in the rule that a lowercase English name in tag + // position is never a single sufficient signal. + assert_accepted("x\ny"); + } + + // ═══════════════════════════════════════════════════════════════ + // C. Extractors + // ═══════════════════════════════════════════════════════════════ + + #[test] + fn c_opt_string_array_absent_null_and_empty_all_yield_empty() { + assert_eq!( + opt_string_array(&json!({}), "tags").unwrap(), + Vec::::new() + ); + assert_eq!( + opt_string_array(&json!({ "tags": null }), "tags").unwrap(), + Vec::::new() + ); + assert_eq!( + opt_string_array(&json!({ "tags": [] }), "tags").unwrap(), + Vec::::new() + ); + } + + #[test] + fn c_opt_string_array_accepts_strings() { + assert_eq!( + opt_string_array(&json!({ "tags": ["a", "b"] }), "tags").unwrap(), + vec!["a".to_string(), "b".to_string()] + ); + } + + #[test] + fn c_opt_string_array_rejects_a_bare_string() { + // The exact shape that silently stored zero tags before. + let err = opt_string_array(&json!({ "tags": "a,b" }), "tags").unwrap_err(); + assert_eq!(err.field, "tags"); + assert_eq!( + err.message, + "Parameter 'tags' must be an array of strings (got string)" + ); + } + + #[test] + fn c_opt_string_array_rejects_an_object() { + let err = opt_string_array(&json!({ "tags": { "0": "a" } }), "tags").unwrap_err(); + assert_eq!( + err.message, + "Parameter 'tags' must be an array of strings (got object)" + ); + } + + #[test] + fn c_opt_string_array_names_the_bad_element_index() { + let err = opt_string_array(&json!({ "tags": ["a", 3] }), "tags").unwrap_err(); + assert_eq!( + err.message, + "Parameter 'tags' must be an array of strings (element at index 1 is a number)" + ); + + let err = opt_string_array(&json!({ "tags": [null] }), "tags").unwrap_err(); + assert_eq!( + err.message, + "Parameter 'tags' must be an array of strings (element at index 0 is a null)" + ); + + let err = opt_string_array(&json!({ "tags": ["a", "b", 3] }), "tags").unwrap_err(); + assert!( + err.message.contains("element at index 2 is a number"), + "{err}" + ); + } + + #[test] + fn c_opt_str_treats_null_as_absent_and_keeps_empty_strings() { + assert_eq!(opt_str(&json!({}), "fullText").unwrap(), None); + assert_eq!( + opt_str(&json!({ "fullText": null }), "fullText").unwrap(), + None + ); + assert_eq!( + opt_str(&json!({ "fullText": "" }), "fullText").unwrap(), + Some(String::new()) + ); + } + + #[test] + fn c_opt_str_rejects_wrong_types() { + let err = opt_str(&json!({ "fullText": 5 }), "fullText").unwrap_err(); + assert_eq!( + err.message, + "Parameter 'fullText' must be a string (got number)" + ); + let err = opt_str(&json!({ "fullText": ["x"] }), "fullText").unwrap_err(); + assert_eq!( + err.message, + "Parameter 'fullText' must be a string (got array)" + ); + assert_eq!(err.field, "fullText"); + } + + #[test] + fn c_require_str_treats_null_as_missing() { + let absent = require_str(&json!({}), "summary").unwrap_err(); + assert_eq!(absent.message, "Missing required parameter: summary"); + let null = require_str(&json!({ "summary": null }), "summary").unwrap_err(); + assert_eq!(null, absent); + } + + #[test] + fn c_require_str_reports_wrong_types() { + let err = require_str(&json!({ "summary": 5 }), "summary").unwrap_err(); + assert_eq!( + err.message, + "Parameter 'summary' must be a string (got number)" + ); + } + + #[test] + fn c_opt_uuid_absent_and_null_yield_none() { + assert_eq!(opt_uuid(&json!({}), "supersedes").unwrap(), None); + assert_eq!( + opt_uuid(&json!({ "supersedes": null }), "supersedes").unwrap(), + None + ); + } + + #[test] + fn c_opt_uuid_accepts_the_lenient_forms_uuid_parse_str_allows() { + let canonical = "a1b2c3d4-e5f6-7890-abcd-ef1234567890"; + let expected = uuid::Uuid::parse_str(canonical).unwrap(); + // Documents that `Uuid::parse_str` is lenient: all four + // spellings are accepted and parse to the same value. + for form in [ + canonical.to_string(), + canonical.replace('-', ""), + format!("{{{canonical}}}"), + format!("urn:uuid:{canonical}"), + ] { + let got = opt_uuid(&json!({ "supersedes": form }), "supersedes") + .unwrap_or_else(|e| panic!("{form} rejected: {e}")) + .unwrap_or_else(|| panic!("{form} yielded None")); + assert_eq!(got.into_inner(), expected, "{form}"); + } + } + + #[test] + fn c_opt_uuid_rejects_unparseable_and_wrong_typed_values() { + let err = opt_uuid(&json!({ "supersedes": "mem-1234" }), "supersedes").unwrap_err(); + assert_eq!(err.field, "supersedes"); + assert_eq!( + err.message, + "Invalid UUID in parameter 'supersedes': \"mem-1234\"" + ); + + let err = opt_uuid(&json!({ "supersedes": 123 }), "supersedes").unwrap_err(); + assert_eq!( + err.message, + "Parameter 'supersedes' must be a string containing a UUID (got number)" + ); + + let err = opt_uuid(&json!({ "supersedes": "" }), "supersedes").unwrap_err(); + assert!( + err.message.starts_with("Invalid UUID in parameter"), + "{err}" + ); + } + + #[test] + fn c_opt_uuid_elides_a_long_bad_value() { + let long = "z".repeat(200); + let err = opt_uuid(&json!({ "parentId": long }), "parentId").unwrap_err(); + assert!(err.message.ends_with("...\""), "{err}"); + assert!(err.message.len() < 140, "{err}"); + } + + #[test] + fn c_require_uuid_reports_missing_and_bad_values() { + let err = require_uuid(&json!({}), "id").unwrap_err(); + assert_eq!(err.message, "Missing required parameter: id"); + let err = require_uuid(&json!({ "id": null }), "id").unwrap_err(); + assert_eq!(err.message, "Missing required parameter: id"); + let err = require_uuid(&json!({ "id": "nope" }), "id").unwrap_err(); + assert!(err.message.starts_with("Invalid UUID"), "{err}"); + let id = "a1b2c3d4-e5f6-7890-abcd-ef1234567890"; + assert_eq!( + require_uuid(&json!({ "id": id }), "id") + .unwrap() + .into_inner(), + uuid::Uuid::parse_str(id).unwrap() + ); + } + + #[test] + fn c_elide_never_splits_a_character() { + let text = "—".repeat(40); // 120 bytes, 3 bytes per character + let out = elide(&text, 10); + assert_eq!(out, "———..."); + assert!(elide("short", 64).ends_with('t')); + assert_eq!(elide("short", 64), "short"); + assert_eq!(elide("", 4), ""); + } + + #[test] + fn c_json_type_names() { + assert_eq!(json_type_name(&json!(null)), "null"); + assert_eq!(json_type_name(&json!(true)), "boolean"); + assert_eq!(json_type_name(&json!(1)), "number"); + assert_eq!(json_type_name(&json!("s")), "string"); + assert_eq!(json_type_name(&json!([])), "array"); + assert_eq!(json_type_name(&json!({})), "object"); + } + + // ═══════════════════════════════════════════════════════════════ + // D. parse_store_input + // ═══════════════════════════════════════════════════════════════ + + const NS: &str = "work"; + + #[test] + fn d1_happy_path_populates_every_field() { + let parent = "a1b2c3d4-e5f6-7890-abcd-ef1234567890"; + let older = "b2c3d4e5-f6a7-8901-bcde-f12345678901"; + let item = json!({ + "summary": "Storage layer rewritten", + "fullText": "We replaced the four-file engine with a per-namespace layout.", + "tags": ["type/decision", "tech/rust"], + "entities": ["recalld"], + "topics": ["storage"], + "emotions": ["confident"], + "namespace": "engineering", + "parentId": parent, + "supersedes": older, + }); + let input = parse_store_input(&item, NS).unwrap(); + assert_eq!(input.summary, "Storage layer rewritten"); + assert_eq!( + input.full_text.as_deref(), + Some("We replaced the four-file engine with a per-namespace layout.") + ); + assert_eq!(input.tags, vec!["type/decision", "tech/rust"]); + assert_eq!(input.entities, vec!["recalld"]); + assert_eq!(input.topics, vec!["storage"]); + assert_eq!(input.emotions, vec!["confident"]); + assert_eq!(input.namespace, "engineering"); + assert_eq!( + input.parent_id.map(|id| id.into_inner()), + Some(uuid::Uuid::parse_str(parent).unwrap()) + ); + assert_eq!( + input.supersedes.map(|id| id.into_inner()), + Some(uuid::Uuid::parse_str(older).unwrap()) + ); + assert!(input.embedding.is_none()); + assert!(input.initial_stability.is_none()); + } + + #[test] + fn d2_minimal_input_uses_the_default_namespace_argument() { + let input = parse_store_input(&json!({ "summary": "just this" }), NS).unwrap(); + assert_eq!(input.namespace, NS); + assert!(input.full_text.is_none()); + assert!(input.tags.is_empty()); + assert!(input.entities.is_empty()); + assert!(input.topics.is_empty()); + assert!(input.emotions.is_empty()); + assert!(input.parent_id.is_none()); + assert!(input.supersedes.is_none()); + } + + #[test] + fn d2b_explicit_nulls_behave_exactly_like_absent_keys() { + // THE highest-risk compatibility detail: SDKs that serialize + // `None` as `null` must keep working. + let item = json!({ + "summary": "just this", + "fullText": null, + "tags": null, + "entities": null, + "topics": null, + "emotions": null, + "namespace": null, + "parentId": null, + "supersedes": null, + }); + let from_nulls = parse_store_input(&item, NS).unwrap(); + let from_absent = parse_store_input(&json!({ "summary": "just this" }), NS).unwrap(); + assert_eq!(from_nulls.namespace, from_absent.namespace); + assert_eq!(from_nulls.namespace, NS); + assert_eq!(from_nulls.full_text, from_absent.full_text); + assert_eq!(from_nulls.tags, from_absent.tags); + assert_eq!(from_nulls.entities, from_absent.entities); + assert_eq!(from_nulls.topics, from_absent.topics); + assert_eq!(from_nulls.emotions, from_absent.emotions); + assert_eq!(from_nulls.parent_id, from_absent.parent_id); + assert_eq!(from_nulls.supersedes, from_absent.supersedes); + } + + #[test] + fn d3_wrong_typed_namespace_fails_instead_of_falling_back() { + // Reproduces "namespace silently wrong": the old code took the + // default here and stored the memory in the wrong partition. + let err = + parse_store_input(&json!({ "summary": "s", "namespace": ["a"] }), NS).unwrap_err(); + assert_eq!(err.field, "namespace"); + assert_eq!( + err.message, + "Parameter 'namespace' must be a string (got array)" + ); + assert!( + !err.message.contains(NS), + "the default namespace must not appear in the rejection: {err}" + ); + } + + #[test] + fn d4_wrong_typed_tags_fail_instead_of_vanishing() { + // Reproduces "tags empty". + let err = + parse_store_input(&json!({ "summary": "s", "tags": "topic/rust" }), NS).unwrap_err(); + assert_eq!(err.field, "tags"); + assert!(err.message.contains("must be an array of strings"), "{err}"); + } + + #[test] + fn d5_unparseable_link_ids_fail_instead_of_being_dropped() { + // Reproduces "supersedes ignored": previously the memory was + // stored with no link at all. + for field in ["supersedes", "parentId"] { + let item = json!({ "summary": "s", field: "not-a-uuid" }); + let err = parse_store_input(&item, NS).unwrap_err(); + assert_eq!(err.field, field); + assert_eq!( + err.message, + format!("Invalid UUID in parameter '{field}': \"not-a-uuid\"") + ); + } + } + + #[test] + fn d6_markup_in_full_text_is_rejected_with_the_full_message() { + let item = json!({ "summary": "Design notes", "fullText": report_artifact() }); + let err = parse_store_input(&item, NS).unwrap_err(); + assert_eq!(err.field, "fullText"); + let offset = report_artifact().find("").unwrap(); + assert_eq!( + err.message, + format!( + "Parameter 'fullText' contains literal tool-call markup \ + (\"\" at byte offset {offset}). {MARKUP_ADVICE}" + ) + ); + assert!(err.message.contains("nothing was stored"), "{err}"); + } + + #[test] + fn d6b_markup_in_summary_is_rejected() { + let item = json!({ "summary": "notes\n[]" }); + let err = parse_store_input(&item, NS).unwrap_err(); + assert_eq!(err.field, "summary"); + assert!( + err.message + .starts_with("Parameter 'summary' contains literal tool-call markup") + ); + } + + #[test] + fn d6c_legitimate_technical_prose_still_stores() { + let item = json!({ + "summary": "C# doc comments use elements", + "fullText": "/// Returns the widget count.\n\ + /// int\nThe DocBook timeout \ + element is unrelated.", + }); + assert!(parse_store_input(&item, NS).is_ok()); + } + + #[test] + fn d7_failure_precedence_is_deterministic() { + // Broken in five ways at once; peel them off one at a time and + // the reported field walks the documented order. + let mut item = json!({ + "fullText": ["x"], + "tags": "a,b", + "namespace": 7, + "supersedes": "mem-1", + }); + // summary missing -> summary wins over everything else. + assert_eq!(parse_store_input(&item, NS).unwrap_err().field, "summary"); + + item["summary"] = json!("ok"); + assert_eq!(parse_store_input(&item, NS).unwrap_err().field, "fullText"); + + item["fullText"] = json!("fine"); + assert_eq!(parse_store_input(&item, NS).unwrap_err().field, "tags"); + + item["tags"] = json!([]); + assert_eq!(parse_store_input(&item, NS).unwrap_err().field, "namespace"); + + item["namespace"] = json!("work"); + assert_eq!( + parse_store_input(&item, NS).unwrap_err().field, + "supersedes" + ); + + item["supersedes"] = json!(null); + assert!(parse_store_input(&item, NS).is_ok()); + } + + #[test] + fn d7b_markup_is_checked_before_the_length_limit() { + // A 3000-byte summary that also carries markup reports the + // markup: the type/markup boundary runs before the shared + // content limits. + let summary = format!("{}\n[]", "a".repeat(3000)); + let err = parse_store_input(&json!({ "summary": summary }), NS).unwrap_err(); + assert!(err.message.contains("tool-call markup"), "{err}"); + } + + #[test] + fn d8_content_limits_still_fire_with_byte_based_messages() { + // summary over the byte limit + let err = parse_store_input(&json!({ "summary": "a".repeat(2001) }), NS).unwrap_err(); + assert_eq!(err.field, "summary"); + assert!(err.message.contains("2001 bytes"), "{err}"); + assert!(err.message.contains("2000-byte limit"), "{err}"); + assert!( + !err.message.contains("maximum length of 2000 characters"), + "the old character-based wording must be gone: {err}" + ); + + // 65 tags + let tags: Vec = (0..65).map(|i| format!("t{i}")).collect(); + let err = parse_store_input(&json!({ "summary": "s", "tags": tags }), NS).unwrap_err(); + assert_eq!(err.field, "tags"); + assert_eq!( + err.message, + "tags has 65 items, which exceeds the maximum of 64" + ); + + // 33 entities + let entities: Vec = (0..33).map(|i| format!("e{i}")).collect(); + let err = + parse_store_input(&json!({ "summary": "s", "entities": entities }), NS).unwrap_err(); + assert_eq!(err.field, "entities"); + assert!(err.message.contains("exceeds the maximum of 32"), "{err}"); + + // fullText over 1 MiB + let big = "a".repeat(1_048_577); + let err = parse_store_input(&json!({ "summary": "s", "fullText": big }), NS).unwrap_err(); + assert_eq!(err.field, "fullText"); + assert!(err.message.contains("1048576-byte limit"), "{err}"); + + // empty summary + let err = parse_store_input(&json!({ "summary": "" }), NS).unwrap_err(); + assert_eq!(err.field, "summary"); + assert_eq!(err.message, "summary must not be empty"); + } + + #[test] + fn d9_non_object_items_are_rejected_by_shape() { + for value in [ + json!("hello"), + json!([]), + json!(3), + json!(null), + json!(true), + ] { + let err = parse_store_input(&value, NS).unwrap_err(); + assert_eq!(err.field, ""); + assert!( + err.message.starts_with("Item must be an object (got "), + "{err}" + ); + } + assert_eq!( + parse_store_input(&json!("hello"), NS).unwrap_err().message, + "Item must be an object (got string)" + ); + } + + #[test] + fn d10_batch_items_name_their_index_when_not_an_object() { + let err = parse_store_item(&json!("hello"), 3, NS).unwrap_err(); + assert_eq!(err.field, ""); + assert_eq!( + err.message, + "Item at index 3 must be an object (got string)" + ); + + // Everything else is byte-identical to the single-call path. + let item = json!({ "summary": "s", "tags": "a,b" }); + assert_eq!( + parse_store_item(&item, 0, NS).unwrap_err(), + parse_store_input(&item, NS).unwrap_err() + ); + assert_eq!( + parse_store_item(&json!({ "summary": "s" }), 0, NS) + .unwrap() + .namespace, + NS + ); + } + + #[test] + fn d11_error_display_is_exactly_the_message() { + let err = parse_store_input(&json!({}), NS).unwrap_err(); + assert_eq!(err.to_string(), err.message); + } +} diff --git a/src/mcp/mod.rs b/src/mcp/mod.rs index 4d9bdbc..1b080ac 100644 --- a/src/mcp/mod.rs +++ b/src/mcp/mod.rs @@ -9,6 +9,7 @@ pub mod protocol; pub mod server; pub mod transport; +pub mod args; pub mod bridge; pub mod bridge_adapters; pub mod resources; diff --git a/src/mcp/tools.rs b/src/mcp/tools.rs index 995ee39..a72c68b 100644 --- a/src/mcp/tools.rs +++ b/src/mcp/tools.rs @@ -6,13 +6,14 @@ use serde_json::json; +use crate::mcp::args; use crate::mcp::bridge::McpBridge; use crate::mcp::protocol::{ToolAnnotations, ToolCallResult, ToolInfo}; use crate::model::constants::{ FULL_TEXT_MAX_BYTES, MAX_BATCH_MEMORIES, MAX_EMOTIONS, MAX_ENTITIES, MAX_TAGS, MAX_TOPICS, SUMMARY_MAX_BYTES, }; -use crate::model::validation::{MemoryInputRef, validate_memory_input, validate_namespace_name}; +use crate::model::validation::validate_namespace_name; /// Schema description for the `summary` field. /// @@ -148,74 +149,24 @@ fn store_memory_def() -> ToolInfo { } async fn handle_store_memory(bridge: &McpBridge, arguments: serde_json::Value) -> ToolCallResult { - let summary = match arguments.get("summary").and_then(|v| v.as_str()) { - Some(s) => s.to_string(), - None => return ToolCallResult::error("Missing required parameter: summary"), - }; - - let full_text = arguments - .get("fullText") - .and_then(|v| v.as_str()) - .map(String::from); - - let tags: Vec = arguments - .get("tags") - .and_then(|v| serde_json::from_value(v.clone()).ok()) - .unwrap_or_default(); - let entities: Vec = arguments - .get("entities") - .and_then(|v| serde_json::from_value(v.clone()).ok()) - .unwrap_or_default(); - let topics: Vec = arguments - .get("topics") - .and_then(|v| serde_json::from_value(v.clone()).ok()) - .unwrap_or_default(); - let emotions: Vec = arguments - .get("emotions") - .and_then(|v| serde_json::from_value(v.clone()).ok()) - .unwrap_or_default(); - - // Content limits — shared with the HTTP API so the two paths cannot - // disagree about what fits. - if let Err(e) = validate_memory_input(MemoryInputRef { - summary: &summary, - full_text: full_text.as_deref(), - tags: &tags, - entities: &entities, - topics: &topics, - emotions: &emotions, - }) { - return ToolCallResult::error(e.to_string()); - } - - let namespace = arguments - .get("namespace") - .and_then(|v| v.as_str()) - .map(String::from) - .unwrap_or_else(|| bridge.default_namespace().to_string()); - let parent_id = arguments - .get("parentId") - .and_then(|v| v.as_str()) - .and_then(|s| uuid::Uuid::parse_str(s).ok()) - .map(crate::model::MemoryId::from_uuid); - let supersedes = arguments - .get("supersedes") - .and_then(|v| v.as_str()) - .and_then(|s| uuid::Uuid::parse_str(s).ok()) - .map(crate::model::MemoryId::from_uuid); - - let input = crate::mcp::bridge::StoreInput { - summary, - full_text, - tags, - entities, - topics, - emotions, - namespace, - embedding: None, - initial_stability: None, - parent_id, - supersedes, + // Extraction, type strictness, markup detection, and the shared + // content limits all live in `args`, so this handler and the batch + // one below cannot drift apart. + let input = match args::parse_store_input(&arguments, bridge.default_namespace()) { + Ok(input) => input, + Err(e) => { + // Logged, never silently defaulted: a client whose + // serializer mangles arguments shows up in the operator's + // logs on the first call, not after ten bad writes. The + // offending text body is deliberately not logged. + tracing::warn!( + tool = "store_memory", + field = %e.field, + error = %e.message, + "Rejected malformed MCP tool arguments" + ); + return ToolCallResult::error(e.to_string()); + } }; match bridge.storage.store_memory(input).await { @@ -336,84 +287,27 @@ async fn handle_store_memories(bridge: &McpBridge, arguments: serde_json::Value) let mut results: Vec = Vec::with_capacity(memories_arr.len()); for (index, item) in memories_arr.iter().enumerate() { - let summary = match item.get("summary").and_then(|v| v.as_str()) { - Some(s) => s.to_string(), - None => { + let input = match args::parse_store_item(item, index, bridge.default_namespace()) { + Ok(input) => input, + Err(e) => { + tracing::warn!( + tool = "store_memories", + index, + field = %e.field, + error = %e.message, + "Rejected malformed MCP tool arguments" + ); + // `field` is additive; existing consumers already + // handle a per-item `error`. results.push(json!({ "index": index, - "error": "Missing required parameter: summary" + "error": e.to_string(), + "field": e.field, })); continue; } }; - let full_text = item - .get("fullText") - .and_then(|v| v.as_str()) - .map(String::from); - - let tags: Vec = item - .get("tags") - .and_then(|v| serde_json::from_value(v.clone()).ok()) - .unwrap_or_default(); - let entities: Vec = item - .get("entities") - .and_then(|v| serde_json::from_value(v.clone()).ok()) - .unwrap_or_default(); - let topics: Vec = item - .get("topics") - .and_then(|v| serde_json::from_value(v.clone()).ok()) - .unwrap_or_default(); - let emotions: Vec = item - .get("emotions") - .and_then(|v| serde_json::from_value(v.clone()).ok()) - .unwrap_or_default(); - - if let Err(e) = validate_memory_input(MemoryInputRef { - summary: &summary, - full_text: full_text.as_deref(), - tags: &tags, - entities: &entities, - topics: &topics, - emotions: &emotions, - }) { - results.push(json!({ - "index": index, - "error": e.to_string() - })); - continue; - } - - let namespace = item - .get("namespace") - .and_then(|v| v.as_str()) - .map(String::from) - .unwrap_or_else(|| bridge.default_namespace().to_string()); - let parent_id = item - .get("parentId") - .and_then(|v| v.as_str()) - .and_then(|s| uuid::Uuid::parse_str(s).ok()) - .map(crate::model::MemoryId::from_uuid); - let supersedes = item - .get("supersedes") - .and_then(|v| v.as_str()) - .and_then(|s| uuid::Uuid::parse_str(s).ok()) - .map(crate::model::MemoryId::from_uuid); - - let input = crate::mcp::bridge::StoreInput { - summary, - full_text, - tags, - entities, - topics, - emotions, - namespace, - embedding: None, - initial_stability: None, - parent_id, - supersedes, - }; - match bridge.storage.store_memory(input).await { Ok(stored) => { results.push(json!({ From 0dc46fb3c539b90c4e43656f4e5b807a30b4811a Mon Sep 17 00:00:00 2001 From: Caleb Evans Date: Tue, 11 Aug 2026 01:24:21 -0600 Subject: [PATCH 4/8] feat: serve MCP 2026-07-28 alongside the legacy handshake Closes #2. The issue proposed bumping the protocol version as a backward-compatible first step. It is not. PROTOCOL_VERSION is echoed verbatim into the `initialize` response, so bumping it alone tells a legacy client it is talking to a revision where `initialize` does not exist, `ping` is removed, and resultType/ttlMs/cacheScope are mandatory -- none of which this server emitted. Under 2025-06-18 lifecycle rules such a client should disconnect. The bump is therefore landed together with the stateless path, never before it. The spec explicitly sanctions serving both eras on one endpoint, so the era is decided per message rather than per connection: a `params._meta` carrying io.modelcontextprotocol/protocolVersion is modern, an `initialize` is legacy. The `_meta` key wins over the method name, since a legacy client can send `_meta` for progressToken but never that reverse-DNS key. - Split McpServer into a stateless McpDispatcher plus a thin wrapper holding only legacy lifecycle state. Modern HTTP requests take no lock at all. Merely bypassing the session map would have funnelled every concurrent modern request through one mutex, which is worse than today's per-session locks and would not have fixed the scaling complaint the issue actually raises. - Add server/discover, which the revision requires and the issue does not mention. It is also the era probe, so it and `ping` are answered without `_meta` -- demanding a protocol version in order to discover which versions are supported is circular, and any error there would push a modern client into the legacy fallback. - Add the required resultType to every result, and SEP-2549 caching hints. The catalogs are compile-time constants, so they advertise public/1h; resources/read is live instance state, so it advertises private/0. Emitted on both eras: the 2025-06-18 Result type is an open index signature, so the extra keys are schema-legal there. - Validate the SEP-2243 headers against the body, including the base64 sentinel, and answer GET/DELETE with 405 on the modern path. - Fix resources/templates/list, which emitted resource_templates instead of resourceTemplates. That was broken against every MCP revision, so no compliant client could ever read the result. - Validate Origin (a spec MUST this server never honoured) and apply a body limit to /mcp, which sits outside the tower stack and had none. Loopback is allowed on any scheme or port, since DNS rebinding needs an attacker-controlled hostname. - Fix three session-map defects: entries never expired, a repeat initialize orphaned the previous entry, and a failed initialize still inserted an unusable one. Unknown sessions now answer 404 with a JSON-RPC body saying to re-initialize, so expiry is recoverable. - Stop discarding the real request id on parse and session errors. Uncertainties are preserved as code comments rather than silently resolved: 2025-11-25 is deliberately absent from the legacy list pending a wire diff, the base64 alphabet needs confirming against SEP-2243, and notification header requirements are undefined by the revision. Breaking: mcp_router takes a config argument; /mcp enforces a 10 MiB body limit where it had none; non-loopback browser origins get 403; legacy HTTP sessions expire after 30 minutes idle. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01JNhUehmChiQKJUjv3QPnkH --- Cargo.lock | 1 + Cargo.toml | 4 +- README.md | 2 +- docs/architecture.md | 66 ++- docs/guide.md | 5 +- docs/mcp.md | 32 +- src/config/loader.rs | 18 + src/config/types.rs | 9 + src/main.rs | 10 +- src/mcp/http_transport.rs | 1158 ++++++++++++++++++++++++++++++++++--- src/mcp/mod.rs | 8 +- src/mcp/protocol.rs | 771 +++++++++++++++++++++++- src/mcp/server.rs | 900 +++++++++++++++++++++++++--- src/mcp/transport.rs | 173 +++++- 14 files changed, 2956 insertions(+), 201 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 2fc3069..15eaec4 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2393,6 +2393,7 @@ dependencies = [ "aws-config", "aws-sdk-bedrockruntime", "axum", + "base64", "bincode", "bytemuck", "chrono", diff --git a/Cargo.toml b/Cargo.toml index 1da9e2f..bba2066 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -28,12 +28,14 @@ bench = [] # ── Async runtime & networking ──────────────────────────────────── tokio = { version = "1", features = ["full"] } axum = { version = "0.8", features = ["macros"] } -tower = "0.5" +tower = { version = "0.5", features = ["util"] } tower-http = { version = "0.6", features = ["trace", "cors", "timeout"] } # ── Serialization ───────────────────────────────────────────────── serde = { version = "1", features = ["derive"] } serde_json = "1" +# Required by the MCP 2026-07-28 SEP-2243 `=?base64?...?=` header sentinel. +base64 = "0.22" # ── Configuration ───────────────────────────────────────────────── toml = "0.8" diff --git a/README.md b/README.md index ec88268..771f356 100644 --- a/README.md +++ b/README.md @@ -179,7 +179,7 @@ See [docs/benchmark.md](docs/benchmark.md) for full methodology, per-category br recalld mcp ``` -**HTTP API** -- Runs a standalone HTTP server (default `127.0.0.1:7680`). Also exposes an MCP endpoint at `/mcp` using the streamable HTTP transport, so MCP clients can connect via URL. +**HTTP API** -- Runs a standalone HTTP server (default `127.0.0.1:7680`). Also exposes an MCP endpoint at `/mcp`, so MCP clients can connect via URL. Both MCP transports are dual-era: stateless MCP 2026-07-28 and the legacy `initialize`/`Mcp-Session-Id` handshake (2025-06-18 and earlier) are served on the same endpoint, chosen per message. ```sh recalld serve diff --git a/docs/architecture.md b/docs/architecture.md index 767c507..a3a6884 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -674,16 +674,78 @@ recalld exposes the same core functionality through three transport layers: Two transports expose the same 9 MCP tools (`store_memory`, `store_memories`, `recall_memories`, `get_memory`, `reinforce_memory`, `forget_memory`, `find_similar_memories`, `create_namespace`, `list_memories`) and memory resources: +Both transports are **dual-era**: MCP 2026-07-28 (stateless) and 2025-06-18 +(legacy `initialize` handshake) are served on the same endpoint and process. +The era is decided **per message**, not per connection: + +- A message whose `params._meta` carries + `io.modelcontextprotocol/protocolVersion` is served on the stateless modern + path. The key wins over the method name, so an `initialize` carrying modern + `_meta` gets `-32601` — that revision removed `initialize`. +- An `initialize` or `notifications/initialized` message takes the legacy path. +- `server/discover` and `ping` are answered on either path, without `_meta`, so + a modern client's era probe (`server/discover` first) works before it knows + which era the server speaks. +- Anything else with neither marker is a legacy message: over HTTP it needs an + `Mcp-Session-Id`, over stdio it needs a prior `initialize`. Without either it + is rejected with `-32602` naming both remedies. + +Ten methods are dispatched: `server/discover`, `initialize` (legacy only), +`ping` (legacy only), `tools/list`, `tools/call`, `resources/list`, +`resources/templates/list`, `resources/read`, plus the +`notifications/initialized` and `notifications/cancelled` notifications. + +Every result carries `resultType: "complete"` and +`_meta["io.modelcontextprotocol/serverInfo"]`. The three catalogs and +`server/discover` also carry SEP-2549 caching hints (`ttlMs` 3600000, +`cacheScope: "public"` — recalld's catalogs are compile-time constants); +`resources/read` carries `ttlMs: 0`, `cacheScope: "private"` because its +contents are live instance state. These are top-level fields, siblings of +`tools`/`resources`/`contents`, never nested in `_meta`. They are emitted on +both eras: the 2025-06-18 `Result` type is an open index signature, so the +extra keys are schema-legal there, not merely tolerated. + +`initialize` never echoes `2026-07-28`. It echoes the version the client +asked for when we recognise it (`2025-06-18`, `2025-03-26`, `2024-11-05`), +otherwise `2025-06-18`. Claiming the modern version in an `initialize` +response would advertise a revision in which `initialize` does not exist. + **Stdio** (`recalld mcp`): - Runs as a subprocess of an AI agent (Claude Code, Cursor, etc.). - Communicates via stdin/stdout using newline-delimited JSON-RPC 2.0. +- A process is **not** a session for a modern client: the single `McpServer` + held for the process lifetime carries legacy lifecycle state only, and + modern requests never consult it. **HTTP** (`recalld serve`, endpoint `/mcp`): - Runs alongside the REST API on the same port. -- Implements the MCP streamable HTTP transport (spec 2025-03-26). -- Session management via `Mcp-Session-Id` header. - Responds with `application/json` for requests, `202 Accepted` for notifications. - Recommended transport for Docker containers and remote servers. +- **Modern requests are stateless**: no session is minted, any inbound + `Mcp-Session-Id` is ignored, and dispatch takes no lock — concurrent clients + do not serialize behind one mutex. Clients must send the SEP-2243 headers + `Mcp-Method`, `MCP-Protocol-Version` and (on `tools/call`, `prompts/get`, + `resources/read`) `Mcp-Name`; each is validated against the body and a + mismatch is `400` + `-32020`. Values may use the `=?base64?…?=` sentinel. + Header validation is skipped on notification POSTs, which the revision + leaves undefined. +- **`Mcp-Session-Id` is legacy-only.** A session is created only by a + *successful* `initialize`; a repeat `initialize` bearing a live session id + reuses it rather than orphaning it. Sessions expire after 30 minutes idle + (swept lazily, at most once a minute) and are capped at 1024 (`503` + + `-32603` beyond that). An unknown or expired session gets a `404` with a + JSON-RPC body telling the client to `initialize` again. +- **`Origin` is validated** on every request (a spec MUST, and a real DNS + rebinding exposure for a loopback-bound server): absent → allowed + (non-browser client); `null` → `403`; any loopback host on any scheme or + port → allowed; otherwise it must match `server.mcp_allowed_origins`. +- **10 MiB body limit** applied on the MCP router itself. `/mcp` is merged in + outside the REST tower stack, so it inherits neither that stack's body limit + (previously unlimited here) nor its 30s timeout — the timeout is omitted + deliberately, since tool calls may run long. +- `GET /mcp`, and `DELETE` without a session id, return `405` with + `Allow: POST, DELETE`. The 2026-07-28 transport removed the GET stream + endpoint; recalld never had one. ### Daemon (Unix socket) diff --git a/docs/guide.md b/docs/guide.md index d1b0335..c509eb9 100644 --- a/docs/guide.md +++ b/docs/guide.md @@ -188,9 +188,10 @@ HTTP API server settings. | `bind_address` | string | `"127.0.0.1"` | IP address to bind to. | | `port` | u16 | `7680` | TCP port to listen on. | | `request_timeout_ms` | u64 | `30000` | Maximum request time in milliseconds before abort. | -| `max_body_bytes` | usize | `10485760` | Maximum request body size in bytes (10 MB). | +| `max_body_bytes` | usize | `10485760` | Maximum request body size in bytes (10 MB). Also applied to `/mcp`. | +| `mcp_allowed_origins` | array of string | `[]` | Extra browser origins allowed to call `/mcp`, beyond loopback. Matched case-insensitively and exactly, e.g. `["https://app.example.com"]`. Requests with no `Origin` header are always allowed; any loopback origin is allowed on any scheme or port; `Origin: null` is always rejected. | -Env vars: `RECALLD_SERVER_BIND_ADDRESS`, `RECALLD_SERVER_PORT`, `RECALLD_SERVER_REQUEST_TIMEOUT_MS`, `RECALLD_SERVER_MAX_BODY_BYTES` +Env vars: `RECALLD_SERVER_BIND_ADDRESS`, `RECALLD_SERVER_PORT`, `RECALLD_SERVER_REQUEST_TIMEOUT_MS`, `RECALLD_SERVER_MAX_BODY_BYTES`, `RECALLD_SERVER_MCP_ALLOWED_ORIGINS` (comma-separated) ### `[storage]` diff --git a/docs/mcp.md b/docs/mcp.md index da78e68..4674ba3 100644 --- a/docs/mcp.md +++ b/docs/mcp.md @@ -87,7 +87,37 @@ Optional: add `--log-level ` to `args` for debug logging (logs go to stde } ``` -The HTTP endpoint is available whenever `recalld serve` is running. It implements the MCP streamable HTTP transport with session management via the `Mcp-Session-Id` header. +The HTTP endpoint is available whenever `recalld serve` is running. It is a **dual-era** MCP endpoint: MCP 2026-07-28 (stateless) and MCP 2025-06-18 and earlier (the `initialize` handshake) are served on the same URL, and the era is chosen per message. + +**Modern clients (MCP 2026-07-28)** send no `initialize` and get no session. Every request carries its own identity in `params._meta`: + +```json +"_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientInfo": { "name": "ExampleClient", "version": "1.0.0" }, + "io.modelcontextprotocol/clientCapabilities": {} +} +``` + +`protocolVersion` and `clientCapabilities` are required (`-32602` naming the missing key if absent; `-32022` with `data.supported` if the version is one we do not serve). Every request must also mirror three SEP-2243 headers, which are validated against the body — a mismatch is HTTP 400 with JSON-RPC `-32020`: + +| Header | Value | Required for | +|---|---|---| +| `Mcp-Method` | the `method` field | every request | +| `MCP-Protocol-Version` | `_meta`'s protocol version | every request | +| `Mcp-Name` | `params.name`, or `params.uri` for `resources/read` | `tools/call`, `prompts/get`, `resources/read` | + +Values that are not safe plain ASCII may use the `=?base64?{value}?=` sentinel, which is decoded before comparison. Notification POSTs skip header validation — the revision leaves that case undefined. + +Modern responses never carry `Mcp-Session-Id`, and any session id sent on a modern request is ignored. Every result carries `resultType: "complete"` and `_meta["io.modelcontextprotocol/serverInfo"]`; `server/discover`, `tools/list`, `resources/list` and `resources/templates/list` also carry `ttlMs: 3600000` / `cacheScope: "public"`, and `resources/read` carries `ttlMs: 0` / `cacheScope: "private"`. + +Start with `server/discover`, which is answered on either era and needs no `_meta`. It returns `supportedVersions`, `capabilities` and `instructions`. + +**Legacy clients** work exactly as before: send `initialize`, get an `Mcp-Session-Id`, send it on every subsequent request. `initialize` echoes the version you asked for (`2025-06-18`, `2025-03-26` or `2024-11-05`; anything else falls back to `2025-06-18`) and never `2026-07-28`, since that revision has no `initialize`. Three behaviours are new: sessions expire after 30 minutes idle (the resulting 404 tells you to `initialize` again), at most 1024 may be live at once, and a *failed* `initialize` no longer leaves an unusable session behind. + +**Origin validation.** `/mcp` rejects browser requests from unexpected origins with HTTP 403 (DNS-rebinding protection for a loopback-bound server). Requests with no `Origin` header — curl, SDKs, anything non-browser — are always allowed, as is any loopback origin on any scheme or port. `Origin: null` (a `file://` page or sandboxed iframe) is rejected. To allow a real browser origin, list it in `server.mcp_allowed_origins` (or `RECALLD_SERVER_MCP_ALLOWED_ORIGINS`, comma-separated). + +Request bodies over `server.max_body_bytes` (10 MB by default) are rejected with HTTP 413; `/mcp` previously had no body limit at all. `GET /mcp`, and `DELETE` without a session id, return 405. Configure per-project namespace defaults in a `.recalld.toml` file (see above), not CLI flags. diff --git a/src/config/loader.rs b/src/config/loader.rs index cae0799..22a317c 100644 --- a/src/config/loader.rs +++ b/src/config/loader.rs @@ -315,6 +315,20 @@ fn apply_env_overrides(config: &mut RecalldConfig) -> std::result::Result<(), Ve }; } + // Helper for Vec fields (comma-separated; empty clears). + macro_rules! env_override_string_list { + ($var:expr, $field:expr) => { + if let Ok(val) = std::env::var($var) { + $field = val + .split(',') + .map(str::trim) + .filter(|s| !s.is_empty()) + .map(str::to_string) + .collect(); + } + }; + } + // Helper for Option fields. macro_rules! env_override_opt_string { ($var:expr, $field:expr) => { @@ -337,6 +351,10 @@ fn apply_env_overrides(config: &mut RecalldConfig) -> std::result::Result<(), Ve config.server.max_body_bytes, usize ); + env_override_string_list!( + "RECALLD_SERVER_MCP_ALLOWED_ORIGINS", + config.server.mcp_allowed_origins + ); // --- Storage --- env_override_string!("RECALLD_STORAGE_DATA_DIR", config.storage.data_dir); diff --git a/src/config/types.rs b/src/config/types.rs index 10b5a7e..032c293 100644 --- a/src/config/types.rs +++ b/src/config/types.rs @@ -21,6 +21,14 @@ pub struct ServerConfig { /// Maximum request body size in bytes (prevents OOM from large payloads). pub max_body_bytes: usize, + + /// Extra `Origin` values allowed to call `/mcp`, beyond loopback. + /// + /// Matched case-insensitively and exactly (scheme, host and port all + /// significant), e.g. `"https://app.example.com"`. Requests with no + /// `Origin` header are always allowed; loopback origins are always + /// allowed on any scheme or port. + pub mcp_allowed_origins: Vec, } impl Default for ServerConfig { @@ -30,6 +38,7 @@ impl Default for ServerConfig { port: 7680, request_timeout_ms: 30_000, max_body_bytes: 10 * 1024 * 1024, // 10 MB + mcp_allowed_origins: Vec::new(), } } } diff --git a/src/main.rs b/src/main.rs index 12a0ea3..8fcdac9 100644 --- a/src/main.rs +++ b/src/main.rs @@ -406,6 +406,14 @@ async fn run_serve( ..recalld::api::ApiConfig::from_server_config(&config.server) }; + // `/mcp` is merged in outside the REST tower stack, so it carries its + // own body limit and Origin policy rather than inheriting them. + let mcp_http_config = recalld::mcp::McpHttpConfig { + max_body_bytes: config.server.max_body_bytes, + allowed_origins: config.server.mcp_allowed_origins.clone(), + ..Default::default() + }; + let system = match Recalld::new(config).await { Ok(s) => s, Err(e) => { @@ -456,7 +464,7 @@ async fn run_serve( // Build MCP bridge and router for the /mcp endpoint. let mcp_bridge = create_direct_mcp_bridge(&system, default_namespace); let mcp_handler: Arc = Arc::new(mcp_bridge); - let mcp_router = recalld::mcp::mcp_router(mcp_handler); + let mcp_router = recalld::mcp::mcp_router(mcp_handler, mcp_http_config); // Start the API server (blocks until shutdown signal). match recalld::api::serve(app_state, api_config, Some(mcp_router)).await { diff --git a/src/mcp/http_transport.rs b/src/mcp/http_transport.rs index 5ea80f4..9d92ec9 100644 --- a/src/mcp/http_transport.rs +++ b/src/mcp/http_transport.rs @@ -1,145 +1,620 @@ //! Streamable HTTP transport for MCP. //! -//! Implements the MCP Streamable HTTP transport as an axum router -//! mountable alongside the REST API. Each session gets its own -//! `McpServer` instance; the underlying `McpHandler` is shared. +//! Implements the MCP HTTP transport as an axum router mountable +//! alongside the REST API, serving BOTH protocol eras on `/mcp`: +//! +//! * **Modern** (MCP 2026-07-28): stateless. Identity and protocol version +//! travel in `params._meta` and are mirrored into the SEP-2243 headers. +//! No session is minted, any inbound `Mcp-Session-Id` is ignored, and +//! requests are dispatched with NO lock -- so concurrent clients do not +//! serialize behind one mutex. +//! * **Legacy** (MCP 2025-06-18 and earlier): `initialize` mints an +//! `Mcp-Session-Id`, and each session gets its own [`McpServer`] over the +//! shared [`McpDispatcher`]. +//! +//! The router carries its own state and intentionally applies none of the +//! REST API's tower stack -- in particular no `TimeoutLayer`, since MCP tool +//! calls may legitimately outrun the REST request timeout. The body limit +//! and Origin validation are therefore applied here rather than inherited. use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; use axum::{ Json, Router, - extract::State, - http::{HeaderMap, HeaderValue, StatusCode}, + extract::{DefaultBodyLimit, State}, + http::{HeaderMap, HeaderValue, StatusCode, header}, response::{IntoResponse, Response}, routing::post, }; +use base64::Engine as _; use dashmap::DashMap; use tokio::sync::Mutex; use uuid::Uuid; use crate::mcp::protocol::*; -use crate::mcp::server::{McpHandler, McpServer}; +use crate::mcp::server::{McpDispatcher, McpHandler, McpServer, RequestContext}; + +// ── Headers ───────────────────────────────────────────────────────── +// Header NAMES are case-insensitive (hyper lowercases them); header +// VALUES are case-sensitive. + +const HDR_SESSION_ID: &str = "mcp-session-id"; +const HDR_MCP_METHOD: &str = "mcp-method"; +const HDR_MCP_NAME: &str = "mcp-name"; +const HDR_MCP_PROTOCOL_VER: &str = "mcp-protocol-version"; + +/// SEP-2243 sentinel wrapping a header value that is not safe plain ASCII. +/// The markers are lowercase and exact. +const B64_PREFIX: &str = "=?base64?"; +const B64_SUFFIX: &str = "?="; + +/// How often the opportunistic session sweep may actually run. +const SWEEP_THROTTLE_MS: u64 = 60_000; + +/// Methods allowed on `/mcp`, echoed in `Allow` on a 405. +const ALLOW_METHODS: &str = "POST, DELETE"; + +// ── Configuration ─────────────────────────────────────────────────── + +/// Tuning knobs for the MCP HTTP transport. +#[derive(Debug, Clone)] +pub struct McpHttpConfig { + /// Maximum request body size in bytes. Mirrors the REST API default. + pub max_body_bytes: usize, + /// Extra origins allowed beyond loopback, matched case-insensitively + /// and exactly (scheme, host and port all significant). + pub allowed_origins: Vec, + /// How long a legacy session may sit idle before it is swept. + pub session_idle_timeout: Duration, + /// Hard cap on concurrent legacy sessions. + pub max_sessions: usize, +} + +impl Default for McpHttpConfig { + fn default() -> Self { + Self { + max_body_bytes: 10 * 1024 * 1024, + allowed_origins: Vec::new(), + session_idle_timeout: Duration::from_secs(30 * 60), + max_sessions: 1024, + } + } +} + +// ── Origin policy ─────────────────────────────────────────────────── + +/// Origin validation for `/mcp`. +/// +/// The spec makes this a MUST ("Servers MUST validate the Origin header on +/// all incoming connections to prevent DNS rebinding attacks"), and it +/// matters here because recalld binds to loopback by default. +#[derive(Debug)] +struct OriginPolicy { + allowed: Vec, +} + +impl OriginPolicy { + fn new(allowed: &[String]) -> Self { + Self { + allowed: allowed.iter().map(|o| o.to_ascii_lowercase()).collect(), + } + } + + /// Decide whether a request bearing `origin` may proceed. + fn allows(&self, origin: Option<&str>) -> bool { + let Some(origin) = origin else { + // ABSENT -> allow. Browsers always send Origin on POST; a + // missing one means a non-browser client (curl, an SDK), which + // is not the DNS-rebinding threat model. + return true; + }; + if origin.eq_ignore_ascii_case("null") { + // JUDGMENT CALL (flagged): `Origin: null` is the opaque origin + // used by `file://` documents and sandboxed iframes. It is not + // attributable, so it is denied. + return false; + } + if is_loopback_origin(origin) { + // Port- and scheme-agnostic on purpose: DNS rebinding needs an + // attacker-controlled HOSTNAME, and a rebound name is never + // literally `localhost`. Matching on the port too would break + // every dev tool for no security gain. + return true; + } + let lower = origin.to_ascii_lowercase(); + self.allowed.contains(&lower) + } +} + +/// Whether an origin's host is a loopback name or address. +fn is_loopback_origin(origin: &str) -> bool { + let rest = origin.split_once("://").map_or(origin, |(_, r)| r); + let host = if let Some(end) = rest.find(']') { + &rest[..=end] + } else { + rest.split(':').next().unwrap_or(rest) + }; + matches!(host, "localhost" | "127.0.0.1" | "::1" | "[::1]") +} -const MCP_SESSION_ID: &str = "mcp-session-id"; +// ── State ─────────────────────────────────────────────────────────── + +/// One legacy session: its server plus the last time it was touched. +struct SessionEntry { + server: Arc>, + last_seen: Arc, +} #[derive(Clone)] struct McpHttpState { - handler: Arc, - sessions: Arc>>>, + dispatcher: Arc, + sessions: Arc>, + last_sweep: Arc, + origins: Arc, + config: Arc, +} + +/// Milliseconds since the Unix epoch. +fn now_ms() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_millis() as u64) + .unwrap_or(0) } +// ── Router ────────────────────────────────────────────────────────── + /// Build a self-contained `Router` for the MCP HTTP transport. /// -/// Handles POST and DELETE at `/mcp`. Carries its own state and -/// intentionally applies no middleware (no timeout, no CORS, no -/// request-ID) so MCP tool calls are not subject to REST API limits. -pub fn mcp_router(handler: Arc) -> Router { +/// Handles POST, DELETE and GET at `/mcp`. Applies a body limit (the REST +/// API's limit is not inherited, because `/mcp` is merged in outside that +/// tower stack) but deliberately no timeout. +pub fn mcp_router(handler: Arc, config: McpHttpConfig) -> Router { + build_router(handler, config).0 +} + +fn build_router(handler: Arc, config: McpHttpConfig) -> (Router, McpHttpState) { let state = McpHttpState { - handler, + dispatcher: Arc::new(McpDispatcher::new(handler)), sessions: Arc::new(DashMap::new()), + last_sweep: Arc::new(AtomicU64::new(0)), + origins: Arc::new(OriginPolicy::new(&config.allowed_origins)), + config: Arc::new(config), }; - Router::new() - .route("/mcp", post(handle_post).delete(handle_delete)) - .with_state(state) + let max_body = state.config.max_body_bytes; + let router = Router::new() + .route( + "/mcp", + post(handle_post).delete(handle_delete).get(handle_get), + ) + .layer(DefaultBodyLimit::max(max_body)) + .with_state(state.clone()); + + (router, state) +} + +// ── Response helpers ──────────────────────────────────────────────── + +fn jsonrpc_error( + status: StatusCode, + id: JsonRpcId, + code: i32, + message: impl Into, + data: Option, +) -> Response { + ( + status, + Json(JsonRpcResponse::error_with_data(id, code, message, data)), + ) + .into_response() +} + +fn method_not_allowed() -> Response { + ( + StatusCode::METHOD_NOT_ALLOWED, + [(header::ALLOW, ALLOW_METHODS)], + ) + .into_response() +} + +fn origin_rejected(origin: &str) -> Response { + tracing::warn!(origin = %origin, "rejected /mcp request: Origin not allowed"); + // Deliberately a plain JSON body, not JSON-RPC: the rejection happens + // before the body is parsed, so there is no request id to answer with. + ( + StatusCode::FORBIDDEN, + Json(serde_json::json!({ "error": "Origin not allowed" })), + ) + .into_response() +} + +fn dispatch_response(response: Option) -> Response { + match response { + Some(resp) => Json(resp).into_response(), + None => StatusCode::ACCEPTED.into_response(), + } } +fn header_str<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> { + headers.get(name).and_then(|v| v.to_str().ok()) +} + +// ── POST /mcp ─────────────────────────────────────────────────────── + async fn handle_post( State(state): State, headers: HeaderMap, body: axum::body::Bytes, ) -> Response { + if let Some(origin) = header_str(&headers, header::ORIGIN.as_str()) + && !state.origins.allows(Some(origin)) + { + return origin_rejected(origin); + } + + maybe_sweep(&state); + let message: JsonRpcMessage = match serde_json::from_slice(&body) { Ok(msg) => msg, Err(e) => { - let err_resp = JsonRpcResponse { - jsonrpc: JSONRPC_VERSION.to_string(), - id: JsonRpcId::Number(0), - result: None, - error: Some(JsonRpcError { - code: PARSE_ERROR, - message: format!("Parse error: {e}"), - data: None, - }), - }; - return (StatusCode::BAD_REQUEST, Json(err_resp)).into_response(); + // No usable id: answer with an explicit null rather than + // fabricating `0` and colliding with a real request. + return jsonrpc_error( + StatusCode::BAD_REQUEST, + JsonRpcId::Null, + PARSE_ERROR, + format!("Parse error: {e}"), + None, + ); } }; - if message.method == "initialize" { - handle_initialize(state, message).await - } else { - handle_session_message(state, headers, message).await + match era_hint(&message) { + EraHint::Modern => handle_modern(state, headers, message).await, + EraHint::LegacyHandshake if message.method == "initialize" => { + handle_legacy_initialize(state, headers, message).await + } + EraHint::LegacyHandshake => handle_legacy_session(state, headers, message).await, + EraHint::Ambiguous => { + if header_str(&headers, HDR_SESSION_ID).is_some() { + handle_legacy_session(state, headers, message).await + } else if is_era_neutral_method(&message.method) { + // `server/discover` and `ping` are answered without a + // handshake and without a session -- the stdio-style era + // probe works over HTTP too. + let response = state + .dispatcher + .dispatch(&RequestContext::legacy(), message) + .await; + dispatch_response(response) + } else { + let id = message.id.clone().unwrap_or(JsonRpcId::Null); + jsonrpc_error( + StatusCode::BAD_REQUEST, + id, + INVALID_PARAMS, + MetaError::MissingField(META_PROTOCOL_VERSION).to_string(), + None, + ) + } + } + } +} + +/// Stateless MCP 2026-07-28 request: no session, no lock. +async fn handle_modern( + state: McpHttpState, + headers: HeaderMap, + message: JsonRpcMessage, +) -> Response { + let id = message.id.clone().unwrap_or(JsonRpcId::Null); + let is_notification = message.id.is_none(); + + let meta = match parse_request_meta(message.params.as_ref()) { + Ok(meta) => Some(meta), + Err(e) if is_notification => { + tracing::debug!(error = %e, "modern notification carried invalid _meta; ignoring"); + None + } + Err(e) => { + return jsonrpc_error( + StatusCode::BAD_REQUEST, + id, + e.code(), + e.to_string(), + e.data(), + ); + } + }; + + if let Some(meta) = meta.as_ref() + && !is_notification + && let Err(detail) = validate_modern_headers(&headers, &message, meta) + { + return jsonrpc_error(StatusCode::BAD_REQUEST, id, HEADER_MISMATCH, detail, None); + } + + if header_str(&headers, HDR_SESSION_ID).is_some() { + // Spec MUST: a modern-era request's session id is ignored, and no + // session id is minted or echoed back. + tracing::debug!("ignoring Mcp-Session-Id on a stateless 2026-07-28 request"); + } + + let ctx = match meta { + Some(meta) => RequestContext::modern(meta), + None => RequestContext::legacy(), + }; + dispatch_response(state.dispatcher.dispatch(&ctx, message).await) +} + +/// SEP-2243 header validation. Returns the error detail on mismatch. +fn validate_modern_headers( + headers: &HeaderMap, + message: &JsonRpcMessage, + meta: &RequestMeta, +) -> Result<(), String> { + // NOTIFICATION POSTs are never routed here: whether these headers are + // required on notifications is explicitly "not defined by this + // revision", so we skip validation for them. Skipping is the + // non-breaking choice; revisit if a later revision defines it. + debug_assert!(message.id.is_some()); + + let raw_method = header_str(headers, HDR_MCP_METHOD) + .ok_or_else(|| "Header mismatch: missing required Mcp-Method header".to_string())?; + let method = decode_header_value(raw_method, "Mcp-Method")?; + if method != message.method { + return Err(format!( + "Header mismatch: Mcp-Method header \"{method}\" does not match body method \"{}\"", + message.method + )); + } + + let raw_version = header_str(headers, HDR_MCP_PROTOCOL_VER).ok_or_else(|| { + "Header mismatch: missing required MCP-Protocol-Version header".to_string() + })?; + let version = decode_header_value(raw_version, "MCP-Protocol-Version")?; + if version != meta.protocol_version { + return Err(format!( + "Header mismatch: MCP-Protocol-Version header \"{version}\" does not match \ + _meta protocol version \"{}\"", + meta.protocol_version + )); + } + + // Mcp-Name is required only for the three designated methods. Present + // but unexpected on any other method is ignored, not an error. + let expected = match message.method.as_str() { + "tools/call" | "prompts/get" => Some(("name", "params.name")), + "resources/read" => Some(("uri", "params.uri")), + _ => None, + }; + if let Some((field, label)) = expected { + let want = message + .params + .as_ref() + .and_then(|p| p.get(field)) + .and_then(|v| v.as_str()) + .unwrap_or_default(); + let raw_name = header_str(headers, HDR_MCP_NAME) + .ok_or_else(|| "Header mismatch: missing required Mcp-Name header".to_string())?; + let name = decode_header_value(raw_name, "Mcp-Name")?; + if name != want { + return Err(format!( + "Header mismatch: Mcp-Name header does not match {label}" + )); + } } + + Ok(()) +} + +/// Decode the SEP-2243 `=?base64?{value}?=` sentinel, if present. +/// +/// FLAGGED: the sentinel's base64 alphabet and padding are not pinned down +/// by the material available here. We accept standard base64 and fall back +/// to unpadded standard base64; verify against SEP-2243 or an SDK before +/// relying on anything more exotic. +fn decode_header_value(raw: &str, header_name: &str) -> Result { + let Some(inner) = raw + .strip_prefix(B64_PREFIX) + .and_then(|r| r.strip_suffix(B64_SUFFIX)) + else { + return Ok(raw.to_string()); + }; + let malformed = + || format!("Header mismatch: malformed base64 sentinel in {header_name} header"); + let bytes = base64::engine::general_purpose::STANDARD + .decode(inner) + .or_else(|_| base64::engine::general_purpose::STANDARD_NO_PAD.decode(inner)) + .map_err(|_| malformed())?; + String::from_utf8(bytes).map_err(|_| malformed()) } -async fn handle_initialize(state: McpHttpState, message: JsonRpcMessage) -> Response { - let session_id = Uuid::new_v4().to_string(); - let server = Arc::new(Mutex::new(McpServer::new(state.handler.clone()))); +// ── Legacy era ────────────────────────────────────────────────────── + +async fn handle_legacy_initialize( + state: McpHttpState, + headers: HeaderMap, + message: JsonRpcMessage, +) -> Response { + // Reuse a live session if the client already has one, instead of + // minting a second and orphaning the first. + let existing = header_str(&headers, HDR_SESSION_ID) + .map(str::to_string) + .filter(|id| state.sessions.contains_key(id)); + + let (session_id, server, is_new) = match existing { + Some(id) => { + let server = match state.sessions.get(&id) { + Some(entry) => { + entry.last_seen.store(now_ms(), Ordering::Relaxed); + entry.server.clone() + } + None => return unknown_session(&message), + }; + (id, server, false) + } + None => { + sweep_sessions( + &state.sessions, + now_ms(), + state.config.session_idle_timeout.as_millis() as u64, + ); + if state.sessions.len() >= state.config.max_sessions { + let id = message.id.clone().unwrap_or(JsonRpcId::Null); + return jsonrpc_error( + StatusCode::SERVICE_UNAVAILABLE, + id, + INTERNAL_ERROR, + "Too many active MCP sessions", + None, + ); + } + ( + Uuid::new_v4().to_string(), + Arc::new(Mutex::new(McpServer::from_dispatcher( + state.dispatcher.clone(), + ))), + true, + ) + } + }; let response = { let mut srv = server.lock().await; srv.handle_message(message).await }; - state.sessions.insert(session_id.clone(), server); - tracing::info!(session_id = %session_id, "MCP HTTP session created"); + let succeeded = response.as_ref().is_some_and(|resp| resp.error.is_none()); - match response { - Some(resp) => { - let mut http_resp = Json(resp).into_response(); - if let Ok(val) = HeaderValue::from_str(&session_id) { - http_resp.headers_mut().insert(MCP_SESSION_ID, val); - } - http_resp - } - None => StatusCode::ACCEPTED.into_response(), + // Only a SUCCESSFUL initialize creates a session. A failed one used to + // leave behind an entry that could never become initialized. + if is_new && succeeded { + state.sessions.insert( + session_id.clone(), + SessionEntry { + server, + last_seen: Arc::new(AtomicU64::new(now_ms())), + }, + ); + tracing::info!(session_id = %session_id, "MCP HTTP session created"); + } + + let mut http_resp = dispatch_response(response); + if succeeded && let Ok(value) = HeaderValue::from_str(&session_id) { + http_resp.headers_mut().insert(HDR_SESSION_ID, value); } + http_resp } -async fn handle_session_message( +async fn handle_legacy_session( state: McpHttpState, headers: HeaderMap, message: JsonRpcMessage, ) -> Response { - let session_id = match headers.get(MCP_SESSION_ID).and_then(|v| v.to_str().ok()) { + let session_id = match header_str(&headers, HDR_SESSION_ID) { Some(id) => id.to_string(), None => { - return ( + return jsonrpc_error( StatusCode::BAD_REQUEST, - Json(JsonRpcResponse::error( - JsonRpcId::Number(0), - INVALID_REQUEST, - "Missing Mcp-Session-Id header", - )), - ) - .into_response(); + message.id.clone().unwrap_or(JsonRpcId::Null), + INVALID_REQUEST, + "Missing Mcp-Session-Id header", + None, + ); } }; let server = match state.sessions.get(&session_id) { - Some(entry) => entry.value().clone(), - None => return StatusCode::NOT_FOUND.into_response(), + Some(entry) => { + entry.last_seen.store(now_ms(), Ordering::Relaxed); + entry.server.clone() + } + None => return unknown_session(&message), }; let response = { let mut srv = server.lock().await; srv.handle_message(message).await }; + dispatch_response(response) +} - match response { - Some(resp) => Json(resp).into_response(), - None => StatusCode::ACCEPTED.into_response(), +/// A 404 that a client can actually act on -- what makes idle expiry +/// recoverable rather than a silent hang. +fn unknown_session(message: &JsonRpcMessage) -> Response { + jsonrpc_error( + StatusCode::NOT_FOUND, + message.id.clone().unwrap_or(JsonRpcId::Null), + INVALID_REQUEST, + "Unknown or expired MCP session; send 'initialize' again", + None, + ) +} + +// ── Session lifetime ──────────────────────────────────────────────── + +/// Evict sessions idle for at least `idle_ms`. Returns how many were +/// removed. +/// +/// Pure in `now_ms` so it is unit-testable without a clock abstraction. +fn sweep_sessions(sessions: &DashMap, now_ms: u64, idle_ms: u64) -> usize { + let before = sessions.len(); + sessions.retain(|_, entry| { + let last = entry.last_seen.load(Ordering::Relaxed); + now_ms.saturating_sub(last) < idle_ms + }); + let removed = before - sessions.len(); + if removed > 0 { + tracing::info!(removed, "swept idle MCP HTTP sessions"); } + removed } +/// Run the sweep at most once per [`SWEEP_THROTTLE_MS`]. +/// +/// Lazy rather than a background task: the router constructor has no +/// shutdown handle with which to cancel one, and a lazy sweep is +/// deterministic and testable. +fn maybe_sweep(state: &McpHttpState) { + let now = now_ms(); + let last = state.last_sweep.load(Ordering::Relaxed); + if now.saturating_sub(last) < SWEEP_THROTTLE_MS { + return; + } + if state + .last_sweep + .compare_exchange(last, now, Ordering::Relaxed, Ordering::Relaxed) + .is_err() + { + return; + } + sweep_sessions( + &state.sessions, + now, + state.config.session_idle_timeout.as_millis() as u64, + ); +} + +// ── DELETE / GET ──────────────────────────────────────────────────── + async fn handle_delete(State(state): State, headers: HeaderMap) -> Response { - let session_id = match headers.get(MCP_SESSION_ID).and_then(|v| v.to_str().ok()) { - Some(id) => id.to_string(), - None => return StatusCode::BAD_REQUEST.into_response(), + if let Some(origin) = header_str(&headers, header::ORIGIN.as_str()) + && !state.origins.allows(Some(origin)) + { + return origin_rejected(origin); + } + + // Without a session id there is nothing to delete: DELETE is a + // legacy-only affordance, so the modern-era answer (405) applies. + let Some(session_id) = header_str(&headers, HDR_SESSION_ID) else { + return method_not_allowed(); }; - match state.sessions.remove(&session_id) { + match state.sessions.remove(session_id) { Some(_) => { tracing::info!(session_id = %session_id, "MCP HTTP session terminated"); StatusCode::OK.into_response() @@ -147,3 +622,556 @@ async fn handle_delete(State(state): State, headers: HeaderMap) -> None => StatusCode::NOT_FOUND.into_response(), } } + +/// The 2026-07-28 transport removed the GET stream endpoint; recalld never +/// had one, so 405 is both compliant and non-regressive. +async fn handle_get() -> Response { + method_not_allowed() +} + +// ── Tests ─────────────────────────────────────────────────────────── + +#[cfg(test)] +mod tests { + use super::*; + use crate::mcp::server::tests::{ + initialize_request, modern_request, plain_request, stub_handler, + }; + use axum::body::Body; + use axum::http::Request; + use tower::ServiceExt; + + fn test_state(config: McpHttpConfig) -> (Router, McpHttpState) { + build_router(stub_handler(), config) + } + + fn router() -> Router { + test_state(McpHttpConfig::default()).0 + } + + struct Req { + builder: axum::http::request::Builder, + body: Vec, + } + + impl Req { + fn post(message: &JsonRpcMessage) -> Self { + Self { + builder: Request::builder() + .method("POST") + .uri("/mcp") + .header("content-type", "application/json"), + body: serde_json::to_vec(message).unwrap(), + } + } + + fn raw(body: Vec) -> Self { + Self { + builder: Request::builder() + .method("POST") + .uri("/mcp") + .header("content-type", "application/json"), + body, + } + } + + fn header(mut self, name: &str, value: &str) -> Self { + self.builder = self.builder.header(name, value); + self + } + + /// Attach the three SEP-2243 headers a modern client must send. + fn sep2243(self, method: &str, name: Option<&str>) -> Self { + let mut req = self + .header(HDR_MCP_METHOD, method) + .header(HDR_MCP_PROTOCOL_VER, PROTOCOL_VERSION); + if let Some(name) = name { + req = req.header(HDR_MCP_NAME, name); + } + req + } + + fn build(self) -> Request { + self.builder.body(Body::from(self.body)).unwrap() + } + } + + async fn send(router: &Router, req: Req) -> (StatusCode, HeaderMap, serde_json::Value) { + let resp = router.clone().oneshot(req.build()).await.unwrap(); + let status = resp.status(); + let headers = resp.headers().clone(); + let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX) + .await + .unwrap(); + let json = if bytes.is_empty() { + serde_json::Value::Null + } else { + serde_json::from_slice(&bytes).unwrap_or(serde_json::Value::Null) + }; + (status, headers, json) + } + + fn tools_list() -> JsonRpcMessage { + modern_request(1, "tools/list", serde_json::json!({})) + } + + // ── Modern path ───────────────────────────────────────────────── + + #[tokio::test] + async fn modern_post_succeeds_without_a_session_header() { + let r = router(); + let (status, _, body) = + send(&r, Req::post(&tools_list()).sep2243("tools/list", None)).await; + assert_eq!(status, StatusCode::OK); + assert_eq!(body["result"]["tools"][0]["name"], "store_memory"); + assert_eq!(body["result"]["resultType"], "complete"); + } + + #[tokio::test] + async fn modern_response_carries_no_session_header() { + let r = router(); + let (_, headers, _) = send(&r, Req::post(&tools_list()).sep2243("tools/list", None)).await; + assert!(headers.get(HDR_SESSION_ID).is_none()); + } + + #[tokio::test] + async fn modern_request_ignores_an_inbound_session_id() { + let (r, state) = test_state(McpHttpConfig::default()); + let (status, headers, _) = send( + &r, + Req::post(&tools_list()) + .sep2243("tools/list", None) + .header(HDR_SESSION_ID, "does-not-exist"), + ) + .await; + assert_eq!(status, StatusCode::OK); + assert!(headers.get(HDR_SESSION_ID).is_none()); + assert_eq!(state.sessions.len(), 0); + } + + #[tokio::test] + async fn missing_mcp_method_header_is_header_mismatch() { + let r = router(); + let (status, _, body) = send( + &r, + Req::post(&tools_list()).header(HDR_MCP_PROTOCOL_VER, PROTOCOL_VERSION), + ) + .await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(body["error"]["code"], HEADER_MISMATCH); + assert!( + body["error"]["message"] + .as_str() + .unwrap() + .contains("Mcp-Method") + ); + } + + #[tokio::test] + async fn mismatched_mcp_method_header_is_header_mismatch() { + let r = router(); + let (status, _, body) = + send(&r, Req::post(&tools_list()).sep2243("tools/call", None)).await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(body["error"]["code"], HEADER_MISMATCH); + assert!( + body["error"]["message"] + .as_str() + .unwrap() + .contains("does not match body method") + ); + } + + #[tokio::test] + async fn missing_protocol_version_header_is_header_mismatch() { + let r = router(); + let (status, _, body) = send( + &r, + Req::post(&tools_list()).header(HDR_MCP_METHOD, "tools/list"), + ) + .await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(body["error"]["code"], HEADER_MISMATCH); + assert!( + body["error"]["message"] + .as_str() + .unwrap() + .contains("MCP-Protocol-Version") + ); + } + + #[tokio::test] + async fn mcp_name_mismatch_on_tools_call_is_header_mismatch() { + let r = router(); + let msg = modern_request( + 1, + "tools/call", + serde_json::json!({ "name": "store_memory", "arguments": {} }), + ); + let (status, _, body) = send( + &r, + Req::post(&msg).sep2243("tools/call", Some("recall_memories")), + ) + .await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(body["error"]["code"], HEADER_MISMATCH); + assert!( + body["error"]["message"] + .as_str() + .unwrap() + .contains("params.name") + ); + } + + #[tokio::test] + async fn base64_sentinel_is_decoded_before_comparison() { + let r = router(); + let msg = modern_request( + 1, + "tools/call", + serde_json::json!({ "name": "store_memory", "arguments": {} }), + ); + let encoded = format!( + "{B64_PREFIX}{}{B64_SUFFIX}", + base64::engine::general_purpose::STANDARD.encode("store_memory") + ); + let (status, _, body) = + send(&r, Req::post(&msg).sep2243("tools/call", Some(&encoded))).await; + assert_eq!(status, StatusCode::OK, "body: {body}"); + assert_eq!(body["result"]["content"][0]["text"], "called store_memory"); + + // A malformed sentinel is a mismatch, not a silent pass. + let bad = format!("{B64_PREFIX}!!!not-base64!!!{B64_SUFFIX}"); + let (status, _, body) = send(&r, Req::post(&msg).sep2243("tools/call", Some(&bad))).await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(body["error"]["code"], HEADER_MISMATCH); + } + + #[tokio::test] + async fn modern_notification_skips_header_validation() { + let r = router(); + let mut msg = modern_request(1, "notifications/cancelled", serde_json::json!({})); + msg.id = None; + // No SEP-2243 headers at all. + let (status, _, _) = send(&r, Req::post(&msg)).await; + assert_eq!(status, StatusCode::ACCEPTED); + } + + #[tokio::test] + async fn unsupported_protocol_version_is_rejected_with_supported_list() { + let r = router(); + let msg = plain_request( + 1, + "tools/list", + Some(serde_json::json!({ + "_meta": { + META_PROTOCOL_VERSION: "2027-01-01", + META_CLIENT_CAPABILITIES: {}, + } + })), + ); + let (status, _, body) = send( + &r, + Req::post(&msg) + .header(HDR_MCP_METHOD, "tools/list") + .header(HDR_MCP_PROTOCOL_VER, "2027-01-01"), + ) + .await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(body["error"]["code"], UNSUPPORTED_PROTOCOL_VERSION); + assert_eq!(body["error"]["data"]["requested"], "2027-01-01"); + assert_eq!( + body["error"]["data"]["supported"], + serde_json::json!([PROTOCOL_VERSION]) + ); + } + + // ── Legacy path ───────────────────────────────────────────────── + + #[tokio::test] + async fn legacy_initialize_mints_a_session_id() { + let (r, state) = test_state(McpHttpConfig::default()); + let (status, headers, body) = + send(&r, Req::post(&initialize_request(1, "2025-06-18"))).await; + assert_eq!(status, StatusCode::OK); + assert!(headers.get(HDR_SESSION_ID).is_some()); + assert_eq!(body["result"]["protocolVersion"], "2025-06-18"); + assert_eq!(state.sessions.len(), 1); + } + + #[tokio::test] + async fn failed_initialize_does_not_insert_a_session() { + let (r, state) = test_state(McpHttpConfig::default()); + // `protocolVersion` as a number fails InitializeParams parsing. + let msg = plain_request( + 1, + "initialize", + Some(serde_json::json!({ "protocolVersion": 7 })), + ); + let (_, headers, body) = send(&r, Req::post(&msg)).await; + assert!(body["error"].is_object(), "expected an error: {body}"); + assert!(headers.get(HDR_SESSION_ID).is_none()); + assert_eq!(state.sessions.len(), 0); + } + + #[tokio::test] + async fn repeat_initialize_with_a_live_session_reuses_the_same_id() { + let (r, state) = test_state(McpHttpConfig::default()); + let (_, headers, _) = send(&r, Req::post(&initialize_request(1, "2025-06-18"))).await; + let first = headers + .get(HDR_SESSION_ID) + .unwrap() + .to_str() + .unwrap() + .to_string(); + + let (_, headers, _) = send( + &r, + Req::post(&initialize_request(2, "2025-06-18")).header(HDR_SESSION_ID, &first), + ) + .await; + let second = headers.get(HDR_SESSION_ID).unwrap().to_str().unwrap(); + assert_eq!(first, second); + assert_eq!(state.sessions.len(), 1, "no orphaned session"); + } + + #[tokio::test] + async fn legacy_message_without_session_header_echoes_the_real_id() { + let r = router(); + // `ping` is era-neutral, so use a method that requires a session. + let msg = plain_request(4242, "tools/call", Some(serde_json::json!({ "name": "x" }))); + let (status, _, body) = send(&r, Req::post(&msg)).await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(body["id"], 4242, "must not fabricate id 0"); + assert_eq!(body["error"]["code"], INVALID_PARAMS); + + // With a session header but no such session, the id is echoed too. + let (status, _, body) = send(&r, Req::post(&msg).header(HDR_SESSION_ID, "nope")).await; + assert_eq!(status, StatusCode::NOT_FOUND); + assert_eq!(body["id"], 4242); + } + + #[tokio::test] + async fn unknown_session_returns_404_with_a_jsonrpc_body() { + let r = router(); + let msg = plain_request(9, "tools/list", None); + let (status, _, body) = send(&r, Req::post(&msg).header(HDR_SESSION_ID, "ghost")).await; + assert_eq!(status, StatusCode::NOT_FOUND); + assert_eq!(body["error"]["code"], INVALID_REQUEST); + assert!( + body["error"]["message"] + .as_str() + .unwrap() + .contains("send 'initialize' again") + ); + } + + #[tokio::test] + async fn parse_error_uses_a_null_id() { + let r = router(); + let (status, _, body) = send(&r, Req::raw(b"{not json".to_vec())).await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(body["error"]["code"], PARSE_ERROR); + assert!(body["id"].is_null(), "expected null id, got {}", body["id"]); + } + + // ── DELETE / GET ──────────────────────────────────────────────── + + #[tokio::test] + async fn delete_with_a_session_id_terminates_it() { + let (r, state) = test_state(McpHttpConfig::default()); + let (_, headers, _) = send(&r, Req::post(&initialize_request(1, "2025-06-18"))).await; + let sid = headers + .get(HDR_SESSION_ID) + .unwrap() + .to_str() + .unwrap() + .to_string(); + + let resp = r + .clone() + .oneshot( + Request::builder() + .method("DELETE") + .uri("/mcp") + .header(HDR_SESSION_ID, &sid) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + assert_eq!(state.sessions.len(), 0); + } + + #[tokio::test] + async fn delete_without_a_session_id_is_405() { + let r = router(); + let resp = r + .oneshot( + Request::builder() + .method("DELETE") + .uri("/mcp") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::METHOD_NOT_ALLOWED); + assert_eq!(resp.headers()[header::ALLOW], ALLOW_METHODS); + } + + #[tokio::test] + async fn get_is_405_with_allow_header() { + let r = router(); + let resp = r + .oneshot( + Request::builder() + .method("GET") + .uri("/mcp") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::METHOD_NOT_ALLOWED); + assert_eq!(resp.headers()[header::ALLOW], ALLOW_METHODS); + } + + // ── Session lifetime ──────────────────────────────────────────── + + #[tokio::test] + async fn sweep_evicts_idle_sessions_and_keeps_recent_ones() { + let sessions: DashMap = DashMap::new(); + let make = |last: u64| SessionEntry { + server: Arc::new(Mutex::new(McpServer::new(stub_handler()))), + last_seen: Arc::new(AtomicU64::new(last)), + }; + sessions.insert("stale".to_string(), make(0)); + sessions.insert("fresh".to_string(), make(9_000)); + + let removed = sweep_sessions(&sessions, 10_000, 5_000); + assert_eq!(removed, 1); + assert!(sessions.contains_key("fresh")); + assert!(!sessions.contains_key("stale")); + } + + #[tokio::test] + async fn session_cap_returns_503() { + let (r, state) = test_state(McpHttpConfig { + max_sessions: 1, + ..Default::default() + }); + send(&r, Req::post(&initialize_request(1, "2025-06-18"))).await; + assert_eq!(state.sessions.len(), 1); + + let (status, _, body) = send(&r, Req::post(&initialize_request(2, "2025-06-18"))).await; + assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(body["error"]["code"], INTERNAL_ERROR); + assert!( + body["error"]["message"] + .as_str() + .unwrap() + .contains("Too many active MCP sessions") + ); + } + + // ── Origin validation ─────────────────────────────────────────── + + #[tokio::test] + async fn origin_absent_is_allowed() { + let r = router(); + let (status, _, _) = send(&r, Req::post(&tools_list()).sep2243("tools/list", None)).await; + assert_eq!(status, StatusCode::OK); + } + + #[tokio::test] + async fn loopback_origins_are_allowed_on_any_scheme_or_port() { + let r = router(); + for origin in [ + "http://localhost", + "http://localhost:7680", + "https://127.0.0.1:3000", + "http://[::1]:7680", + ] { + let (status, _, _) = send( + &r, + Req::post(&tools_list()) + .sep2243("tools/list", None) + .header("origin", origin), + ) + .await; + assert_eq!(status, StatusCode::OK, "origin {origin}"); + } + } + + #[tokio::test] + async fn opaque_null_origin_is_forbidden() { + let r = router(); + let (status, _, body) = send( + &r, + Req::post(&tools_list()) + .sep2243("tools/list", None) + .header("origin", "null"), + ) + .await; + assert_eq!(status, StatusCode::FORBIDDEN); + assert_eq!(body["error"], "Origin not allowed"); + } + + #[tokio::test] + async fn foreign_origin_is_forbidden() { + let r = router(); + let (status, _, _) = send( + &r, + Req::post(&tools_list()) + .sep2243("tools/list", None) + .header("origin", "http://evil.example.com"), + ) + .await; + assert_eq!(status, StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn allowlisted_origin_is_permitted() { + let (r, _) = test_state(McpHttpConfig { + allowed_origins: vec!["https://App.Example.com".to_string()], + ..Default::default() + }); + let (status, _, _) = send( + &r, + Req::post(&tools_list()) + .sep2243("tools/list", None) + .header("origin", "https://app.example.com"), + ) + .await; + assert_eq!(status, StatusCode::OK); + } + + // ── Body limit ────────────────────────────────────────────────── + + #[tokio::test] + async fn oversized_body_is_rejected() { + let (r, _) = test_state(McpHttpConfig { + max_body_bytes: 512, + ..Default::default() + }); + let (status, _, _) = send(&r, Req::raw(vec![b'x'; 4096])).await; + assert_eq!(status, StatusCode::PAYLOAD_TOO_LARGE); + } + + // ── Era-neutral over HTTP ─────────────────────────────────────── + + #[tokio::test] + async fn discover_is_answered_without_meta_or_session() { + let r = router(); + let (status, headers, body) = + send(&r, Req::post(&plain_request(1, "server/discover", None))).await; + assert_eq!(status, StatusCode::OK); + assert!(headers.get(HDR_SESSION_ID).is_none()); + assert_eq!( + body["result"]["supportedVersions"], + serde_json::json!([PROTOCOL_VERSION]) + ); + } +} diff --git a/src/mcp/mod.rs b/src/mcp/mod.rs index 1b080ac..7fffaff 100644 --- a/src/mcp/mod.rs +++ b/src/mcp/mod.rs @@ -15,9 +15,9 @@ pub mod bridge_adapters; pub mod resources; pub mod tools; -pub use http_transport::mcp_router; +pub use http_transport::{McpHttpConfig, mcp_router}; pub use protocol::*; -pub use server::{McpHandler, McpServer}; +pub use server::{McpDispatcher, McpHandler, McpServer, RequestContext}; pub use transport::run_stdio; use thiserror::Error; @@ -33,10 +33,6 @@ pub enum McpError { #[error("JSON error: {0}")] Json(#[from] serde_json::Error), - /// The server has not been initialized yet. - #[error("Server not initialized")] - NotInitialized, - /// The requested tool was not found. #[error("Tool not found: {0}")] ToolNotFound(String), diff --git a/src/mcp/protocol.rs b/src/mcp/protocol.rs index f284d26..bfedf95 100644 --- a/src/mcp/protocol.rs +++ b/src/mcp/protocol.rs @@ -2,23 +2,106 @@ //! //! All wire-format types for the Model Context Protocol, including //! lifecycle, tool, and resource messages. +//! +//! Recalld is a **dual-era** MCP server. Two protocol eras share this +//! module and both transports: +//! +//! * **Modern** (MCP `2026-07-28`) is stateless. There is no +//! `initialize` handshake and no session; every request carries its +//! own protocol version, client info and client capabilities in +//! `params._meta`, and every result carries `resultType`, caching +//! hints and `_meta.io.modelcontextprotocol/serverInfo`. +//! * **Legacy** (MCP `2025-06-18` and earlier) keeps the `initialize` +//! handshake and, over HTTP, the `Mcp-Session-Id` session. +//! +//! The era is decided **per message** by [`era_hint`]. use serde::{Deserialize, Serialize}; // ── Constants ─────────────────────────────────────────────────────── -/// MCP protocol version this server implements. -pub const PROTOCOL_VERSION: &str = "2025-06-18"; +/// MCP protocol version this server implements natively (modern era). +/// +/// This value is used for the stateless 2026-07-28 path only: it is the +/// sole version accepted in `_meta.io.modelcontextprotocol/protocolVersion` +/// and the value expected in the `MCP-Protocol-Version` header. +/// +/// It is deliberately **never** echoed from an `initialize` response -- +/// see [`negotiate_legacy_version`]. +pub const PROTOCOL_VERSION: &str = "2026-07-28"; + +/// Modern protocol versions this server accepts on the stateless path. +/// +/// Advertised by `server/discover` as `supportedVersions` and returned in +/// the `data.supported` array of an [`UNSUPPORTED_PROTOCOL_VERSION`] error. +/// +/// This list is intentionally MODERN-ONLY and is **not** the dual-era +/// union: the field feeds a client's fall-forward retry loop, so listing +/// `2025-06-18` here would invite a modern client to retry with +/// `_meta.protocolVersion = "2025-06-18"` on the stateless path -- which we +/// would then have to reject, contradicting our own advertisement. Legacy +/// support stays discoverable the legacy way (send `initialize`). +/// +/// NOTE (flagged inference): the spec does not state how a dual-era server +/// should populate `supportedVersions`; this is inferred from the field's +/// role in the -32022 retry loop. +pub const SUPPORTED_MODERN_VERSIONS: &[&str] = &[PROTOCOL_VERSION]; + +/// Legacy protocol versions accepted -- and echoed -- by `initialize`. +/// +/// NOTE (flagged): `2025-11-25` is deliberately absent. Its wire changes +/// were not verified against recalld's current message shapes; adding it +/// needs a diff review first. +pub const SUPPORTED_LEGACY_VERSIONS: &[&str] = &["2025-06-18", "2025-03-26", "2024-11-05"]; + +/// Version echoed by `initialize` when the client asks for one we do not +/// recognise. +pub const LEGACY_PROTOCOL_VERSION: &str = "2025-06-18"; /// JSON-RPC version string. pub const JSONRPC_VERSION: &str = "2.0"; -/// Server name reported in initialize response. +/// Server name reported in initialize / discover responses. pub const SERVER_NAME: &str = "recalld"; -/// Server version reported in initialize response. +/// Server version reported in initialize / discover responses. pub const SERVER_VERSION: &str = env!("CARGO_PKG_VERSION"); +// ── Modern `_meta` keys (reverse-DNS, MCP 2026-07-28) ─────────────── + +/// `_meta` key carrying the request's protocol version (REQUIRED). +pub const META_PROTOCOL_VERSION: &str = "io.modelcontextprotocol/protocolVersion"; +/// `_meta` key carrying the client's name and version (optional). +pub const META_CLIENT_INFO: &str = "io.modelcontextprotocol/clientInfo"; +/// `_meta` key carrying the client's capability declarations (REQUIRED). +pub const META_CLIENT_CAPABILITIES: &str = "io.modelcontextprotocol/clientCapabilities"; +/// `_meta` key carrying the client's desired log level (optional, unused). +pub const META_LOG_LEVEL: &str = "io.modelcontextprotocol/logLevel"; +/// `_meta` key under which every result reports this server's identity. +pub const META_SERVER_INFO: &str = "io.modelcontextprotocol/serverInfo"; + +// ── Cache hint tuning (SEP-2549) ──────────────────────────────────── + +/// Freshness hint for the tool/resource/template catalogs and +/// `server/discover`, in milliseconds. +/// +/// Recalld's catalogs are compile-time constants, so one hour is +/// conservative rather than optimistic. +pub const CATALOG_TTL_MS: u64 = 3_600_000; + +/// Freshness hint for `resources/read`, in milliseconds. +/// +/// Zero (immediately stale): resource contents are live instance state, +/// not a static catalog. +pub const RESOURCE_READ_TTL_MS: u64 = 0; + +/// Usage instructions surfaced by both `initialize` and `server/discover`. +pub const SERVER_INSTRUCTIONS: &str = "Recalld is an AI memory system with human-like forgetting. \ + Use store_memory to save observations, recall_memories to \ + search by semantic similarity, and reinforce_memory to \ + strengthen useful memories. Memories decay naturally over \ + time unless reinforced."; + // ── JSON-RPC error codes ──────────────────────────────────────────── /// JSON-RPC parse error code. @@ -32,6 +115,23 @@ pub const INVALID_PARAMS: i32 = -32602; /// JSON-RPC internal error code. pub const INTERNAL_ERROR: i32 = -32603; +// ── MCP 2026-07-28 error codes ────────────────────────────────────── + +/// SEP-2243: a required MCP header is missing, or disagrees with the body. +pub const HEADER_MISMATCH: i32 = -32020; + +/// SEP-2575: the server needs a client capability the client did not +/// declare, with `data.requiredCapabilities` naming them. +/// +/// DEFINED BUT NEVER EMITTED. Recalld requires no client capabilities of +/// any kind, so there is no call site -- the constant exists so the error +/// space is complete and a reader does not go hunting for one. +pub const MISSING_REQUIRED_CLIENT_CAPABILITY: i32 = -32021; + +/// SEP-2575: the request's protocol version is not one this server +/// supports. Carries `data.supported` and `data.requested`. +pub const UNSUPPORTED_PROTOCOL_VERSION: i32 = -32022; + // ── Core JSON-RPC Types ───────────────────────────────────────────── /// JSON-RPC 2.0 request or notification. @@ -51,7 +151,14 @@ pub struct JsonRpcMessage { pub params: Option, } -/// JSON-RPC message ID -- either an integer or a string. +/// JSON-RPC message ID -- an integer, a string, or null. +/// +/// `Null` exists so a response to an unparseable request can say "we could +/// not read your id" instead of fabricating id `0` and colliding with a +/// real in-flight request. Variant order matters for `untagged` +/// deserialization: `null` matches neither `Number` nor `String`, and +/// `Option` still resolves an explicit `"id": null` to `None` +/// because serde tries `Option` first. #[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] #[serde(untagged)] pub enum JsonRpcId { @@ -59,6 +166,8 @@ pub enum JsonRpcId { Number(i64), /// String request ID. String(String), + /// Explicit JSON `null` -- used when the request id is unknowable. + Null, } /// JSON-RPC 2.0 success or error response. @@ -99,8 +208,20 @@ impl JsonRpcResponse { } } - /// Build an error response. + /// Build an error response with no `data` member. pub fn error(id: JsonRpcId, code: i32, message: impl Into) -> Self { + Self::error_with_data(id, code, message, None) + } + + /// Build an error response, optionally carrying a `data` member. + /// + /// [`UNSUPPORTED_PROTOCOL_VERSION`] requires `data`. + pub fn error_with_data( + id: JsonRpcId, + code: i32, + message: impl Into, + data: Option, + ) -> Self { Self { jsonrpc: JSONRPC_VERSION.to_string(), id, @@ -108,7 +229,7 @@ impl JsonRpcResponse { error: Some(JsonRpcError { code, message: message.into(), - data: None, + data, }), } } @@ -365,8 +486,644 @@ pub struct ResourcesListResult { } /// Response wrapper for `resources/templates/list`. +/// +/// The `camelCase` rename is load-bearing: every MCP revision names this +/// field `resourceTemplates`. #[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] pub struct ResourceTemplatesListResult { /// List of available resource templates. pub resource_templates: Vec, } + +// ── MCP 2026-07-28: result envelope (SEP-2549, SEP-2575) ──────────── + +/// Whether a result is final or an interim request for more input. +/// +/// REQUIRED on every 2026-07-28 result. Clients speaking older revisions +/// treat its absence as `Complete`, so emitting it is additive. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ResultType { + /// A final result. + Complete, + /// An interim result asking the client for more input. Recalld never + /// produces one; interim results carry no caching hints. + InputRequired, +} + +/// Who may reuse a cached result. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum CacheScope { + /// Contains no user-specific data; any shared proxy may serve it to + /// any user. + Public, + /// Reusable only within the same authorization context. + Private, +} + +/// SEP-2549 caching hints, emitted at the TOP LEVEL of a result (siblings +/// of `tools`/`resources`/`contents`), never inside `_meta`. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct CacheHints { + /// Freshness lifetime in milliseconds, analogous to HTTP + /// `Cache-Control: max-age`. `0` means immediately stale. It is a + /// hint, not a guarantee, and not a polling interval. + pub ttl_ms: u64, + /// Sharing scope for the cached value. + pub cache_scope: CacheScope, +} + +impl CacheHints { + /// Hints for a result containing no user-specific data. + pub const fn public(ttl_ms: u64) -> Self { + Self { + ttl_ms, + cache_scope: CacheScope::Public, + } + } + + /// Hints for a result reusable only within one authorization context. + pub const fn private(ttl_ms: u64) -> Self { + Self { + ttl_ms, + cache_scope: CacheScope::Private, + } + } +} + +/// The `_meta` object attached to every result. +/// +/// The 2026-07-28 spec says servers SHOULD report their identity on every +/// result. `clientInfo`/`serverInfo` are explicitly not verified by the +/// protocol and must not drive behavioral or security decisions. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ResultMeta { + /// This server's name and version. + #[serde(rename = "io.modelcontextprotocol/serverInfo")] + pub server_info: Implementation, +} + +impl ResultMeta { + /// Build the `_meta` block for this build of recalld. + pub fn current() -> Self { + Self { + server_info: Implementation { + name: SERVER_NAME.to_string(), + version: SERVER_VERSION.to_string(), + }, + } + } +} + +impl Default for ResultMeta { + fn default() -> Self { + Self::current() + } +} + +/// Wraps any result payload with the fields 2026-07-28 requires. +/// +/// Applied once, at the single success arm of the dispatcher, rather than +/// as fields on each result struct -- one place to get right, and no result +/// type can forget it. +/// +/// The envelope is emitted on BOTH eras. In the 2025-06-18 schema `Result` +/// is `{ _meta?: {...}; [key: string]: unknown }`, so the extra top-level +/// keys are schema-legal there, not merely tolerated. If a strict +/// `additionalProperties: false` client ever appears, the escape hatch is a +/// single era conditional here. +/// +/// `payload` MUST serialize to a JSON object: `serde(flatten)` over a +/// non-object fails at serialization time (surfacing as an internal error). +/// There is exactly one `flatten` here on purpose -- flattening an +/// `Option` alongside `skip_serializing_if` is a known serde footgun, so +/// the cache hints are two plain `Option` fields instead. +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct ResultEnvelope { + /// The method-specific result body. + #[serde(flatten)] + pub payload: serde_json::Value, + /// Whether this result is final. + pub result_type: ResultType, + /// Freshness lifetime, when the result is cacheable. + #[serde(skip_serializing_if = "Option::is_none")] + pub ttl_ms: Option, + /// Cache sharing scope, when the result is cacheable. + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_scope: Option, + /// Server identity metadata. + #[serde(rename = "_meta")] + pub meta: ResultMeta, +} + +impl ResultEnvelope { + /// Wrap a final result payload, with no caching hints. + pub fn complete(payload: serde_json::Value) -> Self { + debug_assert!( + payload.is_object(), + "ResultEnvelope payload must be a JSON object (serde flatten requirement)" + ); + Self { + payload, + result_type: ResultType::Complete, + ttl_ms: None, + cache_scope: None, + meta: ResultMeta::current(), + } + } + + /// Attach SEP-2549 caching hints. + pub fn with_cache(mut self, hints: CacheHints) -> Self { + self.ttl_ms = Some(hints.ttl_ms); + self.cache_scope = Some(hints.cache_scope); + self + } +} + +// ── MCP 2026-07-28: server/discover ───────────────────────────────── + +/// Result of `server/discover`, the 2026-07-28 replacement for the +/// `initialize` handshake. Servers MUST implement it. +/// +/// `resultType`, `ttlMs`, `cacheScope` and `_meta` are supplied by +/// [`ResultEnvelope`], not by this struct. +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct DiscoverResult { + /// Modern protocol versions this server accepts. + pub supported_versions: Vec, + /// Server capability declarations. + pub capabilities: ServerCapabilities, + /// Optional usage instructions for the client. + #[serde(skip_serializing_if = "Option::is_none")] + pub instructions: Option, +} + +/// The capability set recalld advertises. +/// +/// Shared by `initialize` and `server/discover` so the two cannot drift. +pub fn server_capabilities() -> ServerCapabilities { + ServerCapabilities { + tools: Some(ToolsCapability { + list_changed: false, + }), + resources: Some(ResourcesCapability { + subscribe: false, + list_changed: false, + }), + } +} + +// ── MCP 2026-07-28: per-request `_meta` ───────────────────────────── + +/// Raw, permissive view of `params._meta`. +/// +/// Every field is `Option` on purpose: it lets us emit a -32602 that NAMES +/// the missing key instead of leaking a serde error message. +#[derive(Debug, Clone, Default, Deserialize)] +pub struct RawRequestMeta { + /// The request's protocol version. + #[serde(rename = "io.modelcontextprotocol/protocolVersion")] + pub protocol_version: Option, + /// The client's name and version. + #[serde(rename = "io.modelcontextprotocol/clientInfo")] + pub client_info: Option, + /// The client's capability declarations. + #[serde(rename = "io.modelcontextprotocol/clientCapabilities")] + pub client_capabilities: Option, + /// The client's desired log level. Parsed but unused -- recalld emits + /// no `notifications/message`. + #[serde(rename = "io.modelcontextprotocol/logLevel")] + pub log_level: Option, + /// Progress token for long-running requests. + #[serde(rename = "progressToken")] + pub progress_token: Option, +} + +/// Validated per-request metadata from a modern client. +#[derive(Debug, Clone)] +pub struct RequestMeta { + /// The request's protocol version (already checked against + /// [`SUPPORTED_MODERN_VERSIONS`]). + pub protocol_version: String, + /// The client's name and version, if it sent them. + pub client_info: Option, + /// The client's capability declarations. + pub client_capabilities: ClientCapabilities, + /// Progress token, if any. + /// + /// Parsed and carried but UNUSED: recalld sends no progress + /// notifications. The field exists so it is visible that we saw the + /// token rather than silently dropping it, as the pre-2026 code did. + pub progress_token: Option, +} + +/// Failures parsing or validating `params._meta`. +#[derive(Debug, Clone, thiserror::Error)] +pub enum MetaError { + /// A required `_meta` key was absent. + #[error("{}", missing_field_message(.0))] + MissingField(&'static str), + + /// `_meta` was present but not shaped as the spec requires. + #[error("Invalid params: malformed `_meta`: {0}")] + Malformed(String), + + /// The requested protocol version is not one we serve. + #[error("Unsupported protocol version \"{requested}\"")] + UnsupportedVersion { + /// The version the client asked for. + requested: String, + }, +} + +/// Legacy clients have no way to discover the modern `_meta` requirement, +/// so the protocol-version error also names their remedy. +fn missing_field_message(field: &str) -> String { + let hint = if field == META_PROTOCOL_VERSION { + " (2026-07-28 clients); legacy clients must send `initialize` first" + } else { + "" + }; + format!("Invalid params: missing required `_meta` field \"{field}\"{hint}") +} + +impl MetaError { + /// Map to a JSON-RPC error code. + pub fn code(&self) -> i32 { + match self { + Self::MissingField(_) | Self::Malformed(_) => INVALID_PARAMS, + Self::UnsupportedVersion { .. } => UNSUPPORTED_PROTOCOL_VERSION, + } + } + + /// Structured error payload, where the spec defines one. + pub fn data(&self) -> Option { + match self { + Self::UnsupportedVersion { requested } => Some(serde_json::json!({ + "supported": SUPPORTED_MODERN_VERSIONS, + "requested": requested, + })), + _ => None, + } + } +} + +/// Parse and validate a modern client's `params._meta` block. +pub fn parse_request_meta(params: Option<&serde_json::Value>) -> Result { + let raw_meta = params + .and_then(|p| p.get("_meta")) + .ok_or(MetaError::MissingField(META_PROTOCOL_VERSION))?; + + if !raw_meta.is_object() { + return Err(MetaError::Malformed("`_meta` is not an object".to_string())); + } + + let raw: RawRequestMeta = serde_json::from_value(raw_meta.clone()) + .map_err(|e| MetaError::Malformed(e.to_string()))?; + + let protocol_version = raw + .protocol_version + .ok_or(MetaError::MissingField(META_PROTOCOL_VERSION))?; + + if !SUPPORTED_MODERN_VERSIONS.contains(&protocol_version.as_str()) { + return Err(MetaError::UnsupportedVersion { + requested: protocol_version, + }); + } + + let client_capabilities = raw + .client_capabilities + .ok_or(MetaError::MissingField(META_CLIENT_CAPABILITIES))?; + + Ok(RequestMeta { + protocol_version, + client_info: raw.client_info, + client_capabilities, + progress_token: raw.progress_token, + }) +} + +// ── Era detection ─────────────────────────────────────────────────── + +/// What a single message suggests about which era its sender speaks. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum EraHint { + /// Carries a modern `_meta` protocol version: serve it statelessly. + Modern, + /// A legacy lifecycle message (`initialize` / + /// `notifications/initialized`). + LegacyHandshake, + /// Neither marker is present; the caller decides from context. + Ambiguous, +} + +/// The era a request is actually served under. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Era { + /// Stateless MCP 2026-07-28. + Modern, + /// Session-oriented MCP 2025-06-18 and earlier. + Legacy, +} + +/// Classify a message by era. +/// +/// ORDER MATTERS: the `_meta` protocol-version key WINS over the method +/// name. A legacy client may legitimately send `params._meta` (it carries +/// `progressToken` in 2025-06-18) but never the reverse-DNS +/// `io.modelcontextprotocol/protocolVersion` key. The consequence is +/// deliberate: an `initialize` carrying modern `_meta` classifies as +/// `Modern` and is answered with -32601, because `initialize` was removed +/// in 2026-07-28. +pub fn era_hint(msg: &JsonRpcMessage) -> EraHint { + let has_modern_meta = msg + .params + .as_ref() + .and_then(|p| p.get("_meta")) + .and_then(|m| m.get(META_PROTOCOL_VERSION)) + .is_some(); + + if has_modern_meta { + return EraHint::Modern; + } + if msg.method == "initialize" || msg.method == "notifications/initialized" { + return EraHint::LegacyHandshake; + } + EraHint::Ambiguous +} + +/// Whether a method is answered identically in both eras, without +/// requiring `_meta`. +/// +/// FLAGGED DEVIATION from "every request MUST carry `_meta`": the stdio era +/// probe sends `server/discover` FIRST, so demanding a protocol version in +/// order to discover which versions are supported is circular -- and any +/// error at all (the spec forbids keying the fallback to one code) would +/// push a genuinely modern client into the legacy fallback. `ping` is +/// included because it is tolerated pre-`initialize` today and clients +/// health-check with it. If `_meta` IS supplied on these methods it is +/// still version-validated. The strict alternative is a two-line change. +pub fn is_era_neutral_method(method: &str) -> bool { + matches!(method, "server/discover" | "ping") +} + +/// Choose the version echoed by a legacy `initialize` response. +/// +/// Never returns a modern version -- claiming 2026-07-28 in an `initialize` +/// response is incoherent, since that revision has no `initialize`. Echoing +/// [`PROTOCOL_VERSION`] here would tell a 2025-06-18 client it is talking to +/// a revision where the handshake it just performed does not exist, and +/// under 2025-06-18 lifecycle rules that client SHOULD then disconnect. +pub fn negotiate_legacy_version(requested: &str) -> &'static str { + SUPPORTED_LEGACY_VERSIONS + .iter() + .copied() + .find(|v| *v == requested) + .unwrap_or(LEGACY_PROTOCOL_VERSION) +} + +// ── Tests ─────────────────────────────────────────────────────────── + +#[cfg(test)] +mod tests { + use super::*; + + fn msg(method: &str, params: Option) -> JsonRpcMessage { + JsonRpcMessage { + jsonrpc: JSONRPC_VERSION.to_string(), + id: Some(JsonRpcId::Number(1)), + method: method.to_string(), + params, + } + } + + fn modern_meta() -> serde_json::Value { + serde_json::json!({ + "_meta": { + META_PROTOCOL_VERSION: PROTOCOL_VERSION, + META_CLIENT_CAPABILITIES: {}, + } + }) + } + + /// Every MCP revision names this field `resourceTemplates`. Without + /// `rename_all = "camelCase"` this emitted `resource_templates`, which + /// no compliant client can consume. + #[test] + fn resource_templates_list_serializes_camel_case() { + let result = ResourceTemplatesListResult { + resource_templates: vec![ResourceTemplate { + uri_template: "recalld://namespaces/{name}/stats".to_string(), + name: "namespace-stats".to_string(), + description: None, + mime_type: None, + }], + }; + let json = serde_json::to_value(&result).unwrap(); + assert!( + json.get("resourceTemplates").is_some(), + "expected camelCase `resourceTemplates`, got: {json}" + ); + assert!(json.get("resource_templates").is_none()); + assert!(json["resourceTemplates"][0].get("uriTemplate").is_some()); + } + + #[test] + fn envelope_adds_result_type_and_preserves_payload() { + let env = ResultEnvelope::complete(serde_json::json!({ "tools": [] })); + let json = serde_json::to_value(&env).unwrap(); + assert_eq!(json["resultType"], "complete"); + assert_eq!(json["tools"], serde_json::json!([])); + } + + #[test] + fn envelope_cache_hints_are_top_level_not_in_meta() { + let env = ResultEnvelope::complete(serde_json::json!({ "tools": [] })) + .with_cache(CacheHints::public(CATALOG_TTL_MS)); + let json = serde_json::to_value(&env).unwrap(); + assert_eq!(json["ttlMs"], CATALOG_TTL_MS); + assert_eq!(json["cacheScope"], "public"); + assert!(json["_meta"].get("ttlMs").is_none()); + assert!(json["_meta"].get("cacheScope").is_none()); + } + + #[test] + fn envelope_omits_cache_hints_when_absent() { + let env = ResultEnvelope::complete(serde_json::json!({ "ok": true })); + let json = serde_json::to_value(&env).unwrap(); + assert!(json.get("ttlMs").is_none()); + assert!(json.get("cacheScope").is_none()); + } + + #[test] + fn envelope_meta_uses_reverse_dns_server_info_key() { + let env = ResultEnvelope::complete(serde_json::json!({})); + let json = serde_json::to_value(&env).unwrap(); + let info = &json["_meta"][META_SERVER_INFO]; + assert_eq!(info["name"], SERVER_NAME); + assert_eq!(info["version"], SERVER_VERSION); + assert!(json["_meta"].get("serverInfo").is_none()); + } + + #[test] + fn result_type_and_cache_scope_serialize_as_spec_strings() { + assert_eq!( + serde_json::to_value(ResultType::Complete).unwrap(), + "complete" + ); + assert_eq!( + serde_json::to_value(ResultType::InputRequired).unwrap(), + "input_required" + ); + assert_eq!(serde_json::to_value(CacheScope::Public).unwrap(), "public"); + assert_eq!( + serde_json::to_value(CacheScope::Private).unwrap(), + "private" + ); + } + + #[test] + fn era_hint_detects_modern_meta() { + assert_eq!( + era_hint(&msg("tools/list", Some(modern_meta()))), + EraHint::Modern + ); + } + + #[test] + fn era_hint_detects_legacy_handshake() { + assert_eq!(era_hint(&msg("initialize", None)), EraHint::LegacyHandshake); + assert_eq!( + era_hint(&msg("notifications/initialized", None)), + EraHint::LegacyHandshake + ); + } + + #[test] + fn era_hint_is_ambiguous_for_plain_requests() { + assert_eq!(era_hint(&msg("tools/list", None)), EraHint::Ambiguous); + // A legacy client may send `_meta` for progressToken; that alone is + // not a modern marker. + let legacy_meta = serde_json::json!({ "_meta": { "progressToken": 7 } }); + assert_eq!( + era_hint(&msg("tools/list", Some(legacy_meta))), + EraHint::Ambiguous + ); + } + + #[test] + fn era_hint_meta_key_wins_over_initialize_method() { + // `initialize` was removed in 2026-07-28: a message carrying modern + // `_meta` is modern regardless of its method name. + assert_eq!( + era_hint(&msg("initialize", Some(modern_meta()))), + EraHint::Modern + ); + } + + #[test] + fn parse_meta_missing_protocol_version_names_the_key() { + let err = parse_request_meta(None).unwrap_err(); + assert_eq!(err.code(), INVALID_PARAMS); + let text = err.to_string(); + assert!(text.contains(META_PROTOCOL_VERSION), "{text}"); + assert!(text.contains("initialize"), "{text}"); + + let params = serde_json::json!({ "_meta": { META_CLIENT_CAPABILITIES: {} } }); + let err = parse_request_meta(Some(¶ms)).unwrap_err(); + assert!(err.to_string().contains(META_PROTOCOL_VERSION)); + } + + #[test] + fn parse_meta_missing_client_capabilities_is_invalid_params() { + let params = serde_json::json!({ + "_meta": { META_PROTOCOL_VERSION: PROTOCOL_VERSION } + }); + let err = parse_request_meta(Some(¶ms)).unwrap_err(); + assert_eq!(err.code(), INVALID_PARAMS); + assert!(err.to_string().contains(META_CLIENT_CAPABILITIES)); + assert!(err.data().is_none()); + } + + #[test] + fn parse_meta_unsupported_version_reports_supported_and_requested() { + let params = serde_json::json!({ + "_meta": { + META_PROTOCOL_VERSION: "2027-01-01", + META_CLIENT_CAPABILITIES: {}, + } + }); + let err = parse_request_meta(Some(¶ms)).unwrap_err(); + assert_eq!(err.code(), UNSUPPORTED_PROTOCOL_VERSION); + let data = err.data().unwrap(); + assert_eq!(data["requested"], "2027-01-01"); + assert_eq!(data["supported"], serde_json::json!([PROTOCOL_VERSION])); + // supportedVersions is modern-only. + assert!(!data["supported"].to_string().contains("2025-06-18")); + } + + #[test] + fn parse_meta_accepts_full_modern_block_and_carries_progress_token() { + let params = serde_json::json!({ + "_meta": { + META_PROTOCOL_VERSION: PROTOCOL_VERSION, + META_CLIENT_INFO: { "name": "ExampleClient", "version": "1.0.0" }, + META_CLIENT_CAPABILITIES: {}, + META_LOG_LEVEL: "info", + "progressToken": "tok-1", + } + }); + let meta = parse_request_meta(Some(¶ms)).unwrap(); + assert_eq!(meta.protocol_version, PROTOCOL_VERSION); + assert_eq!(meta.client_info.unwrap().name, "ExampleClient"); + assert_eq!(meta.progress_token.unwrap(), "tok-1"); + } + + #[test] + fn json_rpc_id_null_round_trips() { + let json = serde_json::to_value(JsonRpcId::Null).unwrap(); + assert_eq!(json, serde_json::Value::Null); + assert_eq!( + serde_json::from_value::(serde_json::Value::Null).unwrap(), + JsonRpcId::Null + ); + // An explicit `"id": null` on a message still means "notification". + let msg: JsonRpcMessage = + serde_json::from_str(r#"{"jsonrpc":"2.0","id":null,"method":"x"}"#).unwrap(); + assert!(msg.id.is_none()); + // Number and String still win over Null. + assert_eq!( + serde_json::from_str::("42").unwrap(), + JsonRpcId::Number(42) + ); + assert_eq!( + serde_json::from_str::(r#""abc""#).unwrap(), + JsonRpcId::String("abc".to_string()) + ); + } + + #[test] + fn negotiate_legacy_version_echoes_known_and_falls_back() { + assert_eq!(negotiate_legacy_version("2025-06-18"), "2025-06-18"); + assert_eq!(negotiate_legacy_version("2025-03-26"), "2025-03-26"); + assert_eq!(negotiate_legacy_version("2024-11-05"), "2024-11-05"); + assert_eq!( + negotiate_legacy_version("1999-01-01"), + LEGACY_PROTOCOL_VERSION + ); + // Never the modern version, whatever is asked for. + assert_ne!(negotiate_legacy_version(PROTOCOL_VERSION), PROTOCOL_VERSION); + } + + #[test] + fn era_neutral_methods_are_discover_and_ping() { + assert!(is_era_neutral_method("server/discover")); + assert!(is_era_neutral_method("ping")); + assert!(!is_era_neutral_method("tools/list")); + assert!(!is_era_neutral_method("initialize")); + } +} diff --git a/src/mcp/server.rs b/src/mcp/server.rs index b2fc3b7..6d87d30 100644 --- a/src/mcp/server.rs +++ b/src/mcp/server.rs @@ -1,8 +1,18 @@ -//! MCP server with method dispatch and lifecycle management. +//! MCP server: stateless dispatcher plus a legacy lifecycle wrapper. //! -//! The `McpServer` manages initialization state and dispatches -//! JSON-RPC methods to the `McpHandler` trait, which is implemented -//! by the domain layer (CS-22). +//! Recalld is a dual-era MCP server (see [`crate::mcp::protocol`]): +//! +//! * [`McpDispatcher`] holds no per-connection state and takes `&self`, so +//! modern (MCP 2026-07-28) requests can be served concurrently with no +//! lock at all. +//! * [`McpServer`] wraps a dispatcher with the legacy 2025-06-18 lifecycle +//! state (`initialized`, negotiated version). Its constructor and +//! `handle_message` signatures are unchanged, so stdio and the HTTP +//! session map need no churn. +//! +//! Method dispatch is delegated to the [`McpHandler`] trait, which is +//! implemented by the domain layer (CS-22) and is UNCHANGED by the +//! 2026-07-28 migration. use std::sync::Arc; @@ -34,70 +44,156 @@ pub trait McpHandler: Send + Sync { async fn read_resource(&self, uri: &str) -> Result; } -// ── McpServer ─────────────────────────────────────────────────────── +// ── Request context and outcome ───────────────────────────────────── + +/// Everything the dispatcher needs to know about a single request. +#[derive(Debug, Clone)] +pub struct RequestContext { + /// Which protocol era this request is served under. + pub era: Era, + /// Validated modern `_meta`, when the request carried one. + pub meta: Option, + /// The version a legacy `initialize` should echo, pre-negotiated by + /// [`McpServer`] so the dispatcher stays stateless. + pub legacy_version: Option, +} + +impl RequestContext { + /// Context for a validated modern (2026-07-28) request. + pub fn modern(meta: RequestMeta) -> Self { + Self { + era: Era::Modern, + meta: Some(meta), + legacy_version: None, + } + } + + /// Context for a legacy request, with no `initialize` negotiation. + pub fn legacy() -> Self { + Self { + era: Era::Legacy, + meta: None, + legacy_version: None, + } + } + + /// Context for a legacy `initialize` echoing `version`. + pub fn legacy_initialize(version: impl Into) -> Self { + Self { + era: Era::Legacy, + meta: None, + legacy_version: Some(version.into()), + } + } +} + +/// A successful dispatch: the method-specific payload plus whether it may +/// be cached. +#[derive(Debug, Clone)] +pub struct Outcome { + /// The result payload. MUST serialize to a JSON object. + pub value: serde_json::Value, + /// SEP-2549 caching hints, when the result is cacheable. + pub cache: Option, +} + +impl Outcome { + /// A result with no caching hints. + pub fn plain(value: serde_json::Value) -> Self { + Self { value, cache: None } + } + + /// A result the client may cache for `hints.ttl_ms`. + pub fn cacheable(value: serde_json::Value, hints: CacheHints) -> Self { + Self { + value, + cache: Some(hints), + } + } +} + +// ── McpDispatcher ─────────────────────────────────────────────────── -/// MCP protocol server. +/// Stateless MCP method dispatcher. /// -/// Manages lifecycle state (initialized/not) and dispatches -/// JSON-RPC methods to the appropriate handler. Stateless -/// between requests except for the `initialized` flag. -pub struct McpServer { +/// Takes `&self` only. Modern HTTP requests call it directly, with no +/// mutex, so concurrent clients do not serialize behind one lock. +pub struct McpDispatcher { handler: Arc, - initialized: bool, } -impl McpServer { - /// Create a new MCP server with the given handler. +impl McpDispatcher { + /// Create a dispatcher over the given domain handler. pub fn new(handler: Arc) -> Self { - Self { - handler, - initialized: false, - } + Self { handler } } - /// Dispatch a JSON-RPC message to the appropriate handler. + /// Borrow the underlying domain handler. + pub fn handler(&self) -> &Arc { + &self.handler + } + + /// Dispatch one JSON-RPC message. /// /// Returns `Some(response)` for requests, `None` for notifications. - pub async fn handle_message(&mut self, msg: JsonRpcMessage) -> Option { - let id = msg.id.clone(); - - // Notifications (no id) don't get a response. - let id = match id { + /// Every success is wrapped in a [`ResultEnvelope`] at the single `Ok` + /// arm below, so no result type can forget `resultType` or `_meta`. + pub async fn dispatch( + &self, + ctx: &RequestContext, + msg: JsonRpcMessage, + ) -> Option { + let id = match msg.id.clone() { Some(id) => id, None => { - self.handle_notification(&msg.method, msg.params).await; + self.handle_notification(&msg.method); return None; } }; - // Before initialization, only `initialize` and `ping` are allowed. - if !self.initialized && msg.method != "initialize" && msg.method != "ping" { - return Some(JsonRpcResponse::error( - id, - INVALID_REQUEST, - "Server not initialized. Send 'initialize' first.", - )); - } + let method = msg.method.as_str(); + let result = match (ctx.era, method) { + // Era-neutral: answered identically on both paths. + (_, "server/discover") => self.handle_discover(), - let result = match msg.method.as_str() { - "initialize" => self.handle_initialize(msg.params).await, - "ping" => self.handle_ping().await, - "tools/list" => self.handle_tools_list().await, - "tools/call" => self.handle_tools_call(msg.params).await, - "resources/list" => self.handle_resources_list().await, - "resources/templates/list" => self.handle_resource_templates_list().await, - "resources/read" => self.handle_resources_read(msg.params).await, - other => Err(DispatchError::MethodNotFound(other.to_string())), + // Removed by 2026-07-28 (SEP-2575): no handshake, no ping. + (Era::Modern, "initialize") | (Era::Modern, "ping") => { + Err(DispatchError::MethodNotFound(method.to_string())) + } + + (Era::Legacy, "initialize") => self.handle_initialize(ctx, msg.params), + (Era::Legacy, "ping") => Ok(Outcome::plain(serde_json::json!({}))), + + (_, "tools/list") => self.handle_tools_list(), + (_, "tools/call") => self.handle_tools_call(msg.params).await, + (_, "resources/list") => self.handle_resources_list(), + (_, "resources/templates/list") => self.handle_resource_templates_list(), + (_, "resources/read") => self.handle_resources_read(msg.params).await, + + (_, other) => Err(DispatchError::MethodNotFound(other.to_string())), }; Some(match result { - Ok(value) => JsonRpcResponse::success(id, value), - Err(e) => JsonRpcResponse::error(id, e.code(), e.to_string()), + Ok(outcome) => { + let mut env = ResultEnvelope::complete(outcome.value); + if let Some(hints) = outcome.cache { + env = env.with_cache(hints); + } + match serde_json::to_value(env) { + Ok(value) => JsonRpcResponse::success(id, value), + Err(e) => JsonRpcResponse::error( + id, + INTERNAL_ERROR, + format!("Serialization error: {e}"), + ), + } + } + Err(e) => JsonRpcResponse::error_with_data(id, e.code(), e.to_string(), e.data()), }) } /// Handle a notification (no response expected). - async fn handle_notification(&mut self, method: &str, _params: Option) { + fn handle_notification(&self, method: &str) { match method { "notifications/initialized" => { tracing::info!("Client confirmed initialization"); @@ -111,16 +207,36 @@ impl McpServer { } } - async fn handle_initialize( - &mut self, + /// `server/discover` -- the 2026-07-28 replacement for `initialize`. + /// + /// Answered on both eras and, deliberately, without requiring `_meta`; + /// see [`is_era_neutral_method`] for the reasoning. + fn handle_discover(&self) -> Result { + let result = DiscoverResult { + supported_versions: SUPPORTED_MODERN_VERSIONS + .iter() + .map(|v| (*v).to_string()) + .collect(), + capabilities: server_capabilities(), + instructions: Some(SERVER_INSTRUCTIONS.to_string()), + }; + Ok(Outcome::cacheable( + serde_json::to_value(result)?, + CacheHints::public(CATALOG_TTL_MS), + )) + } + + fn handle_initialize( + &self, + ctx: &RequestContext, params: Option, - ) -> Result { + ) -> Result { let params: InitializeParams = params .map(serde_json::from_value) .transpose() .map_err(|e| DispatchError::InvalidParams(e.to_string()))? .unwrap_or_else(|| InitializeParams { - protocol_version: PROTOCOL_VERSION.to_string(), + protocol_version: LEGACY_PROTOCOL_VERSION.to_string(), capabilities: ClientCapabilities::default(), client_info: Implementation { name: "unknown".to_string(), @@ -132,54 +248,43 @@ impl McpServer { client = %params.client_info.name, version = %params.client_info.version, protocol = %params.protocol_version, - "MCP client initializing" + "MCP client initializing (legacy era)" ); - self.initialized = true; + // NEVER PROTOCOL_VERSION. Echoing 2026-07-28 here would advertise a + // revision in which `initialize` does not exist. + let negotiated = ctx + .legacy_version + .clone() + .unwrap_or_else(|| negotiate_legacy_version(¶ms.protocol_version).to_string()); let result = InitializeResult { - protocol_version: PROTOCOL_VERSION.to_string(), - capabilities: ServerCapabilities { - tools: Some(ToolsCapability { - list_changed: false, - }), - resources: Some(ResourcesCapability { - subscribe: false, - list_changed: false, - }), - }, + protocol_version: negotiated, + capabilities: server_capabilities(), server_info: Implementation { name: SERVER_NAME.to_string(), version: SERVER_VERSION.to_string(), }, - instructions: Some( - "Recalld is an AI memory system with human-like forgetting. \ - Use store_memory to save observations, recall_memories to \ - search by semantic similarity, and reinforce_memory to \ - strengthen useful memories. Memories decay naturally over \ - time unless reinforced." - .to_string(), - ), + instructions: Some(SERVER_INSTRUCTIONS.to_string()), }; - Ok(serde_json::to_value(result)?) + Ok(Outcome::plain(serde_json::to_value(result)?)) } - async fn handle_ping(&self) -> Result { - Ok(serde_json::json!({})) - } - - async fn handle_tools_list(&self) -> Result { + fn handle_tools_list(&self) -> Result { let result = ToolsListResult { tools: self.handler.tools(), }; - Ok(serde_json::to_value(result)?) + Ok(Outcome::cacheable( + serde_json::to_value(result)?, + CacheHints::public(CATALOG_TTL_MS), + )) } async fn handle_tools_call( &self, params: Option, - ) -> Result { + ) -> Result { let params: ToolCallParams = params .map(serde_json::from_value) .transpose() @@ -187,27 +292,33 @@ impl McpServer { .ok_or_else(|| DispatchError::InvalidParams("Missing params".to_string()))?; let result = self.handler.call_tool(¶ms.name, params.arguments).await; - Ok(serde_json::to_value(result)?) + Ok(Outcome::plain(serde_json::to_value(result)?)) } - async fn handle_resources_list(&self) -> Result { + fn handle_resources_list(&self) -> Result { let result = ResourcesListResult { resources: self.handler.resources(), }; - Ok(serde_json::to_value(result)?) + Ok(Outcome::cacheable( + serde_json::to_value(result)?, + CacheHints::public(CATALOG_TTL_MS), + )) } - async fn handle_resource_templates_list(&self) -> Result { + fn handle_resource_templates_list(&self) -> Result { let result = ResourceTemplatesListResult { resource_templates: self.handler.resource_templates(), }; - Ok(serde_json::to_value(result)?) + Ok(Outcome::cacheable( + serde_json::to_value(result)?, + CacheHints::public(CATALOG_TTL_MS), + )) } async fn handle_resources_read( &self, params: Option, - ) -> Result { + ) -> Result { let params: ResourceReadParams = params .map(serde_json::from_value) .transpose() @@ -219,7 +330,135 @@ impl McpServer { .read_resource(¶ms.uri) .await .map_err(DispatchError::Internal)?; - Ok(serde_json::to_value(result)?) + + // Live instance state, not a static catalog: private and + // immediately stale. + Ok(Outcome::cacheable( + serde_json::to_value(result)?, + CacheHints::private(RESOURCE_READ_TTL_MS), + )) + } +} + +// ── McpServer ─────────────────────────────────────────────────────── + +/// MCP protocol server with legacy lifecycle state. +/// +/// One instance per legacy connection (a stdio process, or an HTTP +/// `Mcp-Session-Id` session). Modern requests arriving here are served +/// statelessly and NEVER consult `initialized`. +pub struct McpServer { + dispatcher: Arc, + initialized: bool, + negotiated_version: Option, +} + +impl McpServer { + /// Create a new MCP server with the given handler. + pub fn new(handler: Arc) -> Self { + Self::from_dispatcher(Arc::new(McpDispatcher::new(handler))) + } + + /// Create a legacy session sharing an existing stateless dispatcher. + pub fn from_dispatcher(dispatcher: Arc) -> Self { + Self { + dispatcher, + initialized: false, + negotiated_version: None, + } + } + + /// Whether a legacy `initialize` has succeeded on this connection. + pub fn is_initialized(&self) -> bool { + self.initialized + } + + /// The protocol version echoed to this legacy client, if any. + pub fn negotiated_version(&self) -> Option<&str> { + self.negotiated_version.as_deref() + } + + /// Dispatch a JSON-RPC message to the appropriate handler. + /// + /// Returns `Some(response)` for requests, `None` for notifications. + /// The era is decided PER MESSAGE: an open connection is not a session + /// for a modern client. + pub async fn handle_message(&mut self, msg: JsonRpcMessage) -> Option { + // Notifications never get a response and carry no lifecycle state, + // so they skip era validation entirely and are simply logged. + if msg.id.is_none() { + return self + .dispatcher + .dispatch(&RequestContext::legacy(), msg) + .await; + } + + let ctx = match self.context_for(&msg) { + Ok(ctx) => ctx, + Err(resp) => return Some(*resp), + }; + + let is_legacy_initialize = ctx.era == Era::Legacy && msg.method == "initialize"; + let response = self.dispatcher.dispatch(&ctx, msg).await; + + if is_legacy_initialize && response.as_ref().is_some_and(|resp| resp.error.is_none()) { + self.initialized = true; + self.negotiated_version = ctx.legacy_version; + } + + response + } + + /// Classify a request and build its dispatch context. + /// + /// The error is boxed only to keep the `Result` small. + fn context_for(&self, msg: &JsonRpcMessage) -> Result> { + let id = msg.id.clone().unwrap_or(JsonRpcId::Null); + + match era_hint(msg) { + EraHint::Modern => match parse_request_meta(msg.params.as_ref()) { + Ok(meta) => Ok(RequestContext::modern(meta)), + Err(e) => { + let e = DispatchError::Meta(e); + Err(Box::new(JsonRpcResponse::error_with_data( + id, + e.code(), + e.to_string(), + e.data(), + ))) + } + }, + EraHint::LegacyHandshake => { + let requested = msg + .params + .as_ref() + .and_then(|p| p.get("protocolVersion")) + .and_then(|v| v.as_str()) + .unwrap_or(LEGACY_PROTOCOL_VERSION); + Ok(RequestContext::legacy_initialize(negotiate_legacy_version( + requested, + ))) + } + EraHint::Ambiguous => { + // Era-neutral methods are answered without any handshake, + // and an already-initialized legacy connection keeps + // working exactly as before. + if is_era_neutral_method(&msg.method) || self.initialized { + Ok(RequestContext::legacy()) + } else { + // Replaces the old -32600 "Server not initialized" + // gate: this message names BOTH remedies (modern + // `_meta`, or a legacy `initialize`). + let e = DispatchError::Meta(MetaError::MissingField(META_PROTOCOL_VERSION)); + Err(Box::new(JsonRpcResponse::error_with_data( + id, + e.code(), + e.to_string(), + e.data(), + ))) + } + } + } } } @@ -243,6 +482,10 @@ pub enum DispatchError { /// A serialization or deserialization error occurred. #[error("Serialization error: {0}")] Serialization(#[from] serde_json::Error), + + /// The modern per-request `_meta` block was missing or invalid. + #[error(transparent)] + Meta(#[from] MetaError), } impl DispatchError { @@ -252,6 +495,497 @@ impl DispatchError { Self::MethodNotFound(_) => METHOD_NOT_FOUND, Self::InvalidParams(_) => INVALID_PARAMS, Self::Internal(_) | Self::Serialization(_) => INTERNAL_ERROR, + Self::Meta(e) => e.code(), + } + } + + /// Structured error payload, where one is defined. + pub fn data(&self) -> Option { + match self { + Self::Meta(e) => e.data(), + _ => None, + } + } +} + +// ── Tests ─────────────────────────────────────────────────────────── + +#[cfg(test)] +pub(crate) mod tests { + use super::*; + + /// Minimal in-memory handler, reused by the HTTP transport tests. + pub(crate) struct StubHandler; + + #[async_trait] + impl McpHandler for StubHandler { + fn tools(&self) -> Vec { + vec![ToolInfo { + name: "store_memory".to_string(), + title: None, + description: "Store a memory".to_string(), + input_schema: serde_json::json!({ "type": "object" }), + annotations: None, + }] } + + fn resources(&self) -> Vec { + vec![ResourceInfo { + uri: "recalld://health".to_string(), + name: "health".to_string(), + description: None, + mime_type: Some("application/json".to_string()), + }] + } + + fn resource_templates(&self) -> Vec { + vec![ResourceTemplate { + uri_template: "recalld://namespaces/{name}/stats".to_string(), + name: "namespace-stats".to_string(), + description: None, + mime_type: Some("application/json".to_string()), + }] + } + + async fn call_tool(&self, name: &str, _arguments: serde_json::Value) -> ToolCallResult { + ToolCallResult::text(format!("called {name}")) + } + + async fn read_resource(&self, uri: &str) -> Result { + Ok(ResourceReadResult { + contents: vec![ResourceContent { + uri: uri.to_string(), + mime_type: Some("application/json".to_string()), + text: Some("{}".to_string()), + }], + }) + } + } + + pub(crate) fn stub_handler() -> Arc { + Arc::new(StubHandler) + } + + pub(crate) fn server() -> McpServer { + McpServer::new(stub_handler()) + } + + /// Build a request carrying a valid modern `_meta` block. + pub(crate) fn modern_request( + id: i64, + method: &str, + extra: serde_json::Value, + ) -> JsonRpcMessage { + let mut params = serde_json::json!({ + "_meta": { + META_PROTOCOL_VERSION: PROTOCOL_VERSION, + META_CLIENT_INFO: { "name": "TestClient", "version": "1.0.0" }, + META_CLIENT_CAPABILITIES: {}, + } + }); + if let (Some(dst), Some(src)) = (params.as_object_mut(), extra.as_object()) { + for (k, v) in src { + dst.insert(k.clone(), v.clone()); + } + } + JsonRpcMessage { + jsonrpc: JSONRPC_VERSION.to_string(), + id: Some(JsonRpcId::Number(id)), + method: method.to_string(), + params: Some(params), + } + } + + pub(crate) fn plain_request( + id: i64, + method: &str, + params: Option, + ) -> JsonRpcMessage { + JsonRpcMessage { + jsonrpc: JSONRPC_VERSION.to_string(), + id: Some(JsonRpcId::Number(id)), + method: method.to_string(), + params, + } + } + + pub(crate) fn initialize_request(id: i64, version: &str) -> JsonRpcMessage { + plain_request( + id, + "initialize", + Some(serde_json::json!({ + "protocolVersion": version, + "capabilities": {}, + "clientInfo": { "name": "LegacyClient", "version": "1.0.0" }, + })), + ) + } + + fn ok(resp: &JsonRpcResponse) -> &serde_json::Value { + assert!( + resp.error.is_none(), + "expected success, got {:?}", + resp.error + ); + resp.result.as_ref().unwrap() + } + + // ── server/discover ───────────────────────────────────────────── + + #[tokio::test] + async fn discover_reports_versions_capabilities_and_instructions() { + let mut srv = server(); + let resp = srv + .handle_message(plain_request(1, "server/discover", None)) + .await + .unwrap(); + let result = ok(&resp); + assert_eq!( + result["supportedVersions"], + serde_json::json!([PROTOCOL_VERSION]) + ); + assert_eq!(result["capabilities"]["tools"]["listChanged"], false); + assert_eq!(result["capabilities"]["resources"]["subscribe"], false); + assert!( + result["instructions"] + .as_str() + .unwrap() + .contains("Recalld is an AI memory system") + ); + } + + #[tokio::test] + async fn discover_supported_versions_are_modern_only() { + let mut srv = server(); + let resp = srv + .handle_message(plain_request(1, "server/discover", None)) + .await + .unwrap(); + let versions = ok(&resp)["supportedVersions"].clone(); + for legacy in SUPPORTED_LEGACY_VERSIONS { + assert!( + !versions.to_string().contains(legacy), + "legacy version {legacy} must not appear in supportedVersions" + ); + } + } + + #[tokio::test] + async fn discover_carries_public_cache_hints_and_result_type() { + let mut srv = server(); + let resp = srv + .handle_message(plain_request(1, "server/discover", None)) + .await + .unwrap(); + let result = ok(&resp); + assert_eq!(result["resultType"], "complete"); + assert_eq!(result["ttlMs"], CATALOG_TTL_MS); + assert_eq!(result["cacheScope"], "public"); + } + + /// The stdio era probe sends `server/discover` FIRST, before it knows + /// which era the server speaks -- so it cannot yet supply `_meta`. + #[tokio::test] + async fn discover_is_answered_without_meta() { + let mut srv = server(); + let resp = srv + .handle_message(plain_request(1, "server/discover", None)) + .await + .unwrap(); + assert!(resp.error.is_none()); + } + + #[tokio::test] + async fn discover_with_unsupported_meta_version_is_rejected() { + let mut srv = server(); + let msg = plain_request( + 1, + "server/discover", + Some(serde_json::json!({ + "_meta": { + META_PROTOCOL_VERSION: "2027-01-01", + META_CLIENT_CAPABILITIES: {}, + } + })), + ); + let resp = srv.handle_message(msg).await.unwrap(); + let err = resp.error.unwrap(); + assert_eq!(err.code, UNSUPPORTED_PROTOCOL_VERSION); + assert_eq!(err.data.unwrap()["requested"], "2027-01-01"); + } + + // ── Modern era ────────────────────────────────────────────────── + + #[tokio::test] + async fn modern_tools_list_needs_no_initialize() { + let mut srv = server(); + assert!(!srv.is_initialized()); + let resp = srv + .handle_message(modern_request(1, "tools/list", serde_json::json!({}))) + .await + .unwrap(); + assert_eq!(ok(&resp)["tools"][0]["name"], "store_memory"); + assert!(!srv.is_initialized(), "modern requests never initialize"); + } + + #[tokio::test] + async fn modern_tools_call_dispatches_to_handler() { + let mut srv = server(); + let msg = modern_request( + 1, + "tools/call", + serde_json::json!({ "name": "store_memory", "arguments": {} }), + ); + let resp = srv.handle_message(msg).await.unwrap(); + assert_eq!(ok(&resp)["content"][0]["text"], "called store_memory"); + } + + #[tokio::test] + async fn modern_ping_is_method_not_found() { + let mut srv = server(); + let resp = srv + .handle_message(modern_request(1, "ping", serde_json::json!({}))) + .await + .unwrap(); + let err = resp.error.unwrap(); + assert_eq!(err.code, METHOD_NOT_FOUND); + assert!(err.message.contains("ping")); + } + + #[tokio::test] + async fn modern_initialize_is_method_not_found() { + let mut srv = server(); + let resp = srv + .handle_message(modern_request(1, "initialize", serde_json::json!({}))) + .await + .unwrap(); + assert_eq!(resp.error.unwrap().code, METHOD_NOT_FOUND); + } + + // ── Legacy era ────────────────────────────────────────────────── + + #[tokio::test] + async fn legacy_ping_before_initialize_still_returns_empty_object() { + let mut srv = server(); + let resp = srv + .handle_message(plain_request(1, "ping", None)) + .await + .unwrap(); + let result = ok(&resp); + assert_eq!(result["resultType"], "complete"); + assert!(result.get("ttlMs").is_none()); + } + + #[tokio::test] + async fn legacy_tools_list_before_initialize_names_the_meta_key() { + let mut srv = server(); + let resp = srv + .handle_message(plain_request(1, "tools/list", None)) + .await + .unwrap(); + let err = resp.error.unwrap(); + assert_eq!(err.code, INVALID_PARAMS); + assert!( + err.message.contains(META_PROTOCOL_VERSION), + "{}", + err.message + ); + assert!(err.message.contains("initialize"), "{}", err.message); + } + + #[tokio::test] + async fn legacy_initialize_echoes_the_requested_version() { + for requested in ["2025-06-18", "2025-03-26", "2024-11-05"] { + let mut srv = server(); + let resp = srv + .handle_message(initialize_request(1, requested)) + .await + .unwrap(); + assert_eq!(ok(&resp)["protocolVersion"], requested); + assert!(srv.is_initialized()); + assert_eq!(srv.negotiated_version(), Some(requested)); + } + } + + #[tokio::test] + async fn legacy_initialize_falls_back_for_unknown_version() { + let mut srv = server(); + let resp = srv + .handle_message(initialize_request(1, "1999-01-01")) + .await + .unwrap(); + assert_eq!(ok(&resp)["protocolVersion"], LEGACY_PROTOCOL_VERSION); + } + + /// THE guard against the issue's harmful premise: bumping + /// `PROTOCOL_VERSION` must never leak into an `initialize` response, + /// because 2026-07-28 has no `initialize`. + #[tokio::test] + async fn legacy_initialize_never_echoes_2026_07_28() { + for requested in [ + "2025-06-18", + "2025-03-26", + "2024-11-05", + "1999-01-01", + PROTOCOL_VERSION, + ] { + let mut srv = server(); + let resp = srv + .handle_message(initialize_request(1, requested)) + .await + .unwrap(); + let echoed = ok(&resp)["protocolVersion"].as_str().unwrap().to_string(); + assert_ne!( + echoed, PROTOCOL_VERSION, + "initialize echoed the modern version for request {requested}" + ); + assert!(SUPPORTED_LEGACY_VERSIONS.contains(&echoed.as_str())); + } + } + + #[tokio::test] + async fn legacy_session_works_after_initialize() { + let mut srv = server(); + srv.handle_message(initialize_request(1, "2025-06-18")) + .await + .unwrap(); + let resp = srv + .handle_message(plain_request(2, "tools/list", None)) + .await + .unwrap(); + assert_eq!(ok(&resp)["tools"][0]["name"], "store_memory"); + } + + // ── Envelope invariants across every method ───────────────────── + + async fn all_results() -> Vec<(&'static str, serde_json::Value)> { + let mut out = Vec::new(); + let mut srv = server(); + + // Legacy-era results. + for (method, params) in [ + ("initialize", None), + ("ping", None), + ("server/discover", None), + ] { + let mut s = server(); + let msg = if method == "initialize" { + initialize_request(1, "2025-06-18") + } else { + plain_request(1, method, params) + }; + let resp = s.handle_message(msg).await.unwrap(); + out.push((method, resp.result.unwrap())); + } + + // Modern-era results. + srv.handle_message(initialize_request(1, "2025-06-18")) + .await + .unwrap(); + for (method, extra) in [ + ("tools/list", serde_json::json!({})), + ( + "tools/call", + serde_json::json!({ "name": "store_memory", "arguments": {} }), + ), + ("resources/list", serde_json::json!({})), + ("resources/templates/list", serde_json::json!({})), + ( + "resources/read", + serde_json::json!({ "uri": "recalld://health" }), + ), + ] { + let resp = srv + .handle_message(modern_request(2, method, extra)) + .await + .unwrap(); + out.push((method, resp.result.unwrap())); + } + out + } + + #[tokio::test] + async fn every_result_carries_result_type_complete() { + for (method, result) in all_results().await { + assert_eq!(result["resultType"], "complete", "method {method}"); + } + } + + #[tokio::test] + async fn every_result_carries_server_info_meta() { + for (method, result) in all_results().await { + assert_eq!( + result["_meta"][META_SERVER_INFO]["name"], SERVER_NAME, + "method {method}" + ); + } + } + + #[tokio::test] + async fn cache_hints_match_the_spec_table() { + let expected: &[(&str, Option<(u64, &str)>)] = &[ + ("server/discover", Some((CATALOG_TTL_MS, "public"))), + ("initialize", None), + ("ping", None), + ("tools/list", Some((CATALOG_TTL_MS, "public"))), + ("tools/call", None), + ("resources/list", Some((CATALOG_TTL_MS, "public"))), + ("resources/templates/list", Some((CATALOG_TTL_MS, "public"))), + ("resources/read", Some((RESOURCE_READ_TTL_MS, "private"))), + ]; + let results = all_results().await; + for (method, want) in expected { + let (_, result) = results + .iter() + .find(|(m, _)| m == method) + .unwrap_or_else(|| panic!("no result captured for {method}")); + match want { + Some((ttl, scope)) => { + assert_eq!(result["ttlMs"], *ttl, "method {method}"); + assert_eq!(result["cacheScope"], *scope, "method {method}"); + } + None => { + assert!(result.get("ttlMs").is_none(), "method {method}"); + assert!(result.get("cacheScope").is_none(), "method {method}"); + } + } + } + } + + #[tokio::test] + async fn notifications_produce_no_response() { + let mut srv = server(); + for method in [ + "notifications/initialized", + "notifications/cancelled", + "notifications/unknown", + ] { + let msg = JsonRpcMessage { + jsonrpc: JSONRPC_VERSION.to_string(), + id: None, + method: method.to_string(), + params: None, + }; + assert!(srv.handle_message(msg).await.is_none()); + } + } + + #[tokio::test] + async fn unknown_method_is_method_not_found_in_both_eras() { + let mut srv = server(); + let resp = srv + .handle_message(modern_request(1, "prompts/list", serde_json::json!({}))) + .await + .unwrap(); + assert_eq!(resp.error.unwrap().code, METHOD_NOT_FOUND); + + srv.handle_message(initialize_request(2, "2025-06-18")) + .await + .unwrap(); + let resp = srv + .handle_message(plain_request(3, "prompts/list", None)) + .await + .unwrap(); + assert_eq!(resp.error.unwrap().code, METHOD_NOT_FOUND); } } diff --git a/src/mcp/transport.rs b/src/mcp/transport.rs index 8be99de..fa4128a 100644 --- a/src/mcp/transport.rs +++ b/src/mcp/transport.rs @@ -3,6 +3,15 @@ //! Reads newline-delimited JSON-RPC from stdin, dispatches to the //! MCP server, and writes responses to stdout. All non-protocol //! output (logging, debug) goes to stderr, never stdout. +//! +//! stdio is **dual-era**, and the era is decided PER MESSAGE, not per +//! process. Under MCP 2026-07-28 an open connection is explicitly not a +//! conversation or a session: a stdio process may serve a modern client +//! that never sends `initialize` (it probes `server/discover` first and +//! carries its protocol version in every `params._meta`), a legacy client +//! that does, or both. The single `McpServer` held for the process +//! lifetime therefore holds legacy lifecycle state only; modern requests +//! never consult it. use std::sync::Arc; @@ -12,6 +21,38 @@ use crate::mcp::McpError; use crate::mcp::protocol::*; use crate::mcp::server::McpServer; +/// Handle one line of stdin, returning the line to write to stdout. +/// +/// Returns `Ok(None)` for blank input and for notifications, which produce +/// no response. +pub(crate) async fn process_line( + server: &Arc>, + line: &str, +) -> Result, McpError> { + let line = line.trim(); + if line.is_empty() { + return Ok(None); + } + + let message: JsonRpcMessage = match serde_json::from_str(line) { + Ok(msg) => msg, + Err(e) => { + // We have no id to respond with, but JSON-RPC still wants a + // response -- so say so explicitly with a null id rather than + // fabricating `0` and colliding with a real request. + let err_resp = + JsonRpcResponse::error(JsonRpcId::Null, PARSE_ERROR, format!("Parse error: {e}")); + return Ok(Some(serde_json::to_string(&err_resp)?)); + } + }; + + let mut server = server.lock().await; + match server.handle_message(message).await { + Some(response) => Ok(Some(serde_json::to_string(&response)?)), + None => Ok(None), + } +} + /// Run the MCP server over stdin/stdout. /// /// Reads one JSON-RPC message per line from stdin, dispatches to @@ -28,44 +69,112 @@ pub async fn run_stdio(server: Arc>) -> Result<(), tracing::info!("MCP stdio transport started"); while let Some(line) = lines.next_line().await? { - let line = line.trim().to_string(); - if line.is_empty() { - continue; - } - - let message: JsonRpcMessage = match serde_json::from_str(&line) { - Ok(msg) => msg, - Err(e) => { - // Parse error -- we don't have an ID to respond with, - // but JSON-RPC says we should still respond. - let err_resp = JsonRpcResponse { - jsonrpc: JSONRPC_VERSION.to_string(), - id: JsonRpcId::Number(0), - result: None, - error: Some(JsonRpcError { - code: PARSE_ERROR, - message: format!("Parse error: {e}"), - data: None, - }), - }; - let json = serde_json::to_string(&err_resp)?; - stdout.write_all(json.as_bytes()).await?; - stdout.write_all(b"\n").await?; - stdout.flush().await?; - continue; - } - }; - - let mut server = server.lock().await; - if let Some(response) = server.handle_message(message).await { - let json = serde_json::to_string(&response)?; + if let Some(json) = process_line(&server, &line).await? { stdout.write_all(json.as_bytes()).await?; stdout.write_all(b"\n").await?; stdout.flush().await?; } - // Notifications (id == None) produce no response. } tracing::info!("MCP stdio transport closed (stdin EOF)"); Ok(()) } + +// ── Tests ─────────────────────────────────────────────────────────── + +#[cfg(test)] +mod tests { + use super::*; + use crate::mcp::server::tests::stub_handler; + + fn server() -> Arc> { + Arc::new(tokio::sync::Mutex::new(McpServer::new(stub_handler()))) + } + + async fn line(server: &Arc>, raw: &str) -> serde_json::Value { + let out = process_line(server, raw).await.unwrap().unwrap(); + serde_json::from_str(&out).unwrap() + } + + #[tokio::test] + async fn parse_error_responds_with_a_null_id() { + let srv = server(); + let resp = line(&srv, "{not json").await; + assert_eq!(resp["error"]["code"], PARSE_ERROR); + assert!(resp["id"].is_null(), "expected null id, got {}", resp["id"]); + } + + /// A modern client probes `server/discover` first, then works without + /// ever sending `initialize` -- a stdio process is not a session. + #[tokio::test] + async fn modern_client_discovers_then_lists_tools_without_initialize() { + let srv = server(); + + let discover = line( + &srv, + r#"{"jsonrpc":"2.0","id":1,"method":"server/discover"}"#, + ) + .await; + assert_eq!( + discover["result"]["supportedVersions"], + serde_json::json!([PROTOCOL_VERSION]) + ); + + let raw = serde_json::json!({ + "jsonrpc": "2.0", + "id": 2, + "method": "tools/list", + "params": { + "_meta": { + META_PROTOCOL_VERSION: PROTOCOL_VERSION, + META_CLIENT_CAPABILITIES: {}, + } + } + }) + .to_string(); + let tools = line(&srv, &raw).await; + assert_eq!(tools["result"]["tools"][0]["name"], "store_memory"); + assert_eq!(tools["result"]["resultType"], "complete"); + assert!(!srv.lock().await.is_initialized()); + } + + #[tokio::test] + async fn legacy_client_initializes_then_lists_tools() { + let srv = server(); + + let raw = serde_json::json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "initialize", + "params": { + "protocolVersion": "2025-06-18", + "capabilities": {}, + "clientInfo": { "name": "LegacyClient", "version": "1.0.0" } + } + }) + .to_string(); + let init = line(&srv, &raw).await; + assert_eq!(init["result"]["protocolVersion"], "2025-06-18"); + assert_ne!(init["result"]["protocolVersion"], PROTOCOL_VERSION); + + // The initialized notification produces no output. + assert!( + process_line( + &srv, + r#"{"jsonrpc":"2.0","method":"notifications/initialized"}"# + ) + .await + .unwrap() + .is_none() + ); + + let tools = line(&srv, r#"{"jsonrpc":"2.0","id":2,"method":"tools/list"}"#).await; + assert_eq!(tools["result"]["tools"][0]["name"], "store_memory"); + } + + #[tokio::test] + async fn blank_lines_are_ignored() { + let srv = server(); + assert!(process_line(&srv, " ").await.unwrap().is_none()); + } +} From 8d78a5d636b1720526dc47cd411cae1e89c30a5a Mon Sep 17 00:00:00 2001 From: Caleb Evans Date: Tue, 11 Aug 2026 01:26:23 -0600 Subject: [PATCH 5/8] fix: apply supersedes and parentId, or fail loudly A bug report found that supersedes never appeared to do anything. The parameter was in fact wired up, but every way it could fail was swallowed: a non-existent target produced GraphError::MemoryNotFound which was logged as "non-fatal" and discarded, a duplicate did the same, and a persistence failure did the same again. The store returned success in all three cases, and nothing in the response ever indicated whether a link had been made. parentId was worse -- it was parsed into StoreInput and then read by nothing at all, on both the MCP and HTTP batch paths. It has never created an edge. The report asked that the recall side be verified independently rather than assumed broken or fine. It was, and it had two real defects. - Validate the target before anything is written: it must exist, live in the same namespace, and not be the memory itself. The record used to be committed before the edge was attempted, so failing afterwards would have meant unwinding six subsystems. Checking first makes failing the whole store free, and costs nothing but a lookup. - Treat a duplicate edge as success, not failure. Supersedes is a desired end state, so asking for it twice is idempotent, which also makes client retries safe. - Report the outcome. The store response now carries the target and one of applied / alreadyApplied / appliedNotDurable / failed, so a caller can tell the difference between a link that was made, one already present, one that will not survive restart, and one that was not made. A persistence failure is not an error: the memory really was stored and the edge really is live this session, and saying otherwise would be a lie about a store that succeeded. - Make parentId actually link, via the same helper, so the two kinds of caller-specified edge cannot drift apart again. - Recall: stop removing a superseded memory before its replacement has been validated. Previously the original was dropped first, and if the replacement had decayed to ghost or tombstone, or failed to load, the result set simply lost it with nothing in its place -- permanently unrecallable. - Recall: run injected replacements through the same filters as every other candidate. They bypassed all of them, so a replacement in a different namespace leaked into results. That is a data-isolation defect, not a ranking nicety. - Guard the supersedes chain walk with a visited set, so a cycle returns None instead of an arbitrary node after ten hops. - HTTP batch: every silent `continue` now records a failure entry. Previously items vanished from the response with no explanation. - Correct the docs, which claimed the old memory is "deprioritized" in search. It is removed and replaced. Breaking: a non-existent or cross-namespace supersedes/parentId now fails the store (MCP isError, HTTP 422) where it previously returned success. Callers passing stale ids will see errors immediately. New response fields are all optional and skipped when empty, so happy-path payloads are unchanged, and StoredMemory.supersedes carries serde(default) so a new client can still decode an old daemon's response. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01JNhUehmChiQKJUjv3QPnkH --- docs/architecture.md | 5 +- docs/benchmark.md | 11 +- docs/guide.md | 36 ++- docs/mcp.md | 38 ++- src/api/adapters.rs | 86 +++++- src/api/handlers.rs | 170 ++++++++-- src/api/models.rs | 26 ++ src/api/state.rs | 27 ++ src/graph/explicit_links.rs | 601 ++++++++++++++++++++++++++++++++++++ src/graph/mod.rs | 6 + src/graph/structure.rs | 183 +++++++++++ src/mcp/bridge.rs | 97 ++++++ src/mcp/bridge_adapters.rs | 152 +++++---- src/mcp/tools.rs | 9 +- src/model/edge.rs | 71 +++++ src/model/mod.rs | 2 +- src/search/adapters.rs | 23 +- src/search/pipeline.rs | 597 +++++++++++++++++++++++++++++++---- src/serialization/json.rs | 8 + 19 files changed, 1959 insertions(+), 189 deletions(-) create mode 100644 src/graph/explicit_links.rs diff --git a/docs/architecture.md b/docs/architecture.md index a3a6884..c7a98a0 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -471,7 +471,10 @@ The `QueryEngine` orchestrates a 9-stage search pipeline. All subsystem dependen v [8a] Compute composite score [8b] Apply temporal boost (Gaussian falloff around query time range) - [8c] Resolve supersedes chains (replace outdated with current version) + [8c] Resolve supersedes chains (replace outdated with current version + WHEN THE REPLACEMENT PASSES THE QUERY'S FILTERS; otherwise the + original is kept, so a result is never dropped with nothing in + its place) [8d] Sort descending, truncate to limit | v diff --git a/docs/benchmark.md b/docs/benchmark.md index 4ccda47..e36ed95 100644 --- a/docs/benchmark.md +++ b/docs/benchmark.md @@ -83,7 +83,12 @@ outputs zero or more structured memories to store. - `entities` (required): people, places, proper nouns - `topics` (required): 1--5 topic keywords - `emotions` (optional): emotional tone -- `supersedes` (optional): ID of a memory this one replaces +- `supersedes` (optional): ID of a memory this one replaces. Over the MCP + and HTTP APIs the target must already exist and live in the same + namespace, or the store is rejected outright. The benchmark harness writes + through its own in-process path rather than those APIs, so an ID the model + invents is dropped rather than rejected — the memory is still stored, just + without the link. **Storage pipeline per memory:** 1. Generate embedding from concatenation of summary + full_text + tags @@ -91,7 +96,9 @@ outputs zero or more structured memories to store. 3. Add to SIMD vector index 4. Add to FTS5 full-text search index 5. Add as a node in the memory graph -6. If `supersedes` is set, add a Supersedes edge +6. If `supersedes` is set, add a Supersedes edge from the new memory to the + old one. Recall then returns the new memory in place of the old one, + provided the new one also passes the query's filters 7. Run automatic graph linking (similarity, entity, temporal) **Graph linking** (three types, all automatic): diff --git a/docs/guide.md b/docs/guide.md index c509eb9..e47c607 100644 --- a/docs/guide.md +++ b/docs/guide.md @@ -762,8 +762,38 @@ Store a new memory. | `topics` | string[] | No | Topic keywords, e.g. `["rust", "cooking"]`. Max 32. | | `emotions` | string[] | No | Emotional tone, e.g. `["happy", "anxious"]`. Max 32. | | `namespace` | string | No | Target namespace (default: `"default"`). | -| `parentId` | string | No | UUID of parent memory to create a hierarchical link. | -| `supersedes` | string | No | UUID of an older memory this one replaces. The old memory is deprioritized in search. | +| `parentId` | string | No | UUID of an existing memory in the same namespace, linked as this memory's parent. The store fails if the target does not exist or is in a different namespace. | +| `supersedes` | string | No | UUID of an existing memory in the same namespace that this one replaces. Recall REMOVES the old memory from results and returns this one in its place. The store fails if the target does not exist or is in a different namespace. | + +#### supersedes semantics + +`supersedes` is a *desired end state*, not an event, and the store result +always says what actually happened to the link. + +- **Preconditions, checked before anything is written.** The target must + exist and must live in the same namespace as the new memory, and a memory + cannot supersede itself. Violating any of these fails the whole store — + no memory is created. Checking up front is what keeps a rejected link + from leaving an orphaned memory behind. A *tombstoned* (deleted) target is + fine: correcting a memory you just deleted is a legitimate thing to do. +- **Idempotent.** Storing the same correction twice is not an error; the + second call reports `alreadyApplied`. +- **Reported, never silent.** When the request names a target, the store + result carries a `supersedes` object: + + | Field | Description | + |---|---| + | `target` | The memory that was superseded. | + | `status` | `applied`, `alreadyApplied`, `appliedNotDurable`, or `failed`. | + | `detail` | Explanation, present for every status except `applied`. | + + `applied` means the edge was created and written to disk. + `alreadyApplied` means it was already in place. `appliedNotDurable` means + the edge is active for the running process but could not be persisted, so + it will be lost on restart. `failed` means no edge exists; the memory was + still stored. +- **Counted.** The new memory's `edgeCount` includes the supersedes edge. + Memories stored before this behaviour existed are not backfilled. ### `store_memories` @@ -771,7 +801,7 @@ Store multiple memories in a single call. Each item has the same schema as `stor | Parameter | Type | Required | Description | |---|---|---|---| -| `memories` | array | Yes | Array of memory objects (max 100 per call). Each object has the same fields as `store_memory`. | +| `memories` | array | Yes | Array of memory objects (max 100 per call). Each object has the same fields as `store_memory`, including `supersedes` and its preconditions. | ### `recall_memories` diff --git a/docs/mcp.md b/docs/mcp.md index 4674ba3..0d37fcd 100644 --- a/docs/mcp.md +++ b/docs/mcp.md @@ -185,7 +185,7 @@ Store a new observation, fact, or piece of context. The system automatically gen | `emotions` | string[] | no | `[]` | Emotional tone, e.g. `["happy", "anxious"]`. Max 32. | | `namespace` | string | no | `"default"` | Memory partition | | `parentId` | string | no | -- | UUID of parent memory for hierarchical linking | -| `supersedes` | string | no | -- | UUID of an older memory this one replaces. The old memory is deprioritized in search. | +| `supersedes` | string | no | -- | UUID of an existing memory in the same namespace that this one replaces. Recall returns this memory **in place of** the old one. The store fails if the target does not exist or is in a different namespace. See [supersedes semantics](#supersedes-semantics). | #### Example @@ -214,6 +214,40 @@ Response: --- +#### supersedes semantics + +`supersedes` is a *desired end state*, not an event, and the store result always +says what actually happened to the link. + +**Preconditions, checked before anything is written.** The target must exist and +must live in the same namespace as the new memory, and a memory cannot supersede +itself. Violating any of these fails the whole store -- no memory is created. +Checking up front is what keeps a rejected link from leaving an orphaned memory +behind. A *tombstoned* (deleted) target is fine: correcting a memory you just +deleted is legitimate. + +**Idempotent.** Superseding the same memory twice is not an error; the second +call reports `alreadyApplied`. + +**The result reports the outcome.** Every successful store carries a +`supersedes` object: + +| Field | Description | +|---|---| +| `target` | The memory that was superseded. | +| `status` | `applied`, `alreadyApplied`, `appliedNotDurable`, or `failed`. | +| `detail` | Explanation, present for every status except `applied`. | + +`applied` means the edge was created and written to disk. `alreadyApplied` means +it was already in place. `appliedNotDurable` means the edge is active for the +running process but could not be persisted, so it will be lost on restart. +`failed` means no edge exists; the memory was still stored. + +At recall, a superseded memory is replaced by its replacement **only when the +replacement passes the query's own filters** -- namespace above all. If it does +not, the original is returned instead, so a result is never dropped with nothing +in its place. + ### store_memories Store multiple memories in a single call. Each item has the same schema as `store_memory`. Returns an array of results, one per input memory. @@ -799,7 +833,7 @@ Reinforce memories when they prove useful: - **After successful recall:** If you searched for something and the result was exactly what you needed, reinforce with quality 3 or 4. - **After applying context:** If a recalled memory helped you give better advice, reinforce it. - **Correct wrong memories:** If a recalled memory was inaccurate, either reinforce with quality 1 (weakens it) or use `forget_memory` to delete it outright, then store the corrected version. -- **Use `supersedes`:** When storing a correction, pass the old memory's ID as `supersedes` so the old one is deprioritized rather than deleted. +- **Use `supersedes`:** When storing a correction, pass the old memory's ID as `supersedes` so recall returns the correction instead of the stale memory. The old memory is kept on disk, not deleted. Memories that reach high stability (>1500 days) enter permastore and stop decaying. diff --git a/src/api/adapters.rs b/src/api/adapters.rs index 2bc02fd..9884170 100644 --- a/src/api/adapters.rs +++ b/src/api/adapters.rs @@ -24,6 +24,34 @@ use crate::storage::RedbStorageEngine; use super::state; use crate::storage::StorageEngine as StorageEngineTrait; +// ═══════════════════════════════════════════════════════════════════════ +// LinkError -> AppError +// ═══════════════════════════════════════════════════════════════════════ + +/// A rejected `supersedes`/`parentId` target becomes a 422 naming the +/// offending field — 422 is exactly the "syntactically fine, semantically +/// impossible" case. The two internal-inconsistency variants are ours, not +/// the caller's, so they become 500s instead. +impl From for super::errors::AppError { + fn from(e: crate::graph::LinkError) -> Self { + use crate::graph::LinkError; + let field = e.field().to_string(); + match e { + LinkError::TargetNotFound { .. } + | LinkError::NamespaceMismatch { .. } + | LinkError::SelfReference { .. } => super::errors::AppError::UnprocessableEntity { + message: e.to_string(), + field: Some(field), + }, + LinkError::TargetMissingFromGraph { .. } | LinkError::Lookup { .. } => { + super::errors::AppError::Internal { + source: Box::new(e), + } + } + } + } +} + // ═══════════════════════════════════════════════════════════════════════ // SearchPipelineAdapter // ═══════════════════════════════════════════════════════════════════════ @@ -844,16 +872,66 @@ impl state::RelationshipGraph for RelationshipGraphAdapter { drop(graph); // Release graph lock before storage let storage = self.storage.clone(); - let _ = tokio::task::spawn_blocking(move || { - if let Ok(storage_r) = storage.read() { - let _ = storage_r.batch_add_edges(&[persisted]); - } + let persist = tokio::task::spawn_blocking(move || { + let storage_r = storage + .read() + .map_err(|e| format!("storage lock poisoned: {e}"))?; + storage_r + .batch_add_edges(&[persisted]) + .map_err(|e| e.to_string()) }) .await; + // The edge is already live in the graph, so a persistence failure + // is not worth failing the caller over — but it does mean the edge + // disappears on restart, which is worth saying out loud. + let persist_error = match persist { + Ok(Ok(())) => None, + Ok(Err(e)) => Some(e), + Err(e) => Some(format!("blocking task join error: {e}")), + }; + if let Some(error) = persist_error { + tracing::warn!( + source_id = %from, + target_id = %to, + edge_type, + %error, + "edge is live in the graph but could not be persisted to edges.db; \ + it will be lost on restart" + ); + } + Ok(()) } + async fn check_link_target( + &self, + field: &'static str, + new_id: Option, + target: MemoryId, + expected_ns: NamespaceId, + ) -> Result<(), crate::graph::LinkError> { + crate::graph::check_link_target( + field, + new_id, + target, + expected_ns, + &self.graph, + &self.storage, + &self.cache, + ) + .await + } + + async fn apply_supersedes_edge( + &self, + new_id: MemoryId, + target: MemoryId, + ) -> crate::model::SupersedesOutcome { + crate::graph::apply_supersedes_edge(new_id, target, &self.graph, &self.storage, &self.cache) + .await + } + async fn tombstone_node(&self, id: MemoryId) -> Result<(), crate::graph::GraphError> { let mut graph = self.graph.write().await; graph.update_node_state(id, DecayPhase::Tombstone, 0.0)?; diff --git a/src/api/handlers.rs b/src/api/handlers.rs index 8862c59..0a755b0 100644 --- a/src/api/handlers.rs +++ b/src/api/handlers.rs @@ -55,12 +55,14 @@ fn health_report_cache() /// 1. Validate request content limits (see [`validate_memory_input`]); /// all length limits are UTF-8 bytes, not characters. /// 2. Resolve namespace by name -> NamespaceId. -/// 3. Generate embedding if not provided (calls embedding provider). -/// 4. Validate embedding dimensionality against namespace config. -/// 5. Persist to storage (meta.db, fulltext.dat, vectors.dat). -/// 6. Insert into RAM cache and vector index. -/// 7. If `parent_id` provided, create parent->child edge in graph. -/// 8. Return 201 with the created memory. +/// 3. Validate `parent_id` / `supersedes` targets, BEFORE any write, so a +/// bad link returns 422 without leaving an orphaned memory behind. +/// 4. Generate embedding if not provided (calls embedding provider). +/// 5. Validate embedding dimensionality against namespace config. +/// 6. Persist to storage (meta.db, fulltext.dat, vectors.dat). +/// 7. Insert into RAM cache and vector index. +/// 8. Create the `parent_id` and `supersedes` edges in the graph. +/// 9. Return 201 with the created memory. pub async fn create_memory( State(state): State, Json(req): Json, @@ -86,6 +88,24 @@ pub async fn create_memory( id: req.namespace.clone(), })?; + // --- Validate caller-specified links BEFORE anything is written --- + // + // The memory is persisted further down; a link failure discovered after + // that would leave an orphaned record behind a 4xx (which is exactly + // what the parent edge used to do). Checking here also spares an + // embedding round-trip on a request that cannot succeed. + // + // `None` for the new ID: it does not exist until `create_memory` + // returns, so the self-reference check has nothing to compare against. + for (field, target) in [("parentId", req.parent_id), ("supersedes", req.supersedes)] { + if let Some(target) = target { + state + .graph + .check_link_target(field, None, target, ns.id) + .await?; + } + } + // --- Merge entities/topics/emotions into tags (Issue 8) --- let mut merged_tags = req.tags.clone(); for entity in &req.entities { @@ -192,22 +212,26 @@ pub async fn create_memory( .await; // --- Parent edge --- + // Both explicit edges were validated before the record was written, so + // neither can fail the request from here — which is the point. A `?` at + // this line would return a 4xx with the memory already created. if let Some(parent_id) = req.parent_id { - state.graph.add_edge(parent_id, memory.id, "parent").await?; - } - - // --- Supersedes edge (Issue 8) --- - if let Some(old_id) = req.supersedes { - if let Err(e) = state.graph.add_edge(memory.id, old_id, "supersedes").await { - tracing::warn!( + if let Err(e) = state.graph.add_edge(parent_id, memory.id, "parent").await { + tracing::error!( memory_id = %memory.id, - superseded = %old_id, + parent = %parent_id, %e, - "supersedes edge failed (non-fatal)" + "parent edge could not be created after pre-validation passed" ); } } + // --- Supersedes edge (Issue 8) --- + let supersedes_outcome = match req.supersedes { + Some(old_id) => Some(state.graph.apply_supersedes_edge(memory.id, old_id).await), + None => None, + }; + // --- Autolink, entity-link, temporal-link (Issue 12) --- let created_at = req .created_at @@ -226,6 +250,7 @@ pub async fn create_memory( let mut response = MemoryResponse::from_cached(&memory, ns.name.clone()); response.full_text = req.full_text; + response.supersedes = supersedes_outcome; let took = start.elapsed().as_micros() as u64; Ok(( @@ -1227,31 +1252,68 @@ pub async fn batch_store( } let mut created = Vec::with_capacity(req.memories.len()); - - for mem_req in req.memories { - // Validate. This endpoint silently skips invalid items rather - // than failing the whole batch, so an error here is a `continue` - // — but the limits themselves are the shared ones, which keeps - // an oversized summary from reaching the u16-prefixed on-disk - // record encoder. - if validate_memory_input(MemoryInputRef { + let mut failed: Vec = Vec::new(); + + for (index, mem_req) in req.memories.into_iter().enumerate() { + // A bad item skips rather than failing the whole batch, but it no + // longer skips silently: every `continue` below records why. The + // limits themselves are the shared ones, which keeps an oversized + // summary from reaching the u16-prefixed on-disk record encoder. + if let Err(e) = validate_memory_input(MemoryInputRef { summary: &mem_req.summary, full_text: mem_req.full_text.as_deref(), tags: &mem_req.tags, entities: &mem_req.entities, topics: &mem_req.topics, emotions: &mem_req.emotions, - }) - .is_err() - { + }) { + failed.push(BatchStoreFailure { + index, + error: e.to_string(), + field: None, + }); continue; } let ns = match state.namespaces.resolve(&mem_req.namespace) { Some(ns) => ns, - None => continue, + None => { + failed.push(BatchStoreFailure { + index, + error: format!("namespace '{}' not found", mem_req.namespace), + field: Some("namespace".into()), + }); + continue; + } }; + // Validate caller-specified links before creating anything, so a + // rejected link leaves no orphaned memory behind. + let mut link_error: Option = None; + for (field, target) in [ + ("parentId", mem_req.parent_id), + ("supersedes", mem_req.supersedes), + ] { + if let Some(target) = target { + if let Err(e) = state + .graph + .check_link_target(field, None, target, ns.id) + .await + { + link_error = Some(e); + break; + } + } + } + if let Some(e) = link_error { + failed.push(BatchStoreFailure { + index, + error: e.to_string(), + field: Some(e.field().into()), + }); + continue; + } + // Merge entities/topics/emotions into tags let mut merged_tags = mem_req.tags.clone(); for entity in &mem_req.entities { @@ -1277,6 +1339,16 @@ pub async fn batch_store( let embedding = match mem_req.embedding { Some(ref vec) => { if vec.len() != ns.embedding_dim as usize { + failed.push(BatchStoreFailure { + index, + error: format!( + "embedding has {} dimensions, namespace '{}' requires {}", + vec.len(), + ns.name, + ns.embedding_dim + ), + field: Some("embedding".into()), + }); continue; } vec.clone() @@ -1291,7 +1363,14 @@ pub async fn batch_store( } match state.search.embed_text(&embed_text, ns.id).await { Ok(emb) => emb, - Err(_) => continue, + Err(e) => { + failed.push(BatchStoreFailure { + index, + error: format!("failed to generate embedding: {e}"), + field: None, + }); + continue; + } } } }; @@ -1312,7 +1391,14 @@ pub async fn batch_store( .await { Ok(m) => m, - Err(_) => continue, + Err(e) => { + failed.push(BatchStoreFailure { + index, + error: e.to_string(), + field: None, + }); + continue; + } }; // Cache + index @@ -1352,11 +1438,24 @@ pub async fn batch_store( ) .await; - // Supersedes edge - if let Some(old_id) = mem_req.supersedes { - let _ = state.graph.add_edge(memory.id, old_id, "supersedes").await; + // Explicit links. Both targets were validated before this item was + // created, so failures here are durability details, not rejections. + if let Some(parent_id) = mem_req.parent_id { + if let Err(e) = state.graph.add_edge(parent_id, memory.id, "parent").await { + tracing::error!( + memory_id = %memory.id, + parent = %parent_id, + %e, + "parent edge could not be created after pre-validation passed" + ); + } } + let supersedes_outcome = match mem_req.supersedes { + Some(old_id) => Some(state.graph.apply_supersedes_edge(memory.id, old_id).await), + None => None, + }; + // Post-creation links let created_at = mem_req .created_at @@ -1376,6 +1475,7 @@ pub async fn batch_store( created.push(BatchStoreResult { id: memory.id, namespace: ns.name.clone(), + supersedes: supersedes_outcome, }); } @@ -1385,7 +1485,11 @@ pub async fn batch_store( Ok(( StatusCode::CREATED, Json(ApiResponse { - data: BatchStoreResponse { created, total }, + data: BatchStoreResponse { + created, + total, + failed, + }, took_us: Some(took), }), )) diff --git a/src/api/models.rs b/src/api/models.rs index 3549e39..cd1fe60 100644 --- a/src/api/models.rs +++ b/src/api/models.rs @@ -554,6 +554,12 @@ pub struct BatchStoreResponse { pub created: Vec, /// Total number successfully created. pub total: u64, + /// Items that were not created, in request order. + /// + /// Omitted when everything succeeded, so a fully successful response + /// body is byte-identical to what earlier versions returned. + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub failed: Vec, } /// Result for a single batch store item. @@ -564,6 +570,26 @@ pub struct BatchStoreResult { pub id: MemoryId, /// Namespace name. pub namespace: String, + /// What became of the requested `supersedes` link, if one was asked for. + #[serde(skip_serializing_if = "Option::is_none")] + pub supersedes: Option, +} + +/// A batch item that was rejected, and why. +/// +/// Batch store previously dropped every failure on the floor: a `continue` +/// with no record of what happened. `index` refers to the item's position +/// in the request's `memories` array. +#[derive(Debug, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct BatchStoreFailure { + /// Position of the rejected item in the request array. + pub index: usize, + /// Human-readable reason. + pub error: String, + /// The request field at fault, when the failure names one. + #[serde(skip_serializing_if = "Option::is_none")] + pub field: Option, } // ═══════════════════════════════════════════════════════════════════════ diff --git a/src/api/state.rs b/src/api/state.rs index f0fdc2c..bb303a7 100644 --- a/src/api/state.rs +++ b/src/api/state.rs @@ -197,6 +197,33 @@ pub trait RelationshipGraph: Send + Sync { edge_type: &str, ) -> Result<(), GraphError>; + /// Validate a caller-specified link target (`supersedes`, `parentId`) + /// before the new memory is written. + /// + /// Called pre-write so a bad target rejects the request outright rather + /// than leaving an orphaned record behind a 4xx. Pass `None` for + /// `new_id` when the ID has not been minted yet; that skips the + /// self-reference check, which is vacuous in that case. + async fn check_link_target( + &self, + field: &'static str, + new_id: Option, + target: MemoryId, + expected_ns: NamespaceId, + ) -> Result<(), crate::graph::LinkError>; + + /// Create and persist the `Supersedes` edge for a memory that has + /// already been written. + /// + /// Cannot fail: the record is committed by this point, so the outcome + /// is reported rather than raised. See + /// [`SupersedesStatus`](crate::model::SupersedesStatus). + async fn apply_supersedes_edge( + &self, + new_id: MemoryId, + target: MemoryId, + ) -> crate::model::SupersedesOutcome; + /// Tombstone a memory's graph node (set phase to Tombstone, strength to 0). /// Unlike full removal, this preserves the node and edges for graph traversal. async fn tombstone_node(&self, id: MemoryId) -> Result<(), GraphError>; diff --git a/src/graph/explicit_links.rs b/src/graph/explicit_links.rs new file mode 100644 index 0000000..eac0858 --- /dev/null +++ b/src/graph/explicit_links.rs @@ -0,0 +1,601 @@ +//! Caller-specified graph edges: validation and application. +//! +//! Named for contrast with [`autolink`](super::autolink): the edges here are +//! the ones a client explicitly asked for (`supersedes`, `parentId`), not the +//! ones the system derives from similarity, shared entities, or time. +//! +//! # Why validation is separate from application +//! +//! Both the MCP and HTTP store paths commit the memory record to storage +//! (meta.db + fulltext.dat + vectors.dat + FTS + cache) *before* they touch +//! the graph. Failing after that write would require unwinding all of it. +//! So the flow is split in two: +//! +//! 1. [`check_link_target`] runs **before** anything is written. A missing, +//! cross-namespace, or self-referential target fails the whole store, and +//! because nothing has been written yet there is no rollback to get wrong. +//! 2. [`apply_supersedes_edge`] runs **after** the record is committed and +//! never returns an error — a duplicate edge or a persistence failure must +//! not report a store that plainly succeeded as a failure. The +//! [`SupersedesOutcome`] it returns is the truthful report instead. + +use std::sync::Arc; + +use crate::cache::CacheManager; +use crate::graph::SharedGraph; +use crate::model::{EdgeType, MemoryId, NamespaceId, SupersedesOutcome, SupersedesStatus}; +use crate::storage::engine::RedbStorageEngine; +use crate::storage::{PersistedEdge, StorageEngine as StorageEngineTrait}; + +// ═══════════════════════════════════════════════════════════════════════ +// LinkError +// ═══════════════════════════════════════════════════════════════════════ + +/// A caller-specified link target that cannot be honoured. +/// +/// Every variant is a *pre-write* failure: it is raised before the new +/// memory exists, so the correct response is to reject the store outright. +#[derive(Debug, thiserror::Error)] +pub enum LinkError { + /// No memory with this ID exists in storage. + #[error("{field} target {target} does not exist")] + TargetNotFound { + /// The request field naming the target (`supersedes`, `parentId`). + field: &'static str, + /// The target ID the caller supplied. + target: MemoryId, + }, + + /// The target exists but lives in a different namespace. + #[error( + "{field} target {target} belongs to namespace {actual} but the new memory is being \ + stored in namespace {expected}; {field} links must stay within a single namespace" + )] + NamespaceMismatch { + /// The request field naming the target. + field: &'static str, + /// The target ID the caller supplied. + target: MemoryId, + /// The namespace the target actually lives in. + actual: NamespaceId, + /// The namespace the new memory is being stored in. + expected: NamespaceId, + }, + + /// The target is the memory being stored — a degenerate 1-cycle. + #[error("{field} target {target} is the memory being stored; a memory cannot link to itself")] + SelfReference { + /// The request field naming the target. + field: &'static str, + /// The target ID the caller supplied. + target: MemoryId, + }, + + /// The target has a storage record but no graph node. The edge cannot + /// be created, and the inconsistency is ours, not the caller's. + #[error( + "{field} target {target} exists in storage but is missing from the relationship graph \ + (internal inconsistency)" + )] + TargetMissingFromGraph { + /// The request field naming the target. + field: &'static str, + /// The target ID the caller supplied. + target: MemoryId, + }, + + /// Looking the target up failed (storage error, poisoned lock, panic). + #[error("failed to look up {field} target {target}: {message}")] + Lookup { + /// The request field naming the target. + field: &'static str, + /// The target ID the caller supplied. + target: MemoryId, + /// The underlying failure. + message: String, + }, +} + +impl LinkError { + /// The request field this error is about, for `field`-scoped error + /// responses (HTTP 422 bodies carry it; MCP folds it into the message). + pub fn field(&self) -> &'static str { + match self { + LinkError::TargetNotFound { field, .. } + | LinkError::NamespaceMismatch { field, .. } + | LinkError::SelfReference { field, .. } + | LinkError::TargetMissingFromGraph { field, .. } + | LinkError::Lookup { field, .. } => field, + } + } +} + +// ═══════════════════════════════════════════════════════════════════════ +// validate_link_target (pure) +// ═══════════════════════════════════════════════════════════════════════ + +/// Decide whether a caller-specified link target is usable. +/// +/// Pure: all I/O is done by the caller and handed in. [`check_link_target`] +/// is the async wrapper that performs those lookups. +/// +/// # Arguments +/// +/// * `field` — the request field naming the target, used in messages. +/// * `new_memory_id` — the ID of the memory being stored, or `None` when the +/// caller has not minted one yet (the HTTP create path). `None` skips the +/// self-reference check, which is vacuous when the ID does not exist. +/// * `target_namespace` — the namespace of the target's storage record, or +/// `None` when no record was found. +/// * `target_in_graph` — whether the target has a node in the graph. +/// +/// # Deliberately allowed +/// +/// A tombstoned target passes. Tombstoned memories keep their graph node by +/// design and recall never surfaces them, so rejecting them would break the +/// legitimate "correct a memory I just deleted" flow. +pub fn validate_link_target( + field: &'static str, + new_memory_id: Option, + target: MemoryId, + target_namespace: Option, + target_in_graph: bool, + expected_namespace: NamespaceId, +) -> Result<(), LinkError> { + if new_memory_id == Some(target) { + return Err(LinkError::SelfReference { field, target }); + } + + let Some(actual) = target_namespace else { + return Err(LinkError::TargetNotFound { field, target }); + }; + + if actual != expected_namespace { + return Err(LinkError::NamespaceMismatch { + field, + target, + actual, + expected: expected_namespace, + }); + } + + if !target_in_graph { + return Err(LinkError::TargetMissingFromGraph { field, target }); + } + + Ok(()) +} + +// ═══════════════════════════════════════════════════════════════════════ +// check_link_target (async I/O wrapper) +// ═══════════════════════════════════════════════════════════════════════ + +/// Look up a caller-specified link target and validate it. +/// +/// Call this **before** writing the new memory: a failure here should reject +/// the store, and doing so pre-write means there is nothing to unwind. +/// +/// The lookup reads the RAM cache first and falls back to a blocking +/// `get_record` on a spawned thread, mirroring the rest of the store path. +pub async fn check_link_target( + field: &'static str, + new_memory_id: Option, + target: MemoryId, + expected_namespace: NamespaceId, + graph: &SharedGraph, + storage: &Arc>, + cache: &Arc, +) -> Result<(), LinkError> { + // Cheapest check first: it needs no I/O at all. + if new_memory_id == Some(target) { + return Err(LinkError::SelfReference { field, target }); + } + + let target_namespace = match cache.get(target).await { + Some(record) => Some(record.namespace_id), + None => { + let storage = storage.clone(); + let record = tokio::task::spawn_blocking(move || { + let storage_r = storage + .read() + .map_err(|e| format!("storage lock poisoned: {e}"))?; + storage_r.get_record(target).map_err(|e| e.to_string()) + }) + .await + .map_err(|e| LinkError::Lookup { + field, + target, + message: format!("blocking task join error: {e}"), + })? + .map_err(|message| LinkError::Lookup { + field, + target, + message, + })?; + record.map(|r| NamespaceId::new(r.namespace_id)) + } + }; + + let target_in_graph = graph.read().await.get_node(&target).is_some(); + + validate_link_target( + field, + new_memory_id, + target, + target_namespace, + target_in_graph, + expected_namespace, + ) +} + +// ═══════════════════════════════════════════════════════════════════════ +// Applying the edge +// ═══════════════════════════════════════════════════════════════════════ + +/// What happened when a caller-specified edge was inserted and persisted. +#[derive(Debug)] +enum EdgeApplication { + /// Inserted into the graph and written to `edges.db`. + Created, + /// An edge of this type already joined the two memories. + AlreadyPresent, + /// The graph refused the edge (carries the reason). + Rejected(String), + /// Live in this process's graph but not written to disk. + NotDurable(String), +} + +/// Insert one caller-specified edge into the graph and persist it. +/// +/// Shared by every explicit link field so `supersedes` and `parentId` +/// cannot drift apart in durability, locking, or `edge_count` behaviour — +/// the inconsistency between them was itself part of the reported bug. +async fn apply_explicit_edge( + source: MemoryId, + target: MemoryId, + edge_type: EdgeType, + graph: &SharedGraph, + storage: &Arc>, + cache: &Arc, +) -> EdgeApplication { + // Insert into the in-memory graph, then release the write lock before + // touching storage. Holding the graph lock across a storage acquisition + // would invert the lock order used everywhere else. + let insert_result = { + let mut graph_w = graph.write().await; + graph_w.add_edge(source, target, edge_type, 1.0, false) + }; + + match insert_result { + Ok(_) => {} + Err(crate::graph::GraphError::EdgeExists(_, _)) => { + return EdgeApplication::AlreadyPresent; + } + Err(e) => return EdgeApplication::Rejected(e.to_string()), + } + + let persisted = PersistedEdge { + source, + target, + edge_type, + weight: 1.0, + auto_created: false, + created_at: std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_millis() as u64, + }; + + let persist_result = { + let storage = storage.clone(); + tokio::task::spawn_blocking(move || { + let storage_r = storage + .read() + .map_err(|e| format!("storage lock poisoned: {e}"))?; + storage_r + .batch_add_edges(&[persisted]) + .map_err(|e| e.to_string()) + }) + .await + }; + + let persist_error = match persist_result { + Ok(Ok(())) => None, + Ok(Err(e)) => Some(e), + Err(e) => Some(format!("blocking task join error: {e}")), + }; + + if let Some(message) = persist_error { + return EdgeApplication::NotDurable(message); + } + + // `edge_count` counts OUTGOING edges, so only the source's count moves. + // The update is additive and runs before autolink, which re-reads the + // cache, so the two stay in agreement. + bump_edge_count(source, storage, cache).await; + + EdgeApplication::Created +} + +/// Create and persist the `Supersedes` edge for a memory that has already +/// been written. +/// +/// Never returns an error — by the time this runs the record is committed, +/// so reporting a failure would lie about a store that succeeded. The +/// returned [`SupersedesOutcome`] is the report instead; see +/// [`SupersedesStatus`] for what each variant means. +/// +/// Assumes [`check_link_target`] already ran. If the target vanished in +/// between, the outcome is [`SupersedesStatus::Failed`] and an `error!` is +/// logged, because that is a real (if unreachable) internal inconsistency. +pub async fn apply_supersedes_edge( + new_memory_id: MemoryId, + target: MemoryId, + graph: &SharedGraph, + storage: &Arc>, + cache: &Arc, +) -> SupersedesOutcome { + match apply_explicit_edge( + new_memory_id, + target, + EdgeType::Supersedes, + graph, + storage, + cache, + ) + .await + { + EdgeApplication::Created => SupersedesOutcome::new(target, SupersedesStatus::Applied), + + EdgeApplication::AlreadyPresent => SupersedesOutcome::with_detail( + target, + SupersedesStatus::AlreadyApplied, + "a supersedes edge already connected these two memories; \ + the requested end state was already in place", + ), + + EdgeApplication::Rejected(message) => { + tracing::error!( + memory_id = %new_memory_id, + superseded = %target, + error = %message, + "supersedes edge could not be added to the graph after pre-validation passed" + ); + SupersedesOutcome::with_detail( + target, + SupersedesStatus::Failed, + format!("the supersedes edge could not be created: {message}"), + ) + } + + EdgeApplication::NotDurable(message) => { + tracing::error!( + memory_id = %new_memory_id, + superseded = %target, + error = %message, + "supersedes edge is live in the graph but could not be persisted to edges.db; \ + it will be lost on restart" + ); + SupersedesOutcome::with_detail( + target, + SupersedesStatus::AppliedNotDurable, + format!( + "the supersedes edge is active for this process but was not written to \ + disk and will be lost on restart: {message}" + ), + ) + } + } +} + +/// Create and persist the `ParentChild` edge from `parent` to `child`. +/// +/// Direction follows [`EdgeType::ParentChild`]: source is the parent. +/// Like the supersedes edge this runs after the record is committed and so +/// cannot fail the store; problems are logged. Assumes +/// [`check_link_target`] already ran with field `"parentId"`. +pub async fn apply_parent_edge( + parent: MemoryId, + child: MemoryId, + graph: &SharedGraph, + storage: &Arc>, + cache: &Arc, +) { + match apply_explicit_edge(parent, child, EdgeType::ParentChild, graph, storage, cache).await { + EdgeApplication::Created | EdgeApplication::AlreadyPresent => {} + EdgeApplication::Rejected(message) => tracing::error!( + memory_id = %child, + parent = %parent, + error = %message, + "parent edge could not be added to the graph after pre-validation passed" + ), + EdgeApplication::NotDurable(message) => tracing::error!( + memory_id = %child, + parent = %parent, + error = %message, + "parent edge is live in the graph but could not be persisted to edges.db; \ + it will be lost on restart" + ), + } +} + +/// Additively increment a memory's cached outgoing-edge count by one, +/// in both storage and the RAM cache. Mirrors `autolink::persist_edges`: +/// a failure here is a stale counter, not a lost edge, so it only warns. +async fn bump_edge_count( + memory_id: MemoryId, + storage: &Arc>, + cache: &Arc, +) { + let new_count = cache + .get(memory_id) + .await + .map(|r| r.edge_count) + .unwrap_or(0) + .saturating_add(1); + + { + let storage = storage.clone(); + match tokio::task::spawn_blocking(move || { + let storage_r = storage + .read() + .map_err(|e| format!("storage lock poisoned: {e}"))?; + storage_r + .update_edge_count(memory_id, new_count) + .map_err(|e| e.to_string()) + }) + .await + { + Ok(Ok(())) => {} + Ok(Err(e)) => tracing::warn!( + memory_id = %memory_id, + error = %e, + "Failed to update edge count in storage; cache is correct but on-disk record may be stale" + ), + Err(e) => tracing::warn!( + memory_id = %memory_id, + error = %e, + "Edge count update task panicked; cache is correct but on-disk record may be stale" + ), + } + } + + cache.update_edge_count(memory_id, new_count).await; +} + +// ═══════════════════════════════════════════════════════════════════════ +// Tests +// ═══════════════════════════════════════════════════════════════════════ + +#[cfg(test)] +mod tests { + use super::*; + + /// The namespace the new memory is being stored in. + fn ns() -> NamespaceId { + NamespaceId::new(1) + } + + /// Any other namespace. + fn other_ns() -> NamespaceId { + NamespaceId::new(2) + } + + /// A missing target must error rather than silently no-op — the whole + /// point of the fix. + #[test] + fn missing_target_errors() { + let target = MemoryId::new(); + let err = validate_link_target( + "supersedes", + Some(MemoryId::new()), + target, + None, + false, + ns(), + ) + .unwrap_err(); + + assert!(matches!(err, LinkError::TargetNotFound { .. })); + assert!(err.to_string().contains("does not exist")); + } + + /// A target in another namespace is rejected: honouring it would create + /// exactly the cross-namespace edge that recall must not follow. + #[test] + fn cross_namespace_target_errors() { + let target = MemoryId::new(); + let err = validate_link_target( + "supersedes", + Some(MemoryId::new()), + target, + Some(other_ns()), + true, + ns(), + ) + .unwrap_err(); + + match err { + LinkError::NamespaceMismatch { + actual, expected, .. + } => { + assert_eq!(actual, other_ns()); + assert_eq!(expected, ns()); + } + other => panic!("expected NamespaceMismatch, got {other:?}"), + } + } + + /// The happy path: record present, node present, same namespace. + #[test] + fn valid_target_is_accepted() { + let result = validate_link_target( + "supersedes", + Some(MemoryId::new()), + MemoryId::new(), + Some(ns()), + true, + ns(), + ); + + assert!(result.is_ok()); + } + + /// A record with no graph node cannot receive an edge, and that is an + /// internal inconsistency rather than caller error. + #[test] + fn record_without_graph_node_errors() { + let err = validate_link_target( + "supersedes", + Some(MemoryId::new()), + MemoryId::new(), + Some(ns()), + false, + ns(), + ) + .unwrap_err(); + + assert!(matches!(err, LinkError::TargetMissingFromGraph { .. })); + } + + /// Self-reference is a degenerate 1-cycle and is rejected — but only + /// when the caller has an ID to compare against. + #[test] + fn self_reference_errors_only_when_id_is_known() { + let id = MemoryId::new(); + + let err = + validate_link_target("supersedes", Some(id), id, Some(ns()), true, ns()).unwrap_err(); + assert!(matches!(err, LinkError::SelfReference { .. })); + + // The HTTP create path has no ID yet; the check is vacuous there. + assert!(validate_link_target("supersedes", None, id, Some(ns()), true, ns()).is_ok()); + } + + /// The same validator serves both explicit link fields, so the field + /// name has to survive into both the message and `field()`. + #[test] + fn field_name_propagates_to_message_and_accessor() { + for field in ["supersedes", "parentId"] { + let target = MemoryId::new(); + let err = validate_link_target(field, Some(MemoryId::new()), target, None, false, ns()) + .unwrap_err(); + + assert_eq!(err.field(), field); + assert!(err.to_string().starts_with(field)); + + let mismatch = validate_link_target( + field, + Some(MemoryId::new()), + target, + Some(other_ns()), + true, + ns(), + ) + .unwrap_err(); + + assert_eq!(mismatch.field(), field); + assert!(mismatch.to_string().contains(field)); + } + } +} diff --git a/src/graph/mod.rs b/src/graph/mod.rs index a857b90..cbefbad 100644 --- a/src/graph/mod.rs +++ b/src/graph/mod.rs @@ -11,6 +11,7 @@ pub mod activation; pub mod autolink; +pub mod explicit_links; pub mod rebuild; mod structure; @@ -35,6 +36,11 @@ pub use crate::rif::rif_edge_factor; // Re-export PersistedEdge from storage (the canonical definition) pub use crate::storage::PersistedEdge; +// Re-export explicit (caller-specified) link validation and application +pub use explicit_links::{ + LinkError, apply_parent_edge, apply_supersedes_edge, check_link_target, validate_link_target, +}; + // Re-export CS-11 auto-link types and functions pub use autolink::{ AutoLinkCandidate, AutoLinkError, DEFAULT_MAX_LINKS, THRESHOLD_HARD_FLOOR, auto_link, diff --git a/src/graph/structure.rs b/src/graph/structure.rs index 75a3e55..25e5e3c 100644 --- a/src/graph/structure.rs +++ b/src/graph/structure.rs @@ -526,6 +526,46 @@ impl RelationshipGraph { self.edges.len() } + /// Follow incoming `Supersedes` edges to the newest version of `id`. + /// + /// A `Supersedes` edge runs source (new) -> target (old), so the memory + /// that replaces `id` is the *source* of an incoming edge. The walk + /// repeats until it reaches a memory nothing supersedes. + /// + /// Returns `None` when `id` is already current, and also when the chain + /// contains a cycle. A cycle has no newest version, so substituting an + /// arbitrary node from inside it would make recall return a different + /// memory depending on hash iteration order; returning the original + /// memory unchanged is the only stable answer. + pub fn supersedes_terminal(&self, id: &MemoryId) -> Option { + /// Hop ceiling, a second guard behind the cycle check. Real + /// correction chains are a handful of links at most. + const MAX_HOPS: usize = 10; + + let mut visited: std::collections::HashSet = std::collections::HashSet::new(); + visited.insert(*id); + + let mut current = *id; + for _ in 0..MAX_HOPS { + let replacement = self.incoming_edges(¤t).into_iter().find_map(|edge| { + if edge.edge_type == EdgeType::Supersedes { + self.nodes.get(edge.source).map(|n| n.memory_id) + } else { + None + } + }); + + match replacement { + // Revisiting a node means the chain loops. + Some(next) if !visited.insert(next) => return None, + Some(next) => current = next, + None => break, + } + } + + if current == *id { None } else { Some(current) } + } + /// Read-only access to an edge by key. pub fn get_edge(&self, key: EdgeKey) -> Option<&GraphEdge> { self.edges.get(key) @@ -771,3 +811,146 @@ impl RelationshipGraph { self.id_index.get(memory_id).copied() } } + +// ═══════════════════════════════════════════════════════════════════════ +// Tests +// ═══════════════════════════════════════════════════════════════════════ + +#[cfg(test)] +mod tests { + use super::*; + + /// Build a graph holding `n` fresh Full-phase nodes in one namespace. + fn graph_with_nodes(n: usize) -> (RelationshipGraph, Vec) { + let mut graph = RelationshipGraph::new(); + let ns = NamespaceId::new(1); + let ids: Vec = (0..n).map(|_| MemoryId::new()).collect(); + for id in &ids { + graph + .add_node(*id, ns, DecayPhase::Full, 1.0, 0) + .expect("fresh id"); + } + (graph, ids) + } + + /// Add a Supersedes edge meaning "`new` replaces `old`". + fn supersede( + graph: &mut RelationshipGraph, + new: MemoryId, + old: MemoryId, + ) -> Result { + graph.add_edge(new, old, EdgeType::Supersedes, 1.0, false) + } + + /// An edge to a memory that is not in the graph is refused, which is + /// what makes pre-validating the target worth doing. + #[test] + fn supersedes_edge_to_missing_target_errors() { + let (mut graph, ids) = graph_with_nodes(1); + let missing = MemoryId::new(); + + let err = supersede(&mut graph, ids[0], missing).unwrap_err(); + + assert!(matches!(err, GraphError::MemoryNotFound(id) if id == missing)); + assert_eq!(graph.edge_count(), 0); + } + + /// The exact edge shape `supersedes_terminal` depends on: the new + /// memory is the SOURCE of an edge incoming to the old one. + #[test] + fn valid_target_creates_edge_with_new_memory_as_source() { + let (mut graph, ids) = graph_with_nodes(2); + let (new, old) = (ids[0], ids[1]); + + supersede(&mut graph, new, old).expect("both nodes present"); + + let incoming = graph.incoming_edges(&old); + assert_eq!(incoming.len(), 1); + assert_eq!(incoming[0].edge_type, EdgeType::Supersedes); + assert_eq!( + graph.nodes.get(incoming[0].source).map(|n| n.memory_id), + Some(new) + ); + } + + /// Storing the same correction twice is idempotent at the graph level. + /// The store path reports this as `alreadyApplied`, not as an error: + /// supersedes is a desired end state, and the end state is in place. + #[test] + fn duplicate_supersedes_edge_is_rejected_as_already_existing() { + let (mut graph, ids) = graph_with_nodes(2); + let (new, old) = (ids[0], ids[1]); + + supersede(&mut graph, new, old).expect("first insert"); + let err = supersede(&mut graph, new, old).unwrap_err(); + + assert!(matches!(err, GraphError::EdgeExists(_, _))); + assert_eq!(graph.edge_count(), 1); + } + + /// A chain of corrections resolves to its newest member, not to the + /// first hop. + #[test] + fn supersedes_terminal_follows_the_chain_to_the_newest_version() { + let (mut graph, ids) = graph_with_nodes(3); + let (a, b, c) = (ids[0], ids[1], ids[2]); + + supersede(&mut graph, a, b).expect("a replaces b"); + supersede(&mut graph, c, a).expect("c replaces a"); + + assert_eq!(graph.supersedes_terminal(&b), Some(c)); + assert_eq!(graph.supersedes_terminal(&a), Some(c)); + } + + /// A memory nothing has replaced resolves to nothing, so recall leaves + /// it alone. + #[test] + fn supersedes_terminal_is_none_for_a_current_memory() { + let (mut graph, ids) = graph_with_nodes(2); + supersede(&mut graph, ids[0], ids[1]).expect("a replaces b"); + + assert_eq!(graph.supersedes_terminal(&ids[0]), None); + } + + /// Regression guard for the cycle detection. A two-node cycle cannot be + /// built through `add_edge` (its duplicate check spans both directions) + /// but `rebuild_from_storage` inserts edges.db rows without that check, + /// so a restart can materialise one. Before cycle detection this walked + /// the full hop cap and substituted whichever node it happened to land + /// on; now it terminates with "no newest version". + #[test] + fn supersedes_terminal_returns_none_on_a_cycle() { + let (mut graph, ids) = graph_with_nodes(2); + let (a, b) = (ids[0], ids[1]); + + let edge = |source, target| crate::storage::PersistedEdge { + source, + target, + edge_type: EdgeType::Supersedes, + weight: 1.0, + auto_created: false, + created_at: 0, + }; + let loaded = + crate::graph::rebuild_from_storage(&mut graph, [edge(a, b), edge(b, a)].into_iter()); + assert_eq!(loaded, 2); + + assert_eq!(graph.supersedes_terminal(&a), None); + assert_eq!(graph.supersedes_terminal(&b), None); + } + + /// A longer cycle terminates too, and reports the same "no answer". + #[test] + fn supersedes_terminal_returns_none_on_a_three_node_cycle() { + let (mut graph, ids) = graph_with_nodes(3); + let (a, b, c) = (ids[0], ids[1], ids[2]); + + supersede(&mut graph, a, b).expect("a replaces b"); + supersede(&mut graph, b, c).expect("b replaces c"); + supersede(&mut graph, c, a).expect("c replaces a"); + + assert_eq!(graph.supersedes_terminal(&a), None); + assert_eq!(graph.supersedes_terminal(&b), None); + assert_eq!(graph.supersedes_terminal(&c), None); + } +} diff --git a/src/mcp/bridge.rs b/src/mcp/bridge.rs index 76a29e9..051745d 100644 --- a/src/mcp/bridge.rs +++ b/src/mcp/bridge.rs @@ -283,6 +283,15 @@ pub struct StoredMemory { pub stability: f32, /// Creation timestamp as ISO 8601 string. pub created_at: String, + /// What became of the requested `supersedes` link, when one was asked + /// for. Absent when the store named no target. + /// + /// `#[serde(default)]` is load-bearing, not decoration: this struct + /// round-trips through the daemon protocol, so a new client talking to + /// an older daemon receives a payload without this key and must still + /// decode it. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub supersedes: Option, } /// A full memory record returned by get. @@ -607,3 +616,91 @@ impl crate::mcp::server::McpHandler for McpBridge { crate::mcp::resources::dispatch_resource(self, uri).await } } + +// ═══════════════════════════════════════════════════════════════════════ +// Tests +// ═══════════════════════════════════════════════════════════════════════ + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::*; + use crate::model::{MemoryId, SupersedesOutcome, SupersedesStatus}; + + fn stored_memory_json() -> serde_json::Value { + json!({ + "id": "0195e2c0-0000-7000-8000-000000000001", + "namespace": "default", + "phase": "Full", + "strength": 1.0, + "stability": 3.7145, + "createdAt": "2026-08-11T12:00:00Z", + }) + } + + /// Back-compat guard for the daemon protocol. `StoredMemory` is + /// round-tripped through serde between client and daemon, so a new + /// client must still decode a response from a daemon that predates the + /// `supersedes` field. This is what `#[serde(default)]` buys. + #[test] + fn stored_memory_decodes_a_payload_without_the_supersedes_field() { + let decoded: StoredMemory = + serde_json::from_value(stored_memory_json()).expect("old daemon payload decodes"); + + assert!(decoded.supersedes.is_none()); + assert_eq!(decoded.namespace, "default"); + } + + /// A store that named no target must serialise exactly as before, so + /// existing clients see an unchanged response shape. + #[test] + fn absent_supersedes_outcome_is_omitted_from_the_output() { + let stored: StoredMemory = + serde_json::from_value(stored_memory_json()).expect("payload decodes"); + + let encoded = serde_json::to_value(&stored).expect("encodes"); + assert!(encoded.get("supersedes").is_none()); + } + + /// The status names are part of the wire contract in both directions. + #[test] + fn supersedes_status_round_trips_through_its_camel_case_names() { + let cases = [ + (SupersedesStatus::Applied, "applied"), + (SupersedesStatus::AlreadyApplied, "alreadyApplied"), + (SupersedesStatus::AppliedNotDurable, "appliedNotDurable"), + (SupersedesStatus::Failed, "failed"), + ]; + + for (status, name) in cases { + assert_eq!(serde_json::to_value(status).unwrap(), json!(name)); + assert_eq!( + serde_json::from_value::(json!(name)).unwrap(), + status + ); + } + } + + /// A populated outcome survives the daemon round trip intact. + #[test] + fn supersedes_outcome_round_trips_on_a_stored_memory() { + let target = MemoryId::new(); + let mut stored: StoredMemory = + serde_json::from_value(stored_memory_json()).expect("payload decodes"); + stored.supersedes = Some(SupersedesOutcome::with_detail( + target, + SupersedesStatus::AppliedNotDurable, + "disk write failed", + )); + + let encoded = serde_json::to_value(&stored).expect("encodes"); + assert_eq!(encoded["supersedes"]["status"], json!("appliedNotDurable")); + + let decoded: StoredMemory = serde_json::from_value(encoded).expect("decodes"); + let outcome = decoded.supersedes.expect("outcome preserved"); + assert_eq!(outcome.target, target); + assert_eq!(outcome.status, SupersedesStatus::AppliedNotDurable); + assert_eq!(outcome.detail.as_deref(), Some("disk write failed")); + } +} diff --git a/src/mcp/bridge_adapters.rs b/src/mcp/bridge_adapters.rs index 0c3303c..7ae4002 100644 --- a/src/mcp/bridge_adapters.rs +++ b/src/mcp/bridge_adapters.rs @@ -20,6 +20,25 @@ use crate::time::format_timestamp; use super::bridge; use crate::storage::StorageEngine as StorageEngineTrait; +/// Map an explicit-link validation failure onto the bridge error the tool +/// layer already knows how to render. +/// +/// The distinction that matters to a caller: `NotFound` and `InvalidInput` +/// mean "fix your request", while `Internal` and `Storage` mean "this one is +/// ours". `tools.rs` surfaces all four through its existing +/// `Failed to store memory: {e}` path, so no error handling changes there. +fn map_link_error(e: crate::graph::LinkError) -> bridge::BridgeError { + use crate::graph::LinkError; + match e { + LinkError::TargetNotFound { .. } => bridge::BridgeError::NotFound(e.to_string()), + LinkError::NamespaceMismatch { .. } | LinkError::SelfReference { .. } => { + bridge::BridgeError::InvalidInput(e.to_string()) + } + LinkError::TargetMissingFromGraph { .. } => bridge::BridgeError::Internal(e.to_string()), + LinkError::Lookup { .. } => bridge::BridgeError::Storage(e.to_string()), + } +} + // ═══════════════════════════════════════════════════════════════════════ // McpSearchAdapter // ═══════════════════════════════════════════════════════════════════════ @@ -601,6 +620,45 @@ impl bridge::StorageEngine for McpStorageAdapter { })?? }; + // The ID is minted here rather than just before the insert so the + // link pre-validation below can reject a self-reference. Minting is + // pure (UUIDv7 from the clock), so an early mint costs nothing. + let memory_id = MemoryId::new(); + + // Validate caller-specified links BEFORE anything is written. + // + // The record is committed to meta.db, fulltext.dat, vectors.dat, FTS + // and the cache further down; a link failure discovered after that + // would have to unwind all of it. Checking here makes "reject the + // whole store" free, and it also avoids paying for an embedding + // round-trip on a store that cannot succeed. + if let Some(parent_id) = input.parent_id { + crate::graph::check_link_target( + "parentId", + Some(memory_id), + parent_id, + ns_config.id, + &self.graph, + &self.storage, + &self.cache, + ) + .await + .map_err(map_link_error)?; + } + if let Some(old_id) = input.supersedes { + crate::graph::check_link_target( + "supersedes", + Some(memory_id), + old_id, + ns_config.id, + &self.graph, + &self.storage, + &self.cache, + ) + .await + .map_err(map_link_error)?; + } + // Convert structured metadata fields to tags and merge with explicit tags. let mut merged_tags = input.tags.clone(); for entity in &input.entities { @@ -640,7 +698,6 @@ impl bridge::StorageEngine for McpStorageAdapter { } }; - let memory_id = MemoryId::new(); let now = chrono::Utc::now().timestamp_millis(); let parsed_tags: Vec = merged_tags .iter() @@ -731,9 +788,8 @@ impl bridge::StorageEngine for McpStorageAdapter { } } - // Add node to the relationship graph (must happen BEFORE autolink). - // Track whether we need to persist a supersedes edge after releasing graph lock. - let mut supersedes_edge: Option = None; + // Add node to the relationship graph (must happen BEFORE autolink, + // and before either explicit edge — both endpoints must exist). { let mut graph_w = self.graph.write().await; // Silently ignore DuplicateNode (should not happen for a new memory). @@ -744,73 +800,36 @@ impl bridge::StorageEngine for McpStorageAdapter { 1.0, record.vector_slot, ); - - // Create supersedes edge: new memory → old memory. - if let Some(old_id) = input.supersedes { - if let Err(e) = graph_w.add_edge( - memory_id, - old_id, - crate::model::EdgeType::Supersedes, - 1.0, - false, - ) { - tracing::warn!( - memory_id = %memory_id, - superseded = %old_id, - %e, - "supersedes edge failed (non-fatal)" - ); - } else { - let now_ms = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_millis() as u64; - supersedes_edge = Some(crate::storage::PersistedEdge { - source: memory_id, - target: old_id, - edge_type: crate::model::EdgeType::Supersedes, - weight: 1.0, - auto_created: false, - created_at: now_ms, - }); - } - } } // graph write lock released - // Persist supersedes edge AFTER graph lock is released to avoid lock ordering inversion. - if let Some(persisted) = supersedes_edge { - let old_id = persisted.target; - let storage = self.storage.clone(); - let persist_result = tokio::task::spawn_blocking(move || { - let storage_r = storage.read().map_err(|e| { - bridge::BridgeError::Internal(format!("storage lock poisoned: {e}")) - })?; - storage_r - .batch_add_edges(&[persisted]) - .map_err(|e| bridge::BridgeError::Storage(e.to_string())) - }) + // Explicit links. Both targets were validated before the record was + // written, so anything that goes wrong here is a durability or + // idempotency detail, not a reason to fail a store that succeeded. + if let Some(parent_id) = input.parent_id { + crate::graph::apply_parent_edge( + parent_id, + memory_id, + &self.graph, + &self.storage, + &self.cache, + ) .await; - match persist_result { - Ok(Err(e)) => { - tracing::warn!( - memory_id = %memory_id, - superseded = %old_id, - error = %e, - "supersedes edge persistence failed (non-fatal)" - ); - } - Err(e) => { - tracing::warn!( - memory_id = %memory_id, - superseded = %old_id, - error = %e, - "supersedes edge persistence task panicked (non-fatal)" - ); - } - _ => {} - } } + let supersedes_outcome = match input.supersedes { + Some(old_id) => Some( + crate::graph::apply_supersedes_edge( + memory_id, + old_id, + &self.graph, + &self.storage, + &self.cache, + ) + .await, + ), + None => None, + }; + // Auto-link: discover and create edges to similar existing memories. if self.config.graph.autolink_enabled { let tag_strings: Vec = merged_tags.clone(); @@ -912,6 +931,7 @@ impl bridge::StorageEngine for McpStorageAdapter { strength: record.strength, stability: record.stability, created_at: format_timestamp(record.created_at, self.timezone), + supersedes: supersedes_outcome, }) } diff --git a/src/mcp/tools.rs b/src/mcp/tools.rs index a72c68b..9dbcdeb 100644 --- a/src/mcp/tools.rs +++ b/src/mcp/tools.rs @@ -130,11 +130,11 @@ fn store_memory_def() -> ToolInfo { }, "parentId": { "type": "string", - "description": "UUID of parent memory to create a hierarchical link" + "description": "UUID of an existing memory in the same namespace, linked as this memory's parent. The store fails if the target does not exist or is in a different namespace." }, "supersedes": { "type": "string", - "description": "UUID of an older memory this one replaces. The old memory will be deprioritized in search results in favor of this one." + "description": "UUID of an existing memory in the same namespace that this one replaces. Recall returns this memory in place of the old one. The store fails if the target does not exist or is in a different namespace." } }, "required": ["summary"] @@ -241,11 +241,11 @@ fn store_memories_def() -> ToolInfo { }, "parentId": { "type": "string", - "description": "UUID of parent memory to create a hierarchical link" + "description": "UUID of an existing memory in the same namespace, linked as this memory's parent. The store fails if the target does not exist or is in a different namespace." }, "supersedes": { "type": "string", - "description": "UUID of an older memory this one replaces" + "description": "UUID of an existing memory in the same namespace that this one replaces. The store fails if the target does not exist or is in a different namespace." } }, "required": ["summary"] @@ -318,6 +318,7 @@ async fn handle_store_memories(bridge: &McpBridge, arguments: serde_json::Value) "strength": stored.strength, "stability": stored.stability, "createdAt": stored.created_at, + "supersedes": stored.supersedes, })); } Err(e) => { diff --git a/src/model/edge.rs b/src/model/edge.rs index 169fb45..a952070 100644 --- a/src/model/edge.rs +++ b/src/model/edge.rs @@ -2,6 +2,8 @@ use serde::{Deserialize, Serialize}; +use super::id::MemoryId; + /// The type of directed relationship between two memories. /// /// Direction semantics: @@ -52,3 +54,72 @@ impl EdgeType { self as u8 } } + +// ═══════════════════════════════════════════════════════════════════════ +// Supersedes outcome +// ═══════════════════════════════════════════════════════════════════════ + +/// What actually happened to the `Supersedes` edge requested by a store. +/// +/// A store that names a `supersedes` target is pre-validated *before* +/// anything is written, so a missing, cross-namespace, or self-referential +/// target fails the whole store instead of producing an outcome here. Every +/// variant below therefore describes a store that succeeded; the variant +/// says how durable and how new the edge is. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub enum SupersedesStatus { + /// The edge was created in the graph and persisted to `edges.db`. + Applied, + /// An edge of this type already connected the two memories. Supersedes + /// is a desired end state rather than an event, so a repeat store is + /// idempotent rather than an error. + AlreadyApplied, + /// The edge is live in the in-memory graph for this process, but + /// writing it to `edges.db` failed. Recall honours it until restart. + AppliedNotDurable, + /// The edge could not be created at all. Only reachable if the target + /// disappeared between pre-validation and the write. + Failed, +} + +/// The `supersedes` half of a store's result, echoed back to the caller. +/// +/// Present only when the store named a `supersedes` target. Making the +/// outcome explicit is what keeps "silently succeeded while doing nothing" +/// unrepresentable. +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct SupersedesOutcome { + /// The memory this store was asked to supersede. + pub target: MemoryId, + /// What happened to the edge. + pub status: SupersedesStatus, + /// Human-readable explanation, present for the non-`Applied` statuses. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub detail: Option, +} + +impl SupersedesOutcome { + /// Build an outcome with no detail message. + pub fn new(target: MemoryId, status: SupersedesStatus) -> Self { + Self { + target, + status, + detail: None, + } + } + + /// Build an outcome carrying an explanatory detail message. + pub fn with_detail( + target: MemoryId, + status: SupersedesStatus, + detail: impl Into, + ) -> Self { + Self { + target, + status, + detail: Some(detail.into()), + } + } +} diff --git a/src/model/mod.rs b/src/model/mod.rs index 83d882e..525637b 100644 --- a/src/model/mod.rs +++ b/src/model/mod.rs @@ -29,7 +29,7 @@ pub mod validation; // Re-export primary types at module level for convenience. pub use self::decay::DecayPhase; -pub use self::edge::EdgeType; +pub use self::edge::{EdgeType, SupersedesOutcome, SupersedesStatus}; pub use self::error::{DecodeError, TagError, ValidationError}; pub use self::id::{MemoryId, NamespaceId}; pub use self::memory::{AccessEvent, AccessKind, Memory}; diff --git a/src/search/adapters.rs b/src/search/adapters.rs index fcbc783..5d39f4a 100644 --- a/src/search/adapters.rs +++ b/src/search/adapters.rs @@ -472,29 +472,12 @@ impl GraphReader for SharedGraphReader { } fn superseded_by(&self, id: &MemoryId) -> Option { - use crate::model::EdgeType; let graph = tokio::task::block_in_place(|| { tokio::runtime::Handle::current().block_on(self.graph.read()) }); - // Follow the supersedes chain to the latest version. - // Supersedes edge: source (new) → target (old). - // So if `id` is the target of a Supersedes edge, the source is the replacement. - let mut current = *id; - for _ in 0..10 { - let edges = graph.incoming_edges(¤t); - let replacement = edges.iter().find_map(|edge| { - if edge.edge_type == EdgeType::Supersedes { - graph.nodes.get(edge.source).map(|n| n.memory_id) - } else { - None - } - }); - match replacement { - Some(next) => current = next, - None => break, - } - } - if current == *id { None } else { Some(current) } + // Chain-walking (and its cycle guard) lives on the graph itself, + // where it can be unit-tested against real nodes and edges. + graph.supersedes_terminal(id) } } diff --git a/src/search/pipeline.rs b/src/search/pipeline.rs index a1ed063..6646036 100644 --- a/src/search/pipeline.rs +++ b/src/search/pipeline.rs @@ -519,7 +519,7 @@ impl QueryEngine { timings.rank_us = stage_start.elapsed().as_micros() as u64; // -- Stage 8c: Supersedes Resolution ------------------------------ - self.resolve_supersedes(&mut candidates, &query, &query_entities) + self.resolve_supersedes(&mut candidates, &query, &query_entities, ns_config.id) .await; candidates.sort_by(|a, b| { @@ -729,81 +729,140 @@ impl QueryEngine { } } - /// Stage 8c: Resolve superseded memories by replacing them with current versions. + /// Stage 8c: Replace superseded memories with the versions that + /// replace them. + /// + /// Runs as PLAN -> VALIDATE -> COMMIT, and the ordering is the point: + /// a superseded memory is dropped **only** once its replacement has been + /// loaded and cleared every filter the query applies. Removing it any + /// earlier can leave the caller with neither memory — the correction + /// invisible and the original silently deleted from the results. + /// + /// Replacements are gated by the same [`Self::passes_filters`] that + /// stage 5 applies to every other candidate. That is what keeps a + /// replacement in another namespace out of these results: the graph + /// stores no namespace, so without this gate a cross-namespace + /// `Supersedes` edge would pull a foreign memory into the response. + /// + /// `min_score` still does not bind injected replacements, because their + /// `relevance_score` is `None` — identical to graph-expanded candidates, + /// and intentional: they are included for being current, not for + /// matching the query text. async fn resolve_supersedes( &self, candidates: &mut Vec, query: &SearchQuery, query_entities: &[String], + namespace_id: NamespaceId, ) { - use std::collections::HashSet; + use std::collections::{HashMap, HashSet}; - let candidate_id_set: HashSet = candidates.iter().map(|c| c.record.id).collect(); + // -- Plan: which candidates are superseded, and by what. Nothing is + // removed here; that decision waits until the replacement is known + // to be usable. + let planned: Vec<(usize, MemoryId)> = candidates + .iter() + .enumerate() + .filter_map(|(idx, candidate)| { + self.graph + .superseded_by(&candidate.record.id) + .map(|replacement_id| (idx, replacement_id)) + }) + .collect(); - let mut to_remove: HashSet = HashSet::new(); - let mut replacements_to_inject: Vec<(MemoryId, f32)> = Vec::new(); + if planned.is_empty() { + return; + } - for (idx, candidate) in candidates.iter().enumerate() { - if let Some(replacement_id) = self.graph.superseded_by(&candidate.record.id) { - to_remove.insert(idx); - if !candidate_id_set.contains(&replacement_id) { - if let Some(entry) = replacements_to_inject - .iter_mut() - .find(|(id, _)| *id == replacement_id) - { - entry.1 = entry.1.max(candidate.composite_score); - } else { - replacements_to_inject.push((replacement_id, candidate.composite_score)); - } - } - } + // Several superseded candidates can share one replacement; it + // inherits the best of their scores so it ranks at least as high as + // the memory it stands in for. + let mut inherited_scores: HashMap = HashMap::new(); + for (idx, replacement_id) in &planned { + let score = candidates[*idx].composite_score; + inherited_scores + .entry(*replacement_id) + .and_modify(|best| *best = best.max(score)) + .or_insert(score); } - if !replacements_to_inject.is_empty() { - let load_ids: Vec = - replacements_to_inject.iter().map(|(id, _)| *id).collect(); + let candidate_id_set: HashSet = candidates.iter().map(|c| c.record.id).collect(); - if let Ok(records) = self.load_records(&load_ids).await { - let score_map: std::collections::HashMap = - replacements_to_inject.iter().cloned().collect(); + // A replacement already among the candidates passed stage 5 on its + // own, so it is eligible by construction. + let mut eligible: HashSet = inherited_scores + .keys() + .copied() + .filter(|id| candidate_id_set.contains(id)) + .collect(); - for record in records { - // Skip tombstoned and ghost memories. - if record.phase == DecayPhase::Tombstone { - continue; - } - if record.phase == DecayPhase::Ghost && !query.include_ghosts { - continue; - } + let to_load: Vec = inherited_scores + .keys() + .copied() + .filter(|id| !candidate_id_set.contains(id)) + .collect(); - let inherited_score = score_map.get(&record.id).copied().unwrap_or(0.0); - let effective_r = Self::compute_effective_r(&record); + // -- Validate: build each replacement, then subject it to the + // query's filters before letting it displace anything. + let mut injected: Vec = Vec::with_capacity(to_load.len()); + if !to_load.is_empty() { + let records = match self.load_records(&to_load).await { + Ok(records) => records, + Err(e) => { + // Without the replacements there is nothing to put in + // the originals' place, so the originals stay. + warn!( + error = %e, + count = to_load.len(), + "failed to load supersedes replacements; keeping the superseded memories in the results" + ); + return; + } + }; - let entity_score = if !query_entities.is_empty() { - crate::model::entity_overlap(query_entities, &record.entities) - } else { - 0.0 - }; + for record in records { + let inherited_score = inherited_scores.get(&record.id).copied().unwrap_or(0.0); + let effective_r = Self::compute_effective_r(&record); + let entity_score = if query_entities.is_empty() { + 0.0 + } else { + crate::model::entity_overlap(query_entities, &record.entities) + }; - let mut replacement = Candidate::new(record); - replacement.effective_r = effective_r; - replacement.entity_score = entity_score; - replacement.composite_score = Self::compute_composite_score(&replacement); - replacement.composite_score = replacement.composite_score.max(inherited_score); + let mut replacement = Candidate::new(record); + replacement.effective_r = effective_r; + replacement.entity_score = entity_score; + replacement.composite_score = + Self::compute_composite_score(&replacement).max(inherited_score); - candidates.push(replacement); + // Tombstone, ghost, namespace, phase, strength, permastore + // and tag rules all live here — the same gate stage 5 uses. + if !Self::passes_filters(&replacement, query, namespace_id) { + continue; } + + eligible.insert(replacement.record.id); + injected.push(replacement); } } + // -- Commit: drop only the candidates whose replacement made it. + let to_remove: HashSet = planned + .iter() + .filter(|(_, replacement_id)| eligible.contains(replacement_id)) + .map(|(idx, _)| *idx) + .collect(); + if !to_remove.is_empty() { - let mut keep_idx = 0; - candidates.retain(|_| { - let keep = !to_remove.contains(&keep_idx); - keep_idx += 1; - keep - }); + *candidates = std::mem::take(candidates) + .into_iter() + .enumerate() + .filter(|(idx, _)| !to_remove.contains(idx)) + .map(|(_, candidate)| candidate) + .collect(); } + + candidates.extend(injected); } /// Retrieve a single memory by ID. @@ -1174,3 +1233,435 @@ impl QueryEngine { }); } } + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use serde_json::json; + + use super::*; + use crate::model::Tag; + + // -- Stubs ------------------------------------------------------------ + // + // `resolve_supersedes` touches exactly three of the ten injected + // subsystems. The rest exist so a `QueryEngine` can be constructed and + // are wired to panic: if one ever starts being called, the test should + // say so rather than quietly return an empty answer. + + /// Supersedes chains, pre-resolved to their terminal memory. + #[derive(Default)] + struct StubGraph { + terminal: HashMap, + } + + impl GraphReader for StubGraph { + fn neighbors(&self, _id: &MemoryId) -> Vec { + Vec::new() + } + + fn superseded_by(&self, id: &MemoryId) -> Option { + self.terminal.get(id).copied() + } + } + + /// Always misses, so every load goes through the metadata store. + struct StubCache; + + impl RecordCache for StubCache { + fn get(&self, _id: &MemoryId) -> Option { + None + } + } + + /// Backing store for replacement records, with a switch for making the + /// batch load fail. + #[derive(Default)] + struct StubMetaStore { + records: HashMap, + fail: bool, + } + + #[async_trait::async_trait] + impl MetadataStore for StubMetaStore { + async fn get(&self, id: &MemoryId) -> Result> { + Ok(self.records.get(id).cloned()) + } + + async fn get_batch(&self, ids: &[MemoryId]) -> Result> { + if self.fail { + return Err(SearchError::MetadataError("meta.db unavailable".into())); + } + // Missing IDs are skipped, matching the real adapter. + Ok(ids + .iter() + .filter_map(|id| self.records.get(id).cloned()) + .collect()) + } + } + + struct Unused; + + #[async_trait::async_trait] + impl NamespaceResolver for Unused { + async fn resolve(&self, _name: &str) -> Result { + unreachable!("stage 8c does not resolve namespaces") + } + } + + #[async_trait::async_trait] + impl EmbeddingProviderRegistry for Unused { + async fn embed(&self, _namespace: &str, _text: &str) -> Result> { + unreachable!("stage 8c does not embed") + } + } + + impl VectorIndexRegistry for Unused { + fn search(&self, _ns: NamespaceId, _q: &[f32], _k: usize) -> Result> { + unreachable!("stage 8c does not search vectors") + } + + fn get_vector(&self, _id: MemoryId) -> Option> { + unreachable!("stage 8c does not read vectors") + } + } + + impl FtsIndexRegistry for Unused { + fn search(&self, _ns: NamespaceId, _q: &str, _k: usize) -> Result> { + unreachable!("stage 8c does not search FTS") + } + } + + impl EntityIndexReader for Unused { + fn find_by_entities( + &self, + _ns: NamespaceId, + _entities: &[String], + _exclude: MemoryId, + _k: usize, + ) -> Result> { + unreachable!("stage 8c does not use entity recall") + } + } + + impl RifProcessor for Unused { + fn compute_suppressions( + &self, + _retrieved: &[MemoryId], + _neighbors: &[MemoryId], + ) -> Vec { + unreachable!("stage 8c does not apply RIF") + } + } + + #[async_trait::async_trait] + impl AccessRecorder for Unused { + async fn record_access(&self, _id: MemoryId, _kind: AccessKind) -> Result<()> { + unreachable!("stage 8c does not record access") + } + + async fn record_access_batch(&self, _a: &[(MemoryId, AccessKind)]) -> Result<()> { + unreachable!("stage 8c does not record access") + } + } + + // -- Fixtures --------------------------------------------------------- + + fn ns() -> NamespaceId { + NamespaceId::new(1) + } + + fn other_ns() -> NamespaceId { + NamespaceId::new(2) + } + + /// A Full-phase, strong, untagged record in namespace 1. + fn rec(id: MemoryId) -> CachedRecord { + CachedRecord { + id, + namespace_id: ns(), + created_at: 1_700_000_000_000, + last_accessed_at: 1_700_000_000_000, + phase: DecayPhase::Full, + strength: 0.9, + decay_strength: 0.9, + stability: 10.0, + difficulty: 5.0, + is_permastore: false, + summary: "a memory".to_string(), + tags: Vec::new(), + edge_count: 0, + vector_slot: 0, + entities: Vec::new(), + } + } + + /// A candidate as stage 8 would leave it: scored and ranked. + fn candidate(record: CachedRecord, score: f32) -> Candidate { + let mut c = Candidate::new(record); + c.relevance_score = Some(score); + c.raw_vector_score = Some(score); + c.composite_score = score; + c + } + + fn query(value: serde_json::Value) -> SearchQuery { + serde_json::from_value(value).expect("valid query") + } + + fn engine(graph: StubGraph, meta: StubMetaStore) -> QueryEngine { + QueryEngine::new( + Arc::new(Unused), + Arc::new(Unused), + Arc::new(Unused), + Arc::new(Unused), + Arc::new(Unused), + Arc::new(StubCache), + Arc::new(meta), + Arc::new(graph), + Arc::new(Unused), + Arc::new(Unused), + ) + } + + /// Wire up "B is superseded by A", where A's record is `replacement`. + fn superseded_setup(old: MemoryId, replacement: CachedRecord) -> (StubGraph, StubMetaStore) { + let mut graph = StubGraph::default(); + graph.terminal.insert(old, replacement.id); + + let mut meta = StubMetaStore::default(); + meta.records.insert(replacement.id, replacement); + + (graph, meta) + } + + fn ids(candidates: &[Candidate]) -> Vec { + candidates.iter().map(|c| c.record.id).collect() + } + + // -- Tests ------------------------------------------------------------ + + /// The headline behaviour: recall hands back the correction instead of + /// the memory it corrects, ranked no lower than what it replaced. + #[tokio::test] + async fn substitutes_an_eligible_replacement_for_the_superseded_memory() { + let (old, new) = (MemoryId::new(), MemoryId::new()); + let (graph, meta) = superseded_setup(old, rec(new)); + let engine = engine(graph, meta); + + let mut candidates = vec![candidate(rec(old), 0.8)]; + engine + .resolve_supersedes(&mut candidates, &query(json!({"text": "q"})), &[], ns()) + .await; + + assert_eq!(ids(&candidates), vec![new]); + assert!(candidates[0].composite_score >= 0.8); + } + + /// When the replacement is already a hit in its own right, the old + /// memory still goes, and the new one is not duplicated. + #[tokio::test] + async fn replacement_already_among_the_candidates_appears_exactly_once() { + let (old, new) = (MemoryId::new(), MemoryId::new()); + let (graph, meta) = superseded_setup(old, rec(new)); + let engine = engine(graph, meta); + + let mut candidates = vec![candidate(rec(old), 0.8), candidate(rec(new), 0.9)]; + engine + .resolve_supersedes(&mut candidates, &query(json!({"text": "q"})), &[], ns()) + .await; + + assert_eq!(ids(&candidates), vec![new]); + } + + /// A tombstoned replacement cannot be returned — and because it cannot, + /// the memory it replaces must survive. Dropping both would lose the + /// information entirely. + #[tokio::test] + async fn tombstoned_replacement_does_not_drop_the_original() { + let (old, new) = (MemoryId::new(), MemoryId::new()); + let mut replacement = rec(new); + replacement.phase = DecayPhase::Tombstone; + let (graph, meta) = superseded_setup(old, replacement); + let engine = engine(graph, meta); + + let mut candidates = vec![candidate(rec(old), 0.8)]; + engine + .resolve_supersedes(&mut candidates, &query(json!({"text": "q"})), &[], ns()) + .await; + + assert_eq!(ids(&candidates), vec![old]); + } + + /// Same reasoning for a ghost the query did not ask for. + #[tokio::test] + async fn ghost_replacement_is_kept_out_unless_ghosts_were_requested() { + let (old, new) = (MemoryId::new(), MemoryId::new()); + let mut replacement = rec(new); + replacement.phase = DecayPhase::Ghost; + let (graph, meta) = superseded_setup(old, replacement); + let engine = engine(graph, meta); + + let mut candidates = vec![candidate(rec(old), 0.8)]; + engine + .resolve_supersedes(&mut candidates, &query(json!({"text": "q"})), &[], ns()) + .await; + + assert_eq!(ids(&candidates), vec![old]); + } + + /// ...and when they were requested, the ghost does substitute. + #[tokio::test] + async fn ghost_replacement_substitutes_when_ghosts_are_included() { + let (old, new) = (MemoryId::new(), MemoryId::new()); + let mut replacement = rec(new); + replacement.phase = DecayPhase::Ghost; + let (graph, meta) = superseded_setup(old, replacement); + let engine = engine(graph, meta); + + let mut candidates = vec![candidate(rec(old), 0.8)]; + engine + .resolve_supersedes( + &mut candidates, + &query(json!({"text": "q", "includeGhosts": true})), + &[], + ns(), + ) + .await; + + assert_eq!(ids(&candidates), vec![new]); + } + + /// Namespace isolation regression. A `Supersedes` edge can point across + /// namespaces (older data, or a hand-written edges.db), and the graph + /// stores no namespace to catch it. Before replacements were run through + /// `passes_filters`, this stage injected the foreign memory straight + /// into the response — a cross-namespace data leak. + #[tokio::test] + async fn replacement_from_another_namespace_is_not_leaked() { + let (old, new) = (MemoryId::new(), MemoryId::new()); + let mut replacement = rec(new); + replacement.namespace_id = other_ns(); + let (graph, meta) = superseded_setup(old, replacement); + let engine = engine(graph, meta); + + let mut candidates = vec![candidate(rec(old), 0.8)]; + engine + .resolve_supersedes(&mut candidates, &query(json!({"text": "q"})), &[], ns()) + .await; + + assert_eq!(ids(&candidates), vec![old]); + } + + /// The rest of the filter set binds replacements too: required tags, + /// minimum strength, and the phase whitelist. + #[tokio::test] + async fn replacement_failing_metadata_filters_is_not_substituted() { + /// (label, query, how to make the replacement fail that query). + type FilterCase = (&'static str, serde_json::Value, fn(&mut CachedRecord)); + + let cases: Vec = vec![ + ( + "required tag missing", + json!({"text": "q", "filter": {"requireTags": ["project/recalld"]}}), + |_r| {}, + ), + ( + "below minimum strength", + json!({"text": "q", "filter": {"minStrength": 0.95}}), + |r| r.strength = 0.2, + ), + ( + "phase not in whitelist", + json!({"text": "q", "filter": {"phases": ["summary"]}}), + |r| r.phase = DecayPhase::Full, + ), + ( + "not permastore", + json!({"text": "q", "filter": {"permastoreOnly": true}}), + |r| r.is_permastore = false, + ), + ]; + + for (label, q, mutate) in cases { + let (old, new) = (MemoryId::new(), MemoryId::new()); + let mut replacement = rec(new); + mutate(&mut replacement); + let (graph, meta) = superseded_setup(old, replacement); + let engine = engine(graph, meta); + + // The superseded memory itself satisfies the filter, so only the + // replacement's eligibility is under test. + let mut original = rec(old); + original.tags = vec![Tag::new("project/recalld").unwrap()]; + original.phase = DecayPhase::Summary; + original.is_permastore = true; + + let mut candidates = vec![candidate(original, 0.8)]; + engine + .resolve_supersedes(&mut candidates, &query(q), &[], ns()) + .await; + + assert_eq!(ids(&candidates), vec![old], "{label}"); + } + } + + /// If the replacement cannot be loaded, the original is all we have. + /// The old code swallowed the error and removed it anyway. + #[tokio::test] + async fn metadata_store_failure_keeps_the_superseded_memory() { + let (old, new) = (MemoryId::new(), MemoryId::new()); + let (graph, mut meta) = superseded_setup(old, rec(new)); + meta.fail = true; + let engine = engine(graph, meta); + + let mut candidates = vec![candidate(rec(old), 0.8)]; + engine + .resolve_supersedes(&mut candidates, &query(json!({"text": "q"})), &[], ns()) + .await; + + assert_eq!(ids(&candidates), vec![old]); + } + + /// Likewise when the replacement's record simply is not there. + #[tokio::test] + async fn replacement_with_no_record_anywhere_keeps_the_superseded_memory() { + let old = MemoryId::new(); + let mut graph = StubGraph::default(); + graph.terminal.insert(old, MemoryId::new()); + let engine = engine(graph, StubMetaStore::default()); + + let mut candidates = vec![candidate(rec(old), 0.8)]; + engine + .resolve_supersedes(&mut candidates, &query(json!({"text": "q"})), &[], ns()) + .await; + + assert_eq!(ids(&candidates), vec![old]); + } + + /// Two memories superseded by the same correction collapse into one + /// copy of it, carrying the better of the two scores. + #[tokio::test] + async fn many_superseded_memories_share_one_replacement() { + let (old_a, old_b, new) = (MemoryId::new(), MemoryId::new(), MemoryId::new()); + let mut graph = StubGraph::default(); + graph.terminal.insert(old_a, new); + graph.terminal.insert(old_b, new); + let mut meta = StubMetaStore::default(); + meta.records.insert(new, rec(new)); + let engine = engine(graph, meta); + + let mut candidates = vec![candidate(rec(old_a), 0.5), candidate(rec(old_b), 0.85)]; + engine + .resolve_supersedes(&mut candidates, &query(json!({"text": "q"})), &[], ns()) + .await; + + assert_eq!(ids(&candidates), vec![new]); + assert!(candidates[0].composite_score >= 0.85); + } +} diff --git a/src/serialization/json.rs b/src/serialization/json.rs index c0a2945..1685f79 100644 --- a/src/serialization/json.rs +++ b/src/serialization/json.rs @@ -130,6 +130,13 @@ pub struct MemoryResponse { /// Access history (included only when explicitly requested). #[serde(skip_serializing_if = "Option::is_none")] pub access_history: Option>, + + /// What became of the requested `supersedes` link. + /// + /// Only the create endpoint populates this, and only when the request + /// named a target — so GET, search and list bodies are unchanged. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub supersedes: Option, } impl MemoryResponse { @@ -158,6 +165,7 @@ impl MemoryResponse { edge_count: record.edge_count, embedding: None, access_history: None, + supersedes: None, } } } From 16e78cf77def4eec119b10eaad1f02cd53bd97cb Mon Sep 17 00:00:00 2001 From: Caleb Evans Date: Tue, 11 Aug 2026 01:39:51 -0600 Subject: [PATCH 6/8] fix: bound label length and stop discarding tags in silence Entities, topics and emotions had a count limit but no per-string length limit anywhere. Two consequences, both silent. First, the derived tag. Those labels are merged into tags as "entity/" after validation, so a 128-byte entity became a 135-byte tag, which Tag::new rejected, which filter_map(..ok()) discarded without an error or a log. The caller sent 20 tags, 17 were stored, and nothing said so. The limits now reserve the prefix -- 121 bytes for entities, 122 for topics, 120 for emotions -- so a request that validates cannot produce a tag that gets thrown away, and a test asserts the derived tag survives Tag::new at exactly the limit. Second, the frame budget. The 8 MiB daemon frame limit was justified by arithmetic that assumed 128-byte labels; without a per-string limit the envelope was unbounded and that justification did not hold. There is now a test that builds the worst-case valid request out of control bytes, which JSON escapes at 6x, and asserts it serializes within MAX_MESSAGE_SIZE. It fails the moment anyone raises a content limit or lowers the frame limit. - Validate every element of tags, entities, topics and emotions in the shared validator, naming the field, the index, the measured bytes and the limit. Counts are still reported first when both fail. - Replace the silent tag drops with parse_tags_lossy, which logs each drop with the field, the value and the reason. It stays a drop rather than an error because the reachable case is the tag alphabet, not length -- an entity named Jose with an accent is legitimate input that should not fail a store. - Close the daemon RPC store_memory path, which deserialized StoreInput straight off the socket with no limit checks at all. It was the last door through which an over-long label reached storage and was dropped while the call reported success. - Advertise maxLength on the label arrays in both tool schemas, and document the limits and the prefix reservation. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01JNhUehmChiQKJUjv3QPnkH --- docs/guide.md | 8 +- docs/mcp.md | 20 +++-- src/api/adapters.rs | 21 +++-- src/api/handlers.rs | 27 +++--- src/daemon/protocol.rs | 83 +++++++++++++++++- src/daemon/server.rs | 13 +++ src/mcp/args.rs | 28 ++++++ src/mcp/bridge_adapters.rs | 36 ++++---- src/mcp/tools.rs | 37 +++++--- src/model/constants.rs | 44 ++++++++++ src/model/error.rs | 21 ++++- src/model/mod.rs | 4 +- src/model/tag.rs | 36 ++++++++ src/model/validation.rs | 171 ++++++++++++++++++++++++++++++++++++- 14 files changed, 476 insertions(+), 73 deletions(-) diff --git a/docs/guide.md b/docs/guide.md index e47c607..467bfc9 100644 --- a/docs/guide.md +++ b/docs/guide.md @@ -757,10 +757,10 @@ Store a new memory. |---|---|---|---| | `summary` | string | Yes | Short description (max 2000 bytes of UTF-8; ~2000 ASCII characters, fewer with em dashes/curly quotes/emoji). | | `fullText` | string | No | Detailed content (max 1 MB = 1 048 576 bytes of UTF-8, counted on the raw text before JSON escaping). Dropped as memory decays to ghost phase. | -| `tags` | string[] | No | Categorization tags, e.g. `["topic/rust", "type/observation"]`. Max 64. | -| `entities` | string[] | No | Named entities (people, places, orgs). Used for search indexing and graph linking. Max 32. | -| `topics` | string[] | No | Topic keywords, e.g. `["rust", "cooking"]`. Max 32. | -| `emotions` | string[] | No | Emotional tone, e.g. `["happy", "anxious"]`. Max 32. | +| `tags` | string[] | No | Categorization tags, e.g. `["topic/rust", "type/observation"]`. Max 64 items, 128 bytes of UTF-8 each. | +| `entities` | string[] | No | Named entities (people, places, orgs). Used for search indexing and graph linking. Max 32 items, 121 bytes of UTF-8 each (the derived `entity/