diff --git a/.gitignore b/.gitignore index a72f6d6..a19bf28 100644 --- a/.gitignore +++ b/.gitignore @@ -48,3 +48,4 @@ nohup.out sdk/python/.venv/ sdk/python/uv.lock .DS_Store +memoria/.DS_Store diff --git a/memoria/crates/memoria-api/src/models.rs b/memoria/crates/memoria-api/src/models.rs index 3630918..d53824e 100644 --- a/memoria/crates/memoria-api/src/models.rs +++ b/memoria/crates/memoria-api/src/models.rs @@ -19,6 +19,10 @@ pub struct StoreRequest { pub observed_at: Option, pub source: Option, pub branch: Option, + /// 任意业务元数据(如 scene/agent)。透传落库到 memories.extra_metadata,并在读取时原样 + /// 返回给调用方;Memoria 本身不对其做检索/打分逻辑(下游消费者如 matrixflow 的 decay 可自行使用)。 + #[serde(default)] + pub extra_metadata: Option>, } fn default_memory_type() -> String { "semantic".to_string() @@ -281,6 +285,10 @@ pub struct MemoryResponse { pub observed_at: Option, pub created_at: Option, pub retrieval_score: Option, + /// 业务元数据(如 scene/agent)从 memories.extra_metadata 原样透传回给调用方;Memoria 本身 + /// 不对其做检索/打分逻辑(下游消费者如 matrixflow 的 decay 可自行使用)。 + #[serde(skip_serializing_if = "Option::is_none")] + pub extra_metadata: Option>, } impl From for MemoryResponse { @@ -299,6 +307,9 @@ impl From for MemoryResponse { observed_at: m.observed_at.map(|dt| dt.to_rfc3339()), created_at: m.created_at.map(|dt| dt.to_rfc3339()), retrieval_score: m.retrieval_score, + // 空 map 在响应里归一为 None(省略),与读取侧「"{}" → None」一致:新写记录直接 + // 返回内存对象时也不会出现「POST 带 {} 而后续 list/get 无该字段」的不一致。 + extra_metadata: m.extra_metadata.filter(|md| !md.is_empty()), } } } diff --git a/memoria/crates/memoria-api/src/routes/memory.rs b/memoria/crates/memoria-api/src/routes/memory.rs index b0ec2c0..eb0a0a6 100644 --- a/memoria/crates/memoria-api/src/routes/memory.rs +++ b/memoria/crates/memoria-api/src/routes/memory.rs @@ -220,7 +220,7 @@ pub async fn store_memory( }; let m = state .service - .store_memory_on_branch( + .store_memory_with_metadata_on_branch( auth.scope_id(), branch_param(req.branch.as_deref()), &req.content, @@ -231,6 +231,7 @@ pub async fn store_memory( req.initial_confidence, author, req.subject_id, + req.extra_metadata, ) .await .map_err(|e| { @@ -284,16 +285,24 @@ pub async fn batch_store( .map_err(|e| (StatusCode::UNPROCESSABLE_ENTITY, e)); // item-level subject_id takes priority; fall back to batch-level let subject_id = r.subject_id.or_else(|| batch_subject_id.clone()); - Ok((r.content, mt, tier, r.session_id, top_branch.clone(), subject_id)) + Ok(( + r.content, + mt, + tier, + r.session_id, + top_branch.clone(), + subject_id, + r.extra_metadata, + )) }) .collect::, _>>()?; // Validate all types upfront let mut validated = Vec::with_capacity(items.len()); - for (content, mt_result, tier_result, session_id, branch, subject_id) in items { + for (content, mt_result, tier_result, session_id, branch, subject_id, extra_metadata) in items { let mt = mt_result?; let tier = tier_result?; - validated.push((content, mt, session_id, tier, branch, subject_id)); + validated.push((content, mt, session_id, tier, branch, subject_id, extra_metadata)); } let author = if auth.group_id.is_some() { @@ -303,13 +312,13 @@ pub async fn batch_store( }; let batch_items = validated .into_iter() - .map(|(content, mt, session_id, tier, _branch, subject_id)| { - (content, mt, session_id, tier, subject_id) + .map(|(content, mt, session_id, tier, _branch, subject_id, extra_metadata)| { + (content, mt, session_id, tier, subject_id, extra_metadata) }) .collect(); let results = state .service - .store_batch_on_branch( + .store_batch_with_metadata_on_branch( auth.scope_id(), branch_param(top_branch.as_deref()), batch_items, diff --git a/memoria/crates/memoria-api/tests/api_e2e.rs b/memoria/crates/memoria-api/tests/api_e2e.rs index 55c9071..edb24a9 100644 --- a/memoria/crates/memoria-api/tests/api_e2e.rs +++ b/memoria/crates/memoria-api/tests/api_e2e.rs @@ -259,6 +259,141 @@ async fn test_api_store_and_list() { ); } +/// #224: metadata must survive every REST write/read path. This exercises the +/// MatrixOne JSON binding and the lightweight list mapper as well as the full +/// mappers used by get/retrieve/search. +#[tokio::test] +async fn test_api_extra_metadata_round_trip_and_dedup() { + let (base, client, _server) = + spawn_server_with_custom_embedder_and_pool(Arc::new(SessionScopeTestEmbedder), test_dim()) + .await; + let user_id = uid(); + let metadata = json!({"scene": "incident", "agent": "triage", "rank": 2}); + + let response = client + .post(format!("{base}/v1/memories")) + .header("X-User-Id", &user_id) + .json(&json!({ + "content": "metadata round trip needle", + "memory_type": "semantic", + "extra_metadata": metadata, + })) + .send() + .await + .expect("store with metadata"); + assert_eq!(response.status(), 201); + let stored: Value = response.json().await.unwrap(); + let memory_id = stored["memory_id"].as_str().unwrap().to_string(); + assert_eq!(stored["extra_metadata"], metadata); + + // Same-content dedup returns the persisted survivor, not a phantom ID, and + // refreshes its metadata. + let updated_metadata = json!({"scene": "resolved", "agent": "review"}); + let response = client + .post(format!("{base}/v1/memories")) + .header("X-User-Id", &user_id) + .json(&json!({ + "content": "metadata round trip needle", + "memory_type": "semantic", + "extra_metadata": updated_metadata, + })) + .send() + .await + .expect("dedup metadata update"); + assert_eq!(response.status(), 201); + let deduped: Value = response.json().await.unwrap(); + assert_eq!(deduped["memory_id"], memory_id); + assert_eq!(deduped["extra_metadata"], updated_metadata); + + let response = client + .post(format!("{base}/v1/memories/batch")) + .header("X-User-Id", &user_id) + .json(&json!({"memories": [{ + "content": "batch metadata round trip", + "memory_type": "semantic", + "extra_metadata": {"source": "batch"} + }]})) + .send() + .await + .expect("batch store with metadata"); + assert_eq!(response.status(), 201); + let batch: Value = response.json().await.unwrap(); + assert_eq!(batch[0]["extra_metadata"], json!({"source": "batch"})); + + let response = client + .get(format!("{base}/v1/memories")) + .header("X-User-Id", &user_id) + .send() + .await + .expect("list"); + let list: Value = response.json().await.unwrap(); + assert_eq!( + list["items"] + .as_array() + .unwrap() + .iter() + .find(|item| item["memory_id"] == memory_id) + .unwrap()["extra_metadata"], + updated_metadata + ); + + let response = client + .get(format!("{base}/v1/memories/{memory_id}")) + .header("X-User-Id", &user_id) + .send() + .await + .expect("get"); + let got: Value = response.json().await.unwrap(); + assert_eq!(got["extra_metadata"], updated_metadata); + + for endpoint in ["retrieve", "search"] { + let response = client + .post(format!("{base}/v1/memories/{endpoint}")) + .header("X-User-Id", &user_id) + .json(&json!({"query": "metadata round trip needle", "top_k": 10})) + .send() + .await + .expect("retrieve/search"); + assert_eq!(response.status(), 200); + let results: Value = response.json().await.unwrap(); + assert_eq!( + results + .as_array() + .unwrap() + .iter() + .find(|item| item["memory_id"] == memory_id) + .unwrap()["extra_metadata"], + updated_metadata, + "{endpoint} must return metadata" + ); + } + + // `{}` is the explicit clear value but is normalized to an omitted response + // field, while scalars/arrays are rejected by StoreRequest deserialization. + let response = client + .post(format!("{base}/v1/memories")) + .header("X-User-Id", &user_id) + .json(&json!({"content": "empty metadata", "extra_metadata": {}})) + .send() + .await + .expect("store empty metadata"); + assert_eq!(response.status(), 201); + assert!(response + .json::() + .await + .unwrap() + .get("extra_metadata") + .is_none()); + let response = client + .post(format!("{base}/v1/memories")) + .header("X-User-Id", &user_id) + .json(&json!({"content": "bad metadata", "extra_metadata": ["not", "an", "object"]})) + .send() + .await + .expect("reject non-object metadata"); + assert_eq!(response.status(), 422); +} + // ── 2b. list response is lightweight (no embedding) and respects limit ──────── #[tokio::test] @@ -313,7 +448,7 @@ async fn test_api_list_no_embedding_and_limit() { ); assert!( item.get("extra_metadata").is_none(), - "must not contain extra_metadata" + "metadata is omitted when no metadata was stored" ); // Must contain core fields assert!(item["memory_id"].as_str().is_some(), "memory_id required"); @@ -1895,6 +2030,39 @@ async fn test_remote_store_retrieve() { println!("✅ remote retrieve: {}", &text[..text.len().min(80)]); } +#[tokio::test] +async fn test_remote_store_metadata_round_trip() { + use memoria_mcp::remote::RemoteClient; + + let (base, _, server) = spawn_api_for_remote().await; + let user_id = uid(); + let remote = RemoteClient::new(&base, None, user_id.clone(), None); + remote + .call( + "memory_store", + json!({ + "content": "remote metadata memory", + "extra_metadata": {"scene": "remote", "agent": "mcp"} + }), + ) + .await + .expect("remote store"); + + let stored = server + .service() + .list_active(&user_id, 10) + .await + .unwrap() + .into_iter() + .find(|memory| memory.content == "remote metadata memory") + .expect("stored remote memory"); + assert_eq!( + serde_json::to_value(stored.extra_metadata).unwrap(), + json!({"scene": "remote", "agent": "mcp"}), + "remote MCP must forward metadata to REST" + ); +} + #[tokio::test] async fn test_remote_retrieve_session_scope_only_includes_unscoped() { use memoria_mcp::remote::RemoteClient; @@ -8703,6 +8871,108 @@ async fn test_mcp_tools_call_memory_store() { println!("✅ POST /mcp tools/call memory_store: {text}"); } +#[tokio::test] +async fn test_mcp_memory_store_rejects_empty_content() { + let (base, client, _server) = spawn_server().await; + let uid = uid(); + + for (id, args) in [ + (4, json!({})), + (5, json!({"content": ""})), + (6, json!({"content": " \t"})), + ] { + let resp = mcp_post_with_headers( + &client, + &base, + json!({ + "jsonrpc": "2.0", + "id": id, + "method": "tools/call", + "params": { + "name": "memory_store", + "arguments": args + } + }), + &[("X-User-Id", uid.as_str())], + ) + .await; + + assert_eq!(resp["jsonrpc"], "2.0"); + assert_eq!(resp["id"], id); + assert!( + resp["error"].is_null(), + "validation should return tool text, not RPC error: {}", + resp["error"] + ); + let text = resp["result"]["content"][0]["text"].as_str().unwrap_or(""); + assert!( + text.contains("content is required"), + "unexpected response for args={args}: {text}" + ); + } + + let list = client + .get(format!("{base}/v1/memories")) + .header("X-User-Id", &uid) + .send() + .await + .unwrap(); + assert_eq!(list.status(), 200); + assert!( + list.json::().await.unwrap()["items"] + .as_array() + .unwrap() + .is_empty(), + "empty MCP store must not create memories" + ); + println!("✅ POST /mcp memory_store rejects empty content"); +} + +#[tokio::test] +async fn test_mcp_memory_retrieve_and_search_reject_missing_query() { + let (base, client, _server) = spawn_server().await; + let uid = uid(); + + for (id, tool, args) in [ + (7, "memory_retrieve", json!({})), + (8, "memory_retrieve", json!({"query": ""})), + (9, "memory_retrieve", json!({"query": " \t"})), + (10, "memory_search", json!({})), + (11, "memory_search", json!({"query": ""})), + (12, "memory_search", json!({"query": " \t"})), + ] { + let resp = mcp_post_with_headers( + &client, + &base, + json!({ + "jsonrpc": "2.0", + "id": id, + "method": "tools/call", + "params": { + "name": tool, + "arguments": args + } + }), + &[("X-User-Id", uid.as_str())], + ) + .await; + + assert_eq!(resp["jsonrpc"], "2.0"); + assert_eq!(resp["id"], id); + assert!( + resp["error"].is_null(), + "validation should return tool text, not RPC error: {}", + resp["error"] + ); + let text = resp["result"]["content"][0]["text"].as_str().unwrap_or(""); + assert!( + text.contains("query is required"), + "{tool} unexpected response for args={args}: {text}" + ); + } + println!("✅ POST /mcp memory_retrieve/search reject missing query"); +} + #[tokio::test] async fn test_mcp_memory_retrieve_session_scope_end_to_end() { let (base, client, _server) = diff --git a/memoria/crates/memoria-mcp/src/remote.rs b/memoria/crates/memoria-mcp/src/remote.rs index 2b3f1c3..c4ce383 100644 --- a/memoria/crates/memoria-mcp/src/remote.rs +++ b/memoria/crates/memoria-mcp/src/remote.rs @@ -120,6 +120,12 @@ impl RemoteClient { } pub async fn call(&self, name: &str, args: Value) -> Result { + // remote 模式直接拼 REST payload,会绕过 embedded handler 的必填校验;这里前置调用 + // 与 embedded 共享的校验。失败返回**软** tool result 文本(error=null),与 embedded + // 契约一致(不再转成 JSON-RPC -32000)。 + if let Err(e) = crate::tools::validate_tool_args(name, &args) { + return Ok(Self::mcp_text(e)); + } match name { "memory_store" => { let mut payload = json!({ @@ -137,6 +143,9 @@ impl RemoteClient { { payload["subject_id"] = json!(sid); } + if args.get("extra_metadata").map(Value::is_object).unwrap_or(false) { + payload["extra_metadata"] = args["extra_metadata"].clone(); + } let r = self .client .post(self.url("/v1/memories")) @@ -200,11 +209,17 @@ impl RemoteClient { let text = mems .iter() .map(|m| { + let meta = m + .get("extra_metadata") + .filter(|v| v.is_object()) + .map(|v| format!(" | metadata: {v}")) + .unwrap_or_default(); format!( - "[{}] ({}) {}", + "[{}] ({}) {}{}", m["memory_id"].as_str().unwrap_or(""), m["memory_type"].as_str().unwrap_or(""), - m["content"].as_str().unwrap_or("") + m["content"].as_str().unwrap_or(""), + meta ) }) .collect::>() @@ -389,11 +404,17 @@ impl RemoteClient { let text = items .iter() .map(|m| { + let meta = m + .get("extra_metadata") + .filter(|v| v.is_object()) + .map(|v| format!(" | metadata: {v}")) + .unwrap_or_default(); format!( - "[{}] ({}) {}", + "[{}] ({}) {}{}", m["memory_id"].as_str().unwrap_or(""), m["memory_type"].as_str().unwrap_or(""), - m["content"].as_str().unwrap_or("") + m["content"].as_str().unwrap_or(""), + meta ) }) .collect::>() diff --git a/memoria/crates/memoria-mcp/src/tools.rs b/memoria/crates/memoria-mcp/src/tools.rs index cd12621..6dd17f5 100644 --- a/memoria/crates/memoria-mcp/src/tools.rs +++ b/memoria/crates/memoria-mcp/src/tools.rs @@ -94,6 +94,45 @@ fn branch_arg(args: &Value) -> Option<&str> { .filter(|branch| !branch.is_empty()) } +fn parse_required_str(args: &Value, key: &str, err: &'static str) -> Result { + let raw = args[key].as_str().unwrap_or(""); + if raw.trim().is_empty() { + Err(err) + } else { + Ok(raw.to_string()) + } +} + +fn parse_store_content(args: &Value) -> Result { + parse_required_str(args, "content", "content is required") +} + +/// 校验各工具的必填非空字符串参数。embedded 分发在各 handler 内部已隐式校验(parse_required_str), +/// 但 remote 模式直接拼 REST payload 会绕过校验——remote::call 前置调用本函数以保持一致。 +pub fn validate_tool_args(name: &str, args: &Value) -> Result<(), &'static str> { + // extra_metadata 若存在必须是 object(与 REST StoreRequest 一致);非 object 明确拒绝, + // 不能静默丢弃当作缺失。 + if let Some(v) = args.get("extra_metadata") { + if !v.is_null() && !v.is_object() { + return Err("extra_metadata must be an object"); + } + } + match name { + "memory_store" => parse_required_str(args, "content", "content is required").map(|_| ()), + "memory_retrieve" | "memory_search" => { + parse_required_str(args, "query", "query is required").map(|_| ()) + } + "memory_correct" => { + parse_required_str(args, "new_content", "new_content is required").map(|_| ()) + } + _ => Ok(()), + } +} + +fn parse_retrieve_query(args: &Value) -> Result { + parse_required_str(args, "query", "query is required") +} + enum ToolCallName { MemoryStore, MemoryRetrieve, @@ -156,6 +195,7 @@ pub fn list() -> Value { "content": {"type": "string"}, "memory_type": {"type": "string", "default": "semantic"}, "session_id": {"type": "string"}, + "extra_metadata": {"type": "object", "description": "Optional business metadata (e.g. scene, agent). Persisted with the memory and returned verbatim on read. Memoria itself does not score by it; downstream consumers may use it for their own retrieval-time ranking."}, "subject_id": {"type": "string", "description": "Stable business ID of the memory subject (e.g. end-user ID). Set by the integration layer; optional."}, "branch": {"type": "string", "description": "Optional branch to read/write without changing the active checkout"}, "trust_tier": { @@ -335,6 +375,11 @@ pub async fn call( user_id: &str, ) -> Result { tracing::debug!(tool = name, user_id, "MCP tool call"); + // 与 remote 模式共享的参数校验(必填非空 + extra_metadata 类型)。失败返回软 tool result + // 文本(error=null),与既有 embedded 错误契约一致。 + if let Err(e) = validate_tool_args(name, &args) { + return Ok(mcp_text(e)); + } let tool = match name { "memory_store" => ToolCallName::MemoryStore, "memory_retrieve" => ToolCallName::MemoryRetrieve, @@ -358,7 +403,10 @@ pub async fn call( }; match tool { ToolCallName::MemoryStore => { - let content = args["content"].as_str().unwrap_or("").to_string(); + let content = match parse_store_content(&args) { + Ok(content) => content, + Err(msg) => return Ok(mcp_text(msg)), + }; let memory_type = args["memory_type"].as_str().unwrap_or("semantic"); let session_id = args["session_id"].as_str().map(String::from); let trust_tier = args["trust_tier"] @@ -372,8 +420,17 @@ pub async fn call( .map(str::trim) .filter(|s| !s.is_empty()) .map(String::from); + let extra_metadata = args + .get("extra_metadata") + .filter(|v| v.is_object()) + .and_then(|v| { + serde_json::from_value::>( + v.clone(), + ) + .ok() + }); let m = match service - .store_memory_on_branch( + .store_memory_with_metadata_on_branch( user_id, branch_arg(&args), &content, @@ -384,6 +441,7 @@ pub async fn call( None, None, subject_id, + extra_metadata, ) .await { @@ -447,7 +505,10 @@ pub async fn call( } ToolCallName::MemoryRetrieve | ToolCallName::MemorySearch => { - let query = args["query"].as_str().unwrap_or("").to_string(); + let query = match parse_retrieve_query(&args) { + Ok(query) => query, + Err(msg) => return Ok(mcp_text(msg)), + }; let top_k = if matches!(tool, ToolCallName::MemorySearch) { args["top_k"].as_i64().unwrap_or(10) } else { @@ -481,7 +542,7 @@ pub async fn call( } let text = results .iter() - .map(|m| format!("[{}] ({}) {}", m.memory_id, m.memory_type, m.content)) + .map(|m| format!("[{}] ({}) {}{}", m.memory_id, m.memory_type, m.content, m.extra_metadata.as_ref().filter(|md| !md.is_empty()).map(|md| format!(" | metadata: {}", serde_json::to_string(md).unwrap_or_default())).unwrap_or_default())) .collect::>() .join("\n"); let explain_json = serde_json::to_string_pretty(&stats).unwrap_or_default(); @@ -503,7 +564,7 @@ pub async fn call( } let text = results .iter() - .map(|m| format!("[{}] ({}) {}", m.memory_id, m.memory_type, m.content)) + .map(|m| format!("[{}] ({}) {}{}", m.memory_id, m.memory_type, m.content, m.extra_metadata.as_ref().filter(|md| !md.is_empty()).map(|md| format!(" | metadata: {}", serde_json::to_string(md).unwrap_or_default())).unwrap_or_default())) .collect::>() .join("\n"); Ok(mcp_text(&text)) @@ -511,10 +572,10 @@ pub async fn call( } ToolCallName::MemoryCorrect => { - let new_content = args["new_content"].as_str().unwrap_or(""); - if new_content.is_empty() { - return Ok(mcp_text("new_content is required")); - } + let new_content = match parse_required_str(&args, "new_content", "new_content is required") { + Ok(s) => s, + Err(msg) => return Ok(mcp_text(msg)), + }; let memory_id = args["memory_id"].as_str().unwrap_or(""); let query = args["query"].as_str().unwrap_or(""); @@ -541,7 +602,7 @@ pub async fn call( }; let m = service - .correct_on_branch(user_id, branch_arg(&args), &old_mid, new_content) + .correct_on_branch(user_id, branch_arg(&args), &old_mid, &new_content) .await?; Ok(mcp_text(&format!( @@ -652,7 +713,7 @@ pub async fn call( } let text = memories .iter() - .map(|m| format!("[{}] ({}) {}", m.memory_id, m.memory_type, m.content)) + .map(|m| format!("[{}] ({}) {}{}", m.memory_id, m.memory_type, m.content, m.extra_metadata.as_ref().filter(|md| !md.is_empty()).map(|md| format!(" | metadata: {}", serde_json::to_string(md).unwrap_or_default())).unwrap_or_default())) .collect::>() .join("\n"); Ok(mcp_text(&text)) @@ -1436,6 +1497,64 @@ pub fn entity_extract_prompt(text: &str) -> String { mod tests { use super::*; + #[test] + fn parse_required_str_rejects_missing_and_blank() { + for val in [json!({}), json!({"k": ""}), json!({"k": " \t\n"})] { + assert!(parse_required_str(&val, "k", "k is required").is_err()); + } + } + + #[test] + fn parse_required_str_preserves_whitespace_when_valid() { + assert_eq!( + parse_required_str(&json!({"k": " hello "}), "k", "k is required").unwrap(), + " hello " + ); + } + + #[test] + fn parse_store_content_rejects_missing_and_blank() { + assert_eq!(parse_store_content(&json!({})), Err("content is required")); + assert_eq!( + parse_store_content(&json!({"content": ""})), + Err("content is required") + ); + assert_eq!( + parse_store_content(&json!({"content": " \t\n"})), + Err("content is required") + ); + } + + #[test] + fn parse_retrieve_query_rejects_missing_and_blank() { + assert_eq!(parse_retrieve_query(&json!({})), Err("query is required")); + assert_eq!( + parse_retrieve_query(&json!({"query": ""})), + Err("query is required") + ); + assert_eq!( + parse_retrieve_query(&json!({"query": " \t\n"})), + Err("query is required") + ); + } + + #[test] + fn parse_correct_new_content_rejects_blank() { + // memory_correct now uses parse_required_str — verify whitespace is rejected + assert_eq!( + parse_required_str(&json!({"new_content": ""}), "new_content", "new_content is required"), + Err("new_content is required") + ); + assert_eq!( + parse_required_str(&json!({"new_content": " "}), "new_content", "new_content is required"), + Err("new_content is required") + ); + assert_eq!( + parse_required_str(&json!({"new_content": "ok"}), "new_content", "new_content is required").unwrap(), + "ok" + ); + } + #[test] fn entity_extract_prompt_truncates_at_char_boundary() { // 2000 ASCII bytes then a 3-byte Chinese char — must not panic diff --git a/memoria/crates/memoria-mcp/tests/core_tools_e2e.rs b/memoria/crates/memoria-mcp/tests/core_tools_e2e.rs index 32e8d92..1fa7c8a 100644 --- a/memoria/crates/memoria-mcp/tests/core_tools_e2e.rs +++ b/memoria/crates/memoria-mcp/tests/core_tools_e2e.rs @@ -123,6 +123,52 @@ async fn test_store_rejects_invalid_trust_tier() { println!("✅ invalid trust_tier rejected explicitly"); } +// ── 2b. memory_store: reject missing/blank content ─────────────────────────── + +#[tokio::test] +async fn test_store_rejects_empty_content() { + let (svc, uid, _ctx) = setup().await; + for args in [json!({}), json!({"content": ""}), json!({"content": " \t"})] { + let r = call("memory_store", args, &svc, &uid).await; + assert!( + text(&r).contains("content is required"), + "unexpected response: {}", + text(&r) + ); + } + assert!( + svc.list_active(&uid, 10).await.unwrap().is_empty(), + "empty content must not create memories" + ); + println!("✅ store rejects empty content"); +} + +#[tokio::test] +async fn test_embedded_store_metadata_round_trip() { + let (svc, uid, _ctx) = setup().await; + let metadata = json!({"scene": "embedded", "agent": "mcp"}); + let result = call( + "memory_store", + json!({"content": "embedded metadata memory", "extra_metadata": metadata}), + &svc, + &uid, + ) + .await; + let memory_id = text(&result) + .split_whitespace() + .nth(2) + .unwrap() + .trim_end_matches(':'); + let stored = svc.get_for_user(&uid, memory_id).await.unwrap().unwrap(); + assert_eq!( + serde_json::to_value(stored.extra_metadata).unwrap(), + metadata, + "embedded MCP must persist metadata" + ); + let listed = call("memory_list", json!({}), &svc, &uid).await; + assert!(text(&listed).contains("\"scene\":\"embedded\"")); +} + // ── 3. memory_retrieve: returns relevant memories ──────────────────────────── #[tokio::test] @@ -170,6 +216,30 @@ async fn test_retrieve_empty() { println!("✅ retrieve empty: {}", text(&r)); } +// ── 3b. memory_retrieve/search: reject missing/blank query ─────────────────── + +#[tokio::test] +async fn test_retrieve_and_search_reject_missing_query() { + let (svc, uid, _ctx) = setup().await; + + for (tool, args) in [ + ("memory_retrieve", json!({})), + ("memory_retrieve", json!({"query": ""})), + ("memory_retrieve", json!({"query": " \t"})), + ("memory_search", json!({})), + ("memory_search", json!({"query": ""})), + ("memory_search", json!({"query": " \t"})), + ] { + let r = call(tool, args, &svc, &uid).await; + assert!( + text(&r).contains("query is required"), + "{tool} unexpected response: {}", + text(&r) + ); + } + println!("✅ retrieve/search reject missing query"); +} + #[tokio::test] async fn test_retrieve_session_scope_only() { let (svc, uid, _ctx) = setup().await; diff --git a/memoria/crates/memoria-mcp/tests/edit_log_e2e.rs b/memoria/crates/memoria-mcp/tests/edit_log_e2e.rs index a209887..d7f4931 100644 --- a/memoria/crates/memoria-mcp/tests/edit_log_e2e.rs +++ b/memoria/crates/memoria-mcp/tests/edit_log_e2e.rs @@ -491,6 +491,7 @@ async fn test_store_batch_all_fields() { None, None, None, + None, ), ( "batch beta".to_string(), @@ -498,9 +499,13 @@ async fn test_store_batch_all_fields() { None, None, None, + None, ), ]; - let results = svc.store_batch(&uid, items, None).await.unwrap(); + let results = svc + .store_batch_with_metadata_on_branch(&uid, None, items, None) + .await + .unwrap(); assert_eq!(results.len(), 2); svc.flush_edit_log().await; @@ -1058,8 +1063,9 @@ async fn test_edit_id_globally_unique() { ) .await; // store_batch - svc.store_batch( + svc.store_batch_with_metadata_on_branch( &uid, + None, vec![ ( "batch1".into(), @@ -1067,6 +1073,7 @@ async fn test_edit_id_globally_unique() { None, None, None, + None, ), ( "batch2".into(), @@ -1074,6 +1081,7 @@ async fn test_edit_id_globally_unique() { None, None, None, + None, ), ], None, diff --git a/memoria/crates/memoria-mcp/tests/tools_unit.rs b/memoria/crates/memoria-mcp/tests/tools_unit.rs index b0123c8..1ed12db 100644 --- a/memoria/crates/memoria-mcp/tests/tools_unit.rs +++ b/memoria/crates/memoria-mcp/tests/tools_unit.rs @@ -156,6 +156,26 @@ async fn test_tool_memory_store() { println!("✅ tool memory_store: {text}"); } +#[tokio::test] +async fn test_tool_memory_store_rejects_empty_content() { + let svc = make_service(); + for args in [json!({}), json!({"content": ""}), json!({"content": " "})] { + let result = memoria_mcp::tools::call("memory_store", args, &svc, "u1") + .await + .unwrap(); + let text = result["content"][0]["text"].as_str().unwrap(); + assert!( + text.contains("content is required"), + "unexpected response: {text}" + ); + } + assert!( + svc.list_active("u1", 10).await.unwrap().is_empty(), + "empty content must not create memories" + ); + println!("✅ tool memory_store rejects empty content"); +} + #[tokio::test] async fn test_tool_memory_retrieve_empty() { let svc = make_service(); @@ -172,6 +192,30 @@ async fn test_tool_memory_retrieve_empty() { println!("✅ tool memory_retrieve empty: {text}"); } +#[tokio::test] +async fn test_tool_memory_retrieve_rejects_missing_query() { + let svc = make_service(); + + for (tool, args) in [ + ("memory_retrieve", json!({})), + ("memory_retrieve", json!({"query": ""})), + ("memory_retrieve", json!({"query": " "})), + ("memory_search", json!({})), + ("memory_search", json!({"query": ""})), + ("memory_search", json!({"query": " "})), + ] { + let result = memoria_mcp::tools::call(tool, args, &svc, "u1") + .await + .unwrap(); + let text = result["content"][0]["text"].as_str().unwrap(); + assert!( + text.contains("query is required"), + "{tool} unexpected response: {text}" + ); + } + println!("✅ tool memory_retrieve/search reject missing query"); +} + #[tokio::test] async fn test_tool_memory_retrieve_finds() { let svc = make_service(); diff --git a/memoria/crates/memoria-service/src/service.rs b/memoria/crates/memoria-service/src/service.rs index 580d5cd..458c712 100644 --- a/memoria/crates/memoria-service/src/service.rs +++ b/memoria/crates/memoria-service/src/service.rs @@ -496,6 +496,9 @@ impl EditLogBuffer { /// Single item for batch memory storage. /// Tuple fields: `(content, memory_type, session_id, trust_tier, subject_id)`. +/// +/// This alias is part of the public Rust API. Keep it stable; callers that need +/// metadata should use [`MetadataBatchStoreItem`] with `store_batch_with_metadata`. pub type BatchStoreItem = ( String, MemoryType, @@ -504,6 +507,17 @@ pub type BatchStoreItem = ( Option, ); +/// Metadata-aware item for batch memory storage. +/// Tuple fields: `(content, memory_type, session_id, trust_tier, subject_id, extra_metadata)`. +pub type MetadataBatchStoreItem = ( + String, + MemoryType, + Option, + Option, + Option, + Option>, +); + pub struct MemoryService { /// Trait-based store for generic ops (used by tests with MockStore) pub store: Arc, @@ -1233,6 +1247,39 @@ impl MemoryService { initial_confidence: Option, author_id: Option, subject_id: Option, + ) -> Result { + self.store_memory_with_metadata_on_branch( + user_id, + branch, + content, + memory_type, + session_id, + trust_tier, + observed_at, + initial_confidence, + author_id, + subject_id, + None, + ) + .await + } + + /// Store a memory on an optional branch, including caller-owned metadata. + #[allow(clippy::too_many_arguments)] + #[tracing::instrument(skip(self, content), fields(user_id, branch))] + pub async fn store_memory_with_metadata_on_branch( + &self, + user_id: &str, + branch: Option<&str>, + content: &str, + memory_type: MemoryType, + session_id: Option, + trust_tier: Option, + observed_at: Option>, + initial_confidence: Option, + author_id: Option, + subject_id: Option, + extra_metadata: Option>, ) -> Result { let subject_id = normalize_opt_string(subject_id); if let Some(ref sid) = subject_id { @@ -1253,6 +1300,9 @@ impl MemoryService { } let content = sensitivity.redacted_content.as_deref().unwrap_or(content); + // 保留 extra_metadata 的三态语义(None=不更新 / 非空=替换 / {}=清空)。不在入口把 {} + // 归一为 None——否则去重同内容路径无法用 {} 清空旧 metadata。响应/读取一致性改由 + // MemoryResponse 与读取侧的「"{}" → None」约定统一处理。 let effective_tier = trust_tier.unwrap_or(TrustTier::T1Verified); let embedding = self.embed(content).await?; let t_embed = t0.elapsed(); @@ -1274,7 +1324,7 @@ impl MemoryService { observed_at: Some(observed_at.unwrap_or_else(Utc::now)), created_at: None, updated_at: None, - extra_metadata: None, + extra_metadata, trust_tier: effective_tier, retrieval_score: None, }; @@ -1347,7 +1397,15 @@ impl MemoryService { }; return Ok(memory); } - // Same content — skip storing duplicate + // Same content — a near-duplicate already exists. Do NOT create/return a + // phantom, never-inserted record. Refresh the survivor's extra_metadata with + // the caller's new metadata (so an updated scene/agent isn't silently dropped), + // write an audit edit-log entry, then return the ACTUAL persisted record + // (real id + real author/session/trust/timestamps), not the new in-memory object. + if let Some(ref meta) = memory.extra_metadata { + sql.update_extra_metadata(&table, &old_id, meta).await?; + } + let existing = sql.get_from(&table, &old_id).await?; if t0.elapsed().as_secs() >= 1 { tracing::warn!( embed_ms = t_embed.as_millis() as u64, @@ -1356,7 +1414,53 @@ impl MemoryService { "store_memory slow (skip dup)" ); }; - return Ok(memory); + match existing { + Some(m) => { + // 只有确认幸存者仍 active 且是本次返回的记录时,才记 metadata 刷新审计。 + // 若竞态下幸存者已失效(update 命中 0 行、get_from 返回 None),会走下方 + // race-insert,此时记 update_metadata 会误导(那次更新并未生效/未返回)。 + if memory.extra_metadata.is_some() { + // 用**实际读到的** m.extra_metadata(而非请求 meta)记审计:并发下 + // 若本请求的更新被他人覆盖,get_from 读到的才是返回值,审计与响应保持一致。 + // (完整的“每请求返回自己的更新”需 CAS/事务,属更大改动,此处先对齐审计与响应。) + let payload = serde_json::json!({"memory_id": &old_id, "extra_metadata": m.extra_metadata}).to_string(); + self.send_edit_log( + user_id, + "update_metadata", + Some(&old_id), + Some(&payload), + "store_memory:dedup_metadata_refresh", + None, + ); + } + return Ok(m); + } + None => { + // Race: the survivor was deactivated between the dedup check and the + // fetch, so it is no longer a duplicate. Insert the new memory normally + // and return the real, persisted record (never a phantom id). + sql.insert_into(&table, &memory).await?; + let payload = serde_json::json!({"content": &memory.content, "type": memory.memory_type.to_string()}).to_string(); + self.send_edit_log( + user_id, + "inject", + Some(&memory.memory_id), + Some(&payload), + "store_memory:dedup_race_insert", + None, + ); + // 与正常插入路径一致的插入后副作用:统计事件 + 实体抽取入队, + // 否则该记忆不进实体图、活跃记忆指标偏低。 + self.report(StatsEvent::MemoryStored { + user_id: user_id.to_string(), + memory_type: memory.memory_type.to_string(), + trust_tier: memory.trust_tier.to_string(), + }); + self.enqueue_entity_extraction(user_id, &memory.memory_id, &memory.content) + .await; + return Ok(memory); + } + } } let t_dedup = t1.elapsed(); let t2 = std::time::Instant::now(); @@ -2531,6 +2635,17 @@ impl MemoryService { .await } + /// Batch store with caller-owned metadata. + pub async fn store_batch_with_metadata( + &self, + user_id: &str, + items: Vec, + author_id: Option, + ) -> Result, MemoriaError> { + self.store_batch_with_metadata_on_branch(user_id, None, items, author_id) + .await + } + /// Batch store with a single embedding API call, targeting an optional branch. /// Item tuple: (content, memory_type, session_id, trust_tier, subject_id) pub async fn store_batch_on_branch( @@ -2539,6 +2654,29 @@ impl MemoryService { branch: Option<&str>, items: Vec, author_id: Option, + ) -> Result, MemoriaError> { + self.store_batch_with_metadata_on_branch( + user_id, + branch, + items + .into_iter() + .map(|(content, mt, session_id, tier, subject_id)| { + (content, mt, session_id, tier, subject_id, None) + }) + .collect(), + author_id, + ) + .await + } + + /// Batch store with a single embedding API call and caller-owned metadata, + /// targeting an optional branch. + pub async fn store_batch_with_metadata_on_branch( + &self, + user_id: &str, + branch: Option<&str>, + items: Vec, + author_id: Option, ) -> Result, MemoriaError> { if items.is_empty() { return Ok(vec![]); @@ -2547,7 +2685,7 @@ impl MemoryService { // Sensitivity check + collect contents let mut contents = Vec::with_capacity(items.len()); let mut checked_items = Vec::with_capacity(items.len()); - for (content, mt, session_id, tier, subject_id) in items { + for (content, mt, session_id, tier, subject_id, extra_metadata) in items { let subject_id = normalize_opt_string(subject_id); if let Some(ref sid) = subject_id { if sid.len() > 128 { @@ -2565,14 +2703,21 @@ impl MemoryService { } let final_content = sensitivity.redacted_content.unwrap_or(content); contents.push(final_content.clone()); - checked_items.push((final_content, mt, session_id, tier, subject_id)); + checked_items.push(( + final_content, + mt, + session_id, + tier, + subject_id, + extra_metadata, + )); } // Batch embed let embeddings = self.embed_batch(&contents).await?; let mut results = Vec::with_capacity(checked_items.len()); - for (i, (content, mt, session_id, tier, subject_id)) in + for (i, (content, mt, session_id, tier, subject_id, extra_metadata)) in checked_items.into_iter().enumerate() { let effective_tier = tier.unwrap_or(TrustTier::T1Verified); @@ -2594,7 +2739,7 @@ impl MemoryService { observed_at: Some(Utc::now()), created_at: None, updated_at: None, - extra_metadata: None, + extra_metadata, trust_tier: effective_tier, retrieval_score: None, }; diff --git a/memoria/crates/memoria-service/tests/subject_id_mo_e2e.rs b/memoria/crates/memoria-service/tests/subject_id_mo_e2e.rs index ce30df8..3f887c8 100644 --- a/memoria/crates/memoria-service/tests/subject_id_mo_e2e.rs +++ b/memoria/crates/memoria-service/tests/subject_id_mo_e2e.rs @@ -321,7 +321,7 @@ async fn test_service_batch_store_subject_ids() { let subject_b = "batch-bob"; let stored = svc - .store_batch_on_branch( + .store_batch_with_metadata_on_branch( &uid, None, vec![ @@ -331,6 +331,7 @@ async fn test_service_batch_store_subject_ids() { None, None, Some(subject_a.to_string()), + None, ), ( "batch bob profile".to_string(), @@ -338,6 +339,7 @@ async fn test_service_batch_store_subject_ids() { None, None, Some(subject_b.to_string()), + None, ), ], None, diff --git a/memoria/crates/memoria-storage/src/store.rs b/memoria/crates/memoria-storage/src/store.rs index f0f3ab9..2c99fd2 100644 --- a/memoria/crates/memoria-storage/src/store.rs +++ b/memoria/crates/memoria-storage/src/store.rs @@ -5054,6 +5054,32 @@ impl SqlMemoryStore { Ok(()) } + /// Refresh only the extra_metadata JSON of an existing memory. Used by the single-store + /// dedup path: when a same-content near-duplicate already exists, we update the survivor's + /// metadata (so the caller's newer scene/agent isn't silently dropped) instead of creating + /// a phantom, never-inserted record. + pub async fn update_extra_metadata( + &self, + table: &str, + memory_id: &str, + extra_metadata: &std::collections::HashMap, + ) -> Result<(), MemoriaError> { + let table = self.t(table); + let json = serde_json::to_string(extra_metadata)?; + // is_active = 1:与读取侧(get_from)「幸存者必须 active」契约对齐——若该记忆在去重 + // 检查与本次更新之间被竞态置为 inactive,则不改动已失效记录(UPDATE 命中 0 行, + // 上层 get_from 返回 None 后走 race-insert)。 + sqlx::query(&format!( + "UPDATE {table} SET extra_metadata = ?, updated_at = NOW() WHERE memory_id = ? AND is_active = 1" + )) + .bind(json) + .bind(memory_id) + .execute(&self.pool) + .await + .map_err(db_err)?; + Ok(()) + } + #[tracing::instrument(skip(self, memory), fields(memory_id = %memory.memory_id))] pub async fn insert_into(&self, table: &str, memory: &Memory) -> Result<(), MemoriaError> { let now = Utc::now().naive_utc(); @@ -5200,6 +5226,9 @@ impl SqlMemoryStore { /// Lightweight list for API responses — skips embedding, source_event_ids, /// extra_metadata to reduce I/O and deserialization cost. #[allow(clippy::too_many_arguments)] + /// "Lite" list: skips the heavy `embedding` and `source_event_ids` columns for performance, + /// but DOES select `extra_metadata` (small JSON — needed by callers for scene/agent display), + /// mapped via `row_to_memory_lite`. Do not assume extra_metadata is omitted here. pub async fn list_active_lite( &self, table: &str, @@ -5237,7 +5266,8 @@ impl SqlMemoryStore { let sql = format!( "SELECT memory_id, user_id, author_id, subject_id, memory_type, content, \ session_id, is_active, superseded_by, trust_tier, \ - initial_confidence, observed_at, created_at, updated_at \ + initial_confidence, observed_at, created_at, updated_at, \ + CAST(extra_metadata AS CHAR) AS extra_meta \ FROM {table} WHERE memory_id IN ({inner}) \ ORDER BY memory_id DESC" ); @@ -6326,7 +6356,17 @@ fn row_to_memory(row: &sqlx::mysql::MySqlRow) -> Result { Ok(m) } -/// Lightweight row mapper — skips embedding, source_event_ids, extra_metadata. +/// Lightweight row mapper — skips embedding and source_event_ids, but DOES read +/// extra_metadata (small JSON; needed by list callers for scene/agent display). +/// Requires the query to SELECT `CAST(extra_metadata AS CHAR) AS extra_meta`. fn row_to_memory_lite(row: &sqlx::mysql::MySqlRow) -> Result { - row_to_memory_base(row) + let mut m = row_to_memory_base(row)?; + m.extra_metadata = { + let s: Option = row.try_get("extra_meta").map_err(db_err)?; + // MO#23859: we store "{}" instead of NULL; treat empty object as None. + s.filter(|v| v != "{}") + .map(|v| serde_json::from_str(&v)) + .transpose()? + }; + Ok(m) } diff --git a/memoria/crates/memoria-storage/tests/store_crud.rs b/memoria/crates/memoria-storage/tests/store_crud.rs index c7fbb44..a09fa56 100644 --- a/memoria/crates/memoria-storage/tests/store_crud.rs +++ b/memoria/crates/memoria-storage/tests/store_crud.rs @@ -720,11 +720,17 @@ async fn test_list_active_lite() { let (store, uid) = setup().await; // Insert 3 memories with embeddings for i in 0..3 { - let m = make_memory( + let mut m = make_memory( &format!("lite-{i}-{uid}"), &format!("lite memory {i}"), &uid, ); + if i == 0 { + m.extra_metadata = Some(std::collections::HashMap::from([( + "mapper".to_string(), + serde_json::json!("lite"), + )])); + } store.insert(&m).await.expect("insert"); } store @@ -744,12 +750,20 @@ async fn test_list_active_lite() { m.source_event_ids.is_empty(), "lite should skip source_event_ids" ); - assert!( - m.extra_metadata.is_none(), - "lite should skip extra_metadata" - ); assert!(!m.content.is_empty(), "content must be present"); } + let metadata_memory = results + .iter() + .find(|m| m.memory_id == format!("lite-0-{uid}")) + .expect("metadata memory"); + assert_eq!( + metadata_memory + .extra_metadata + .as_ref() + .and_then(|metadata| metadata.get("mapper")), + Some(&serde_json::json!("lite")), + "lite mapper must preserve extra_metadata" + ); // Verify ordering: newest first assert!(results[0].created_at >= results[1].created_at); println!(