diff --git a/bin/memayu/src/cli.rs b/bin/memayu/src/cli.rs index 7bc2251..305dd8b 100644 --- a/bin/memayu/src/cli.rs +++ b/bin/memayu/src/cli.rs @@ -6,11 +6,14 @@ //! single-user by design: the user id comes from the config (default `default`) //! and no auth is required for a local instance. +use crate::service::build_service; use memayu_config::Config; +use memayu_core::{ + BatchItemResult, BatchMemory, Memory, MemoryService, MetadataFilter, MAX_PAGE_SIZE, +}; use std::collections::{HashMap, HashSet}; - -use crate::service::build_service; -use memayu_core::{Memory, MetadataFilter, MAX_PAGE_SIZE}; +use std::fs::File; +use std::io::{self, BufRead, BufReader}; /// Default user id for the local single-user CLI. pub const DEFAULT_USER_ID: &str = "default"; @@ -53,7 +56,7 @@ fn parse_args(args: impl Iterator) -> Parsed { if let Some(rest) = a.strip_prefix("--") { if let Some((k, v)) = rest.split_once('=') { insert_flag(&mut p, k, v.to_string()); - } else if matches!(rest, "limit" | "cursor" | "filter") { + } else if matches!(rest, "limit" | "cursor" | "filter" | "batch") { let v = it.next().unwrap_or_default(); insert_flag(&mut p, rest, v); } else { @@ -158,13 +161,21 @@ fn reject_unknown( /// `memayu add ""` — run ADD/UPDATE extraction and store the memory. pub async fn cmd_add(config: &Config, args: impl Iterator) -> Result<(), String> { let p = parse_args(args); - reject_unknown(&p, &[], &[])?; + reject_unknown(&p, &["batch"], &[])?; + guard_local(config)?; + let (service, _) = build_service(config).await.map_err(|e| e.to_string())?; + + if let Some(path) = p.opts.get("batch") { + if !p.positionals.is_empty() { + return Err("--batch cannot be combined with positional content".to_string()); + } + return cmd_add_batch(&service, path).await; + } + let content = p.positionals.join(" ").trim().to_string(); if content.is_empty() { - return Err("usage: memayu add \"\"".to_string()); + return Err("usage: memayu add \"\" | --batch ".to_string()); } - guard_local(config)?; - let (service, _) = build_service(config).await.map_err(|e| e.to_string())?; let mem = service .add_memory(DEFAULT_USER_ID, &content, &Default::default()) .await @@ -173,6 +184,79 @@ pub async fn cmd_add(config: &Config, args: impl Iterator) -> Res Ok(()) } +/// `memayu add --batch ` — store many memories at once. +/// Each line is a JSON object with a `content` field (and optional `metadata`). +/// A `-` path reads from stdin. Per-item failures don't abort the batch. +async fn cmd_add_batch(service: &MemoryService, path: &str) -> Result<(), String> { + let reader: Box = if path == "-" { + Box::new(BufReader::new(io::stdin())) + } else { + let file = File::open(path).map_err(|e| format!("cannot open {path}: {e}"))?; + Box::new(BufReader::new(file)) + }; + + let mut items = Vec::new(); + for (idx, line) in reader.lines().enumerate() { + let line_no = idx + 1; + let line = line + .map_err(|e| format!("read {path} line {line_no}: {e}"))? + .trim() + .to_string(); + if line.is_empty() { + continue; + } + let value: serde_json::Value = serde_json::from_str(&line) + .map_err(|e| format!("{path}:{line_no}: invalid JSON: {e}"))?; + let content = value + .get("content") + .and_then(|c| c.as_str()) + .unwrap_or("") + .to_string(); + let metadata = value + .get("metadata") + .and_then(|m| m.as_object()) + .map(|obj| { + obj.iter() + .filter_map(|(k, val)| val.as_str().map(|s| (k.clone(), s.to_string()))) + .collect() + }) + .unwrap_or_default(); + items.push(BatchMemory { content, metadata }); + } + + if items.is_empty() { + return Err(format!( + "{path}: no memories found (file is empty or all blank)" + )); + } + + let results = service + .add_memories_batch(DEFAULT_USER_ID, &items) + .await + .map_err(|e| e.to_string())?; + + let mut added = 0usize; + let mut failed = 0usize; + for result in &results { + match result { + BatchItemResult::Stored { memory_id } => { + added += 1; + println!("stored: {memory_id}"); + } + BatchItemResult::Failed { error } => { + failed += 1; + eprintln!("failed: {error}"); + } + } + } + println!("added: {added} memory(ies)"); + if failed > 0 { + Err(format!("{failed} item(s) failed")) + } else { + Ok(()) + } +} + /// `memayu search "" [--limit N] [--filter key=value] [--json]` — /// ranked semantic results. pub async fn cmd_search(config: &Config, args: impl Iterator) -> Result<(), String> { diff --git a/bin/memayu/tests/cli.rs b/bin/memayu/tests/cli.rs index 0c2e39e..5e4a9fa 100644 --- a/bin/memayu/tests/cli.rs +++ b/bin/memayu/tests/cli.rs @@ -50,6 +50,73 @@ fn start_mock_embedder() -> u16 { port } +/// Like [`start_mock_embedder`] but derives a distinct embedding from each +/// request's full body. Used by batch tests where every item must get its own +/// vector so the storage layer doesn't treat them as duplicates and merge them. +fn start_distinct_embedder() -> u16 { + use std::hash::{Hash, Hasher}; + let listener = TcpListener::bind("127.0.0.1:0").expect("bind distinct embedder"); + let port = listener.local_addr().unwrap().port(); + std::thread::spawn(move || { + for stream in listener.incoming() { + let Ok(mut stream) = stream else { break }; + // Read the full request: headers plus Content-Length bytes of body. + let mut buf = Vec::new(); + let mut tmp = [0u8; 2048]; + loop { + match stream.read(&mut tmp) { + Ok(0) => break, + Ok(n) => { + buf.extend_from_slice(&tmp[..n]); + if let Some(header_end) = buf.windows(4).position(|w| w == b"\r\n\r\n") { + let header = String::from_utf8_lossy(&buf[..header_end]).into_owned(); + let content_length = header + .lines() + .find_map(|line| { + let mut parts = line.splitn(2, ':'); + if parts + .next() + .map(|k| k.trim().eq_ignore_ascii_case("content-length")) + == Some(true) + { + parts.next()?.trim().parse::().ok() + } else { + None + } + }) + .unwrap_or(0); + if buf.len() >= header_end + 4 + content_length { + break; + } + } + } + Err(_) => break, + } + } + let mut h = std::collections::hash_map::DefaultHasher::new(); + buf.hash(&mut h); + let v = h.finish(); + let emb = [ + ((v & 0xffff) as f32) / 65535.0, + (((v >> 16) & 0xffff) as f32) / 65535.0, + (((v >> 32) & 0xffff) as f32) / 65535.0, + ]; + let body = format!( + r#"{{"data":[{{"embedding":[{:.6},{:.6},{:.6}]}}]}}"#, + emb[0], emb[1], emb[2] + ); + let resp = format!( + "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", + body.len(), + body + ); + let _ = stream.write_all(resp.as_bytes()); + let _ = stream.flush(); + } + }); + port +} + /// A temp libsql db path unique per test process. fn temp_db_path() -> std::path::PathBuf { static COUNTER: AtomicU16 = AtomicU16::new(0); @@ -229,3 +296,60 @@ fn list_exposes_paging_contract() { assert_eq!(obj["memories"].as_array().unwrap().len(), 0); assert_eq!(obj["total"].as_u64(), Some(0)); } + +#[test] +fn add_batch_stores_all_items() { + let port = start_distinct_embedder(); + let db = temp_db_path(); + let batch = temp_db_path().with_extension("batch.jsonl"); + std::fs::write( + &batch, + concat!( + "{\"content\":\"alpha\",\"metadata\":{\"source\":\"batch\"}}\n", + "{\"content\":\"beta\"}\n", + "{\"content\":\"gamma\"}\n", + ), + ) + .unwrap(); + + let (code, stdout) = run_memayu(&["add", "--batch", batch.to_str().unwrap()], port, &db); + assert_eq!(code, 0, "batch add should succeed"); + assert_eq!(stdout.matches("stored:").count(), 3); + assert!(stdout.contains("added: 3 memory(ies)"), "stdout: {stdout}"); + + // All three are persisted. + let (code, stdout) = run_memayu(&["list", "--json"], port, &db); + assert_eq!(code, 0); + let obj: serde_json::Value = serde_json::from_str(stdout.trim()).unwrap(); + let contents: Vec = obj["memories"] + .as_array() + .unwrap() + .iter() + .map(|m| m["content"].as_str().unwrap().to_string()) + .collect(); + for c in ["alpha", "beta", "gamma"] { + assert!( + contents.iter().any(|x| x == c), + "missing {c} in {contents:?}" + ); + } +} + +#[test] +fn add_batch_reports_partial_failures() { + let port = start_distinct_embedder(); + let db = temp_db_path(); + let batch = temp_db_path().with_extension("batch.jsonl"); + std::fs::write(&batch, "{\"content\":\"ok\"}\n{\"content\":\" \"}\n").unwrap(); + + let (code, stdout, stderr) = + run_memayu_full(&["add", "--batch", batch.to_str().unwrap()], port, &db); + // The valid item is stored, but the blank one fails and we exit non-zero. + assert_ne!(code, 0, "partial failure should exit non-zero"); + assert_eq!(stdout.matches("stored:").count(), 1); + assert!( + stderr.contains("failed: memory content is required"), + "stderr: {stderr}" + ); + assert!(stderr.contains("1 item(s) failed"), "stderr: {stderr}"); +} diff --git a/crates/memayu-api/src/error.rs b/crates/memayu-api/src/error.rs index e1bda2a..c263d23 100644 --- a/crates/memayu-api/src/error.rs +++ b/crates/memayu-api/src/error.rs @@ -54,6 +54,7 @@ impl From for ApiError { let status = match &e { CoreError::DimensionMismatch { .. } => 422, CoreError::InvalidExtraction(_) => 422, + CoreError::InvalidInput(_) => 400, CoreError::NotFound(_) => 404, CoreError::InvalidCursor(_) => 400, CoreError::LimitExceeded { .. } => 400, diff --git a/crates/memayu-api/src/modules/memory/dto.rs b/crates/memayu-api/src/modules/memory/dto.rs index e097ddc..8a27f0a 100644 --- a/crates/memayu-api/src/modules/memory/dto.rs +++ b/crates/memayu-api/src/modules/memory/dto.rs @@ -10,6 +10,39 @@ pub struct AddMemoryRequest { pub metadata: Metadata, } +/// One memory within a batch add request. +#[derive(Debug, Deserialize, ToSchema)] +pub struct AddMemoryItem { + pub content: String, + #[serde(default)] + pub metadata: Metadata, +} + +/// Batch add request — `{"memories": [{"content": ..., "metadata": ...}, ...]}`. +#[derive(Debug, Deserialize, ToSchema)] +pub struct AddMemoriesBatchRequest { + pub memories: Vec, +} + +/// A single failed item within a batch response. +#[derive(Debug, Serialize, ToSchema)] +pub struct BatchMemoryError { + /// Zero-based index of the failed item in the request's `memories` array. + pub index: usize, + pub error: String, +} + +#[derive(Debug, Serialize, ToSchema)] +pub struct AddMemoriesBatchResponse { + pub status: String, + /// Number of memories actually stored (successes). + pub added: usize, + /// Memory ids of the successfully stored items, in request order. + pub memory_ids: Vec, + /// Per-item failures. Empty when every item was stored. + pub errors: Vec, +} + #[derive(Debug, Serialize, ToSchema)] pub struct AddMemoryResponse { pub status: String, diff --git a/crates/memayu-api/src/transport/handlers/memory.rs b/crates/memayu-api/src/transport/handlers/memory.rs index 487c9e7..a2607fe 100644 --- a/crates/memayu-api/src/transport/handlers/memory.rs +++ b/crates/memayu-api/src/transport/handlers/memory.rs @@ -2,15 +2,15 @@ /// the core MemoryService. use crate::error::{ApiError, ApiErrorBody, ApiResult}; use crate::modules::memory::dto::{ - AddMemoryRequest, AddMemoryResponse, ListMemoryResponse, ListQuery, ListedMemory, - SearchMemoryRequest, SearchMemoryResponse, SearchResult, UpdateMemoryRequest, - UpdateMemoryResponse, + AddMemoriesBatchRequest, AddMemoriesBatchResponse, AddMemoryRequest, AddMemoryResponse, + BatchMemoryError, ListMemoryResponse, ListQuery, ListedMemory, SearchMemoryRequest, + SearchMemoryResponse, SearchResult, UpdateMemoryRequest, UpdateMemoryResponse, }; use crate::transport::middleware::{AccountId, ApiState}; use axum::extract::{Path, Query, State}; use axum::http::StatusCode; use axum::Json; -use memayu_core::MetadataFilter; +use memayu_core::{BatchItemResult, BatchMemory, MetadataFilter}; /// Add a new memory for a user #[utoipa::path( @@ -49,6 +49,62 @@ pub async fn add_memory( )) } +/// Add many memories for a user in a single call. A failure on one item does +/// not fail the rest; successes are reported via `added`/`memory_ids` and +/// failures via `errors`. +#[utoipa::path( + post, + path = "/api/memories/batch", + request_body = AddMemoriesBatchRequest, + responses( + (status = 200, description = "Memories added", body = ApiResult), + (status = 400, description = "Bad request", body = ApiErrorBody), + ) +)] +pub async fn add_memories_batch( + State(state): State, + _account: AccountId, + Json(req): Json, +) -> Result>, ApiError> { + let user_id = &_account.0; + if req.memories.is_empty() { + return Err(ApiError::bad_request("memories must not be empty")); + } + + let items: Vec = req + .memories + .into_iter() + .map(|m| BatchMemory { + content: m.content, + metadata: m.metadata, + }) + .collect(); + + let results = state.service.add_memories_batch(user_id, &items).await?; + + let mut added = 0usize; + let mut memory_ids = Vec::new(); + let mut errors = Vec::new(); + for (index, outcome) in results.into_iter().enumerate() { + match outcome { + BatchItemResult::Stored { memory_id } => { + added += 1; + memory_ids.push(memory_id); + } + BatchItemResult::Failed { error } => errors.push(BatchMemoryError { index, error }), + } + } + + Ok(Json(ApiResult { + result: AddMemoriesBatchResponse { + status: "success".into(), + added, + memory_ids, + errors, + }, + })) +} + /// Search memories by semantic similarity #[utoipa::path( post, diff --git a/crates/memayu-api/src/transport/routes.rs b/crates/memayu-api/src/transport/routes.rs index e739895..55dcd84 100644 --- a/crates/memayu-api/src/transport/routes.rs +++ b/crates/memayu-api/src/transport/routes.rs @@ -15,6 +15,7 @@ use utoipa_axum::{router::OpenApiRouter, routes}; info(title = "Memayu API", version = "0.1.0"), paths( handlers::memory::add_memory, + handlers::memory::add_memories_batch, handlers::memory::search_memory, handlers::memory::list_memories, handlers::memory::delete_memory, @@ -23,6 +24,10 @@ use utoipa_axum::{router::OpenApiRouter, routes}; components(schemas( crate::modules::memory::dto::AddMemoryRequest, crate::modules::memory::dto::AddMemoryResponse, + crate::modules::memory::dto::AddMemoryItem, + crate::modules::memory::dto::AddMemoriesBatchRequest, + crate::modules::memory::dto::AddMemoriesBatchResponse, + crate::modules::memory::dto::BatchMemoryError, crate::modules::memory::dto::SearchMemoryRequest, crate::modules::memory::dto::SearchMemoryResponse, crate::modules::memory::dto::SearchResult, @@ -74,6 +79,7 @@ pub fn build( let (memory_routes, openapi_spec) = { let (router, api) = OpenApiRouter::with_openapi(ApiDoc::openapi()) .routes(routes!(handlers::memory::add_memory)) + .routes(routes!(handlers::memory::add_memories_batch)) .routes(routes!(handlers::memory::search_memory)) .routes(routes!(handlers::memory::list_memories)) .routes(routes!(handlers::memory::delete_memory)) diff --git a/crates/memayu-api/tests/api.rs b/crates/memayu-api/tests/api.rs index 9f282d5..fcb9e0c 100644 --- a/crates/memayu-api/tests/api.rs +++ b/crates/memayu-api/tests/api.rs @@ -248,6 +248,100 @@ mod tests { assert!(parsed["memory_id"].is_string()); } + #[tokio::test] + async fn add_memories_batch_stores_all() { + let (app, cookie) = setup_and_login().await; + let resp = app + .oneshot( + Request::builder() + .method("POST") + .uri("/api/memories/batch") + .header("content-type", "application/json") + .header("cookie", &cookie) + .body(Body::from( + r#"{ + "memories": [ + {"content":"loves hiking","metadata":{"source":"batch"}}, + {"content":"prefers coffee"}, + {"content":"works remote"} + ] + }"#, + )) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), usize::MAX) + .await + .unwrap(); + let parsed: serde_json::Value = serde_json::from_slice(&body).unwrap(); + let parsed = parsed["result"].clone(); + assert_eq!(parsed["status"], "success"); + assert_eq!(parsed["added"], 3); + assert_eq!(parsed["memory_ids"].as_array().unwrap().len(), 3); + assert!(parsed["errors"].as_array().unwrap().is_empty()); + } + + #[tokio::test] + async fn add_memories_batch_reports_per_item_errors() { + let (app, cookie) = setup_and_login().await; + let resp = app + .oneshot( + Request::builder() + .method("POST") + .uri("/api/memories/batch") + .header("content-type", "application/json") + .header("cookie", &cookie) + .body(Body::from( + r#"{ + "memories": [ + {"content":"valid one"}, + {"content":" "}, + {"content":"valid two"} + ] + }"#, + )) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), usize::MAX) + .await + .unwrap(); + let parsed: serde_json::Value = serde_json::from_slice(&body).unwrap(); + let parsed = parsed["result"].clone(); + assert_eq!(parsed["status"], "success"); + assert_eq!(parsed["added"], 2, "valid items must be stored"); + assert_eq!(parsed["memory_ids"].as_array().unwrap().len(), 2); + let errors = parsed["errors"].as_array().unwrap(); + assert_eq!(errors.len(), 1, "blank item must be reported as a failure"); + assert_eq!(errors[0]["index"], 1); + assert!(errors[0]["error"] + .as_str() + .unwrap() + .contains("content is required")); + } + + #[tokio::test] + async fn add_memories_batch_rejects_empty() { + let (app, cookie) = setup_and_login().await; + let resp = app + .oneshot( + Request::builder() + .method("POST") + .uri("/api/memories/batch") + .header("content-type", "application/json") + .header("cookie", &cookie) + .body(Body::from(r#"{"memories":[]}"#)) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::BAD_REQUEST); + } + #[tokio::test] async fn metadata_round_trips_in_search_and_list() { let (app, cookie) = setup_and_login().await; diff --git a/crates/memayu-core/src/error.rs b/crates/memayu-core/src/error.rs index 7d536f5..b208a92 100644 --- a/crates/memayu-core/src/error.rs +++ b/crates/memayu-core/src/error.rs @@ -14,6 +14,8 @@ pub enum CoreError { DimensionMismatch { expected: usize, got: usize }, #[error("extraction result invalid: {0}")] InvalidExtraction(String), + #[error("invalid input: {0}")] + InvalidInput(String), #[error("not found: {0}")] NotFound(String), #[error("limit {limit} exceeds the maximum of {max}")] diff --git a/crates/memayu-core/src/lib.rs b/crates/memayu-core/src/lib.rs index a70d305..d29a751 100644 --- a/crates/memayu-core/src/lib.rs +++ b/crates/memayu-core/src/lib.rs @@ -15,4 +15,4 @@ pub use ports::{ EmbedError, EmbedderProvider, ExtractionDecision, ExtractionMode, ExtractionResult, LlmError, LlmProvider, Message, Metadata, StorageError, StorageProvider, }; -pub use service::{AddMemoryOutcome, MemoryService}; +pub use service::{AddMemoryOutcome, BatchItemResult, BatchMemory, MemoryService}; diff --git a/crates/memayu-core/src/service.rs b/crates/memayu-core/src/service.rs index 3aba983..42c53e2 100644 --- a/crates/memayu-core/src/service.rs +++ b/crates/memayu-core/src/service.rs @@ -12,6 +12,22 @@ use uuid::Uuid; /// as a duplicate of an existing one. Above this value, ingestion is a NOOP. pub const RAW_DEDUPE_THRESHOLD: f32 = 0.98; +/// One memory to insert as part of a batch. +#[derive(Debug, Clone)] +pub struct BatchMemory { + pub content: String, + pub metadata: Metadata, +} + +/// Per-item outcome of [`MemoryService::add_memories_batch`]. A single failed +/// item does not abort the rest of the batch; failures are reported alongside +/// the stored ids. +#[derive(Debug)] +pub enum BatchItemResult { + Stored { memory_id: String }, + Failed { error: String }, +} + pub struct MemoryService { storage: Arc, embedder: Arc, @@ -221,6 +237,40 @@ impl MemoryService { } } + /// Insert many memories in a single call. Each item is ingested through the + /// same ADD/UPDATE pipeline as [`MemoryService::add_memory`], so per-item + /// failures (empty content, embed/LLM errors) are collected rather than + /// aborting the whole batch. The outer `Err` is reserved for request-level + /// errors, such as an empty batch. + pub async fn add_memories_batch( + &self, + user_id: &str, + items: &[BatchMemory], + ) -> Result, CoreError> { + if items.is_empty() { + return Err(CoreError::InvalidInput( + "batch must contain at least one memory".into(), + )); + } + let mut results = Vec::with_capacity(items.len()); + for item in items { + let content = item.content.trim(); + if content.is_empty() { + results.push(BatchItemResult::Failed { + error: "memory content is required".into(), + }); + continue; + } + match self.add_memory(user_id, content, &item.metadata).await { + Ok(mem) => results.push(BatchItemResult::Stored { memory_id: mem.id }), + Err(e) => results.push(BatchItemResult::Failed { + error: e.to_string(), + }), + } + } + Ok(results) + } + pub async fn search_memory( &self, user_id: &str, @@ -1074,4 +1124,100 @@ mod tests { assert_eq!(ids, vec!["a", "c"]); assert!(page.next_cursor.is_none()); } + + // ── issue #53: batch add ── + + #[tokio::test] + async fn add_memories_batch_stores_multiple_items() { + let storage = MockStorage { + rows: Mutex::new(vec![]), + score: 0.0, + }; + let svc = raw_service_with(storage); + + let items = vec![ + BatchMemory { + content: "loves hiking".into(), + metadata: HashMap::from([("source".into(), "batch".into())]), + }, + BatchMemory { + content: "prefers coffee".into(), + metadata: HashMap::new(), + }, + BatchMemory { + content: "works remote".into(), + metadata: HashMap::new(), + }, + ]; + + let results = svc.add_memories_batch("u1", &items).await.unwrap(); + + assert_eq!(results.len(), 3); + assert!(results + .iter() + .all(|r| matches!(r, BatchItemResult::Stored { .. }))); + // All three rows landed, with the batch metadata attached. + let rows = svc.list_memories("u1", 10).await.unwrap(); + assert_eq!(rows.len(), 3); + assert!(rows + .iter() + .any(|m| m.metadata.get("source") == Some(&"batch".to_string()))); + } + + #[tokio::test] + async fn add_memories_batch_collects_per_item_errors() { + let storage = MockStorage { + rows: Mutex::new(vec![]), + score: 0.0, + }; + let svc = raw_service_with(storage); + + let items = vec![ + BatchMemory { + content: "valid one".into(), + metadata: HashMap::new(), + }, + BatchMemory { + content: " ".into(), // blank -> per-item failure + metadata: HashMap::new(), + }, + BatchMemory { + content: "valid two".into(), + metadata: HashMap::new(), + }, + ]; + + let results = svc.add_memories_batch("u1", &items).await.unwrap(); + + assert_eq!(results.len(), 3); + let stored: Vec<&BatchItemResult> = results + .iter() + .filter(|r| matches!(r, BatchItemResult::Stored { .. })) + .collect(); + let failed: Vec<&BatchItemResult> = results + .iter() + .filter(|r| matches!(r, BatchItemResult::Failed { .. })) + .collect(); + assert_eq!(stored.len(), 2, "valid items must still be stored"); + assert_eq!(failed.len(), 1, "blank item must be reported as failed"); + match failed[0] { + BatchItemResult::Failed { error } => { + assert!(error.contains("content is required"), "error: {error}") + } + _ => unreachable!(), + } + assert_eq!(svc.list_memories("u1", 10).await.unwrap().len(), 2); + } + + #[tokio::test] + async fn add_memories_batch_rejects_empty_batch() { + let storage = MockStorage { + rows: Mutex::new(vec![]), + score: 0.0, + }; + let svc = raw_service_with(storage); + + let err = svc.add_memories_batch("u1", &[]).await.unwrap_err(); + assert!(matches!(err, CoreError::InvalidInput(_))); + } } diff --git a/crates/memayu-mcp/src/tools/add_memory.rs b/crates/memayu-mcp/src/tools/add_memory.rs index 5bdc428..41cc4be 100644 --- a/crates/memayu-mcp/src/tools/add_memory.rs +++ b/crates/memayu-mcp/src/tools/add_memory.rs @@ -1,4 +1,4 @@ -//! `add_memory` tool — store a new memory. +//! `add_memory` tool — store a new memory (single or batch). use crate::types::ToolDefinition; use crate::{McpError, MemoryBackend}; @@ -9,16 +9,33 @@ pub fn definition() -> ToolDefinition { ToolDefinition { name: "add_memory", description: - "Store a new memory or update existing ones. The system will deduplicate and merge similar memories.", + "Store new memories or update existing ones. The system will deduplicate and merge similar memories. Pass a single `content` string, or a `memories` array to store several in one call.", input_schema: serde_json::json!({ "type": "object", "properties": { "content": { "type": "string", - "description": "The memory content to store" + "description": "The memory content to store (single add)" + }, + "memories": { + "type": "array", + "description": "Batch of memories to store in one call", + "items": { + "type": "object", + "properties": { + "content": { + "type": "string", + "description": "The memory content to store" + }, + "metadata": { + "type": "object", + "description": "Optional key/value metadata" + } + }, + "required": ["content"] + } } - }, - "required": ["content"] + } }), } } @@ -27,16 +44,55 @@ pub async fn call( args: &HashMap, backend: &dyn MemoryBackend, ) -> Result { - let content = args - .get("content") - .and_then(|v| v.as_str()) - .ok_or_else(|| McpError::Api("Missing 'content' argument".into()))?; - let user_id = args .get("user_id") .and_then(|v| v.as_str()) .unwrap_or("default"); + // Batch path: `memories` is an array of {content, metadata}. A failure on + // one item does not abort the rest; successes and failures are summarized. + if let Some(items) = args.get("memories").and_then(|v| v.as_array()) { + let mut added = 0usize; + let mut ids = Vec::new(); + let mut errors = Vec::new(); + for (i, item) in items.iter().enumerate() { + let content = item + .get("content") + .and_then(|v| v.as_str()) + .unwrap_or("") + .trim(); + if content.is_empty() { + errors.push(format!("item {i}: content is required")); + continue; + } + match backend.add_memory(user_id, content).await { + Ok(mem) => { + added += 1; + ids.push(mem.id); + } + Err(e) => errors.push(format!("item {i}: {e}")), + } + } + let summary = if errors.is_empty() { + format!("Stored {added} memories (ids: {})", ids.join(", ")) + } else { + format!( + "Stored {added} memories (ids: {}); failures: {}", + ids.join(", "), + errors.join("; ") + ) + }; + return Ok(serde_json::json!({ + "content": [{ "type": "text", "text": summary }] + })); + } + + // Single path: `content` string (backward compatible). + let content = args + .get("content") + .and_then(|v| v.as_str()) + .ok_or_else(|| McpError::Api("Missing 'content' or 'memories' argument".into()))?; + let mem = backend.add_memory(user_id, content).await?; Ok(serde_json::json!({ "content": [{ @@ -45,3 +101,110 @@ pub async fn call( }] })) } + +#[cfg(test)] +mod tests { + use super::*; + use async_trait::async_trait; + use chrono::Utc; + use memayu_core::{Memory, MetadataFilter}; + + /// Backend that stores each memory verbatim with a deterministic id. + struct MockBackend; + + #[async_trait] + impl MemoryBackend for MockBackend { + async fn add_memory(&self, user_id: &str, content: &str) -> Result { + Ok(Memory { + id: format!("id-{content}"), + user_id: user_id.to_string(), + content: content.to_string(), + vector: vec![], + metadata: HashMap::new(), + created_at: Utc::now(), + updated_at: Utc::now(), + }) + } + async fn search_memory( + &self, + _user_id: &str, + _query: &str, + _limit: usize, + _metadata_filter: Option, + ) -> Result, McpError> { + Ok(vec![]) + } + async fn list_memories( + &self, + _user_id: &str, + _limit: usize, + ) -> Result, McpError> { + Ok(vec![]) + } + async fn delete_memory(&self, _memory_id: &str) -> Result<(), McpError> { + Ok(()) + } + async fn update_memory( + &self, + _memory_id: &str, + _content: &str, + ) -> Result { + unreachable!() + } + } + + fn args(map: serde_json::Map) -> HashMap { + HashMap::from_iter(map) + } + + #[tokio::test] + async fn single_content_is_backward_compatible() { + let args = args( + serde_json::json!({ "content": "hello" }) + .as_object() + .unwrap() + .clone(), + ); + let out = call(&args, &MockBackend).await.unwrap(); + let text = out["content"][0]["text"].as_str().unwrap(); + assert!(text.contains("id-hello"), "text: {text}"); + } + + #[tokio::test] + async fn batch_stores_all_items() { + let args = args( + serde_json::json!({ "memories": [ + {"content":"a"}, + {"content":"b"}, + {"content":"c"} + ] }) + .as_object() + .unwrap() + .clone(), + ); + let out = call(&args, &MockBackend).await.unwrap(); + let text = out["content"][0]["text"].as_str().unwrap(); + assert!(text.contains("Stored 3 memories"), "text: {text}"); + assert!(text.contains("id-a"), "text: {text}"); + assert!(text.contains("id-b"), "text: {text}"); + assert!(text.contains("id-c"), "text: {text}"); + } + + #[tokio::test] + async fn batch_reports_per_item_failures() { + let args = args( + serde_json::json!({ "memories": [ + {"content":"ok"}, + {"content":""}, + {"content":"ok2"} + ] }) + .as_object() + .unwrap() + .clone(), + ); + let out = call(&args, &MockBackend).await.unwrap(); + let text = out["content"][0]["text"].as_str().unwrap(); + assert!(text.contains("Stored 2 memories"), "text: {text}"); + assert!(text.contains("item 1: content is required"), "text: {text}"); + } +} diff --git a/crates/memayu-mcp/src/types.rs b/crates/memayu-mcp/src/types.rs index 382dde0..c4697cc 100644 --- a/crates/memayu-mcp/src/types.rs +++ b/crates/memayu-mcp/src/types.rs @@ -61,6 +61,7 @@ impl JsonRpcResponse { // ── MCP Initialize ── #[derive(Debug, Serialize)] +#[serde(rename_all = "camelCase")] pub struct InitializeResult { pub protocol_version: &'static str, pub server_info: ServerInfo, @@ -68,12 +69,14 @@ pub struct InitializeResult { } #[derive(Debug, Serialize)] +#[serde(rename_all = "camelCase")] pub struct ServerInfo { pub name: &'static str, pub version: &'static str, } #[derive(Debug, Serialize)] +#[serde(rename_all = "camelCase")] pub struct ServerCapabilities { pub tools: ToolsCapability, } @@ -117,3 +120,64 @@ pub struct ToolContent { pub content_type: &'static str, pub text: String, } + +#[cfg(test)] +mod tests { + use super::*; + + /// The MCP 2024-11-05 `InitializeResult` must serialize with camelCase keys + /// (`protocolVersion`, `serverInfo`) so strict clients (e.g. Hermes Agent's + /// pydantic-validated client) accept the handshake. + #[test] + fn initialize_result_uses_camel_case_keys() { + let result = InitializeResult { + protocol_version: "2024-11-05", + server_info: ServerInfo { + name: "memayu-mcp", + version: "0.1.0", + }, + capabilities: ServerCapabilities { + tools: ToolsCapability { + list_changed: false, + }, + }, + }; + + let value = serde_json::to_value(&result).unwrap(); + let obj = value.as_object().unwrap(); + + assert!( + obj.contains_key("protocolVersion"), + "missing protocolVersion: {value}" + ); + assert!( + obj.contains_key("serverInfo"), + "missing serverInfo: {value}" + ); + assert!( + obj.contains_key("capabilities"), + "missing capabilities: {value}" + ); + // snake_case variants must NOT leak into the wire format. + assert!( + !obj.contains_key("protocol_version"), + "snake_case protocol_version present" + ); + assert!( + !obj.contains_key("server_info"), + "snake_case server_info present" + ); + + let server_info = obj["serverInfo"].as_object().unwrap(); + assert_eq!(server_info["name"], "memayu-mcp"); + assert_eq!(server_info["version"], "0.1.0"); + + // `tools.listChanged` (nested) is also camelCase per spec. + let capabilities = obj["capabilities"].as_object().unwrap(); + let tools = capabilities["tools"].as_object().unwrap(); + assert!( + tools.contains_key("listChanged"), + "missing listChanged: {value}" + ); + } +}