Skip to main content

apollo/providers/
codex.rs

1use async_trait::async_trait;
2use rs_ai_oauth::codex::{codex_request_body, ChatGptCodexClient};
3use serde_json::{json, Value};
4
5use super::traits::{
6    ChatMessage, ChatRequest, ChatResponse, Provider, ProviderCapabilities, ToolCall,
7};
8
9pub struct CodexProvider {
10    client: ChatGptCodexClient,
11}
12
13impl CodexProvider {
14    pub fn new(access_token: impl Into<String>) -> Self {
15        Self {
16            client: ChatGptCodexClient::new(access_token).with_originator("apollo"),
17        }
18    }
19}
20
21#[async_trait]
22impl Provider for CodexProvider {
23    fn name(&self) -> &str {
24        "chatgpt"
25    }
26
27    fn capabilities(&self) -> ProviderCapabilities {
28        ProviderCapabilities {
29            native_tools: true,
30            streaming: false,
31            vision: true,
32            max_context: 272_000,
33            native_web_search: false,
34        }
35    }
36
37    async fn chat(&self, request: &ChatRequest<'_>) -> anyhow::Result<ChatResponse> {
38        let input = messages_to_responses_input(request.messages);
39        let tools = request
40            .tools
41            .unwrap_or(&[])
42            .iter()
43            .map(|tool| {
44                json!({
45                    "type": "function",
46                    "name": tool.name,
47                    "description": tool.description,
48                    "parameters": tool.parameters,
49                    "strict": null,
50                })
51            })
52            .collect();
53        let instructions = request
54            .messages
55            .iter()
56            .find(|message| message.role == "system")
57            .map(|message| message.content.as_str())
58            .unwrap_or("You are a helpful assistant.");
59        let body = codex_request_body(request.model, instructions, input, tools, None);
60        let response = self.client.complete(body, None).await?;
61        Ok(ChatResponse {
62            text: (!response.text.is_empty()).then_some(response.text),
63            tool_calls: response
64                .tool_calls
65                .into_iter()
66                .map(|call| ToolCall {
67                    id: call.id,
68                    name: call.name,
69                    arguments: call.arguments,
70                })
71                .collect(),
72            usage: Some(super::traits::Usage {
73                input_tokens: response.input_tokens as u32,
74                output_tokens: response.output_tokens as u32,
75            }),
76        })
77    }
78}
79
80fn messages_to_responses_input(messages: &[ChatMessage]) -> Vec<Value> {
81    messages
82        .iter()
83        .filter_map(|message| match message.role.as_str() {
84            "system" => None,
85            "tool_result" => Some(json!({
86                "type": "function_call_output",
87                "call_id": message.tool_use_id.clone().unwrap_or_default(),
88                "output": message.content,
89            })),
90            "assistant" => Some(json!({
91                "type": "message",
92                "role": "assistant",
93                "content": [{"type": "output_text", "text": message.content, "annotations": []}],
94                "status": "completed",
95            })),
96            _ => Some(json!({
97                "role": "user",
98                "content": [{"type": "input_text", "text": message.content}],
99            })),
100        })
101        .collect()
102}
103
104#[cfg(test)]
105mod tests {
106    use super::*;
107
108    #[test]
109    fn converts_messages_to_responses_input() {
110        let messages = [ChatMessage::user("hello")];
111        let input = messages_to_responses_input(&messages);
112        assert_eq!(input[0]["role"], "user");
113        assert_eq!(input[0]["content"][0]["type"], "input_text");
114    }
115}