use std::env;
use std::time::{Duration, Instant};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use tokio::sync::RwLock;
use tokio::time::timeout;
use crate::agent::AgentContext;
use crate::config::InferenceConfig;
use crate::memory::Memory;
use crate::{OxydeError, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderType {
Local,
Cloud,
}
#[derive(Debug, Clone, Serialize)]
pub struct InferenceRequest {
pub input: String,
pub system_prompt: String,
pub memories: Vec<Memory>,
pub context: AgentContext,
pub max_tokens: usize,
pub temperature: f32,
}
#[derive(Debug, Clone, Deserialize)]
pub struct InferenceResponse {
pub text: String,
pub time_ms: u64,
pub provider_name: String,
pub tokens: usize,
}
#[derive(Debug)]
pub struct InferenceEngine {
config: InferenceConfig,
provider_type: RwLock<ProviderType>,
stats: RwLock<InferenceStats>,
}
#[derive(Debug, Default, Clone)]
pub struct InferenceStats {
pub total_requests: usize,
pub successful_requests: usize,
pub failed_requests: usize,
pub avg_latency_ms: f64,
pub avg_tokens: f64,
}
#[async_trait]
pub trait InferenceProvider {
async fn generate(&self, request: InferenceRequest) -> Result<InferenceResponse>;
}
pub struct LocalInferenceProvider {
model_path: String,
}
#[async_trait]
impl InferenceProvider for LocalInferenceProvider {
async fn generate(&self, request: InferenceRequest) -> Result<InferenceResponse> {
log::info!("Generating response with local model: {}", self.model_path);
let start_time = Instant::now();
let mut prompt = String::new();
prompt.push_str(&request.system_prompt);
prompt.push_str("\n\n");
if !request.memories.is_empty() {
prompt.push_str("Relevant context:\n");
for memory in &request.memories {
prompt.push_str(&format!("- {}\n", memory.content));
}
prompt.push_str("\n");
}
prompt.push_str(&format!("User: {}\n", request.input));
prompt.push_str("Assistant: ");
let response = format!("This is a simulated response to: {}", request.input);
let token_count = response.split_whitespace().count();
let elapsed = start_time.elapsed();
Ok(InferenceResponse {
text: response,
time_ms: elapsed.as_millis() as u64,
provider_name: "local".to_string(),
tokens: token_count,
})
}
}
pub struct CloudInferenceProvider {
api_endpoint: String,
api_key: String,
}
#[async_trait]
impl InferenceProvider for CloudInferenceProvider {
async fn generate(&self, request: InferenceRequest) -> Result<InferenceResponse> {
log::info!("Generating response with cloud API: {}", self.api_endpoint);
let start_time = Instant::now();
let system_message = serde_json::json!({
"role": "system",
"content": request.system_prompt,
});
let mut messages = vec![system_message];
if !request.memories.is_empty() {
let memories_content = request.memories.iter()
.map(|m| format!("- {}", m.content))
.collect::<Vec<_>>()
.join("\n");
let context_message = serde_json::json!({
"role": "system",
"content": format!("Relevant context:\n{}", memories_content),
});
messages.push(context_message);
}
let user_message = serde_json::json!({
"role": "user",
"content": request.input,
});
messages.push(user_message);
let client = reqwest::Client::new();
let model_name = if self.api_endpoint.contains("openai") {
"gpt-3.5-turbo"
} else {
"llama-2-7b"
};
let api_request = serde_json::json!({
"model": model_name,
"messages": messages,
"temperature": request.temperature,
"max_tokens": request.max_tokens,
});
let duration = Duration::from_millis(request.context.get("timeout_ms")
.and_then(|v| v.as_u64())
.unwrap_or(5000));
let api_response = timeout(duration, async {
client.post(&self.api_endpoint)
.header("Content-Type", "application/json")
.header("Authorization", format!("Bearer {}", self.api_key))
.json(&api_request)
.send()
.await
.map_err(|e| OxydeError::InferenceError(format!("API request failed: {}", e)))?
.json::<serde_json::Value>()
.await
.map_err(|e| OxydeError::InferenceError(format!("Failed to parse API response: {}", e)))
}).await.map_err(|_| OxydeError::InferenceError("API request timed out".to_string()))??;
let response_text = api_response["choices"][0]["message"]["content"]
.as_str()
.ok_or_else(|| OxydeError::InferenceError("Invalid API response format".to_string()))?
.to_string();
let token_count = response_text.split_whitespace().count();
let elapsed = start_time.elapsed();
Ok(InferenceResponse {
text: response_text,
time_ms: elapsed.as_millis() as u64,
provider_name: "cloud".to_string(),
tokens: token_count,
})
}
}
impl InferenceEngine {
pub fn new(config: &InferenceConfig) -> Self {
let provider_type = if config.use_local {
ProviderType::Local
} else {
ProviderType::Cloud
};
Self {
config: config.clone(),
provider_type: RwLock::new(provider_type),
stats: RwLock::new(InferenceStats::default()),
}
}
pub async fn generate_response(
&self,
input: &str,
memories: &[Memory],
context: &AgentContext,
) -> Result<String> {
let request = self.prepare_request(input, memories, context);
let provider_type = *self.provider_type.read().await;
let response = self.generate_with_provider(provider_type, request.clone()).await;
if response.is_err() && self.config.fallback_api.is_some() {
log::warn!("Primary inference provider failed, trying fallback");
let fallback_provider = match provider_type {
ProviderType::Local => ProviderType::Cloud,
ProviderType::Cloud => ProviderType::Local,
};
{
let mut stats = self.stats.write().await;
stats.total_requests += 1;
stats.failed_requests += 1;
}
return self.generate_with_provider(fallback_provider, request).await
.map(|response| response.text);
}
response.map(|response| response.text)
}
fn prepare_request(
&self,
input: &str,
memories: &[Memory],
context: &AgentContext,
) -> InferenceRequest {
let system_prompt = format!(
"You are an NPC named {} who is a {}. \
Respond in character with brief, concise answers.",
context.get("name").and_then(|v| v.as_str()).unwrap_or("Unknown"),
context.get("role").and_then(|v| v.as_str()).unwrap_or("character"),
);
InferenceRequest {
input: input.to_string(),
system_prompt,
memories: memories.to_vec(),
context: context.clone(),
max_tokens: self.config.max_tokens,
temperature: self.config.temperature,
}
}
async fn generate_with_provider(
&self,
provider_type: ProviderType,
request: InferenceRequest,
) -> Result<InferenceResponse> {
let response = match provider_type {
ProviderType::Local => {
if let Some(model_path) = &self.config.local_model_path {
let local_provider = LocalInferenceProvider {
model_path: model_path.clone(),
};
local_provider.generate(request).await
} else {
return Err(OxydeError::InferenceError(
"No local model path configured".to_string()
));
}
},
ProviderType::Cloud => {
let api_endpoint = self.config.api_endpoint.clone()
.ok_or_else(|| OxydeError::InferenceError(
"No API endpoint configured".to_string()
))?;
let api_key = self.config.api_key.clone()
.or_else(|| env::var("OXYDE_API_KEY").ok())
.ok_or_else(|| OxydeError::InferenceError(
"No API key configured. Set OXYDE_API_KEY environment variable or configure in InferenceConfig".to_string()
))?;
let cloud_provider = CloudInferenceProvider {
api_endpoint,
api_key,
};
cloud_provider.generate(request).await
}
};
if let Ok(ref resp) = response {
let mut stats = self.stats.write().await;
stats.total_requests += 1;
stats.successful_requests += 1;
let count = stats.successful_requests as f64;
stats.avg_latency_ms = (stats.avg_latency_ms * (count - 1.0) + resp.time_ms as f64) / count;
stats.avg_tokens = (stats.avg_tokens * (count - 1.0) + resp.tokens as f64) / count;
}
response
}
pub async fn switch_provider(&self, provider_type: ProviderType) {
let mut current = self.provider_type.write().await;
*current = provider_type;
log::info!("Switched to {:?} inference provider", provider_type);
}
pub async fn get_stats(&self) -> InferenceStats {
self.stats.read().await.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_inference_engine_creation() {
let config = InferenceConfig::default();
let engine = InferenceEngine::new(&config);
let provider_type = *engine.provider_type.read().await;
assert_eq!(provider_type, ProviderType::Cloud);
let stats = engine.get_stats().await;
assert_eq!(stats.total_requests, 0);
}
}