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)
This commit is contained in:
@@ -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<Self, LearnParamsError> {
|
||||||
|
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);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -5,6 +5,8 @@
|
|||||||
|
|
||||||
pub mod query;
|
pub mod query;
|
||||||
pub mod ingest;
|
pub mod ingest;
|
||||||
|
pub mod learn;
|
||||||
|
|
||||||
pub use query::*;
|
pub use query::*;
|
||||||
pub use ingest::*;
|
pub use ingest::*;
|
||||||
|
pub use learn::*;
|
||||||
|
|||||||
+128
-136
@@ -19,7 +19,7 @@ use crate::gateway_queue_adapter::GatewayQueueAdapter;
|
|||||||
use crate::queue_worker::{QueueWorker, QueueWorkerConfig};
|
use crate::queue_worker::{QueueWorker, QueueWorkerConfig};
|
||||||
use crate::queue_adapter::QueueAdapter;
|
use crate::queue_adapter::QueueAdapter;
|
||||||
use crate::rbac::{AccessGuard, Claims as RbacClaims, builtin_role_provider, ResourceMeta, ResourceType, Verb, Visibility};
|
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
|
/// Server state with database and workers
|
||||||
pub struct AppState {
|
pub struct AppState {
|
||||||
@@ -426,55 +426,69 @@ pub async fn ingest_handler(
|
|||||||
body: web::Json<IngestRequest>,
|
body: web::Json<IngestRequest>,
|
||||||
state: web::Data<AppState>,
|
state: web::Data<AppState>,
|
||||||
) -> HttpResponse {
|
) -> HttpResponse {
|
||||||
|
// Auth + capability check
|
||||||
let (claims, _token) = match validate_auth(&req, &state).await {
|
let (claims, _token) = match validate_auth(&req, &state).await {
|
||||||
Ok(c) => c,
|
Ok(c) => c,
|
||||||
Err(e) => return e,
|
Err(e) => return e,
|
||||||
};
|
};
|
||||||
|
|
||||||
// Check write capability
|
|
||||||
if !has_capability(&claims, "memory:write") {
|
if !has_capability(&claims, "memory:write") {
|
||||||
return HttpResponse::Forbidden().json(json!({
|
return HttpResponse::Forbidden().json(json!({
|
||||||
"error": "forbidden",
|
"error": "forbidden",
|
||||||
"reason": "missing capability: memory:write"
|
"reason": "missing capability: memory:write"
|
||||||
}));
|
}));
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Err(e) = check_rate_limit(&claims, &state, "/memory/ingest") {
|
if let Err(e) = check_rate_limit(&claims, &state, "/memory/ingest") {
|
||||||
return e;
|
return e;
|
||||||
}
|
}
|
||||||
|
|
||||||
let project = body.project.clone();
|
|
||||||
|
|
||||||
// RBAC: Check project-level write access
|
// RBAC: Check project-level write access
|
||||||
if let Some(guard) = &state.access_guard {
|
if let Err(e) = check_project_write_access(&state, &claims, &body.project).await {
|
||||||
let rbac_claims = to_rbac_claims(&claims);
|
return e;
|
||||||
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)
|
|
||||||
}));
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
let ingest_id = body.ingest_id.clone();
|
|
||||||
let records: Vec<(String, String)> = body
|
// Check idempotency
|
||||||
.records
|
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<AppState>,
|
||||||
|
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<AppState>,
|
||||||
|
body: &IngestRequest,
|
||||||
|
) -> HttpResponse {
|
||||||
|
let records: Vec<(String, String)> = body.records
|
||||||
.iter()
|
.iter()
|
||||||
.map(|r| (r.text.clone(), body.source.clone()))
|
.map(|r| (r.text.clone(), body.source.clone()))
|
||||||
.collect();
|
.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(
|
let job_result = sqlx::query(
|
||||||
"INSERT INTO ingest_jobs (id, project, ingest_id, status, created_at)
|
"INSERT INTO ingest_jobs (id, project, ingest_id, status, created_at)
|
||||||
VALUES ($1, $2, $3, 'pending', NOW())
|
VALUES ($1, $2, $3, 'pending', NOW())
|
||||||
@@ -482,49 +496,39 @@ pub async fn ingest_handler(
|
|||||||
RETURNING id",
|
RETURNING id",
|
||||||
)
|
)
|
||||||
.bind(uuid::Uuid::new_v4())
|
.bind(uuid::Uuid::new_v4())
|
||||||
.bind(&project)
|
.bind(&body.project)
|
||||||
.bind(&ingest_id)
|
.bind(&body.ingest_id)
|
||||||
.fetch_optional(&state.pool)
|
.fetch_optional(&state.pool)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
|
let response = json!({
|
||||||
|
"ingest_id": body.ingest_id,
|
||||||
|
"status": "pending",
|
||||||
|
"status_url": format!("/memory/ingest/{}", body.ingest_id)
|
||||||
|
});
|
||||||
|
|
||||||
match job_result {
|
match job_result {
|
||||||
Ok(Some(_)) => {
|
Ok(Some(_)) => {
|
||||||
// Spawn async ingest task
|
// Spawn async ingest task
|
||||||
let worker = state.ingest_worker.clone();
|
let worker = state.ingest_worker.clone();
|
||||||
let proj = project.clone();
|
let project = body.project.clone();
|
||||||
let id = ingest_id.clone();
|
let ingest_id = body.ingest_id.clone();
|
||||||
tokio::spawn(async move {
|
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);
|
tracing::error!("Ingest failed: {}", e);
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
|
state.idempotency_store.set(body.ingest_id.clone(), response.clone());
|
||||||
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());
|
|
||||||
|
|
||||||
HttpResponse::Accepted().json(response)
|
HttpResponse::Accepted().json(response)
|
||||||
}
|
}
|
||||||
Ok(None) => {
|
Ok(None) => {
|
||||||
// Already exists in DB (was inserted concurrently)
|
// Already exists (concurrent insert)
|
||||||
let response = json!({
|
state.idempotency_store.set(body.ingest_id.clone(), response.clone());
|
||||||
"ingest_id": ingest_id,
|
|
||||||
"status": "pending",
|
|
||||||
"status_url": format!("/memory/ingest/{}", ingest_id)
|
|
||||||
});
|
|
||||||
state.idempotency_store.set(ingest_id.clone(), response.clone());
|
|
||||||
HttpResponse::Accepted().json(response)
|
HttpResponse::Accepted().json(response)
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
tracing::error!("DB error: {}", e);
|
tracing::error!("DB error: {}", e);
|
||||||
HttpResponse::InternalServerError().json(json!({
|
HttpResponse::InternalServerError().json(json!({"error": "database_error"}))
|
||||||
"error": "database_error"
|
|
||||||
}))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -642,56 +646,34 @@ pub async fn learn_handler(
|
|||||||
body: web::Json<serde_json::Value>,
|
body: web::Json<serde_json::Value>,
|
||||||
state: web::Data<AppState>,
|
state: web::Data<AppState>,
|
||||||
) -> HttpResponse {
|
) -> HttpResponse {
|
||||||
|
// Auth + capability check
|
||||||
let (claims, _token) = match validate_auth(&req, &state).await {
|
let (claims, _token) = match validate_auth(&req, &state).await {
|
||||||
Ok(c) => c,
|
Ok(c) => c,
|
||||||
Err(e) => return e,
|
Err(e) => return e,
|
||||||
};
|
};
|
||||||
|
|
||||||
if !has_capability(&claims, "memory:write") {
|
if !has_capability(&claims, "memory:write") {
|
||||||
return HttpResponse::Forbidden().json(json!({
|
return HttpResponse::Forbidden().json(json!({
|
||||||
"error": "forbidden",
|
"error": "forbidden",
|
||||||
"reason": "missing capability: memory:write"
|
"reason": "missing capability: memory:write"
|
||||||
}));
|
}));
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Err(e) = check_rate_limit(&claims, &state, "/memory/ingest") {
|
if let Err(e) = check_rate_limit(&claims, &state, "/memory/ingest") {
|
||||||
return e;
|
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
|
// RBAC: Check project-level write access
|
||||||
if let Some(guard) = &state.access_guard {
|
if let Err(e) = check_project_write_access(&state, &claims, ¶ms.project).await {
|
||||||
let rbac_claims = to_rbac_claims(&claims);
|
return e;
|
||||||
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)
|
|
||||||
}));
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
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
|
// 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() {
|
if chunks.is_empty() {
|
||||||
return HttpResponse::BadRequest().json(json!({
|
return HttpResponse::BadRequest().json(json!({
|
||||||
"error": "bad_request",
|
"error": "bad_request",
|
||||||
@@ -709,7 +691,7 @@ pub async fn learn_handler(
|
|||||||
text: text.clone(),
|
text: text.clone(),
|
||||||
timestamp: time::OffsetDateTime::now_utc(),
|
timestamp: time::OffsetDateTime::now_utc(),
|
||||||
provenance: mem_core::Provenance {
|
provenance: mem_core::Provenance {
|
||||||
source_id: format!("learn://{}:{}", project, i),
|
source_id: format!("learn://{}:{}", params.project, i),
|
||||||
offset: i as u64,
|
offset: i as u64,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
@@ -720,14 +702,14 @@ pub async fn learn_handler(
|
|||||||
// Build query for gated loop
|
// Build query for gated loop
|
||||||
let query = mem_core::query::Query {
|
let query = mem_core::query::Query {
|
||||||
id: format!("learn-{}", uuid::Uuid::new_v4()),
|
id: format!("learn-{}", uuid::Uuid::new_v4()),
|
||||||
question,
|
question: params.question.clone(),
|
||||||
exit_gate: false,
|
exit_gate: false,
|
||||||
};
|
};
|
||||||
|
|
||||||
let config = mem_core::gated_loop::LoopConfig {
|
let config = mem_core::gated_loop::LoopConfig {
|
||||||
level: mem_core::Level::L1,
|
level: mem_core::Level::L1,
|
||||||
query,
|
query,
|
||||||
memory_budget,
|
memory_budget: params.memory_budget,
|
||||||
use_exit_gate: false,
|
use_exit_gate: false,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -737,7 +719,7 @@ pub async fn learn_handler(
|
|||||||
let llm_key = std::env::var("LLM_API_KEY")
|
let llm_key = std::env::var("LLM_API_KEY")
|
||||||
.or_else(|_| std::env::var("MEM_API_KEY"))
|
.or_else(|_| std::env::var("MEM_API_KEY"))
|
||||||
.unwrap_or_default();
|
.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,
|
Ok(c) => c,
|
||||||
Err(e) => return HttpResponse::InternalServerError().json(json!({
|
Err(e) => return HttpResponse::InternalServerError().json(json!({
|
||||||
"error": "llm_init_failed",
|
"error": "llm_init_failed",
|
||||||
@@ -754,55 +736,65 @@ pub async fn learn_handler(
|
|||||||
})),
|
})),
|
||||||
};
|
};
|
||||||
|
|
||||||
// Store compacted memory in pgvector if non-empty
|
// Store compacted memory
|
||||||
let mut stored = false;
|
let stored = store_compacted_memory(&state, ¶ms.project, &outcome.final_memory).await;
|
||||||
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;
|
|
||||||
|
|
||||||
match result {
|
build_learn_response(
|
||||||
Ok(_) => { stored = true; }
|
¶ms.project,
|
||||||
Err(e) => {
|
¶ms.model,
|
||||||
tracing::error!("Failed to store compacted memory: {}", e);
|
outcome.chunks_seen,
|
||||||
}
|
outcome.chunks_used,
|
||||||
}
|
&outcome.final_memory,
|
||||||
}
|
stored,
|
||||||
Err(e) => {
|
)
|
||||||
tracing::error!("Failed to embed compacted memory: {}", e);
|
}
|
||||||
}
|
|
||||||
}
|
/// Store compacted memory in pgvector
|
||||||
|
async fn store_compacted_memory(
|
||||||
|
state: &web::Data<AppState>,
|
||||||
|
project: &str,
|
||||||
|
memory: &str,
|
||||||
|
) -> bool {
|
||||||
|
if memory.is_empty() {
|
||||||
|
return false;
|
||||||
}
|
}
|
||||||
|
|
||||||
HttpResponse::Ok().json(json!({
|
let embedding = match state.embeddings.embed_one(memory).await {
|
||||||
"project": project,
|
Ok(e) => e,
|
||||||
"status": "completed",
|
Err(e) => {
|
||||||
"chunks_seen": outcome.chunks_seen,
|
tracing::error!("Failed to embed compacted memory: {}", e);
|
||||||
"chunks_used": outcome.chunks_used,
|
return false;
|
||||||
"memory": outcome.final_memory,
|
}
|
||||||
"memory_tokens": outcome.final_memory.len() / 4,
|
};
|
||||||
"stored": stored,
|
|
||||||
"model": model,
|
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.
|
/// Split markdown on ## headings for learn endpoint.
|
||||||
|
|||||||
Reference in New Issue
Block a user