mod common;
use std::sync::Arc;
use common::{MockResponse, MockServer};
use hey::config::ProviderConfig;
use hey::llm::build_provider;
use hey::llm::ir::{ChatRequest, EffortLevel, Message, Role, ToolCall};
use hey::llm::{Delta, LlmError};
fn sse_part(parts: &[serde_json::Value]) -> String {
format!(
"data: {}\n\n",
serde_json::json!({"candidates": [{"content": {"role": "model", "parts": parts}}]})
)
}
fn text_part(t: &str) -> serde_json::Value {
serde_json::json!({"text": t})
}
fn thought_part(t: &str) -> serde_json::Value {
serde_json::json!({"text": t, "thought": true})
}
fn fn_call_part(name: &str, args: serde_json::Value) -> serde_json::Value {
serde_json::json!({"functionCall": {"name": name, "args": args}})
}
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://{}/v1beta", server.addr),
api_key: api_key.map(str::to_string),
models: vec!["gemini-2.5-flash".into()],
protocol: Some("google".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_and_thinking_from_sse() {
let body = format!(
"{}{}{}data: [DONE]\n\n",
sse_part(&[thought_part("reasoning step")]),
sse_part(&[text_part("Hel")]),
sse_part(&[text_part("lo")]),
);
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, "Hello");
assert_eq!(
deltas,
vec![
Delta::Thinking("reasoning step".into()),
Delta::Text("Hel".into()),
Delta::Text("lo".into())
]
);
}
#[tokio::test]
async fn no_api_key_sends_no_google_key_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!(
req.path.contains(":streamGenerateContent?alt=sse"),
"端点: {}",
req.path
);
assert!(
req.path.contains("gemini-2.5-flash"),
"模型在路径: {}",
req.path
);
assert!(req.header("x-goog-api-key").is_none(), "无认证时不发 key");
}
#[tokio::test]
async fn api_key_sends_google_key_header() {
let server = MockServer::start(vec![MockResponse::sse("data: [DONE]\n\n")]).await;
let provider = provider_at(&server, Some("GOOG_KEY_1"));
provider.stream(&request(), &mut |_| {}).await.unwrap();
let req = server.request(0).unwrap();
assert_eq!(req.header("x-goog-api-key"), Some("GOOG_KEY_1"));
}
#[tokio::test]
async fn function_call_streamed_and_merged() {
let body = format!(
"{}{}data: [DONE]\n\n",
sse_part(&[fn_call_part(
"bash",
serde_json::json!({"cmd": "ls", "flag": "-la"})
)]),
sse_part(&[fn_call_part(
"bash",
serde_json::json!({"cmd": "ls", "extra": true})
)]),
);
let server = MockServer::start(vec![MockResponse::sse(&body)]).await;
let provider = provider_at(&server, Some("k"));
let completion = provider.stream(&request(), &mut |_| {}).await.unwrap();
assert_eq!(completion.tool_calls.len(), 1);
let call = &completion.tool_calls[0];
assert_eq!(call.name, "bash");
let args: serde_json::Value = serde_json::from_str(&call.arguments).unwrap();
assert_eq!(args["cmd"], "ls");
assert_eq!(args["flag"], "-la");
assert_eq!(args["extra"], true);
assert!(!call.id.is_empty(), "适配器生成稳定 id");
}
#[tokio::test]
async fn multi_turn_tool_result_uses_function_response_with_name() {
let turn1 = sse_part(&[fn_call_part("bash", serde_json::json!({"cmd": "ls"}))]);
let server = MockServer::start(vec![
MockResponse::sse(&format!("{turn1}data: [DONE]\n\n")),
MockResponse::sse("data: [DONE]\n\n"),
])
.await;
let provider = provider_at(&server, Some("k"));
let completion = provider.stream(&request(), &mut |_| {}).await.unwrap();
assert_eq!(completion.tool_calls.len(), 1);
let call = completion.tool_calls[0].clone();
let req2 = ChatRequest {
messages: vec![
Message::text(Role::User, "list files"),
Message {
id: String::new(),
role: Role::Assistant,
content: vec![],
tool_calls: vec![ToolCall {
id: call.id.clone(),
name: call.name.clone(),
arguments: call.arguments.clone(),
}],
tool_call_id: None,
},
Message::tool_result(&call.id, "src/"),
],
tools: vec![],
effort: EffortLevel::Off,
};
provider.stream(&req2, &mut |_| {}).await.unwrap();
let req = server.request(1).unwrap();
let body: serde_json::Value = serde_json::from_str(&req.body).unwrap();
let last = body["contents"].as_array().unwrap().last().unwrap();
assert_eq!(last["role"], "user");
assert_eq!(last["parts"][0]["functionResponse"]["name"], "bash");
assert_eq!(
last["parts"][0]["functionResponse"]["response"]["result"],
"src/"
);
}
#[tokio::test]
async fn request_body_carries_system_tools_and_thinking() {
let server = MockServer::start(vec![MockResponse::sse("data: [DONE]\n\n")]).await;
let provider = provider_at(&server, Some("k"));
let req = ChatRequest {
messages: vec![
Message::text(Role::System, "SYS"),
Message::text(Role::User, "hi"),
],
tools: vec![hey::llm::ir::ToolSchema {
name: "bash".into(),
description: "Run a command".into(),
parameters: serde_json::json!({"type": "object"}),
}],
effort: EffortLevel::High,
};
provider.stream(&req, &mut |_| {}).await.unwrap();
let body: serde_json::Value = serde_json::from_str(&server.request(0).unwrap().body).unwrap();
assert_eq!(body["systemInstruction"]["parts"][0]["text"], "SYS");
assert_eq!(body["tools"][0]["functionDeclarations"][0]["name"], "bash");
assert_eq!(
body["generationConfig"]["thinkingConfig"]["thinkingBudget"],
8192 );
}
#[tokio::test]
async fn rate_limit_returns_error_with_retry_after() {
let mut resp = MockResponse::status(
429,
r#"{"error":{"code":429,"message":"RESOURCE_EXHAUSTED","status":"RESOURCE_EXHAUSTED"}}"#,
);
resp.headers.push(("Retry-After".into(), "5".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(5));
}
other => panic!("expected RateLimit, got {other:?}"),
}
}
#[tokio::test]
async fn auth_error_classified_from_body() {
let server = MockServer::start(vec![MockResponse::status(
403,
r#"{"error":{"code":403,"message":"API key not valid","status":"PERMISSION_DENIED"}}"#,
)])
.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 partial_stream_without_done_google_errors() {
let body = format!(
"data: {}\n\n",
serde_json::json!({"candidates":[{"content":{"parts":[{"text":"hello"}]}}]})
);
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_google_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 body = format!("{}data: [DONE]\n\n", sse_part(&[text_part("hi")]));
let server = MockServer::start(vec![MockResponse::sse(&body)]).await;
let cfg = ProviderConfig {
base_url: format!("http://{}/v1beta/", server.addr),
protocol: Some("google".into()),
models: vec!["gemini-test".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!(
req.path.starts_with("/v1beta/models/"),
"无双斜杠: {}",
req.path
);
assert!(
!req.path.contains("//models"),
"尾斜杠未规范化: {}",
req.path
);
}