use std::pin::Pin;
use std::time::Duration;
use futures::{Stream, StreamExt};
use crate::cancel::CancellationToken;
use crate::error::ProviderError;
use crate::providers::sse::SseLineBuffer;
pub async fn build_request(
client: &reqwest::Client,
url: &str,
body: serde_json::Value,
headers: Vec<(String, String)>,
timeout: Option<Duration>,
) -> Result<reqwest::Response, ProviderError> {
let timeout_ms = timeout.map(|d| d.as_millis() as u64);
let mut req_builder = client.post(url);
for (key, value) in &headers {
req_builder = req_builder.header(key.as_str(), value.as_str());
}
if let Some(t) = timeout {
req_builder = req_builder.timeout(t);
}
req_builder.json(&body).send().await.map_err(|e| {
if e.is_timeout() {
let ms = timeout_ms.unwrap_or(120_000);
tracing::warn!(
%ms,
per_request_timeout = ?timeout,
error = %e,
url = %url,
"LLM request timed out (generic http)"
);
ProviderError::Timeout { timeout_ms: ms }
} else {
ProviderError::Network { message: e.to_string(), detail: None }
}
})
}
pub fn handle_error_response(
status: u16,
headers: &reqwest::header::HeaderMap,
body: &str,
) -> ProviderError {
let retry_after = headers
.get(reqwest::header::RETRY_AFTER)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.trim().parse::<f64>().ok())
.map(|secs| (secs * 1000.0) as u64)
.unwrap_or(1000);
match status {
401 => ProviderError::Auth("Unauthorized: check API key".to_owned()),
403 => ProviderError::Auth("Forbidden: check API key permissions".to_owned()),
429 => ProviderError::RateLimit { retry_after_ms: retry_after },
503 => ProviderError::Overloaded,
529 => ProviderError::Overloaded,
_ => ProviderError::Internal { status, message: body.to_owned() },
}
}
pub fn create_sse_stream(
response: reqwest::Response,
cancel: Option<CancellationToken>,
) -> Pin<Box<dyn Stream<Item = Result<String, ProviderError>> + Send>> {
let bytes_stream = response.bytes_stream();
let mut buffer = SseLineBuffer::new();
let line_stream = bytes_stream.flat_map(move |chunk_result| {
let bytes = match chunk_result {
Ok(b) => b,
Err(e) => {
return futures::stream::iter(vec![Err(ProviderError::Network {
message: e.to_string(),
detail: None,
})]);
}
};
let lines = buffer.feed(&bytes);
let results: Vec<Result<String, ProviderError>> = lines.into_iter().map(Ok).collect();
futures::stream::iter(results)
});
let stream: Pin<Box<dyn Stream<Item = Result<String, ProviderError>> + Send>> = match cancel {
Some(ct) => {
let inner: tokio_util::sync::CancellationToken = ct.into();
Box::pin(line_stream.take_until(inner.cancelled_owned()))
}
None => Box::pin(line_stream),
};
stream
}
#[cfg(test)]
mod tests {
use super::*;
use reqwest::header::{HeaderMap, HeaderValue, RETRY_AFTER};
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
#[test]
fn test_handle_error_401() {
let err = handle_error_response(401, &HeaderMap::new(), "");
assert!(matches!(err, ProviderError::Auth(m) if m.contains("Unauthorized")));
}
#[test]
fn test_handle_error_403() {
let err = handle_error_response(403, &HeaderMap::new(), "");
assert!(matches!(err, ProviderError::Auth(m) if m.contains("Forbidden")));
}
#[test]
fn test_handle_error_429_with_retry_after() {
let mut headers = HeaderMap::new();
headers.insert(RETRY_AFTER, HeaderValue::from_static("30"));
let err = handle_error_response(429, &headers, "");
assert!(
matches!(err, ProviderError::RateLimit { retry_after_ms } if retry_after_ms == 30_000)
);
}
#[test]
fn test_handle_error_429_with_float_retry_after() {
let mut headers = HeaderMap::new();
headers.insert(RETRY_AFTER, HeaderValue::from_static("2.5"));
let err = handle_error_response(429, &headers, "");
assert!(
matches!(err, ProviderError::RateLimit { retry_after_ms } if retry_after_ms == 2_500)
);
}
#[test]
fn test_handle_error_429_without_retry_after() {
let err = handle_error_response(429, &HeaderMap::new(), "");
assert!(
matches!(err, ProviderError::RateLimit { retry_after_ms } if retry_after_ms == 1000)
);
}
#[test]
fn test_handle_error_503() {
let err = handle_error_response(503, &HeaderMap::new(), "");
assert!(matches!(err, ProviderError::Overloaded));
}
#[test]
fn test_handle_error_500_with_body() {
let err = handle_error_response(500, &HeaderMap::new(), "Internal Server Error");
match err {
ProviderError::Internal { status, message } => {
assert_eq!(status, 500);
assert_eq!(message, "Internal Server Error");
}
_ => panic!("Expected Internal error, got {:?}", err),
}
}
#[test]
fn test_handle_error_502() {
let err = handle_error_response(502, &HeaderMap::new(), "Bad Gateway");
match err {
ProviderError::Internal { status, message } => {
assert_eq!(status, 502);
assert_eq!(message, "Bad Gateway");
}
_ => panic!("Expected Internal error, got {:?}", err),
}
}
#[test]
fn test_handle_error_400_with_empty_body() {
let err = handle_error_response(400, &HeaderMap::new(), "");
match err {
ProviderError::Internal { status, message } => {
assert_eq!(status, 400);
assert!(message.is_empty());
}
_ => panic!("Expected Internal error, got {:?}", err),
}
}
#[tokio::test]
async fn test_build_request_success() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/test"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({"ok": true})))
.mount(&mock_server)
.await;
let client = reqwest::Client::new();
let body = serde_json::json!({"key": "value"});
let headers = vec![("X-Custom".to_owned(), "test-value".to_owned())];
let resp =
build_request(&client, &format!("{}/test", mock_server.uri()), body, headers, None)
.await
.unwrap();
assert!(resp.status().is_success());
}
#[tokio::test]
async fn test_build_request_sends_headers() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/test"))
.and(wiremock::matchers::header("Authorization", "Bearer sk-test"))
.respond_with(ResponseTemplate::new(200))
.mount(&mock_server)
.await;
let client = reqwest::Client::new();
let body = serde_json::json!({"prompt": "hello"});
let headers = vec![("Authorization".to_owned(), "Bearer sk-test".to_owned())];
let resp =
build_request(&client, &format!("{}/test", mock_server.uri()), body, headers, None)
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 200);
}
#[tokio::test]
async fn test_build_request_returns_error_on_401() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/test"))
.respond_with(ResponseTemplate::new(401).set_body_string("Unauthorized"))
.mount(&mock_server)
.await;
let client = reqwest::Client::new();
let body = serde_json::json!({"test": true});
let resp =
build_request(&client, &format!("{}/test", mock_server.uri()), body, vec![], None)
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 401);
}
#[tokio::test]
async fn test_build_request_returns_error_on_429() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/test"))
.respond_with(ResponseTemplate::new(429).set_body_string("Rate limited"))
.mount(&mock_server)
.await;
let client = reqwest::Client::new();
let body = serde_json::json!({"test": true});
let resp =
build_request(&client, &format!("{}/test", mock_server.uri()), body, vec![], None)
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 429);
}
#[tokio::test]
async fn test_build_request_invalid_url() {
let client = reqwest::Client::new();
let body = serde_json::json!({"test": true});
let result = build_request(
&client,
"http://invalid-host-that-does-not-exist.local/test",
body,
vec![],
None,
)
.await;
assert!(result.is_err());
assert!(
matches!(result.unwrap_err(), ProviderError::Network { .. }),
"Expected Network error"
);
}
#[tokio::test]
async fn test_sse_stream_basic_lines() {
let mock_server = MockServer::start().await;
let sse_body = "data: {\"delta\":\"Hello\"}\n\ndata: [DONE]\n\n";
Mock::given(method("POST"))
.and(path("/stream"))
.respond_with(ResponseTemplate::new(200).set_body_string(sse_body.to_owned()))
.mount(&mock_server)
.await;
let client = reqwest::Client::new();
let resp = client.post(format!("{}/stream", mock_server.uri())).send().await.unwrap();
assert!(resp.status().is_success());
let stream = create_sse_stream(resp, None);
let lines: Vec<String> =
stream.filter_map(|r| futures::future::ready(r.ok())).collect().await;
assert_eq!(lines, vec!["data: {\"delta\":\"Hello\"}", "", "data: [DONE]", ""]);
}
#[tokio::test]
async fn test_sse_stream_cancellation() {
let mock_server = MockServer::start().await;
let mut sse_body = String::new();
for _ in 0..100 {
sse_body.push_str("data: ping\n\n");
}
Mock::given(method("POST"))
.and(path("/stream-cancel"))
.respond_with(ResponseTemplate::new(200).set_body_string(sse_body))
.mount(&mock_server)
.await;
let client = reqwest::Client::new();
let resp =
client.post(format!("{}/stream-cancel", mock_server.uri())).send().await.unwrap();
let cancel = CancellationToken::new();
cancel.cancel();
let stream = create_sse_stream(resp, Some(cancel));
let lines: Vec<String> =
stream.filter_map(|r| futures::future::ready(r.ok())).collect().await;
assert!(
lines.len() < 5,
"Expected fewer than 5 lines with immediate cancel, got {}",
lines.len()
);
}
#[tokio::test]
async fn test_sse_stream_no_cancel() {
let mock_server = MockServer::start().await;
let sse_body = "data: {\"a\":1}\n\ndata: {\"a\":2}\n\n";
Mock::given(method("POST"))
.and(path("/stream-nocancel"))
.respond_with(ResponseTemplate::new(200).set_body_string(sse_body.to_owned()))
.mount(&mock_server)
.await;
let client = reqwest::Client::new();
let resp =
client.post(format!("{}/stream-nocancel", mock_server.uri())).send().await.unwrap();
let stream = create_sse_stream(resp, None);
let lines: Vec<String> =
stream.filter_map(|r| futures::future::ready(r.ok())).collect().await;
assert_eq!(lines, vec!["data: {\"a\":1}", "", "data: {\"a\":2}", ""]);
}
#[tokio::test]
async fn test_sse_stream_empty_body() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/stream-empty"))
.respond_with(ResponseTemplate::new(200))
.mount(&mock_server)
.await;
let client = reqwest::Client::new();
let resp = client.post(format!("{}/stream-empty", mock_server.uri())).send().await.unwrap();
let stream = create_sse_stream(resp, None);
let lines: Vec<String> =
stream.filter_map(|r| futures::future::ready(r.ok())).collect().await;
assert!(lines.is_empty());
}
}