1use std::env;
7use std::time::{Duration, Instant};
8
9use async_trait::async_trait;
10use serde::{Deserialize, Serialize};
11use tokio::sync::RwLock;
12use tokio::time::timeout;
13
14use crate::agent::AgentContext;
15use crate::config::InferenceConfig;
16use crate::memory::Memory;
17use crate::{OxydeError, Result};
18
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub enum ProviderType {
22 Local,
24 Cloud,
26}
27
28#[derive(Debug, Clone, Serialize)]
30pub struct InferenceRequest {
31 pub input: String,
33
34 pub system_prompt: String,
36
37 pub memories: Vec<Memory>,
39
40 pub context: AgentContext,
42
43 pub max_tokens: usize,
45
46 pub temperature: f32,
48}
49
50#[derive(Debug, Clone, Deserialize)]
52pub struct InferenceResponse {
53 pub text: String,
55
56 pub time_ms: u64,
58
59 pub provider_name: String,
61
62 pub tokens: usize,
64}
65
66#[derive(Debug)]
68pub struct InferenceEngine {
69 config: InferenceConfig,
71
72 provider_type: RwLock<ProviderType>,
74
75 stats: RwLock<InferenceStats>,
77}
78
79#[derive(Debug, Default, Clone)]
81pub struct InferenceStats {
82 pub total_requests: usize,
84
85 pub successful_requests: usize,
87
88 pub failed_requests: usize,
90
91 pub avg_latency_ms: f64,
93
94 pub avg_tokens: f64,
96}
97
98#[async_trait]
100pub trait InferenceProvider {
101 async fn generate(&self, request: InferenceRequest) -> Result<InferenceResponse>;
103}
104
105pub struct LocalInferenceProvider {
107 model_path: String,
108}
109
110#[async_trait]
111impl InferenceProvider for LocalInferenceProvider {
112 async fn generate(&self, request: InferenceRequest) -> Result<InferenceResponse> {
113 log::info!("Generating response with local model: {}", self.model_path);
117
118 let start_time = Instant::now();
119
120 let mut prompt = String::new();
122
123 prompt.push_str(&request.system_prompt);
125 prompt.push_str("\n\n");
126
127 if !request.memories.is_empty() {
129 prompt.push_str("Relevant context:\n");
130 for memory in &request.memories {
131 prompt.push_str(&format!("- {}\n", memory.content));
132 }
133 prompt.push_str("\n");
134 }
135
136 prompt.push_str(&format!("User: {}\n", request.input));
138 prompt.push_str("Assistant: ");
139
140 let response = format!("This is a simulated response to: {}", request.input);
142 let token_count = response.split_whitespace().count();
143
144 let elapsed = start_time.elapsed();
145
146 Ok(InferenceResponse {
147 text: response,
148 time_ms: elapsed.as_millis() as u64,
149 provider_name: "local".to_string(),
150 tokens: token_count,
151 })
152 }
153}
154
155pub struct CloudInferenceProvider {
157 api_endpoint: String,
158 api_key: String,
159}
160
161#[async_trait]
162impl InferenceProvider for CloudInferenceProvider {
163 async fn generate(&self, request: InferenceRequest) -> Result<InferenceResponse> {
164 log::info!("Generating response with cloud API: {}", self.api_endpoint);
165
166 let start_time = Instant::now();
167
168 let system_message = serde_json::json!({
170 "role": "system",
171 "content": request.system_prompt,
172 });
173
174 let mut messages = vec![system_message];
175
176 if !request.memories.is_empty() {
178 let memories_content = request.memories.iter()
179 .map(|m| format!("- {}", m.content))
180 .collect::<Vec<_>>()
181 .join("\n");
182
183 let context_message = serde_json::json!({
184 "role": "system",
185 "content": format!("Relevant context:\n{}", memories_content),
186 });
187
188 messages.push(context_message);
189 }
190
191 let user_message = serde_json::json!({
193 "role": "user",
194 "content": request.input,
195 });
196
197 messages.push(user_message);
198
199 let client = reqwest::Client::new();
201 let model_name = if self.api_endpoint.contains("openai") {
202 "gpt-3.5-turbo"
203 } else {
204 "llama-2-7b"
205 };
206 let api_request = serde_json::json!({
207 "model": model_name,
208 "messages": messages,
209 "temperature": request.temperature,
210 "max_tokens": request.max_tokens,
211 });
212
213 let duration = Duration::from_millis(request.context.get("timeout_ms")
215 .and_then(|v| v.as_u64())
216 .unwrap_or(5000));
217
218 let api_response = timeout(duration, async {
220 client.post(&self.api_endpoint)
221 .header("Content-Type", "application/json")
222 .header("Authorization", format!("Bearer {}", self.api_key))
223 .json(&api_request)
224 .send()
225 .await
226 .map_err(|e| OxydeError::InferenceError(format!("API request failed: {}", e)))?
227 .json::<serde_json::Value>()
228 .await
229 .map_err(|e| OxydeError::InferenceError(format!("Failed to parse API response: {}", e)))
230 }).await.map_err(|_| OxydeError::InferenceError("API request timed out".to_string()))??;
231
232 let response_text = api_response["choices"][0]["message"]["content"]
234 .as_str()
235 .ok_or_else(|| OxydeError::InferenceError("Invalid API response format".to_string()))?
236 .to_string();
237
238 let token_count = response_text.split_whitespace().count();
240
241 let elapsed = start_time.elapsed();
242
243 Ok(InferenceResponse {
244 text: response_text,
245 time_ms: elapsed.as_millis() as u64,
246 provider_name: "cloud".to_string(),
247 tokens: token_count,
248 })
249 }
250}
251
252impl InferenceEngine {
253 pub fn new(config: &InferenceConfig) -> Self {
263 let provider_type = if config.use_local {
264 ProviderType::Local
265 } else {
266 ProviderType::Cloud
267 };
268
269 Self {
270 config: config.clone(),
271 provider_type: RwLock::new(provider_type),
272 stats: RwLock::new(InferenceStats::default()),
273 }
274 }
275
276 pub async fn generate_response(
288 &self,
289 input: &str,
290 memories: &[Memory],
291 context: &AgentContext,
292 ) -> Result<String> {
293 let request = self.prepare_request(input, memories, context);
294
295 let provider_type = *self.provider_type.read().await;
297 let response = self.generate_with_provider(provider_type, request.clone()).await;
298
299 if response.is_err() && self.config.fallback_api.is_some() {
301 log::warn!("Primary inference provider failed, trying fallback");
302
303 let fallback_provider = match provider_type {
304 ProviderType::Local => ProviderType::Cloud,
305 ProviderType::Cloud => ProviderType::Local,
306 };
307
308 {
310 let mut stats = self.stats.write().await;
311 stats.total_requests += 1;
312 stats.failed_requests += 1;
313 }
314
315 return self.generate_with_provider(fallback_provider, request).await
316 .map(|response| response.text);
317 }
318
319 response.map(|response| response.text)
320 }
321
322 fn prepare_request(
324 &self,
325 input: &str,
326 memories: &[Memory],
327 context: &AgentContext,
328 ) -> InferenceRequest {
329 let system_prompt = format!(
331 "You are an NPC named {} who is a {}. \
332 Respond in character with brief, concise answers.",
333 context.get("name").and_then(|v| v.as_str()).unwrap_or("Unknown"),
334 context.get("role").and_then(|v| v.as_str()).unwrap_or("character"),
335 );
336
337 InferenceRequest {
338 input: input.to_string(),
339 system_prompt,
340 memories: memories.to_vec(),
341 context: context.clone(),
342 max_tokens: self.config.max_tokens,
343 temperature: self.config.temperature,
344 }
345 }
346
347 async fn generate_with_provider(
349 &self,
350 provider_type: ProviderType,
351 request: InferenceRequest,
352 ) -> Result<InferenceResponse> {
353 let response = match provider_type {
354 ProviderType::Local => {
355 if let Some(model_path) = &self.config.local_model_path {
356 let local_provider = LocalInferenceProvider {
357 model_path: model_path.clone(),
358 };
359 local_provider.generate(request).await
360 } else {
361 return Err(OxydeError::InferenceError(
362 "No local model path configured".to_string()
363 ));
364 }
365 },
366 ProviderType::Cloud => {
367 let api_endpoint = self.config.api_endpoint.clone()
368 .ok_or_else(|| OxydeError::InferenceError(
369 "No API endpoint configured".to_string()
370 ))?;
371
372 let api_key = self.config.api_key.clone()
373 .or_else(|| env::var("OXYDE_API_KEY").ok())
374 .ok_or_else(|| OxydeError::InferenceError(
375 "No API key configured. Set OXYDE_API_KEY environment variable or configure in InferenceConfig".to_string()
376 ))?;
377
378 let cloud_provider = CloudInferenceProvider {
379 api_endpoint,
380 api_key,
381 };
382
383 cloud_provider.generate(request).await
384 }
385 };
386
387 if let Ok(ref resp) = response {
389 let mut stats = self.stats.write().await;
390 stats.total_requests += 1;
391 stats.successful_requests += 1;
392
393 let count = stats.successful_requests as f64;
395 stats.avg_latency_ms = (stats.avg_latency_ms * (count - 1.0) + resp.time_ms as f64) / count;
396 stats.avg_tokens = (stats.avg_tokens * (count - 1.0) + resp.tokens as f64) / count;
397 }
398
399 response
400 }
401
402 pub async fn switch_provider(&self, provider_type: ProviderType) {
408 let mut current = self.provider_type.write().await;
409 *current = provider_type;
410
411 log::info!("Switched to {:?} inference provider", provider_type);
412 }
413
414 pub async fn get_stats(&self) -> InferenceStats {
416 self.stats.read().await.clone()
417 }
418}
419
420#[cfg(test)]
421mod tests {
422 use super::*;
423
424 #[tokio::test]
425 async fn test_inference_engine_creation() {
426 let config = InferenceConfig::default();
427 let engine = InferenceEngine::new(&config);
428
429 let provider_type = *engine.provider_type.read().await;
430 assert_eq!(provider_type, ProviderType::Cloud);
431
432 let stats = engine.get_stats().await;
433 assert_eq!(stats.total_requests, 0);
434 }
435}