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};
use hey::llm::{Delta, LlmError};
fn ev(line: &str) -> String {
format!("data: {line}\n\n")
}
fn request() -> ChatRequest {
ChatRequest {
messages: vec![Message::text(Role::User, "hi")],
tools: vec![],
effort: EffortLevel::Medium,
}
}
fn provider_at(server: &MockServer) -> Arc<dyn hey::llm::Provider> {
let cfg = ProviderConfig {
protocol: Some("responses".into()),
base_url: format!("http://{}/v1", server.addr),
api_key: Some("test-key".into()),
models: vec!["o3".into()],
effort: Some(EffortLevel::High),
..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_thinking_and_tool_call() {
let body = ev(r#"{"type":"response.created","response":{}}"#)
+ &ev(
r#"{"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","call_id":"fc_1","name":"bash","arguments":""}}"#,
)
+ &ev(
r#"{"type":"response.function_call_arguments.delta","output_index":0,"delta":"{\"cmd\""}"#,
)
+ &ev(
r#"{"type":"response.function_call_arguments.delta","output_index":0,"delta":":\"ls\"}"}"#,
)
+ &ev(
r#"{"type":"response.reasoning_summary_text.delta","item_id":"r1","output_index":0,"delta":"thinking-out"}"#,
)
+ &ev(
r#"{"type":"response.output_text.delta","item_id":"m1","output_index":1,"delta":"answer"}"#,
)
+ &ev(r#"{"type":"response.completed","response":{"id":"r1"}}"#);
let server = MockServer::start(vec![MockResponse::sse(&body)]).await;
let provider = provider_at(&server);
let mut deltas = Vec::new();
let completion = provider
.stream(&request(), &mut |d| deltas.push(d))
.await
.unwrap();
assert_eq!(completion.text, "answer");
assert_eq!(completion.thinking, "thinking-out");
assert_eq!(completion.tool_calls.len(), 1);
let tc = &completion.tool_calls[0];
assert_eq!(tc.id, "fc_1");
assert_eq!(tc.name, "bash");
assert_eq!(tc.arguments, r#"{"cmd":"ls"}"#);
assert_eq!(
deltas,
vec![
Delta::Thinking("thinking-out".into()),
Delta::Text("answer".into()),
]
);
let reqs = server.requests.lock().unwrap();
assert_eq!(reqs[0].path, "/v1/responses");
let body: serde_json::Value = serde_json::from_str(&reqs[0].body).unwrap();
assert_eq!(body["model"], "o3");
assert_eq!(
body["reasoning"]["effort"], "medium",
"effort 直通(请求级)"
);
assert_eq!(body["input"][0]["role"], "user");
assert_eq!(body["input"][0]["content"][0]["type"], "input_text");
}
#[tokio::test]
async fn streams_error_on_missing_completed() {
let body = ev(r#"{"type":"response.output_text.delta","delta":"half"}"#);
let server = MockServer::start(vec![MockResponse::sse(&body)]).await;
let provider = provider_at(&server);
let e = provider
.stream(&request(), &mut |_| {})
.await
.expect_err("断流应报错");
assert!(matches!(e, LlmError::StreamInterrupted(_)), "{e}");
}
#[tokio::test]
async fn streams_error_on_non_sse_response() {
let server = MockServer::start(vec![MockResponse::status(
200,
"<html>gateway error</html>",
)])
.await;
let provider = provider_at(&server);
let e = provider
.stream(&request(), &mut |_| {})
.await
.expect_err("非 SSE 应报错");
assert!(matches!(e, LlmError::StreamInterrupted(_)), "{e}");
}
#[tokio::test]
async fn streams_error_on_http_429() {
let server = MockServer::start(vec![MockResponse::status(
429,
r#"{"error":{"message":"rate limited"}}"#,
)])
.await;
let provider = provider_at(&server);
let e = provider
.stream(&request(), &mut |_| {})
.await
.expect_err("429 应报 RateLimit");
assert!(matches!(e, LlmError::RateLimit { .. }), "{e}");
}