From 3c2b96a4d11d4ce142f5aebc4930056fc1914c51 Mon Sep 17 00:00:00 2001 From: SidneyZhang Date: Mon, 17 Aug 2026 16:04:26 +0800 Subject: [PATCH] =?UTF-8?q?refactor(llm):=20=E5=88=A0=E9=99=A4=E6=89=8B?= =?UTF-8?q?=E5=86=99=20provider=20=E5=AE=9E=E7=8E=B0=E4=B8=8E=E5=86=97?= =?UTF-8?q?=E4=BD=99=E4=BE=9D=E8=B5=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 删除 6 个手写 provider 文件、LlmProvider trait、动态分发、HTTP 客户端工厂与门面死代码 - llm 模块收敛为 rig 门面 + 提示词/解析/思考状态四个模块 - 移除 reqwest(0.12) 与 async-trait 直接依赖(futures-util 由流式消费保留) - 依赖树确认无 rig-agent/fastembed/lancedb/milvus 等组件 --- Cargo.toml | 3 - src/llm/anthropic.rs | 655 ----------------------------------------- src/llm/deepseek.rs | 622 --------------------------------------- src/llm/kimi.rs | 587 ------------------------------------- src/llm/mod.rs | 341 +--------------------- src/llm/ollama.rs | 229 --------------- src/llm/openai.rs | 659 ------------------------------------------ src/llm/openrouter.rs | 286 ------------------ 8 files changed, 14 insertions(+), 3368 deletions(-) delete mode 100644 src/llm/anthropic.rs delete mode 100644 src/llm/deepseek.rs delete mode 100644 src/llm/kimi.rs delete mode 100644 src/llm/ollama.rs delete mode 100644 src/llm/openai.rs delete mode 100644 src/llm/openrouter.rs diff --git a/Cargo.toml b/Cargo.toml index d5ebbcf..9461166 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -32,8 +32,6 @@ dirs = "5.0" git2 = "0.20.3" 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", "macros", "rt-multi-thread"] } # Error handling @@ -57,7 +55,6 @@ tempfile = "3.9" sha2 = "0.10" hex = "0.4" textwrap = "0.16" -async-trait = "0.1" futures-util = "0.3" serde_json = "1.0" atty = "0.2" diff --git a/src/llm/anthropic.rs b/src/llm/anthropic.rs deleted file mode 100644 index 0779a5e..0000000 --- a/src/llm/anthropic.rs +++ /dev/null @@ -1,655 +0,0 @@ -use super::thinking::ThinkingStateManager; -use super::{LlmProvider, create_http_client}; -use anyhow::{Context, Result, bail}; -use async_trait::async_trait; -use serde::{Deserialize, Serialize}; -use std::sync::Arc; -use std::time::Duration; - -/// Anthropic Claude API client -pub struct AnthropicClient { - api_key: String, - model: String, - client: reqwest::Client, - thinking_enabled: bool, - thinking_budget_tokens: u32, - max_tokens: u32, - temperature: f32, - top_p: Option, - thinking_state: Option>, -} - -#[derive(Debug, Serialize)] -struct MessagesRequest { - model: String, - max_tokens: u32, - #[serde(skip_serializing_if = "Option::is_none")] - temperature: Option, - #[serde(skip_serializing_if = "Option::is_none")] - top_p: Option, - messages: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - system: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - thinking: Option, - stream: bool, -} - -#[derive(Debug, Serialize, Clone)] -struct SystemContent { - #[serde(rename = "type")] - content_type: String, - text: String, -} - -#[derive(Debug, Serialize)] -struct ThinkingConfig { - #[serde(rename = "type")] - thinking_type: String, - #[serde(skip_serializing_if = "Option::is_none")] - budget_tokens: Option, -} - -#[derive(Debug, Serialize, Deserialize, Clone)] -struct AnthropicMessage { - role: String, - content: AnthropicContent, -} - -#[derive(Debug, Serialize, Deserialize, Clone)] -#[serde(untagged)] -enum AnthropicContent { - Text(String), - Blocks(Vec), -} - -#[derive(Debug, Serialize, Deserialize, Clone)] -struct ContentBlock { - #[serde(rename = "type")] - content_type: String, - #[serde(skip_serializing_if = "Option::is_none")] - text: Option, -} - -#[derive(Debug, Deserialize)] -struct MessagesResponse { - content: Vec, -} - -#[derive(Debug, Deserialize)] -struct ResponseContentBlock { - #[serde(rename = "type")] - content_type: String, - text: String, -} - -#[derive(Debug, Deserialize)] -struct ErrorResponse { - error: AnthropicError, -} - -#[derive(Debug, Deserialize)] -struct AnthropicError { - #[serde(rename = "type")] - error_type: String, - message: String, -} - -// --- Streaming SSE event structures --- - -#[derive(Debug, Deserialize)] -struct SseEvent { - #[serde(rename = "type")] - event_type: String, - #[serde(default)] - message: Option, - #[serde(default)] - index: Option, - #[serde(default)] - content_block: Option, - #[serde(default)] - delta: Option, - #[serde(default)] - usage: Option, -} - -#[derive(Debug, Deserialize)] -struct SseMessage { - #[serde(default)] - content: Option>, -} - -#[derive(Debug, Deserialize)] -struct SseContentBlock { - #[serde(rename = "type")] - content_type: String, - #[serde(default)] - thinking: Option, - #[serde(default)] - text: Option, -} - -#[derive(Debug, Deserialize)] -struct SseDelta { - #[serde(rename = "type")] - delta_type: Option, - #[serde(default)] - thinking: Option, - #[serde(default)] - text: Option, -} - -#[derive(Debug, Deserialize)] -struct SseUsage { - #[serde(default)] - output_tokens: Option, -} - -impl AnthropicClient { - pub fn new(api_key: &str, model: &str) -> Result { - let client = create_http_client(Duration::from_secs(60))?; - - Ok(Self { - api_key: api_key.to_string(), - model: model.to_string(), - client, - thinking_enabled: false, - thinking_budget_tokens: 1024, - max_tokens: 500, - temperature: 0.7, - top_p: None, - thinking_state: None, - }) - } - - pub fn with_timeout(mut self, timeout: Duration) -> Result { - self.client = create_http_client(timeout)?; - Ok(self) - } - - pub fn with_thinking(mut self, enabled: bool) -> Self { - self.thinking_enabled = enabled; - self - } - - pub fn with_thinking_budget_tokens(mut self, budget_tokens: u32) -> Self { - self.thinking_budget_tokens = budget_tokens; - self - } - - pub fn with_max_tokens(mut self, max_tokens: u32) -> Self { - self.max_tokens = max_tokens; - self - } - - pub fn with_temperature(mut self, temperature: f32) -> Self { - self.temperature = temperature; - self - } - - pub fn with_top_p(mut self, top_p: f32) -> Self { - self.top_p = Some(top_p); - self - } - - pub fn with_thinking_state(mut self, state: Arc) -> Self { - self.thinking_state = Some(state); - self - } - - pub async fn list_models(&self) -> Result> { - Ok(ANTHROPIC_MODELS.iter().map(|&m| m.to_string()).collect()) - } - - pub async fn validate_key(&self) -> Result { - let url = "https://api.anthropic.com/v1/messages"; - - let request = MessagesRequest { - model: self.model.clone(), - max_tokens: 5, - temperature: Some(0.0), - top_p: None, - messages: vec![AnthropicMessage { - role: "user".to_string(), - content: AnthropicContent::Text("Hi".to_string()), - }], - system: None, - thinking: None, - stream: false, - }; - - let response = self - .client - .post(url) - .header("x-api-key", &self.api_key) - .header("anthropic-version", "2023-06-01") - .header("Content-Type", "application/json") - .json(&request) - .send() - .await; - - match response { - Ok(resp) => { - if resp.status().is_success() { - Ok(true) - } else { - let status = resp.status(); - if status.as_u16() == 401 { - Ok(false) - } else { - let text = resp.text().await.unwrap_or_default(); - bail!("Anthropic API error: {} - {}", status, text) - } - } - } - Err(e) => Err(e.into()), - } - } -} - -#[async_trait] -impl LlmProvider for AnthropicClient { - async fn generate(&self, prompt: &str) -> Result { - let messages = vec![AnthropicMessage { - role: "user".to_string(), - content: AnthropicContent::Text(prompt.to_string()), - }]; - - self.messages_request_with_retry(messages, None).await - } - - async fn generate_with_system(&self, system: &str, user: &str) -> Result { - let messages = vec![AnthropicMessage { - role: "user".to_string(), - content: AnthropicContent::Text(user.to_string()), - }]; - - let system = if system.is_empty() { - None - } else { - Some(vec![SystemContent { - content_type: "text".to_string(), - text: system.to_string(), - }]) - }; - - self.messages_request_with_retry(messages, system).await - } - - async fn is_available(&self) -> bool { - self.validate_key().await.unwrap_or(false) - } - - fn name(&self) -> &str { - "anthropic" - } -} - -impl AnthropicClient { - async fn messages_request_with_retry( - &self, - messages: Vec, - system: Option>, - ) -> Result { - let mut last_error = None; - - for attempt in 1..=3 { - match self - .messages_request(messages.clone(), system.clone()) - .await - { - Ok(result) => return Ok(result), - Err(e) => { - let err_msg = e.to_string(); - let is_retryable = err_msg.contains("timeout") - || err_msg.contains("connection") - || err_msg.contains("temporary") - || err_msg.contains("5") - && (err_msg.contains("500") - || err_msg.contains("502") - || err_msg.contains("503") - || err_msg.contains("504")); - - if !is_retryable || attempt == 3 { - last_error = Some(e); - break; - } - - tokio::time::sleep(Duration::from_millis(500 * 2u64.pow(attempt - 1))).await; - } - } - } - - Err(last_error.unwrap_or_else(|| anyhow::anyhow!("Request failed after retries"))) - } - - async fn messages_request( - &self, - messages: Vec, - system: Option>, - ) -> Result { - if self.thinking_enabled { - self.streaming_messages_request(messages, system).await - } else { - self.non_streaming_messages_request(messages, system).await - } - } - - async fn non_streaming_messages_request( - &self, - messages: Vec, - system: Option>, - ) -> Result { - let url = "https://api.anthropic.com/v1/messages"; - - let temperature = if self.temperature == 0.0 { - None - } else { - Some(self.temperature) - }; - - let request = MessagesRequest { - model: self.model.clone(), - max_tokens: self.max_tokens, - temperature, - top_p: self.top_p, - messages, - system, - thinking: Some(ThinkingConfig { - thinking_type: "disabled".to_string(), - budget_tokens: None, - }), - stream: false, - }; - - let response = self - .client - .post(url) - .header("x-api-key", &self.api_key) - .header("anthropic-version", "2023-06-01") - .header("Content-Type", "application/json") - .json(&request) - .send() - .await - .context("Failed to send request to Anthropic")?; - - let status = response.status(); - - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - - if let Ok(error) = serde_json::from_str::(&text) { - bail!( - "Anthropic API error: {} ({})", - error.error.message, - error.error.error_type - ); - } - - bail!("Anthropic API error: {} - {}", status, text); - } - - let result: MessagesResponse = response - .json() - .await - .context("Failed to parse Anthropic response")?; - - result - .content - .into_iter() - .find(|c| c.content_type == "text") - .map(|c| c.text.trim().to_string()) - .filter(|s| !s.is_empty()) - .ok_or_else(|| anyhow::anyhow!("No text response from Anthropic")) - } - - /// Streaming request for thinking mode, filters thinking content blocks - async fn streaming_messages_request( - &self, - messages: Vec, - system: Option>, - ) -> Result { - let url = "https://api.anthropic.com/v1/messages"; - - let thinking = ThinkingConfig { - thinking_type: "enabled".to_string(), - budget_tokens: Some(self.thinking_budget_tokens), - }; - - // max_tokens must exceed budget_tokens - let max_tokens = (self.max_tokens).max(self.thinking_budget_tokens + 100); - - let request = MessagesRequest { - model: self.model.clone(), - max_tokens, - temperature: None, // must be omitted for thinking mode - top_p: None, - messages, - system, - thinking: Some(thinking), - stream: true, - }; - - let response = self - .client - .post(url) - .header("x-api-key", &self.api_key) - .header("anthropic-version", "2023-06-01") - .header("Content-Type", "application/json") - .header("Accept", "text/event-stream") - .json(&request) - .send() - .await - .context("Failed to send streaming request to Anthropic")?; - - let status = response.status(); - - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - - if let Ok(error) = serde_json::from_str::(&text) { - bail!( - "Anthropic API error: {} ({})", - error.error.message, - error.error.error_type - ); - } - - bail!("Anthropic API error: {} - {}", status, text); - } - - let mut content_buffer = String::new(); - let mut in_thinking = false; - let mut has_reasoning = false; - let mut has_content = false; - - let thinking_state = self.thinking_state.as_ref(); - - let mut byte_stream = response.bytes_stream(); - let mut line_buffer = String::new(); - - use futures_util::StreamExt; - - while let Some(chunk) = byte_stream.next().await { - let chunk = chunk.context("Failed to read streaming response chunk")?; - let chunk_str = - String::from_utf8(chunk.to_vec()).context("Invalid UTF-8 in stream chunk")?; - - line_buffer.push_str(&chunk_str); - - while let Some(line_end) = line_buffer.find('\n') { - let line = line_buffer[..line_end].trim().to_string(); - line_buffer = line_buffer[line_end + 1..].to_string(); - - if line.is_empty() { - continue; - } - - // Parse SSE event line - if let Some(data) = line.strip_prefix("data: ") { - if let Ok(event) = serde_json::from_str::(data) { - match event.event_type.as_str() { - "content_block_start" => { - if let Some(ref block) = event.content_block { - if block.content_type == "thinking" { - in_thinking = true; - if !has_reasoning { - has_reasoning = true; - if let Some(state) = thinking_state { - state.start_thinking(); - } - } - } - } - } - "content_block_delta" => { - if let Some(ref delta) = event.delta { - // Thinking delta - ignore content but track state - if delta.thinking.is_some() { - continue; - } - - // Text delta - collect - if in_thinking && delta.text.is_some() { - // Transition from thinking to text - if let Some(state) = thinking_state { - state.end_thinking(); - } - in_thinking = false; - } - if let Some(ref text) = delta.text - && !text.is_empty() - { - has_content = true; - content_buffer.push_str(text); - } - } - } - "content_block_stop" => { - if in_thinking { - if let Some(state) = thinking_state { - state.end_thinking(); - } - in_thinking = false; - } - } - _ => {} - } - } - } - } - } - - // Ensure thinking state is ended - if let Some(state) = thinking_state { - state.end_thinking(); - } - - let result = content_buffer.trim().to_string(); - - if result.is_empty() { - if has_reasoning && !has_content { - bail!( - "Anthropic returned thinking content but no final answer. \ - The model may have entered an incomplete thinking state. \ - Please try again or disable thinking mode." - ); - } - bail!( - "No response from Anthropic. \ - If thinking mode is enabled, try disabling it or ensure the model supports it." - ); - } - - Ok(result) - } -} - -/// Available Anthropic models (Claude 4 series with extended thinking) -pub const ANTHROPIC_MODELS: &[&str] = &[ - "claude-opus-4-7", - "claude-sonnet-4-6", - "claude-haiku-4-5", - // Legacy models - "claude-3-opus-20240229", - "claude-3-sonnet-20240229", - "claude-3-haiku-20240307", - "claude-2.1", - "claude-2.0", - "claude-instant-1.2", -]; - -pub fn is_valid_model(model: &str) -> bool { - ANTHROPIC_MODELS.contains(&model) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_model_validation_claude4() { - assert!(is_valid_model("claude-opus-4-7")); - assert!(is_valid_model("claude-sonnet-4-6")); - assert!(is_valid_model("claude-haiku-4-5")); - assert!(is_valid_model("claude-3-sonnet-20240229")); - assert!(!is_valid_model("invalid-model")); - } - - #[test] - fn test_thinking_config_serialization() { - let config = ThinkingConfig { - thinking_type: "enabled".to_string(), - budget_tokens: Some(2048), - }; - let json = serde_json::to_string(&config).unwrap(); - assert!(json.contains(r#""type":"enabled""#)); - assert!(json.contains(r#""budget_tokens":2048"#)); - } - - #[test] - fn test_thinking_config_disabled_serialization() { - let config = ThinkingConfig { - thinking_type: "disabled".to_string(), - budget_tokens: None, - }; - let json = serde_json::to_string(&config).unwrap(); - assert_eq!(json, r#"{"type":"disabled"}"#); - } - - #[test] - fn test_system_content_serialization() { - let content = SystemContent { - content_type: "text".to_string(), - text: "You are helpful.".to_string(), - }; - let json = serde_json::to_string(&content).unwrap(); - assert!(json.contains(r#""type":"text""#)); - } - - #[test] - fn test_sse_event_parsing_content_block_start() { - let json = r#"{"type":"content_block_start","index":0,"content_block":{"type":"thinking","thinking":""}}"#; - let event: SseEvent = serde_json::from_str(json).unwrap(); - assert_eq!(event.event_type, "content_block_start"); - assert_eq!(event.content_block.unwrap().content_type, "thinking"); - } - - #[test] - fn test_sse_event_parsing_text_delta() { - let json = r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}"#; - let event: SseEvent = serde_json::from_str(json).unwrap(); - assert_eq!(event.event_type, "content_block_delta"); - assert_eq!(event.delta.unwrap().text, Some("Hello".to_string())); - } - - #[test] - fn test_anthropic_content_text() { - let msg = AnthropicMessage { - role: "user".to_string(), - content: AnthropicContent::Text("Hello".to_string()), - }; - let json = serde_json::to_string(&msg).unwrap(); - assert!(json.contains(r#""content":"Hello""#)); - } -} diff --git a/src/llm/deepseek.rs b/src/llm/deepseek.rs deleted file mode 100644 index f287917..0000000 --- a/src/llm/deepseek.rs +++ /dev/null @@ -1,622 +0,0 @@ -use super::thinking::ThinkingStateManager; -use super::{LlmProvider, create_http_client}; -use anyhow::{Context, Result, bail}; -use async_trait::async_trait; -use serde::{Deserialize, Serialize}; -use std::sync::Arc; -use std::time::Duration; - -/// DeepSeek API client -pub struct DeepSeekClient { - base_url: String, - api_key: String, - model: String, - client: reqwest::Client, - thinking_enabled: bool, - reasoning_effort: Option, - max_tokens: u32, - temperature: f32, - thinking_state: Option>, -} - -#[derive(Debug, Serialize)] -struct ChatCompletionRequest { - model: String, - messages: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - max_tokens: Option, - #[serde(skip_serializing_if = "Option::is_none")] - temperature: Option, - #[serde(skip_serializing_if = "Option::is_none")] - top_p: Option, - #[serde(skip_serializing_if = "Option::is_none")] - presence_penalty: Option, - #[serde(skip_serializing_if = "Option::is_none")] - frequency_penalty: Option, - stream: bool, - #[serde(skip_serializing_if = "Option::is_none")] - thinking: Option, - #[serde(skip_serializing_if = "Option::is_none")] - reasoning_effort: Option, -} - -#[derive(Debug, Serialize)] -struct ThinkingConfig { - #[serde(rename = "type")] - thinking_type: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -struct Message { - role: String, - content: String, - #[serde(skip_serializing_if = "Option::is_none")] - reasoning_content: Option, -} - -#[derive(Debug, Deserialize)] -struct ChatCompletionResponse { - choices: Vec, -} - -#[derive(Debug, Deserialize)] -struct Choice { - message: Message, - #[serde(default)] - reasoning_content: Option, -} - -// --- Streaming response structures --- - -#[derive(Debug, Deserialize)] -struct StreamChunk { - choices: Vec, -} - -#[derive(Debug, Deserialize)] -struct StreamChoice { - delta: StreamDelta, - #[serde(default)] - finish_reason: Option, - index: Option, -} - -#[derive(Debug, Deserialize, Default)] -struct StreamDelta { - #[serde(default)] - content: Option, - #[serde(default)] - reasoning_content: Option, -} - -#[derive(Debug, Deserialize)] -struct ErrorResponse { - error: ApiError, -} - -#[derive(Debug, Deserialize)] -struct ApiError { - message: String, - #[serde(rename = "type")] - error_type: String, -} - -impl DeepSeekClient { - pub fn new(api_key: &str, model: &str) -> Result { - let client = create_http_client(Duration::from_secs(300))?; - - Ok(Self { - base_url: "https://api.deepseek.com".to_string(), - api_key: api_key.to_string(), - model: model.to_string(), - client, - thinking_enabled: false, - reasoning_effort: None, - max_tokens: 500, - temperature: 0.7, - thinking_state: None, - }) - } - - pub fn with_base_url(api_key: &str, model: &str, base_url: &str) -> Result { - let client = create_http_client(Duration::from_secs(300))?; - - Ok(Self { - base_url: base_url.trim_end_matches('/').to_string(), - api_key: api_key.to_string(), - model: model.to_string(), - client, - thinking_enabled: false, - reasoning_effort: None, - max_tokens: 500, - temperature: 0.7, - thinking_state: None, - }) - } - - pub fn with_timeout(mut self, timeout: Duration) -> Result { - self.client = create_http_client(timeout)?; - Ok(self) - } - - pub fn with_thinking(mut self, enabled: bool) -> Self { - self.thinking_enabled = enabled; - self - } - - pub fn with_reasoning_effort(mut self, effort: Option) -> Self { - self.reasoning_effort = effort; - self - } - - pub fn with_max_tokens(mut self, max_tokens: u32) -> Self { - self.max_tokens = max_tokens; - self - } - - pub fn with_temperature(mut self, temperature: f32) -> Self { - self.temperature = temperature; - self - } - - pub fn with_thinking_state(mut self, state: Arc) -> Self { - self.thinking_state = Some(state); - self - } - - pub async fn list_models(&self) -> Result> { - let url = format!("{}/models", self.base_url); - - let response = self - .client - .get(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .send() - .await - .context("Failed to list DeepSeek models")?; - - if !response.status().is_success() { - let status = response.status(); - let text = response.text().await.unwrap_or_default(); - bail!("DeepSeek API error: {} - {}", status, text); - } - - #[derive(Deserialize)] - struct ModelsResponse { - data: Vec, - } - - #[derive(Deserialize)] - struct ModelId { - id: String, - } - - let result: ModelsResponse = response - .json() - .await - .context("Failed to parse DeepSeek response")?; - - Ok(result.data.into_iter().map(|m| m.id).collect()) - } - - pub async fn validate_key(&self) -> Result { - match self.list_models().await { - Ok(_) => Ok(true), - Err(e) => { - let err_str = e.to_string(); - if err_str.contains("401") || err_str.contains("Unauthorized") { - Ok(false) - } else { - Err(e) - } - } - } - } -} - -#[async_trait] -impl LlmProvider for DeepSeekClient { - async fn generate(&self, prompt: &str) -> Result { - let messages = vec![Message { - role: "user".to_string(), - content: prompt.to_string(), - reasoning_content: None, - }]; - - self.chat_completion_with_retry(messages).await - } - - async fn generate_with_system(&self, system: &str, user: &str) -> Result { - let mut messages = vec![]; - - if !system.is_empty() { - messages.push(Message { - role: "system".to_string(), - content: system.to_string(), - reasoning_content: None, - }); - } - - messages.push(Message { - role: "user".to_string(), - content: user.to_string(), - reasoning_content: None, - }); - - self.chat_completion_with_retry(messages).await - } - - async fn is_available(&self) -> bool { - self.validate_key().await.unwrap_or(false) - } - - fn name(&self) -> &str { - "deepseek" - } -} - -impl DeepSeekClient { - async fn chat_completion_with_retry(&self, messages: Vec) -> Result { - let mut last_error = None; - - for attempt in 1..=3 { - match self.chat_completion(messages.clone()).await { - Ok(result) => return Ok(result), - Err(e) => { - let err_msg = e.to_string(); - // 网络临时错误才重试 - let is_retryable = err_msg.contains("timeout") - || err_msg.contains("connection") - || err_msg.contains("temporary") - || err_msg.contains("5") - && (err_msg.contains("500") - || err_msg.contains("502") - || err_msg.contains("503") - || err_msg.contains("504")); - - if !is_retryable || attempt == 3 { - last_error = Some(e); - break; - } - - // 指数退避 - tokio::time::sleep(Duration::from_millis(500 * 2u64.pow(attempt - 1))).await; - } - } - } - - Err(last_error.unwrap_or_else(|| anyhow::anyhow!("Request failed after retries"))) - } - - async fn chat_completion(&self, messages: Vec) -> Result { - let url = format!("{}/chat/completions", self.base_url); - - let thinking = Some(ThinkingConfig { - thinking_type: if self.thinking_enabled { - "enabled".to_string() - } else { - "disabled".to_string() - }, - }); - - // 思考模式下,temperature/top_p 等参数不应传递 - // 非思考模式下可以正常传递 - let (temperature, max_tokens, top_p, presence_penalty, frequency_penalty) = - if self.thinking_enabled { - (None, Some(self.max_tokens), None, None, None) - } else { - ( - Some(self.temperature), - Some(self.max_tokens), - None, - None, - None, - ) - }; - - let reasoning_effort = if self.thinking_enabled { - self.reasoning_effort.clone() - } else { - None - }; - - let request = ChatCompletionRequest { - model: self.model.clone(), - messages: messages.clone(), - max_tokens, - temperature, - top_p, - presence_penalty, - frequency_penalty, - stream: self.thinking_enabled, - thinking, - reasoning_effort, - }; - - if self.thinking_enabled { - self.streaming_chat_completion(&url, &request).await - } else { - self.non_streaming_chat_completion(&url, &request).await - } - } - - /// 非流式请求(非思考模式) - async fn non_streaming_chat_completion( - &self, - url: &str, - request: &ChatCompletionRequest, - ) -> Result { - let response = self - .client - .post(url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .json(request) - .send() - .await - .context("Failed to send request to DeepSeek")?; - - let status = response.status(); - - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - - if let Ok(error) = serde_json::from_str::(&text) { - bail!( - "DeepSeek API error: {} ({})", - error.error.message, - error.error.error_type - ); - } - - bail!("DeepSeek API error: {} - {}", status, text); - } - - let result: ChatCompletionResponse = response - .json() - .await - .context("Failed to parse DeepSeek response")?; - - result - .choices - .into_iter() - .next() - .map(|c| c.message.content.trim().to_string()) - .filter(|s| !s.is_empty()) - .ok_or_else(|| anyhow::anyhow!("No response from DeepSeek")) - } - - /// 流式请求(思考模式),处理 reasoning_content 和 content - async fn streaming_chat_completion( - &self, - url: &str, - request: &ChatCompletionRequest, - ) -> Result { - let response = self - .client - .post(url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .header("Accept", "text/event-stream") - .json(request) - .send() - .await - .context("Failed to send streaming request to DeepSeek")?; - - let status = response.status(); - - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - - if let Ok(error) = serde_json::from_str::(&text) { - bail!( - "DeepSeek API error: {} ({})", - error.error.message, - error.error.error_type - ); - } - - bail!("DeepSeek API error: {} - {}", status, text); - } - - let mut content_buffer = String::new(); - let mut has_reasoning = false; - let mut has_content = false; - let mut stream_ended = false; - - let thinking_state = self.thinking_state.as_ref(); - - let mut byte_stream = response.bytes_stream(); - let mut line_buffer = String::new(); - - use futures_util::StreamExt; - - while let Some(chunk) = byte_stream.next().await { - let chunk = chunk.context("Failed to read streaming response chunk")?; - let chunk_str = - String::from_utf8(chunk.to_vec()).context("Invalid UTF-8 in stream chunk")?; - - line_buffer.push_str(&chunk_str); - - // 处理完整行 - while let Some(line_end) = line_buffer.find('\n') { - let line = line_buffer[..line_end].trim().to_string(); - line_buffer = line_buffer[line_end + 1..].to_string(); - - if line.is_empty() { - continue; - } - - // SSE 格式:data: {...} 或 data: [DONE] - if line == "data: [DONE]" { - stream_ended = true; - break; - } - - if let Some(json_str) = line.strip_prefix("data: ") { - match serde_json::from_str::(json_str) { - Ok(chunk) => { - for choice in &chunk.choices { - // 处理 reasoning_content - if let Some(ref reasoning) = choice.delta.reasoning_content - && !reasoning.is_empty() - { - if !has_reasoning { - has_reasoning = true; - if let Some(state) = thinking_state { - state.start_thinking(); - } - } - // reasoning_content 不对外输出,仅用于内部状态判断 - continue; - } - - // 处理 content - if let Some(ref content) = choice.delta.content - && !content.is_empty() - { - // reasoning 结束,content 开始出现时移除 thinking 标识 - if has_reasoning - && !has_content - && let Some(state) = thinking_state - { - state.end_thinking(); - } - has_content = true; - content_buffer.push_str(content); - } - - // 检查 finish_reason - if let Some(ref reason) = choice.finish_reason - && reason == "stop" - { - stream_ended = true; - } - } - } - Err(_) => { - // 忽略无法解析的行(可能是心跳或注释) - } - } - } - } - - if stream_ended { - break; - } - } - - // 确保思考状态已结束 - if let Some(state) = thinking_state { - state.end_thinking(); - } - - let result = content_buffer.trim().to_string(); - - if result.is_empty() { - if has_reasoning && !has_content { - bail!( - "DeepSeek returned reasoning content but no final answer. \ - The model may have entered an incomplete thinking state. \ - Please try again or disable thinking mode." - ); - } - bail!( - "No response from DeepSeek. \ - If thinking mode is enabled, try disabling it or ensure the model supports it." - ); - } - - Ok(result) - } -} - -/// 可用 DeepSeek 模型列表 -/// deepseek-chat / deepseek-reasoner 将于 2026-07-24 停用,推荐使用 V4 系列 -pub const DEEPSEEK_MODELS: &[&str] = &[ - "deepseek-v4-flash", - "deepseek-v4-pro", - // 兼容旧版模型 ID(将于 2026-07-24 停用) - "deepseek-chat", - "deepseek-reasoner", -]; - -pub fn is_valid_model(model: &str) -> bool { - DEEPSEEK_MODELS.contains(&model) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_model_validation_v4() { - assert!(is_valid_model("deepseek-v4-flash")); - assert!(is_valid_model("deepseek-v4-pro")); - assert!(is_valid_model("deepseek-chat")); - assert!(is_valid_model("deepseek-reasoner")); - assert!(!is_valid_model("invalid-model")); - assert!(!is_valid_model("deepseek-v3")); - } - - #[test] - fn test_client_builder_defaults() { - let client = DeepSeekClient::new("test-key", "deepseek-v4-flash").unwrap(); - assert!(!client.thinking_enabled); - assert_eq!(client.max_tokens, 500); - assert_eq!(client.temperature, 0.7); - assert!(client.reasoning_effort.is_none()); - assert!(client.thinking_state.is_none()); - } - - #[test] - fn test_client_builder_with_thinking() { - let client = DeepSeekClient::new("test-key", "deepseek-v4-flash") - .unwrap() - .with_thinking(true) - .with_reasoning_effort(Some("high".to_string())) - .with_max_tokens(1000) - .with_temperature(0.5); - - assert!(client.thinking_enabled); - assert_eq!(client.reasoning_effort, Some("high".to_string())); - assert_eq!(client.max_tokens, 1000); - assert_eq!(client.temperature, 0.5); - } - - #[test] - fn test_thinking_config_serialization() { - let config = ThinkingConfig { - thinking_type: "enabled".to_string(), - }; - let json = serde_json::to_string(&config).unwrap(); - assert_eq!(json, r#"{"type":"enabled"}"#); - } - - #[test] - fn test_message_serialization_without_reasoning() { - let msg = Message { - role: "user".to_string(), - content: "Hello".to_string(), - reasoning_content: None, - }; - let json = serde_json::to_string(&msg).unwrap(); - assert!(!json.contains("reasoning_content")); - } - - #[test] - fn test_stream_delta_parsing() { - let json = r#"{"content":"Hello","reasoning_content":null}"#; - let delta: StreamDelta = serde_json::from_str(json).unwrap(); - assert_eq!(delta.content, Some("Hello".to_string())); - assert!(delta.reasoning_content.is_none()); - } - - #[test] - fn test_stream_delta_reasoning_only() { - let json = r#"{"content":null,"reasoning_content":"Let me think..."}"#; - let delta: StreamDelta = serde_json::from_str(json).unwrap(); - assert!(delta.content.is_none()); - assert_eq!(delta.reasoning_content, Some("Let me think...".to_string())); - } -} diff --git a/src/llm/kimi.rs b/src/llm/kimi.rs deleted file mode 100644 index f382a76..0000000 --- a/src/llm/kimi.rs +++ /dev/null @@ -1,587 +0,0 @@ -use super::thinking::ThinkingStateManager; -use super::{LlmProvider, create_http_client}; -use anyhow::{Context, Result, bail}; -use async_trait::async_trait; -use serde::{Deserialize, Serialize}; -use std::sync::Arc; -use std::time::Duration; - -/// Kimi API client (Moonshot AI) -pub struct KimiClient { - base_url: String, - api_key: String, - model: String, - client: reqwest::Client, - thinking_enabled: bool, - max_tokens: u32, - temperature: f32, - thinking_state: Option>, -} - -#[derive(Debug, Serialize)] -struct ChatCompletionRequest { - model: String, - messages: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - max_tokens: Option, - #[serde(skip_serializing_if = "Option::is_none")] - temperature: Option, - stream: bool, - #[serde(skip_serializing_if = "Option::is_none")] - thinking: Option, -} - -#[derive(Debug, Serialize)] -struct ThinkingConfig { - #[serde(rename = "type")] - thinking_type: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -struct Message { - role: String, - content: String, - #[serde(skip_serializing_if = "Option::is_none")] - reasoning_content: Option, -} - -#[derive(Debug, Deserialize)] -struct ChatCompletionResponse { - choices: Vec, -} - -#[derive(Debug, Deserialize)] -struct Choice { - message: Message, - #[serde(default)] - reasoning_content: Option, -} - -// --- Streaming response structures --- - -#[derive(Debug, Deserialize)] -struct StreamChunk { - choices: Vec, -} - -#[derive(Debug, Deserialize)] -struct StreamChoice { - delta: StreamDelta, - #[serde(default)] - finish_reason: Option, - index: Option, -} - -#[derive(Debug, Deserialize, Default)] -struct StreamDelta { - #[serde(default)] - content: Option, - #[serde(default)] - reasoning_content: Option, -} - -#[derive(Debug, Deserialize)] -struct ErrorResponse { - error: ApiError, -} - -#[derive(Debug, Deserialize)] -struct ApiError { - message: String, - #[serde(rename = "type")] - error_type: String, -} - -impl KimiClient { - pub fn new(api_key: &str, model: &str) -> Result { - let client = create_http_client(Duration::from_secs(300))?; - - Ok(Self { - base_url: "https://api.moonshot.cn/v1".to_string(), - api_key: api_key.to_string(), - model: model.to_string(), - client, - thinking_enabled: false, - max_tokens: 500, - temperature: 1.0, - thinking_state: None, - }) - } - - pub fn with_base_url(api_key: &str, model: &str, base_url: &str) -> Result { - let client = create_http_client(Duration::from_secs(300))?; - - Ok(Self { - base_url: base_url.trim_end_matches('/').to_string(), - api_key: api_key.to_string(), - model: model.to_string(), - client, - thinking_enabled: false, - max_tokens: 500, - temperature: 1.0, - thinking_state: None, - }) - } - - pub fn with_timeout(mut self, timeout: Duration) -> Result { - self.client = create_http_client(timeout)?; - Ok(self) - } - - pub fn with_thinking(mut self, enabled: bool) -> Self { - self.thinking_enabled = enabled; - self - } - - pub fn with_max_tokens(mut self, max_tokens: u32) -> Self { - self.max_tokens = max_tokens; - self - } - - pub fn with_temperature(mut self, temperature: f32) -> Self { - self.temperature = temperature; - self - } - - pub fn with_thinking_state(mut self, state: Arc) -> Self { - self.thinking_state = Some(state); - self - } - - pub async fn list_models(&self) -> Result> { - let url = format!("{}/models", self.base_url); - - let response = self - .client - .get(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .send() - .await - .context("Failed to list Kimi models")?; - - if !response.status().is_success() { - let status = response.status(); - let text = response.text().await.unwrap_or_default(); - bail!("Kimi API error: {} - {}", status, text); - } - - #[derive(Deserialize)] - struct ModelsResponse { - data: Vec, - } - - #[derive(Deserialize)] - struct ModelId { - id: String, - } - - let result: ModelsResponse = response - .json() - .await - .context("Failed to parse Kimi response")?; - - Ok(result.data.into_iter().map(|m| m.id).collect()) - } - - pub async fn validate_key(&self) -> Result { - match self.list_models().await { - Ok(_) => Ok(true), - Err(e) => { - let err_str = e.to_string(); - if err_str.contains("401") || err_str.contains("Unauthorized") { - Ok(false) - } else { - Err(e) - } - } - } - } -} - -#[async_trait] -impl LlmProvider for KimiClient { - async fn generate(&self, prompt: &str) -> Result { - let messages = vec![Message { - role: "user".to_string(), - content: prompt.to_string(), - reasoning_content: None, - }]; - - self.chat_completion_with_retry(messages).await - } - - async fn generate_with_system(&self, system: &str, user: &str) -> Result { - let mut messages = vec![]; - - if !system.is_empty() { - messages.push(Message { - role: "system".to_string(), - content: system.to_string(), - reasoning_content: None, - }); - } - - messages.push(Message { - role: "user".to_string(), - content: user.to_string(), - reasoning_content: None, - }); - - self.chat_completion_with_retry(messages).await - } - - async fn is_available(&self) -> bool { - self.validate_key().await.unwrap_or(false) - } - - fn name(&self) -> &str { - "kimi" - } -} - -impl KimiClient { - async fn chat_completion_with_retry(&self, messages: Vec) -> Result { - let mut last_error = None; - - for attempt in 1..=3 { - match self.chat_completion(messages.clone()).await { - Ok(result) => return Ok(result), - Err(e) => { - let err_msg = e.to_string(); - let is_retryable = err_msg.contains("timeout") - || err_msg.contains("connection") - || err_msg.contains("temporary") - || err_msg.contains("5") - && (err_msg.contains("500") - || err_msg.contains("502") - || err_msg.contains("503") - || err_msg.contains("504")); - - if !is_retryable || attempt == 3 { - last_error = Some(e); - break; - } - - tokio::time::sleep(Duration::from_millis(500 * 2u64.pow(attempt - 1))).await; - } - } - } - - Err(last_error.unwrap_or_else(|| anyhow::anyhow!("Request failed after retries"))) - } - - async fn chat_completion(&self, messages: Vec) -> Result { - let url = format!("{}/chat/completions", self.base_url); - - let thinking = Some(ThinkingConfig { - thinking_type: if self.thinking_enabled { - "enabled".to_string() - } else { - "disabled".to_string() - }, - }); - - // Kimi API temperature 要求: - // - 思考模式: temperature 必须为 1.0 - // - 非思考模式: temperature 必须为 0.6 - let temperature = if self.thinking_enabled { - Some(1.0) - } else { - Some(0.6) - }; - - let request = ChatCompletionRequest { - model: self.model.clone(), - messages: messages.clone(), - max_tokens: Some(self.max_tokens), - temperature, - stream: self.thinking_enabled, - thinking, - }; - - if self.thinking_enabled { - self.streaming_chat_completion(&url, &request).await - } else { - self.non_streaming_chat_completion(&url, &request).await - } - } - - /// 非流式请求(非思考模式) - async fn non_streaming_chat_completion( - &self, - url: &str, - request: &ChatCompletionRequest, - ) -> Result { - let response = self - .client - .post(url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .json(request) - .send() - .await - .context("Failed to send request to Kimi")?; - - let status = response.status(); - - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - - if let Ok(error) = serde_json::from_str::(&text) { - bail!( - "Kimi API error: {} ({})", - error.error.message, - error.error.error_type - ); - } - - bail!("Kimi API error: {} - {}", status, text); - } - - let result: ChatCompletionResponse = response - .json() - .await - .context("Failed to parse Kimi response")?; - - result - .choices - .into_iter() - .next() - .map(|c| { - let content = c.message.content.trim().to_string(); - if content.is_empty() { - c.reasoning_content - .or(c.message.reasoning_content) - .map(|r| r.trim().to_string()) - .unwrap_or_default() - } else { - content - } - }) - .filter(|s| !s.is_empty()) - .ok_or_else(|| anyhow::anyhow!("No response from Kimi")) - } - - /// 流式请求(思考模式),处理 reasoning_content 和 content - async fn streaming_chat_completion( - &self, - url: &str, - request: &ChatCompletionRequest, - ) -> Result { - let response = self - .client - .post(url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .header("Accept", "text/event-stream") - .json(request) - .send() - .await - .context("Failed to send streaming request to Kimi")?; - - let status = response.status(); - - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - - if let Ok(error) = serde_json::from_str::(&text) { - bail!( - "Kimi API error: {} ({})", - error.error.message, - error.error.error_type - ); - } - - bail!("Kimi API error: {} - {}", status, text); - } - - let mut content_buffer = String::new(); - let mut has_reasoning = false; - let mut has_content = false; - let mut stream_ended = false; - - let thinking_state = self.thinking_state.as_ref(); - - let mut byte_stream = response.bytes_stream(); - let mut line_buffer = String::new(); - - use futures_util::StreamExt; - - while let Some(chunk) = byte_stream.next().await { - let chunk = chunk.context("Failed to read streaming response chunk")?; - let chunk_str = - String::from_utf8(chunk.to_vec()).context("Invalid UTF-8 in stream chunk")?; - - line_buffer.push_str(&chunk_str); - - while let Some(line_end) = line_buffer.find('\n') { - let line = line_buffer[..line_end].trim().to_string(); - line_buffer = line_buffer[line_end + 1..].to_string(); - - if line.is_empty() { - continue; - } - - if line == "data: [DONE]" { - stream_ended = true; - break; - } - - if let Some(json_str) = line.strip_prefix("data: ") { - match serde_json::from_str::(json_str) { - Ok(chunk) => { - for choice in &chunk.choices { - if let Some(ref reasoning) = choice.delta.reasoning_content - && !reasoning.is_empty() - { - if !has_reasoning { - has_reasoning = true; - if let Some(state) = thinking_state { - state.start_thinking(); - } - } - continue; - } - - if let Some(ref content) = choice.delta.content - && !content.is_empty() - { - if has_reasoning - && !has_content - && let Some(state) = thinking_state - { - state.end_thinking(); - } - has_content = true; - content_buffer.push_str(content); - } - - if let Some(ref reason) = choice.finish_reason - && reason == "stop" - { - stream_ended = true; - } - } - } - Err(_) => { - // 忽略无法解析的行 - } - } - } - } - - if stream_ended { - break; - } - } - - // 确保思考状态已结束 - if let Some(state) = thinking_state { - state.end_thinking(); - } - - let result = content_buffer.trim().to_string(); - - if result.is_empty() { - if has_reasoning && !has_content { - bail!( - "Kimi returned reasoning content but no final answer. \ - The model may have entered an incomplete thinking state. \ - Please try again or disable thinking mode." - ); - } - bail!( - "No response from Kimi. \ - If thinking mode is enabled, try disabling it or ensure the model supports it." - ); - } - - Ok(result) - } -} - -/// 可用 Kimi 模型列表 -pub const KIMI_MODELS: &[&str] = &[ - // K2 系列(推荐) - "kimi-k2.6", - "kimi-k2.5", - "kimi-k2-thinking", - "kimi-k2-thinking-turbo", - "kimi-k2-instruct", - "kimi-k2-instruct-0905", - // 兼容旧版模型 ID - "moonshot-v1-8k", - "moonshot-v1-32k", - "moonshot-v1-128k", -]; - -pub fn is_valid_model(model: &str) -> bool { - KIMI_MODELS.contains(&model) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_model_validation_k2() { - assert!(is_valid_model("kimi-k2.6")); - assert!(is_valid_model("kimi-k2.5")); - assert!(is_valid_model("kimi-k2-thinking")); - assert!(is_valid_model("kimi-k2-thinking-turbo")); - assert!(is_valid_model("moonshot-v1-8k")); - assert!(is_valid_model("moonshot-v1-32k")); - assert!(is_valid_model("moonshot-v1-128k")); - assert!(!is_valid_model("invalid-model")); - assert!(!is_valid_model("kimi-k1.5")); - } - - #[test] - fn test_client_builder_defaults() { - let client = KimiClient::new("test-key", "kimi-k2.6").unwrap(); - assert!(!client.thinking_enabled); - assert_eq!(client.max_tokens, 500); - assert_eq!(client.temperature, 1.0); - assert!(client.thinking_state.is_none()); - } - - #[test] - fn test_client_builder_with_thinking() { - let client = KimiClient::new("test-key", "kimi-k2.6") - .unwrap() - .with_thinking(true) - .with_max_tokens(1000) - .with_temperature(0.5); - - assert!(client.thinking_enabled); - assert_eq!(client.max_tokens, 1000); - assert_eq!(client.temperature, 0.5); - } - - #[test] - fn test_thinking_config_serialization() { - let config = ThinkingConfig { - thinking_type: "enabled".to_string(), - }; - let json = serde_json::to_string(&config).unwrap(); - assert_eq!(json, r#"{"type":"enabled"}"#); - } - - #[test] - fn test_client_new_defaults() { - let client = KimiClient::new("test-key", "kimi-k2.6").unwrap(); - assert_eq!(client.name(), "kimi"); - assert!(!client.thinking_enabled); - } - - #[test] - fn test_message_serialization() { - let msg = Message { - role: "user".to_string(), - content: "Hello".to_string(), - reasoning_content: None, - }; - let json = serde_json::to_string(&msg).unwrap(); - assert!(!json.contains("reasoning_content")); - } -} diff --git a/src/llm/mod.rs b/src/llm/mod.rs index 8dd2724..63e5c3a 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -1,327 +1,14 @@ -use crate::config::Language; -use anyhow::{Context, Result, bail}; -use async_trait::async_trait; -use std::time::Duration; - -pub mod anthropic; -pub mod deepseek; -pub mod kimi; -pub mod ollama; -pub mod openai; -pub mod openrouter; -pub mod parsing; -pub mod prompts; -pub mod rig; -pub mod thinking; - -pub use anthropic::AnthropicClient; -pub use deepseek::DeepSeekClient; -pub use kimi::KimiClient; -pub use ollama::OllamaClient; -pub use openai::OpenAiClient; -pub use openrouter::OpenRouterClient; -pub use parsing::GeneratedCommit; - -/// LLM provider trait -#[async_trait] -pub trait LlmProvider: Send + Sync { - /// Generate text from prompt - async fn generate(&self, prompt: &str) -> Result; - - /// Generate with system prompt - async fn generate_with_system(&self, system: &str, user: &str) -> Result; - - /// Check if provider is available - async fn is_available(&self) -> bool; - - /// Get provider name - fn name(&self) -> &str; -} - -/// LLM client that wraps different providers -pub struct LlmClient { - provider: Box, - config: LlmClientConfig, -} - -#[derive(Debug, Clone)] -pub struct LlmClientConfig { - pub max_tokens: u32, - pub temperature: f32, - pub timeout: Duration, - pub thinking_enabled: bool, -} - -impl Default for LlmClientConfig { - fn default() -> Self { - Self { - max_tokens: 500, - temperature: 0.7, - timeout: Duration::from_secs(30), - thinking_enabled: false, - } - } -} - -impl LlmClient { - /// Create LLM client from configuration manager - pub async fn from_config(manager: &crate::config::manager::ConfigManager) -> Result { - Self::from_config_with_think(manager, manager.config().llm.thinking_enabled).await - } - - /// Create LLM client from configuration with explicit thinking override - pub async fn from_config_with_think( - manager: &crate::config::manager::ConfigManager, - thinking_enabled: bool, - ) -> Result { - let config = manager.config(); - let client_config = LlmClientConfig { - max_tokens: config.llm.max_tokens, - temperature: config.llm.temperature, - timeout: Duration::from_secs(config.llm.timeout), - thinking_enabled, - }; - - let provider = config.llm.provider.as_str(); - let model = config.llm.model.as_str(); - let base_url = manager.llm_base_url(); - let api_key = manager.get_api_key(); - - let provider: Box = match provider { - "ollama" => Box::new( - OllamaClient::new(&base_url, model) - .with_max_tokens(client_config.max_tokens) - .with_temperature(client_config.temperature), - ), - "openai" => { - let key = api_key - .as_ref() - .ok_or_else(|| anyhow::anyhow!("OpenAI API key not configured"))?; - let thinking_state = if thinking_enabled { - Some(thinking::create_console_thinking_state()) - } else { - None - }; - let mut client = OpenAiClient::new(&base_url, key, model)? - .with_thinking(thinking_enabled) - .with_max_tokens(client_config.max_tokens) - .with_temperature(client_config.temperature) - .with_timeout(client_config.timeout)?; - if let Some(state) = thinking_state { - client = client.with_thinking_state(state); - } - Box::new(client) - } - "anthropic" => { - let key = api_key - .as_ref() - .ok_or_else(|| anyhow::anyhow!("Anthropic API key not configured"))?; - let thinking_state = if thinking_enabled { - Some(thinking::create_console_thinking_state()) - } else { - None - }; - let budget = config.llm.thinking_budget_tokens.unwrap_or(1024); - let mut client = AnthropicClient::new(key, model)? - .with_thinking(thinking_enabled) - .with_thinking_budget_tokens(budget) - .with_max_tokens(client_config.max_tokens) - .with_temperature(client_config.temperature) - .with_timeout(client_config.timeout)?; - if let Some(state) = thinking_state { - client = client.with_thinking_state(state); - } - Box::new(client) - } - "kimi" => { - let key = api_key - .as_ref() - .ok_or_else(|| anyhow::anyhow!("Kimi API key not configured"))?; - let thinking_state = if thinking_enabled { - Some(thinking::create_console_thinking_state()) - } else { - None - }; - let mut client = KimiClient::with_base_url(key, model, &base_url)? - .with_thinking(thinking_enabled) - .with_max_tokens(client_config.max_tokens) - .with_temperature(client_config.temperature) - .with_timeout(client_config.timeout)?; - if let Some(state) = thinking_state { - client = client.with_thinking_state(state); - } - Box::new(client) - } - "deepseek" => { - let key = api_key - .as_ref() - .ok_or_else(|| anyhow::anyhow!("DeepSeek API key not configured"))?; - let thinking_state = if thinking_enabled { - Some(thinking::create_console_thinking_state()) - } else { - None - }; - let mut client = DeepSeekClient::with_base_url(key, model, &base_url)? - .with_thinking(thinking_enabled) - .with_max_tokens(client_config.max_tokens) - .with_temperature(client_config.temperature) - .with_timeout(client_config.timeout)?; - if let Some(state) = thinking_state { - client = client.with_thinking_state(state); - } - Box::new(client) - } - "openrouter" => { - let key = api_key - .as_ref() - .ok_or_else(|| anyhow::anyhow!("OpenRouter API key not configured"))?; - Box::new( - OpenRouterClient::with_base_url(key, model, &base_url)? - .with_max_tokens(client_config.max_tokens) - .with_temperature(client_config.temperature) - .with_timeout(client_config.timeout)?, - ) - } - _ => bail!("Unknown LLM provider: {}", provider), - }; - - Ok(Self { - provider, - config: client_config, - }) - } - - /// Create with specific provider - pub fn with_provider(provider: Box) -> Self { - Self { - provider, - config: LlmClientConfig::default(), - } - } - - /// Generate commit message from git diff - pub async fn generate_commit_message( - &self, - diff: &str, - format: crate::config::CommitFormat, - language: Language, - template: Option<&str>, - ) -> Result { - let mut system_prompt = prompts::get_commit_system_prompt(format, language).to_string(); - - if let Some(tmpl) = template { - system_prompt.push_str(&format!( - "\n\n## Commit Message Template\nFollow this template structure:\n{}", - tmpl - )); - } - - // Add language instruction to the prompt - let language_instruction = match language { - Language::Chinese => "\n\n请用中文生成提交消息。", - Language::Japanese => "\n\n日本語でコミットメッセージを生成してください。", - Language::Korean => "\n\n한국어로 커밋 메시지를 생성하세요.", - Language::Spanish => "\n\nPor favor, genera el mensaje de commit en español.", - Language::French => "\n\nVeuillez générer le message de commit en français.", - Language::German => "\n\nBitte generieren Sie die Commit-Nachricht auf Deutsch.", - Language::English => "", - }; - - let prompt = format!("{}{}", diff, language_instruction); - let response = self - .provider - .generate_with_system(&system_prompt, &prompt) - .await?; - - parsing::parse_commit_response(&response, format) - } - - /// Generate tag message from commits - pub async fn generate_tag_message( - &self, - version: &str, - commits: &[String], - language: Language, - ) -> Result { - let system_prompt = prompts::get_tag_system_prompt(language); - let commits_text = commits.join("\n"); - - // Add language instruction to the prompt - let language_instruction = match language { - Language::Chinese => "\n\n请用中文生成标签消息。", - Language::Japanese => "\n\n日本語でタグメッセージを生成してください。", - Language::Korean => "\n\n한국어로 태그 메시지를 생성하세요.", - Language::Spanish => "\n\nPor favor, genera el mensaje de etiqueta en español.", - Language::French => "\n\nVeuillez générer le message de balise en français.", - Language::German => "\n\nBitte generieren Sie die Tag-Nachricht auf Deutsch.", - Language::English => "", - }; - - let prompt = format!( - "Version: {}\n\nCommits:\n{}{}", - version, commits_text, language_instruction - ); - - self.provider - .generate_with_system(system_prompt, &prompt) - .await - } - - /// Generate changelog entry - pub async fn generate_changelog_entry( - &self, - version: &str, - commits: &[(String, String)], // (type, message) - language: Language, - ) -> Result { - let system_prompt = prompts::get_changelog_system_prompt(language); - - let commits_text = commits - .iter() - .map(|(t, m)| format!("- [{}] {}", t, m)) - .collect::>() - .join("\n"); - - // Add language instruction to the prompt - let language_instruction = match language { - Language::Chinese => "\n\n请用中文生成变更日志。", - Language::Japanese => "\n\n日本語で変更ログを生成してください。", - Language::Korean => "\n\n한국어로 변경 로그를 생성하세요.", - Language::Spanish => "\n\nPor favor, genera el registro de cambios en español.", - Language::French => "\n\nVeuillez générer le journal des modifications en français.", - Language::German => "\n\nBitte generieren Sie das Changelog auf Deutsch.", - Language::English => "", - }; - - let prompt = format!( - "Version: {}\n\nCommits:\n{}{}", - version, commits_text, language_instruction - ); - - self.provider - .generate_with_system(system_prompt, &prompt) - .await - } - - /// Check if provider is available - pub async fn is_available(&self) -> bool { - self.provider.is_available().await - } - -} - - -/// HTTP client helper -pub(crate) fn create_http_client(timeout: Duration) -> Result { - reqwest::Client::builder() - .timeout(timeout) - .build() - .context("Failed to create HTTP client") -} - - -/// Test LLM connection -pub async fn test_connection(manager: &crate::config::manager::ConfigManager) -> Result { - let client = crate::llm::rig::LlmClient::from_config(manager).await?; - client.generate(None, "Say 'Hello, World!'").await -} \ No newline at end of file +pub mod parsing; +pub mod prompts; +pub mod rig; +pub mod thinking; + +pub use parsing::GeneratedCommit; + +use anyhow::Result; + +/// Test LLM connection +pub async fn test_connection(manager: &crate::config::manager::ConfigManager) -> Result { + let client = crate::llm::rig::LlmClient::from_config(manager).await?; + client.generate(None, "Say 'Hello, World!'").await +} diff --git a/src/llm/ollama.rs b/src/llm/ollama.rs deleted file mode 100644 index f544ed7..0000000 --- a/src/llm/ollama.rs +++ /dev/null @@ -1,229 +0,0 @@ -use super::{LlmProvider, create_http_client}; -use anyhow::{Context, Result}; -use async_trait::async_trait; -use serde::{Deserialize, Serialize}; -use std::time::Duration; - -/// Ollama API client -pub struct OllamaClient { - base_url: String, - model: String, - client: reqwest::Client, - max_tokens: u32, - temperature: f32, - top_p: Option, -} - -#[derive(Debug, Serialize)] -struct GenerateRequest { - model: String, - prompt: String, - system: Option, - stream: bool, - options: GenerationOptions, -} - -#[derive(Debug, Serialize, Default)] -struct GenerationOptions { - #[serde(skip_serializing_if = "Option::is_none")] - temperature: Option, - #[serde(skip_serializing_if = "Option::is_none")] - num_predict: Option, -} - -#[derive(Debug, Deserialize)] -struct GenerateResponse { - response: String, - done: bool, -} - -#[derive(Debug, Deserialize)] -struct ListModelsResponse { - models: Vec, -} - -#[derive(Debug, Deserialize)] -struct ModelInfo { - name: String, -} - -impl OllamaClient { - /// Create new Ollama client - pub fn new(base_url: &str, model: &str) -> Self { - let client = - create_http_client(Duration::from_secs(120)).expect("Failed to create HTTP client"); - - Self { - base_url: base_url.trim_end_matches('/').to_string(), - model: model.to_string(), - client, - max_tokens: 500, - temperature: 0.7, - top_p: None, - } - } - - /// Set timeout - pub fn with_timeout(mut self, timeout: Duration) -> Self { - self.client = create_http_client(timeout).expect("Failed to create HTTP client"); - self - } - - pub fn with_max_tokens(mut self, max_tokens: u32) -> Self { - self.max_tokens = max_tokens; - self - } - - pub fn with_temperature(mut self, temperature: f32) -> Self { - self.temperature = temperature; - self - } - - pub fn with_top_p(mut self, top_p: f32) -> Self { - self.top_p = Some(top_p); - self - } - - /// List available models - pub async fn list_models(&self) -> Result> { - let url = format!("{}/api/tags", self.base_url); - - let response = self - .client - .get(&url) - .send() - .await - .context("Failed to list Ollama models")?; - - if !response.status().is_success() { - let status = response.status(); - let text = response.text().await.unwrap_or_default(); - anyhow::bail!("Ollama API error: {} - {}", status, text); - } - - let result: ListModelsResponse = response - .json() - .await - .context("Failed to parse Ollama response")?; - - Ok(result.models.into_iter().map(|m| m.name).collect()) - } - - /// Pull a model - pub async fn pull_model(&self, model: &str) -> Result<()> { - let url = format!("{}/api/pull", self.base_url); - - let request = serde_json::json!({ - "name": model, - "stream": false, - }); - - let response = self - .client - .post(&url) - .json(&request) - .send() - .await - .context("Failed to pull Ollama model")?; - - if !response.status().is_success() { - let status = response.status(); - let text = response.text().await.unwrap_or_default(); - anyhow::bail!("Ollama pull error: {} - {}", status, text); - } - - Ok(()) - } - - /// Check if model exists - pub async fn model_exists(&self, model: &str) -> bool { - match self.list_models().await { - Ok(models) => models.contains(&model.to_string()), - Err(_) => false, - } - } -} - -#[async_trait] -impl LlmProvider for OllamaClient { - async fn generate(&self, prompt: &str) -> Result { - self.generate_with_system("", prompt).await - } - - async fn generate_with_system(&self, system: &str, user: &str) -> Result { - let url = format!("{}/api/generate", self.base_url); - - let system = if system.is_empty() { - None - } else { - Some(system.to_string()) - }; - - let request = GenerateRequest { - model: self.model.clone(), - prompt: user.to_string(), - system, - stream: false, - options: GenerationOptions { - temperature: Some(self.temperature), - num_predict: Some(self.max_tokens), - }, - }; - - let response = self - .client - .post(&url) - .json(&request) - .send() - .await - .context("Failed to send request to Ollama")?; - - if !response.status().is_success() { - let status = response.status(); - let text = response.text().await.unwrap_or_default(); - anyhow::bail!("Ollama API error: {} - {}", status, text); - } - - let result: GenerateResponse = response - .json() - .await - .context("Failed to parse Ollama response")?; - - Ok(result.response.trim().to_string()) - } - - async fn is_available(&self) -> bool { - let url = format!("{}/api/tags", self.base_url); - - match self.client.get(&url).send().await { - Ok(response) => response.status().is_success(), - Err(_) => false, - } - } - - fn name(&self) -> &str { - "ollama" - } -} - -#[cfg(test)] -mod tests { - use super::*; - - // These tests require a running Ollama server - #[tokio::test] - #[ignore] - async fn test_ollama_connection() { - let client = OllamaClient::new("http://localhost:11434", "llama2"); - assert!(client.is_available().await); - } - - #[tokio::test] - #[ignore] - async fn test_ollama_generate() { - let client = OllamaClient::new("http://localhost:11434", "llama2"); - let response = client.generate("Hello, how are you?").await; - assert!(response.is_ok()); - println!("Response: {}", response.unwrap()); - } -} diff --git a/src/llm/openai.rs b/src/llm/openai.rs deleted file mode 100644 index 21c7c12..0000000 --- a/src/llm/openai.rs +++ /dev/null @@ -1,659 +0,0 @@ -use super::thinking::ThinkingStateManager; -use super::{LlmProvider, create_http_client}; -use anyhow::{Context, Result, bail}; -use async_trait::async_trait; -use serde::{Deserialize, Serialize}; -use std::sync::Arc; -use std::time::Duration; - -/// OpenAI API client with o-series reasoning support -pub struct OpenAiClient { - base_url: String, - api_key: String, - model: String, - client: reqwest::Client, - thinking_enabled: bool, - reasoning_effort: Option, - max_tokens: u32, - temperature: f32, - top_p: Option, - thinking_state: Option>, -} - -#[derive(Debug, Serialize)] -struct ChatCompletionRequest { - model: String, - messages: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - max_tokens: Option, - #[serde(skip_serializing_if = "Option::is_none")] - temperature: Option, - #[serde(skip_serializing_if = "Option::is_none")] - top_p: Option, - #[serde(skip_serializing_if = "Option::is_none")] - reasoning_effort: Option, - stream: bool, -} - -#[derive(Debug, Serialize, Deserialize, Clone)] -struct Message { - role: String, - content: String, -} - -#[derive(Debug, Deserialize)] -struct ChatCompletionResponse { - choices: Vec, -} - -#[derive(Debug, Deserialize)] -struct Choice { - message: Message, -} - -// --- Streaming response structures --- - -#[derive(Debug, Deserialize)] -struct StreamChunk { - choices: Vec, -} - -#[derive(Debug, Deserialize)] -struct StreamChoice { - delta: StreamDelta, - #[serde(default)] - finish_reason: Option, -} - -#[derive(Debug, Deserialize, Default)] -struct StreamDelta { - #[serde(default)] - content: Option, - #[serde(default)] - reasoning_content: Option, -} - -#[derive(Debug, Deserialize)] -struct ErrorResponse { - error: ApiError, -} - -#[derive(Debug, Deserialize)] -struct ApiError { - message: String, - #[serde(rename = "type")] - error_type: String, -} - -impl OpenAiClient { - /// Create new OpenAI client - pub fn new(base_url: &str, api_key: &str, model: &str) -> Result { - let client = create_http_client(Duration::from_secs(60))?; - - Ok(Self { - base_url: base_url.trim_end_matches('/').to_string(), - api_key: api_key.to_string(), - model: model.to_string(), - client, - thinking_enabled: false, - reasoning_effort: None, - max_tokens: 500, - temperature: 0.7, - top_p: None, - thinking_state: None, - }) - } - - pub fn with_timeout(mut self, timeout: Duration) -> Result { - self.client = create_http_client(timeout)?; - Ok(self) - } - - pub fn with_thinking(mut self, enabled: bool) -> Self { - self.thinking_enabled = enabled; - self - } - - pub fn with_reasoning_effort(mut self, effort: Option) -> Self { - self.reasoning_effort = effort; - self - } - - pub fn with_max_tokens(mut self, max_tokens: u32) -> Self { - self.max_tokens = max_tokens; - self - } - - pub fn with_temperature(mut self, temperature: f32) -> Self { - self.temperature = temperature; - self - } - - pub fn with_top_p(mut self, top_p: f32) -> Self { - self.top_p = Some(top_p); - self - } - - pub fn with_thinking_state(mut self, state: Arc) -> Self { - self.thinking_state = Some(state); - self - } - - pub async fn list_models(&self) -> Result> { - let url = format!("{}/models", self.base_url); - - let response = self - .client - .get(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .send() - .await - .context("Failed to list OpenAI models")?; - - if !response.status().is_success() { - let status = response.status(); - let text = response.text().await.unwrap_or_default(); - bail!("OpenAI API error: {} - {}", status, text); - } - - #[derive(Deserialize)] - struct ModelsResponse { - data: Vec, - } - - #[derive(Deserialize)] - struct Model { - id: String, - } - - let result: ModelsResponse = response - .json() - .await - .context("Failed to parse OpenAI response")?; - - Ok(result.data.into_iter().map(|m| m.id).collect()) - } - - pub async fn validate_key(&self) -> Result { - match self.list_models().await { - Ok(_) => Ok(true), - Err(e) => { - let err_str = e.to_string(); - if err_str.contains("401") || err_str.contains("Unauthorized") { - Ok(false) - } else { - Err(e) - } - } - } - } -} - -#[async_trait] -impl LlmProvider for OpenAiClient { - async fn generate(&self, prompt: &str) -> Result { - let messages = vec![Message { - role: "user".to_string(), - content: prompt.to_string(), - }]; - - self.chat_completion_with_retry(messages).await - } - - async fn generate_with_system(&self, system: &str, user: &str) -> Result { - let mut messages = vec![]; - - if !system.is_empty() { - messages.push(Message { - role: "system".to_string(), - content: system.to_string(), - }); - } - - messages.push(Message { - role: "user".to_string(), - content: user.to_string(), - }); - - self.chat_completion_with_retry(messages).await - } - - async fn is_available(&self) -> bool { - self.validate_key().await.unwrap_or(false) - } - - fn name(&self) -> &str { - "openai" - } -} - -impl OpenAiClient { - async fn chat_completion_with_retry(&self, messages: Vec) -> Result { - let mut last_error = None; - - for attempt in 1..=3 { - match self.chat_completion(messages.clone()).await { - Ok(result) => return Ok(result), - Err(e) => { - let err_msg = e.to_string(); - let is_retryable = err_msg.contains("timeout") - || err_msg.contains("connection") - || err_msg.contains("temporary") - || err_msg.contains("5") - && (err_msg.contains("500") - || err_msg.contains("502") - || err_msg.contains("503") - || err_msg.contains("504")); - - if !is_retryable || attempt == 3 { - last_error = Some(e); - break; - } - - tokio::time::sleep(Duration::from_millis(500 * 2u64.pow(attempt - 1))).await; - } - } - } - - Err(last_error.unwrap_or_else(|| anyhow::anyhow!("Request failed after retries"))) - } - - async fn chat_completion(&self, messages: Vec) -> Result { - if self.thinking_enabled { - self.streaming_chat_completion(messages).await - } else { - self.non_streaming_chat_completion(messages).await - } - } - - async fn non_streaming_chat_completion(&self, messages: Vec) -> Result { - let url = format!("{}/chat/completions", self.base_url); - - let request = ChatCompletionRequest { - model: self.model.clone(), - messages, - max_tokens: Some(self.max_tokens), - temperature: Some(self.temperature), - top_p: self.top_p, - reasoning_effort: if is_reasoning_model(&self.model) { - Some("none".to_string()) - } else { - None - }, - stream: false, - }; - - let response = self - .client - .post(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .json(&request) - .send() - .await - .context("Failed to send request to OpenAI")?; - - let status = response.status(); - - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - - if let Ok(error) = serde_json::from_str::(&text) { - bail!( - "OpenAI API error: {} ({})", - error.error.message, - error.error.error_type - ); - } - - bail!("OpenAI API error: {} - {}", status, text); - } - - let result: ChatCompletionResponse = response - .json() - .await - .context("Failed to parse OpenAI response")?; - - result - .choices - .into_iter() - .next() - .map(|c| c.message.content.trim().to_string()) - .filter(|s| !s.is_empty()) - .ok_or_else(|| anyhow::anyhow!("No response from OpenAI")) - } - - /// Streaming request for reasoning mode, filters reasoning_content from output - async fn streaming_chat_completion(&self, messages: Vec) -> Result { - let url = format!("{}/chat/completions", self.base_url); - - // For reasoning/thinking mode, omit temperature and top_p - let request = ChatCompletionRequest { - model: self.model.clone(), - messages, - max_tokens: Some(self.max_tokens), - temperature: None, - top_p: None, - reasoning_effort: self.reasoning_effort.clone(), - stream: true, - }; - - let response = self - .client - .post(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .header("Accept", "text/event-stream") - .json(&request) - .send() - .await - .context("Failed to send streaming request to OpenAI")?; - - let status = response.status(); - - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - - if let Ok(error) = serde_json::from_str::(&text) { - bail!( - "OpenAI API error: {} ({})", - error.error.message, - error.error.error_type - ); - } - - bail!("OpenAI API error: {} - {}", status, text); - } - - let mut content_buffer = String::new(); - let mut has_reasoning = false; - let mut has_content = false; - - let thinking_state = self.thinking_state.as_ref(); - - let mut byte_stream = response.bytes_stream(); - let mut line_buffer = String::new(); - - use futures_util::StreamExt; - - while let Some(chunk) = byte_stream.next().await { - let chunk = chunk.context("Failed to read streaming response chunk")?; - let chunk_str = - String::from_utf8(chunk.to_vec()).context("Invalid UTF-8 in stream chunk")?; - - line_buffer.push_str(&chunk_str); - - while let Some(line_end) = line_buffer.find('\n') { - let line = line_buffer[..line_end].trim().to_string(); - line_buffer = line_buffer[line_end + 1..].to_string(); - - if line.is_empty() { - continue; - } - - if line == "data: [DONE]" { - break; - } - - if let Some(json_str) = line.strip_prefix("data: ") { - if let Ok(chunk) = serde_json::from_str::(json_str) { - for choice in &chunk.choices { - // Handle reasoning_content (o-series) - if let Some(ref reasoning) = choice.delta.reasoning_content - && !reasoning.is_empty() - { - if !has_reasoning { - has_reasoning = true; - if let Some(state) = thinking_state { - state.start_thinking(); - } - } - continue; - } - - // Handle content - if let Some(ref content) = choice.delta.content - && !content.is_empty() - { - if has_reasoning - && !has_content - && let Some(state) = thinking_state - { - state.end_thinking(); - } - has_content = true; - content_buffer.push_str(content); - } - } - } - } - } - } - - if let Some(state) = thinking_state { - state.end_thinking(); - } - - let result = content_buffer.trim().to_string(); - - if result.is_empty() { - if has_reasoning && !has_content { - bail!( - "OpenAI returned reasoning content but no final answer. \ - The model may have entered an incomplete reasoning state. \ - Please try again or disable thinking mode." - ); - } - bail!( - "No response from OpenAI. \ - If thinking mode is enabled, try disabling it or ensure the model supports reasoning." - ); - } - - Ok(result) - } -} - -/// Azure OpenAI client (extends OpenAI with Azure-specific config) -pub struct AzureOpenAiClient { - endpoint: String, - api_key: String, - deployment: String, - api_version: String, - client: reqwest::Client, - thinking_enabled: bool, - reasoning_effort: Option, - max_tokens: u32, - temperature: f32, - top_p: Option, - thinking_state: Option>, -} - -impl AzureOpenAiClient { - pub fn new(endpoint: &str, api_key: &str, deployment: &str, api_version: &str) -> Result { - let client = create_http_client(Duration::from_secs(60))?; - - Ok(Self { - endpoint: endpoint.trim_end_matches('/').to_string(), - api_key: api_key.to_string(), - deployment: deployment.to_string(), - api_version: api_version.to_string(), - client, - thinking_enabled: false, - reasoning_effort: None, - max_tokens: 500, - temperature: 0.7, - top_p: None, - thinking_state: None, - }) - } - - async fn chat_completion(&self, messages: Vec) -> Result { - let url = format!( - "{}/openai/deployments/{}/chat/completions?api-version={}", - self.endpoint, self.deployment, self.api_version - ); - - let request = ChatCompletionRequest { - model: self.deployment.clone(), - messages, - max_tokens: Some(self.max_tokens), - temperature: Some(self.temperature), - top_p: self.top_p, - reasoning_effort: self.reasoning_effort.clone(), - stream: false, - }; - - let response = self - .client - .post(&url) - .header("api-key", &self.api_key) - .header("Content-Type", "application/json") - .json(&request) - .send() - .await - .context("Failed to send request to Azure OpenAI")?; - - if !response.status().is_success() { - let status = response.status(); - let text = response.text().await.unwrap_or_default(); - bail!("Azure OpenAI API error: {} - {}", status, text); - } - - let result: ChatCompletionResponse = response - .json() - .await - .context("Failed to parse Azure OpenAI response")?; - - result - .choices - .into_iter() - .next() - .map(|c| c.message.content.trim().to_string()) - .filter(|s| !s.is_empty()) - .ok_or_else(|| anyhow::anyhow!("No response from Azure OpenAI")) - } -} - -#[async_trait] -impl LlmProvider for AzureOpenAiClient { - async fn generate(&self, prompt: &str) -> Result { - let messages = vec![Message { - role: "user".to_string(), - content: prompt.to_string(), - }]; - - self.chat_completion(messages).await - } - - async fn generate_with_system(&self, system: &str, user: &str) -> Result { - let mut messages = vec![]; - - if !system.is_empty() { - messages.push(Message { - role: "system".to_string(), - content: system.to_string(), - }); - } - - messages.push(Message { - role: "user".to_string(), - content: user.to_string(), - }); - - self.chat_completion(messages).await - } - - async fn is_available(&self) -> bool { - let url = format!( - "{}/openai/deployments/{}/chat/completions?api-version={}", - self.endpoint, self.deployment, self.api_version - ); - - let request = ChatCompletionRequest { - model: self.deployment.clone(), - messages: vec![Message { - role: "user".to_string(), - content: "Hi".to_string(), - }], - max_tokens: Some(5), - temperature: Some(0.0), - top_p: None, - reasoning_effort: None, - stream: false, - }; - - match self - .client - .post(&url) - .header("api-key", &self.api_key) - .json(&request) - .send() - .await - { - Ok(response) => response.status().is_success(), - Err(_) => false, - } - } - - fn name(&self) -> &str { - "azure-openai" - } -} - -/// Available OpenAI models (including o-series with reasoning) -pub const OPENAI_MODELS: &[&str] = &[ - "o4-mini", - "o3", - "o3-mini", - "o1", - "o1-mini", - "o1-pro", - "gpt-4.1", - "gpt-4.1-mini", - "gpt-4.1-nano", - "gpt-4o", - "gpt-4o-mini", - "gpt-4-turbo", - "gpt-4", - "gpt-3.5-turbo", -]; - -pub fn is_valid_model(model: &str) -> bool { - OPENAI_MODELS.contains(&model) -} - -fn is_reasoning_model(model: &str) -> bool { - model.starts_with("o") -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_model_validation_o_series() { - assert!(is_valid_model("o4-mini")); - assert!(is_valid_model("o3")); - assert!(is_valid_model("o1")); - assert!(is_valid_model("gpt-4o")); - assert!(is_valid_model("gpt-3.5-turbo")); - assert!(!is_valid_model("invalid-model")); - } - - #[test] - fn test_stream_delta_reasoning_parsing() { - let json = r#"{"content":null,"reasoning_content":"Let me think..."}"#; - let delta: StreamDelta = serde_json::from_str(json).unwrap(); - assert!(delta.content.is_none()); - assert_eq!(delta.reasoning_content, Some("Let me think...".to_string())); - } - - #[test] - fn test_stream_delta_content_parsing() { - let json = r#"{"content":"Hello","reasoning_content":null}"#; - let delta: StreamDelta = serde_json::from_str(json).unwrap(); - assert_eq!(delta.content, Some("Hello".to_string())); - assert!(delta.reasoning_content.is_none()); - } -} diff --git a/src/llm/openrouter.rs b/src/llm/openrouter.rs deleted file mode 100644 index 475d6f3..0000000 --- a/src/llm/openrouter.rs +++ /dev/null @@ -1,286 +0,0 @@ -use super::{LlmProvider, create_http_client}; -use anyhow::{Context, Result, bail}; -use async_trait::async_trait; -use serde::{Deserialize, Serialize}; -use std::time::Duration; - -/// OpenRouter API client -pub struct OpenRouterClient { - base_url: String, - api_key: String, - model: String, - client: reqwest::Client, - max_tokens: u32, - temperature: f32, - top_p: Option, -} - -#[derive(Debug, Serialize)] -struct ChatCompletionRequest { - model: String, - messages: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - max_tokens: Option, - #[serde(skip_serializing_if = "Option::is_none")] - temperature: Option, - stream: bool, -} - -#[derive(Debug, Serialize, Deserialize)] -struct Message { - role: String, - content: String, -} - -#[derive(Debug, Deserialize)] -struct ChatCompletionResponse { - choices: Vec, -} - -#[derive(Debug, Deserialize)] -struct Choice { - message: Message, -} - -#[derive(Debug, Deserialize)] -struct ErrorResponse { - error: ApiError, -} - -#[derive(Debug, Deserialize)] -struct ApiError { - message: String, - #[serde(rename = "type")] - error_type: String, -} - -impl OpenRouterClient { - /// Create new OpenRouter client - pub fn new(api_key: &str, model: &str) -> Result { - let client = create_http_client(Duration::from_secs(60))?; - - Ok(Self { - base_url: "https://openrouter.ai/api/v1".to_string(), - api_key: api_key.to_string(), - model: model.to_string(), - client, - max_tokens: 500, - temperature: 0.7, - top_p: None, - }) - } - - /// Create with custom base URL - pub fn with_base_url(api_key: &str, model: &str, base_url: &str) -> Result { - let client = create_http_client(Duration::from_secs(60))?; - - Ok(Self { - base_url: base_url.trim_end_matches('/').to_string(), - api_key: api_key.to_string(), - model: model.to_string(), - client, - max_tokens: 500, - temperature: 0.7, - top_p: None, - }) - } - - /// Set timeout - pub fn with_timeout(mut self, timeout: Duration) -> Result { - self.client = create_http_client(timeout)?; - Ok(self) - } - - pub fn with_max_tokens(mut self, max_tokens: u32) -> Self { - self.max_tokens = max_tokens; - self - } - - pub fn with_temperature(mut self, temperature: f32) -> Self { - self.temperature = temperature; - self - } - - pub fn with_top_p(mut self, top_p: f32) -> Self { - self.top_p = Some(top_p); - self - } - - /// List available models - pub async fn list_models(&self) -> Result> { - let url = format!("{}/models", self.base_url); - - let response = self - .client - .get(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("HTTP-Referer", "https://quicommit.dev") - .header("X-Title", "QuiCommit") - .send() - .await - .context("Failed to list OpenRouter models")?; - - if !response.status().is_success() { - let status = response.status(); - let text = response.text().await.unwrap_or_default(); - bail!("OpenRouter API error: {} - {}", status, text); - } - - #[derive(Deserialize)] - struct ModelsResponse { - data: Vec, - } - - #[derive(Deserialize)] - struct Model { - id: String, - } - - let result: ModelsResponse = response - .json() - .await - .context("Failed to parse OpenRouter response")?; - - Ok(result.data.into_iter().map(|m| m.id).collect()) - } - - /// Validate API key - pub async fn validate_key(&self) -> Result { - match self.list_models().await { - Ok(_) => Ok(true), - Err(e) => { - let err_str = e.to_string(); - if err_str.contains("401") || err_str.contains("Unauthorized") { - Ok(false) - } else { - Err(e) - } - } - } - } -} - -#[async_trait] -impl LlmProvider for OpenRouterClient { - async fn generate(&self, prompt: &str) -> Result { - let messages = vec![Message { - role: "user".to_string(), - content: prompt.to_string(), - }]; - - self.chat_completion(messages).await - } - - async fn generate_with_system(&self, system: &str, user: &str) -> Result { - let mut messages = vec![]; - - if !system.is_empty() { - messages.push(Message { - role: "system".to_string(), - content: system.to_string(), - }); - } - - messages.push(Message { - role: "user".to_string(), - content: user.to_string(), - }); - - self.chat_completion(messages).await - } - - async fn is_available(&self) -> bool { - self.validate_key().await.unwrap_or(false) - } - - fn name(&self) -> &str { - "openrouter" - } -} - -impl OpenRouterClient { - async fn chat_completion(&self, messages: Vec) -> Result { - let url = format!("{}/chat/completions", self.base_url); - - let request = ChatCompletionRequest { - model: self.model.clone(), - messages, - max_tokens: Some(self.max_tokens), - temperature: Some(self.temperature), - stream: false, - }; - - let response = self - .client - .post(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .header("HTTP-Referer", "https://quicommit.dev") - .header("X-Title", "QuiCommit") - .json(&request) - .send() - .await - .context("Failed to send request to OpenRouter")?; - - let status = response.status(); - - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - - // Try to parse error - if let Ok(error) = serde_json::from_str::(&text) { - bail!( - "OpenRouter API error: {} ({})", - error.error.message, - error.error.error_type - ); - } - - bail!("OpenRouter API error: {} - {}", status, text); - } - - let result: ChatCompletionResponse = response - .json() - .await - .context("Failed to parse OpenRouter response")?; - - result - .choices - .into_iter() - .next() - .map(|c| c.message.content.trim().to_string()) - .ok_or_else(|| anyhow::anyhow!("No response from OpenRouter")) - } -} - -/// Popular OpenRouter models -pub const OPENROUTER_MODELS: &[&str] = &[ - "openai/gpt-3.5-turbo", - "openai/gpt-4", - "openai/gpt-4-turbo", - "anthropic/claude-3-opus", - "anthropic/claude-3-sonnet", - "anthropic/claude-3-haiku", - "google/gemini-pro", - "meta-llama/llama-2-70b-chat", - "mistralai/mixtral-8x7b-instruct", - "01-ai/yi-34b-chat", -]; - -/// Check if a model name is valid -pub fn is_valid_model(_model: &str) -> bool { - // Since OpenRouter supports many models, we'll allow any model name - // but provide some popular ones as suggestions - true -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_model_validation() { - assert!(is_valid_model("openai/gpt-4")); - assert!(is_valid_model("custom/model")); - } -}