use std::future::Future;
use bytes::Bytes;
use crate::error::{LiterLlmError, Result};
use crate::http::retry;
pub(crate) fn retry_after_from_response(resp: &reqwest::Response) -> Option<std::time::Duration> {
let value = resp.headers().get(reqwest::header::RETRY_AFTER)?.to_str().ok()?;
retry::parse_retry_after(value)
}
async fn sleep_for_retry(delay: std::time::Duration) {
#[cfg(not(target_arch = "wasm32"))]
tokio::time::sleep(delay).await;
#[cfg(target_arch = "wasm32")]
gloo_timers::future::sleep(std::time::Duration::from_millis(delay.as_millis() as u64)).await;
}
pub(crate) async fn with_retry<F, Fut>(url: &str, max_retries: u32, mut send: F) -> Result<reqwest::Response>
where
F: FnMut() -> Fut,
Fut: Future<Output = std::result::Result<reqwest::Response, reqwest::Error>>,
{
crate::provider::validate_outbound_url(url).await?;
let mut attempt = 0u32;
loop {
let resp = match send().await {
Ok(resp) => resp,
Err(transport_error) => {
if let Some(policy_error) = crate::provider::outbound_forbidden_from_reqwest(&transport_error) {
return Err(policy_error);
}
if let Some(delay) = retry::should_retry_transport_error(attempt, max_retries) {
attempt += 1;
tracing::warn!(
error = %transport_error,
attempt,
max_retries,
"transport-level error sending request; retrying"
);
sleep_for_retry(delay).await;
continue;
}
return Err(LiterLlmError::from(transport_error));
}
};
let status = resp.status().as_u16();
if resp.status().is_success() {
return Ok(resp);
}
let server_retry_after = retry_after_from_response(&resp);
if let Some(delay) = retry::should_retry(status, attempt, max_retries, server_retry_after) {
attempt += 1;
sleep_for_retry(delay).await;
continue;
}
let text = resp
.text()
.await
.unwrap_or_else(|e| format!("(failed to read body: {e})"));
return Err(LiterLlmError::from_status(status, &text, server_retry_after));
}
}
#[tracing::instrument(
level = "debug",
skip_all,
fields(
http.method = "POST",
http.url = %url,
http.status_code = tracing::field::Empty,
http.retry_count = tracing::field::Empty,
)
)]
pub async fn post_json_raw(
client: &reqwest::Client,
url: &str,
auth_header: Option<(&str, &str)>,
extra_headers: &[(&str, &str)],
body: Bytes,
max_retries: u32,
) -> Result<serde_json::Value> {
let mut retry_count = 0u32;
let resp = with_retry(url, max_retries, || {
let mut builder = client
.post(url)
.header(reqwest::header::CONTENT_TYPE, "application/json")
.body(body.clone());
if let Some((name, value)) = auth_header {
builder = builder.header(name, value);
}
for (name, value) in extra_headers {
builder = builder.header(*name, *value);
}
retry_count += 1;
builder.send()
})
.await?;
{
let span = tracing::Span::current();
span.record("http.status_code", resp.status().as_u16());
span.record("http.retry_count", retry_count.saturating_sub(1));
}
resp.json::<serde_json::Value>().await.map_err(LiterLlmError::from)
}
#[tracing::instrument(
level = "debug",
skip_all,
fields(
http.method = "POST",
http.url = %url,
http.status_code = tracing::field::Empty,
http.retry_count = tracing::field::Empty,
)
)]
pub async fn post_binary(
client: &reqwest::Client,
url: &str,
auth_header: Option<(&str, &str)>,
extra_headers: &[(&str, &str)],
body: Bytes,
max_retries: u32,
) -> Result<Bytes> {
let mut retry_count = 0u32;
let resp = with_retry(url, max_retries, || {
let mut builder = client
.post(url)
.header(reqwest::header::CONTENT_TYPE, "application/json")
.body(body.clone());
if let Some((name, value)) = auth_header {
builder = builder.header(name, value);
}
for (name, value) in extra_headers {
builder = builder.header(*name, *value);
}
retry_count += 1;
builder.send()
})
.await?;
{
let span = tracing::Span::current();
span.record("http.status_code", resp.status().as_u16());
span.record("http.retry_count", retry_count.saturating_sub(1));
}
resp.bytes().await.map_err(LiterLlmError::from)
}
#[tracing::instrument(
level = "debug",
skip_all,
fields(
http.method = "POST",
http.url = %url,
http.status_code = tracing::field::Empty,
)
)]
pub async fn post_multipart(
client: &reqwest::Client,
url: &str,
auth_header: Option<(&str, &str)>,
extra_headers: &[(&str, &str)],
form: reqwest::multipart::Form,
) -> Result<serde_json::Value> {
crate::provider::validate_outbound_url(url).await?;
let mut builder = client.post(url).multipart(form);
if let Some((name, value)) = auth_header {
builder = builder.header(name, value);
}
for (name, value) in extra_headers {
builder = builder.header(*name, *value);
}
let resp = builder.send().await?;
{
let span = tracing::Span::current();
span.record("http.status_code", resp.status().as_u16());
}
let status = resp.status().as_u16();
if !resp.status().is_success() {
let server_retry_after = retry_after_from_response(&resp);
let text = resp
.text()
.await
.unwrap_or_else(|e| format!("(failed to read body: {e})"));
return Err(LiterLlmError::from_status(status, &text, server_retry_after));
}
resp.json::<serde_json::Value>().await.map_err(LiterLlmError::from)
}
#[tracing::instrument(
level = "debug",
skip_all,
fields(
http.method = "GET",
http.url = %url,
http.status_code = tracing::field::Empty,
http.retry_count = tracing::field::Empty,
)
)]
pub async fn get_json_raw(
client: &reqwest::Client,
url: &str,
auth_header: Option<(&str, &str)>,
extra_headers: &[(&str, &str)],
max_retries: u32,
) -> Result<serde_json::Value> {
let mut retry_count = 0u32;
let resp = with_retry(url, max_retries, || {
let mut builder = client.get(url);
if let Some((name, value)) = auth_header {
builder = builder.header(name, value);
}
for (name, value) in extra_headers {
builder = builder.header(*name, *value);
}
retry_count += 1;
builder.send()
})
.await?;
{
let span = tracing::Span::current();
span.record("http.status_code", resp.status().as_u16());
span.record("http.retry_count", retry_count.saturating_sub(1));
}
resp.json::<serde_json::Value>().await.map_err(LiterLlmError::from)
}
#[tracing::instrument(
level = "debug",
skip_all,
fields(
http.method = "DELETE",
http.url = %url,
http.status_code = tracing::field::Empty,
http.retry_count = tracing::field::Empty,
)
)]
pub async fn delete_json(
client: &reqwest::Client,
url: &str,
auth_header: Option<(&str, &str)>,
extra_headers: &[(&str, &str)],
max_retries: u32,
) -> Result<serde_json::Value> {
let mut retry_count = 0u32;
let resp = with_retry(url, max_retries, || {
let mut builder = client.delete(url);
if let Some((name, value)) = auth_header {
builder = builder.header(name, value);
}
for (name, value) in extra_headers {
builder = builder.header(*name, *value);
}
retry_count += 1;
builder.send()
})
.await?;
{
let span = tracing::Span::current();
span.record("http.status_code", resp.status().as_u16());
span.record("http.retry_count", retry_count.saturating_sub(1));
}
resp.json::<serde_json::Value>().await.map_err(LiterLlmError::from)
}
#[tracing::instrument(
level = "debug",
skip_all,
fields(
http.method = "GET",
http.url = %url,
http.status_code = tracing::field::Empty,
http.retry_count = tracing::field::Empty,
)
)]
pub async fn get_binary(
client: &reqwest::Client,
url: &str,
auth_header: Option<(&str, &str)>,
extra_headers: &[(&str, &str)],
max_retries: u32,
) -> Result<Bytes> {
let mut retry_count = 0u32;
let resp = with_retry(url, max_retries, || {
let mut builder = client.get(url);
if let Some((name, value)) = auth_header {
builder = builder.header(name, value);
}
for (name, value) in extra_headers {
builder = builder.header(*name, *value);
}
retry_count += 1;
builder.send()
})
.await?;
{
let span = tracing::Span::current();
span.record("http.status_code", resp.status().as_u16());
span.record("http.retry_count", retry_count.saturating_sub(1));
}
resp.bytes().await.map_err(LiterLlmError::from)
}
#[cfg(test)]
mod tests {
use std::io::{Read, Write};
use std::net::{SocketAddr, TcpListener};
use serial_test::serial;
use super::*;
use crate::provider::{OutboundPolicy, set_outbound_policy};
fn closed_port_url() -> String {
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind ephemeral port");
let port = listener.local_addr().expect("local_addr").port();
drop(listener);
format!("http://127.0.0.1:{port}/")
}
fn one_shot_json_server() -> (String, std::thread::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind JSON server");
let address = listener.local_addr().expect("JSON server address");
let handle = std::thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("accept JSON request");
let mut request = [0_u8; 4096];
let _ = stream.read(&mut request).expect("read JSON request");
stream
.write_all(b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 2\r\n\r\n{}")
.expect("write JSON response");
});
(format!("http://{address}/v1/chat/completions"), handle)
}
fn one_shot_server(response: String) -> (SocketAddr, std::thread::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind HTTP server");
let address = listener.local_addr().expect("HTTP server address");
let handle = std::thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("accept HTTP request");
let mut request = [0_u8; 4096];
let _ = stream.read(&mut request).expect("read HTTP request");
stream.write_all(response.as_bytes()).expect("write HTTP response");
});
(address, handle)
}
#[tokio::test]
#[serial(outbound_policy)]
async fn with_retry_retries_transport_errors_before_giving_up() {
let client = reqwest::Client::new();
let url = closed_port_url();
let mut attempts = 0u32;
let result = with_retry(&url, 2, || {
attempts += 1;
client.get(url.as_str()).send()
})
.await;
assert!(
result.is_err(),
"every attempt fails at the transport layer, so the final result must be Err"
);
assert_eq!(
attempts, 3,
"must attempt once plus 2 retries (matching max_retries) before giving up"
);
}
#[tokio::test]
#[serial(outbound_policy)]
async fn with_retry_does_not_retry_transport_errors_when_max_retries_is_zero() {
let client = reqwest::Client::new();
let url = closed_port_url();
let mut attempts = 0u32;
let result = with_retry(&url, 0, || {
attempts += 1;
client.get(url.as_str()).send()
})
.await;
assert!(result.is_err());
assert_eq!(
attempts, 1,
"max_retries = 0 must still make exactly one attempt and no retries"
);
}
#[tokio::test]
#[serial(outbound_policy)]
async fn post_json_raw_rejects_direct_private_url_under_deny_private() {
set_outbound_policy(OutboundPolicy::DenyPrivate);
let result = post_json_raw(
&reqwest::Client::new(),
&closed_port_url(),
None,
&[],
Bytes::from_static(b"{}"),
0,
)
.await;
set_outbound_policy(OutboundPolicy::Off);
assert!(
matches!(result, Err(LiterLlmError::OutboundForbidden { .. })),
"private request URL must be rejected by the policy before reqwest connects: {result:?}"
);
}
#[tokio::test]
#[serial(outbound_policy)]
async fn post_json_raw_allows_local_mock_when_policy_is_off() {
set_outbound_policy(OutboundPolicy::Off);
let (url, server) = one_shot_json_server();
let result = post_json_raw(&reqwest::Client::new(), &url, None, &[], Bytes::from_static(b"{}"), 0).await;
server.join().expect("JSON server thread");
assert_eq!(
result.expect("Off policy must preserve local mock access"),
serde_json::json!({})
);
}
#[tokio::test]
#[serial(outbound_policy)]
async fn post_multipart_rejects_direct_private_url_under_deny_private() {
set_outbound_policy(OutboundPolicy::DenyPrivate);
let result = post_multipart(
&reqwest::Client::new(),
&closed_port_url(),
None,
&[],
reqwest::multipart::Form::new(),
)
.await;
set_outbound_policy(OutboundPolicy::Off);
assert!(
matches!(result, Err(LiterLlmError::OutboundForbidden { .. })),
"private multipart URL must be rejected before reqwest connects: {result:?}"
);
}
#[tokio::test]
#[serial(outbound_policy)]
async fn post_multipart_preserves_redirect_policy_error_classification() {
set_outbound_policy(OutboundPolicy::Off);
let (source_address, source) = one_shot_server(
"HTTP/1.1 302 Found\r\nLocation: http://169.254.169.254/token\r\nContent-Length: 0\r\n\r\n".to_owned(),
);
let source_url = format!("http://multipart.test:{}/upload", source_address.port());
let client = crate::provider::configure_outbound_client_builder(
reqwest::Client::builder().resolve("multipart.test", source_address),
None,
)
.build()
.expect("multipart redirect client");
set_outbound_policy(OutboundPolicy::Allowlist(vec![
url::Url::parse(&source_url).expect("source allowlist URL"),
]));
let result = post_multipart(&client, &source_url, None, &[], reqwest::multipart::Form::new()).await;
set_outbound_policy(OutboundPolicy::Off);
source.join().expect("multipart source server");
assert!(
matches!(result, Err(LiterLlmError::OutboundForbidden { .. })),
"multipart redirect policy errors must remain non-transient: {result:?}"
);
}
}