Skip to main content

apollo/providers/
rs_ai.rs

1//! rs_ai-backed provider adapter for first-class ChatGPT, Gemini, xAI, Claude,
2//! Cloudflare and generic OpenAI-compatible endpoints.
3
4use 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
13/// Provider implementation backed by the `rs_ai` / `rs_ai_core` SDK.
14pub 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}