apollo/providers/
rs_ai.rs1use async_trait::async_trait;
5use rs_ai_core::{
6 GenerateOptions, GenerateResult, Message, Prompt, ToolCallRequest, ToolDefinition,
7};
8
9use crate::providers::traits::{
10 ChatRequest, ChatResponse, Provider, ProviderCapabilities, ToolCall, Usage,
11};
12
13pub struct RsAiProvider {
15 provider_name: String,
16 model_id: String,
17 api_key: String,
18 base_url: Option<String>,
19 account_id: Option<String>,
20}
21
22impl RsAiProvider {
23 pub fn new(
24 provider_name: &str,
25 model_id: &str,
26 api_key: &str,
27 base_url: Option<String>,
28 account_id: Option<String>,
29 ) -> Self {
30 Self {
31 provider_name: provider_name.to_string(),
32 model_id: model_id.to_string(),
33 api_key: api_key.to_string(),
34 base_url,
35 account_id,
36 }
37 }
38
39 fn effective_model_id(&self) -> &str {
40 if !self.model_id.is_empty() {
41 return &self.model_id;
42 }
43 match self.provider_name.as_str() {
44 "chatgpt" | "openai" => "gpt-4o",
45 "gemini" => "gemini-2.5-flash",
46 "xai" | "grok" => "grok-4.20-reasoning",
47 "claude" => "claude-sonnet-4-6",
48 "cloudflare" => "@cf/meta/llama-3.1-8b-instruct",
49 _ => "",
50 }
51 }
52
53 fn build_model(&self) -> anyhow::Result<Box<dyn rs_ai_core::LanguageModel>> {
54 use rs_ai_providers::{
55 ChatGptProvider, ClaudeProvider, CloudflareProvider, GeminiProvider,
56 OpenAiCompatibleConfig, OpenAiCompatibleProvider, XaiProvider,
57 };
58
59 let model: Box<dyn rs_ai_core::LanguageModel> = match self.provider_name.as_str() {
60 "claude" => {
61 Box::new(ClaudeProvider::new(&self.api_key).model(self.effective_model_id()))
62 }
63 "chatgpt" | "openai" => {
64 Box::new(ChatGptProvider::new(&self.api_key).model(self.effective_model_id()))
65 }
66 "gemini" => {
67 Box::new(GeminiProvider::new(&self.api_key).model(self.effective_model_id()))
68 }
69 "xai" | "grok" => {
70 Box::new(XaiProvider::new(&self.api_key).model(self.effective_model_id()))
71 }
72 "cloudflare" => {
73 let account_id = self
74 .account_id
75 .as_deref()
76 .or(self.base_url.as_deref())
77 .unwrap_or("")
78 .to_string();
79 Box::new(
80 CloudflareProvider::new(account_id, &self.api_key)
81 .model(self.effective_model_id()),
82 )
83 }
84 other => {
85 let base_url = self
86 .base_url
87 .as_deref()
88 .unwrap_or("https://api.openai.com/v1");
89 let config = OpenAiCompatibleConfig::new(base_url, &self.api_key);
90 let provider = OpenAiCompatibleProvider::new(config, other, other);
91 provider.language_model(self.effective_model_id())
92 }
93 };
94 Ok(model)
95 }
96}
97
98#[async_trait]
99impl Provider for RsAiProvider {
100 fn name(&self) -> &str {
101 &self.provider_name
102 }
103
104 fn capabilities(&self) -> ProviderCapabilities {
105 ProviderCapabilities {
106 native_tools: true,
107 streaming: true,
108 vision: matches!(
109 self.provider_name.as_str(),
110 "claude" | "chatgpt" | "openai" | "gemini" | "xai" | "grok"
111 ),
112 max_context: 200_000,
113 }
114 }
115
116 async fn chat(&self, request: &ChatRequest<'_>) -> anyhow::Result<ChatResponse> {
117 let model = self.build_model()?;
118
119 let messages: Vec<Message> = request
120 .messages
121 .iter()
122 .map(|m| match m.role.as_str() {
123 "system" => Message::system(&m.content),
124 "assistant" => Message::assistant(&m.content),
125 "tool_result" => {
126 Message::tool_result(m.tool_use_id.as_deref().unwrap_or(""), &m.content)
127 }
128 _ => Message::user(&m.content),
129 })
130 .collect();
131
132 let prompt = Prompt::Messages(messages);
133
134 let mut options = GenerateOptions::default().with_temperature(request.temperature);
135 if let Some(max_tokens) = request.max_tokens {
136 options = options.with_max_tokens(max_tokens);
137 }
138
139 let tools: Vec<ToolDefinition> = request
140 .tools
141 .unwrap_or(&[])
142 .iter()
143 .map(|t| ToolDefinition {
144 name: t.name.clone(),
145 description: t.description.clone(),
146 parameters: t.parameters.clone(),
147 examples: None,
148 })
149 .collect();
150 if !tools.is_empty() {
151 options = options
152 .with_tools(tools)
153 .with_tool_choice(rs_ai_core::ToolChoice::Auto);
154 }
155
156 let result = model
157 .generate(prompt, options)
158 .await
159 .map_err(|e| anyhow::anyhow!("rs_ai provider error: {e}"))?;
160
161 Ok(map_generate_result(result)?)
162 }
163}
164
165fn map_generate_result(result: GenerateResult) -> anyhow::Result<ChatResponse> {
166 let tool_calls = result
167 .tool_calls
168 .iter()
169 .map(|tc: &ToolCallRequest| -> anyhow::Result<ToolCall> {
170 Ok(ToolCall {
171 id: tc.id.clone(),
172 name: tc.name.clone(),
173 arguments: serde_json::to_string(&tc.arguments)?,
174 })
175 })
176 .collect::<Result<Vec<_>, _>>()?;
177
178 let usage = Usage {
179 input_tokens: result.usage.prompt_tokens.unwrap_or(0) as u32,
180 output_tokens: result.usage.completion_tokens.unwrap_or(0) as u32,
181 };
182
183 Ok(ChatResponse {
184 text: result.text,
185 tool_calls,
186 usage: Some(usage),
187 })
188}