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 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}