diff --git a/memoria/crates/memoria-api/src/lib.rs b/memoria/crates/memoria-api/src/lib.rs index dadba4b..230ef93 100644 --- a/memoria/crates/memoria-api/src/lib.rs +++ b/memoria/crates/memoria-api/src/lib.rs @@ -238,6 +238,7 @@ pub fn build_router(state: AppState) -> Router { .route("/metrics", get(routes::metrics::prometheus_metrics)) // Memory reads .route("/v1/memories", get(routes::memory::list_memories)) + .route("/v1/memories/query", post(routes::memory::query_memories)) .route("/v1/memories/retrieve", post(routes::memory::retrieve)) .route("/v1/memories/search", post(routes::memory::search)) .route("/v1/memories/:id", get(routes::memory::get_memory)) diff --git a/memoria/crates/memoria-api/src/models.rs b/memoria/crates/memoria-api/src/models.rs index d53824e..0adfea6 100644 --- a/memoria/crates/memoria-api/src/models.rs +++ b/memoria/crates/memoria-api/src/models.rs @@ -3,6 +3,7 @@ use memoria_core::{Memory, MemoryType, TrustTier}; use serde::{Deserialize, Serialize}; +use std::collections::HashMap; use std::str::FromStr; // ── Memory ──────────────────────────────────────────────────────────────────── @@ -20,7 +21,7 @@ pub struct StoreRequest { pub source: Option, pub branch: Option, /// 任意业务元数据(如 scene/agent)。透传落库到 memories.extra_metadata,并在读取时原样 - /// 返回给调用方;Memoria 本身不对其做检索/打分逻辑(下游消费者如 matrixflow 的 decay 可自行使用)。 + /// 返回给调用方;结构化 query 可做精确过滤,但不参与相关性检索或打分。 #[serde(default)] pub extra_metadata: Option>, } @@ -67,7 +68,7 @@ fn parse_session_scope( fn parse_memory_types_opt( types: Option<&Vec>, ) -> Result>, String> { - types + let mut parsed = types .map(|ts| { ts.iter() .map(|s| s.trim()) @@ -76,7 +77,12 @@ fn parse_memory_types_opt( .collect::, _>>() }) .transpose() - .map(|v| v.filter(|t| !t.is_empty())) + .map(|v| v.filter(|t| !t.is_empty()))?; + if let Some(types) = parsed.as_mut() { + let mut seen = std::collections::HashSet::new(); + types.retain(|memory_type| seen.insert(memory_type.clone())); + } + Ok(parsed) } impl RetrieveRequest { @@ -143,6 +149,102 @@ impl SearchRequest { } } +fn default_structured_query_limit() -> i64 { + 100 +} + +/// A pure structured query. All supplied selectors are combined with AND. +/// Metadata equality preserves JSON type families (for example, string `"2"` +/// does not equal number `2`); JSON numbers `2` and `2.0` may compare equal. +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +pub struct StructuredQueryRequest { + #[serde(default)] + pub extra_metadata_filter: HashMap, + pub subject_id: Option, + pub memory_types: Option>, + pub session_id: Option, + pub trust_tier: Option, + pub branch: Option, + #[serde(default = "default_structured_query_limit")] + pub limit: i64, + pub cursor: Option, +} + +impl StructuredQueryRequest { + pub fn structured_options(&self) -> Result { + if !(1..=500).contains(&self.limit) { + return Err("limit must be between 1 and 500".to_string()); + } + memoria_storage::validate_extra_metadata_filter(&self.extra_metadata_filter) + .map_err(|err| err.to_string())?; + + let subject_id = normalized_filter("subject_id", self.subject_id.as_deref())?; + let session_id = normalized_filter("session_id", self.session_id.as_deref())?; + let normalized_trust_tier = + normalized_filter("trust_tier", self.trust_tier.as_deref())?; + let branch = normalized_filter("branch", self.branch.as_deref())?; + let trust_tier = normalized_trust_tier + .as_deref() + .map(parse_trust_tier) + .transpose()?; + if let Some(memory_types) = self.memory_types.as_ref() { + if memory_types.is_empty() || memory_types.iter().any(|value| value.trim().is_empty()) { + return Err( + "memory_types must contain only non-empty values when provided".to_string(), + ); + } + } + let memory_types = parse_memory_types_opt(self.memory_types.as_ref())?; + if self.extra_metadata_filter.is_empty() + && subject_id.is_none() + && session_id.is_none() + && trust_tier.is_none() + && memory_types.is_none() + && branch.is_none() + { + return Err("structured query requires at least one filter selector".to_string()); + } + + let cursor = normalized(self.cursor.as_deref()); + if let Some(cursor) = cursor.as_deref() { + if cursor.len() != 32 || !cursor.chars().all(|ch| ch.is_ascii_hexdigit()) { + return Err("cursor must be a 32-character hexadecimal memory_id".to_string()); + } + } + + Ok(memoria_service::StructuredQueryOptions { + limit: self.limit, + memory_types, + session_id, + trust_tier, + cursor, + subject_id, + extra_metadata_filter: self.extra_metadata_filter.clone(), + }) + } +} + +fn normalized(value: Option<&str>) -> Option { + value + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) +} + +fn normalized_filter(name: &str, value: Option<&str>) -> Result, String> { + value + .map(|value| { + let value = value.trim(); + if value.is_empty() { + Err(format!("{name} must not be empty when provided")) + } else { + Ok(value.to_string()) + } + }) + .transpose() +} + fn deserialize_explain<'de, D: serde::Deserializer<'de>>(d: D) -> Result { use serde::Deserialize; #[derive(Deserialize)] @@ -285,8 +387,8 @@ pub struct MemoryResponse { pub observed_at: Option, pub created_at: Option, pub retrieval_score: Option, - /// 业务元数据(如 scene/agent)从 memories.extra_metadata 原样透传回给调用方;Memoria 本身 - /// 不对其做检索/打分逻辑(下游消费者如 matrixflow 的 decay 可自行使用)。 + /// 业务元数据(如 scene/agent)从 memories.extra_metadata 原样透传回给调用方;结构化 query + /// 可做精确过滤,但不参与相关性检索或打分。 #[serde(skip_serializing_if = "Option::is_none")] pub extra_metadata: Option>, } @@ -498,7 +600,23 @@ pub fn parse_trust_tier(s: &str) -> Result { #[cfg(test)] mod tests { - use super::{PurgeRequest, PurgeSelector}; + use super::{parse_memory_types_opt, PurgeRequest, PurgeSelector}; + use memoria_core::MemoryType; + + #[test] + fn memory_type_parser_deduplicates_before_sql_option_construction() { + let raw = vec![ + "semantic".to_string(), + " semantic ".to_string(), + "profile".to_string(), + "semantic".to_string(), + ]; + + assert_eq!( + parse_memory_types_opt(Some(&raw)).unwrap(), + Some(vec![MemoryType::Semantic, MemoryType::Profile]) + ); + } #[test] fn purge_selector_ignores_empty_arrays() { diff --git a/memoria/crates/memoria-api/src/routes/memory.rs b/memoria/crates/memoria-api/src/routes/memory.rs index eb0a0a6..a8f8b2c 100644 --- a/memoria/crates/memoria-api/src/routes/memory.rs +++ b/memoria/crates/memoria-api/src/routes/memory.rs @@ -182,6 +182,32 @@ pub async fn list_memories( Ok(Json(ListResponse { items, next_cursor })) } +pub async fn query_memories( + State(state): State, + auth: AuthUser, + Json(req): Json, +) -> ApiResult { + let branch = normalize_branch(req.branch.clone()); + let limit = req.limit; + let mut options = req + .structured_options() + .map_err(|err| (StatusCode::UNPROCESSABLE_ENTITY, err))?; + options.limit = limit + 1; + + let mut memories = state + .service + .query_active_structured_on_branch(auth.scope_id(), branch.as_deref(), &options) + .await + .map_err(api_err_typed)?; + let has_more = memories.len() > limit as usize; + memories.truncate(limit as usize); + let next_cursor = has_more + .then(|| memories.last().map(|memory| memory.memory_id.clone())) + .flatten(); + let items = memories.into_iter().map(Into::into).collect(); + Ok(Json(ListResponse { items, next_cursor })) +} + pub async fn store_memory( State(state): State, auth: AuthUser, diff --git a/memoria/crates/memoria-api/tests/api_e2e.rs b/memoria/crates/memoria-api/tests/api_e2e.rs index edb24a9..d6448f9 100644 --- a/memoria/crates/memoria-api/tests/api_e2e.rs +++ b/memoria/crates/memoria-api/tests/api_e2e.rs @@ -394,6 +394,349 @@ async fn test_api_extra_metadata_round_trip_and_dedup() { assert_eq!(response.status(), 422); } +#[tokio::test] +async fn test_api_structured_query_by_extra_metadata() { + let (base, client, _server) = spawn_server().await; + let user_id = uid(); + + for (content, metadata) in [ + ( + "structured incident", + json!({"scene": "incident", "rank": 2, "urgent": true}), + ), + ( + "structured review", + json!({"scene": "review", "rank": 2, "urgent": true}), + ), + ( + "structured string rank", + json!({"scene": "incident", "rank": "2", "urgent": true}), + ), + ] { + let response = client + .post(format!("{base}/v1/memories")) + .header("X-User-Id", &user_id) + .json(&json!({"content": content, "extra_metadata": metadata})) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 201); + } + + let response = client + .post(format!("{base}/v1/memories/query")) + .header("X-User-Id", &user_id) + .json(&json!({ + "extra_metadata_filter": {"scene": "incident", "rank": 2, "urgent": true}, + "memory_types": ["semantic"], + "limit": 10 + })) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 200); + let body: Value = response.json().await.unwrap(); + let items = body["items"].as_array().unwrap(); + assert_eq!(items.len(), 1); + assert_eq!(items[0]["content"], "structured incident"); + assert_eq!(items[0]["retrieval_score"], Value::Null); + + let response = client + .post(format!("{base}/v1/memories/query")) + .header("X-User-Id", &user_id) + .json(&json!({})) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 422); + + let response = client + .post(format!("{base}/v1/memories/query")) + .header("X-User-Id", &user_id) + .json(&json!({"extra_metadata_filter": {"nested": {"value": 1}}})) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 422); + + for invalid_request in [ + json!({"subject_id": "subject", "limit": 0}), + json!({"subject_id": "subject", "limit": 501}), + json!({"subject_id": "subject", "extra_metadata_filters": {"scene": "incident"}}), + json!({"extra_metadata_filter": {"scene": "incident"}, "subject_id": " "}), + json!({"extra_metadata_filter": {"scene": "incident"}, "session_id": " "}), + json!({"extra_metadata_filter": {"scene": "incident"}, "trust_tier": " "}), + json!({"extra_metadata_filter": {"scene": "incident"}, "branch": " "}), + json!({"extra_metadata_filter": {"scene": "incident"}, "memory_types": []}), + json!({"extra_metadata_filter": {"scene": "incident"}, "memory_types": ["semantic", " "]}), + json!({"extra_metadata_filter": {"1scene": "incident"}}), + json!({"extra_metadata_filter": {"scene": "x".repeat(1025)}}), + ] { + let response = client + .post(format!("{base}/v1/memories/query")) + .header("X-User-Id", &user_id) + .json(&invalid_request) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 422, "request: {invalid_request}"); + } +} + +#[tokio::test] +async fn test_api_structured_query_enforces_scope_active_and_all_predicates() { + let (base, client, _server) = spawn_server().await; + let user_id = uid(); + let other_user_id = uid(); + let marker = format!("structured_filter_{}", uuid::Uuid::new_v4().simple()); + let matching_metadata = json!({"marker": marker, "urgent": true}); + + async fn store( + client: &reqwest::Client, + base: &str, + user_id: &str, + content: &str, + memory_type: &str, + metadata: Value, + ) -> String { + let response = client + .post(format!("{base}/v1/memories")) + .header("X-User-Id", user_id) + .json(&json!({ + "content": content, + "memory_type": memory_type, + "extra_metadata": metadata, + })) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 201); + response.json::().await.unwrap()["memory_id"] + .as_str() + .unwrap() + .to_string() + } + + let expected_id = store( + &client, + &base, + &user_id, + "only eligible structured result", + "semantic", + matching_metadata.clone(), + ) + .await; + store( + &client, + &base, + &user_id, + "wrong metadata boolean", + "semantic", + json!({"marker": marker, "urgent": false}), + ) + .await; + store( + &client, + &base, + &user_id, + "wrong memory type", + "profile", + matching_metadata.clone(), + ) + .await; + store( + &client, + &base, + &other_user_id, + "other tenant", + "semantic", + matching_metadata.clone(), + ) + .await; + let inactive_id = store( + &client, + &base, + &user_id, + "inactive match", + "semantic", + matching_metadata.clone(), + ) + .await; + let response = client + .delete(format!("{base}/v1/memories/{inactive_id}")) + .header("X-User-Id", &user_id) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 204); + + let response = client + .post(format!("{base}/v1/memories/query")) + .header("X-User-Id", &user_id) + .json(&json!({ + "extra_metadata_filter": {"marker": marker, "urgent": true}, + "memory_types": ["semantic"], + "limit": 10, + })) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 200); + let body: Value = response.json().await.unwrap(); + let items = body["items"].as_array().unwrap(); + assert_eq!(items.len(), 1); + assert_eq!(items[0]["memory_id"], expected_id); +} + +#[tokio::test] +async fn test_api_structured_query_cursor_and_subject_isolation() { + let (base, client, _server) = spawn_server().await; + let user_id = uid(); + let matching_subject = format!("subject_{}", uuid::Uuid::new_v4().simple()); + let other_subject = format!("subject_{}", uuid::Uuid::new_v4().simple()); + let marker = format!("marker_{}", uuid::Uuid::new_v4().simple()); + let mut expected_ids = std::collections::HashSet::new(); + + for index in 0..3 { + let response = client + .post(format!("{base}/v1/memories")) + .header("X-User-Id", &user_id) + .json(&json!({ + "content": format!("structured page {index}"), + "subject_id": matching_subject, + "extra_metadata": {"marker": marker} + })) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 201); + expected_ids.insert( + response.json::().await.unwrap()["memory_id"] + .as_str() + .unwrap() + .to_string(), + ); + } + + let response = client + .post(format!("{base}/v1/memories")) + .header("X-User-Id", &user_id) + .json(&json!({ + "content": "same marker, other subject", + "subject_id": other_subject, + "extra_metadata": {"marker": marker} + })) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 201); + + let first = client + .post(format!("{base}/v1/memories/query")) + .header("X-User-Id", &user_id) + .json(&json!({ + "extra_metadata_filter": {"marker": marker}, + "subject_id": matching_subject, + "limit": 2 + })) + .send() + .await + .unwrap(); + assert_eq!(first.status(), 200); + let first: Value = first.json().await.unwrap(); + let cursor = first["next_cursor"].as_str().expect("first page cursor"); + let first_ids: std::collections::HashSet = first["items"] + .as_array() + .unwrap() + .iter() + .map(|item| item["memory_id"].as_str().unwrap().to_string()) + .collect(); + assert_eq!(first_ids.len(), 2); + + let second = client + .post(format!("{base}/v1/memories/query")) + .header("X-User-Id", &user_id) + .json(&json!({ + "extra_metadata_filter": {"marker": marker}, + "subject_id": matching_subject, + "limit": 2, + "cursor": cursor + })) + .send() + .await + .unwrap(); + assert_eq!(second.status(), 200); + let second: Value = second.json().await.unwrap(); + assert!(second["next_cursor"].is_null()); + let second_ids: std::collections::HashSet = second["items"] + .as_array() + .unwrap() + .iter() + .map(|item| item["memory_id"].as_str().unwrap().to_string()) + .collect(); + assert_eq!(second_ids.len(), 1); + assert!(first_ids.is_disjoint(&second_ids)); + assert_eq!( + first_ids + .union(&second_ids) + .cloned() + .collect::>(), + expected_ids + ); +} + +#[tokio::test] +async fn test_api_structured_query_on_branch() { + let (base, client, _server) = spawn_server().await; + let user_id = uid(); + let branch = format!("query_{}", &uuid::Uuid::new_v4().simple().to_string()[..8]); + let marker = format!("branch_{}", uuid::Uuid::new_v4().simple()); + + let response = client + .post(format!("{base}/v1/branches")) + .header("X-User-Id", &user_id) + .json(&json!({"name": branch})) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 201); + + let response = client + .post(format!("{base}/v1/memories")) + .header("X-User-Id", &user_id) + .json(&json!({ + "content": "structured branch only", + "branch": branch, + "extra_metadata": {"marker": marker} + })) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 201); + + let branch_response = client + .post(format!("{base}/v1/memories/query")) + .header("X-User-Id", &user_id) + .json(&json!({"branch": branch})) + .send() + .await + .unwrap(); + assert_eq!(branch_response.status(), 200); + let branch_body: Value = branch_response.json().await.unwrap(); + assert_eq!(branch_body["items"].as_array().unwrap().len(), 1); + + let main_response = client + .post(format!("{base}/v1/memories/query")) + .header("X-User-Id", &user_id) + .json(&json!({"extra_metadata_filter": {"marker": marker}})) + .send() + .await + .unwrap(); + assert_eq!(main_response.status(), 200); + let main_body: Value = main_response.json().await.unwrap(); + assert!(main_body["items"].as_array().unwrap().is_empty()); +} + // ── 2b. list response is lightweight (no embedding) and respects limit ──────── #[tokio::test] diff --git a/memoria/crates/memoria-service/src/lib.rs b/memoria/crates/memoria-service/src/lib.rs index eb3d7e9..6c2cb2b 100644 --- a/memoria/crates/memoria-service/src/lib.rs +++ b/memoria/crates/memoria-service/src/lib.rs @@ -46,7 +46,8 @@ pub use scoring::{ }; pub use service::{ CandidateScore, ExplainLevel, InMemoryFlusher, ListActiveOptions, MemoryService, PurgeResult, - RetrievalExplain, RetrieveOptions, SessionScope, ENTITY_EXTRACTION_DROPS, + RetrievalExplain, RetrieveOptions, SessionScope, StructuredQueryOptions, + ENTITY_EXTRACTION_DROPS, }; pub use stats_reporter::StatsReporter; diff --git a/memoria/crates/memoria-service/src/service.rs b/memoria/crates/memoria-service/src/service.rs index 458c712..b3641f4 100644 --- a/memoria/crates/memoria-service/src/service.rs +++ b/memoria/crates/memoria-service/src/service.rs @@ -8,6 +8,7 @@ use memoria_embedding::llm::ChatMessage; use memoria_embedding::LlmClient; use memoria_storage::{DbRouter, OwnedEditLogEntry, SqlMemoryStore}; use moka::sync::Cache; +use std::collections::HashMap; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, OnceLock}; use std::time::Duration; @@ -162,6 +163,19 @@ pub struct ListActiveOptions<'a> { pub subject_id: Option<&'a str>, } +/// Filters for a structured memory query. This path performs no embedding, +/// keyword, graph, or relevance-scoring work. +#[derive(Debug, Clone)] +pub struct StructuredQueryOptions { + pub limit: i64, + pub memory_types: Option>, + pub session_id: Option, + pub trust_tier: Option, + pub cursor: Option, + pub subject_id: Option, + pub extra_metadata_filter: HashMap, +} + impl ListActiveOptions<'_> { pub fn new(limit: i64) -> Self { Self { @@ -2577,6 +2591,43 @@ impl MemoryService { Ok(mems) } + /// Query active memories using exact structured filters only. + pub async fn query_active_structured_on_branch( + &self, + user_id: &str, + branch: Option<&str>, + options: &StructuredQueryOptions, + ) -> Result, MemoriaError> { + if self.sql_store.is_some() { + let sql = self.user_sql_store(user_id).await?; + let table = sql.table_for_branch(user_id, branch).await?; + return sql + .query_active_structured_lite( + &table, + user_id, + options.limit, + options.memory_types.as_deref(), + options.session_id.as_deref(), + options + .trust_tier + .as_ref() + .map(ToString::to_string) + .as_deref(), + options.cursor.as_deref(), + options.subject_id.as_deref(), + &options.extra_metadata_filter, + ) + .await; + } + + // A bounded trait-level list cannot implement complete filtering or + // cursor pagination. Fail explicitly instead of returning plausible but + // incomplete results. Production configurations always provide SQL. + Err(MemoriaError::Internal( + "structured queries require a SQL-backed memory store".to_string(), + )) + } + pub async fn embed(&self, text: &str) -> Result>, MemoriaError> { match self.embedder.as_ref() { None => Ok(None), diff --git a/memoria/crates/memoria-storage/src/lib.rs b/memoria/crates/memoria-storage/src/lib.rs index ba2d5ff..0687bb2 100644 --- a/memoria/crates/memoria-storage/src/lib.rs +++ b/memoria/crates/memoria-storage/src/lib.rs @@ -22,7 +22,9 @@ pub use pool_config::{ }; pub use router::{DbRouter, UserDatabaseRecord}; pub use store::{ - snapshot_extra_memory_count, snapshot_extra_with_memory_count, FeedbackStats, MemoryFeedback, - OwnedEditLogEntry, PoolHealthLevel, PoolHealthSnapshot, SqlMemoryStore, TierFeedback, - UserRetrievalParams, ACTOR_USER_ID, + snapshot_extra_memory_count, snapshot_extra_with_memory_count, validate_extra_metadata_filter, + FeedbackStats, MemoryFeedback, OwnedEditLogEntry, PoolHealthLevel, PoolHealthSnapshot, + SqlMemoryStore, TierFeedback, UserRetrievalParams, ACTOR_USER_ID, + EXTRA_METADATA_FILTER_MAX_FIELDS, EXTRA_METADATA_FILTER_MAX_KEY_BYTES, + EXTRA_METADATA_FILTER_MAX_VALUE_BYTES, }; diff --git a/memoria/crates/memoria-storage/src/store.rs b/memoria/crates/memoria-storage/src/store.rs index 2c99fd2..eb99df0 100644 --- a/memoria/crates/memoria-storage/src/store.rs +++ b/memoria/crates/memoria-storage/src/store.rs @@ -12,6 +12,47 @@ use std::sync::atomic::{AtomicU64, AtomicU8, Ordering}; use std::sync::Arc; use std::time::{SystemTime, UNIX_EPOCH}; +pub const EXTRA_METADATA_FILTER_MAX_FIELDS: usize = 16; +pub const EXTRA_METADATA_FILTER_MAX_KEY_BYTES: usize = 64; +pub const EXTRA_METADATA_FILTER_MAX_VALUE_BYTES: usize = 1024; + +/// Validate the public structured-query metadata contract at the storage boundary. +/// Keys become JSON paths, so the first character must be an ASCII letter or +/// underscore and remaining characters are limited to ASCII alphanumerics/underscore. +pub fn validate_extra_metadata_filter( + filter: &std::collections::HashMap, +) -> Result<(), MemoriaError> { + if filter.len() > EXTRA_METADATA_FILTER_MAX_FIELDS { + return Err(MemoriaError::Validation(format!( + "extra_metadata_filter must not contain more than {EXTRA_METADATA_FILTER_MAX_FIELDS} fields" + ))); + } + + for (key, value) in filter { + let mut chars = key.chars(); + let valid_first = chars + .next() + .is_some_and(|ch| ch.is_ascii_alphabetic() || ch == '_'); + let valid_rest = chars.all(|ch| ch.is_ascii_alphanumeric() || ch == '_'); + if key.len() > EXTRA_METADATA_FILTER_MAX_KEY_BYTES || !valid_first || !valid_rest { + return Err(MemoriaError::Validation(format!( + "extra_metadata_filter key '{key}' must start with an ASCII letter or underscore, contain only ASCII letters, digits, or underscore, and be at most {EXTRA_METADATA_FILTER_MAX_KEY_BYTES} bytes" + ))); + } + if value.is_null() || value.is_array() || value.is_object() { + return Err(MemoriaError::Validation(format!( + "extra_metadata_filter value for '{key}' must be a string, number, or boolean" + ))); + } + if serde_json::to_string(value)?.len() > EXTRA_METADATA_FILTER_MAX_VALUE_BYTES { + return Err(MemoriaError::Validation(format!( + "extra_metadata_filter value for '{key}' must not exceed {EXTRA_METADATA_FILTER_MAX_VALUE_BYTES} bytes" + ))); + } + } + Ok(()) +} + tokio::task_local! { /// Real user ID for per-user state (active branch). /// In group mode the "user_id" flowing through the service layer is the @@ -5293,6 +5334,92 @@ impl SqlMemoryStore { rows.iter().map(row_to_memory_lite).collect() } + /// Exact structured query over ordinary columns and scalar extra_metadata fields. + /// This intentionally bypasses all vector/fulltext retrieval machinery. + #[allow(clippy::too_many_arguments)] + pub async fn query_active_structured_lite( + &self, + table: &str, + user_id: &str, + limit: i64, + memory_types: Option<&[MemoryType]>, + session_id: Option<&str>, + trust_tier: Option<&str>, + cursor: Option<&str>, + subject_id: Option<&str>, + extra_metadata_filter: &std::collections::HashMap, + ) -> Result, MemoriaError> { + if !(1..=501).contains(&limit) { + return Err(MemoriaError::Validation( + "structured query storage limit must be between 1 and 501".to_string(), + )); + } + validate_extra_metadata_filter(extra_metadata_filter)?; + let table = self.t(table); + let mut metadata_filters: Vec<_> = extra_metadata_filter.iter().collect(); + metadata_filters.sort_by(|(left, _), (right, _)| left.cmp(right)); + + let mut inner = + format!("SELECT memory_id FROM {table} WHERE user_id = ? AND is_active = 1"); + if let Some(types) = memory_types.filter(|types| !types.is_empty()) { + inner.push_str(" AND memory_type IN ("); + inner.push_str(&vec!["?"; types.len()].join(", ")); + inner.push(')'); + } + if session_id.is_some() { + inner.push_str(" AND session_id = ?"); + } + if trust_tier.is_some() { + inner.push_str(" AND trust_tier = ?"); + } + if subject_id.is_some() { + inner.push_str(" AND subject_id = ?"); + } + if cursor.is_some() { + inner.push_str(" AND memory_id < ?"); + } + for (key, _) in &metadata_filters { + inner.push_str(&format!( + " AND json_extract(extra_metadata, '$.{key}') = CAST(? AS JSON)" + )); + } + inner.push_str(" ORDER BY memory_id DESC LIMIT ?"); + + 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, \ + CAST(extra_metadata AS CHAR) AS extra_meta \ + FROM {table} WHERE memory_id IN ({inner}) \ + ORDER BY memory_id DESC" + ); + + let mut query = sqlx::query(&sql).bind(user_id); + if let Some(types) = memory_types.filter(|types| !types.is_empty()) { + for memory_type in types { + query = query.bind(memory_type.to_string()); + } + } + if let Some(value) = session_id { + query = query.bind(value); + } + if let Some(value) = trust_tier { + query = query.bind(value); + } + if let Some(value) = subject_id { + query = query.bind(value); + } + if let Some(value) = cursor { + query = query.bind(value); + } + for (_, value) in metadata_filters { + query = query.bind(serde_json::to_string(value)?); + } + query = query.bind(limit); + let rows = query.fetch_all(&self.pool).await.map_err(db_err)?; + rows.iter().map(row_to_memory_lite).collect() + } + /// Find memory IDs whose content contains `topic` (exact substring match). /// Uses fulltext boolean MUST with LIKE refinement. Requires topic >= 3 chars. pub async fn find_ids_by_topic( @@ -6008,8 +6135,8 @@ fn build_safety_snapshot_name(db_name: Option<&str>, operation: &str) -> String mod tests { use super::{ classify_pool_health, detect_connection_anomaly, should_emit_saturated_warning, - ConnectionAnomalyKind, OwnedEditLogEntry, PoolHealthLevel, PoolHealthSnapshot, - SqlMemoryStore, + validate_extra_metadata_filter, ConnectionAnomalyKind, OwnedEditLogEntry, PoolHealthLevel, + PoolHealthSnapshot, SqlMemoryStore, }; use sqlx::mysql::MySqlPoolOptions; use std::io::{self, Write}; @@ -6017,6 +6144,31 @@ mod tests { static LOG_TEST_LOCK: OnceLock> = OnceLock::new(); + #[test] + fn structured_metadata_filter_validation_is_enforced_at_storage_boundary() { + let valid = std::collections::HashMap::from([ + ("_scene".to_string(), serde_json::json!("incident")), + ("rank2".to_string(), serde_json::json!(2)), + ]); + assert!(validate_extra_metadata_filter(&valid).is_ok()); + + for invalid in [ + std::collections::HashMap::from([("1scene".to_string(), serde_json::json!(true))]), + std::collections::HashMap::from([("scene".to_string(), serde_json::json!([1]))]), + std::collections::HashMap::from([( + "scene".to_string(), + serde_json::json!("x".repeat(1025)), + )]), + ] { + assert!(validate_extra_metadata_filter(&invalid).is_err()); + } + + let too_many = (0..17) + .map(|index| (format!("key_{index}"), serde_json::json!(index))) + .collect(); + assert!(validate_extra_metadata_filter(&too_many).is_err()); + } + #[test] fn saturated_warning_requires_full_delay_and_only_emits_once() { assert!(!should_emit_saturated_warning(299, false)); diff --git a/sdk/python/CHANGELOG.md b/sdk/python/CHANGELOG.md index e843983..4457b03 100644 --- a/sdk/python/CHANGELOG.md +++ b/sdk/python/CHANGELOG.md @@ -2,6 +2,10 @@ ## [Unreleased] +### Added +- Sync and async `memories.query()` for exact structured filtering through the REST API, + including scalar `extra_metadata`, subject, type, session, trust tier, branch, and pagination. + ### Fixed - `ping()` no longer wraps `MemoriaAuthError` / `MemoriaNotFoundError` and other API errors into `MemoriaConnectionError`; callers can now distinguish network failures from API errors. diff --git a/sdk/python/README.md b/sdk/python/README.md index 50d9a24..5167a0a 100644 --- a/sdk/python/README.md +++ b/sdk/python/README.md @@ -75,6 +75,13 @@ result = client.memories.search(query="...", top_k=10) page = client.memories.list(limit=100, cursor=None) # page.next_cursor — pass as cursor= to get the next page +# Exact structured query (no vector/keyword retrieval) +page = client.memories.query( + extra_metadata_filter={"scene": "incident", "rank": 2}, + memory_types=["semantic"], + limit=100, +) + # Correct by ID mem = client.memories.correct("mem_id", new_content="...", reason="...") diff --git a/sdk/python/src/memoria/models.py b/sdk/python/src/memoria/models.py index bb8a27f..d3a6074 100644 --- a/sdk/python/src/memoria/models.py +++ b/sdk/python/src/memoria/models.py @@ -41,6 +41,10 @@ class Memory: observed_at: datetime | None = None created_at: datetime | None = None retrieval_score: float | None = None # populated by retrieve/search, None from list + # New response fields are appended so existing positional construction keeps + # the public dataclass field order released by earlier SDK versions. + subject_id: str | None = None + extra_metadata: dict[str, Any] | None = None @classmethod def from_dict(cls, d: dict[str, Any]) -> Memory: @@ -53,10 +57,12 @@ def from_dict(cls, d: dict[str, Any]) -> Memory: is_active=bool(d.get("is_active", True)), user_id=d.get("user_id", ""), author_id=d.get("author_id"), + subject_id=d.get("subject_id"), session_id=d.get("session_id"), observed_at=_parse_dt(d.get("observed_at")), created_at=_parse_dt(d.get("created_at")), retrieval_score=d.get("retrieval_score"), + extra_metadata=d.get("extra_metadata"), ) diff --git a/sdk/python/src/memoria/resources/memories.py b/sdk/python/src/memoria/resources/memories.py index a1b33a4..9c840dd 100644 --- a/sdk/python/src/memoria/resources/memories.py +++ b/sdk/python/src/memoria/resources/memories.py @@ -2,6 +2,8 @@ from __future__ import annotations +import json +import math from typing import TYPE_CHECKING, Any from ..exceptions import MemoriaValidationError @@ -10,11 +12,144 @@ if TYPE_CHECKING: from .._http import _HttpTransport +_MEMORY_TYPE_NAMES = { + "semantic", + "working", + "episodic", + "profile", + "tool_result", + "procedural", +} + def _strip_none(d: dict[str, Any]) -> dict[str, Any]: return {k: v for k, v in d.items() if v is not None} +def _normalize_query_memory_types( + memory_types: list[str] | None, +) -> list[str] | None: + if memory_types is None: + return None + if not isinstance(memory_types, list): + raise MemoriaValidationError("query: memory_types must be a list") + if not memory_types: + raise MemoriaValidationError( + "query: memory_types must contain at least one value when provided" + ) + normalized: list[str] = [] + seen: set[str] = set() + for value in memory_types: + if not isinstance(value, str): + raise MemoriaValidationError("query: memory_types entries must be strings") + value = value.strip() + if not value: + raise MemoriaValidationError("query: memory_types entries must not be empty") + if value not in _MEMORY_TYPE_NAMES: + raise MemoriaValidationError(f"query: unknown memory type: {value}") + if value not in seen: + seen.add(value) + normalized.append(value) + return normalized or None + + +def _validate_structured_query( + *, + extra_metadata_filter: dict[str, Any] | None, + subject_id: str | None, + memory_types: list[str] | None, + session_id: str | None, + trust_tier: str | None, + branch: str | None, + limit: int, +) -> tuple[ + list[str] | None, + str | None, + str | None, + str | None, + str | None, +]: + if type(limit) is not int: + raise MemoriaValidationError("query: limit must be an integer") + if limit < 1 or limit > 500: + raise MemoriaValidationError("query: limit must be between 1 and 500") + if extra_metadata_filter is not None and not isinstance(extra_metadata_filter, dict): + raise MemoriaValidationError("query: extra_metadata_filter must be a dictionary") + for name, value in [ + ("subject_id", subject_id), + ("session_id", session_id), + ("trust_tier", trust_tier), + ("branch", branch), + ]: + if value is not None and (not isinstance(value, str) or not value.strip()): + raise MemoriaValidationError(f"query: {name} must be a non-empty string when provided") + subject_id = subject_id.strip() if subject_id is not None else None + session_id = session_id.strip() if session_id is not None else None + trust_tier = trust_tier.strip() if trust_tier is not None else None + branch = branch.strip() if branch is not None else None + memory_types = _normalize_query_memory_types(memory_types) + has_selector = any( + [ + bool(extra_metadata_filter), + bool(subject_id and subject_id.strip()), + bool(memory_types), + bool(session_id and session_id.strip()), + bool(trust_tier and trust_tier.strip()), + bool(branch and branch.strip()), + ] + ) + if not has_selector: + raise MemoriaValidationError("query: at least one filter selector is required") + if extra_metadata_filter is not None: + if len(extra_metadata_filter) > 16: + raise MemoriaValidationError( + "query: extra_metadata_filter must not contain more than 16 fields" + ) + for key, value in extra_metadata_filter.items(): + if not isinstance(key, str): + raise MemoriaValidationError("query: extra_metadata_filter keys must be strings") + try: + key_bytes = len(key.encode("utf-8")) + except UnicodeEncodeError as error: + raise MemoriaValidationError( + "query: extra_metadata_filter keys must be valid UTF-8" + ) from error + valid_key = ( + bool(key) + and key_bytes <= 64 + and (key[0].isascii() and (key[0].isalpha() or key[0] == "_")) + and all(char.isascii() and (char.isalnum() or char == "_") for char in key[1:]) + ) + if not valid_key: + raise MemoriaValidationError( + "query: extra_metadata_filter keys must start with an ASCII letter or " + "underscore, contain only ASCII letters, digits, or underscore, and " + "must not exceed 64 bytes" + ) + if not isinstance(value, (str, int, float, bool)): + raise MemoriaValidationError( + "query: extra_metadata_filter values must be strings, numbers, or booleans" + ) + if isinstance(value, float) and not math.isfinite(value): + raise MemoriaValidationError( + "query: extra_metadata_filter numeric values must be finite" + ) + try: + encoded_value = json.dumps(value, ensure_ascii=False, separators=(",", ":")).encode( + "utf-8" + ) + except (ValueError, UnicodeEncodeError) as error: + raise MemoriaValidationError( + "query: extra_metadata_filter values must be valid JSON scalars " + "encoded as UTF-8" + ) from error + if len(encoded_value) > 1024: + raise MemoriaValidationError( + "query: extra_metadata_filter values must not exceed 1024 bytes" + ) + return memory_types, subject_id, session_id, trust_tier, branch + + class MemoriesResource: def __init__(self, client: _HttpTransport) -> None: self._client = client @@ -130,6 +265,43 @@ def list( data = self._client._request("GET", "/v1/memories", params=params) return MemoryPage.from_dict(data) + def query( + self, + *, + extra_metadata_filter: dict[str, Any] | None = None, + subject_id: str | None = None, + memory_types: list[str] | None = None, + session_id: str | None = None, + trust_tier: str | None = None, + branch: str | None = None, + limit: int = 100, + cursor: str | None = None, + ) -> MemoryPage: + """Run an exact structured query without vector or keyword retrieval.""" + memory_types, subject_id, session_id, trust_tier, branch = _validate_structured_query( + extra_metadata_filter=extra_metadata_filter, + subject_id=subject_id, + memory_types=memory_types, + session_id=session_id, + trust_tier=trust_tier, + branch=branch, + limit=limit, + ) + body = _strip_none( + { + "extra_metadata_filter": extra_metadata_filter, + "subject_id": subject_id, + "memory_types": memory_types, + "session_id": session_id, + "trust_tier": trust_tier, + "branch": branch, + "limit": limit, + "cursor": cursor, + } + ) + data = self._client._request("POST", "/v1/memories/query", json=body) + return MemoryPage.from_dict(data) + def correct( self, id: str, @@ -343,6 +515,43 @@ async def list( data = await self._client._arequest("GET", "/v1/memories", params=params) return MemoryPage.from_dict(data) + async def query( + self, + *, + extra_metadata_filter: dict[str, Any] | None = None, + subject_id: str | None = None, + memory_types: list[str] | None = None, + session_id: str | None = None, + trust_tier: str | None = None, + branch: str | None = None, + limit: int = 100, + cursor: str | None = None, + ) -> MemoryPage: + """Run an exact structured query without vector or keyword retrieval.""" + memory_types, subject_id, session_id, trust_tier, branch = _validate_structured_query( + extra_metadata_filter=extra_metadata_filter, + subject_id=subject_id, + memory_types=memory_types, + session_id=session_id, + trust_tier=trust_tier, + branch=branch, + limit=limit, + ) + body = _strip_none( + { + "extra_metadata_filter": extra_metadata_filter, + "subject_id": subject_id, + "memory_types": memory_types, + "session_id": session_id, + "trust_tier": trust_tier, + "branch": branch, + "limit": limit, + "cursor": cursor, + } + ) + data = await self._client._arequest("POST", "/v1/memories/query", json=body) + return MemoryPage.from_dict(data) + async def correct( self, id: str, diff --git a/sdk/python/tests/unit/test_memories.py b/sdk/python/tests/unit/test_memories.py index 72dde43..05989c3 100644 --- a/sdk/python/tests/unit/test_memories.py +++ b/sdk/python/tests/unit/test_memories.py @@ -2,13 +2,21 @@ from __future__ import annotations +import json + import pytest from pytest_httpx import HTTPXMock -from memoria import MemoriaClient, MemoriaAuthError, MemoriaForbiddenError, MemoriaNotFoundError -from memoria import MemoriaUnprocessableError, MemoriaValidationError +from memoria import ( + MemoriaAuthError, + MemoriaClient, + MemoriaForbiddenError, + MemoriaNotFoundError, + MemoriaUnprocessableError, + MemoriaValidationError, +) from memoria.models import Memory, MemoryPage, PurgeResult, RetrieveResult -from tests.conftest import BASE_URL, API_KEY, MEMORY_STUB +from tests.conftest import API_KEY, BASE_URL, MEMORY_STUB @pytest.fixture @@ -29,6 +37,29 @@ def test_store_happy_path(httpx_mock: HTTPXMock, client: MemoriaClient) -> None: assert mem.content == "test content" +def test_memory_preserves_legacy_positional_field_order() -> None: + memory = Memory( + "mem_legacy", + "legacy content", + "semantic", + "T3", + 0.65, + True, + "user_1", + "author_1", + "session_1", + None, + None, + 0.75, + ) + + assert memory.author_id == "author_1" + assert memory.session_id == "session_1" + assert memory.retrieval_score == 0.75 + assert memory.subject_id is None + assert memory.extra_metadata is None + + def test_store_with_all_params(httpx_mock: HTTPXMock, client: MemoriaClient) -> None: httpx_mock.add_response(json=MEMORY_STUB) mem = client.memories.store( @@ -41,7 +72,6 @@ def test_store_with_all_params(httpx_mock: HTTPXMock, client: MemoriaClient) -> assert mem.memory_id == "mem_abc123" req = httpx_mock.get_request() assert req is not None - import json body = json.loads(req.content) assert body["memory_type"] == "profile" assert body["session_id"] == "sess_1" @@ -83,10 +113,9 @@ def test_store_batch_happy_path(httpx_mock: HTTPXMock, client: MemoriaClient) -> ) assert len(mems) == 2 # Verify the request body uses "memories" (not "items") to match the server contract - import json as _json req = httpx_mock.get_request() assert req is not None - body = _json.loads(req.content) + body = json.loads(req.content) assert "memories" in body assert "items" not in body @@ -147,6 +176,134 @@ def test_list_over_max_limit_raises(client: MemoriaClient) -> None: client.memories.list(limit=501) +def test_structured_query(httpx_mock: HTTPXMock, client: MemoriaClient) -> None: + response = { + **MEMORY_STUB, + "subject_id": "subject_1", + "extra_metadata": {"scene": "incident", "rank": 2}, + } + httpx_mock.add_response(json={"items": [response], "next_cursor": "next_id"}) + + page = client.memories.query( + extra_metadata_filter={"scene": "incident", "rank": 2}, + subject_id="subject_1", + memory_types=["semantic", " semantic ", "semantic"], + limit=10, + ) + + assert page.items[0].subject_id == "subject_1" + assert page.items[0].extra_metadata == {"scene": "incident", "rank": 2} + assert page.next_cursor == "next_id" + request = httpx_mock.get_request() + assert request is not None + assert request.url.path == "/v1/memories/query" + body = json.loads(request.content) + assert body["extra_metadata_filter"] == {"scene": "incident", "rank": 2} + assert body["memory_types"] == ["semantic"] + assert "query" not in body + + +def test_structured_query_requires_selector(client: MemoriaClient) -> None: + with pytest.raises(MemoriaValidationError, match="selector"): + client.memories.query() + + +@pytest.mark.parametrize("limit", ["1", 1.5, True, None]) +def test_structured_query_rejects_non_integer_limit(client: MemoriaClient, limit: object) -> None: + with pytest.raises(MemoriaValidationError, match="limit must be an integer"): + client.memories.query(subject_id="subject", limit=limit) # type: ignore[arg-type] + + +@pytest.mark.parametrize("extra_metadata_filter", [[("scene", "incident")], "scene=incident"]) +def test_structured_query_rejects_non_dictionary_metadata_filter( + client: MemoriaClient, extra_metadata_filter: object +) -> None: + with pytest.raises(MemoriaValidationError, match="must be a dictionary"): + client.memories.query( + subject_id="subject", + extra_metadata_filter=extra_metadata_filter, # type: ignore[arg-type] + ) + + +def test_structured_query_accepts_branch_only( + httpx_mock: HTTPXMock, client: MemoriaClient +) -> None: + httpx_mock.add_response(json={"items": [], "next_cursor": None}) + page = client.memories.query(branch="experiment") + assert page.items == [] + request = httpx_mock.get_request() + assert request is not None + assert json.loads(request.content)["branch"] == "experiment" + + +def test_structured_query_rejects_nested_metadata(client: MemoriaClient) -> None: + with pytest.raises(MemoriaValidationError, match="strings, numbers, or booleans"): + client.memories.query(extra_metadata_filter={"nested": {"key": "value"}}) + + +def test_structured_query_rejects_invalid_key_and_oversized_value( + client: MemoriaClient, +) -> None: + with pytest.raises(MemoriaValidationError, match="must start"): + client.memories.query(extra_metadata_filter={"1scene": "incident"}) + with pytest.raises(MemoriaValidationError, match="1024"): + client.memories.query(extra_metadata_filter={"scene": "x" * 1025}) + + +def test_structured_query_reports_metadata_key_byte_limit( + client: MemoriaClient, +) -> None: + with pytest.raises(MemoriaValidationError, match="64 bytes"): + client.memories.query(extra_metadata_filter={"a" * 65: "incident"}) + + +@pytest.mark.parametrize("memory_types", [["unknown"], [1], "semantic"]) +def test_structured_query_rejects_invalid_memory_types( + client: MemoriaClient, memory_types: object +) -> None: + with pytest.raises(MemoriaValidationError, match="memory_types|memory type"): + client.memories.query(memory_types=memory_types) # type: ignore[arg-type] + + +@pytest.mark.parametrize( + "kwargs", + [ + {"subject_id": " "}, + {"session_id": " "}, + {"trust_tier": " "}, + {"branch": " "}, + {"memory_types": []}, + {"memory_types": ["semantic", " "]}, + ], +) +def test_structured_query_rejects_blank_supplied_selector( + client: MemoriaClient, kwargs: dict[str, object] +) -> None: + with pytest.raises(MemoriaValidationError, match="empty|at least one"): + client.memories.query( + extra_metadata_filter={"scene": "incident"}, + **kwargs, # type: ignore[arg-type] + ) + + +@pytest.mark.parametrize( + "value", [10**5000, "\ud800"], ids=["large-integer", "lone-surrogate"] +) +def test_structured_query_translates_metadata_serialization_errors( + client: MemoriaClient, value: object +) -> None: + with pytest.raises(MemoriaValidationError, match="valid JSON scalars"): + client.memories.query(extra_metadata_filter={"value": value}) + + +@pytest.mark.parametrize("value", [float("nan"), float("inf"), float("-inf")]) +def test_structured_query_rejects_non_finite_metadata_number( + client: MemoriaClient, value: float +) -> None: + with pytest.raises(MemoriaValidationError, match="finite"): + client.memories.query(extra_metadata_filter={"rank": value}) + + # --------------------------------------------------------------------------- # correct / correct_by_query # --------------------------------------------------------------------------- diff --git a/sdk/python/tests/unit/test_memories_async.py b/sdk/python/tests/unit/test_memories_async.py index f9065c2..0343c1d 100644 --- a/sdk/python/tests/unit/test_memories_async.py +++ b/sdk/python/tests/unit/test_memories_async.py @@ -7,7 +7,7 @@ from memoria import AsyncMemoriaClient, MemoriaAuthError, MemoriaValidationError from memoria.models import Memory, MemoryPage, PurgeResult, RetrieveResult -from tests.conftest import BASE_URL, API_KEY, MEMORY_STUB +from tests.conftest import API_KEY, BASE_URL, MEMORY_STUB @pytest.fixture @@ -46,6 +46,54 @@ async def test_list_happy_path(httpx_mock: HTTPXMock, client: AsyncMemoriaClient assert page.next_cursor == "cursor_xyz" +@pytest.mark.asyncio +async def test_structured_query(httpx_mock: HTTPXMock, client: AsyncMemoriaClient) -> None: + response = {**MEMORY_STUB, "extra_metadata": {"scene": "incident"}} + httpx_mock.add_response(json={"items": [response], "next_cursor": None}) + page = await client.memories.query( + extra_metadata_filter={"scene": "incident"}, trust_tier="T2" + ) + assert page.items[0].extra_metadata == {"scene": "incident"} + request = httpx_mock.get_request() + assert request is not None + assert request.url.path == "/v1/memories/query" + + +@pytest.mark.asyncio +async def test_structured_query_accepts_branch_only( + httpx_mock: HTTPXMock, client: AsyncMemoriaClient +) -> None: + httpx_mock.add_response(json={"items": [], "next_cursor": None}) + page = await client.memories.query(branch="experiment") + assert page.items == [] + + +@pytest.mark.asyncio +async def test_structured_query_rejects_blank_supplied_selector( + client: AsyncMemoriaClient, +) -> None: + with pytest.raises(MemoriaValidationError, match="subject_id"): + await client.memories.query( + extra_metadata_filter={"scene": "incident"}, subject_id=" " + ) + + +@pytest.mark.asyncio +async def test_structured_query_rejects_invalid_runtime_types( + client: AsyncMemoriaClient, +) -> None: + with pytest.raises(MemoriaValidationError, match="limit must be an integer"): + await client.memories.query( + subject_id="subject", + limit="1", # type: ignore[arg-type] + ) + with pytest.raises(MemoriaValidationError, match="must be a dictionary"): + await client.memories.query( + subject_id="subject", + extra_metadata_filter=[("scene", "incident")], # type: ignore[arg-type] + ) + + @pytest.mark.asyncio async def test_purge_by_ids(httpx_mock: HTTPXMock, client: AsyncMemoriaClient) -> None: httpx_mock.add_response(json={"purged": 1, "snapshot_name": "snap_x"})