use crate::agent::{Agent, ExecutionStats};
use crate::error::CompletionError;
use crate::message::Message;
use crate::provider::MetricsCompletionModel;
use crate::providers::{
provider_type::ProviderType, AnthropicCompletionModelAdapter, CerebrasCompletionModelAdapter,
LmStudioCompletionModelAdapter, MinimaxCompletionModelAdapter, MlxLmCompletionModelAdapter,
OllamaCompletionModelAdapter, OpenAICompletionModelAdapter, OpenRouterCompletionModelAdapter,
ZaiCompletionModelAdapter,
};
pub enum AgentWrapper {
OpenRouter(Agent<MetricsCompletionModel<OpenRouterCompletionModelAdapter>>),
OpenAI(Agent<MetricsCompletionModel<OpenAICompletionModelAdapter>>),
Anthropic(Agent<MetricsCompletionModel<AnthropicCompletionModelAdapter>>),
Minimax(Agent<MetricsCompletionModel<MinimaxCompletionModelAdapter>>),
Cerebras(Agent<MetricsCompletionModel<CerebrasCompletionModelAdapter>>),
Ollama(Agent<MetricsCompletionModel<OllamaCompletionModelAdapter>>),
Zai(Agent<MetricsCompletionModel<ZaiCompletionModelAdapter>>),
MlxLm(Agent<MetricsCompletionModel<MlxLmCompletionModelAdapter>>),
LmStudio(Agent<MetricsCompletionModel<LmStudioCompletionModelAdapter>>),
}
impl AgentWrapper {
pub async fn execute(
&self,
prompt: &str,
history: &[Message],
) -> Result<(String, ExecutionStats), CompletionError> {
match self {
AgentWrapper::OpenRouter(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats)),
AgentWrapper::OpenAI(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats)),
AgentWrapper::Anthropic(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats)),
AgentWrapper::Minimax(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats)),
AgentWrapper::Cerebras(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats)),
AgentWrapper::Ollama(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats)),
AgentWrapper::Zai(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats)),
AgentWrapper::MlxLm(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats)),
AgentWrapper::LmStudio(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats)),
}
}
pub async fn execute_with_messages(
&self,
prompt: &str,
history: &[Message],
) -> Result<(String, ExecutionStats, Vec<Message>), CompletionError> {
match self {
AgentWrapper::OpenRouter(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats, r.messages)),
AgentWrapper::OpenAI(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats, r.messages)),
AgentWrapper::Anthropic(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats, r.messages)),
AgentWrapper::Minimax(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats, r.messages)),
AgentWrapper::Cerebras(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats, r.messages)),
AgentWrapper::Ollama(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats, r.messages)),
AgentWrapper::Zai(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats, r.messages)),
AgentWrapper::MlxLm(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats, r.messages)),
AgentWrapper::LmStudio(agent) => agent
.prompt_with_history(prompt, history)
.await
.map(|r| (r.content(), r.stats, r.messages)),
}
}
pub fn provider_type(&self) -> ProviderType {
match self {
AgentWrapper::OpenRouter(_) => ProviderType::OpenRouter,
AgentWrapper::OpenAI(_) => ProviderType::OpenAI,
AgentWrapper::Anthropic(_) => ProviderType::Anthropic,
AgentWrapper::Minimax(_) => ProviderType::Minimax,
AgentWrapper::Cerebras(_) => ProviderType::Cerebras,
AgentWrapper::Ollama(_) => ProviderType::Ollama,
AgentWrapper::Zai(_) => ProviderType::Zai,
AgentWrapper::MlxLm(_) => ProviderType::MlxLm,
AgentWrapper::LmStudio(_) => ProviderType::LmStudio,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_provider_type_mapping() {
assert_eq!(ProviderType::OpenRouter.to_string(), "openrouter");
assert_eq!(ProviderType::OpenAI.to_string(), "openai");
assert_eq!(ProviderType::Anthropic.to_string(), "anthropic");
assert_eq!(ProviderType::Minimax.to_string(), "minimax");
assert_eq!(ProviderType::Cerebras.to_string(), "cerebras");
assert_eq!(ProviderType::Ollama.to_string(), "ollama");
assert_eq!(ProviderType::Zai.to_string(), "zai");
assert_eq!(ProviderType::MlxLm.to_string(), "mlx");
}
}