1use std::{collections::BTreeMap, sync::Arc};
4
5use async_trait::async_trait;
6use serde_json::Value;
7use starweaver_oauth::OAuthTokenSource;
8
9use crate::{
10 ModelError,
11 oauth::headers::{
12 CODEX_USER_AGENT_HEADER, build_codex_headers, patch_codex_responses_body,
13 patch_codex_websocket_request, trace_session_headers, validate_safe_extra_headers,
14 },
15 transport::{
16 DynHttpClient, HttpRequest, HttpResponse, ModelEventStream, ModelHttpClient,
17 ModelWebSocketEventSession, ReqwestHttpClient, extend_headers_case_insensitive,
18 },
19};
20
21pub struct OAuthBearerHttpClient {
23 inner: DynHttpClient,
24 token_source: Arc<dyn OAuthTokenSource>,
25 provider_name: String,
26 extra_headers: BTreeMap<String, String>,
27}
28
29impl OAuthBearerHttpClient {
30 pub fn new(
36 inner: DynHttpClient,
37 token_source: Arc<dyn OAuthTokenSource>,
38 provider_name: impl Into<String>,
39 extra_headers: BTreeMap<String, String>,
40 ) -> Result<Self, ModelError> {
41 validate_safe_extra_headers(&extra_headers)?;
42 Ok(Self {
43 inner,
44 token_source,
45 provider_name: provider_name.into(),
46 extra_headers,
47 })
48 }
49
50 pub fn with_default_http_client(
56 token_source: Arc<dyn OAuthTokenSource>,
57 provider_name: impl Into<String>,
58 extra_headers: BTreeMap<String, String>,
59 ) -> Result<Self, ModelError> {
60 Self::new(
61 Arc::new(ReqwestHttpClient::new()?),
62 token_source,
63 provider_name,
64 extra_headers,
65 )
66 }
67
68 fn prepare_request(
69 &self,
70 mut request: HttpRequest,
71 snapshot: &starweaver_oauth::TokenSnapshot,
72 ) -> Result<HttpRequest, ModelError> {
73 let explicit_codex_routing_headers = if self.provider_name == "codex" {
74 extract_case_insensitive_headers(&request.headers, CODEX_ROUTING_HEADER_NAMES)
75 } else {
76 BTreeMap::new()
77 };
78 request.headers.insert(
79 "Authorization".to_string(),
80 format!("Bearer {}", snapshot.access_token),
81 );
82 if self.provider_name == "codex" {
83 insert_header_if_absent_case_insensitive(
84 &mut request.headers,
85 CODEX_USER_AGENT_HEADER,
86 codex_user_agent(),
87 );
88 let explicit_extra_routing_headers =
89 extract_case_insensitive_headers(&self.extra_headers, CODEX_ROUTING_HEADER_NAMES);
90 let mut extra_headers = trace_session_headers(&request);
91 extend_headers_case_insensitive(&mut extra_headers, self.extra_headers.clone());
92 restore_case_insensitive_headers(&mut extra_headers, explicit_extra_routing_headers);
93 extend_headers_case_insensitive(
94 &mut request.headers,
95 build_codex_headers(&snapshot.account, Some(&extra_headers))?,
96 );
97 restore_case_insensitive_headers(&mut request.headers, explicit_codex_routing_headers);
98 patch_codex_responses_body(&mut request);
99 patch_codex_websocket_request(&mut request);
100 } else {
101 extend_headers_case_insensitive(&mut request.headers, self.extra_headers.clone());
102 }
103 Ok(request)
104 }
105}
106
107#[async_trait]
108impl ModelHttpClient for OAuthBearerHttpClient {
109 async fn send(&self, request: HttpRequest) -> Result<HttpResponse, ModelError> {
110 let snapshot = self
111 .token_source
112 .get_token()
113 .await
114 .map_err(|error| ModelError::Transport(error.to_string()))?;
115 let request_with_auth = self.prepare_request(request.clone(), &snapshot)?;
116 match self.inner.send(request_with_auth).await {
117 Err(ModelError::ProviderStatus { status: 401, .. }) => {
118 let refreshed = self
119 .token_source
120 .refresh_token()
121 .await
122 .map_err(|error| ModelError::Transport(error.to_string()))?;
123 self.inner
124 .send(self.prepare_request(request, &refreshed)?)
125 .await
126 }
127 result => result,
128 }
129 }
130
131 async fn send_event_stream(&self, request: HttpRequest) -> Result<Vec<Value>, ModelError> {
132 let snapshot = self
133 .token_source
134 .get_token()
135 .await
136 .map_err(|error| ModelError::Transport(error.to_string()))?;
137 let request_with_auth = self.prepare_request(request.clone(), &snapshot)?;
138 match self.inner.send_event_stream(request_with_auth).await {
139 Err(ModelError::ProviderStatus { status: 401, .. }) => {
140 let refreshed = self
141 .token_source
142 .refresh_token()
143 .await
144 .map_err(|error| ModelError::Transport(error.to_string()))?;
145 self.inner
146 .send_event_stream(self.prepare_request(request, &refreshed)?)
147 .await
148 }
149 result => result,
150 }
151 }
152
153 async fn send_event_stream_incremental(
154 &self,
155 request: HttpRequest,
156 ) -> Result<ModelEventStream, ModelError> {
157 let snapshot = self
158 .token_source
159 .get_token()
160 .await
161 .map_err(|error| ModelError::Transport(error.to_string()))?;
162 let request_with_auth = self.prepare_request(request.clone(), &snapshot)?;
163 match self
164 .inner
165 .send_event_stream_incremental(request_with_auth)
166 .await
167 {
168 Err(ModelError::ProviderStatus { status: 401, .. }) => {
169 let refreshed = self
170 .token_source
171 .refresh_token()
172 .await
173 .map_err(|error| ModelError::Transport(error.to_string()))?;
174 self.inner
175 .send_event_stream_incremental(self.prepare_request(request, &refreshed)?)
176 .await
177 }
178 result => result,
179 }
180 }
181
182 async fn send_websocket_event_stream_incremental(
183 &self,
184 mut request: HttpRequest,
185 ) -> Result<ModelEventStream, ModelError> {
186 request
187 .metadata
188 .entry("starweaver.response_stream_transport".to_string())
189 .or_insert_with(|| Value::String("websocket".to_string()));
190 let snapshot = self
191 .token_source
192 .get_token()
193 .await
194 .map_err(|error| ModelError::Transport(error.to_string()))?;
195 let request_with_auth = self.prepare_request(request.clone(), &snapshot)?;
196 match self
197 .inner
198 .send_websocket_event_stream_incremental(request_with_auth)
199 .await
200 {
201 Err(ModelError::ProviderStatus { status: 401, .. }) => {
202 let refreshed = self
203 .token_source
204 .refresh_token()
205 .await
206 .map_err(|error| ModelError::Transport(error.to_string()))?;
207 self.inner
208 .send_websocket_event_stream_incremental(
209 self.prepare_request(request, &refreshed)?,
210 )
211 .await
212 }
213 result => result,
214 }
215 }
216
217 fn websocket_event_session(&self) -> Box<dyn ModelWebSocketEventSession + '_> {
218 Box::new(OAuthWebSocketEventSession {
219 client: self,
220 inner: self.inner.websocket_event_session(),
221 })
222 }
223}
224
225struct OAuthWebSocketEventSession<'a> {
226 client: &'a OAuthBearerHttpClient,
227 inner: Box<dyn ModelWebSocketEventSession + 'a>,
228}
229
230#[async_trait]
231impl ModelWebSocketEventSession for OAuthWebSocketEventSession<'_> {
232 async fn send_websocket_event_stream_incremental(
233 &mut self,
234 mut request: HttpRequest,
235 ) -> Result<ModelEventStream, ModelError> {
236 request
237 .metadata
238 .entry("starweaver.response_stream_transport".to_string())
239 .or_insert_with(|| Value::String("websocket".to_string()));
240 let snapshot = self
241 .client
242 .token_source
243 .get_token()
244 .await
245 .map_err(|error| ModelError::Transport(error.to_string()))?;
246 let request_with_auth = self.client.prepare_request(request.clone(), &snapshot)?;
247 match self
248 .inner
249 .send_websocket_event_stream_incremental(request_with_auth)
250 .await
251 {
252 Err(ModelError::ProviderStatus { status: 401, .. }) => {
253 self.inner.reset().await;
254 let refreshed = self
255 .client
256 .token_source
257 .refresh_token()
258 .await
259 .map_err(|error| ModelError::Transport(error.to_string()))?;
260 self.inner
261 .send_websocket_event_stream_incremental(
262 self.client.prepare_request(request, &refreshed)?,
263 )
264 .await
265 }
266 result => result,
267 }
268 }
269
270 async fn reset(&mut self) {
271 self.inner.reset().await;
272 }
273}
274
275fn codex_user_agent() -> String {
276 format!(
277 "{}/{}",
278 starweaver_core::sdk_name(),
279 env!("CARGO_PKG_VERSION")
280 )
281}
282
283const CODEX_SESSION_ROUTING_HEADER_NAMES: &[&str] = &["session_id", "session-id"];
284const CODEX_THREAD_ROUTING_HEADER_NAMES: &[&str] =
285 &["thread_id", "thread-id", "x-client-request-id"];
286const CODEX_ROUTING_HEADER_NAMES: &[&str] = &[
287 "session_id",
288 "session-id",
289 "thread_id",
290 "thread-id",
291 "x-client-request-id",
292];
293
294fn insert_header_if_absent_case_insensitive(
295 headers: &mut BTreeMap<String, String>,
296 name: &str,
297 value: String,
298) {
299 if !headers.keys().any(|key| key.eq_ignore_ascii_case(name)) {
300 headers.insert(name.to_string(), value);
301 }
302}
303
304fn extract_case_insensitive_headers(
305 headers: &BTreeMap<String, String>,
306 names: &[&str],
307) -> BTreeMap<String, String> {
308 headers
309 .iter()
310 .filter(|(key, _)| names.iter().any(|name| key.eq_ignore_ascii_case(name)))
311 .map(|(key, value)| (key.clone(), value.clone()))
312 .collect()
313}
314
315fn restore_case_insensitive_headers(
316 headers: &mut BTreeMap<String, String>,
317 explicit_headers: BTreeMap<String, String>,
318) {
319 if explicit_headers
320 .keys()
321 .any(|key| key_matches_any(key, CODEX_SESSION_ROUTING_HEADER_NAMES))
322 {
323 remove_case_insensitive_headers(headers, CODEX_SESSION_ROUTING_HEADER_NAMES);
324 }
325 if explicit_headers
326 .keys()
327 .any(|key| key_matches_any(key, CODEX_THREAD_ROUTING_HEADER_NAMES))
328 {
329 remove_case_insensitive_headers(headers, CODEX_THREAD_ROUTING_HEADER_NAMES);
330 }
331 for (explicit_key, explicit_value) in explicit_headers {
332 headers.retain(|key, _| !key.eq_ignore_ascii_case(&explicit_key));
333 headers.insert(explicit_key, explicit_value);
334 }
335}
336
337fn remove_case_insensitive_headers(headers: &mut BTreeMap<String, String>, names: &[&str]) {
338 headers.retain(|key, _| !key_matches_any(key, names));
339}
340
341fn key_matches_any(key: &str, names: &[&str]) -> bool {
342 names.iter().any(|name| key.eq_ignore_ascii_case(name))
343}