From 0296cae6f44e014265262b487567106a4a6632a6 Mon Sep 17 00:00:00 2001 From: rock Date: Thu, 3 Sep 2026 09:15:35 -0700 Subject: [PATCH] refactor(handlers): extract LearnParams + reusable RBAC helpers learn_handler refactored: - Extract LearnParams struct with validation + bounds clamping - Extract store_compacted_memory helper - Extract build_learn_response helper - Reuse check_project_write_access for RBAC ingest_handler refactored: - Extract check_project_write_access (reusable) - Extract execute_ingest helper New tests (6 total): - LearnParams validation tests Total tests: 694 (was 688) --- crates/mem-cli/src/handlers/learn.rs | 172 +++++++++++++++++ crates/mem-cli/src/handlers/mod.rs | 2 + crates/mem-cli/src/http_server.rs | 264 +++++++++++++-------------- 3 files changed, 302 insertions(+), 136 deletions(-) create mode 100644 crates/mem-cli/src/handlers/learn.rs diff --git a/crates/mem-cli/src/handlers/learn.rs b/crates/mem-cli/src/handlers/learn.rs new file mode 100644 index 0000000..b9f7643 --- /dev/null +++ b/crates/mem-cli/src/handlers/learn.rs @@ -0,0 +1,172 @@ +/// Learn Handler Helpers +/// +/// Extracted to reduce learn_handler complexity. + +use actix_web::HttpResponse; +use serde_json::{json, Value}; + +// ============================================================================ +// Learn Parameters +// ============================================================================ + +/// Validated learn request parameters +#[derive(Debug, Clone)] +pub struct LearnParams { + pub project: String, + pub text: String, + pub question: String, + pub memory_budget: u32, + pub chunk_size: usize, + pub model: String, +} + +impl LearnParams { + /// Parse and validate learn request body + pub fn from_body(body: &Value) -> Result { + let project = body["project"] + .as_str() + .unwrap_or("knowledge") + .to_string(); + + let text = body["text"] + .as_str() + .filter(|t| !t.trim().is_empty()) + .ok_or(LearnParamsError::MissingText)? + .to_string(); + + let question = body["query"] + .as_str() + .unwrap_or("What are the key facts, patterns, and practices in this knowledge?") + .to_string(); + + let memory_budget = body["memory_budget"] + .as_u64() + .unwrap_or(4096) + .clamp(256, 32768) as u32; + + let chunk_size = body["chunk_size"] + .as_u64() + .unwrap_or(2000) + .clamp(500, 10000) as usize; + + let model = body["model"] + .as_str() + .unwrap_or("qwen2.5:3b-instruct") + .to_string(); + + Ok(Self { project, text, question, memory_budget, chunk_size, model }) + } +} + +/// Learn parameter validation errors +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum LearnParamsError { + MissingText, +} + +impl LearnParamsError { + pub fn to_response(&self) -> HttpResponse { + let reason = match self { + Self::MissingText => "missing required field: text", + }; + HttpResponse::BadRequest().json(json!({ + "error": "bad_request", + "reason": reason + })) + } +} + +// ============================================================================ +// Learn Response Builder +// ============================================================================ + +/// Build learn response JSON +pub fn build_learn_response( + project: &str, + model: &str, + chunks_seen: u32, + chunks_used: u32, + memory: &str, + stored: bool, +) -> HttpResponse { + HttpResponse::Ok().json(json!({ + "project": project, + "status": "completed", + "chunks_seen": chunks_seen, + "chunks_used": chunks_used, + "memory": memory, + "memory_tokens": memory.len() / 4, + "stored": stored, + "model": model, + })) +} + +// ============================================================================ +// Tests +// ============================================================================ + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_learn_params_valid() { + let body = json!({ + "project": "homelab", + "text": "Some knowledge to learn" + }); + let params = LearnParams::from_body(&body).unwrap(); + + assert_eq!(params.project, "homelab"); + assert_eq!(params.text, "Some knowledge to learn"); + assert_eq!(params.memory_budget, 4096); + assert_eq!(params.chunk_size, 2000); + } + + #[test] + fn test_learn_params_defaults() { + let body = json!({"text": "content"}); + let params = LearnParams::from_body(&body).unwrap(); + + assert_eq!(params.project, "knowledge"); + assert_eq!(params.model, "qwen2.5:3b-instruct"); + } + + #[test] + fn test_learn_params_custom_values() { + let body = json!({ + "text": "content", + "memory_budget": 8192, + "chunk_size": 3000, + "model": "gpt-4" + }); + let params = LearnParams::from_body(&body).unwrap(); + + assert_eq!(params.memory_budget, 8192); + assert_eq!(params.chunk_size, 3000); + assert_eq!(params.model, "gpt-4"); + } + + #[test] + fn test_learn_params_clamped_budget() { + let body = json!({"text": "x", "memory_budget": 999999}); + let params = LearnParams::from_body(&body).unwrap(); + assert_eq!(params.memory_budget, 32768); + + let body = json!({"text": "x", "memory_budget": 10}); + let params = LearnParams::from_body(&body).unwrap(); + assert_eq!(params.memory_budget, 256); + } + + #[test] + fn test_learn_params_missing_text() { + let body = json!({"project": "test"}); + assert_eq!(LearnParams::from_body(&body).unwrap_err(), LearnParamsError::MissingText); + } + + #[test] + fn test_learn_params_empty_text() { + let body = json!({"text": " "}); + assert_eq!(LearnParams::from_body(&body).unwrap_err(), LearnParamsError::MissingText); + } +} diff --git a/crates/mem-cli/src/handlers/mod.rs b/crates/mem-cli/src/handlers/mod.rs index f416234..119f0fb 100644 --- a/crates/mem-cli/src/handlers/mod.rs +++ b/crates/mem-cli/src/handlers/mod.rs @@ -5,6 +5,8 @@ pub mod query; pub mod ingest; +pub mod learn; pub use query::*; pub use ingest::*; +pub use learn::*; diff --git a/crates/mem-cli/src/http_server.rs b/crates/mem-cli/src/http_server.rs index 98c7338..1a4bd80 100644 --- a/crates/mem-cli/src/http_server.rs +++ b/crates/mem-cli/src/http_server.rs @@ -19,7 +19,7 @@ use crate::gateway_queue_adapter::GatewayQueueAdapter; use crate::queue_worker::{QueueWorker, QueueWorkerConfig}; use crate::queue_adapter::QueueAdapter; use crate::rbac::{AccessGuard, Claims as RbacClaims, builtin_role_provider, ResourceMeta, ResourceType, Verb, Visibility}; -use crate::handlers::{QueryParams, QueryParamsError, SearchMethod, build_search_response}; +use crate::handlers::{QueryParams, QueryParamsError, SearchMethod, build_search_response, LearnParams, LearnParamsError, build_learn_response}; /// Server state with database and workers pub struct AppState { @@ -426,55 +426,69 @@ pub async fn ingest_handler( body: web::Json, state: web::Data, ) -> HttpResponse { + // Auth + capability check let (claims, _token) = match validate_auth(&req, &state).await { Ok(c) => c, Err(e) => return e, }; - - // Check write capability if !has_capability(&claims, "memory:write") { return HttpResponse::Forbidden().json(json!({ "error": "forbidden", "reason": "missing capability: memory:write" })); } - if let Err(e) = check_rate_limit(&claims, &state, "/memory/ingest") { return e; } - let project = body.project.clone(); - // RBAC: Check project-level write access - if let Some(guard) = &state.access_guard { - let rbac_claims = to_rbac_claims(&claims); - let project_resource = ResourceMeta::new(&project, ResourceType::Project, &project); - - if !guard.can_write(&rbac_claims, &project_resource).await { - tracing::warn!( - "RBAC denied write access to project '{}' for user '{}'", - project, claims.sub - ); - return HttpResponse::Forbidden().json(json!({ - "error": "forbidden", - "reason": format!("write access denied to project '{}'", project) - })); - } + if let Err(e) = check_project_write_access(&state, &claims, &body.project).await { + return e; } - let ingest_id = body.ingest_id.clone(); - let records: Vec<(String, String)> = body - .records + + // Check idempotency + if let Some(cached) = state.idempotency_store.get(&body.ingest_id) { + tracing::info!("Returning cached response for ingest_id: {}", body.ingest_id); + return HttpResponse::Accepted().json(cached); + } + + // Execute ingest + execute_ingest(&state, &body).await +} + +/// Check RBAC project write access +async fn check_project_write_access( + state: &web::Data, + claims: &JwtClaims, + project: &str, +) -> Result<(), HttpResponse> { + let Some(guard) = &state.access_guard else { + return Ok(()); + }; + + let rbac_claims = to_rbac_claims(claims); + let resource = ResourceMeta::new(project, ResourceType::Project, project); + + if !guard.can_write(&rbac_claims, &resource).await { + tracing::warn!("RBAC denied write access to project '{}' for user '{}'", project, claims.sub); + return Err(HttpResponse::Forbidden().json(json!({ + "error": "forbidden", + "reason": format!("write access denied to project '{}'", project) + }))); + } + Ok(()) +} + +/// Execute ingest job creation and spawn worker +async fn execute_ingest( + state: &web::Data, + body: &IngestRequest, +) -> HttpResponse { + let records: Vec<(String, String)> = body.records .iter() .map(|r| (r.text.clone(), body.source.clone())) .collect(); - // Check idempotency cache first - if let Some(cached_response) = state.idempotency_store.get(&ingest_id) { - tracing::info!("Returning cached response for ingest_id: {}", ingest_id); - return HttpResponse::Accepted().json(cached_response); - } - - // Create ingest job in DB let job_result = sqlx::query( "INSERT INTO ingest_jobs (id, project, ingest_id, status, created_at) VALUES ($1, $2, $3, 'pending', NOW()) @@ -482,49 +496,39 @@ pub async fn ingest_handler( RETURNING id", ) .bind(uuid::Uuid::new_v4()) - .bind(&project) - .bind(&ingest_id) + .bind(&body.project) + .bind(&body.ingest_id) .fetch_optional(&state.pool) .await; + let response = json!({ + "ingest_id": body.ingest_id, + "status": "pending", + "status_url": format!("/memory/ingest/{}", body.ingest_id) + }); + match job_result { Ok(Some(_)) => { // Spawn async ingest task let worker = state.ingest_worker.clone(); - let proj = project.clone(); - let id = ingest_id.clone(); + let project = body.project.clone(); + let ingest_id = body.ingest_id.clone(); tokio::spawn(async move { - if let Err(e) = worker.process_ingest(&proj, &id, records).await { + if let Err(e) = worker.process_ingest(&project, &ingest_id, records).await { tracing::error!("Ingest failed: {}", e); } }); - - let response = json!({ - "ingest_id": ingest_id, - "status": "pending", - "status_url": format!("/memory/ingest/{}", ingest_id) - }); - - // Cache the response for idempotency - state.idempotency_store.set(ingest_id.clone(), response.clone()); - + state.idempotency_store.set(body.ingest_id.clone(), response.clone()); HttpResponse::Accepted().json(response) } Ok(None) => { - // Already exists in DB (was inserted concurrently) - let response = json!({ - "ingest_id": ingest_id, - "status": "pending", - "status_url": format!("/memory/ingest/{}", ingest_id) - }); - state.idempotency_store.set(ingest_id.clone(), response.clone()); + // Already exists (concurrent insert) + state.idempotency_store.set(body.ingest_id.clone(), response.clone()); HttpResponse::Accepted().json(response) } Err(e) => { tracing::error!("DB error: {}", e); - HttpResponse::InternalServerError().json(json!({ - "error": "database_error" - })) + HttpResponse::InternalServerError().json(json!({"error": "database_error"})) } } } @@ -642,56 +646,34 @@ pub async fn learn_handler( body: web::Json, state: web::Data, ) -> HttpResponse { + // Auth + capability check let (claims, _token) = match validate_auth(&req, &state).await { Ok(c) => c, Err(e) => return e, }; - if !has_capability(&claims, "memory:write") { return HttpResponse::Forbidden().json(json!({ "error": "forbidden", "reason": "missing capability: memory:write" })); } - if let Err(e) = check_rate_limit(&claims, &state, "/memory/ingest") { return e; } - let project = body["project"].as_str().unwrap_or("knowledge").to_string(); + // Parse + validate params + let params = match LearnParams::from_body(&body) { + Ok(p) => p, + Err(e) => return e.to_response(), + }; // RBAC: Check project-level write access - if let Some(guard) = &state.access_guard { - let rbac_claims = to_rbac_claims(&claims); - let project_resource = ResourceMeta::new(&project, ResourceType::Project, &project); - - if !guard.can_write(&rbac_claims, &project_resource).await { - tracing::warn!( - "RBAC denied write access to project '{}' for user '{}'", - project, claims.sub - ); - return HttpResponse::Forbidden().json(json!({ - "error": "forbidden", - "reason": format!("write access denied to project '{}'", project) - })); - } + if let Err(e) = check_project_write_access(&state, &claims, ¶ms.project).await { + return e; } - let text = match body["text"].as_str() { - Some(t) => t.to_string(), - None => return HttpResponse::BadRequest().json(json!({ - "error": "bad_request", - "reason": "missing required field: text" - })), - }; - let question = body["query"].as_str() - .unwrap_or("What are the key facts, patterns, and practices in this knowledge?") - .to_string(); - let memory_budget = body["memory_budget"].as_u64().unwrap_or(4096) as u32; - let chunk_size = body["chunk_size"].as_u64().unwrap_or(2000) as usize; - let model = body["model"].as_str().unwrap_or("qwen2.5:3b-instruct").to_string(); // Chunk the markdown - let chunks = chunk_markdown_text(&text, chunk_size); + let chunks = chunk_markdown_text(¶ms.text, params.chunk_size); if chunks.is_empty() { return HttpResponse::BadRequest().json(json!({ "error": "bad_request", @@ -709,7 +691,7 @@ pub async fn learn_handler( text: text.clone(), timestamp: time::OffsetDateTime::now_utc(), provenance: mem_core::Provenance { - source_id: format!("learn://{}:{}", project, i), + source_id: format!("learn://{}:{}", params.project, i), offset: i as u64, }, }; @@ -720,14 +702,14 @@ pub async fn learn_handler( // Build query for gated loop let query = mem_core::query::Query { id: format!("learn-{}", uuid::Uuid::new_v4()), - question, + question: params.question.clone(), exit_gate: false, }; let config = mem_core::gated_loop::LoopConfig { level: mem_core::Level::L1, query, - memory_budget, + memory_budget: params.memory_budget, use_exit_gate: false, }; @@ -737,7 +719,7 @@ pub async fn learn_handler( let llm_key = std::env::var("LLM_API_KEY") .or_else(|_| std::env::var("MEM_API_KEY")) .unwrap_or_default(); - let llm = match mem_llm::ChatClient::new(&llm_base, &llm_key, &model) { + let llm = match mem_llm::ChatClient::new(&llm_base, &llm_key, ¶ms.model) { Ok(c) => c, Err(e) => return HttpResponse::InternalServerError().json(json!({ "error": "llm_init_failed", @@ -754,55 +736,65 @@ pub async fn learn_handler( })), }; - // Store compacted memory in pgvector if non-empty - let mut stored = false; - if !outcome.final_memory.is_empty() { - match state.embeddings.embed_one(&outcome.final_memory).await { - Ok(embedding) => { - let chunk_id = uuid::Uuid::new_v4(); - let sha = { - use sha2::{Digest, Sha256}; - let mut h = Sha256::new(); - h.update(outcome.final_memory.as_bytes()); - format!("{:x}", h.finalize()) - }; - let result = sqlx::query( - "INSERT INTO memory_chunks (id, project, level, text, embedding, sha256, source, created_at) - VALUES ($1, $2, 'L1', $3, $4, $5, $6, NOW()) - ON CONFLICT (sha256) DO UPDATE SET text = $3, embedding = $4", - ) - .bind(chunk_id) - .bind(&project) - .bind(&outcome.final_memory) - .bind(embedding.to_vec()) - .bind(&sha) - .bind(format!("learn://{}", project)) - .fetch_optional(&state.pool) - .await; + // Store compacted memory + let stored = store_compacted_memory(&state, ¶ms.project, &outcome.final_memory).await; - match result { - Ok(_) => { stored = true; } - Err(e) => { - tracing::error!("Failed to store compacted memory: {}", e); - } - } - } - Err(e) => { - tracing::error!("Failed to embed compacted memory: {}", e); - } - } + build_learn_response( + ¶ms.project, + ¶ms.model, + outcome.chunks_seen, + outcome.chunks_used, + &outcome.final_memory, + stored, + ) +} + +/// Store compacted memory in pgvector +async fn store_compacted_memory( + state: &web::Data, + project: &str, + memory: &str, +) -> bool { + if memory.is_empty() { + return false; } - HttpResponse::Ok().json(json!({ - "project": project, - "status": "completed", - "chunks_seen": outcome.chunks_seen, - "chunks_used": outcome.chunks_used, - "memory": outcome.final_memory, - "memory_tokens": outcome.final_memory.len() / 4, - "stored": stored, - "model": model, - })) + let embedding = match state.embeddings.embed_one(memory).await { + Ok(e) => e, + Err(e) => { + tracing::error!("Failed to embed compacted memory: {}", e); + return false; + } + }; + + let sha = { + use sha2::{Digest, Sha256}; + let mut h = Sha256::new(); + h.update(memory.as_bytes()); + format!("{:x}", h.finalize()) + }; + + let result = sqlx::query( + "INSERT INTO memory_chunks (id, project, level, text, embedding, sha256, source, created_at) + VALUES ($1, $2, 'L1', $3, $4, $5, $6, NOW()) + ON CONFLICT (sha256) DO UPDATE SET text = $3, embedding = $4", + ) + .bind(uuid::Uuid::new_v4()) + .bind(project) + .bind(memory) + .bind(embedding.to_vec()) + .bind(&sha) + .bind(format!("learn://{}", project)) + .fetch_optional(&state.pool) + .await; + + match result { + Ok(_) => true, + Err(e) => { + tracing::error!("Failed to store compacted memory: {}", e); + false + } + } } /// Split markdown on ## headings for learn endpoint.