diff --git a/src/llm/rig/mod.rs b/src/llm/rig/mod.rs index a46a61e..7ab64d5 100644 --- a/src/llm/rig/mod.rs +++ b/src/llm/rig/mod.rs @@ -4,14 +4,17 @@ //! 全部 provider 接入并切换后,旧实现将被删除。 use crate::config::manager::ConfigManager; +use crate::llm::thinking::{create_console_thinking_state, ThinkingStateManager}; use anyhow::{Context, Result, bail}; use rig_core::{ client::{CompletionClient, Nothing, VerifyClient}, completion::{AssistantContent, CompletionError, CompletionModel}, http_client::{HttpClientExt, ReqwestClient}, - providers::ollama, + providers::{deepseek, moonshot, ollama}, + streaming::StreamedAssistantContent, wasm_compat::{WasmCompatSend, WasmCompatSync}, }; +use std::sync::Arc; use std::time::Duration; /// LLM 客户端运行时配置(与用户配置解耦的参数)。 @@ -39,6 +42,8 @@ impl Default for LlmClientConfig { #[derive(Clone)] pub enum Backend { Ollama(ollama::Client), + DeepSeek(deepseek::Client), + Kimi(moonshot::Client), } /// 基于 rig 的 LLM 客户端门面。 @@ -49,6 +54,19 @@ pub struct LlmClient { provider: String, config: LlmClientConfig, thinking_enabled: bool, + thinking_state: Option>, +} + +/// 一次生成请求的 provider 相关参数。 +#[derive(Debug, Clone, PartialEq)] +struct RequestOptions { + /// `None` 表示省略该参数(如 thinking 模式下不传温度)。 + temperature: Option, + max_tokens: Option, + /// 透传到请求体的 provider 专属参数(thinking 开关等)。 + additional_params: Option, + /// 是否走流式路径。 + stream: bool, } impl LlmClient { @@ -79,23 +97,47 @@ impl LlmClient { let backend = match provider.as_str() { "ollama" => Backend::Ollama(build_ollama_client(&base_url, http)?), + "deepseek" => { + let key = required_api_key(&provider, manager.get_api_key())?; + Backend::DeepSeek(build_deepseek_client(&key, &base_url, http)?) + } + "kimi" => { + let key = required_api_key(&provider, manager.get_api_key())?; + Backend::Kimi(build_kimi_client(&key, &base_url, http)?) + } // 其余 provider 在后续 ticket 接入,扩张期内新门面尚未被应用使用。 - "openai" | "anthropic" | "kimi" | "deepseek" | "openrouter" => { + "openai" | "anthropic" | "openrouter" => { bail!("Provider '{}' is not available in the new LLM backend yet", provider) } _ => bail!("Unknown LLM provider: {}", provider), }; + let thinking_state = if thinking_enabled && supports_thinking(&provider) { + Some(create_console_thinking_state()) + } else { + None + }; + Ok(Self { backend, model, provider, config: cfg, thinking_enabled, + thinking_state, }) } } +/// 该 provider 是否支持 thinking 模式(与旧实现的白名单一致)。 +pub(crate) fn supports_thinking(provider: &str) -> bool { + matches!(provider, "deepseek" | "kimi" | "anthropic" | "openai") +} + +fn required_api_key(provider: &str, key: Option) -> Result { + key.ok_or_else(|| anyhow::anyhow!("{} API key not configured", provider_display_name(provider))) +} + impl LlmClient { /// 内部与测试构造入口。 pub(crate) fn new( @@ -104,6 +146,7 @@ impl LlmClient { provider: impl Into, config: LlmClientConfig, thinking_enabled: bool, + thinking_state: Option>, ) -> Self { Self { backend, @@ -111,12 +154,42 @@ impl LlmClient { provider: provider.into(), config, thinking_enabled, + thinking_state, } } - /// 当前是否处于 thinking 模式。 - pub(crate) fn thinking_enabled(&self) -> bool { - self.thinking_enabled + /// 按 provider × thinking 计算请求参数。 + fn request_options(&self) -> RequestOptions { + match &self.backend { + Backend::Ollama(_) => RequestOptions { + temperature: Some(self.config.temperature), + max_tokens: Some(self.config.max_tokens), + additional_params: None, + stream: false, + }, + // DeepSeek:thinking 时省略温度;thinking 参数始终显式传(与旧实现一致)。 + Backend::DeepSeek(_) => RequestOptions { + temperature: if self.thinking_enabled { + None + } else { + Some(self.config.temperature) + }, + max_tokens: Some(self.config.max_tokens), + additional_params: Some(serde_json::json!({ + "thinking": { "type": if self.thinking_enabled { "enabled" } else { "disabled" } } + })), + stream: self.thinking_enabled, + }, + // Kimi:Moonshot API 要求 thinking 恒 1.0、普通恒 0.6(忽略用户配置)。 + Backend::Kimi(_) => RequestOptions { + temperature: Some(if self.thinking_enabled { 1.0 } else { 0.6 }), + max_tokens: Some(self.config.max_tokens), + additional_params: Some(serde_json::json!({ + "thinking": { "type": if self.thinking_enabled { "enabled" } else { "disabled" } } + })), + stream: self.thinking_enabled, + }, + } } } @@ -132,12 +205,21 @@ where + WasmCompatSync + 'static, { - /// 生成文本(thinking=false 的非流式路径)。 + /// 生成文本;thinking 模式下走流式路径并驱动思考状态显示。 pub async fn generate(&self, system: Option<&str>, user: &str) -> Result { + let options = self.request_options(); match &self.backend { Backend::Ollama(client) => { let model = client.completion_model(self.model.as_str()); - complete(&model, &self.provider, system, user, &self.config).await + self.run(&model, system, user, &options).await + } + Backend::DeepSeek(client) => { + let model = client.completion_model(self.model.as_str()); + self.run(&model, system, user, &options).await + } + Backend::Kimi(client) => { + let model = client.completion_model(self.model.as_str()); + self.run(&model, system, user, &options).await } } } @@ -146,8 +228,92 @@ where pub async fn is_available(&self) -> bool { match &self.backend { Backend::Ollama(client) => client.verify().await.is_ok(), + Backend::DeepSeek(client) => client.verify().await.is_ok(), + Backend::Kimi(client) => client.verify().await.is_ok(), } } + + /// 按请求参数选择流式/非流式路径。 + async fn run( + &self, + model: &M, + system: Option<&str>, + user: &str, + options: &RequestOptions, + ) -> Result { + if options.stream { + self.stream_completion(model, system, user, options).await + } else { + complete(model, &self.provider, system, user, options).await + } + } + + /// 流式补全:聚合正文,驱动思考状态显示(与旧实现行为一致)。 + async fn stream_completion( + &self, + model: &M, + system: Option<&str>, + user: &str, + options: &RequestOptions, + ) -> Result { + let request = build_request(model, system, user, options); + let mut response = model + .stream(request) + .await + .map_err(|e| map_completion_error(&self.provider, e))?; + + use futures_util::StreamExt; + + let state = self.thinking_state.as_deref(); + let mut text = String::new(); + let mut has_reasoning = false; + let mut has_content = false; + + while let Some(event) = response.next().await { + let event = event.map_err(|e| map_completion_error(&self.provider, e))?; + match event { + StreamedAssistantContent::Text(chunk) => { + if has_reasoning && !has_content + && let Some(state) = state + { + state.end_thinking(); + } + has_content = true; + text.push_str(&chunk.text); + } + StreamedAssistantContent::Reasoning(_) + | StreamedAssistantContent::ReasoningDelta { .. } if !has_reasoning => { + has_reasoning = true; + if let Some(state) = state { + state.start_thinking(); + } + } + _ => {} + } + } + + if let Some(state) = state { + state.end_thinking(); + } + + let text = text.trim().to_string(); + if text.is_empty() { + if has_reasoning && !has_content { + bail!( + "{} returned reasoning content but no final answer. \ + The model may have entered an incomplete reasoning state. \ + Please try again or disable thinking mode.", + provider_display_name(&self.provider) + ); + } + bail!( + "No response from {}. \ + If thinking mode is enabled, try disabling it or ensure the model supports it.", + provider_display_name(&self.provider) + ); + } + Ok(text) + } } /// 构建 Ollama 客户端(无 API key;支持自定义 base_url 与注入的 HTTP 后端)。 @@ -160,26 +326,58 @@ fn build_ollama_client(base_url: &str, http: ReqwestClient) -> Result Result { + deepseek::Client::builder() + .api_key(key) + .base_url(base_url) + .http_client(http) + .build() + .map_err(|e| anyhow::anyhow!("Failed to build DeepSeek client: {}", e)) +} + +/// 构建 Kimi(Moonshot)客户端(OpenAI 兼容 API)。 +fn build_kimi_client(key: &str, base_url: &str, http: ReqwestClient) -> Result { + moonshot::Client::builder() + .api_key(key) + .base_url(base_url) + .http_client(http) + .build() + .map_err(|e| anyhow::anyhow!("Failed to build Kimi client: {}", e)) +} + +/// 组装统一的补全请求(preamble/temperature/max_tokens/additional_params)。 +fn build_request( + model: &M, + system: Option<&str>, + user: &str, + options: &RequestOptions, +) -> rig_core::completion::CompletionRequest { + let mut builder = model.completion_request(user); + if let Some(sys) = system { + builder = builder.preamble(sys.to_string()); + } + builder = builder + .temperature_opt(options.temperature) + .max_tokens_opt(options.max_tokens); + if let Some(params) = &options.additional_params { + builder = builder.additional_params(params.clone()); + } + builder.build() +} + +/// 非流式补全:聚合正文文本。 async fn complete( model: &M, provider: &str, system: Option<&str>, user: &str, - config: &LlmClientConfig, + options: &RequestOptions, ) -> Result where M: CompletionModel, { - let mut builder = model.completion_request(user); - if let Some(sys) = system { - builder = builder.preamble(sys.to_string()); - } - let request = builder - .temperature(config.temperature) - .max_tokens(config.max_tokens) - .build(); - + let request = build_request(model, system, user, options); let response = model .completion(request) .await @@ -234,7 +432,8 @@ pub(crate) fn map_completion_error(provider: &str, e: CompletionError) -> anyhow #[cfg(test)] mod tests { use super::*; - use rig_core::test_utils::{MockHttpResponse, RecordingHttpClient}; + use rig_core::test_utils::{MockHttpResponse, MockStreamingClient, RecordingHttpClient}; + use std::sync::atomic::{AtomicUsize, Ordering}; const OLLAMA_OK: &str = r#"{ "model": "llama3.2", @@ -243,40 +442,59 @@ mod tests { "done": true }"#; - /// 用注入的 HTTP 后端构建 Ollama 客户端(与生产构造路径一致)。 - fn build_backend(recorder: RecordingHttpClient) -> Backend { - Backend::Ollama( + const DEEPSEEK_OK: &str = r#"{ + "id": "cmpl-1", + "model": "deepseek-v4-flash", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hello from deepseek"}, "finish_reason": "stop"}], + "usage": {"completion_tokens": 3, "prompt_tokens": 4, "prompt_cache_hit_tokens": 0, "prompt_cache_miss_tokens": 0, "total_tokens": 7} + }"#; + + fn test_config() -> LlmClientConfig { + LlmClientConfig { + max_tokens: 123, + temperature: 0.5, + timeout: Duration::from_secs(30), + } + } + + fn ollama_client( + recorder: RecordingHttpClient, + ) -> (LlmClient, RecordingHttpClient) { + let backend = Backend::Ollama( ollama::Client::builder() .api_key(Nothing) .base_url("http://localhost:11434") - .http_client(recorder) + .http_client(recorder.clone()) .build() .expect("build ollama client with mock backend"), - ) - } - - fn client_with( - recorder: RecordingHttpClient, - ) -> (LlmClient, RecordingHttpClient) { - let backend = build_backend(recorder.clone()); - let client = LlmClient::new( - backend, - "llama3.2", - "ollama", - LlmClientConfig { - max_tokens: 123, - temperature: 0.5, - timeout: Duration::from_secs(30), - }, - false, ); + let client = LlmClient::new(backend, "llama3.2", "ollama", test_config(), false, None); (client, recorder) } + fn deepseek_client( + recorder: RecordingHttpClient, + thinking: bool, + ) -> (LlmClient, RecordingHttpClient) { + let backend = Backend::DeepSeek( + deepseek::Client::builder() + .api_key("sk-test") + .base_url("https://api.deepseek.com/v1") + .http_client(recorder.clone()) + .build() + .expect("build deepseek client with mock backend"), + ); + let client = + LlmClient::new(backend, "deepseek-v4-flash", "deepseek", test_config(), thinking, None); + (client, recorder) + } + + // ---- Ollama 基本路径 ---- + #[tokio::test] async fn generate_maps_request_params_and_returns_text() { let recorder = RecordingHttpClient::new(OLLAMA_OK); - let (client, recorder) = client_with(recorder); + let (client, recorder) = ollama_client(recorder); let text = client.generate(Some("be helpful"), "say hi").await.unwrap(); assert_eq!(text, "hello from ollama"); @@ -299,7 +517,7 @@ mod tests { #[tokio::test] async fn generate_without_system_omits_preamble() { let recorder = RecordingHttpClient::new(OLLAMA_OK); - let (client, recorder) = client_with(recorder); + let (client, recorder) = ollama_client(recorder); client.generate(None, "say hi").await.unwrap(); @@ -316,7 +534,7 @@ mod tests { http::StatusCode::BAD_GATEWAY, "upstream exploded".into(), )); - let (client, _) = client_with(recorder); + let (client, _) = ollama_client(recorder); let err = client.generate(None, "hi").await.unwrap_err(); let msg = err.to_string(); @@ -330,7 +548,7 @@ mod tests { let recorder = RecordingHttpClient::new( r#"{"model":"llama3.2","created_at":"x","message":{"role":"assistant","content":""},"done":true}"#, ); - let (client, _) = client_with(recorder); + let (client, _) = ollama_client(recorder); let err = client.generate(None, "hi").await.unwrap_err(); assert_eq!(err.to_string(), "No response from Ollama"); @@ -339,7 +557,7 @@ mod tests { #[tokio::test] async fn is_available_uses_verify_endpoint() { let recorder = RecordingHttpClient::new(""); - let (client, recorder) = client_with(recorder); + let (client, recorder) = ollama_client(recorder); assert!(client.is_available().await); let captured = recorder.requests(); @@ -354,25 +572,189 @@ mod tests { http::StatusCode::UNAUTHORIZED, "nope".into(), )); - let (client, _) = client_with(recorder); + let (client, _) = ollama_client(recorder); assert!(!client.is_available().await); } + // ---- DeepSeek/Kimi 请求参数形状 ---- + + #[tokio::test] + async fn deepseek_normal_request_carries_thinking_disabled() { + let recorder = RecordingHttpClient::new(DEEPSEEK_OK); + let (client, recorder) = deepseek_client(recorder, false); + + let text = client.generate(Some("sys"), "hi").await.unwrap(); + assert_eq!(text, "hello from deepseek"); + + let captured = recorder.requests(); + assert_eq!(captured[0].uri, "https://api.deepseek.com/v1/chat/completions"); + let body: serde_json::Value = serde_json::from_slice(&captured[0].body).unwrap(); + assert_eq!(body["model"], "deepseek-v4-flash"); + // 非流式请求省略 stream 字段(API 默认 false) + assert_ne!(body["stream"], true); + assert_eq!(body["temperature"], 0.5); + assert_eq!(body["thinking"]["type"], "disabled"); + assert_eq!(body["max_tokens"], 123); + } + + #[tokio::test] + async fn kimi_normal_request_uses_kimi_temperature() { + let recorder = RecordingHttpClient::new(DEEPSEEK_OK); + let backend = Backend::Kimi( + moonshot::Client::builder() + .api_key("sk-test") + .base_url("https://api.moonshot.cn/v1") + .http_client(recorder.clone()) + .build() + .expect("build kimi client with mock backend"), + ); + let client = LlmClient::new(backend, "kimi-k2.6", "kimi", test_config(), false, None); + + client.generate(None, "hi").await.unwrap(); + + let captured = recorder.requests(); + assert_eq!(captured[0].uri, "https://api.moonshot.cn/v1/chat/completions"); + let body: serde_json::Value = serde_json::from_slice(&captured[0].body).unwrap(); + assert_eq!(body["temperature"], 0.6); + assert_eq!(body["thinking"]["type"], "disabled"); + } + + // ---- 流式与 thinking 事件 ---- + + /// 构造 thinking 模式下的 DeepSeek 客户端(流式 mock 后端 + 记录回调的思考状态)。 + fn deepseek_streaming_client( + sse: &str, + ) -> ( + LlmClient, + Arc, + Arc, + Arc, + ) { + let start_count = Arc::new(AtomicUsize::new(0)); + let end_count = Arc::new(AtomicUsize::new(0)); + let state = Arc::new( + ThinkingStateManager::new() + .on_thinking_start({ + let c = start_count.clone(); + move || { + c.fetch_add(1, Ordering::SeqCst); + } + }) + .on_thinking_end({ + let c = end_count.clone(); + move || { + c.fetch_add(1, Ordering::SeqCst); + } + }), + ); + + let backend = Backend::DeepSeek( + deepseek::Client::builder() + .api_key("sk-test") + .base_url("https://api.deepseek.com/v1") + .http_client(MockStreamingClient { sse_bytes: sse.to_string().into() }) + .build() + .expect("build deepseek client with streaming mock"), + ); + let client = LlmClient::new( + backend, + "deepseek-v4-flash", + "deepseek", + test_config(), + true, + Some(state.clone()), + ); + (client, state, start_count, end_count) + } + + const SSE_THINK_THEN_TEXT: &str = concat!( + "data: {\"choices\":[{\"index\":0,\"delta\":{\"reasoning_content\":\"let me think\"},\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"hello\"},\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n", + "data: [DONE]\n\n", + ); + + const SSE_REASONING_ONLY: &str = concat!( + "data: {\"choices\":[{\"index\":0,\"delta\":{\"reasoning_content\":\"only thoughts\"},\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n", + "data: [DONE]\n\n", + ); + + #[tokio::test] + async fn streaming_aggregates_text_and_drives_thinking_state() { + let (client, _, start_count, end_count) = deepseek_streaming_client(SSE_THINK_THEN_TEXT); + + let text = client.generate(None, "hi").await.unwrap(); + assert_eq!(text, "hello"); + assert_eq!(start_count.load(Ordering::SeqCst), 1); + assert_eq!(end_count.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn streaming_reasoning_only_errors_with_hint() { + let (client, _, _, _) = deepseek_streaming_client(SSE_REASONING_ONLY); + + let err = client.generate(None, "hi").await.unwrap_err(); + let msg = err.to_string(); + assert!( + msg.contains("DeepSeek returned reasoning content but no final answer"), + "got: {msg}" + ); + assert!(msg.contains("disable thinking mode"), "got: {msg}"); + } + + #[tokio::test] + async fn streaming_without_reasoning_skips_state_start() { + let (client, _, start_count, _) = deepseek_streaming_client(concat!( + "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"plain\"},\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n", + "data: [DONE]\n\n", + )); + + let text = client.generate(None, "hi").await.unwrap(); + assert_eq!(text, "plain"); + assert_eq!(start_count.load(Ordering::SeqCst), 0); + } + + // ---- 错误映射 ---- + #[test] fn map_error_prefers_status_and_body() { let err = map_completion_error( - "Ollama", + "DeepSeek", CompletionError::from_http_response(http::StatusCode::TOO_MANY_REQUESTS, "slow down"), ); let msg = err.to_string(); - assert!(msg.contains("Ollama API error: 429"), "got: {msg}"); + assert!(msg.contains("DeepSeek API error: 429"), "got: {msg}"); assert!(msg.contains("slow down"), "got: {msg}"); } #[test] fn map_error_without_body_keeps_provider_prefix() { - let err = map_completion_error("Ollama", CompletionError::ProviderError("boom".into())); - assert!(err.to_string().contains("Ollama API request failed: ProviderError: boom")); + let err = map_completion_error("DeepSeek", CompletionError::ProviderError("boom".into())); + assert!(err + .to_string() + .contains("DeepSeek API request failed: ProviderError: boom")); + } + + #[test] + fn provider_display_names_match_legacy() { + assert_eq!(provider_display_name("ollama"), "Ollama"); + assert_eq!(provider_display_name("openai"), "OpenAI"); + assert_eq!(provider_display_name("anthropic"), "Anthropic"); + assert_eq!(provider_display_name("kimi"), "Kimi"); + assert_eq!(provider_display_name("deepseek"), "DeepSeek"); + assert_eq!(provider_display_name("openrouter"), "OpenRouter"); + } + + #[test] + fn supports_thinking_whitelist_matches_legacy() { + for provider in ["deepseek", "kimi", "anthropic", "openai"] { + assert!(supports_thinking(provider), "{provider} should support thinking"); + } + for provider in ["ollama", "openrouter"] { + assert!(!supports_thinking(provider), "{provider} should not support thinking"); + } } } \ No newline at end of file