Skip to main content

oxyde/
inference.rs

1//! Inference engine for the Oxyde SDK
2//!
3//! This module provides the inference capabilities for generating NPC responses
4//! using either local models (via llm crate) or cloud API services.
5
6use 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/// Inference provider types
20#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub enum ProviderType {
22    /// Local model inference
23    Local,
24    /// Cloud API inference
25    Cloud,
26}
27
28/// Request to the inference engine
29#[derive(Debug, Clone, Serialize)]
30pub struct InferenceRequest {
31    /// Input text
32    pub input: String,
33    
34    /// System prompt
35    pub system_prompt: String,
36    
37    /// Relevant memories
38    pub memories: Vec<Memory>,
39    
40    /// Context data
41    pub context: AgentContext,
42    
43    /// Maximum tokens to generate
44    pub max_tokens: usize,
45    
46    /// Temperature
47    pub temperature: f32,
48}
49
50/// Response from the inference engine
51#[derive(Debug, Clone, Deserialize)]
52pub struct InferenceResponse {
53    /// Generated text
54    pub text: String,
55    
56    /// Time taken for inference in milliseconds
57    pub time_ms: u64,
58    
59    /// Provider name or identifier
60    pub provider_name: String,
61    
62    /// Tokens generated
63    pub tokens: usize,
64}
65
66/// Inference engine for generating NPC responses
67#[derive(Debug)]
68pub struct InferenceEngine {
69    /// Configuration for the inference engine
70    config: InferenceConfig,
71    
72    /// Current inference provider type
73    provider_type: RwLock<ProviderType>,
74    
75    /// Statistics about inference
76    stats: RwLock<InferenceStats>,
77}
78
79/// Statistics about inference operations
80#[derive(Debug, Default, Clone)]
81pub struct InferenceStats {
82    /// Total number of requests
83    pub total_requests: usize,
84    
85    /// Number of successful requests
86    pub successful_requests: usize,
87    
88    /// Number of failed requests
89    pub failed_requests: usize,
90    
91    /// Average latency in milliseconds
92    pub avg_latency_ms: f64,
93    
94    /// Average tokens generated
95    pub avg_tokens: f64,
96}
97
98/// Trait for inference providers
99#[async_trait]
100pub trait InferenceProvider {
101    /// Generate a response for the given request
102    async fn generate(&self, request: InferenceRequest) -> Result<InferenceResponse>;
103}
104
105/// Local model inference provider
106pub 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        // Simulate local model inference for now
114        // In a real implementation, this would use llm crate to load and run the model
115        
116        log::info!("Generating response with local model: {}", self.model_path);
117        
118        let start_time = Instant::now();
119        
120        // Prepare the prompt
121        let mut prompt = String::new();
122        
123        // Add system prompt
124        prompt.push_str(&request.system_prompt);
125        prompt.push_str("\n\n");
126        
127        // Add memories as context
128        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        // Add user input
137        prompt.push_str(&format!("User: {}\n", request.input));
138        prompt.push_str("Assistant: ");
139        
140        // TODO: Replace with actual local model inference
141        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
155/// Cloud API inference provider
156pub 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        // Prepare the messages for the API
169        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        // Add memories as context if available
177        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        // Add user message
192        let user_message = serde_json::json!({
193            "role": "user",
194            "content": request.input,
195        });
196        
197        messages.push(user_message);
198        
199        // Prepare the API request
200        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        // Set timeout for the request
214        let duration = Duration::from_millis(request.context.get("timeout_ms")
215            .and_then(|v| v.as_u64())
216            .unwrap_or(5000));
217        
218        // Send the request to the API
219        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        // Extract the response text
233        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        // Count tokens before moving the string
239        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    /// Create a new inference engine with the given configuration
254    ///
255    /// # Arguments
256    ///
257    /// * `config` - Inference engine configuration
258    ///
259    /// # Returns
260    ///
261    /// A new InferenceEngine instance
262    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    /// Generate a response for the given input
277    ///
278    /// # Arguments
279    ///
280    /// * `input` - User input to respond to
281    /// * `memories` - Relevant memories for context
282    /// * `context` - Additional context data
283    ///
284    /// # Returns
285    ///
286    /// The generated response text
287    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        // Try primary provider first
296        let provider_type = *self.provider_type.read().await;
297        let response = self.generate_with_provider(provider_type, request.clone()).await;
298        
299        // If primary fails and fallback is available, try fallback
300        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            // Update stats for the failed request
309            {
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    /// Prepare an inference request
323    fn prepare_request(
324        &self,
325        input: &str,
326        memories: &[Memory],
327        context: &AgentContext,
328    ) -> InferenceRequest {
329        // Create system prompt for the agent
330        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    /// Generate a response with the specified provider type
348    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        // Update stats on success
388        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            // Update moving average for latency and tokens
394            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    /// Switch to a different inference provider type
403    ///
404    /// # Arguments
405    ///
406    /// * `provider_type` - The provider type to switch to
407    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    /// Get current inference statistics
415    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}