Skip to main content

code_repo_wiki/generate/
llm.rs

1use std::time::Duration;
2
3use anyhow::{Context, Result};
4use reqwest::Client;
5
6use crate::config::schema::LlmSection;
7
8/// LLM 对话消息
9#[derive(Debug, Clone)]
10pub struct Message {
11    pub role: String,
12    pub content: String,
13}
14
15impl Message {
16    pub fn system(content: impl Into<String>) -> Self {
17        Self {
18            role: "system".into(),
19            content: content.into(),
20        }
21    }
22
23    pub fn user(content: impl Into<String>) -> Self {
24        Self {
25            role: "user".into(),
26            content: content.into(),
27        }
28    }
29
30    pub fn assistant(content: impl Into<String>) -> Self {
31        Self {
32            role: "assistant".into(),
33            content: content.into(),
34        }
35    }
36}
37
38/// LLM Provider 抽象 trait
39///
40/// Rust 2024 支持在 trait 中使用 async fn,无需 async-trait crate。
41/// 注意:async fn in trait 不满足 dyn 安全性,请通过泛型或 Provider 枚举使用。
42///
43/// 契约(t09 真流式接线):**生产路径统一走流式**——`complete` 的默认
44/// 实现 = `complete_stream` 收集并拼接 chunks,实现方只需实现
45/// `complete_stream` 即获得完整语义;需要自定义非流式行为的实现
46/// (如 Mock 返回固定 JSON)显式覆盖 `complete`。未实现 `complete_stream`
47/// 的实现调用 `complete` 会得到显式错误,不静默退化。
48#[allow(async_fn_in_trait)]
49pub trait LlmProvider: Send + Sync {
50    async fn complete(&self, messages: &[Message]) -> Result<String> {
51        let chunks = self.complete_stream(messages).await?;
52        Ok(chunks.concat())
53    }
54    async fn complete_stream(&self, messages: &[Message]) -> Result<Vec<String>> {
55        let _ = messages;
56        Err(anyhow::anyhow!("streaming not supported"))
57    }
58    /// 带输出预算上限的完整(非流式)调用。
59    ///
60    /// 评测裁判(rubrics / TQS 叶子判定)等长结构化输出场景必须显式
61    /// 传预算:推理型模型(如 deepseek-v4-flash)的 reasoning 会消耗
62    /// 输出预算,预算不足时响应可能只有 reasoning 块没有 message
63    /// (实测 v22 rubrics 首跑 3+3 轮全败,max=4000 复现只有
64    /// reasoning、max=8192 才出现 message)。默认实现不带预算
65    /// (等价 complete),需要预算的 Provider 自行覆盖。
66    async fn complete_with_budget(
67        &self,
68        messages: &[Message],
69        _max_output_tokens: Option<u32>,
70    ) -> Result<String> {
71        self.complete(messages).await
72    }
73    /// 返回已完成的 LLM 调用次数
74    fn call_count(&self) -> usize {
75        0
76    }
77}
78
79/// 统一的 Provider 枚举,包装所有 Provider 实现
80///
81/// 通过此枚举可以在需要动态分发时避免 dyn trait 的限制。
82pub enum Provider {
83    OpenAi(OpenAiProvider),
84    Anthropic(AnthropicProvider),
85    Mock(MockProvider),
86}
87
88impl LlmProvider for Provider {
89    async fn complete(&self, messages: &[Message]) -> Result<String> {
90        match self {
91            Provider::OpenAi(p) => p.complete(messages).await,
92            Provider::Anthropic(p) => p.complete(messages).await,
93            Provider::Mock(p) => p.complete(messages).await,
94        }
95    }
96
97    async fn complete_stream(&self, messages: &[Message]) -> Result<Vec<String>> {
98        match self {
99            Provider::OpenAi(p) => p.complete_stream(messages).await,
100            Provider::Anthropic(p) => p.complete_stream(messages).await,
101            Provider::Mock(p) => p.complete_stream(messages).await,
102        }
103    }
104
105    async fn complete_with_budget(
106        &self,
107        messages: &[Message],
108        max_output_tokens: Option<u32>,
109    ) -> Result<String> {
110        match self {
111            Provider::OpenAi(p) => p.complete_with_budget(messages, max_output_tokens).await,
112            Provider::Anthropic(p) => p.complete_with_budget(messages, max_output_tokens).await,
113            Provider::Mock(p) => p.complete_with_budget(messages, max_output_tokens).await,
114        }
115    }
116
117    fn call_count(&self) -> usize {
118        match self {
119            Provider::OpenAi(p) => p.call_count(),
120            Provider::Anthropic(p) => p.call_count(),
121            Provider::Mock(p) => p.call_count(),
122        }
123    }
124}
125
126// ============ 共享骨架:重试 + SSE 解析(OpenAI 与 Anthropic 共用) ============
127
128/// 统一的重试上限(总尝试次数),OpenAiProvider/AnthropicProvider 均从该常量取值
129pub(crate) const MAX_RETRIES: u32 = 3;
130
131/// 指数退避:500ms * 2^attempt + 随机抖动 0-250ms。
132/// 抖动用系统时钟纳秒取模实现,避免为单点功能引入 rand 依赖。
133fn backoff_delay(attempt: u32) -> Duration {
134    let base = 500u64 * 2u64.pow(attempt);
135    let jitter = std::time::SystemTime::now()
136        .duration_since(std::time::UNIX_EPOCH)
137        .map(|d| u64::from(d.subsec_nanos()) % 251)
138        .unwrap_or(0);
139    Duration::from_millis(base + jitter)
140}
141
142/// 可重试的 HTTP 状态码白名单:429 限流 + 5xx 服务端错误;
143/// 其余 4xx 业务错误(400/401/403/404/422 等)不在白名单内,立即失败。
144fn is_retryable_status(status: reqwest::StatusCode) -> bool {
145    status == reqwest::StatusCode::TOO_MANY_REQUESTS || status.is_server_error()
146}
147
148/// 统一重试骨架:可重试错误(429/5xx/超时/连接失败)按指数退避重试,
149/// 其余 4xx 立即失败。send_fn 每轮重新构建请求(请求构建仍是协议差异点,留在调用方)。
150///
151/// 可观测性(v16 B 组):每次可重试失败记录 attempt/原因与下次退避延迟,
152/// 错误响应体保留(截断 2000 字符防日志爆炸);重试耗尽时汇总 last_error。
153/// 此前本函数零日志——生产排查重试风暴只能看到最终错误,过程全黑。
154pub(crate) async fn retry_with_backoff<F, Fut>(
155    max_retries: u32,
156    send_fn: F,
157) -> Result<reqwest::Response>
158where
159    F: Fn() -> Fut,
160    Fut: std::future::Future<Output = Result<reqwest::Response, reqwest::Error>> + Send,
161{
162    let mut last_error = None;
163    for attempt in 0..max_retries {
164        if attempt > 0 {
165            let delay = backoff_delay(attempt - 1);
166            tracing::info!(
167                "LLM 请求重试(第 {}/{} 次),退避 {}ms",
168                attempt + 1,
169                max_retries,
170                delay.as_millis()
171            );
172            tokio::time::sleep(delay).await;
173        }
174        match send_fn().await {
175            Ok(resp) if is_retryable_status(resp.status()) => {
176                let status = resp.status();
177                let text = resp.text().await.unwrap_or_default();
178                tracing::warn!(
179                    "LLM API 返回可重试状态 {}(第 {} 次尝试): {}",
180                    status,
181                    attempt + 1,
182                    text.chars().take(2000).collect::<String>()
183                );
184                last_error = Some(anyhow::anyhow!("API 返回错误 ({}): {}", status, text));
185            }
186            Ok(resp) => return Ok(resp),
187            Err(e) if e.is_timeout() || e.is_connect() => {
188                tracing::warn!(
189                    "LLM 请求超时/连接失败(第 {} 次尝试,将重试): {}",
190                    attempt + 1,
191                    e
192                );
193                last_error = Some(anyhow::anyhow!("请求失败: {}", e));
194            }
195            Err(e) => return Err(anyhow::anyhow!("请求失败: {}", e)),
196        }
197    }
198    tracing::error!("LLM API 调用重试 {} 次后全部失败: {:?}", max_retries, last_error);
199    Err(anyhow::anyhow!(
200        "LLM API 调用重试 {} 次后全部失败: {:?}",
201        max_retries,
202        last_error
203    ))
204}
205
206/// 共享 SSE 行解析:按 \n 切分、跳过空行、剥离行前缀、解析 JSON,
207/// 文本增量交给 extract 提取。OpenAI 与 Anthropic 的 SSE 行均为 `data: ` 前缀,
208/// 真实差异在 JSON 字段路径(OpenAI: choices[0].delta.content;
209/// Anthropic: type=content_block_delta 事件的 delta.text),故用提取闭包参数化。
210/// `data: [DONE]` 等非 JSON 行解析失败自然跳过。
211/// 仅测试使用(流式路径已内联同样的按行解析,此处保留整块解析供 mock 断言)
212#[cfg(test)]
213fn parse_sse_stream(
214    bytes: &[u8],
215    line_prefix: &str,
216    extract: impl Fn(&serde_json::Value) -> Option<String>,
217) -> Vec<String> {
218    let mut chunks = Vec::new();
219    for line in String::from_utf8_lossy(bytes).split('\n') {
220        let line = line.trim();
221        if line.is_empty() {
222            continue;
223        }
224        let Some(json_str) = line.strip_prefix(line_prefix) else { continue };
225        if let Ok(val) = serde_json::from_str::<serde_json::Value>(json_str)
226            && let Some(text) = extract(&val)
227        {
228            chunks.push(text);
229        }
230    }
231    chunks
232}
233
234/// 消费流式响应并走共享 SSE 解析(逐 chunk 流式,t09)
235///
236/// 原实现 `resp.bytes().await` 全量收包后统一解析:长生成(真实 LLM
237/// 10min 超时前科)受 client 总超时(120s)**整体截断**,重试又从头
238/// 再来——进度全部丢失。流式方案:`bytes_stream` 逐块读取,每次读块
239/// 用**空闲超时**保护(60s 无数据才判超时):只要模型还在产出就不会
240/// 超时,真正停止产出才失败,长生成不再受总超时限制。
241///
242/// SSE 行解析复用 parse_sse_stream 的语义(`data: ` 前缀 + 提取闭包),
243/// 逐行增量处理;跨 chunk 的残行保留到下一块。
244const SSE_IDLE_TIMEOUT: Duration = Duration::from_secs(60);
245
246async fn collect_sse(
247    resp: reqwest::Response,
248    line_prefix: &str,
249    extract: impl Fn(&serde_json::Value) -> Option<String>,
250) -> Result<Vec<String>> {
251    use futures::StreamExt;
252    let mut stream = resp.bytes_stream();
253    let mut buf: Vec<u8> = Vec::new();
254    let mut chunks = Vec::new();
255
256    // v16 B 组:流式消费的可观测性——长生成场景(分钟级)若无日志,
257    // 用户无法区分"模型还在产出"与"卡死无响应"。记录流开始与结束
258    // 统计,空闲超时单独 warn(含已收 chunk 数,便于判断进度丢失量)。
259    tracing::info!("SSE 流开始消费(空闲超时保护 {}s)", SSE_IDLE_TIMEOUT.as_secs());
260    loop {
261        let item = match tokio::time::timeout(SSE_IDLE_TIMEOUT, stream.next()).await {
262            Ok(Some(Ok(item))) => item,
263            Ok(Some(Err(e))) => return Err(e.into()),
264            // 流正常结束:处理尾部残行后返回
265            Ok(None) => break,
266            Err(_) => {
267                tracing::warn!(
268                    "SSE 流读取空闲超时({}s 无数据,已收 {} 个 chunk,模型可能已停止产出)",
269                    SSE_IDLE_TIMEOUT.as_secs(),
270                    chunks.len()
271                );
272                anyhow::bail!(
273                    "SSE 流读取空闲超时({}s 无数据,模型可能已停止产出)",
274                    SSE_IDLE_TIMEOUT.as_secs()
275                )
276            }
277        };
278        buf.extend_from_slice(&item);
279        // 按行切分处理,保留未完成的尾部残行
280        let mut consumed = 0usize;
281        for (idx, b) in buf.iter().enumerate() {
282            if *b != b'\n' {
283                continue;
284            }
285            let line = String::from_utf8_lossy(&buf[consumed..idx]);
286            let line = line.trim();
287            if !line.is_empty()
288                && let Some(json_str) = line.strip_prefix(line_prefix)
289                && let Ok(val) = serde_json::from_str::<serde_json::Value>(json_str)
290                && let Some(text) = extract(&val)
291            {
292                chunks.push(text);
293            }
294            consumed = idx + 1;
295        }
296        buf.drain(..consumed);
297    }
298    // 尾部残行(流结束时最后一行可能无换行符)
299    if !buf.is_empty() {
300        let line = String::from_utf8_lossy(&buf);
301        if let Some(json_str) = line.trim().strip_prefix(line_prefix)
302            && let Ok(val) = serde_json::from_str::<serde_json::Value>(json_str)
303            && let Some(text) = extract(&val)
304        {
305            chunks.push(text);
306        }
307    }
308    tracing::info!(
309        "SSE 流消费完成: {} 个 chunk, {} 字符",
310        chunks.len(),
311        chunks.iter().map(|c| c.len()).sum::<usize>()
312    );
313    Ok(chunks)
314}
315
316/// OpenAI 协议形态(v17 t02 拆分:协议按 provider 类型显式绑定)
317///
318/// - `Responses`:OpenAI **Responses API**(POST /responses;DeepSeek 等
319///   支持 Responses 的服务经 base_url 接入)
320/// - `Chat`:**chat/completions**(OpenAI 兼容端点:阿里云/自建等)
321///
322/// 拆分原因:两协议的请求体(input/instructions vs messages[])、响应
323/// 解析(output.items vs choices[])、SSE 事件(语义化事件 vs
324/// choices[].delta)差异大,且不是所有兼容端点都提供 /responses。
325#[derive(Debug, Clone, Copy, PartialEq, Eq)]
326pub enum OpenAiProtocol {
327    Responses,
328    Chat,
329}
330
331/// OpenAI 兼容 API 的 LLM Provider 实现(v17 t02 起支持双协议)
332pub struct OpenAiProvider {
333    client: Client,
334    api_key: String,
335    model: String,
336    base_url: String,
337    /// 协议形态(构造时按 provider 类型绑定,见 create_provider)
338    protocol: OpenAiProtocol,
339    max_retries: u32,
340    max_tokens: Option<u32>,
341    temperature: Option<f32>,
342    call_count: std::sync::atomic::AtomicUsize,
343}
344
345impl OpenAiProvider {
346    /// 从配置创建 OpenAI Provider(v17 起按 provider 类型绑定协议)
347    ///
348    /// 优先使用 api_key 字段,其次从环境变量读取。支持自定义 base_url。
349    pub fn new(config: &LlmSection, protocol: OpenAiProtocol) -> Result<Self> {
350        let api_key = config.api_key.clone()
351            .or_else(|| std::env::var(&config.api_key_env).ok())
352            .with_context(|| {
353                // v17 t04:错误消息附加可操作引导——新用户只需设置环境变量
354                // 或编辑配置文件的 [llm] 段,不猜
355                format!(
356                    "LLM API Key 未设置(api_key 为空且环境变量 {} 未定义)。请设置环境变量 {},或编辑配置文件的 [llm] 段填入 api_key",
357                    config.api_key_env, config.api_key_env
358                )
359            })?;
360        let base_url = config
361            .base_url
362            .clone()
363            .unwrap_or_else(|| "https://api.openai.com/v1".to_string());
364
365        let client = Client::builder()
366            // 不设总超时(t09 真流式):长生成由流式路径的 SSE_IDLE_TIMEOUT
367            // 空闲超时保护——模型持续产出即不超时,真正停止产出 60s 才失败;
368            // 总超时会在长生成中途整体截断已产出内容(v12 实测 10min 超时前科)
369            .build()
370            .context("创建 HTTP 客户端失败")?;
371
372        Ok(Self {
373            client,
374            api_key,
375            model: config.model.clone(),
376            base_url,
377            protocol,
378            max_retries: MAX_RETRIES,
379            max_tokens: None,
380            temperature: None,
381            call_count: std::sync::atomic::AtomicUsize::new(0),
382        })
383    }
384}
385
386impl OpenAiProvider {
387    /// 构建 chat/completions 请求体;stream 决定是否追加流式标记
388    ///
389    /// max_tokens_override 为显式预算覆盖(v22 起构造时为 None,交模型
390    /// 默认);评测裁判的长结构化输出经 complete_with_budget 传入。
391    fn build_chat_body(
392        &self,
393        messages: &[Message],
394        stream: bool,
395        max_tokens_override: Option<u32>,
396    ) -> serde_json::Value {
397        let mut body = serde_json::json!({
398            "model": self.model,
399            "messages": messages.iter().map(|m| {
400                serde_json::json!({"role": m.role, "content": m.content})
401            }).collect::<Vec<_>>(),
402        });
403        if stream {
404            body["stream"] = serde_json::json!(true);
405        }
406        // 可选参数:显式覆盖优先,回退构造时的 max_tokens(均为 None 时省略)
407        if let Some(maxt) = max_tokens_override.or(self.max_tokens) {
408            body["max_tokens"] = serde_json::json!(maxt);
409        }
410        if let Some(temp) = self.temperature {
411            body["temperature"] = serde_json::json!(temp);
412        }
413        body
414    }
415
416    /// 构建 Responses API 请求体(v17 B4)
417    ///
418    /// 协议差异(t01 查证,OpenAI 官方迁移指南):请求从 messages[]
419    /// 改为 `input`(typed items 数组),system 消息分离到顶层
420    /// `instructions` 字段;token 上限参数名从 max_tokens 改为
421    /// `max_output_tokens`(DeepSeek 对不支持的参数静默忽略——参数名
422    /// 写错会静默失效,必须按协议用正确名称)。
423    fn build_responses_body(
424        &self,
425        messages: &[Message],
426        stream: bool,
427        max_output_tokens_override: Option<u32>,
428    ) -> serde_json::Value {
429        // system 消息 → 顶层 instructions;user/assistant → input items
430        let system = messages.iter().find(|m| m.role == "system").map(|m| &m.content);
431        let input: Vec<serde_json::Value> = messages
432            .iter()
433            .filter(|m| m.role != "system")
434            .map(|m| {
435                serde_json::json!({
436                    "role": if m.role == "user" { "user" } else { "assistant" },
437                    "content": serde_json::json!([{ "type": "input_text", "text": m.content }]),
438                })
439            })
440            .collect();
441        let mut body = serde_json::json!({
442            "model": self.model,
443            "input": input,
444        });
445        if let Some(s) = system {
446            body["instructions"] = serde_json::json!(s);
447        }
448        if stream {
449            body["stream"] = serde_json::json!(true);
450        }
451        // 可选参数:显式覆盖优先,回退构造时的 max_tokens(均为 None 时省略)
452        if let Some(maxt) = max_output_tokens_override.or(self.max_tokens) {
453            body["max_output_tokens"] = serde_json::json!(maxt);
454        }
455        if let Some(temp) = self.temperature {
456            body["temperature"] = serde_json::json!(temp);
457        }
458        body
459    }
460}
461
462impl LlmProvider for OpenAiProvider {
463    async fn complete_stream(&self, messages: &[Message]) -> Result<Vec<String>> {
464        self.call_count
465            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
466        match self.protocol {
467            OpenAiProtocol::Chat => self.chat_complete_stream(messages, None).await,
468            OpenAiProtocol::Responses => self.responses_complete_stream(messages, None).await,
469        }
470    }
471
472    /// 带输出预算的完整调用(评测裁判用):流式路径 + 显式预算
473    /// (推理型模型 reasoning 吞预算,见 trait 文档)
474    async fn complete_with_budget(
475        &self,
476        messages: &[Message],
477        max_output_tokens: Option<u32>,
478    ) -> Result<String> {
479        self.call_count
480            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
481        let chunks = match self.protocol {
482            OpenAiProtocol::Chat => self.chat_complete_stream(messages, max_output_tokens).await?,
483            OpenAiProtocol::Responses => {
484                self.responses_complete_stream(messages, max_output_tokens).await?
485            }
486        };
487        Ok(chunks.concat())
488    }
489
490    // complete 走 trait 默认实现(complete_stream 收集拼接)——
491    // 生产路径统一流式,无整读分支
492    fn call_count(&self) -> usize {
493        self.call_count.load(std::sync::atomic::Ordering::Relaxed)
494    }
495}
496
497impl OpenAiProvider {
498    /// chat/completions 协议路径(OpenAI 兼容端点;v17 起为
499    /// openai-compatible provider 的协议)
500    ///
501    /// max_tokens_override:流式请求体中的显式预算(complete_with_budget
502    /// 传入;常规路径 None)
503    async fn chat_complete_stream(
504        &self,
505        messages: &[Message],
506        max_tokens_override: Option<u32>,
507    ) -> Result<Vec<String>> {
508        let url = format!("{}/chat/completions", self.base_url);
509        let body = self.build_chat_body(messages, true, max_tokens_override);
510
511        let resp = retry_with_backoff(self.max_retries, || {
512            self.client
513                .post(&url)
514                .bearer_auth(&self.api_key)
515                .json(&body)
516                .send()
517        })
518        .await?;
519
520        if !resp.status().is_success() {
521            let status = resp.status();
522            let text = resp.text().await.unwrap_or_default();
523            anyhow::bail!("API 返回错误 ({}): {}", status, text);
524        }
525
526        // chat/completions SSE:data: 行内 choices[0].delta.content
527        collect_sse(resp, "data: ", |v| {
528            v["choices"][0]["delta"]["content"]
529                .as_str()
530                .map(|s| s.to_string())
531        })
532        .await
533    }
534
535    /// Responses API 协议路径(openai provider 的协议,v17 B4)
536    ///
537    /// 端点不支持信号(404/400,如服务未提供 /responses)→ 自动回退
538    /// chat/completions 重发一次(t02 拍板;429/5xx 由 retry_with_backoff
539    /// 处理,不触发回退——回退只针对"端点不支持",不掩盖限流/服务端错误)。
540    async fn responses_complete_stream(
541        &self,
542        messages: &[Message],
543        max_output_tokens_override: Option<u32>,
544    ) -> Result<Vec<String>> {
545        let url = format!("{}/responses", self.base_url);
546        let body = self.build_responses_body(messages, true, max_output_tokens_override);
547
548        let resp = retry_with_backoff(self.max_retries, || {
549            self.client
550                .post(&url)
551                .bearer_auth(&self.api_key)
552                .json(&body)
553                .send()
554        })
555        .await?;
556
557        if resp.status() == reqwest::StatusCode::NOT_FOUND
558            || resp.status() == reqwest::StatusCode::BAD_REQUEST
559        {
560            // 端点不支持(404)/参数被拒(400):服务未实现 Responses 协议,
561            // 回退 chat/completions 重发(仅一次——chat 失败按既有错误传播)
562            let status = resp.status();
563            let text = resp.text().await.unwrap_or_default();
564            tracing::warn!(
565                "Responses 端点不支持 ({}: {}),自动回退 chat/completions 重发",
566                status,
567                text.chars().take(500).collect::<String>()
568            );
569            return self.chat_complete_stream(messages, max_output_tokens_override).await;
570        }
571        if !resp.status().is_success() {
572            let status = resp.status();
573            let text = resp.text().await.unwrap_or_default();
574            anyhow::bail!("API 返回错误 ({}): {}", status, text);
575        }
576
577        // Responses SSE:语义化事件流——data: 行内 type=response.output_text.delta
578        // 事件的 delta 字段(无 [DONE] 终止符,以流结束为终止,collect_sse 兼容)
579        collect_sse(resp, "data: ", |v| {
580            if v["type"].as_str() == Some("response.output_text.delta") {
581                v["delta"].as_str().map(|s| s.to_string())
582            } else {
583                None
584            }
585        })
586        .await
587    }
588}
589
590/// Anthropic Claude API LLM Provider 实现
591///
592/// 通过 Anthropic Messages API 调用 Claude 系列模型。
593pub struct AnthropicProvider {
594    client: Client,
595    api_key: String,
596    model: String,
597    /// API 根地址(含 /v1 前缀,与 OpenAiProvider 语义一致);默认官方地址。
598    /// 可自定义以接入网关/本地代理,同时使请求构建可被测试。
599    base_url: String,
600    max_retries: u32,
601    max_tokens: Option<u32>,
602    temperature: Option<f32>,
603    call_count: std::sync::atomic::AtomicUsize,
604}
605
606impl AnthropicProvider {
607    /// 从配置创建 Anthropic Provider
608    ///
609    /// 优先使用 api_key 字段,其次从环境变量读取。
610    pub fn new(config: &LlmSection) -> Result<Self> {
611        let api_key = config.api_key.clone()
612            .or_else(|| std::env::var(&config.api_key_env).ok())
613            .with_context(|| format!("Anthropic API Key 未设置(api_key 为空且环境变量 {} 未定义)。请设置环境变量 {},或编辑配置文件的 [llm] 段填入 api_key", config.api_key_env, config.api_key_env))?;
614        let base_url = config
615            .base_url
616            .clone()
617            .unwrap_or_else(|| "https://api.anthropic.com/v1".to_string());
618
619        let client = Client::builder()
620            // 不设总超时(同 OpenAiProvider:长生成由流式路径 SSE_IDLE_TIMEOUT 保护)
621            .build()
622            .context("创建 HTTP 客户端失败")?;
623
624        Ok(Self {
625            client,
626            api_key,
627            model: config.model.clone(),
628            base_url,
629            max_retries: MAX_RETRIES,
630            max_tokens: None,
631            temperature: None,
632            call_count: std::sync::atomic::AtomicUsize::new(0),
633        })
634    }
635}
636
637impl AnthropicProvider {
638    /// 构建 messages API 请求体:system 消息分离到顶层字段、
639    /// 非 system 消息进 messages、max_tokens 未配置时默认 4096;
640    /// stream 决定是否追加流式标记
641    fn build_messages_body(
642        &self,
643        messages: &[Message],
644        stream: bool,
645        max_tokens_override: Option<u32>,
646    ) -> serde_json::Value {
647        // 分离 system 消息与用户/助手消息
648        let system = messages.iter().find(|m| m.role == "system").map(|m| &m.content);
649        let non_system: Vec<&Message> = messages.iter().filter(|m| m.role != "system").collect();
650
651        // 将非 system 消息转换为 Anthropic 格式
652        let anthropic_messages: Vec<serde_json::Value> = non_system
653            .iter()
654            .map(|m| {
655                serde_json::json!({
656                    "role": if m.role == "user" { "user" } else { "assistant" },
657                    "content": m.content
658                })
659            })
660            .collect();
661
662        let mut body = serde_json::json!({
663            "model": self.model,
664            // 显式覆盖优先,回退构造时的 max_tokens(默认 4096)
665            "max_tokens": max_tokens_override.or(self.max_tokens).unwrap_or(4096),
666            "messages": anthropic_messages,
667        });
668        if let Some(s) = system {
669            body["system"] = serde_json::json!(s);
670        }
671        if let Some(temp) = self.temperature {
672            body["temperature"] = serde_json::json!(temp);
673        }
674        if stream {
675            body["stream"] = serde_json::json!(true);
676        }
677        body
678    }
679}
680
681impl LlmProvider for AnthropicProvider {
682    async fn complete_stream(&self, messages: &[Message]) -> Result<Vec<String>> {
683        self.call_count
684            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
685
686        let url = format!("{}/messages", self.base_url);
687        let body = self.build_messages_body(messages, true, None);
688
689        let resp = retry_with_backoff(self.max_retries, || {
690            self.client
691                .post(&url)
692                .header("x-api-key", &self.api_key)
693                .header("anthropic-version", "2023-06-01")
694                .json(&body)
695                .send()
696        })
697        .await?;
698
699        if !resp.status().is_success() {
700            let status = resp.status();
701            let text = resp.text().await.unwrap_or_default();
702            anyhow::bail!("Anthropic API 返回错误 ({}): {}", status, text);
703        }
704
705        // Anthropic SSE:data: 行内 type=content_block_delta 事件的 delta.text
706        collect_sse(resp, "data: ", |v| {
707            if v["type"] == "content_block_delta" {
708                v["delta"]["text"].as_str().map(|s| s.to_string())
709            } else {
710                None
711            }
712        })
713        .await
714    }
715
716    // complete 走 trait 默认实现(complete_stream 收集拼接)——
717    // 生产路径统一流式,无整读分支
718    async fn complete_with_budget(
719        &self,
720        messages: &[Message],
721        max_output_tokens: Option<u32>,
722    ) -> Result<String> {
723        self.call_count
724            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
725
726        let url = format!("{}/messages", self.base_url);
727        let body = self.build_messages_body(messages, true, max_output_tokens);
728
729        let resp = retry_with_backoff(self.max_retries, || {
730            self.client
731                .post(&url)
732                .header("x-api-key", &self.api_key)
733                .header("anthropic-version", "2023-06-01")
734                .json(&body)
735                .send()
736        })
737        .await?;
738
739        if !resp.status().is_success() {
740            let status = resp.status();
741            let text = resp.text().await.unwrap_or_default();
742            anyhow::bail!("Anthropic API 返回错误 ({}): {}", status, text);
743        }
744
745        // Anthropic SSE:data: 行内 type=content_block_delta 事件的 delta.text
746        let chunks = collect_sse(resp, "data: ", |v| {
747            if v["type"] == "content_block_delta" {
748                v["delta"]["text"].as_str().map(|s| s.to_string())
749            } else {
750                None
751            }
752        })
753        .await?;
754        Ok(chunks.concat())
755    }
756
757    fn call_count(&self) -> usize {
758        self.call_count.load(std::sync::atomic::Ordering::Relaxed)
759    }
760}
761
762/// Mock LLM Provider(用于测试和离线模式)
763///
764/// 不发起真实网络请求,返回固定的模拟响应。
765pub struct MockProvider {
766    call_count: std::sync::atomic::AtomicUsize,
767}
768
769impl MockProvider {
770    pub fn new() -> Self {
771        Self {
772            call_count: std::sync::atomic::AtomicUsize::new(0),
773        }
774    }
775}
776
777impl Default for MockProvider {
778    fn default() -> Self {
779        Self::new()
780    }
781}
782
783impl LlmProvider for MockProvider {
784    async fn complete(&self, _messages: &[Message]) -> Result<String> {
785        self.call_count
786            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
787        Ok(
788            r#"{"summary": "这是 Mock Provider 生成的模拟摘要", "key_entities": []}"#
789                .to_string(),
790        )
791    }
792
793    async fn complete_stream(&self, _messages: &[Message]) -> Result<Vec<String>> {
794        self.call_count
795            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
796        Ok(vec!["模拟流式响应 chunk".to_string()])
797    }
798
799    fn call_count(&self) -> usize {
800        self.call_count.load(std::sync::atomic::Ordering::Relaxed)
801    }
802}
803
804#[cfg(test)]
805mod tests {
806    use super::*;
807    use crate::config::schema::{LlmProviderType, LlmSection};
808    use std::io::{Read, Write};
809    use std::net::{TcpListener, TcpStream};
810    use std::sync::atomic::{AtomicUsize, Ordering};
811    use std::sync::{Arc, Mutex};
812
813    // ============ 本地 mock HTTP server ============
814    // 用 std 线程 + std::net 起阻塞式 mock,避免依赖 tokio net 特性,
815    // 与 reqwest 的异步请求天然解耦(无 runtime 饥饿/死锁问题)。
816
817    /// 捕获的 HTTP 请求(请求路径、请求头、请求体)
818    struct MockRequest {
819        path: String,
820        headers: Vec<(String, String)>,
821        body: String,
822    }
823
824    /// mock 服务器响应
825    struct MockResponse {
826        status: u16,
827        body: String,
828    }
829
830    /// 缓冲区中是否已出现完整请求头(含 \r\n\r\n 分隔符,其后可能还有请求体)
831    fn header_complete(buf: &[u8]) -> bool {
832        buf.windows(4).any(|w| w == b"\r\n\r\n")
833    }
834
835    /// 读取一个完整 HTTP 请求:请求行 + 请求头 + Content-Length 指定的请求体
836    fn read_request(stream: &mut TcpStream) -> MockRequest {
837        let mut buf = Vec::new();
838        let mut tmp = [0u8; 4096];
839        while !header_complete(&buf) {
840            match stream.read(&mut tmp) {
841                Ok(0) | Err(_) => break,
842                Ok(n) => buf.extend_from_slice(&tmp[..n]),
843            }
844        }
845        let head_end = buf.windows(4).position(|w| w == b"\r\n\r\n").unwrap_or(buf.len());
846        let head = String::from_utf8_lossy(&buf[..head_end]).to_string();
847        let mut lines = head.split("\r\n");
848        let path = lines.next().unwrap_or("").split_whitespace().nth(1).unwrap_or("").to_string();
849        let headers: Vec<(String, String)> = lines
850            .filter_map(|l| l.split_once(':'))
851            .map(|(k, v)| (k.trim().to_string(), v.trim().to_string()))
852            .collect();
853        let content_length = headers.iter()
854            .find(|(k, _)| k.eq_ignore_ascii_case("content-length"))
855            .and_then(|(_, v)| v.parse::<usize>().ok())
856            .unwrap_or(0);
857        // 头部结束标记 \r\n\r\n 之后才是请求体(head_end 指向标记起始,
858        // 体从 head_end + 4 开始;此前漏加 4 导致 body 前缀残留 \r\n\r\n)
859        const HEADER_SEP: usize = 4;
860        while buf.len() < head_end + HEADER_SEP + content_length {
861            match stream.read(&mut tmp) {
862                Ok(0) | Err(_) => break,
863                Ok(n) => buf.extend_from_slice(&tmp[..n]),
864            }
865        }
866        let body =
867            String::from_utf8_lossy(&buf[head_end + HEADER_SEP..head_end + HEADER_SEP + content_length])
868                .to_string();
869        MockRequest { path, headers, body }
870    }
871
872    /// 启动本地 mock HTTP server:每连接处理一个请求,由 handler 生成响应。
873    /// 响应带 Connection: close,迫使 reqwest 每次请求都新建连接。
874    /// 返回形如 http://127.0.0.1:<port> 的 base_url。
875    fn spawn_mock_server(
876        handler: impl Fn(MockRequest) -> MockResponse + Send + Sync + 'static,
877    ) -> String {
878        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
879        let base_url = format!("http://{}", listener.local_addr().unwrap());
880        let handler = Arc::new(handler);
881        std::thread::spawn(move || {
882            for stream in listener.incoming() {
883                let Ok(mut stream) = stream else { break };
884                let handler = handler.clone();
885                std::thread::spawn(move || {
886                    let req = read_request(&mut stream);
887                    let resp = handler(req);
888                    let reason = if resp.status == 200 { "OK" } else { "Error" };
889                    let raw = format!(
890                        "HTTP/1.1 {} {}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
891                        resp.status, reason, resp.body.len(), resp.body
892                    );
893                    let _ = stream.write_all(raw.as_bytes());
894                });
895            }
896        });
897        base_url
898    }
899
900    /// 构造指向本地 mock 的 OpenAI 配置(base_url 带 /v1 前缀,与生产默认一致)
901    fn openai_config(base_url: &str) -> LlmSection {
902        LlmSection {
903            provider: LlmProviderType::OpenAI,
904            model: "gpt-test".into(),
905            base_url: Some(format!("{}/v1", base_url)),
906            api_key: Some("test-key".into()),
907            api_key_env: "OPENAI_API_KEY".into(),
908        }
909    }
910
911    // ============ 原有测试 ============
912
913    #[tokio::test]
914    async fn test_mock_provider() {
915        let provider = MockProvider::new();
916
917        let messages = vec![Message::user("测试消息")];
918        let result = provider.complete(&messages).await;
919        assert!(result.is_ok());
920        assert!(result.unwrap().contains("模拟摘要"));
921        assert_eq!(provider.call_count(), 1);
922    }
923
924    #[tokio::test]
925    async fn test_message_constructors() {
926        let sys = Message::system("你好");
927        assert_eq!(sys.role, "system");
928        assert_eq!(sys.content, "你好");
929
930        let user = Message::user("测试");
931        assert_eq!(user.role, "user");
932
933        let asst = Message::assistant("回复");
934        assert_eq!(asst.role, "assistant");
935    }
936
937    // ============ 新增:OpenAI 请求构建、SSE 流式解析与重试 ============
938
939    #[tokio::test]
940    async fn test_openai_request_builds_correct_payload() {
941        // mock 服务器:捕获请求并返回 SSE 流(生产路径统一流式,
942        // complete() 走 trait 默认实现收集流式 chunks 拼接)
943        let captured = Arc::new(Mutex::new(None::<MockRequest>));
944        let captured_server = captured.clone();
945        let base_url = spawn_mock_server(move |req| {
946            *captured_server.lock().unwrap() = Some(req);
947            MockResponse {
948                status: 200,
949                body: r#"data: {"choices":[{"delta":{"content":"你好,这是 mock 回复"}}]}
950
951data: [DONE]
952
953"#
954                .into(),
955            }
956        });
957
958        let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
959        let messages = vec![Message::system("你是测试助手"), Message::user("你好")];
960        let reply = provider.complete(&messages).await.unwrap();
961
962        // 回复来自流式 chunks 的拼接(delta.content)
963        assert_eq!(reply, "你好,这是 mock 回复");
964
965        // 请求路径:base_url + /chat/completions
966        let req = captured.lock().unwrap().take().expect("应收到一次请求");
967
968        assert_eq!(req.path, "/v1/chat/completions");
969        // Authorization: Bearer <api_key>
970        let auth = req.headers.iter()
971            .find(|(k, _)| k.eq_ignore_ascii_case("authorization"))
972            .expect("应携带 Authorization 头");
973        assert_eq!(auth.1, "Bearer test-key");
974        // JSON body:model / messages(含 system 角色)/ stream:true(生产路径流式)/ 可选参数
975        let body: serde_json::Value = serde_json::from_str(&req.body).unwrap();
976        assert_eq!(body["model"], "gpt-test");
977        assert_eq!(body["messages"][0]["role"], "system");
978        assert_eq!(body["messages"][0]["content"], "你是测试助手");
979        assert_eq!(body["messages"][1]["role"], "user");
980        // v22 起 max_tokens/temperature 硬编码为 None(模型默认),断言不写入
981        assert!(body.get("max_tokens").is_none(), "硬编码后不应写 max_tokens");
982        assert!(body.get("temperature").is_none(), "硬编码后不应写 temperature");
983        assert_eq!(
984            body["stream"].as_bool(),
985            Some(true),
986            "生产路径必须请求流式响应(stream:true)"
987        );
988    }
989
990    #[tokio::test]
991    async fn test_openai_stream_parses_sse() {
992        // mock 返回 SSE 流:两条 delta 增量 + [DONE] 结束标记
993        let sse = concat!(
994            "data: {\"choices\":[{\"delta\":{\"content\":\"你\"}}]}\n\n",
995            "data: {\"choices\":[{\"delta\":{\"content\":\"好\"}}]}\n\n",
996            "data: [DONE]\n\n",
997        );
998        let base_url = spawn_mock_server(move |_req| MockResponse {
999            status: 200,
1000            body: sse.to_string(),
1001        });
1002
1003        let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
1004        let messages = vec![Message::user("你好")];
1005        let chunks = provider.complete_stream(&messages).await.unwrap();
1006
1007        // 增量按到达顺序拼接
1008        assert_eq!(chunks, vec!["你", "好"]);
1009        assert_eq!(chunks.join(""), "你好");
1010    }
1011
1012    /// A4:慢流响应(两段 SSE 之间间隔 300ms)不被总超时截断——
1013    /// client 不再设总超时,长生成由 60s 空闲超时保护,只要模型
1014    /// 持续产出就不会中途失败。响应分两次写入同一连接。
1015    #[tokio::test]
1016    async fn test_slow_stream_not_truncated_by_total_timeout() {
1017        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1018        let base_url = format!("http://{}", listener.local_addr().unwrap());
1019        std::thread::spawn(move || {
1020            for stream in listener.incoming() {
1021                let Ok(mut stream) = stream else { break };
1022                std::thread::spawn(move || {
1023                    let _req = read_request(&mut stream);
1024                    // 第一段立即写出,间隔 300ms 后再写第二段
1025                    let head = "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: close\r\n\r\n";
1026                    let _ = stream.write_all(head.as_bytes());
1027                    let _ = stream.write_all("data: {\"choices\":[{\"delta\":{\"content\":\"第一段\"}}]}\n\n".as_bytes());
1028                    let _ = stream.flush();
1029                    std::thread::sleep(Duration::from_millis(300));
1030                    let _ = stream.write_all("data: {\"choices\":[{\"delta\":{\"content\":\"第二段\"}}]}\n\n".as_bytes());
1031                    let _ = stream.write_all(b"data: [DONE]\n\n");
1032                    let _ = stream.flush();
1033                });
1034            }
1035        });
1036
1037        let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
1038        let messages = vec![Message::user("你好")];
1039        let reply = provider.complete(&messages).await.unwrap();
1040
1041        assert_eq!(reply, "第一段第二段", "慢流两段必须完整拼接(无总超时截断)");
1042    }
1043
1044    #[tokio::test]
1045    async fn test_retry_on_server_error() {
1046        // 第一次返回 500、第二次返回 200;5xx 在可重试白名单内,
1047        // 退避 500ms×2^n + 抖动 0-250ms,本测试约耗时 0.5-0.75s
1048        let attempts = Arc::new(AtomicUsize::new(0));
1049        let attempts_server = attempts.clone();
1050        let base_url = spawn_mock_server(move |_req| {
1051            let n = attempts_server.fetch_add(1, Ordering::Relaxed);
1052            if n == 0 {
1053                MockResponse { status: 500, body: "internal error".into() }
1054            } else {
1055                // 成功响应为 SSE 流(流式路径的输入格式)
1056                MockResponse { status: 200, body: "data: {\"choices\":[{\"delta\":{\"content\":\"重试成功\"}}]}\n\ndata: [DONE]\n\n".into() }
1057            }
1058        });
1059
1060        let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
1061        let messages = vec![Message::user("你好")];
1062        let reply = provider.complete(&messages).await.unwrap();
1063
1064        assert_eq!(reply, "重试成功");
1065        assert_eq!(attempts.load(Ordering::Relaxed), 2);
1066    }
1067
1068    #[tokio::test]
1069    async fn test_retry_on_429() {
1070        // 第一次返回 429(限流)、第二次返回 200:429 在可重试白名单内,断言请求次数 = 2
1071        let attempts = Arc::new(AtomicUsize::new(0));
1072        let attempts_server = attempts.clone();
1073        let base_url = spawn_mock_server(move |_req| {
1074            let n = attempts_server.fetch_add(1, Ordering::Relaxed);
1075            if n == 0 {
1076                MockResponse { status: 429, body: "rate limited".into() }
1077            } else {
1078                MockResponse { status: 200, body: "data: {\"choices\":[{\"delta\":{\"content\":\"限流后成功\"}}]}\n\ndata: [DONE]\n\n".into() }
1079            }
1080        });
1081
1082        let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
1083        let messages = vec![Message::user("你好")];
1084        let reply = provider.complete(&messages).await.unwrap();
1085
1086        assert_eq!(reply, "限流后成功");
1087        assert_eq!(attempts.load(Ordering::Relaxed), 2);
1088    }
1089
1090    #[tokio::test]
1091    async fn test_no_retry_on_401() {
1092        // 401 不在可重试白名单内:立即失败,断言仅 1 次请求
1093        let attempts = Arc::new(AtomicUsize::new(0));
1094        let attempts_server = attempts.clone();
1095        let base_url = spawn_mock_server(move |_req| {
1096            attempts_server.fetch_add(1, Ordering::Relaxed);
1097            MockResponse { status: 401, body: "unauthorized".into() }
1098        });
1099
1100        let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
1101        let messages = vec![Message::user("你好")];
1102        let result = provider.complete(&messages).await;
1103
1104        assert!(result.is_err());
1105        assert_eq!(attempts.load(Ordering::Relaxed), 1);
1106    }
1107
1108    #[tokio::test]
1109    async fn test_retry_exhausted_on_5xx() {
1110        // 永远返回 500:重试到上限(MAX_RETRIES=3 次尝试)后返回 Err
1111        let attempts = Arc::new(AtomicUsize::new(0));
1112        let attempts_server = attempts.clone();
1113        let base_url = spawn_mock_server(move |_req| {
1114            attempts_server.fetch_add(1, Ordering::Relaxed);
1115            MockResponse { status: 500, body: "internal error".into() }
1116        });
1117
1118        let provider = OpenAiProvider::new(&openai_config(&base_url), OpenAiProtocol::Chat).unwrap();
1119        let messages = vec![Message::user("你好")];
1120        let result = provider.complete(&messages).await;
1121
1122        assert!(result.is_err());
1123        assert_eq!(attempts.load(Ordering::Relaxed), MAX_RETRIES as usize);
1124    }
1125
1126    #[tokio::test]
1127    async fn test_retry_on_timeout() {
1128        // 第一次响应慢于客户端超时(触发超时重试),第二次立即成功。
1129        // 直接测共享骨架:短超时 client + 慢响应 mock,断言请求次数 = 2
1130        let attempts = Arc::new(AtomicUsize::new(0));
1131        let attempts_server = attempts.clone();
1132        let base_url = spawn_mock_server(move |_req| {
1133            let n = attempts_server.fetch_add(1, Ordering::Relaxed);
1134            if n == 0 {
1135                std::thread::sleep(Duration::from_millis(500));
1136            }
1137            MockResponse { status: 200, body: "{}".into() }
1138        });
1139
1140        let client = Client::builder()
1141            .timeout(Duration::from_millis(200))
1142            .build()
1143            .unwrap();
1144
1145        let resp = retry_with_backoff(MAX_RETRIES, || {
1146            client.get(format!("{}/t", base_url)).send()
1147        })
1148        .await
1149        .unwrap();
1150
1151        assert_eq!(resp.status(), 200);
1152        assert_eq!(attempts.load(Ordering::Relaxed), 2);
1153    }
1154
1155    #[test]
1156    fn test_parse_sse_openai_format() {
1157        // OpenAI 格式:data: 行内 choices[0].delta.content;[DONE] 行被跳过
1158        let sse = concat!(
1159            "data: {\"choices\":[{\"delta\":{\"content\":\"你\"}}]}\n\n",
1160            "data: {\"choices\":[{\"delta\":{\"content\":\"好\"}}]}\n\n",
1161            "data: [DONE]\n\n",
1162        );
1163        let chunks = parse_sse_stream(sse.as_bytes(), "data: ", |v| {
1164            v["choices"][0]["delta"]["content"]
1165                .as_str()
1166                .map(|s| s.to_string())
1167        });
1168        assert_eq!(chunks, vec!["你", "好"]);
1169    }
1170
1171    #[test]
1172    fn test_parse_sse_anthropic_format() {
1173        // Anthropic 格式:data: 行内 type=content_block_delta 事件的 delta.text,
1174        // 其余事件(message_start/message_stop)被跳过
1175        let sse = concat!(
1176            "data: {\"type\":\"message_start\",\"message\":{\"id\":\"m1\"}}\n\n",
1177            "data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"你\"}}\n\n",
1178            "data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"好\"}}\n\n",
1179            "data: {\"type\":\"message_stop\"}\n\n",
1180        );
1181        let chunks = parse_sse_stream(sse.as_bytes(), "data: ", |v| {
1182            if v["type"] == "content_block_delta" {
1183                v["delta"]["text"].as_str().map(|s| s.to_string())
1184            } else {
1185                None
1186            }
1187        });
1188        assert_eq!(chunks, vec!["你", "好"]);
1189    }
1190
1191    #[tokio::test]
1192    async fn test_anthropic_stream_parses_sse() {
1193        // Anthropic 流式:mock base_url 生效(修复 stream 硬编码 URL)+ SSE 解析
1194        let sse = concat!(
1195            "data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"克\"}}\n\n",
1196            "data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"劳\"}}\n\n",
1197            "data: {\"type\":\"message_stop\"}\n\n",
1198        );
1199        let base_url = spawn_mock_server(move |_req| MockResponse {
1200            status: 200,
1201            body: sse.to_string(),
1202        });
1203
1204        let config = LlmSection {
1205            provider: LlmProviderType::Anthropic,
1206            model: "claude-test".into(),
1207            base_url: Some(base_url),
1208            api_key: Some("sk-ant-test".into()),
1209            api_key_env: "ANTHROPIC_API_KEY".into(),
1210        };
1211        let provider = AnthropicProvider::new(&config).unwrap();
1212        let messages = vec![Message::user("你好")];
1213        let chunks = provider.complete_stream(&messages).await.unwrap();
1214
1215        assert_eq!(chunks, vec!["克", "劳"]);
1216    }
1217
1218    #[test]
1219    fn test_anthropic_provider_construction() {
1220        let config = LlmSection {
1221            provider: LlmProviderType::Anthropic,
1222            model: "claude-test".into(),
1223            base_url: None,
1224            api_key: Some("sk-ant-test".into()),
1225            api_key_env: "ANTHROPIC_API_KEY".into(),
1226        };
1227        let provider = AnthropicProvider::new(&config).unwrap();
1228        assert_eq!(provider.call_count(), 0);
1229    }
1230
1231    /// Anthropic 请求构建:base_url 可配置后(与 OpenAiProvider 对齐),
1232    /// 用本地 mock 断言 x-api-key / anthropic-version 头与请求体
1233    /// (system 消息分离到顶层字段、非 system 消息进 messages、max_tokens 默认 4096)
1234    #[tokio::test]
1235    async fn test_anthropic_request_builds_correct_payload() {
1236        let captured = Arc::new(Mutex::new(None::<MockRequest>));
1237        let captured_server = captured.clone();
1238        let base_url = spawn_mock_server(move |req| {
1239            *captured_server.lock().unwrap() = Some(req);
1240            MockResponse {
1241                status: 200,
1242                // Anthropic 流式格式(complete 走流式默认实现后 mock 必须返回 SSE)
1243                body: concat!(
1244                    "data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"claude 回复\"}}\n\n",
1245                    "data: {\"type\":\"message_stop\"}\n\n",
1246                )
1247                .into(),
1248            }
1249        });
1250
1251        let config = LlmSection {
1252            provider: LlmProviderType::Anthropic,
1253            model: "claude-test".into(),
1254            base_url: Some(base_url),
1255            api_key: Some("sk-ant-test".into()),
1256            api_key_env: "ANTHROPIC_API_KEY".into(),
1257        };
1258        let provider = AnthropicProvider::new(&config).unwrap();
1259        let messages = vec![
1260            Message::system("你是助手"),
1261            Message::user("你好"),
1262            Message::assistant("在的"),
1263        ];
1264        let reply = provider.complete(&messages).await.unwrap();
1265        assert_eq!(reply, "claude 回复");
1266
1267        let req = captured.lock().unwrap().take().expect("应收到一次请求");
1268        assert_eq!(req.path, "/messages");
1269        let api_key_header = req
1270            .headers
1271            .iter()
1272            .find(|(k, _)| k.eq_ignore_ascii_case("x-api-key"))
1273            .expect("应携带 x-api-key 头");
1274        assert_eq!(api_key_header.1, "sk-ant-test");
1275        let version_header = req
1276            .headers
1277            .iter()
1278            .find(|(k, _)| k.eq_ignore_ascii_case("anthropic-version"))
1279            .expect("应携带 anthropic-version 头");
1280        assert_eq!(version_header.1, "2023-06-01");
1281        let body: serde_json::Value = serde_json::from_str(&req.body).unwrap();
1282        assert_eq!(body["model"], "claude-test");
1283        assert_eq!(body["max_tokens"].as_u64(), Some(4096), "max_tokens 未配置时默认 4096");
1284        // system 消息分离到顶层 system 字段
1285        assert_eq!(body["system"], "你是助手");
1286        let msgs = body["messages"].as_array().unwrap();
1287        assert_eq!(msgs.len(), 2, "非 system 消息才进 messages");
1288        assert_eq!(msgs[0]["role"], "user");
1289        assert_eq!(msgs[0]["content"], "你好");
1290        assert_eq!(msgs[1]["role"], "assistant");
1291        assert_eq!(msgs[1]["content"], "在的");
1292    }
1293
1294    // ============ v17 B4/B5:Responses 协议(openai provider) ============
1295
1296    /// Responses 流式 SSE 解析:语义化事件(response.created →
1297    /// output_text.delta ×2 → response.completed),无 [DONE] 终止符,
1298    /// 流结束即终止(collect_sse 兼容)
1299    #[tokio::test]
1300    async fn test_responses_stream_parses_semantic_sse() {
1301        let sse = concat!(
1302            "data: {\"type\":\"response.created\",\"response\":{\"id\":\"r1\"}}\n\n",
1303            "data: {\"type\":\"response.output_text.delta\",\"sequence_number\":0,\"delta\":\"你\"}\n\n",
1304            "data: {\"type\":\"response.output_text.delta\",\"sequence_number\":1,\"delta\":\"好\"}\n\n",
1305            "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"r1\"}}\n\n",
1306        );
1307        let base_url = spawn_mock_server(move |_req| MockResponse {
1308            status: 200,
1309            body: sse.to_string(),
1310        });
1311
1312        let config = LlmSection {
1313            provider: LlmProviderType::OpenAI,
1314            model: "deepseek-v4-flash".into(),
1315            base_url: Some(format!("{}/v1", base_url)),
1316            api_key: Some("test-key".into()),
1317            api_key_env: "DEEPSEEK_API_KEY".into(),
1318        };
1319        let provider = OpenAiProvider::new(&config, OpenAiProtocol::Responses).unwrap();
1320        let chunks = provider.complete_stream(&[Message::user("你好")]).await.unwrap();
1321        assert_eq!(chunks, vec!["你", "好"], "语义化事件应提取 delta 文本");
1322        assert_eq!(provider.call_count(), 1);
1323    }
1324
1325    /// Responses 请求体:input/instructions/max_output_tokens(v17 B4 协议差异)
1326    #[tokio::test]
1327    async fn test_responses_request_builds_correct_payload() {
1328        let captured = Arc::new(Mutex::new(None::<MockRequest>));
1329        let captured_server = captured.clone();
1330        let base_url = spawn_mock_server(move |req| {
1331            *captured_server.lock().unwrap() = Some(req);
1332            MockResponse {
1333                status: 200,
1334                body: "data: {\"type\":\"response.completed\"}\n\n".into(),
1335            }
1336        });
1337
1338        let config = LlmSection {
1339            provider: LlmProviderType::OpenAI,
1340            model: "deepseek-v4-flash".into(),
1341            base_url: Some(format!("{}/v1", base_url)),
1342            api_key: Some("test-key".into()),
1343            api_key_env: "DEEPSEEK_API_KEY".into(),
1344        };
1345        let provider = OpenAiProvider::new(&config, OpenAiProtocol::Responses).unwrap();
1346        let messages = vec![Message::system("你是助手"), Message::user("你好")];
1347        let _ = provider.complete_stream(&messages).await.unwrap();
1348
1349        let req = captured.lock().unwrap().take().expect("应收到一次请求");
1350        assert_eq!(req.path, "/v1/responses", "Responses 协议应请求 /responses 端点");
1351        let body: serde_json::Value = serde_json::from_str(&req.body).unwrap();
1352        assert_eq!(body["model"], "deepseek-v4-flash");
1353        // system 消息分离到顶层 instructions;user 进 input items
1354        assert_eq!(body["instructions"], "你是助手");
1355        let input = body["input"].as_array().unwrap();
1356        assert_eq!(input.len(), 1, "非 system 消息才进 input");
1357        assert_eq!(input[0]["role"], "user");
1358        assert_eq!(input[0]["content"][0]["type"], "input_text");
1359        assert_eq!(input[0]["content"][0]["text"], "你好");
1360        // token 上限与温度参数:v22 起硬编码为 None(交给模型默认),
1361        // 测试断言"既不写 max_output_tokens 也不写 max_tokens/temperature"
1362        assert!(body.get("max_output_tokens").is_none(), "硬编码后不应写 max_output_tokens");
1363        assert!(body.get("max_tokens").is_none(), "Responses 不得用 max_tokens 参数名");
1364        assert!(body.get("temperature").is_none(), "硬编码后不应写 temperature");
1365        assert_eq!(body["stream"].as_bool(), Some(true));
1366    }
1367
1368    /// v17 B5:Responses 端点不支持(404)→ 自动回退 chat/completions 重发成功
1369    #[tokio::test]
1370    async fn test_responses_falls_back_to_chat_on_404() {
1371        let requests: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
1372        let requests_server = requests.clone();
1373        let base_url = spawn_mock_server(move |req| {
1374            requests_server.lock().unwrap().push(req.path.clone());
1375            if req.path.ends_with("/responses") {
1376                MockResponse { status: 404, body: "not found".into() }
1377            } else {
1378                MockResponse {
1379                    status: 200,
1380                    body: "data: {\"choices\":[{\"delta\":{\"content\":\"回退成功\"}}]}\n\ndata: [DONE]\n\n".into(),
1381                }
1382            }
1383        });
1384
1385        let config = LlmSection {
1386            provider: LlmProviderType::OpenAI,
1387            model: "deepseek-v4-flash".into(),
1388            base_url: Some(format!("{}/v1", base_url)),
1389            api_key: Some("test-key".into()),
1390            api_key_env: "DEEPSEEK_API_KEY".into(),
1391        };
1392        let provider = OpenAiProvider::new(&config, OpenAiProtocol::Responses).unwrap();
1393        let chunks = provider.complete_stream(&[Message::user("你好")]).await.unwrap();
1394        assert_eq!(chunks.join(""), "回退成功");
1395        let paths = requests.lock().unwrap();
1396        assert_eq!(paths.len(), 2, "应请求 responses + chat 两次");
1397        assert!(paths[0].ends_with("/responses"), "第一次应请求 responses: {:?}", paths);
1398        assert!(paths[1].ends_with("/chat/completions"), "回退应请求 chat/completions: {:?}", paths);
1399    }
1400
1401    /// v22 修复:评测裁判完整调用带显式输出预算——请求体必须写入
1402    /// max_output_tokens=16384(reasoning 型模型预算不足时只有
1403    /// reasoning 块没有 message,见 BENCH_MAX_OUTPUT_TOKENS 文档)
1404    #[tokio::test]
1405    async fn test_complete_with_budget_sets_max_output_tokens() {
1406        let captured = Arc::new(Mutex::new(None::<MockRequest>));
1407        let captured_server = captured.clone();
1408        let base_url = spawn_mock_server(move |req| {
1409            *captured_server.lock().unwrap() = Some(req);
1410            MockResponse {
1411                status: 200,
1412                body: "data: {\"type\":\"response.output_text.delta\",\"delta\":\"{\\\"rubrics\\\":[]}\"}\n\n"
1413                    .into(),
1414            }
1415        });
1416
1417        let config = LlmSection {
1418            provider: LlmProviderType::OpenAI,
1419            model: "deepseek-v4-flash".into(),
1420            base_url: Some(format!("{}/v1", base_url)),
1421            api_key: Some("test-key".into()),
1422            api_key_env: "DEEPSEEK_API_KEY".into(),
1423        };
1424        let provider = OpenAiProvider::new(&config, OpenAiProtocol::Responses).unwrap();
1425        let out = provider
1426            .complete_with_budget(&[Message::user("你好")], Some(16384))
1427            .await
1428            .unwrap();
1429        assert_eq!(out, "{\"rubrics\":[]}", "带预算调用应返回完整文本");
1430
1431        let req = captured.lock().unwrap().take().expect("应收到一次请求");
1432        assert_eq!(req.path, "/v1/responses");
1433        let body: serde_json::Value = serde_json::from_str(&req.body).unwrap();
1434        assert_eq!(body["max_output_tokens"].as_u64(), Some(16384), "预算应写入 max_output_tokens");
1435        assert_eq!(body["stream"].as_bool(), Some(true), "带预算路径仍走流式");
1436    }
1437}