// src/clients/llm/chat.rs // // ChatCompleter trait + LlmClient 实现(对话补全、流式对话、图片分析)。 use async_trait::async_trait; use futures_util::StreamExt; use reqwest::Client; use serde::Deserialize; use std::sync::Arc; use tracing::error; use super::types::{ ChatMessage, CompletionResult, FunctionCall, StreamEvent, TokenUsage, ToolCall, ToolDefinition, }; /// LLM 对话补全的抽象 trait,支持真实客户端和测试 mock。 #[async_trait] pub trait ChatCompleter: Send + Sync { async fn complete(&self, system_prompt: &str, user_content: &str) -> anyhow::Result; } #[derive(Clone, Debug)] pub struct LlmClient { api_key: String, api_base: String, model: Arc>, client: Client, } #[async_trait] impl ChatCompleter for LlmClient { async fn complete(&self, system_prompt: &str, user_content: &str) -> anyhow::Result { self.chat_completion(system_prompt, user_content).await } } impl LlmClient { pub fn new(api_key: String, api_base: String, model: String) -> Self { LlmClient { api_key, api_base, model: Arc::new(std::sync::RwLock::new(model)), client: Client::new(), } } pub fn model(&self) -> String { self.model.read().unwrap_or_else(|e| e.into_inner()).clone() } /// 运行时切换模型(用于故障转移)。 pub fn set_model(&self, model: String) { if let Ok(mut m) = self.model.write() { *m = model; tracing::info!("[LlmClient] 模型切换为: {}", m); } } pub fn api_base(&self) -> &str { &self.api_base } pub fn api_key(&self) -> &str { &self.api_key } pub async fn chat_completion( &self, system_prompt: &str, user_content: &str, ) -> anyhow::Result { let url = format!("{}/chat/completions", self.api_base); let payload = serde_json::json!({ "model": self.model(), "messages": [ { "role": "system", "content": system_prompt }, { "role": "user", "content": user_content } ], "temperature": 0.3 }); let response = self .client .post(&url) .header("Authorization", format!("Bearer {}", self.api_key)) .header("Content-Type", "application/json") .json(&payload) .send() .await?; if !response.status().is_success() { let status = response.status(); let body = response.text().await.unwrap_or_default(); error!("LLM 接口调用失败: 状态码={}, 报错={}", status, body); return Err(anyhow::anyhow!("大模型接口返回错误状态: {}", status)); } #[derive(Deserialize)] struct Message { content: String, } #[derive(Deserialize)] struct Choice { message: Message, } #[derive(Deserialize)] struct LLMResponse { choices: Vec, } let res_data: LLMResponse = response.json().await?; if let Some(choice) = res_data.choices.first() { Ok(choice.message.content.clone()) } else { Err(anyhow::anyhow!("大模型返回空翻译选项集")) } } /// 多轮对话补全(含原生 Tool Calling 支持),非流式 pub async fn chat( &self, messages: &[ChatMessage], tools: &[ToolDefinition], enable_thinking: bool, ) -> anyhow::Result { let url = format!("{}/chat/completions", self.api_base); let mut payload = serde_json::json!({ "model": self.model(), "messages": messages, "temperature": 0.3 }); // 前端可控的思考模式开关(仅对千问/DashScope 启用 enable_thinking 参数)(仅对千问/DashScope 启用 enable_thinking 参数) if enable_thinking && (self.api_base.contains("dashscope.aliyuncs.com") || self.model().to_lowercase().contains("qwen")) { if let Some(obj) = payload.as_object_mut() { obj.insert("enable_thinking".to_string(), serde_json::json!(true)); } } if !tools.is_empty() { payload["tools"] = serde_json::to_value(tools)?; } let response = self .client .post(&url) .header("Authorization", format!("Bearer {}", self.api_key)) .header("Content-Type", "application/json") .json(&payload) .send() .await?; if !response.status().is_success() { let status = response.status(); let body = response.text().await.unwrap_or_default(); error!("LLM chat 接口调用失败: 状态码={}, 报错={}", status, body); return Err(anyhow::anyhow!( "大模型 chat 接口返回错误状态: {} - {}", status, body )); } #[derive(Deserialize)] struct ResponseMessage { content: Option, tool_calls: Option>, reasoning_content: Option, } #[derive(Deserialize)] struct Choice { message: ResponseMessage, } #[derive(Deserialize)] struct ChatResponse { choices: Vec, usage: Option, } let res_data: ChatResponse = response.json().await?; if let Some(choice) = res_data.choices.into_iter().next() { Ok(CompletionResult { content: choice.message.content, reasoning_content: choice.message.reasoning_content, tool_calls: choice.message.tool_calls.unwrap_or_default(), usage: res_data.usage.unwrap_or_default(), }) } else { Err(anyhow::anyhow!("大模型返回空对话选项集")) } } /// 流式分析图片,通过 `on_delta` 回调推送增量文本,返回完整结果。 pub async fn analyze_image_stream( &self, system_prompt: &str, question: &str, image_base64: &str, mime_type: &str, mut on_delta: impl FnMut(&str), ) -> anyhow::Result { let user_content = serde_json::json!([ { "type": "text", "text": question }, { "type": "image_url", "image_url": { "url": format!("data:{};base64,{}", mime_type, image_base64) } } ]); let payload = serde_json::json!({ "model": self.model(), "messages": [ { "role": "system", "content": system_prompt }, { "role": "user", "content": user_content } ], "temperature": 0.3, "stream": true, "stream_options": { "include_usage": true } }); let url = format!("{}/chat/completions", self.api_base); let response = self .client .post(&url) .header("Authorization", format!("Bearer {}", self.api_key)) .header("Content-Type", "application/json") .json(&payload) .send() .await?; if !response.status().is_success() { let status = response.status(); let body = response.text().await.unwrap_or_default(); return Err(anyhow::anyhow!( "视觉模型调用失败 (HTTP {}): {}", status, body )); } use futures_util::StreamExt; let mut stream = response.bytes_stream(); let mut accumulated = String::new(); let mut buffer = String::new(); while let Some(chunk) = stream.next().await { let chunk = chunk?; buffer.push_str(&String::from_utf8_lossy(&chunk)); while let Some(pos) = buffer.find('\n') { let line = buffer[..pos].trim().to_string(); buffer = buffer[pos + 1..].to_string(); if line.is_empty() || !line.starts_with("data: ") { continue; } let data = line[6..].trim().to_string(); if data == "[DONE]" { continue; } if let Ok(event) = serde_json::from_str::(&data) { if let Some(delta) = event["choices"][0]["delta"]["content"].as_str() { accumulated.push_str(delta); on_delta(delta); } } } } if accumulated.is_empty() { return Err(anyhow::anyhow!("视觉模型返回空结果")); } Ok(accumulated) } /// 流式多轮对话补全(含 Tool Call 碎片累积还原),返回 SSE 事件流 pub async fn chat_stream( &self, messages: &[ChatMessage], tools: &[ToolDefinition], enable_thinking: bool, ) -> anyhow::Result> { let url = format!("{}/chat/completions", self.api_base); let mut payload = serde_json::json!({ "model": self.model(), "messages": messages, "temperature": 0.3, "stream": true, "stream_options": { "include_usage": true } }); // 前端可控的思考模式开关(仅对千问/DashScope 启用 enable_thinking 参数) if enable_thinking && (self.api_base.contains("dashscope.aliyuncs.com") || self.model().to_lowercase().contains("qwen")) { if let Some(obj) = payload.as_object_mut() { obj.insert("enable_thinking".to_string(), serde_json::json!(true)); } } if !tools.is_empty() { payload["tools"] = serde_json::to_value(tools)?; } let response = self .client .post(&url) .header("Authorization", format!("Bearer {}", self.api_key)) .header("Content-Type", "application/json") .json(&payload) .send() .await?; if !response.status().is_success() { let status = response.status(); let status_code = status.as_u16(); let retry_after = response .headers() .get("retry-after") .and_then(|v| v.to_str().ok()) .and_then(|v| v.parse::().ok()); let body = response.text().await.unwrap_or_default(); error!("LLM stream 接口调用失败: 状态码={}, 报错={}", status, body); return Err(anyhow::anyhow!( "HTTP {}: {} | retry_after={:?}", status_code, body, retry_after )); } let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); let mut byte_stream = response.bytes_stream(); tokio::spawn(async move { // 流式 Tool Call 碎片累积器 let mut tool_call_accumulators: std::collections::HashMap< usize, (String, String, String), > = std::collections::HashMap::new(); let mut buffer = String::new(); while let Some(chunk_result) = byte_stream.next().await { let chunk = match chunk_result { Ok(c) => c, Err(e) => { let _ = tx.send(StreamEvent::Error(format!("流式读取错误: {}", e))); break; } }; buffer.push_str(&String::from_utf8_lossy(&chunk)); // 按行解析 SSE 事件 while let Some(line_end) = buffer.find('\n') { let line = buffer[..line_end].trim().to_string(); buffer = buffer[line_end + 1..].to_string(); if line.is_empty() || !line.starts_with("data: ") { continue; } let data = &line[6..]; if data == "[DONE]" { // 流结束前,将累积的 tool calls 还原并发送 if !tool_call_accumulators.is_empty() { let mut indices: Vec = tool_call_accumulators.keys().cloned().collect(); indices.sort(); let tool_calls: Vec = indices .into_iter() .map(|idx| { let (id, name, args) = tool_call_accumulators.remove(&idx).unwrap(); ToolCall { id: if id.is_empty() { format!( "call_{}", &uuid::Uuid::new_v4().to_string()[..8] ) } else { id }, call_type: "function".to_string(), function: FunctionCall { name, arguments: args, }, } }) .collect(); let _ = tx.send(StreamEvent::ToolCallsComplete(tool_calls)); } let _ = tx.send(StreamEvent::Done); return; } // 解析 JSON delta let parsed: serde_json::Value = match serde_json::from_str(data) { Ok(v) => v, Err(_) => continue, }; // 提取 usage(在最后一条 chunk 中 stream_options 返回) if let Some(usage_val) = parsed.get("usage") { if !usage_val.is_null() { if let Ok(usage) = serde_json::from_value::(usage_val.clone()) { let _ = tx.send(StreamEvent::Usage(usage)); } } } if let Some(choices) = parsed.get("choices").and_then(|c| c.as_array()) { if let Some(choice) = choices.first() { let delta = &choice["delta"]; // 文本增量 if let Some(content) = delta.get("content").and_then(|c| c.as_str()) { if !content.is_empty() { let _ = tx.send(StreamEvent::TextDelta(content.to_string())); } } // 推理内容增量 if let Some(reasoning) = delta.get("reasoning_content").and_then(|c| c.as_str()) { if !reasoning.is_empty() { let _ = tx.send(StreamEvent::ReasoningDelta(reasoning.to_string())); } } // 工具调用增量 if let Some(tool_calls) = delta.get("tool_calls").and_then(|t| t.as_array()) { for tc in tool_calls { let index = tc.get("index").and_then(|i| i.as_u64()).unwrap_or(0) as usize; let id = tc .get("id") .and_then(|i| i.as_str()) .map(|s| s.to_string()); let fn_name = tc .get("function") .and_then(|f| f.get("name")) .and_then(|n| n.as_str()) .map(|s| s.to_string()); let fn_args = tc .get("function") .and_then(|f| f.get("arguments")) .and_then(|a| a.as_str()) .unwrap_or(""); let entry = tool_call_accumulators.entry(index).or_insert_with(|| { (String::new(), String::new(), String::new()) }); if let Some(ref id_str) = id { entry.0 = id_str.clone(); } if let Some(ref name_str) = fn_name { entry.1 = name_str.clone(); } entry.2.push_str(fn_args); let _ = tx.send(StreamEvent::ToolCallDelta { index, id, name: fn_name, arguments_delta: fn_args.to_string(), }); } } // 检查是否 finish_reason == "tool_calls" if let Some(finish_reason) = choice.get("finish_reason").and_then(|f| f.as_str()) { if finish_reason == "tool_calls" && !tool_call_accumulators.is_empty() { let mut indices: Vec = tool_call_accumulators.keys().cloned().collect(); indices.sort(); let tool_calls: Vec = indices .into_iter() .map(|idx| { let (id, name, args) = tool_call_accumulators.remove(&idx).unwrap(); ToolCall { id, call_type: "function".to_string(), function: FunctionCall { name, arguments: args, }, } }) .collect(); let _ = tx.send(StreamEvent::ToolCallsComplete(tool_calls)); } } } } } } // 异常结束(连接断开等) if !tool_call_accumulators.is_empty() { let mut indices: Vec = tool_call_accumulators.keys().cloned().collect(); indices.sort(); let tool_calls: Vec = indices .into_iter() .map(|idx| { let (id, name, args) = tool_call_accumulators.remove(&idx).unwrap(); ToolCall { id: if id.is_empty() { format!("call_{}", &uuid::Uuid::new_v4().to_string()[..8]) } else { id }, call_type: "function".to_string(), function: FunctionCall { name, arguments: args, }, } }) .collect(); let _ = tx.send(StreamEvent::ToolCallsComplete(tool_calls)); } let _ = tx.send(StreamEvent::Done); }); Ok(rx) } } #[cfg(test)] mod tests { use super::*; use crate::clients::llm::types::{ ChatMessage, FunctionCall, MessageRole, ToolCall, ToolDefinition, }; #[test] fn test_llm_client_initialization() { let client = LlmClient::new("key".to_string(), "base".to_string(), "model".to_string()); assert_eq!(client.api_key(), "key"); assert_eq!(client.api_base(), "base"); assert_eq!(client.model(), "model"); } #[test] fn test_chat_message_construction() { // 测试系统消息构造 let sys = ChatMessage::system("你是一个助手"); assert_eq!(sys.role, MessageRole::System); assert_eq!(sys.text(), Some("你是一个助手")); assert!(sys.tool_calls.is_none()); // 测试用户消息构造 let user = ChatMessage::user("你好"); assert_eq!(user.role, MessageRole::User); assert_eq!(user.text(), Some("你好")); // 测试助手消息构造 let assistant = ChatMessage::assistant("你好!有什么可以帮助你的吗?"); assert_eq!(assistant.role, MessageRole::Assistant); assert_eq!(assistant.text(), Some("你好!有什么可以帮助你的吗?")); // 测试工具结果消息构造 let tool_result = ChatMessage::tool_result("call_123", r#"{"result": 42}"#); assert_eq!(tool_result.role, MessageRole::Tool); assert_eq!(tool_result.tool_call_id.as_deref(), Some("call_123")); assert_eq!(tool_result.text(), Some(r#"{"result": 42}"#)); // 测试带工具调用的助手消息 let tool_calls = vec![ToolCall { id: "call_abc".to_string(), call_type: "function".to_string(), function: FunctionCall { name: "get_weather".to_string(), arguments: r#"{"city": "北京"}"#.to_string(), }, }]; let assistant_tc = ChatMessage::assistant_with_tool_calls(None, tool_calls.clone()); assert_eq!(assistant_tc.role, MessageRole::Assistant); assert!(assistant_tc.text().is_none()); assert_eq!(assistant_tc.tool_calls.as_ref().unwrap().len(), 1); assert_eq!( assistant_tc.tool_calls.as_ref().unwrap()[0].function.name, "get_weather" ); } #[test] fn test_tool_definition_construction() { let tool = ToolDefinition::new( "search", "搜索互联网内容", serde_json::json!({ "type": "object", "properties": { "query": { "type": "string", "description": "搜索关键词" } }, "required": ["query"] }), ); assert_eq!(tool.tool_type, "function"); assert_eq!(tool.function.name, "search"); assert_eq!(tool.function.description, "搜索互联网内容"); } #[test] fn test_chat_message_serialization() { // 测试序列化时 skip_serializing_if 生效 let msg = ChatMessage::user("测试消息"); let json = serde_json::to_value(&msg).unwrap(); assert_eq!(json.get("role").unwrap(), "user"); assert_eq!(json.get("content").unwrap(), "测试消息"); // None 字段不应出现在 JSON 中 assert!(json.get("tool_call_id").is_none()); assert!(json.get("tool_calls").is_none()); assert!(json.get("name").is_none()); assert!(json.get("reasoning_content").is_none()); } #[test] fn test_chat_message_reasoning_serialization() { // 测试带有 reasoning_content 的消息序列化和反序列化 let msg = ChatMessage::assistant_with_reasoning( Some("回答内容".to_string()), Some("这是思考过程".to_string()), None, ); let json = serde_json::to_value(&msg).unwrap(); assert_eq!(json.get("role").unwrap(), "assistant"); assert_eq!(json.get("content").unwrap(), "回答内容"); assert_eq!(json.get("reasoning_content").unwrap(), "这是思考过程"); assert!(json.get("tool_calls").is_none()); // 测试反序列化 let deserialized: ChatMessage = serde_json::from_value(json).unwrap(); assert_eq!(deserialized.role, MessageRole::Assistant); assert_eq!(deserialized.text(), Some("回答内容")); assert_eq!( deserialized.reasoning_content.as_deref(), Some("这是思考过程") ); } #[test] fn test_message_role_serialization() { // 测试角色枚举的 rename_all = "lowercase" 序列化 let role_json = serde_json::to_string(&MessageRole::System).unwrap(); assert_eq!(role_json, r#""system""#); let role_json = serde_json::to_string(&MessageRole::Assistant).unwrap(); assert_eq!(role_json, r#""assistant""#); let role_json = serde_json::to_string(&MessageRole::Tool).unwrap(); assert_eq!(role_json, r#""tool""#); // 测试反序列化 let role: MessageRole = serde_json::from_str(r#""user""#).unwrap(); assert_eq!(role, MessageRole::User); } /// 真实网络集成测试 —— 需要配置 LLM_API_KEY 和 EMBEDDING_API_KEY #[tokio::test] #[ignore] async fn test_live_llm_and_embedding() -> anyhow::Result<()> { let config = crate::Config::from_env(); println!("================= 开始大模型与向量模型真实网络集成测试 ================="); // 1. 测试 LlmClient if config.llm_api_key.is_empty() { println!("警告: 未在环境配置中检测到 LLM_API_KEY,跳过 LlmClient 集成测试。"); } else { println!( "测试大模型: {} (API Base: {})", config.llm_model, config.llm_api_base ); let llm = LlmClient::new( config.llm_api_key.clone(), config.llm_api_base.clone(), config.llm_model.clone(), ); match llm .chat_completion("You are a helpful assistant.", "Say Hello!") .await { Ok(reply) => { println!("LlmClient 响应成功: {}", reply.trim()); assert!(!reply.trim().is_empty(), "错误: 大模型返回了空响应"); } Err(e) => panic!("LlmClient 接口调用失败: {}", e), } } // 2. 测试 EmbeddingClient if config.embedding_api_key.is_empty() { println!( "警告: 未在环境配置中检测到 EMBEDDING_API_KEY,跳过 EmbeddingClient 集成测试。" ); } else { println!( "测试向量模型: {} (API Base: {})", config.embedding_model, config.embedding_api_base ); let embedding_client = crate::clients::llm::EmbeddingClient::new( config.embedding_api_key.clone(), config.embedding_api_base.clone(), config.embedding_model.clone(), ); let test_text = "active galactic nucleus"; match embedding_client.create_embedding(test_text).await { Ok(vector) => { println!("EmbeddingClient 响应成功!向量维度: {}", vector.len()); assert!(!vector.is_empty(), "错误: 向量数据为空"); let preview_len = std::cmp::min(5, vector.len()); println!( "前 {} 个向量数值样例: {:?}", preview_len, &vector[..preview_len] ); } Err(e) => panic!("EmbeddingClient 接口调用失败: {}", e), } } println!("================= 大模型与向量模型真实网络集成测试完成 ================="); Ok(()) } }