Skip to main content

starweaver_model/oauth/
http_client.rs

1//! OAuth bearer HTTP client wrapper.
2
3use 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
21/// HTTP client wrapper that attaches OAuth bearer headers and refreshes once on 401.
22pub 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    /// Create a wrapper around an injected model HTTP client.
31    ///
32    /// # Errors
33    ///
34    /// Returns an error when `extra_headers` contains reserved OAuth/Codex headers.
35    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    /// Create a wrapper using the default reqwest HTTP client.
51    ///
52    /// # Errors
53    ///
54    /// Returns an error when reqwest client construction fails or headers are invalid.
55    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}