1use std::sync::Arc;
8
9use async_trait::async_trait;
10use serde_json::Value;
11
12use crate::llm::{BaseLlm, LlmError, LlmRequest};
13use crate::state::State;
14
15use super::phase::Phase;
16use super::transcript::TranscriptTurn;
17
18#[derive(Debug, Clone, PartialEq, Eq)]
24pub enum ExtractionTrigger {
25 EveryTurn,
27 Interval(u32),
29 AfterToolCall,
31 OnPhaseChange,
33 OnGenerationComplete,
38}
39
40#[derive(Debug, Clone, Copy, PartialEq, Eq)]
42pub enum MergePolicy {
43 KeepKnown,
45 Overwrite,
47}
48
49pub type PromotionPredicate = Arc<dyn Fn(&State, &Value) -> bool + Send + Sync>;
51
52#[derive(Clone)]
54pub struct FieldPromotion {
55 pub field: String,
57 pub state_key: String,
59 pub merge: MergePolicy,
61 pub accept: Option<PromotionPredicate>,
63}
64
65impl FieldPromotion {
66 pub fn keep_known(field: impl Into<String>) -> Self {
68 let field = field.into();
69 Self {
70 state_key: field.clone(),
71 field,
72 merge: MergePolicy::KeepKnown,
73 accept: None,
74 }
75 }
76
77 pub fn overwrite(field: impl Into<String>) -> Self {
79 let field = field.into();
80 Self {
81 state_key: field.clone(),
82 field,
83 merge: MergePolicy::Overwrite,
84 accept: None,
85 }
86 }
87
88 pub fn true_only(field: impl Into<String>) -> Self {
90 Self::overwrite(field).accept_when(|_, value| value.as_bool() == Some(true))
91 }
92
93 pub fn non_empty(field: impl Into<String>) -> Self {
95 Self::overwrite(field)
96 .accept_when(|_, value| value.as_str().is_some_and(|s| !s.trim().is_empty()))
97 }
98
99 pub fn to(mut self, state_key: impl Into<String>) -> Self {
101 self.state_key = state_key.into();
102 self
103 }
104
105 pub fn accept_when(
110 mut self,
111 predicate: impl Fn(&State, &Value) -> bool + Send + Sync + 'static,
112 ) -> Self {
113 self.accept = Some(Arc::new(predicate));
114 self
115 }
116
117 pub fn and_accept_when(
119 mut self,
120 predicate: impl Fn(&State, &Value) -> bool + Send + Sync + 'static,
121 ) -> Self {
122 let previous = self.accept.take();
123 self.accept = Some(Arc::new(move |state, value| {
124 previous.as_ref().is_none_or(|accept| accept(state, value)) && predicate(state, value)
125 }));
126 self
127 }
128
129 pub fn after_presented(self, concept: impl Into<String>) -> Self {
131 let concept = concept.into();
132 self.and_accept_when(move |state, _| Phase::is_presented(state, &concept))
133 }
134}
135
136fn strip_code_fences(text: &str) -> &str {
140 let trimmed = text.trim();
141 if let Some(rest) = trimmed.strip_prefix("```") {
142 let rest = rest.trim_start_matches(|c: char| c != '\n');
144 let rest = rest.strip_prefix('\n').unwrap_or(rest);
145 let rest = rest.trim_end();
147 rest.strip_suffix("```").unwrap_or(rest).trim()
148 } else {
149 trimmed
150 }
151}
152
153#[async_trait]
159pub trait TurnExtractor: Send + Sync {
160 fn name(&self) -> &str;
162
163 fn window_size(&self) -> usize;
165
166 fn should_extract(&self, window: &[TranscriptTurn]) -> bool {
174 let _ = window;
175 true
176 }
177
178 fn trigger(&self) -> ExtractionTrigger {
182 ExtractionTrigger::EveryTurn
183 }
184
185 fn promotion_rules(&self) -> &[FieldPromotion] {
192 &[]
193 }
194
195 async fn extract(&self, window: &[TranscriptTurn]) -> Result<Value, LlmError>;
197
198 async fn extract_with_state(
204 &self,
205 window: &[TranscriptTurn],
206 state: &State,
207 ) -> Result<Value, LlmError> {
208 let _ = state;
209 self.extract(window).await
210 }
211
212 fn on_complete(&self) -> Option<OnComplete> {
216 None
217 }
218}
219
220#[derive(Clone)]
222pub struct OnComplete {
223 pub agent: Arc<dyn crate::text::TextAgent>,
225 pub mode: crate::orchestration::AgentMode,
227}
228
229pub struct LlmExtractor {
232 name: String,
233 llm: Arc<dyn BaseLlm>,
234 prompt: String,
235 window_size: usize,
236 schema: Option<Value>,
237 schema_str: Option<String>,
239 min_words: usize,
241 trigger: ExtractionTrigger,
243 promotion_rules: Vec<FieldPromotion>,
246}
247
248impl LlmExtractor {
249 pub fn new(
256 name: impl Into<String>,
257 llm: Arc<dyn BaseLlm>,
258 prompt: impl Into<String>,
259 window_size: usize,
260 ) -> Self {
261 Self {
262 name: name.into(),
263 llm,
264 prompt: prompt.into(),
265 window_size,
266 schema: None,
267 schema_str: None,
268 min_words: 0,
269 trigger: ExtractionTrigger::EveryTurn,
270 promotion_rules: Vec::new(),
271 }
272 }
273
274 pub fn with_min_words(mut self, n: usize) -> Self {
279 self.min_words = n;
280 self
281 }
282
283 pub fn with_schema(mut self, schema: Value) -> Self {
288 self.schema_str = serde_json::to_string_pretty(&schema).ok();
289 self.schema = Some(schema);
290 self
291 }
292
293 pub fn with_trigger(mut self, trigger: ExtractionTrigger) -> Self {
295 self.trigger = trigger;
296 self
297 }
298
299 pub fn with_promotions(mut self, rules: Vec<FieldPromotion>) -> Self {
304 self.promotion_rules = rules;
305 self
306 }
307
308 fn format_transcript(window: &[TranscriptTurn]) -> String {
310 let mut out = String::new();
311 for turn in window {
312 if !turn.user.is_empty() {
313 out.push_str("User: ");
314 out.push_str(turn.user.trim());
315 out.push('\n');
316 }
317 if !turn.model.is_empty() {
318 out.push_str("Assistant: ");
319 out.push_str(turn.model.trim());
320 out.push('\n');
321 }
322 out.push('\n');
323 }
324 out
325 }
326}
327
328#[async_trait]
329impl TurnExtractor for LlmExtractor {
330 fn name(&self) -> &str {
331 &self.name
332 }
333
334 fn window_size(&self) -> usize {
335 self.window_size
336 }
337
338 fn should_extract(&self, window: &[TranscriptTurn]) -> bool {
339 if self.min_words == 0 {
340 return true;
341 }
342 window
344 .iter()
345 .rev()
346 .find(|t| !t.user.is_empty())
347 .is_some_and(|t| t.user.split_whitespace().count() >= self.min_words)
348 }
349
350 fn trigger(&self) -> ExtractionTrigger {
351 self.trigger.clone()
352 }
353
354 fn promotion_rules(&self) -> &[FieldPromotion] {
355 &self.promotion_rules
356 }
357
358 async fn extract(&self, window: &[TranscriptTurn]) -> Result<Value, LlmError> {
359 let transcript = Self::format_transcript(window);
360
361 let mut request = LlmRequest::from_text(format!(
362 "Transcript:\n{transcript}\nExtract the requested information."
363 ));
364 request.system_instruction = Some(self.prompt.clone());
365
366 if let Some(ref schema) = self.schema {
370 request.response_mime_type = Some("application/json".to_string());
371 request.response_json_schema = Some(schema.clone());
372 } else {
373 request.response_mime_type = Some("application/json".to_string());
374 }
375
376 let response = self.llm.generate(request).await?;
377 let text = response.text();
378
379 let cleaned = strip_code_fences(&text);
381
382 serde_json::from_str(cleaned).map_err(|e| {
383 LlmError::Other(format!(
384 "Failed to parse extraction result as JSON: {e}. Raw: {text}"
385 ))
386 })
387 }
388}
389
390#[cfg(test)]
391mod tests {
392 use super::*;
393 use crate::llm::LlmResponse;
394 use gemini_genai_rs::prelude::{Content, Part, Role};
395 use std::time::Instant;
396
397 struct MockLlm {
398 response: String,
399 }
400
401 #[async_trait]
402 impl BaseLlm for MockLlm {
403 fn model_id(&self) -> &str {
404 "mock"
405 }
406 async fn generate(&self, _request: LlmRequest) -> Result<LlmResponse, LlmError> {
407 Ok(LlmResponse {
408 content: Content {
409 role: Some(Role::Model),
410 parts: vec![Part::Text {
411 text: self.response.clone(),
412 }],
413 },
414 finish_reason: Some("STOP".into()),
415 usage: None,
416 })
417 }
418 }
419
420 fn make_turns(pairs: &[(&str, &str)]) -> Vec<TranscriptTurn> {
421 pairs
422 .iter()
423 .enumerate()
424 .map(|(i, (user, model))| TranscriptTurn {
425 turn_number: i as u32,
426 user: user.to_string(),
427 model: model.to_string(),
428 tool_calls: Vec::new(),
429 timestamp: Instant::now(),
430 })
431 .collect()
432 }
433
434 #[tokio::test]
435 async fn llm_extractor_produces_json() {
436 let llm = Arc::new(MockLlm {
437 response: r#"{"phase": "ordering", "items": ["pizza"]}"#.to_string(),
438 });
439
440 let extractor = LlmExtractor::new("OrderState", llm, "Extract order state", 3);
441
442 let turns = make_turns(&[
443 ("I'd like a pizza", "Great! What size?"),
444 ("Large please", "Coming right up!"),
445 ]);
446
447 let result = extractor.extract(&turns).await.unwrap();
448 assert_eq!(result["phase"], "ordering");
449 assert_eq!(result["items"][0], "pizza");
450 }
451
452 #[tokio::test]
453 async fn llm_extractor_with_schema() {
454 let llm = Arc::new(MockLlm {
455 response: r#"{"sentiment": "positive", "score": 0.9}"#.to_string(),
456 });
457
458 let schema = serde_json::json!({
459 "type": "object",
460 "properties": {
461 "sentiment": {"type": "string", "enum": ["positive", "neutral", "negative"]},
462 "score": {"type": "number"}
463 }
464 });
465
466 let extractor =
467 LlmExtractor::new("Sentiment", llm, "Rate sentiment", 1).with_schema(schema);
468
469 let turns = make_turns(&[("This is great!", "Glad you think so!")]);
470 let result = extractor.extract(&turns).await.unwrap();
471 assert_eq!(result["sentiment"], "positive");
472 }
473
474 #[tokio::test]
475 async fn llm_extractor_invalid_json_returns_error() {
476 let llm = Arc::new(MockLlm {
477 response: "not json at all".to_string(),
478 });
479
480 let extractor = LlmExtractor::new("Bad", llm, "Extract", 1);
481 let turns = make_turns(&[("hi", "hello")]);
482 let result = extractor.extract(&turns).await;
483 assert!(result.is_err());
484 }
485
486 #[test]
487 fn format_transcript_readable() {
488 let turns = make_turns(&[("Hello", "Hi there!"), ("How are you?", "I'm doing well")]);
489 let formatted = LlmExtractor::format_transcript(&turns);
490 assert!(formatted.contains("User: Hello"));
491 assert!(formatted.contains("Assistant: Hi there!"));
492 assert!(formatted.contains("User: How are you?"));
493 }
494
495 #[tokio::test]
496 async fn llm_extractor_handles_markdown_fenced_json() {
497 let llm = Arc::new(MockLlm {
498 response: "```json\n{\"status\": \"ok\"}\n```".to_string(),
499 });
500
501 let extractor = LlmExtractor::new("Fenced", llm, "Extract", 1);
502 let turns = make_turns(&[("test", "reply")]);
503 let result = extractor.extract(&turns).await.unwrap();
504 assert_eq!(result["status"], "ok");
505 }
506
507 #[test]
508 fn strip_code_fences_variants() {
509 assert_eq!(super::strip_code_fences("```json\n{}\n```"), "{}");
510 assert_eq!(super::strip_code_fences("```\n{}\n```"), "{}");
511 assert_eq!(
512 super::strip_code_fences(" ```json\n{\"a\":1}\n``` "),
513 "{\"a\":1}"
514 );
515 assert_eq!(
516 super::strip_code_fences("{\"bare\":true}"),
517 "{\"bare\":true}"
518 );
519 }
520
521 #[test]
522 fn extractor_name_and_window_size() {
523 let llm = Arc::new(MockLlm {
524 response: "{}".to_string(),
525 });
526 let ext = LlmExtractor::new("TestExtractor", llm, "test", 5);
527 assert_eq!(ext.name(), "TestExtractor");
528 assert_eq!(ext.window_size(), 5);
529 }
530
531 #[test]
532 fn extractor_default_trigger_is_every_turn() {
533 let llm = Arc::new(MockLlm {
534 response: "{}".to_string(),
535 });
536 let ext = LlmExtractor::new("Test", llm, "test", 5);
537 assert_eq!(ext.trigger(), ExtractionTrigger::EveryTurn);
538 }
539
540 #[test]
541 fn extractor_with_trigger() {
542 let llm = Arc::new(MockLlm {
543 response: "{}".to_string(),
544 });
545 let ext = LlmExtractor::new("Test", llm, "test", 5)
546 .with_trigger(ExtractionTrigger::AfterToolCall);
547 assert_eq!(ext.trigger(), ExtractionTrigger::AfterToolCall);
548 }
549}