1use std::sync::Arc;
7use tokio::sync::RwLock;
8
9use super::config::BehaviorModelConfig;
10use super::types::LlmGenerationRequest;
11use mockforge_foundation::Result;
12
13pub struct LlmClient {
15 rag_engine: Arc<RwLock<Option<Box<dyn LlmProvider>>>>,
17 config: BehaviorModelConfig,
19}
20
21impl LlmClient {
22 pub fn new(config: BehaviorModelConfig) -> Self {
24 Self {
25 rag_engine: Arc::new(RwLock::new(None)),
26 config,
27 }
28 }
29
30 async fn ensure_initialized(&self) -> Result<()> {
32 let mut engine = self.rag_engine.write().await;
33
34 if engine.is_none() {
35 let provider = self.create_provider()?;
37 *engine = Some(provider);
38 }
39
40 Ok(())
41 }
42
43 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 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 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 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 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 match serde_json::from_str::<serde_json::Value>(&response_text) {
99 Ok(json) => Ok(json),
100 Err(_) => {
101 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 Ok(serde_json::json!({
113 "response": response_text,
114 "note": "Response was not valid JSON, wrapped in object"
115 }))
116 }
117 }
118 }
119
120 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 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 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 let json_value = match serde_json::from_str::<serde_json::Value>(&response_text) {
156 Ok(json) => json,
157 Err(_) => {
158 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 pub fn config(&self) -> &BehaviorModelConfig {
190 &self.config
191 }
192}
193
194#[derive(Debug, Clone)]
196struct ChatMessage {
197 role: String,
198 content: String,
199}
200
201#[derive(Debug, Clone, Default)]
203pub struct LlmUsage {
204 pub prompt_tokens: u64,
206 pub completion_tokens: u64,
208 pub total_tokens: u64,
210}
211
212impl LlmUsage {
213 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#[async_trait::async_trait]
225trait LlmProvider: Send + Sync {
226 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 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 let response = self.generate_chat(messages, temperature, max_tokens, seed).await?;
245 let estimated_tokens = (response.len() as f64 / 4.0) as u64;
247 Ok((response, LlmUsage::new(estimated_tokens, estimated_tokens)))
248 }
249}
250
251struct 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 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 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 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 let estimated = (content.len() as f64 / 4.0) as u64;
399 LlmUsage::new(estimated, estimated)
400 };
401
402 Ok((content, usage))
403 }
404}
405
406struct 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 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 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
488struct 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 _seed: Option<i64>,
527 ) -> Result<String> {
528 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 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
597struct 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 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 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}