mod common;
use std::sync::Arc;
use common::{MockResponse, MockServer};
use hey::config::ProviderConfig;
use hey::llm::build_provider;
use hey::llm::ir::{ChatRequest, ContentBlock, 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 provider_at(server: &MockServer, api_key: Option<&str>) -> Arc<dyn hey::llm::Provider> {
let cfg = ProviderConfig {
base_url: format!("http://{}/v1", server.addr),
api_key: api_key.map(str::to_string),
models: vec!["test-model".into()],
..Default::default()
};
let retry = hey::config::RetryConfig {
enabled: false,
..Default::default()
};
build_provider(&cfg, reqwest::Client::new(), &retry).unwrap()
}
#[tokio::test]
async fn streams_text_from_sse() {
let server = MockServer::start(vec![MockResponse::sse(
&format!("{}data: [DONE]\n\n", sse_chunk("Hel") + &sse_chunk("lo"))[..],
)])
.await;
let provider = provider_at(&server, None);
let mut deltas = Vec::new();
let completion = provider
.stream(&request(), &mut |d| deltas.push(d))
.await
.unwrap();
assert_eq!(completion.text, "Hello");
assert_eq!(
deltas,
vec![Delta::Text("Hel".into()), Delta::Text("lo".into())]
);
}
#[tokio::test]
async fn no_api_key_sends_no_auth_header() {
let server = MockServer::start(vec![MockResponse::sse("data: [DONE]\n\n")]).await;
let provider = provider_at(&server, None);
provider.stream(&request(), &mut |_| {}).await.unwrap();
assert_eq!(server.count(), 1);
let req = server.request(0).unwrap();
assert_eq!(req.path, "/v1/chat/completions");
assert!(
req.header("authorization").is_none(),
"no auth header expected"
);
assert!(req.body.contains("\"model\":\"test-model\""));
assert!(req.body.contains("\"stream\":true"));
}
#[tokio::test]
async fn api_key_sends_bearer() {
let server = MockServer::start(vec![MockResponse::sse("data: [DONE]\n\n")]).await;
let provider = provider_at(&server, Some("sk-123"));
provider.stream(&request(), &mut |_| {}).await.unwrap();
let req = server.request(0).unwrap();
assert_eq!(req.header("authorization"), Some("Bearer sk-123"));
}
#[tokio::test]
async fn reasoning_content_is_sent_back_on_next_turn() {
let reasoning = "thinking step-by-step";
let turn1 = format!(
"data: {}\n\ndata: {}\n\ndata: [DONE]\n\n",
serde_json::json!({"choices": [{"delta": {"reasoning_content": reasoning}}]}),
serde_json::json!({"choices": [{"delta": {"tool_calls": [{"index": 0, "id": "c1", "function": {"name": "bash", "arguments": "{\"cmd\":\"ls\"}"}}]}}]}),
);
let turn2 = format!("{}data: [DONE]\n\n", sse_chunk("done"));
let server =
MockServer::start(vec![MockResponse::sse(&turn1), MockResponse::sse(&turn2)]).await;
let provider = provider_at(&server, None);
let mut req = request();
req.effort = EffortLevel::Medium;
let c = provider.stream(&req, &mut |_| {}).await.unwrap();
assert_eq!(c.thinking, reasoning, "reasoning 必须累积进 completion");
assert_eq!(c.tool_calls.len(), 1);
let mut assistant = Message::text(Role::Assistant, c.text.clone());
assistant
.content
.insert(0, ContentBlock::Thinking(c.thinking.clone()));
assistant.tool_calls = c.tool_calls.clone();
let tool_result = Message::tool_result(&c.tool_calls[0].id, "ok");
let req2 = ChatRequest {
messages: vec![assistant, tool_result],
tools: vec![],
effort: EffortLevel::Medium,
};
provider.stream(&req2, &mut |_| {}).await.unwrap();
let body2 = &server.request(1).unwrap().body;
assert!(
body2.contains(&format!("\"reasoning_content\":\"{reasoning}\"")),
"第二轮必须回传 reasoning_content: {body2}"
);
}
#[tokio::test]
async fn no_thinking_sends_no_reasoning_field() {
let server = MockServer::start(vec![MockResponse::sse("data: [DONE]\n\n")]).await;
let provider = provider_at(&server, None);
let req = ChatRequest {
messages: vec![Message::text(Role::Assistant, "plain reply")],
tools: vec![],
effort: EffortLevel::Off,
};
provider.stream(&req, &mut |_| {}).await.unwrap();
assert!(
!server
.request(0)
.unwrap()
.body
.contains("reasoning_content")
);
}
#[tokio::test]
async fn effort_passthrough_sets_reasoning_effort() {
let server = MockServer::start(vec![MockResponse::sse("data: [DONE]\n\n")]).await;
let provider = provider_at(&server, None);
let req = ChatRequest {
messages: vec![Message::text(Role::User, "hi")],
tools: vec![],
effort: EffortLevel::Minimal,
};
provider.stream(&req, &mut |_| {}).await.unwrap();
let body = &server.request(0).unwrap().body;
assert!(
body.contains("\"reasoning_effort\":\"minimal\""),
"effort 原值透传: {body}"
);
let server2 = MockServer::start(vec![MockResponse::sse("data: [DONE]\n\n")]).await;
let provider2 = provider_at(&server2, None);
let req2 = ChatRequest {
effort: EffortLevel::Off,
..req
};
provider2.stream(&req2, &mut |_| {}).await.unwrap();
assert!(
server2
.request(0)
.unwrap()
.body
.contains("\"reasoning_effort\":\"none\"")
);
}
#[tokio::test]
async fn content_filter_finish_reason_is_an_error() {
let filtered = format!(
"{}data: {}\n\n[DONE]",
sse_chunk("part of a bad answer"),
serde_json::json!({"choices": [{"delta": {}, "finish_reason": "content_filter"}]}),
);
let server = MockServer::start(vec![MockResponse::sse(&filtered)]).await;
let provider = provider_at(&server, None);
let err = provider.stream(&request(), &mut |_| {}).await.unwrap_err();
assert!(
err.to_string().contains("content_filter"),
"content_filter 必须转错误: {err}"
);
}
#[tokio::test]
async fn rate_limit_returns_error_with_retry_after() {
let mut resp = MockResponse::status(429, "{\"error\":\"rate limit\"}");
resp.headers.push(("Retry-After".into(), "7".into()));
let server = MockServer::start(vec![resp]).await;
let provider = provider_at(&server, Some("k"));
let err = provider.stream(&request(), &mut |_| {}).await.unwrap_err();
match err {
LlmError::RateLimit { retry_after } => {
assert_eq!(retry_after.map(|d| d.as_secs()), Some(7));
}
other => panic!("expected RateLimit, got {other:?}"),
}
}
#[tokio::test]
async fn auth_error_classified() {
let server = MockServer::start(vec![MockResponse::status(
401,
"{\"error\":\"invalid key\"}",
)])
.await;
let provider = provider_at(&server, Some("bad"));
let err = provider.stream(&request(), &mut |_| {}).await.unwrap_err();
assert!(matches!(err, LlmError::Auth(_)));
}
#[tokio::test]
async fn thinking_delta_forwarded() {
let body = format!(
"data: {}\n\ndata: {}\n\ndata: [DONE]\n\n",
serde_json::json!({"choices":[{"delta":{"reasoning_content":"step1"}}]}),
serde_json::json!({"choices":[{"delta":{"content":"answer"}}]})
);
let server = MockServer::start(vec![MockResponse::sse(&body)]).await;
let provider = provider_at(&server, None);
let mut deltas = Vec::new();
let completion = provider
.stream(&request(), &mut |d| deltas.push(d))
.await
.unwrap();
assert_eq!(completion.text, "answer");
assert_eq!(
deltas,
vec![
Delta::Thinking("step1".into()),
Delta::Text("answer".into())
]
);
}
#[tokio::test]
async fn partial_stream_without_done_errors() {
let body = format!(
"data: {}\n\ndata: {}\n\n",
serde_json::json!({"choices":[{"delta":{"content":"hello"}}]}),
serde_json::json!({"choices":[{"delta":{"content":" world"}}]})
);
let server = MockServer::start(vec![MockResponse::sse(&body)]).await;
let provider = provider_at(&server, None);
let err = provider.stream(&request(), &mut |_| {}).await.unwrap_err();
assert!(
matches!(err, LlmError::StreamInterrupted(_)),
"expected StreamInterrupted: {err}"
);
}
#[tokio::test]
async fn empty_stream_no_done_ok() {
let server = MockServer::start(vec![MockResponse::sse("")]).await;
let provider = provider_at(&server, None);
let completion = provider.stream(&request(), &mut |_| {}).await.unwrap();
assert!(completion.text.is_empty());
assert!(completion.tool_calls.is_empty());
}
#[tokio::test]
async fn html_gateway_response_errors() {
let server = MockServer::start(vec![MockResponse::sse(
"<html><body>502 Bad Gateway</body></html>",
)])
.await;
let provider = provider_at(&server, None);
let err = provider.stream(&request(), &mut |_| {}).await.unwrap_err();
assert!(
err.to_string().contains("not an SSE event stream"),
"错误应指明清:{}",
err
);
}
#[tokio::test]
async fn base_url_trailing_slash_normalized() {
let server = MockServer::start(vec![MockResponse::sse("data: [DONE]\n\n")]).await;
let cfg = ProviderConfig {
base_url: format!("http://{}/v1/", server.addr),
models: vec!["test-model".into()],
..Default::default()
};
let provider = build_provider(
&cfg,
reqwest::Client::new(),
&hey::config::RetryConfig {
enabled: false,
..Default::default()
},
)
.unwrap();
let _ = provider.stream(&request(), &mut |_| {}).await;
let req = server.request(0).unwrap();
assert_eq!(
req.path, "/v1/chat/completions",
"尾斜杠规范化: {}",
req.path
);
}