From 29c6ff393510e523c7fb373b8ff36998abb56fdd Mon Sep 17 00:00:00 2001 From: SidneyZhang Date: Mon, 17 Aug 2026 15:33:00 +0800 Subject: [PATCH] =?UTF-8?q?feat(llm):=20=E5=BC=95=E5=85=A5=20rig-core=20?= =?UTF-8?q?=E5=B9=B6=E5=BB=BA=E7=AB=8B=E6=96=B0=20LLM=20=E9=97=A8=E9=9D=A2?= =?UTF-8?q?=EF=BC=88Ollama=20=E9=A6=96=E4=B8=AA=E6=89=93=E9=80=9A=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 依赖改用 rig-core 0.41(仅 reqwest+rustls,不引入 agent/derive) - 新门面:provider 枚举 + 统一配置 + 单一生成入口 + 凭据校验 + 错误映射 - Ollama 走 rig /api/chat,参数映射与错误文案与旧实现一致 - 用 rig 请求录制 mock 断言请求形状;8 个新单元测试 --- Cargo.toml | 7 +- src/llm/mod.rs | 1 + src/llm/rig/mod.rs | 378 +++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 385 insertions(+), 1 deletion(-) create mode 100644 src/llm/rig/mod.rs diff --git a/Cargo.toml b/Cargo.toml index fc193ad..d5ebbcf 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -34,7 +34,7 @@ which = "6.0" # HTTP client for LLM APIs reqwest = { version = "0.12", features = ["json", "rustls-tls", "stream"], default-features = false } -tokio = { version = "1.35", features = ["full"] } +tokio = { version = "1.35", features = ["full", "macros", "rt-multi-thread"] } # Error handling thiserror = "1.0" @@ -76,6 +76,8 @@ edit = "0.1" # Shell completion generation shell-words = "1.1" +# LLM integration (rig-core only: HTTP backend + rustls; no agent/derive) +rig-core = { version = "0.41", default-features = false, features = ["reqwest", "rustls"] } [dev-dependencies] assert_cmd = "2.0" @@ -83,6 +85,9 @@ predicates = "3.1" tempfile = "3.9" mockall = "0.12" wiremock = "0.6" +# rig mock HTTP backends for LLM layer tests +rig-core = { version = "0.41", default-features = false, features = ["test-utils"] } +http = "1" [profile.release] opt-level = "s" diff --git a/src/llm/mod.rs b/src/llm/mod.rs index 4a6d656..c917455 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -11,6 +11,7 @@ pub mod openai; pub mod openrouter; pub mod parsing; pub mod prompts; +pub mod rig; pub mod thinking; pub use anthropic::AnthropicClient; diff --git a/src/llm/rig/mod.rs b/src/llm/rig/mod.rs new file mode 100644 index 0000000..a46a61e --- /dev/null +++ b/src/llm/rig/mod.rs @@ -0,0 +1,378 @@ +//! 基于 rig-core 的新 LLM 门面。 +//! +//! 迁移扩张期:本模块与旧的手写实现并存,按 ticket 逐个接入 provider。 +//! 全部 provider 接入并切换后,旧实现将被删除。 + +use crate::config::manager::ConfigManager; +use anyhow::{Context, Result, bail}; +use rig_core::{ + client::{CompletionClient, Nothing, VerifyClient}, + completion::{AssistantContent, CompletionError, CompletionModel}, + http_client::{HttpClientExt, ReqwestClient}, + providers::ollama, + wasm_compat::{WasmCompatSend, WasmCompatSync}, +}; +use std::time::Duration; + +/// LLM 客户端运行时配置(与用户配置解耦的参数)。 +#[derive(Debug, Clone)] +pub struct LlmClientConfig { + pub max_tokens: u64, + pub temperature: f64, + pub timeout: Duration, +} + +impl Default for LlmClientConfig { + fn default() -> Self { + Self { + max_tokens: 500, + temperature: 0.7, + timeout: Duration::from_secs(30), + } + } +} + +/// Provider 后端枚举。 +/// +/// 泛型 `H` 是 HTTP 后端(生产环境为 reqwest 客户端),作为测试注入点: +/// 单测注入 rig 提供的 mock 后端以断言请求形状与响应处理。 +#[derive(Clone)] +pub enum Backend { + Ollama(ollama::Client), +} + +/// 基于 rig 的 LLM 客户端门面。 +#[derive(Clone)] +pub struct LlmClient { + backend: Backend, + model: String, + provider: String, + config: LlmClientConfig, + thinking_enabled: bool, +} + +impl LlmClient { + /// 从配置管理器构建(默认 reqwest 后端)。 + pub async fn from_config(manager: &ConfigManager) -> Result { + Self::from_config_with_think(manager, manager.config().llm.thinking_enabled).await + } + + /// 从配置构建,thinking 由参数显式指定。 + pub async fn from_config_with_think( + manager: &ConfigManager, + thinking_enabled: bool, + ) -> Result { + let config = manager.config(); + let cfg = LlmClientConfig { + max_tokens: config.llm.max_tokens as u64, + temperature: config.llm.temperature as f64, + timeout: Duration::from_secs(config.llm.timeout), + }; + let provider = manager.llm_provider().to_string(); + let model = manager.llm_model().to_string(); + let base_url = manager.llm_base_url(); + + let http = ReqwestClient::builder() + .timeout(cfg.timeout) + .build() + .context("Failed to create HTTP client")?; + + let backend = match provider.as_str() { + "ollama" => Backend::Ollama(build_ollama_client(&base_url, http)?), + // 其余 provider 在后续 ticket 接入,扩张期内新门面尚未被应用使用。 + "openai" | "anthropic" | "kimi" | "deepseek" | "openrouter" => { + bail!("Provider '{}' is not available in the new LLM backend yet", provider) + } + _ => bail!("Unknown LLM provider: {}", provider), + }; + + Ok(Self { + backend, + model, + provider, + config: cfg, + thinking_enabled, + }) + } +} + +impl LlmClient { + /// 内部与测试构造入口。 + pub(crate) fn new( + backend: Backend, + model: impl Into, + provider: impl Into, + config: LlmClientConfig, + thinking_enabled: bool, + ) -> Self { + Self { + backend, + model: model.into(), + provider: provider.into(), + config, + thinking_enabled, + } + } + + /// 当前是否处于 thinking 模式。 + pub(crate) fn thinking_enabled(&self) -> bool { + self.thinking_enabled + } +} + +impl LlmClient +where + H: HttpClientExt + + Clone + + Default + + std::fmt::Debug + + Send + + Sync + + WasmCompatSend + + WasmCompatSync + + 'static, +{ + /// 生成文本(thinking=false 的非流式路径)。 + pub async fn generate(&self, system: Option<&str>, user: &str) -> Result { + match &self.backend { + Backend::Ollama(client) => { + let model = client.completion_model(self.model.as_str()); + complete(&model, &self.provider, system, user, &self.config).await + } + } + } + + /// 检查 provider 是否可用(走 rig 凭据校验接口)。 + pub async fn is_available(&self) -> bool { + match &self.backend { + Backend::Ollama(client) => client.verify().await.is_ok(), + } + } +} + +/// 构建 Ollama 客户端(无 API key;支持自定义 base_url 与注入的 HTTP 后端)。 +fn build_ollama_client(base_url: &str, http: ReqwestClient) -> Result { + ollama::Client::builder() + .api_key(Nothing) + .base_url(base_url) + .http_client(http) + .build() + .map_err(|e| anyhow::anyhow!("Failed to build Ollama client: {}", e)) +} + +/// 非流式补全:系统提示、温度、max_tokens 映射到统一请求,聚合正文文本。 +async fn complete( + model: &M, + provider: &str, + system: Option<&str>, + user: &str, + config: &LlmClientConfig, +) -> 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 response = model + .completion(request) + .await + .map_err(|e| map_completion_error(provider, e))?; + + let text = collect_text(response.choice.iter()); + let text = text.trim().to_string(); + if text.is_empty() { + bail!("No response from {}", provider_display_name(provider)); + } + Ok(text) +} + +/// 从响应内容中聚合正文文本(忽略 reasoning/tool call 等块)。 +fn collect_text<'a>(items: impl Iterator) -> String { + items + .filter_map(|item| match item { + AssistantContent::Text(text) => Some(text.text.as_str()), + _ => None, + }) + .collect() +} + +/// provider 用户可见展示名(与旧实现文案一致,首字母大写)。 +pub(crate) fn provider_display_name(provider: &str) -> &str { + match provider { + "ollama" => "Ollama", + "openai" => "OpenAI", + "anthropic" => "Anthropic", + "kimi" => "Kimi", + "deepseek" => "DeepSeek", + "openrouter" => "OpenRouter", + other => other, + } +} + +/// 把 rig 的类型化补全错误映射为与旧实现一致的用户可见文案。 +pub(crate) fn map_completion_error(provider: &str, e: CompletionError) -> anyhow::Error { + // rig 在 provider 响应没有任何正文内容时的统一错误 → 映射为旧文案风格。 + if matches!(&e, CompletionError::ResponseError(msg) if msg == "No content provided") { + return anyhow::anyhow!("No response from {}", provider_display_name(provider)); + } + let name = provider_display_name(provider); + match (e.provider_response_status(), e.provider_response_body()) { + (Some(status), Some(body)) => anyhow::anyhow!("{} API error: {} - {}", name, status, body), + (Some(status), None) => anyhow::anyhow!("{} API error: {}", name, status), + (None, Some(body)) => anyhow::anyhow!("{} API error: {}", name, body), + (None, None) => anyhow::anyhow!("{} API request failed: {}", name, e), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use rig_core::test_utils::{MockHttpResponse, RecordingHttpClient}; + + const OLLAMA_OK: &str = r#"{ + "model": "llama3.2", + "created_at": "2024-01-01T00:00:00Z", + "message": {"role": "assistant", "content": "hello from ollama"}, + "done": true + }"#; + + /// 用注入的 HTTP 后端构建 Ollama 客户端(与生产构造路径一致)。 + fn build_backend(recorder: RecordingHttpClient) -> Backend { + Backend::Ollama( + ollama::Client::builder() + .api_key(Nothing) + .base_url("http://localhost:11434") + .http_client(recorder) + .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, + ); + (client, recorder) + } + + #[tokio::test] + async fn generate_maps_request_params_and_returns_text() { + let recorder = RecordingHttpClient::new(OLLAMA_OK); + let (client, recorder) = client_with(recorder); + + let text = client.generate(Some("be helpful"), "say hi").await.unwrap(); + assert_eq!(text, "hello from ollama"); + + let captured = recorder.requests(); + assert_eq!(captured.len(), 1); + assert_eq!(captured[0].uri, "http://localhost:11434/api/chat"); + + let body: serde_json::Value = serde_json::from_slice(&captured[0].body).unwrap(); + assert_eq!(body["model"], "llama3.2"); + assert_eq!(body["stream"], false); + assert_eq!(body["messages"][0]["role"], "system"); + assert_eq!(body["messages"][0]["content"], "be helpful"); + assert_eq!(body["messages"][1]["role"], "user"); + assert_eq!(body["messages"][1]["content"], "say hi"); + assert_eq!(body["options"]["temperature"], 0.5); + assert_eq!(body["options"]["num_predict"], 123); + } + + #[tokio::test] + async fn generate_without_system_omits_preamble() { + let recorder = RecordingHttpClient::new(OLLAMA_OK); + let (client, recorder) = client_with(recorder); + + client.generate(None, "say hi").await.unwrap(); + + let captured = recorder.requests(); + let body: serde_json::Value = serde_json::from_slice(&captured[0].body).unwrap(); + assert_eq!(body["messages"][0]["role"], "user"); + assert_eq!(body["messages"].as_array().unwrap().len(), 1); + } + + #[tokio::test] + async fn http_error_maps_to_provider_error_message() { + let recorder = RecordingHttpClient::new("ignored"); + recorder.set_response(MockHttpResponse::ErrorResponse( + http::StatusCode::BAD_GATEWAY, + "upstream exploded".into(), + )); + let (client, _) = client_with(recorder); + + let err = client.generate(None, "hi").await.unwrap_err(); + let msg = err.to_string(); + assert!(msg.contains("Ollama API error:"), "got: {msg}"); + assert!(msg.contains("502"), "got: {msg}"); + assert!(msg.contains("upstream exploded"), "got: {msg}"); + } + + #[tokio::test] + async fn empty_response_errors() { + let recorder = RecordingHttpClient::new( + r#"{"model":"llama3.2","created_at":"x","message":{"role":"assistant","content":""},"done":true}"#, + ); + let (client, _) = client_with(recorder); + + let err = client.generate(None, "hi").await.unwrap_err(); + assert_eq!(err.to_string(), "No response from Ollama"); + } + + #[tokio::test] + async fn is_available_uses_verify_endpoint() { + let recorder = RecordingHttpClient::new(""); + let (client, recorder) = client_with(recorder); + + assert!(client.is_available().await); + let captured = recorder.requests(); + assert_eq!(captured.len(), 1); + assert_eq!(captured[0].uri, "http://localhost:11434/api/tags"); + } + + #[tokio::test] + async fn is_available_false_on_authentication_error() { + let recorder = RecordingHttpClient::new(""); + recorder.set_response(MockHttpResponse::ErrorResponse( + http::StatusCode::UNAUTHORIZED, + "nope".into(), + )); + let (client, _) = client_with(recorder); + + assert!(!client.is_available().await); + } + + #[test] + fn map_error_prefers_status_and_body() { + let err = map_completion_error( + "Ollama", + 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("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")); + } +} \ No newline at end of file