Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -48,3 +48,4 @@ nohup.out
sdk/python/.venv/
sdk/python/uv.lock
.DS_Store
memoria/.DS_Store
11 changes: 11 additions & 0 deletions memoria/crates/memoria-api/src/models.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,10 @@ pub struct StoreRequest {
pub observed_at: Option<String>,
pub source: Option<String>,
pub branch: Option<String>,
/// 任意业务元数据(如 scene/agent)。透传落库到 memories.extra_metadata,并在读取时原样
/// 返回给调用方;Memoria 本身不对其做检索/打分逻辑(下游消费者如 matrixflow 的 decay 可自行使用)。
#[serde(default)]
pub extra_metadata: Option<std::collections::HashMap<String, serde_json::Value>>,
}
fn default_memory_type() -> String {
"semantic".to_string()
Expand Down Expand Up @@ -281,6 +285,10 @@ pub struct MemoryResponse {
pub observed_at: Option<String>,
pub created_at: Option<String>,
pub retrieval_score: Option<f64>,
/// 业务元数据(如 scene/agent)从 memories.extra_metadata 原样透传回给调用方;Memoria 本身
/// 不对其做检索/打分逻辑(下游消费者如 matrixflow 的 decay 可自行使用)。
#[serde(skip_serializing_if = "Option::is_none")]
pub extra_metadata: Option<std::collections::HashMap<String, serde_json::Value>>,
}

impl From<Memory> for MemoryResponse {
Expand All @@ -299,6 +307,9 @@ impl From<Memory> 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()),
}
}
}
Expand Down
23 changes: 16 additions & 7 deletions memoria/crates/memoria-api/src/routes/memory.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -231,6 +231,7 @@ pub async fn store_memory(
req.initial_confidence,
author,
req.subject_id,
req.extra_metadata,
)
.await
.map_err(|e| {
Expand Down Expand Up @@ -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::<Result<Vec<_>, _>>()?;

// 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() {
Expand All @@ -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,
Expand Down
272 changes: 271 additions & 1 deletion memoria/crates/memoria-api/tests/api_e2e.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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::<Value>()
.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]
Expand Down Expand Up @@ -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");
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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::<Value>().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) =
Expand Down
Loading
Loading