use std::time::Duration;
use anyhow::{Context, Result};
use reqwest::Client;
use crate::config::schema::LlmSection;
#[derive(Debug, Clone)]
pub struct Message {
pub role: String,
pub content: String,
}
impl Message {
pub fn system(content: impl Into<String>) -> Self {
Self {
role: "system".into(),
content: content.into(),
}
}
pub fn user(content: impl Into<String>) -> Self {
Self {
role: "user".into(),
content: content.into(),
}
}
pub fn assistant(content: impl Into<String>) -> Self {
Self {
role: "assistant".into(),
content: content.into(),
}
}
}
#[allow(async_fn_in_trait)]
pub trait LlmProvider: Send + Sync {
async fn complete(&self, messages: &[Message]) -> Result<String> {
let chunks = self.complete_stream(messages).await?;
Ok(chunks.concat())
}
async fn complete_stream(&self, messages: &[Message]) -> Result<Vec<String>> {
let _ = messages;
Err(anyhow::anyhow!("streaming not supported"))
}
async fn complete_with_budget(
&self,
messages: &[Message],
_max_output_tokens: Option<u32>,
) -> Result<String> {
self.complete(messages).await
}
fn call_count(&self) -> usize {
0
}
}
pub enum Provider {
OpenAi(OpenAiProvider),
Anthropic(AnthropicProvider),
Mock(MockProvider),
}
impl LlmProvider for Provider {
async fn complete(&self, messages: &[Message]) -> Result<String> {
match self {
Provider::OpenAi(p) => p.complete(messages).await,
Provider::Anthropic(p) => p.complete(messages).await,
Provider::Mock(p) => p.complete(messages).await,
}
}
async fn complete_stream(&self, messages: &[Message]) -> Result<Vec<String>> {
match self {
Provider::OpenAi(p) => p.complete_stream(messages).await,
Provider::Anthropic(p) => p.complete_stream(messages).await,
Provider::Mock(p) => p.complete_stream(messages).await,
}
}
async fn complete_with_budget(
&self,
messages: &[Message],
max_output_tokens: Option<u32>,
) -> Result<String> {
match self {
Provider::OpenAi(p) => p.complete_with_budget(messages, max_output_tokens).await,
Provider::Anthropic(p) => p.complete_with_budget(messages, max_output_tokens).await,
Provider::Mock(p) => p.complete_with_budget(messages, max_output_tokens).await,
}
}
fn call_count(&self) -> usize {
match self {
Provider::OpenAi(p) => p.call_count(),
Provider::Anthropic(p) => p.call_count(),
Provider::Mock(p) => p.call_count(),
}
}
}
pub(crate) const MAX_RETRIES: u32 = 3;
fn backoff_delay(attempt: u32) -> Duration {
let base = 500u64 * 2u64.pow(attempt);
let jitter = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| u64::from(d.subsec_nanos()) % 251)
.unwrap_or(0);
Duration::from_millis(base + jitter)
}
fn is_retryable_status(status: reqwest::StatusCode) -> bool {
status == reqwest::StatusCode::TOO_MANY_REQUESTS || status.is_server_error()
}
pub(crate) async fn retry_with_backoff<F, Fut>(
max_retries: u32,
send_fn: F,
) -> Result<reqwest::Response>
where
F: Fn() -> Fut,
Fut: std::future::Future<Output = Result<reqwest::Response, reqwest::Error>> + Send,
{
let mut last_error = None;
for attempt in 0..max_retries {
if attempt > 0 {
let delay = backoff_delay(attempt - 1);
tracing::info!(
"LLM 请求重试(第 {}/{} 次),退避 {}ms",
attempt + 1,
max_retries,
delay.as_millis()
);
tokio::time::sleep(delay).await;
}
match send_fn().await {
Ok(resp) if is_retryable_status(resp.status()) => {
let status = resp.status();
let text = resp.text().await.unwrap_or_default();
tracing::warn!(
"LLM API 返回可重试状态 {}(第 {} 次尝试): {}",
status,
attempt + 1,
text.chars().take(2000).collect::<String>()
);
last_error = Some(anyhow::anyhow!("API 返回错误 ({}): {}", status, text));
}
Ok(resp) => return Ok(resp),
Err(e) if e.is_timeout() || e.is_connect() => {
tracing::warn!(
"LLM 请求超时/连接失败(第 {} 次尝试,将重试): {}",
attempt + 1,
e
);
last_error = Some(anyhow::anyhow!("请求失败: {}", e));
}
Err(e) => return Err(anyhow::anyhow!("请求失败: {}", e)),
}
}
tracing::error!("LLM API 调用重试 {} 次后全部失败: {:?}", max_retries, last_error);
Err(anyhow::anyhow!(
"LLM API 调用重试 {} 次后全部失败: {:?}",
max_retries,
last_error
))
}
#[cfg(test)]
fn parse_sse_stream(
bytes: &[u8],
line_prefix: &str,
extract: impl Fn(&serde_json::Value) -> Option<String>,
) -> Vec<String> {
let mut chunks = Vec::new();
for line in String::from_utf8_lossy(bytes).split('\n') {
let line = line.trim();
if line.is_empty() {
continue;
}
let Some(json_str) = line.strip_prefix(line_prefix) else { continue };
if let Ok(val) = serde_json::from_str::<serde_json::Value>(json_str)
&& let Some(text) = extract(&val)
{
chunks.push(text);
}
}
chunks
}
const SSE_IDLE_TIMEOUT: Duration = Duration::from_secs(60);
async fn collect_sse(
resp: reqwest::Response,
line_prefix: &str,
extract: impl Fn(&serde_json::Value) -> Option<String>,
) -> Result<Vec<String>> {
use futures::StreamExt;
let mut stream = resp.bytes_stream();
let mut buf: Vec<u8> = Vec::new();
let mut chunks = Vec::new();
tracing::info!("SSE 流开始消费(空闲超时保护 {}s)", SSE_IDLE_TIMEOUT.as_secs());
loop {
let item = match tokio::time::timeout(SSE_IDLE_TIMEOUT, stream.next()).await {
Ok(Some(Ok(item))) => item,
Ok(Some(Err(e))) => return Err(e.into()),
Ok(None) => break,
Err(_) => {
tracing::warn!(
"SSE 流读取空闲超时({}s 无数据,已收 {} 个 chunk,模型可能已停止产出)",
SSE_IDLE_TIMEOUT.as_secs(),
chunks.len()
);
anyhow::bail!(
"SSE 流读取空闲超时({}s 无数据,模型可能已停止产出)",
SSE_IDLE_TIMEOUT.as_secs()
)
}
};
buf.extend_from_slice(&item);
let mut consumed = 0usize;
for (idx, b) in buf.iter().enumerate() {
if *b != b'\n' {
continue;
}
let line = String::from_utf8_lossy(&buf[consumed..idx]);
let line = line.trim();
if !line.is_empty()
&& let Some(json_str) = line.strip_prefix(line_prefix)
&& let Ok(val) = serde_json::from_str::<serde_json::Value>(json_str)
&& let Some(text) = extract(&val)
{
chunks.push(text);
}
consumed = idx + 1;
}
buf.drain(..consumed);
}
if !buf.is_empty() {
let line = String::from_utf8_lossy(&buf);
if let Some(json_str) = line.trim().strip_prefix(line_prefix)
&& let Ok(val) = serde_json::from_str::<serde_json::Value>(json_str)
&& let Some(text) = extract(&val)
{
chunks.push(text);
}
}
tracing::info!(
"SSE 流消费完成: {} 个 chunk, {} 字符",
chunks.len(),
chunks.iter().map(|c| c.len()).sum::<usize>()
);
Ok(chunks)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OpenAiProtocol {
Responses,
Chat,
}
pub struct OpenAiProvider {
client: Client,
api_key: String,
model: String,
base_url: String,
protocol: OpenAiProtocol,
max_retries: u32,
max_tokens: Option<u32>,
temperature: Option<f32>,
call_count: std::sync::atomic::AtomicUsize,
}
impl OpenAiProvider {
pub fn new(config: &LlmSection, protocol: OpenAiProtocol) -> Result<Self> {
let api_key = config.api_key.clone()
.or_else(|| std::env::var(&config.api_key_env).ok())
.with_context(|| {
format!(
"LLM API Key 未设置(api_key 为空且环境变量 {} 未定义)。请设置环境变量 {},或编辑配置文件的 [llm] 段填入 api_key",
config.api_key_env, config.api_key_env
)
})?;
let base_url = config
.base_url
.clone()
.unwrap_or_else(|| "https://api.openai.com/v1".to_string());
let client = Client::builder()
.build()
.context("创建 HTTP 客户端失败")?;
Ok(Self {
client,
api_key,
model: config.model.clone(),
base_url,
protocol,
max_retries: MAX_RETRIES,
max_tokens: None,
temperature: None,
call_count: std::sync::atomic::AtomicUsize::new(0),
})
}
}
impl OpenAiProvider {
fn build_chat_body(
&self,
messages: &[Message],
stream: bool,
max_tokens_override: Option<u32>,
) -> serde_json::Value {
let mut body = serde_json::json!({
"model": self.model,
"messages": messages.iter().map(|m| {
serde_json::json!({"role": m.role, "content": m.content})
}).collect::<Vec<_>>(),
});
if stream {
body["stream"] = serde_json::json!(true);
}
if let Some(maxt) = max_tokens_override.or(self.max_tokens) {
body["max_tokens"] = serde_json::json!(maxt);
}
if let Some(temp) = self.temperature {
body["temperature"] = serde_json::json!(temp);
}
body
}
fn build_responses_body(
&self,
messages: &[Message],
stream: bool,
max_output_tokens_override: Option<u32>,
) -> serde_json::Value {
let system = messages.iter().find(|m| m.role == "system").map(|m| &m.content);
let input: Vec<serde_json::Value> = messages
.iter()
.filter(|m| m.role != "system")
.map(|m| {
serde_json::json!({
"role": if m.role == "user" { "user" } else { "assistant" },
"content": serde_json::json!([{ "type": "input_text", "text": m.content }]),
})
})
.collect();
let mut body = serde_json::json!({
"model": self.model,
"input": input,
});
if let Some(s) = system {
body["instructions"] = serde_json::json!(s);
}
if stream {
body["stream"] = serde_json::json!(true);
}
if let Some(maxt) = max_output_tokens_override.or(self.max_tokens) {
body["max_output_tokens"] = serde_json::json!(maxt);
}
if let Some(temp) = self.temperature {
body["temperature"] = serde_json::json!(temp);
}
body
}
}
impl LlmProvider for OpenAiProvider {
async fn complete_stream(&self, messages: &[Message]) -> Result<Vec<String>> {
self.call_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
match self.protocol {
OpenAiProtocol::Chat => self.chat_complete_stream(messages, None).await,
OpenAiProtocol::Responses => self.responses_complete_stream(messages, None).await,
}
}
async fn complete_with_budget(
&self,
messages: &[Message],
max_output_tokens: Option<u32>,
) -> Result<String> {
self.call_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let chunks = match self.protocol {
OpenAiProtocol::Chat => self.chat_complete_stream(messages, max_output_tokens).await?,
OpenAiProtocol::Responses => {
self.responses_complete_stream(messages, max_output_tokens).await?
}
};
Ok(chunks.concat())
}
fn call_count(&self) -> usize {
self.call_count.load(std::sync::atomic::Ordering::Relaxed)
}
}
impl OpenAiProvider {
async fn chat_complete_stream(
&self,
messages: &[Message],
max_tokens_override: Option<u32>,
) -> Result<Vec<String>> {
let url = format!("{}/chat/completions", self.base_url);
let body = self.build_chat_body(messages, true, max_tokens_override);
let resp = retry_with_backoff(self.max_retries, || {
self.client
.post(&url)
.bearer_auth(&self.api_key)
.json(&body)
.send()
})
.await?;
if !resp.status().is_success() {
let status = resp.status();
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("API 返回错误 ({}): {}", status, text);
}
collect_sse(resp, "data: ", |v| {
v["choices"][0]["delta"]["content"]
.as_str()
.map(|s| s.to_string())
})
.await
}
async fn responses_complete_stream(
&self,
messages: &[Message],
max_output_tokens_override: Option<u32>,
) -> Result<Vec<String>> {
let url = format!("{}/responses", self.base_url);
let body = self.build_responses_body(messages, true, max_output_tokens_override);
let resp = retry_with_backoff(self.max_retries, || {
self.client
.post(&url)
.bearer_auth(&self.api_key)
.json(&body)
.send()
})
.await?;
if resp.status() == reqwest::StatusCode::NOT_FOUND
|| resp.status() == reqwest::StatusCode::BAD_REQUEST
{
let status = resp.status();
let text = resp.text().await.unwrap_or_default();
tracing::warn!(
"Responses 端点不支持 ({}: {}),自动回退 chat/completions 重发",
status,
text.chars().take(500).collect::<String>()
);
return self.chat_complete_stream(messages, max_output_tokens_override).await;
}
if !resp.status().is_success() {
let status = resp.status();
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("API 返回错误 ({}): {}", status, text);
}
collect_sse(resp, "data: ", |v| {
if v["type"].as_str() == Some("response.output_text.delta") {
v["delta"].as_str().map(|s| s.to_string())
} else {
None
}
})
.await
}
}
pub struct AnthropicProvider {
client: Client,
api_key: String,
model: String,
base_url: String,
max_retries: u32,
max_tokens: Option<u32>,
temperature: Option<f32>,
call_count: std::sync::atomic::AtomicUsize,
}
impl AnthropicProvider {
pub fn new(config: &LlmSection) -> Result<Self> {
let api_key = config.api_key.clone()
.or_else(|| std::env::var(&config.api_key_env).ok())
.with_context(|| format!("Anthropic API Key 未设置(api_key 为空且环境变量 {} 未定义)。请设置环境变量 {},或编辑配置文件的 [llm] 段填入 api_key", config.api_key_env, config.api_key_env))?;
let base_url = config
.base_url
.clone()
.unwrap_or_else(|| "https://api.anthropic.com/v1".to_string());
let client = Client::builder()
.build()
.context("创建 HTTP 客户端失败")?;
Ok(Self {
client,
api_key,
model: config.model.clone(),
base_url,
max_retries: MAX_RETRIES,
max_tokens: None,
temperature: None,
call_count: std::sync::atomic::AtomicUsize::new(0),
})
}
}
impl AnthropicProvider {
fn build_messages_body(
&self,
messages: &[Message],
stream: bool,
max_tokens_override: Option<u32>,
) -> serde_json::Value {
let system = messages.iter().find(|m| m.role == "system").map(|m| &m.content);
let non_system: Vec<&Message> = messages.iter().filter(|m| m.role != "system").collect();
let anthropic_messages: Vec<serde_json::Value> = non_system
.iter()
.map(|m| {
serde_json::json!({
"role": if m.role == "user" { "user" } else { "assistant" },
"content": m.content
})
})
.collect();
let mut body = serde_json::json!({
"model": self.model,
"max_tokens": max_tokens_override.or(self.max_tokens).unwrap_or(4096),
"messages": anthropic_messages,
});
if let Some(s) = system {
body["system"] = serde_json::json!(s);
}
if let Some(temp) = self.temperature {
body["temperature"] = serde_json::json!(temp);
}
if stream {
body["stream"] = serde_json::json!(true);
}
body
}
}
impl LlmProvider for AnthropicProvider {
async fn complete_stream(&self, messages: &[Message]) -> Result<Vec<String>> {
self.call_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let url = format!("{}/messages", self.base_url);
let body = self.build_messages_body(messages, true, None);
let resp = retry_with_backoff(self.max_retries, || {
self.client
.post(&url)
.header("x-api-key", &self.api_key)
.header("anthropic-version", "2023-06-01")
.json(&body)
.send()
})
.await?;
if !resp.status().is_success() {
let status = resp.status();
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("Anthropic API 返回错误 ({}): {}", status, text);
}
collect_sse(resp, "data: ", |v| {
if v["type"] == "content_block_delta" {
v["delta"]["text"].as_str().map(|s| s.to_string())
} else {
None
}
})
.await
}
async fn complete_with_budget(
&self,
messages: &[Message],
max_output_tokens: Option<u32>,
) -> Result<String> {
self.call_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let url = format!("{}/messages", self.base_url);
let body = self.build_messages_body(messages, true, max_output_tokens);
let resp = retry_with_backoff(self.max_retries, || {
self.client
.post(&url)
.header("x-api-key", &self.api_key)
.header("anthropic-version", "2023-06-01")
.json(&body)
.send()
})
.await?;
if !resp.status().is_success() {
let status = resp.status();
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("Anthropic API 返回错误 ({}): {}", status, text);
}
let chunks = collect_sse(resp, "data: ", |v| {
if v["type"] == "content_block_delta" {
v["delta"]["text"].as_str().map(|s| s.to_string())
} else {
None
}
})
.await?;
Ok(chunks.concat())
}
fn call_count(&self) -> usize {
self.call_count.load(std::sync::atomic::Ordering::Relaxed)
}
}
pub struct MockProvider {
call_count: std::sync::atomic::AtomicUsize,
}
impl MockProvider {
pub fn new() -> Self {
Self {
call_count: std::sync::atomic::AtomicUsize::new(0),
}
}
}
impl Default for MockProvider {
fn default() -> Self {
Self::new()
}
}
impl LlmProvider for MockProvider {
async fn complete(&self, _messages: &[Message]) -> Result<String> {
self.call_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Ok(
r#"{"summary": "这是 Mock Provider 生成的模拟摘要", "key_entities": []}"#
.to_string(),
)
}
async fn complete_stream(&self, _messages: &[Message]) -> Result<Vec<String>> {
self.call_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Ok(vec!["模拟流式响应 chunk".to_string()])
}
fn call_count(&self) -> usize {
self.call_count.load(std::sync::atomic::Ordering::Relaxed)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::schema::{LlmProviderType, LlmSection};
use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
struct MockRequest {
path: String,
headers: Vec<(String, String)>,
body: String,
}
struct MockResponse {
status: u16,
body: String,
}
fn header_complete(buf: &[u8]) -> bool {
buf.windows(4).any(|w| w == b"\r\n\r\n")
}
fn read_request(stream: &mut TcpStream) -> MockRequest {
let mut buf = Vec::new();
let mut tmp = [0u8; 4096];
while !header_complete(&buf) {
match stream.read(&mut tmp) {
Ok(0) | Err(_) => break,
Ok(n) => buf.extend_from_slice(&tmp[..n]),
}
}
let head_end = buf.windows(4).position(|w| w == b"\r\n\r\n").unwrap_or(buf.len());
let head = String::from_utf8_lossy(&buf[..head_end]).to_string();
let mut lines = head.split("\r\n");
let path = lines.next().unwrap_or("").split_whitespace().nth(1).unwrap_or("").to_string();
let headers: Vec<(String, String)> = lines
.filter_map(|l| l.split_once(':'))
.map(|(k, v)| (k.trim().to_string(), v.trim().to_string()))
.collect();
let content_length = headers.iter()
.find(|(k, _)| k.eq_ignore_ascii_case("content-length"))
.and_then(|(_, v)| v.parse::<usize>().ok())
.unwrap_or(0);
const HEADER_SEP: usize = 4;
while buf.len() < head_end + HEADER_SEP + content_length {
match stream.read(&mut tmp) {
Ok(0) | Err(_) => break,
Ok(n) => buf.extend_from_slice(&tmp[..n]),
}
}
let body =
String::from_utf8_lossy(&buf[head_end + HEADER_SEP..head_end + HEADER_SEP + content_length])
.to_string();
MockRequest { path, headers, body }
}
fn spawn_mock_server(
handler: impl Fn(MockRequest) -> MockResponse + Send + Sync + 'static,
) -> String {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let base_url = format!("http://{}", listener.local_addr().unwrap());
let handler = Arc::new(handler);
std::thread::spawn(move || {
for stream in listener.incoming() {
let Ok(mut stream) = stream else { break };
let handler = handler.clone();
std::thread::spawn(move || {
let req = read_request(&mut stream);
let resp = handler(req);
let reason = if resp.status == 200 { "OK" } else { "Error" };
let raw = format!(
"HTTP/1.1 {} {}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
resp.status, reason, resp.body.len(), resp.body
);
let _ = stream.write_all(raw.as_bytes());
});
}
});
base_url
}
fn openai_config(base_url: &str) -> LlmSection {
LlmSection {
provider: LlmProviderType::OpenAI,
model: "gpt-test".into(),
base_url: Some(format!("{}/v1", base_url)),
api_key: Some("test-key".into()),
api_key_env: "OPENAI_API_KEY".into(),
}
}
#[tokio::test]
async fn test_mock_provider() {
let provider = MockProvider::new();
let messages = vec![Message::user("测试消息")];
let result = provider.complete(&messages).await;
assert!(result.is_ok());
assert!(result.unwrap().contains("模拟摘要"));
assert_eq!(provider.call_count(), 1);
}
#[tokio::test]
async fn test_message_constructors() {
let sys = Message::system("你好");
assert_eq!(sys.role, "system");
assert_eq!(sys.content, "你好");
let user = Message::user("测试");
assert_eq!(user.role, "user");
let asst = Message::assistant("回复");
assert_eq!(asst.role, "assistant");
}
#[tokio::test]
async fn test_openai_request_builds_correct_payload() {
let captured = Arc::new(Mutex::new(None::<MockRequest>));
let captured_server = captured.clone();
let base_url = spawn_mock_server(move |req| {
*captured_server.lock().unwrap() = Some(req);
MockResponse {
status: 200,
body: r#"data: {"choices":[{"delta":{"content":"你好,这是 mock 回复"}}]}
data: [DONE]
"#
.into(),
}
});
let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
let messages = vec![Message::system("你是测试助手"), Message::user("你好")];
let reply = provider.complete(&messages).await.unwrap();
assert_eq!(reply, "你好,这是 mock 回复");
let req = captured.lock().unwrap().take().expect("应收到一次请求");
assert_eq!(req.path, "/v1/chat/completions");
let auth = req.headers.iter()
.find(|(k, _)| k.eq_ignore_ascii_case("authorization"))
.expect("应携带 Authorization 头");
assert_eq!(auth.1, "Bearer test-key");
let body: serde_json::Value = serde_json::from_str(&req.body).unwrap();
assert_eq!(body["model"], "gpt-test");
assert_eq!(body["messages"][0]["role"], "system");
assert_eq!(body["messages"][0]["content"], "你是测试助手");
assert_eq!(body["messages"][1]["role"], "user");
assert!(body.get("max_tokens").is_none(), "硬编码后不应写 max_tokens");
assert!(body.get("temperature").is_none(), "硬编码后不应写 temperature");
assert_eq!(
body["stream"].as_bool(),
Some(true),
"生产路径必须请求流式响应(stream:true)"
);
}
#[tokio::test]
async fn test_openai_stream_parses_sse() {
let sse = concat!(
"data: {\"choices\":[{\"delta\":{\"content\":\"你\"}}]}\n\n",
"data: {\"choices\":[{\"delta\":{\"content\":\"好\"}}]}\n\n",
"data: [DONE]\n\n",
);
let base_url = spawn_mock_server(move |_req| MockResponse {
status: 200,
body: sse.to_string(),
});
let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
let messages = vec![Message::user("你好")];
let chunks = provider.complete_stream(&messages).await.unwrap();
assert_eq!(chunks, vec!["你", "好"]);
assert_eq!(chunks.join(""), "你好");
}
#[tokio::test]
async fn test_slow_stream_not_truncated_by_total_timeout() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let base_url = format!("http://{}", listener.local_addr().unwrap());
std::thread::spawn(move || {
for stream in listener.incoming() {
let Ok(mut stream) = stream else { break };
std::thread::spawn(move || {
let _req = read_request(&mut stream);
let head = "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: close\r\n\r\n";
let _ = stream.write_all(head.as_bytes());
let _ = stream.write_all("data: {\"choices\":[{\"delta\":{\"content\":\"第一段\"}}]}\n\n".as_bytes());
let _ = stream.flush();
std::thread::sleep(Duration::from_millis(300));
let _ = stream.write_all("data: {\"choices\":[{\"delta\":{\"content\":\"第二段\"}}]}\n\n".as_bytes());
let _ = stream.write_all(b"data: [DONE]\n\n");
let _ = stream.flush();
});
}
});
let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
let messages = vec![Message::user("你好")];
let reply = provider.complete(&messages).await.unwrap();
assert_eq!(reply, "第一段第二段", "慢流两段必须完整拼接(无总超时截断)");
}
#[tokio::test]
async fn test_retry_on_server_error() {
let attempts = Arc::new(AtomicUsize::new(0));
let attempts_server = attempts.clone();
let base_url = spawn_mock_server(move |_req| {
let n = attempts_server.fetch_add(1, Ordering::Relaxed);
if n == 0 {
MockResponse { status: 500, body: "internal error".into() }
} else {
MockResponse { status: 200, body: "data: {\"choices\":[{\"delta\":{\"content\":\"重试成功\"}}]}\n\ndata: [DONE]\n\n".into() }
}
});
let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
let messages = vec![Message::user("你好")];
let reply = provider.complete(&messages).await.unwrap();
assert_eq!(reply, "重试成功");
assert_eq!(attempts.load(Ordering::Relaxed), 2);
}
#[tokio::test]
async fn test_retry_on_429() {
let attempts = Arc::new(AtomicUsize::new(0));
let attempts_server = attempts.clone();
let base_url = spawn_mock_server(move |_req| {
let n = attempts_server.fetch_add(1, Ordering::Relaxed);
if n == 0 {
MockResponse { status: 429, body: "rate limited".into() }
} else {
MockResponse { status: 200, body: "data: {\"choices\":[{\"delta\":{\"content\":\"限流后成功\"}}]}\n\ndata: [DONE]\n\n".into() }
}
});
let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
let messages = vec![Message::user("你好")];
let reply = provider.complete(&messages).await.unwrap();
assert_eq!(reply, "限流后成功");
assert_eq!(attempts.load(Ordering::Relaxed), 2);
}
#[tokio::test]
async fn test_no_retry_on_401() {
let attempts = Arc::new(AtomicUsize::new(0));
let attempts_server = attempts.clone();
let base_url = spawn_mock_server(move |_req| {
attempts_server.fetch_add(1, Ordering::Relaxed);
MockResponse { status: 401, body: "unauthorized".into() }
});
let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
let messages = vec![Message::user("你好")];
let result = provider.complete(&messages).await;
assert!(result.is_err());
assert_eq!(attempts.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn test_retry_exhausted_on_5xx() {
let attempts = Arc::new(AtomicUsize::new(0));
let attempts_server = attempts.clone();
let base_url = spawn_mock_server(move |_req| {
attempts_server.fetch_add(1, Ordering::Relaxed);
MockResponse { status: 500, body: "internal error".into() }
});
let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
let messages = vec![Message::user("你好")];
let result = provider.complete(&messages).await;
assert!(result.is_err());
assert_eq!(attempts.load(Ordering::Relaxed), MAX_RETRIES as usize);
}
#[tokio::test]
async fn test_retry_on_timeout() {
let attempts = Arc::new(AtomicUsize::new(0));
let attempts_server = attempts.clone();
let base_url = spawn_mock_server(move |_req| {
let n = attempts_server.fetch_add(1, Ordering::Relaxed);
if n == 0 {
std::thread::sleep(Duration::from_millis(500));
}
MockResponse { status: 200, body: "{}".into() }
});
let client = Client::builder()
.timeout(Duration::from_millis(200))
.build()
.unwrap();
let resp = retry_with_backoff(MAX_RETRIES, || {
client.get(format!("{}/t", base_url)).send()
})
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(attempts.load(Ordering::Relaxed), 2);
}
#[test]
fn test_parse_sse_openai_format() {
let sse = concat!(
"data: {\"choices\":[{\"delta\":{\"content\":\"你\"}}]}\n\n",
"data: {\"choices\":[{\"delta\":{\"content\":\"好\"}}]}\n\n",
"data: [DONE]\n\n",
);
let chunks = parse_sse_stream(sse.as_bytes(), "data: ", |v| {
v["choices"][0]["delta"]["content"]
.as_str()
.map(|s| s.to_string())
});
assert_eq!(chunks, vec!["你", "好"]);
}
#[test]
fn test_parse_sse_anthropic_format() {
let sse = concat!(
"data: {\"type\":\"message_start\",\"message\":{\"id\":\"m1\"}}\n\n",
"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"你\"}}\n\n",
"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"好\"}}\n\n",
"data: {\"type\":\"message_stop\"}\n\n",
);
let chunks = parse_sse_stream(sse.as_bytes(), "data: ", |v| {
if v["type"] == "content_block_delta" {
v["delta"]["text"].as_str().map(|s| s.to_string())
} else {
None
}
});
assert_eq!(chunks, vec!["你", "好"]);
}
#[tokio::test]
async fn test_anthropic_stream_parses_sse() {
let sse = concat!(
"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"克\"}}\n\n",
"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"劳\"}}\n\n",
"data: {\"type\":\"message_stop\"}\n\n",
);
let base_url = spawn_mock_server(move |_req| MockResponse {
status: 200,
body: sse.to_string(),
});
let config = LlmSection {
provider: LlmProviderType::Anthropic,
model: "claude-test".into(),
base_url: Some(base_url),
api_key: Some("sk-ant-test".into()),
api_key_env: "ANTHROPIC_API_KEY".into(),
};
let provider = AnthropicProvider::new(&config).unwrap();
let messages = vec![Message::user("你好")];
let chunks = provider.complete_stream(&messages).await.unwrap();
assert_eq!(chunks, vec!["克", "劳"]);
}
#[test]
fn test_anthropic_provider_construction() {
let config = LlmSection {
provider: LlmProviderType::Anthropic,
model: "claude-test".into(),
base_url: None,
api_key: Some("sk-ant-test".into()),
api_key_env: "ANTHROPIC_API_KEY".into(),
};
let provider = AnthropicProvider::new(&config).unwrap();
assert_eq!(provider.call_count(), 0);
}
#[tokio::test]
async fn test_anthropic_request_builds_correct_payload() {
let captured = Arc::new(Mutex::new(None::<MockRequest>));
let captured_server = captured.clone();
let base_url = spawn_mock_server(move |req| {
*captured_server.lock().unwrap() = Some(req);
MockResponse {
status: 200,
body: concat!(
"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"claude 回复\"}}\n\n",
"data: {\"type\":\"message_stop\"}\n\n",
)
.into(),
}
});
let config = LlmSection {
provider: LlmProviderType::Anthropic,
model: "claude-test".into(),
base_url: Some(base_url),
api_key: Some("sk-ant-test".into()),
api_key_env: "ANTHROPIC_API_KEY".into(),
};
let provider = AnthropicProvider::new(&config).unwrap();
let messages = vec![
Message::system("你是助手"),
Message::user("你好"),
Message::assistant("在的"),
];
let reply = provider.complete(&messages).await.unwrap();
assert_eq!(reply, "claude 回复");
let req = captured.lock().unwrap().take().expect("应收到一次请求");
assert_eq!(req.path, "/messages");
let api_key_header = req
.headers
.iter()
.find(|(k, _)| k.eq_ignore_ascii_case("x-api-key"))
.expect("应携带 x-api-key 头");
assert_eq!(api_key_header.1, "sk-ant-test");
let version_header = req
.headers
.iter()
.find(|(k, _)| k.eq_ignore_ascii_case("anthropic-version"))
.expect("应携带 anthropic-version 头");
assert_eq!(version_header.1, "2023-06-01");
let body: serde_json::Value = serde_json::from_str(&req.body).unwrap();
assert_eq!(body["model"], "claude-test");
assert_eq!(body["max_tokens"].as_u64(), Some(4096), "max_tokens 未配置时默认 4096");
assert_eq!(body["system"], "你是助手");
let msgs = body["messages"].as_array().unwrap();
assert_eq!(msgs.len(), 2, "非 system 消息才进 messages");
assert_eq!(msgs[0]["role"], "user");
assert_eq!(msgs[0]["content"], "你好");
assert_eq!(msgs[1]["role"], "assistant");
assert_eq!(msgs[1]["content"], "在的");
}
#[tokio::test]
async fn test_responses_stream_parses_semantic_sse() {
let sse = concat!(
"data: {\"type\":\"response.created\",\"response\":{\"id\":\"r1\"}}\n\n",
"data: {\"type\":\"response.output_text.delta\",\"sequence_number\":0,\"delta\":\"你\"}\n\n",
"data: {\"type\":\"response.output_text.delta\",\"sequence_number\":1,\"delta\":\"好\"}\n\n",
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"r1\"}}\n\n",
);
let base_url = spawn_mock_server(move |_req| MockResponse {
status: 200,
body: sse.to_string(),
});
let config = LlmSection {
provider: LlmProviderType::OpenAI,
model: "deepseek-v4-flash".into(),
base_url: Some(format!("{}/v1", base_url)),
api_key: Some("test-key".into()),
api_key_env: "DEEPSEEK_API_KEY".into(),
};
let provider = OpenAiProvider::new(&config, OpenAiProtocol::Responses).unwrap();
let chunks = provider.complete_stream(&[Message::user("你好")]).await.unwrap();
assert_eq!(chunks, vec!["你", "好"], "语义化事件应提取 delta 文本");
assert_eq!(provider.call_count(), 1);
}
#[tokio::test]
async fn test_responses_request_builds_correct_payload() {
let captured = Arc::new(Mutex::new(None::<MockRequest>));
let captured_server = captured.clone();
let base_url = spawn_mock_server(move |req| {
*captured_server.lock().unwrap() = Some(req);
MockResponse {
status: 200,
body: "data: {\"type\":\"response.completed\"}\n\n".into(),
}
});
let config = LlmSection {
provider: LlmProviderType::OpenAI,
model: "deepseek-v4-flash".into(),
base_url: Some(format!("{}/v1", base_url)),
api_key: Some("test-key".into()),
api_key_env: "DEEPSEEK_API_KEY".into(),
};
let provider = OpenAiProvider::new(&config, OpenAiProtocol::Responses).unwrap();
let messages = vec![Message::system("你是助手"), Message::user("你好")];
let _ = provider.complete_stream(&messages).await.unwrap();
let req = captured.lock().unwrap().take().expect("应收到一次请求");
assert_eq!(req.path, "/v1/responses", "Responses 协议应请求 /responses 端点");
let body: serde_json::Value = serde_json::from_str(&req.body).unwrap();
assert_eq!(body["model"], "deepseek-v4-flash");
assert_eq!(body["instructions"], "你是助手");
let input = body["input"].as_array().unwrap();
assert_eq!(input.len(), 1, "非 system 消息才进 input");
assert_eq!(input[0]["role"], "user");
assert_eq!(input[0]["content"][0]["type"], "input_text");
assert_eq!(input[0]["content"][0]["text"], "你好");
assert!(body.get("max_output_tokens").is_none(), "硬编码后不应写 max_output_tokens");
assert!(body.get("max_tokens").is_none(), "Responses 不得用 max_tokens 参数名");
assert!(body.get("temperature").is_none(), "硬编码后不应写 temperature");
assert_eq!(body["stream"].as_bool(), Some(true));
}
#[tokio::test]
async fn test_responses_falls_back_to_chat_on_404() {
let requests: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
let requests_server = requests.clone();
let base_url = spawn_mock_server(move |req| {
requests_server.lock().unwrap().push(req.path.clone());
if req.path.ends_with("/responses") {
MockResponse { status: 404, body: "not found".into() }
} else {
MockResponse {
status: 200,
body: "data: {\"choices\":[{\"delta\":{\"content\":\"回退成功\"}}]}\n\ndata: [DONE]\n\n".into(),
}
}
});
let config = LlmSection {
provider: LlmProviderType::OpenAI,
model: "deepseek-v4-flash".into(),
base_url: Some(format!("{}/v1", base_url)),
api_key: Some("test-key".into()),
api_key_env: "DEEPSEEK_API_KEY".into(),
};
let provider = OpenAiProvider::new(&config, OpenAiProtocol::Responses).unwrap();
let chunks = provider.complete_stream(&[Message::user("你好")]).await.unwrap();
assert_eq!(chunks.join(""), "回退成功");
let paths = requests.lock().unwrap();
assert_eq!(paths.len(), 2, "应请求 responses + chat 两次");
assert!(paths[0].ends_with("/responses"), "第一次应请求 responses: {:?}", paths);
assert!(paths[1].ends_with("/chat/completions"), "回退应请求 chat/completions: {:?}", paths);
}
#[tokio::test]
async fn test_complete_with_budget_sets_max_output_tokens() {
let captured = Arc::new(Mutex::new(None::<MockRequest>));
let captured_server = captured.clone();
let base_url = spawn_mock_server(move |req| {
*captured_server.lock().unwrap() = Some(req);
MockResponse {
status: 200,
body: "data: {\"type\":\"response.output_text.delta\",\"delta\":\"{\\\"rubrics\\\":[]}\"}\n\n"
.into(),
}
});
let config = LlmSection {
provider: LlmProviderType::OpenAI,
model: "deepseek-v4-flash".into(),
base_url: Some(format!("{}/v1", base_url)),
api_key: Some("test-key".into()),
api_key_env: "DEEPSEEK_API_KEY".into(),
};
let provider = OpenAiProvider::new(&config, OpenAiProtocol::Responses).unwrap();
let out = provider
.complete_with_budget(&[Message::user("你好")], Some(16384))
.await
.unwrap();
assert_eq!(out, "{\"rubrics\":[]}", "带预算调用应返回完整文本");
let req = captured.lock().unwrap().take().expect("应收到一次请求");
assert_eq!(req.path, "/v1/responses");
let body: serde_json::Value = serde_json::from_str(&req.body).unwrap();
assert_eq!(body["max_output_tokens"].as_u64(), Some(16384), "预算应写入 max_output_tokens");
assert_eq!(body["stream"].as_bool(), Some(true), "带预算路径仍走流式");
}
}