Skip to main content

codex_api/endpoint/
models.rs

1use crate::auth::SharedAuthProvider;
2use crate::endpoint::session::EndpointSession;
3use crate::error::ApiError;
4use crate::provider::Provider;
5use codex_client::HttpTransport;
6use codex_client::RequestTelemetry;
7use codex_protocol::openai_models::ModelInfo;
8use codex_protocol::openai_models::ModelsResponse;
9use http::HeaderMap;
10use http::Method;
11use http::header::ETAG;
12use std::sync::Arc;
13
14pub struct ModelsClient<T: HttpTransport> {
15    session: EndpointSession<T>,
16}
17
18impl<T: HttpTransport> ModelsClient<T> {
19    pub fn new(transport: T, provider: Provider, auth: SharedAuthProvider) -> Self {
20        Self {
21            session: EndpointSession::new(transport, provider, auth),
22        }
23    }
24
25    pub fn with_telemetry(self, request: Option<Arc<dyn RequestTelemetry>>) -> Self {
26        Self {
27            session: self.session.with_request_telemetry(request),
28        }
29    }
30
31    fn path() -> &'static str {
32        "models"
33    }
34
35    fn append_client_version_query(req: &mut codex_client::Request, client_version: &str) {
36        let separator = if req.url.contains('?') { '&' } else { '?' };
37        req.url = format!("{}{}client_version={client_version}", req.url, separator);
38    }
39
40    pub fn request_url(provider: &Provider, client_version: &str) -> String {
41        let mut request = provider.build_request(Method::GET, Self::path());
42        Self::append_client_version_query(&mut request, client_version);
43        request.url
44    }
45
46    pub async fn list_models(
47        &self,
48        request_url: String,
49        extra_headers: HeaderMap,
50    ) -> Result<(Vec<ModelInfo>, Option<String>), ApiError> {
51        let resp = self
52            .session
53            .execute_with(
54                Method::GET,
55                Self::path(),
56                extra_headers,
57                /*body*/ None,
58                move |req| {
59                    req.url.clone_from(&request_url);
60                },
61            )
62            .await?;
63
64        let header_etag = resp
65            .headers
66            .get(ETAG)
67            .and_then(|value| value.to_str().ok())
68            .map(ToString::to_string);
69
70        let ModelsResponse { models } = serde_json::from_slice::<ModelsResponse>(&resp.body)
71            .map_err(|e| {
72                ApiError::Stream(format!(
73                    "failed to decode models response: {e}; body: {}",
74                    String::from_utf8_lossy(&resp.body)
75                ))
76            })?;
77
78        Ok((models, header_etag))
79    }
80}
81
82#[cfg(test)]
83mod tests {
84    use super::*;
85    use crate::auth::AuthProvider;
86    use crate::provider::RetryConfig;
87    use codex_client::Request;
88    use codex_client::Response;
89    use codex_client::StreamResponse;
90    use codex_client::TransportError;
91    use http::HeaderMap;
92    use http::StatusCode;
93    use pretty_assertions::assert_eq;
94    use serde_json::json;
95    use std::sync::Arc;
96    use std::sync::Mutex;
97    use std::time::Duration;
98
99    #[derive(Clone)]
100    struct CapturingTransport {
101        last_request: Arc<Mutex<Option<Request>>>,
102        body: Arc<ModelsResponse>,
103        etag: Option<String>,
104    }
105
106    impl Default for CapturingTransport {
107        fn default() -> Self {
108            Self {
109                last_request: Arc::new(Mutex::new(None)),
110                body: Arc::new(ModelsResponse { models: Vec::new() }),
111                etag: None,
112            }
113        }
114    }
115
116    impl HttpTransport for CapturingTransport {
117        async fn execute(&self, req: Request) -> Result<Response, TransportError> {
118            *self.last_request.lock().unwrap() = Some(req);
119            let body = serde_json::to_vec(&*self.body).unwrap();
120            let mut headers = HeaderMap::new();
121            if let Some(etag) = &self.etag {
122                headers.insert(ETAG, etag.parse().unwrap());
123            }
124            Ok(Response {
125                status: StatusCode::OK,
126                headers,
127                body: body.into(),
128            })
129        }
130
131        async fn stream(&self, _req: Request) -> Result<StreamResponse, TransportError> {
132            Err(TransportError::Build("stream should not run".to_string()))
133        }
134    }
135
136    #[derive(Clone, Default)]
137    struct DummyAuth;
138
139    impl AuthProvider for DummyAuth {
140        fn add_auth_headers(&self, _headers: &mut HeaderMap) {}
141    }
142
143    fn provider(base_url: &str) -> Provider {
144        Provider {
145            name: "test".to_string(),
146            base_url: base_url.to_string(),
147            query_params: None,
148            headers: HeaderMap::new(),
149            retry: RetryConfig {
150                max_attempts: 1,
151                base_delay: Duration::from_millis(1),
152                retry_429: false,
153                retry_5xx: true,
154                retry_transport: true,
155            },
156            stream_idle_timeout: Duration::from_secs(1),
157        }
158    }
159
160    #[tokio::test]
161    async fn appends_client_version_query() {
162        let response = ModelsResponse { models: Vec::new() };
163
164        let transport = CapturingTransport {
165            last_request: Arc::new(Mutex::new(None)),
166            body: Arc::new(response),
167            etag: None,
168        };
169
170        let provider = provider("https://example.com/api/codex");
171        let request_url = ModelsClient::<CapturingTransport>::request_url(&provider, "0.99.0");
172        let client = ModelsClient::new(transport.clone(), provider, Arc::new(DummyAuth));
173
174        let (models, _) = client
175            .list_models(request_url, HeaderMap::new())
176            .await
177            .expect("request should succeed");
178
179        assert_eq!(models.len(), 0);
180
181        let url = transport
182            .last_request
183            .lock()
184            .unwrap()
185            .as_ref()
186            .unwrap()
187            .url
188            .clone();
189        assert_eq!(
190            url,
191            "https://example.com/api/codex/models?client_version=0.99.0"
192        );
193    }
194
195    #[tokio::test]
196    async fn parses_models_response() {
197        let response = ModelsResponse {
198            models: vec![
199                serde_json::from_value(json!({
200                    "slug": "gpt-test",
201                    "display_name": "gpt-test",
202                    "description": "desc",
203                    "default_reasoning_level": "medium",
204                    "supported_reasoning_levels": [{"effort": "low", "description": "low"}, {"effort": "medium", "description": "medium"}, {"effort": "high", "description": "high"}],
205                    "shell_type": "shell_command",
206                    "visibility": "list",
207                    "minimal_client_version": [0, 99, 0],
208                    "supported_in_api": true,
209                    "priority": 1,
210                    "upgrade": null,
211                    "base_instructions": "base instructions",
212                    "support_verbosity": false,
213                    "default_verbosity": null,
214                    "apply_patch_tool_type": null,
215                    "truncation_policy": {"mode": "bytes", "limit": 10_000},
216                    "supports_parallel_tool_calls": false,
217                    "supports_image_detail_original": false,
218                    "context_window": 272_000,
219                    "experimental_supported_tools": [],
220                }))
221                .unwrap(),
222            ],
223        };
224
225        let transport = CapturingTransport {
226            last_request: Arc::new(Mutex::new(None)),
227            body: Arc::new(response),
228            etag: None,
229        };
230
231        let provider = provider("https://example.com/api/codex");
232        let request_url = ModelsClient::<CapturingTransport>::request_url(&provider, "0.99.0");
233        let client = ModelsClient::new(transport, provider, Arc::new(DummyAuth));
234
235        let (models, _) = client
236            .list_models(request_url, HeaderMap::new())
237            .await
238            .expect("request should succeed");
239
240        assert_eq!(models.len(), 1);
241        assert_eq!(models[0].slug, "gpt-test");
242        assert_eq!(models[0].supported_in_api, true);
243        assert_eq!(models[0].priority, 1);
244    }
245
246    #[tokio::test]
247    async fn list_models_includes_etag() {
248        let response = ModelsResponse { models: Vec::new() };
249
250        let transport = CapturingTransport {
251            last_request: Arc::new(Mutex::new(None)),
252            body: Arc::new(response),
253            etag: Some("\"abc\"".to_string()),
254        };
255
256        let provider = provider("https://example.com/api/codex");
257        let request_url = ModelsClient::<CapturingTransport>::request_url(&provider, "0.1.0");
258        let client = ModelsClient::new(transport, provider, Arc::new(DummyAuth));
259
260        let (models, etag) = client
261            .list_models(request_url, HeaderMap::new())
262            .await
263            .expect("request should succeed");
264
265        assert_eq!(models.len(), 0);
266        assert_eq!(etag, Some("\"abc\"".to_string()));
267    }
268}