Skip to main content

starweaver_model/transport/
reqwest_client.rs

1use std::collections::BTreeMap;
2
3use async_trait::async_trait;
4use futures_util::StreamExt;
5use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
6use serde_json::Value;
7
8use crate::{ModelError, allow_real_model_requests};
9
10use super::sse::{SseJsonParser, StreamSendError, push_sse_utf8_buffer, send_sse_parser_events};
11use super::{HttpMethod, HttpRequest, HttpResponse, ModelEventStream, ModelHttpClient};
12use crate::transport::{is_retryable_status, websocket};
13
14/// Reqwest-backed HTTP client.
15#[derive(Clone, Debug)]
16pub struct ReqwestHttpClient {
17    client: reqwest::Client,
18}
19
20impl ReqwestHttpClient {
21    /// Create a reqwest-backed client with rustls TLS.
22    ///
23    /// # Errors
24    ///
25    /// Returns an error when reqwest client construction fails.
26    pub fn new() -> Result<Self, ModelError> {
27        let client = reqwest::Client::builder()
28            .build()
29            .map_err(|err| ModelError::Transport(err.to_string()))?;
30        Ok(Self { client })
31    }
32
33    async fn send_request(&self, request: &HttpRequest) -> Result<reqwest::Response, ModelError> {
34        if !allow_real_model_requests() {
35            return Err(ModelError::RealModelRequestBlocked {
36                url: request.url.clone(),
37            });
38        }
39
40        let mut builder = match request.method {
41            HttpMethod::Post => self.client.post(&request.url),
42        }
43        .headers(Self::header_map(&request.headers)?)
44        .json(&request.body);
45
46        if let Some(timeout) = request.timeout {
47            builder = builder.timeout(timeout);
48        }
49
50        builder
51            .send()
52            .await
53            .map_err(|err| ModelError::Transport(err.to_string()))
54    }
55
56    fn header_map(headers: &BTreeMap<String, String>) -> Result<HeaderMap, ModelError> {
57        let mut map = HeaderMap::new();
58        for (name, value) in headers {
59            let name = HeaderName::from_bytes(name.as_bytes()).map_err(|err| {
60                ModelError::Transport(format!("invalid header name {name}: {err}"))
61            })?;
62            let value = HeaderValue::from_str(value).map_err(|err| {
63                ModelError::Transport(format!("invalid header value for {name}: {err}"))
64            })?;
65            map.insert(name, value);
66        }
67        Ok(map)
68    }
69}
70
71#[async_trait]
72impl ModelHttpClient for ReqwestHttpClient {
73    async fn send(&self, request: HttpRequest) -> Result<HttpResponse, ModelError> {
74        let cancellation_token = request.cancellation_token.clone();
75        if cancellation_token.is_cancelled() {
76            return Err(ModelError::Cancelled {
77                reason: "model HTTP request cancellation requested".to_string(),
78            });
79        }
80        let response = tokio::select! {
81            biased;
82            () = cancellation_token.cancelled() => {
83                return Err(ModelError::Cancelled {
84                    reason: "model HTTP request cancellation requested".to_string(),
85                });
86            }
87            response = self.send_request(&request) => response?,
88        };
89        let status = response.status().as_u16();
90        let headers = response_headers(&response);
91        let body = tokio::select! {
92            biased;
93            () = cancellation_token.cancelled() => {
94                return Err(ModelError::Cancelled {
95                    reason: "model HTTP request cancellation requested".to_string(),
96                });
97            }
98            body = response.json::<Value>() => {
99                body.map_err(|err| ModelError::Transport(err.to_string()))?
100            }
101        };
102
103        if (200..300).contains(&status) {
104            Ok(HttpResponse {
105                status,
106                headers,
107                body,
108            })
109        } else {
110            Err(ModelError::ProviderStatus {
111                status,
112                body,
113                retryable: is_retryable_status(status),
114            })
115        }
116    }
117
118    async fn send_event_stream_incremental(
119        &self,
120        request: HttpRequest,
121    ) -> Result<ModelEventStream, ModelError> {
122        let cancellation_token = request.cancellation_token.clone();
123        if cancellation_token.is_cancelled() {
124            return Err(ModelError::Cancelled {
125                reason: "model event stream cancellation requested".to_string(),
126            });
127        }
128        let response = tokio::select! {
129            biased;
130            () = cancellation_token.cancelled() => {
131                return Err(ModelError::Cancelled {
132                    reason: "model event stream cancellation requested".to_string(),
133                });
134            }
135            response = self.send_request(&request) => response?,
136        };
137        let status = response.status().as_u16();
138        if !(200..300).contains(&status) {
139            let text = response
140                .text()
141                .await
142                .map_err(|err| ModelError::Transport(err.to_string()))?;
143            let body = serde_json::from_str(&text).unwrap_or(Value::String(text));
144            return Err(ModelError::ProviderStatus {
145                status,
146                body,
147                retryable: is_retryable_status(status),
148            });
149        }
150        let (sender, receiver) = tokio::sync::mpsc::channel(32);
151        let worker_cancellation_token = cancellation_token.clone();
152        tokio::spawn(async move {
153            let mut parser = SseJsonParser::default();
154            let mut bytes = response.bytes_stream();
155            let mut utf8_buffer = Vec::new();
156            loop {
157                let chunk = tokio::select! {
158                    biased;
159                    () = worker_cancellation_token.cancelled() => {
160                        let _ = sender
161                            .send(Err(ModelError::Cancelled {
162                                reason: "model event stream cancellation requested".to_string(),
163                            }))
164                            .await;
165                        return;
166                    }
167                    chunk = bytes.next() => chunk,
168                };
169                let Some(chunk) = chunk else {
170                    break;
171                };
172                match chunk {
173                    Ok(bytes) => {
174                        utf8_buffer.extend_from_slice(&bytes);
175                        match push_sse_utf8_buffer(&sender, &mut parser, &mut utf8_buffer).await {
176                            Ok(()) => {}
177                            Err(StreamSendError::Closed) => return,
178                            Err(StreamSendError::InvalidUtf8(error)) => {
179                                let _ = sender
180                                    .send(Err(ModelError::ResponseParsing(format!(
181                                        "invalid server-sent event UTF-8: {error}"
182                                    ))))
183                                    .await;
184                                return;
185                            }
186                        }
187                    }
188                    Err(error) => {
189                        let _ = sender
190                            .send(Err(ModelError::Transport(error.to_string())))
191                            .await;
192                        return;
193                    }
194                }
195            }
196            if !utf8_buffer.is_empty() {
197                match std::str::from_utf8(&utf8_buffer) {
198                    Ok(text) => {
199                        if !send_sse_parser_events(&sender, parser.push_str(text)).await {
200                            return;
201                        }
202                    }
203                    Err(error) => {
204                        let _ = sender
205                            .send(Err(ModelError::ResponseParsing(format!(
206                                "invalid server-sent event UTF-8: {error}"
207                            ))))
208                            .await;
209                        return;
210                    }
211                }
212            }
213            let _ = send_sse_parser_events(&sender, parser.finish()).await;
214        });
215        Ok(ModelEventStream::new_with_cancellation(
216            receiver,
217            cancellation_token,
218        ))
219    }
220
221    async fn send_websocket_event_stream_incremental(
222        &self,
223        request: HttpRequest,
224    ) -> Result<ModelEventStream, ModelError> {
225        Box::pin(websocket::send_websocket_event_stream_incremental(request)).await
226    }
227
228    fn websocket_event_session(&self) -> Box<dyn super::ModelWebSocketEventSession + '_> {
229        Box::new(websocket::ReusableWebSocketEventSession::default())
230    }
231}
232
233fn response_headers(response: &reqwest::Response) -> BTreeMap<String, String> {
234    response
235        .headers()
236        .iter()
237        .filter_map(|(name, value)| {
238            value
239                .to_str()
240                .ok()
241                .map(|value| (name.as_str().to_string(), value.to_string()))
242        })
243        .collect()
244}