1use std::time::Duration;
2
3use anyhow::{Context, Result};
4use reqwest::Client;
5
6use crate::config::schema::LlmSection;
7
8#[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#[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 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 fn call_count(&self) -> usize {
75 0
76 }
77}
78
79pub 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
126pub(crate) const MAX_RETRIES: u32 = 3;
130
131fn 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
142fn is_retryable_status(status: reqwest::StatusCode) -> bool {
145 status == reqwest::StatusCode::TOO_MANY_REQUESTS || status.is_server_error()
146}
147
148pub(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#[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
234const 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 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 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 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 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
326pub enum OpenAiProtocol {
327 Responses,
328 Chat,
329}
330
331pub struct OpenAiProvider {
333 client: Client,
334 api_key: String,
335 model: String,
336 base_url: String,
337 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 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 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 .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 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 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 fn build_responses_body(
424 &self,
425 messages: &[Message],
426 stream: bool,
427 max_output_tokens_override: Option<u32>,
428 ) -> serde_json::Value {
429 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 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 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 fn call_count(&self) -> usize {
493 self.call_count.load(std::sync::atomic::Ordering::Relaxed)
494 }
495}
496
497impl OpenAiProvider {
498 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 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 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 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 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
590pub struct AnthropicProvider {
594 client: Client,
595 api_key: String,
596 model: String,
597 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 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 .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 fn build_messages_body(
642 &self,
643 messages: &[Message],
644 stream: bool,
645 max_tokens_override: Option<u32>,
646 ) -> serde_json::Value {
647 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 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": 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 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 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 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
762pub 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 struct MockRequest {
819 path: String,
820 headers: Vec<(String, String)>,
821 body: String,
822 }
823
824 struct MockResponse {
826 status: u16,
827 body: String,
828 }
829
830 fn header_complete(buf: &[u8]) -> bool {
832 buf.windows(4).any(|w| w == b"\r\n\r\n")
833 }
834
835 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 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 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 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 #[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 #[tokio::test]
940 async fn test_openai_request_builds_correct_payload() {
941 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 assert_eq!(reply, "你好,这是 mock 回复");
964
965 let req = captured.lock().unwrap().take().expect("应收到一次请求");
967
968 assert_eq!(req.path, "/v1/chat/completions");
969 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 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 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 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 assert_eq!(chunks, vec!["你", "好"]);
1009 assert_eq!(chunks.join(""), "你好");
1010 }
1011
1012 #[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 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 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 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 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 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 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 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 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 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 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 #[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 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 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 #[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 #[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 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 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 #[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 #[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}