use std::collections::BTreeMap;
use async_trait::async_trait;
use futures_util::StreamExt;
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use serde_json::Value;
use crate::{ModelError, allow_real_model_requests};
use super::sse::{SseJsonParser, StreamSendError, push_sse_utf8_buffer, send_sse_parser_events};
use super::{HttpMethod, HttpRequest, HttpResponse, ModelEventStream, ModelHttpClient};
use crate::transport::{is_retryable_status, websocket};
#[derive(Clone, Debug)]
pub struct ReqwestHttpClient {
client: reqwest::Client,
}
impl ReqwestHttpClient {
pub fn new() -> Result<Self, ModelError> {
let client = reqwest::Client::builder()
.build()
.map_err(|err| ModelError::Transport(err.to_string()))?;
Ok(Self { client })
}
async fn send_request(&self, request: &HttpRequest) -> Result<reqwest::Response, ModelError> {
if !allow_real_model_requests() {
return Err(ModelError::RealModelRequestBlocked {
url: request.url.clone(),
});
}
let mut builder = match request.method {
HttpMethod::Post => self.client.post(&request.url),
}
.headers(Self::header_map(&request.headers)?)
.json(&request.body);
if let Some(timeout) = request.timeout {
builder = builder.timeout(timeout);
}
builder
.send()
.await
.map_err(|err| ModelError::Transport(err.to_string()))
}
fn header_map(headers: &BTreeMap<String, String>) -> Result<HeaderMap, ModelError> {
let mut map = HeaderMap::new();
for (name, value) in headers {
let name = HeaderName::from_bytes(name.as_bytes()).map_err(|err| {
ModelError::Transport(format!("invalid header name {name}: {err}"))
})?;
let value = HeaderValue::from_str(value).map_err(|err| {
ModelError::Transport(format!("invalid header value for {name}: {err}"))
})?;
map.insert(name, value);
}
Ok(map)
}
}
#[async_trait]
impl ModelHttpClient for ReqwestHttpClient {
async fn send(&self, request: HttpRequest) -> Result<HttpResponse, ModelError> {
let cancellation_token = request.cancellation_token.clone();
if cancellation_token.is_cancelled() {
return Err(ModelError::Cancelled {
reason: "model HTTP request cancellation requested".to_string(),
});
}
let response = tokio::select! {
biased;
() = cancellation_token.cancelled() => {
return Err(ModelError::Cancelled {
reason: "model HTTP request cancellation requested".to_string(),
});
}
response = self.send_request(&request) => response?,
};
let status = response.status().as_u16();
let headers = response_headers(&response);
let body = tokio::select! {
biased;
() = cancellation_token.cancelled() => {
return Err(ModelError::Cancelled {
reason: "model HTTP request cancellation requested".to_string(),
});
}
body = response.json::<Value>() => {
body.map_err(|err| ModelError::Transport(err.to_string()))?
}
};
if (200..300).contains(&status) {
Ok(HttpResponse {
status,
headers,
body,
})
} else {
Err(ModelError::ProviderStatus {
status,
body,
retryable: is_retryable_status(status),
})
}
}
async fn send_event_stream_incremental(
&self,
request: HttpRequest,
) -> Result<ModelEventStream, ModelError> {
let cancellation_token = request.cancellation_token.clone();
if cancellation_token.is_cancelled() {
return Err(ModelError::Cancelled {
reason: "model event stream cancellation requested".to_string(),
});
}
let response = tokio::select! {
biased;
() = cancellation_token.cancelled() => {
return Err(ModelError::Cancelled {
reason: "model event stream cancellation requested".to_string(),
});
}
response = self.send_request(&request) => response?,
};
let status = response.status().as_u16();
if !(200..300).contains(&status) {
let text = response
.text()
.await
.map_err(|err| ModelError::Transport(err.to_string()))?;
let body = serde_json::from_str(&text).unwrap_or(Value::String(text));
return Err(ModelError::ProviderStatus {
status,
body,
retryable: is_retryable_status(status),
});
}
let (sender, receiver) = tokio::sync::mpsc::channel(32);
let worker_cancellation_token = cancellation_token.clone();
tokio::spawn(async move {
let mut parser = SseJsonParser::default();
let mut bytes = response.bytes_stream();
let mut utf8_buffer = Vec::new();
loop {
let chunk = tokio::select! {
biased;
() = worker_cancellation_token.cancelled() => {
let _ = sender
.send(Err(ModelError::Cancelled {
reason: "model event stream cancellation requested".to_string(),
}))
.await;
return;
}
chunk = bytes.next() => chunk,
};
let Some(chunk) = chunk else {
break;
};
match chunk {
Ok(bytes) => {
utf8_buffer.extend_from_slice(&bytes);
match push_sse_utf8_buffer(&sender, &mut parser, &mut utf8_buffer).await {
Ok(()) => {}
Err(StreamSendError::Closed) => return,
Err(StreamSendError::InvalidUtf8(error)) => {
let _ = sender
.send(Err(ModelError::ResponseParsing(format!(
"invalid server-sent event UTF-8: {error}"
))))
.await;
return;
}
}
}
Err(error) => {
let _ = sender
.send(Err(ModelError::Transport(error.to_string())))
.await;
return;
}
}
}
if !utf8_buffer.is_empty() {
match std::str::from_utf8(&utf8_buffer) {
Ok(text) => {
if !send_sse_parser_events(&sender, parser.push_str(text)).await {
return;
}
}
Err(error) => {
let _ = sender
.send(Err(ModelError::ResponseParsing(format!(
"invalid server-sent event UTF-8: {error}"
))))
.await;
return;
}
}
}
let _ = send_sse_parser_events(&sender, parser.finish()).await;
});
Ok(ModelEventStream::new_with_cancellation(
receiver,
cancellation_token,
))
}
async fn send_websocket_event_stream_incremental(
&self,
request: HttpRequest,
) -> Result<ModelEventStream, ModelError> {
Box::pin(websocket::send_websocket_event_stream_incremental(request)).await
}
fn websocket_event_session(&self) -> Box<dyn super::ModelWebSocketEventSession + '_> {
Box::new(websocket::ReusableWebSocketEventSession::default())
}
}
fn response_headers(response: &reqwest::Response) -> BTreeMap<String, String> {
response
.headers()
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.as_str().to_string(), value.to_string()))
})
.collect()
}