Skip to main content

mockforge_intelligence/intelligent_behavior/
llm_client.rs

1//! LLM client wrapper for intelligent behavior
2//!
3//! This module provides a simplified interface to the RAG engine for
4//! intelligent mock behavior generation.
5
6use std::sync::Arc;
7use tokio::sync::RwLock;
8
9use super::config::BehaviorModelConfig;
10use super::types::LlmGenerationRequest;
11use mockforge_foundation::Result;
12
13/// LLM client for generating intelligent responses
14pub struct LlmClient {
15    /// RAG engine (lazily initialized)
16    rag_engine: Arc<RwLock<Option<Box<dyn LlmProvider>>>>,
17    /// Configuration
18    config: BehaviorModelConfig,
19}
20
21impl LlmClient {
22    /// Create a new LLM client
23    pub fn new(config: BehaviorModelConfig) -> Self {
24        Self {
25            rag_engine: Arc::new(RwLock::new(None)),
26            config,
27        }
28    }
29
30    /// Initialize the RAG engine (lazy initialization)
31    async fn ensure_initialized(&self) -> Result<()> {
32        let mut engine = self.rag_engine.write().await;
33
34        if engine.is_none() {
35            // Create provider based on configuration
36            let provider = self.create_provider()?;
37            *engine = Some(provider);
38        }
39
40        Ok(())
41    }
42
43    /// Create LLM provider based on configuration
44    fn create_provider(&self) -> Result<Box<dyn LlmProvider>> {
45        match self.config.llm_provider.to_lowercase().as_str() {
46            "openai" => Ok(Box::new(OpenAIProvider::new(&self.config)?)),
47            "anthropic" => Ok(Box::new(AnthropicProvider::new(&self.config)?)),
48            "ollama" => Ok(Box::new(OllamaProvider::new(&self.config)?)),
49            "openai-compatible" => Ok(Box::new(OpenAICompatibleProvider::new(&self.config)?)),
50            _ => Err(mockforge_foundation::Error::internal(format!(
51                "Unsupported LLM provider: {}",
52                self.config.llm_provider
53            ))),
54        }
55    }
56
57    /// Resolve the effective sampling seed (#852): per-request override,
58    /// then the behavior-model config, then the MOCKFORGE_AI_SEED env var.
59    fn resolve_seed(&self, request: &LlmGenerationRequest) -> Option<i64> {
60        request
61            .seed
62            .or(self.config.seed)
63            .or_else(|| std::env::var("MOCKFORGE_AI_SEED").ok().and_then(|s| s.trim().parse().ok()))
64    }
65
66    /// Generate a response from a prompt
67    pub async fn generate(&self, request: &LlmGenerationRequest) -> Result<serde_json::Value> {
68        self.ensure_initialized().await?;
69
70        let engine = self.rag_engine.read().await;
71        let provider = engine
72            .as_ref()
73            .ok_or_else(|| mockforge_foundation::Error::internal("LLM provider not initialized"))?;
74
75        // Build messages
76        let messages = vec![
77            ChatMessage {
78                role: "system".to_string(),
79                content: request.system_prompt.clone(),
80            },
81            ChatMessage {
82                role: "user".to_string(),
83                content: request.user_prompt.clone(),
84            },
85        ];
86
87        // Generate response
88        let response_text = provider
89            .generate_chat(
90                messages,
91                request.temperature,
92                request.max_tokens,
93                self.resolve_seed(request),
94            )
95            .await?;
96
97        // Try to parse as JSON
98        match serde_json::from_str::<serde_json::Value>(&response_text) {
99            Ok(json) => Ok(json),
100            Err(_) => {
101                // Try to extract JSON from response
102                if let Some(start) = response_text.find('{') {
103                    if let Some(end) = response_text.rfind('}') {
104                        let json_str = &response_text[start..=end];
105                        if let Ok(json) = serde_json::from_str::<serde_json::Value>(json_str) {
106                            return Ok(json);
107                        }
108                    }
109                }
110
111                // Fallback: wrap in object
112                Ok(serde_json::json!({
113                    "response": response_text,
114                    "note": "Response was not valid JSON, wrapped in object"
115                }))
116            }
117        }
118    }
119
120    /// Generate a response and return usage information
121    pub async fn generate_with_usage(
122        &self,
123        request: &LlmGenerationRequest,
124    ) -> Result<(serde_json::Value, LlmUsage)> {
125        self.ensure_initialized().await?;
126
127        let engine = self.rag_engine.read().await;
128        let provider = engine
129            .as_ref()
130            .ok_or_else(|| mockforge_foundation::Error::internal("LLM provider not initialized"))?;
131
132        // Build messages
133        let messages = vec![
134            ChatMessage {
135                role: "system".to_string(),
136                content: request.system_prompt.clone(),
137            },
138            ChatMessage {
139                role: "user".to_string(),
140                content: request.user_prompt.clone(),
141            },
142        ];
143
144        // Generate response with usage tracking
145        let (response_text, usage) = provider
146            .generate_chat_with_usage(
147                messages,
148                request.temperature,
149                request.max_tokens,
150                self.resolve_seed(request),
151            )
152            .await?;
153
154        // Try to parse as JSON
155        let json_value = match serde_json::from_str::<serde_json::Value>(&response_text) {
156            Ok(json) => json,
157            Err(_) => {
158                // Try to extract JSON from response
159                if let Some(start) = response_text.find('{') {
160                    if let Some(end) = response_text.rfind('}') {
161                        let json_str = &response_text[start..=end];
162                        if let Ok(json) = serde_json::from_str::<serde_json::Value>(json_str) {
163                            json
164                        } else {
165                            serde_json::json!({
166                                "response": response_text,
167                                "note": "Response was not valid JSON, wrapped in object"
168                            })
169                        }
170                    } else {
171                        serde_json::json!({
172                            "response": response_text,
173                            "note": "Response was not valid JSON, wrapped in object"
174                        })
175                    }
176                } else {
177                    serde_json::json!({
178                        "response": response_text,
179                        "note": "Response was not valid JSON, wrapped in object"
180                    })
181                }
182            }
183        };
184
185        Ok((json_value, usage))
186    }
187
188    /// Get configuration
189    pub fn config(&self) -> &BehaviorModelConfig {
190        &self.config
191    }
192}
193
194/// Chat message for LLM
195#[derive(Debug, Clone)]
196struct ChatMessage {
197    role: String,
198    content: String,
199}
200
201/// LLM usage information
202#[derive(Debug, Clone, Default)]
203pub struct LlmUsage {
204    /// Prompt tokens used
205    pub prompt_tokens: u64,
206    /// Completion tokens used
207    pub completion_tokens: u64,
208    /// Total tokens used
209    pub total_tokens: u64,
210}
211
212impl LlmUsage {
213    /// Create new usage info
214    pub fn new(prompt_tokens: u64, completion_tokens: u64) -> Self {
215        Self {
216            prompt_tokens,
217            completion_tokens,
218            total_tokens: prompt_tokens + completion_tokens,
219        }
220    }
221}
222
223/// LLM provider trait
224#[async_trait::async_trait]
225trait LlmProvider: Send + Sync {
226    /// Generate chat completion
227    async fn generate_chat(
228        &self,
229        messages: Vec<ChatMessage>,
230        temperature: f64,
231        max_tokens: usize,
232        seed: Option<i64>,
233    ) -> Result<String>;
234
235    /// Generate chat completion with usage tracking
236    async fn generate_chat_with_usage(
237        &self,
238        messages: Vec<ChatMessage>,
239        temperature: f64,
240        max_tokens: usize,
241        seed: Option<i64>,
242    ) -> Result<(String, LlmUsage)> {
243        // Default implementation: call generate_chat and estimate tokens
244        let response = self.generate_chat(messages, temperature, max_tokens, seed).await?;
245        // Rough estimation: ~4 characters per token
246        let estimated_tokens = (response.len() as f64 / 4.0) as u64;
247        Ok((response, LlmUsage::new(estimated_tokens, estimated_tokens)))
248    }
249}
250
251/// OpenAI provider implementation
252struct OpenAIProvider {
253    client: reqwest::Client,
254    api_key: String,
255    model: String,
256    endpoint: String,
257}
258
259impl OpenAIProvider {
260    fn new(config: &BehaviorModelConfig) -> Result<Self> {
261        let api_key = config
262            .api_key
263            .clone()
264            .or_else(|| std::env::var("OPENAI_API_KEY").ok())
265            .ok_or_else(|| mockforge_foundation::Error::internal("OpenAI API key not found"))?;
266
267        let endpoint = config
268            .api_endpoint
269            .clone()
270            .unwrap_or_else(|| "https://api.openai.com/v1/chat/completions".to_string());
271
272        Ok(Self {
273            client: reqwest::Client::new(),
274            api_key,
275            model: config.model.clone(),
276            endpoint,
277        })
278    }
279}
280
281#[async_trait::async_trait]
282impl LlmProvider for OpenAIProvider {
283    async fn generate_chat(
284        &self,
285        messages: Vec<ChatMessage>,
286        temperature: f64,
287        max_tokens: usize,
288        seed: Option<i64>,
289    ) -> Result<String> {
290        let mut request_body = serde_json::json!({
291            "model": self.model,
292            "messages": messages.iter().map(|m| {
293                serde_json::json!({
294                    "role": m.role,
295                    "content": m.content
296                })
297            }).collect::<Vec<_>>(),
298            "temperature": temperature,
299            "max_tokens": max_tokens,
300        });
301        if let Some(seed) = seed {
302            request_body["seed"] = serde_json::json!(seed);
303        }
304
305        let response = self
306            .client
307            .post(&self.endpoint)
308            .header("Authorization", format!("Bearer {}", self.api_key))
309            .header("Content-Type", "application/json")
310            .json(&request_body)
311            .send()
312            .await
313            .map_err(|e| {
314                mockforge_foundation::Error::internal(format!("OpenAI API request failed: {}", e))
315            })?;
316
317        if !response.status().is_success() {
318            let error_text = response.text().await.unwrap_or_default();
319            return Err(mockforge_foundation::Error::internal(format!(
320                "OpenAI API error: {}",
321                error_text
322            )));
323        }
324
325        let response_json: serde_json::Value = response.json().await.map_err(|e| {
326            mockforge_foundation::Error::internal(format!("Failed to parse OpenAI response: {}", e))
327        })?;
328
329        // Extract content from response
330        let content = response_json["choices"][0]["message"]["content"]
331            .as_str()
332            .ok_or_else(|| mockforge_foundation::Error::internal("Invalid OpenAI response format"))?
333            .to_string();
334
335        Ok(content)
336    }
337
338    async fn generate_chat_with_usage(
339        &self,
340        messages: Vec<ChatMessage>,
341        temperature: f64,
342        max_tokens: usize,
343        seed: Option<i64>,
344    ) -> Result<(String, LlmUsage)> {
345        let mut request_body = serde_json::json!({
346            "model": self.model,
347            "messages": messages.iter().map(|m| {
348                serde_json::json!({
349                    "role": m.role,
350                    "content": m.content
351                })
352            }).collect::<Vec<_>>(),
353            "temperature": temperature,
354            "max_tokens": max_tokens,
355        });
356        if let Some(seed) = seed {
357            request_body["seed"] = serde_json::json!(seed);
358        }
359
360        let response = self
361            .client
362            .post(&self.endpoint)
363            .header("Authorization", format!("Bearer {}", self.api_key))
364            .header("Content-Type", "application/json")
365            .json(&request_body)
366            .send()
367            .await
368            .map_err(|e| {
369                mockforge_foundation::Error::internal(format!("OpenAI API request failed: {}", e))
370            })?;
371
372        if !response.status().is_success() {
373            let error_text = response.text().await.unwrap_or_default();
374            return Err(mockforge_foundation::Error::internal(format!(
375                "OpenAI API error: {}",
376                error_text
377            )));
378        }
379
380        let response_json: serde_json::Value = response.json().await.map_err(|e| {
381            mockforge_foundation::Error::internal(format!("Failed to parse OpenAI response: {}", e))
382        })?;
383
384        // Extract content from response
385        let content = response_json["choices"][0]["message"]["content"]
386            .as_str()
387            .ok_or_else(|| mockforge_foundation::Error::internal("Invalid OpenAI response format"))?
388            .to_string();
389
390        // Extract usage information
391        let usage = if let Some(usage_obj) = response_json.get("usage") {
392            LlmUsage::new(
393                usage_obj["prompt_tokens"].as_u64().unwrap_or(0),
394                usage_obj["completion_tokens"].as_u64().unwrap_or(0),
395            )
396        } else {
397            // Fallback: estimate tokens
398            let estimated = (content.len() as f64 / 4.0) as u64;
399            LlmUsage::new(estimated, estimated)
400        };
401
402        Ok((content, usage))
403    }
404}
405
406/// Ollama provider implementation
407struct OllamaProvider {
408    client: reqwest::Client,
409    model: String,
410    endpoint: String,
411}
412
413impl OllamaProvider {
414    fn new(config: &BehaviorModelConfig) -> Result<Self> {
415        let endpoint = config
416            .api_endpoint
417            .clone()
418            .unwrap_or_else(|| "http://localhost:11434/api/chat".to_string());
419
420        Ok(Self {
421            client: reqwest::Client::new(),
422            model: config.model.clone(),
423            endpoint,
424        })
425    }
426}
427
428#[async_trait::async_trait]
429impl LlmProvider for OllamaProvider {
430    async fn generate_chat(
431        &self,
432        messages: Vec<ChatMessage>,
433        temperature: f64,
434        max_tokens: usize,
435        seed: Option<i64>,
436    ) -> Result<String> {
437        let mut request_body = serde_json::json!({
438            "model": self.model,
439            "messages": messages.iter().map(|m| {
440                serde_json::json!({
441                    "role": m.role,
442                    "content": m.content
443                })
444            }).collect::<Vec<_>>(),
445            "options": {
446                "temperature": temperature,
447                "num_predict": max_tokens,
448            },
449            "stream": false,
450        });
451        // Ollama takes the seed inside `options` (#852).
452        if let Some(seed) = seed {
453            request_body["options"]["seed"] = serde_json::json!(seed);
454        }
455        let response = self
456            .client
457            .post(&self.endpoint)
458            .header("Content-Type", "application/json")
459            .json(&request_body)
460            .send()
461            .await
462            .map_err(|e| {
463                mockforge_foundation::Error::internal(format!("Ollama API request failed: {}", e))
464            })?;
465
466        if !response.status().is_success() {
467            let error_text = response.text().await.unwrap_or_default();
468            return Err(mockforge_foundation::Error::internal(format!(
469                "Ollama API error: {}",
470                error_text
471            )));
472        }
473
474        let response_json: serde_json::Value = response.json().await.map_err(|e| {
475            mockforge_foundation::Error::internal(format!("Failed to parse Ollama response: {}", e))
476        })?;
477
478        // Extract content from response
479        let content = response_json["message"]["content"]
480            .as_str()
481            .ok_or_else(|| mockforge_foundation::Error::internal("Invalid Ollama response format"))?
482            .to_string();
483
484        Ok(content)
485    }
486}
487
488/// Anthropic provider implementation
489struct AnthropicProvider {
490    client: reqwest::Client,
491    api_key: String,
492    model: String,
493    endpoint: String,
494}
495
496impl AnthropicProvider {
497    fn new(config: &BehaviorModelConfig) -> Result<Self> {
498        let api_key = config
499            .api_key
500            .clone()
501            .or_else(|| std::env::var("ANTHROPIC_API_KEY").ok())
502            .ok_or_else(|| mockforge_foundation::Error::internal("Anthropic API key not found"))?;
503
504        let endpoint = config
505            .api_endpoint
506            .clone()
507            .unwrap_or_else(|| "https://api.anthropic.com/v1/messages".to_string());
508
509        Ok(Self {
510            client: reqwest::Client::new(),
511            api_key,
512            model: config.model.clone(),
513            endpoint,
514        })
515    }
516}
517
518#[async_trait::async_trait]
519impl LlmProvider for AnthropicProvider {
520    async fn generate_chat(
521        &self,
522        messages: Vec<ChatMessage>,
523        temperature: f64,
524        max_tokens: usize,
525        // Anthropic's Messages API has no sampling-seed parameter; ignored.
526        _seed: Option<i64>,
527    ) -> Result<String> {
528        // Separate system message from other messages
529        let system_message =
530            messages.iter().find(|m| m.role == "system").map(|m| m.content.clone());
531
532        let chat_messages: Vec<_> = messages
533            .iter()
534            .filter(|m| m.role != "system")
535            .map(|m| {
536                serde_json::json!({
537                    "role": m.role,
538                    "content": m.content
539                })
540            })
541            .collect();
542
543        let mut request_body = serde_json::json!({
544            "model": self.model,
545            "messages": chat_messages,
546            "temperature": temperature,
547            "max_tokens": max_tokens,
548        });
549
550        if let Some(system) = system_message {
551            request_body["system"] = serde_json::Value::String(system);
552        }
553
554        let response = self
555            .client
556            .post(&self.endpoint)
557            .header("x-api-key", &self.api_key)
558            .header("anthropic-version", "2023-06-01")
559            .header("Content-Type", "application/json")
560            .json(&request_body)
561            .send()
562            .await
563            .map_err(|e| {
564                mockforge_foundation::Error::internal(format!(
565                    "Anthropic API request failed: {}",
566                    e
567                ))
568            })?;
569
570        if !response.status().is_success() {
571            let error_text = response.text().await.unwrap_or_default();
572            return Err(mockforge_foundation::Error::internal(format!(
573                "Anthropic API error: {}",
574                error_text
575            )));
576        }
577
578        let response_json: serde_json::Value = response.json().await.map_err(|e| {
579            mockforge_foundation::Error::internal(format!(
580                "Failed to parse Anthropic response: {}",
581                e
582            ))
583        })?;
584
585        // Extract content from response
586        let content = response_json["content"][0]["text"]
587            .as_str()
588            .ok_or_else(|| {
589                mockforge_foundation::Error::internal("Invalid Anthropic response format")
590            })?
591            .to_string();
592
593        Ok(content)
594    }
595}
596
597/// OpenAI-compatible provider (generic)
598struct OpenAICompatibleProvider {
599    client: reqwest::Client,
600    api_key: Option<String>,
601    model: String,
602    endpoint: String,
603}
604
605impl OpenAICompatibleProvider {
606    fn new(config: &BehaviorModelConfig) -> Result<Self> {
607        let endpoint = config.api_endpoint.clone().ok_or_else(|| {
608            mockforge_foundation::Error::internal(
609                "API endpoint required for OpenAI-compatible provider",
610            )
611        })?;
612
613        Ok(Self {
614            client: reqwest::Client::new(),
615            api_key: config.api_key.clone(),
616            model: config.model.clone(),
617            endpoint,
618        })
619    }
620}
621
622#[async_trait::async_trait]
623impl LlmProvider for OpenAICompatibleProvider {
624    async fn generate_chat(
625        &self,
626        messages: Vec<ChatMessage>,
627        temperature: f64,
628        max_tokens: usize,
629        seed: Option<i64>,
630    ) -> Result<String> {
631        let mut request_body = serde_json::json!({
632            "model": self.model,
633            "messages": messages.iter().map(|m| {
634                serde_json::json!({
635                    "role": m.role,
636                    "content": m.content
637                })
638            }).collect::<Vec<_>>(),
639            "temperature": temperature,
640            "max_tokens": max_tokens,
641        });
642        // Most OpenAI-compatible servers honour the OpenAI `seed` field
643        // (#852); those that don't ignore it.
644        if let Some(seed) = seed {
645            request_body["seed"] = serde_json::json!(seed);
646        }
647        let mut request =
648            self.client.post(&self.endpoint).header("Content-Type", "application/json");
649
650        if let Some(api_key) = &self.api_key {
651            request = request.header("Authorization", format!("Bearer {}", api_key));
652        }
653
654        let response = request.json(&request_body).send().await.map_err(|e| {
655            mockforge_foundation::Error::internal(format!("API request failed: {}", e))
656        })?;
657
658        if !response.status().is_success() {
659            let error_text = response.text().await.unwrap_or_default();
660            return Err(mockforge_foundation::Error::internal(format!(
661                "API error: {}",
662                error_text
663            )));
664        }
665
666        let response_json: serde_json::Value = response.json().await.map_err(|e| {
667            mockforge_foundation::Error::internal(format!("Failed to parse API response: {}", e))
668        })?;
669
670        // Extract content (try both OpenAI and Ollama formats)
671        let content = response_json["choices"][0]["message"]["content"]
672            .as_str()
673            .or_else(|| response_json["message"]["content"].as_str())
674            .ok_or_else(|| mockforge_foundation::Error::internal("Invalid API response format"))?
675            .to_string();
676
677        Ok(content)
678    }
679}
680
681#[cfg(test)]
682mod tests {
683    use super::*;
684
685    #[test]
686    fn test_llm_client_creation() {
687        let config = BehaviorModelConfig::default();
688        let client = LlmClient::new(config);
689        assert_eq!(client.config().llm_provider, "openai");
690    }
691}