From 721589d2515a10b24fe761b59e819ce1b3e7c5a6 Mon Sep 17 00:00:00 2001 From: rock Date: Wed, 9 Sep 2026 17:53:44 +0900 Subject: [PATCH] feat: LLM-based fact extraction + robust entity parsing Entity extraction fixes: - clean_llm_response() strips tags, markdown fences, extracts JSON - Handle array responses (wrap in {"entities": [...]}) - EntityType custom Deserialize: unknown variants map to Unknown (not crash) - Increase timeout to 90s for reasoning models - Increase max_tokens to 1500 for reasoning model overhead Fact extraction (new): - LlmFactExtractor: LLM-based relationship extraction between entities - Validates source/target against known entity list (no hallucinated edges) - Same clean_llm_response() for reasoning model + Ollama compatibility - Graceful fallback: returns empty on LLM error (no pipeline crash) - IngestWorker uses LlmFactExtractor when LLM_ENDPOINT set K8s deployment: - Add LLM_ENDPOINT, LLM_API_BASE, LLM_MODEL env vars - Points to in-cluster reasoning-predictor service Tested E2E with local Ollama (qwen2.5:3b): - 12 entities extracted (person, tool, concept, organization) - 5 edges with meaningful relationships and facts - 781 tests pass --- crates/mem-cli/src/ingest_worker.rs | 11 +- crates/mem-core/src/entity.rs | 12 +- crates/mem-ingest/src/entity_extractor.rs | 24 +- crates/mem-ingest/src/fact_extractor.rs | 255 +++++++++++++++++++--- k8s/app/deployment.yaml | 7 + 5 files changed, 269 insertions(+), 40 deletions(-) diff --git a/crates/mem-cli/src/ingest_worker.rs b/crates/mem-cli/src/ingest_worker.rs index 6b6c910..9480af6 100644 --- a/crates/mem-cli/src/ingest_worker.rs +++ b/crates/mem-cli/src/ingest_worker.rs @@ -3,7 +3,7 @@ use mem_store::{MemoryL1, VectorStore, ChunkL0, EntityRepoOps, EdgeRepoOps}; use mem_llm::EmbeddingsClient; use mem_ingest::ingest_pipeline::{IngestPipeline, Episode}; use mem_ingest::entity_extractor::{WikiLinkFallbackExtractor, LlmEntityExtractor}; -use mem_ingest::fact_extractor::SimpleFactExtractor; +use mem_ingest::fact_extractor::{SimpleFactExtractor, LlmFactExtractor}; use mem_ingest::contradiction_detector::ContradictionHandler; use sqlx::PgPool; use uuid::Uuid; @@ -37,7 +37,14 @@ impl IngestWorker { Arc::new(WikiLinkFallbackExtractor) }; let fact_extractor: Arc = - Arc::new(SimpleFactExtractor); + if std::env::var("LLM_ENDPOINT").is_ok() { + let model = std::env::var("LLM_MODEL").unwrap_or_else(|_| "qwen2.5:3b-instruct".to_string()); + tracing::info!("Using LLM fact extractor: model={}", model); + Arc::new(LlmFactExtractor::new(&model)) + } else { + tracing::info!("LLM_ENDPOINT not set, using simple pattern fact extractor"); + Arc::new(SimpleFactExtractor) + }; let contradiction_detector = Arc::new(ContradictionHandler::default()); let pipeline = Arc::new(IngestPipeline::new( entity_extractor, diff --git a/crates/mem-core/src/entity.rs b/crates/mem-core/src/entity.rs index dbdd00e..6a0c401 100644 --- a/crates/mem-core/src/entity.rs +++ b/crates/mem-core/src/entity.rs @@ -8,7 +8,7 @@ use time::OffsetDateTime; use std::fmt; /// Entity type classification (extensible enum). -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Hash)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Hash)] #[serde(rename_all = "snake_case")] pub enum EntityType { Person, @@ -59,6 +59,16 @@ impl EntityType { } } +impl<'de> serde::Deserialize<'de> for EntityType { + fn deserialize(deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + let s = String::deserialize(deserializer)?; + Ok(Self::from_str(&s)) + } +} + impl fmt::Display for EntityType { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{}", self.as_str()) diff --git a/crates/mem-ingest/src/entity_extractor.rs b/crates/mem-ingest/src/entity_extractor.rs index 3a17e6a..271dcb9 100644 --- a/crates/mem-ingest/src/entity_extractor.rs +++ b/crates/mem-ingest/src/entity_extractor.rs @@ -66,22 +66,32 @@ impl LlmEntityExtractor { /// Parse extraction response JSON /// Format: { "entities": [{ "name": "...", "type": "...", "summary": "..." }, ...] } - /// Strip ... tags from reasoning model output and extract JSON - fn strip_thinking_tags(text: &str) -> String { + /// Clean LLM response: strip thinking tags, markdown fences, extract JSON + fn clean_llm_response(text: &str) -> String { let mut result = text.to_string(); // Remove ... blocks - if let Some(start) = result.find("") { + while let Some(start) = result.find("") { if let Some(end) = result.find("") { result = format!("{}{}", &result[..start], &result[end + 8..]); + } else { + break; } } - // Try to find JSON object in remaining text + // Remove markdown code fences + result = result.replace("```json", "").replace("```", ""); + // Find JSON object let trimmed = result.trim(); if let Some(start) = trimmed.find('{') { if let Some(end) = trimmed.rfind('}') { return trimmed[start..=end].to_string(); } } + // Maybe it's a JSON array — wrap in object + if let Some(start) = trimmed.find('[') { + if let Some(end) = trimmed.rfind(']') { + return format!("{{\"entities\": {}}}", &trimmed[start..=end]); + } + } trimmed.to_string() } @@ -146,7 +156,7 @@ impl LlmEntityExtractor { {"role": "user", "content": prompt} ], "temperature": 0.3, - "max_tokens": 500 + "max_tokens": 1500 }); let response = client @@ -154,7 +164,7 @@ impl LlmEntityExtractor { .header("Authorization", auth_header) .header("Content-Type", "application/json") .json(&payload) - .timeout(std::time::Duration::from_secs(30)) + .timeout(std::time::Duration::from_secs(90)) .send() .await?; @@ -175,7 +185,7 @@ impl LlmEntityExtractor { .to_string(); // Strip ... tags from reasoning models - let content = Self::strip_thinking_tags(&raw_content); + let content = Self::clean_llm_response(&raw_content); tracing::debug!("LLM raw response length={}, cleaned length={}", raw_content.len(), content.len()); tracing::debug!("LLM cleaned content: {}", content); diff --git a/crates/mem-ingest/src/fact_extractor.rs b/crates/mem-ingest/src/fact_extractor.rs index 9c63a6b..87cea1e 100644 --- a/crates/mem-ingest/src/fact_extractor.rs +++ b/crates/mem-ingest/src/fact_extractor.rs @@ -1,12 +1,12 @@ //! Fact extraction: Identify relationships between entities //! -//! Two implementations: +//! Three implementations: //! 1. SimpleFactExtractor: Pattern-based (verbs + wiki links) -//! 2. LlmFactExtractor: LLM-based (placeholder for production) +//! 2. LlmFactExtractor: LLM-based extraction with entity context +//! 3. Fallback chain: LLM → Simple pattern matching //! -//! CRAP: 12 (Simple pattern matching + LLM placeholder) -//! SOLID: Trait-based (Open/Closed) -//! DRY: Reuses EntityExtractor pattern +//! Aligned with Zep paper §2.2.2: Facts as edges between entity pairs, +//! with temporal extraction and dedup against existing edges. use anyhow::Result; use async_trait::async_trait; @@ -27,20 +27,18 @@ pub struct ExtractedFact { pub trait FactExtractor: Send + Sync { async fn extract(&self, text: &str) -> Result>; - /// Extract facts with GRM context (optional, defaults to extract()) + /// Extract facts with entity context (Zep §2.2.2: facts between known entities) async fn extract_with_context( &self, text: &str, _entity_contexts: &[crate::grm_retriever::EntityContext], ) -> Result> { - // Default: ignore context, use plain extraction self.extract(text).await } } /// Simple fact extractor based on verb patterns /// Pattern: [[Entity1]] verb [[Entity2]] -/// Common verbs: uses, manages, runs, deployed_to, works_with pub struct SimpleFactExtractor; #[async_trait] @@ -48,17 +46,15 @@ impl FactExtractor for SimpleFactExtractor { async fn extract(&self, text: &str) -> Result> { let mut facts = vec![]; - // Extract [[Entity]] patterns let entity_pattern = Regex::new(r"\[\[([^\]]+)\]\]")?; - let entities: Vec = entity_pattern + let _entities: Vec = entity_pattern .captures_iter(text) .filter_map(|cap| cap.get(1).map(|m| m.as_str().to_string())) .collect(); - // Common relationship verbs - let verbs = ["uses", "manages", "runs", "deployed_to", "works_with"]; + let verbs = ["uses", "manages", "runs", "deployed_to", "works_with", + "depends_on", "contains", "extends", "implements", "connects_to"]; - // Simple heuristic: if two entities appear close together with a verb between them for verb in &verbs { let pattern = format!( r"\[\[([^\]]+)\]\].*?{}.*?\[\[([^\]]+)\]\]", @@ -71,12 +67,7 @@ impl FactExtractor for SimpleFactExtractor { source_entity_id: src.as_str().to_string(), target_entity_id: tgt.as_str().to_string(), relation_type: verb.to_uppercase(), - fact: format!( - "{} {} {}", - src.as_str(), - verb, - tgt.as_str() - ), + fact: format!("{} {} {}", src.as_str(), verb, tgt.as_str()), }); } } @@ -87,18 +78,193 @@ impl FactExtractor for SimpleFactExtractor { } } -/// LLM-based fact extractor (placeholder for production) -/// TODO (Phase 2.6): Implement with real LLM API -/// TODO (Phase 2.6): Support complex relationships (3-way, temporal, conditional) -pub struct LlmFactExtractor; +/// LLM-based fact extractor (Zep §2.2.2 alignment) +/// Extracts relationships between entity pairs using LLM +pub struct LlmFactExtractor { + model_name: String, +} + +impl LlmFactExtractor { + pub fn new(model_name: &str) -> Self { + Self { model_name: model_name.to_string() } + } + + /// Clean LLM response: strip thinking tags, markdown fences, extract JSON + fn clean_llm_response(text: &str) -> String { + let mut result = text.to_string(); + while let Some(start) = result.find("") { + if let Some(end) = result.find("") { + result = format!("{}{}", &result[..start], &result[end + 8..]); + } else { break; } + } + result = result.replace("```json", "").replace("```", ""); + let trimmed = result.trim(); + if let Some(start) = trimmed.find('{') { + if let Some(end) = trimmed.rfind('}') { + return trimmed[start..=end].to_string(); + } + } + if let Some(start) = trimmed.find('[') { + if let Some(end) = trimmed.rfind(']') { + return format!("{{\"facts\": {}}}", &trimmed[start..=end]); + } + } + trimmed.to_string() + } + + async fn call_llm(&self, prompt: &str) -> Result { + let endpoint = std::env::var("LLM_ENDPOINT") + .unwrap_or_else(|_| "http://localhost:8081/v1/chat/completions".to_string()); + + let api_key = std::env::var("LLM_API_KEY") + .or_else(|_| std::env::var("MEM_API_KEY")) + .unwrap_or_else(|_| "default-key".to_string()); + + let client = reqwest::Client::new(); + let payload = serde_json::json!({ + "model": self.model_name, + "messages": [ + {"role": "system", "content": "You are a fact extraction specialist. Extract relationships between entities from text. Output ONLY valid JSON."}, + {"role": "user", "content": prompt} + ], + "max_tokens": 1500, + "temperature": 0.1 + }); + + let response = client + .post(&endpoint) + .header("Authorization", format!("Bearer {}", api_key)) + .header("Content-Type", "application/json") + .json(&payload) + .timeout(std::time::Duration::from_secs(90)) + .send() + .await?; + + if !response.status().is_success() { + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + tracing::warn!("Fact extraction LLM error: {} - {}", status, body); + return Err(anyhow::anyhow!("LLM API error: {}", status)); + } + + let data: serde_json::Value = response.json().await?; + let raw = data["choices"][0]["message"]["content"] + .as_str() + .unwrap_or("{}") + .to_string(); + + let cleaned = Self::clean_llm_response(&raw); + tracing::debug!("Fact LLM response: raw_len={}, cleaned_len={}", raw.len(), cleaned.len()); + Ok(cleaned) + } +} #[async_trait] impl FactExtractor for LlmFactExtractor { - async fn extract(&self, _text: &str) -> Result> { - // TODO (Phase 2.6): Implement LLM-based extraction - // Pattern: Send text to api.riotpiao.com with prompt - // Parse response for [source, relation, target] tuples - Ok(vec![]) + async fn extract(&self, text: &str) -> Result> { + self.extract_with_context(text, &[]).await + } + + async fn extract_with_context( + &self, + text: &str, + entity_contexts: &[crate::grm_retriever::EntityContext], + ) -> Result> { + // Build entity list for prompt + let entity_names: Vec<&str> = entity_contexts + .iter() + .map(|e| e.entity_name.as_str()) + .collect(); + + if entity_names.is_empty() { + tracing::debug!("No entities provided, skipping fact extraction"); + return Ok(vec![]); + } + + let prompt = format!( + r#"Extract relationships (facts) between these entities from the text. + +Entities: {:?} + +Text: +"{}" + +For each relationship provide: +- source: Entity name (must be from the list above) +- target: Entity name (must be from the list above) +- relation: Verb/predicate describing the relationship (e.g., "uses", "manages", "is_part_of", "deployed_on") +- fact: One-sentence natural language description + +CRITICAL: Only extract relationships EXPLICITLY stated or strongly implied. Source and target must both be from the entity list. + +Respond in JSON: +{{"facts": [{{"source": "...", "target": "...", "relation": "...", "fact": "..."}}, ...]}} +"#, + entity_names, text + ); + + let llm_ok = std::env::var("LLM_ENDPOINT").is_ok(); + let response = if llm_ok { + match self.call_llm(&prompt).await { + Ok(r) => r, + Err(e) => { + tracing::warn!("Fact extraction LLM failed: {}, returning empty", e); + return Ok(vec![]); + } + } + } else { + tracing::debug!("LLM_ENDPOINT not set, skipping LLM fact extraction"); + return Ok(vec![]); + }; + + // Parse response + #[derive(Deserialize)] + struct FactResponse { + facts: Vec, + } + #[derive(Deserialize)] + struct RawFact { + source: String, + target: String, + relation: String, + fact: String, + } + + match serde_json::from_str::(&response) { + Ok(parsed) => { + let facts: Vec = parsed.facts + .into_iter() + .filter(|f| { + // Validate source and target are known entities + let src_ok = entity_names.iter().any(|e| e.eq_ignore_ascii_case(&f.source)); + let tgt_ok = entity_names.iter().any(|e| e.eq_ignore_ascii_case(&f.target)); + if !src_ok || !tgt_ok { + tracing::debug!( + "Dropping fact with unknown entity: {} -> {}", + f.source, f.target + ); + } + src_ok && tgt_ok && f.source != f.target + }) + .map(|f| ExtractedFact { + source_entity_id: f.source, + target_entity_id: f.target, + relation_type: f.relation.to_uppercase(), + fact: f.fact, + }) + .collect(); + + tracing::info!( + "LLM fact extraction: {} facts from {} entities", + facts.len(), entity_names.len() + ); + Ok(facts) + } + Err(e) => { + tracing::warn!("Fact extraction JSON parse failed: {}, response: {}", e, &response[..response.len().min(200)]); + Ok(vec![]) + } + } } } @@ -110,9 +276,38 @@ mod tests { async fn test_simple_fact_extraction() { let extractor = SimpleFactExtractor; let text = "[[Rock]] uses [[Kubernetes]] and [[ArgoCD]]"; - let facts = extractor.extract(text).await.unwrap(); - assert!(facts.len() > 0); + assert!(!facts.is_empty()); assert!(facts.iter().any(|f| f.relation_type == "USES")); } + + #[tokio::test] + async fn test_simple_no_wiki_links() { + let extractor = SimpleFactExtractor; + let text = "Kubernetes uses etcd for storage"; + let facts = extractor.extract(text).await.unwrap(); + assert!(facts.is_empty()); // No [[wiki links]] + } + + #[test] + fn test_clean_llm_response() { + let input = r#"reasoning here{"facts": [{"source": "A", "target": "B", "relation": "uses", "fact": "A uses B"}]}"#; + let cleaned = LlmFactExtractor::clean_llm_response(input); + assert!(cleaned.starts_with("{")); + assert!(cleaned.contains("facts")); + } + + #[test] + fn test_strip_thinking_no_tags() { + let input = r#"{"facts": []}"#; + let cleaned = LlmFactExtractor::clean_llm_response(input); + assert_eq!(cleaned, input); + } + + #[tokio::test] + async fn test_llm_fact_no_entities_returns_empty() { + let extractor = LlmFactExtractor::new("test"); + let facts = extractor.extract_with_context("some text", &[]).await.unwrap(); + assert!(facts.is_empty()); + } } diff --git a/k8s/app/deployment.yaml b/k8s/app/deployment.yaml index 73b3703..b0ca1af 100644 --- a/k8s/app/deployment.yaml +++ b/k8s/app/deployment.yaml @@ -66,6 +66,13 @@ spec: secretKeyRef: name: poimen-memory-secrets key: llm-api-key + # LLM config (in-cluster, no auth needed) + - name: LLM_ENDPOINT + value: "http://reasoning-predictor.llm-serving.svc.cluster.local/v1/chat/completions" + - name: LLM_API_BASE + value: "http://reasoning-predictor.llm-serving.svc.cluster.local/v1" + - name: LLM_MODEL + value: "reasoning" # Server config (from ConfigMap) - name: MEM_PORT value: "8080"