mod common;
use std::sync::Arc;
use common::{MockResponse, MockServer};
use hey::config::{ProviderConfig, RetryConfig};
use hey::llm::build_provider;
use hey::llm::ir::{ChatRequest, EffortLevel, Message, Role};
use hey::llm::{Delta, LlmError};
fn sse_chunk(text: &str) -> String {
format!(
"data: {}\n\n",
serde_json::json!({"choices":[{"delta":{"content":text}}]})
)
}
fn request() -> ChatRequest {
ChatRequest {
messages: vec![Message::text(Role::User, "hi")],
tools: vec![],
effort: EffortLevel::Off,
}
}
fn fast_retry() -> RetryConfig {
RetryConfig {
enabled: true,
max_retries: 3,
base_delay_ms: 1,
max_delay_ms: 5,
jitter: false,
respect_retry_after: true,
}
}
fn provider_at(server: &MockServer, retry: &RetryConfig) -> Arc<dyn hey::llm::Provider> {
let cfg = ProviderConfig {
base_url: format!("http://{}/v1", server.addr),
api_key: None,
models: vec!["test-model".into()],
..Default::default()
};
build_provider(&cfg, reqwest::Client::new(), retry).unwrap()
}
#[tokio::test]
async fn rate_limit_then_success_retries() {
let server = MockServer::start(vec![
MockResponse::status(429, "{\"error\":\"rate limit\"}"),
MockResponse::sse(&format!("{}data: [DONE]\n\n", sse_chunk("recovered"))),
])
.await;
let provider = provider_at(&server, &fast_retry());
let mut deltas = Vec::new();
let completion = provider
.stream(&request(), &mut |d| deltas.push(d))
.await
.unwrap();
assert_eq!(completion.text, "recovered");
assert_eq!(deltas, vec![Delta::Text("recovered".into())]);
assert_eq!(server.count(), 2, "must hit the server twice");
}
#[tokio::test]
async fn persistent_rate_limit_gives_up() {
let server = MockServer::start(vec![MockResponse::status(429, "{}")]).await;
let provider = provider_at(&server, &fast_retry());
let err = provider.stream(&request(), &mut |_| {}).await.unwrap_err();
assert!(matches!(err, LlmError::RateLimit { .. }));
assert_eq!(server.count(), 1 + 3);
}
#[tokio::test]
async fn auth_error_never_retries() {
let server = MockServer::start(vec![MockResponse::status(401, "{}")]).await;
let provider = provider_at(&server, &fast_retry());
let err = provider.stream(&request(), &mut |_| {}).await.unwrap_err();
assert!(matches!(err, LlmError::Auth(_)));
assert_eq!(server.count(), 1, "401 must not be retried");
}
#[tokio::test]
async fn disabled_retry_single_attempt() {
let retry = RetryConfig {
enabled: false,
..fast_retry()
};
let server = MockServer::start(vec![MockResponse::status(429, "{}")]).await;
let provider = provider_at(&server, &retry);
let err = provider.stream(&request(), &mut |_| {}).await.unwrap_err();
assert!(matches!(err, LlmError::RateLimit { .. }));
assert_eq!(server.count(), 1);
}
#[tokio::test]
async fn partial_output_prevents_retry() {
let body = format!("{}{}", sse_chunk("partial "), "data: [DONE]\n\n");
let server = MockServer::start(vec![MockResponse::sse(&body), MockResponse::sse(&body)]).await;
let provider = provider_at(&server, &fast_retry());
let completion = provider.stream(&request(), &mut |_| {}).await.unwrap();
assert_eq!(completion.text, "partial ");
assert_eq!(server.count(), 1);
}