starweaver_model/transport/
reqwest_client.rs1use 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#[derive(Clone, Debug)]
16pub struct ReqwestHttpClient {
17 client: reqwest::Client,
18}
19
20impl ReqwestHttpClient {
21 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}