From 92cc00835919e948602f95687b6a446a4d569c50 Mon Sep 17 00:00:00 2001 From: wenzr <282277167@qq.com> Date: Tue, 21 Jul 2026 22:16:32 +0800 Subject: [PATCH 1/4] =?UTF-8?q?Provider:=E6=94=AF=E6=8C=81OpenAI=E5=85=BC?= =?UTF-8?q?=E5=AE=B9=E6=8E=A5=E5=8F=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 55 +++- crates/kuncode-cli/src/runtime.rs | 52 +++- crates/kuncode-cli/src/settings.rs | 137 ++++++++- crates/kuncode-core/src/json_utils.rs | 11 + crates/kuncode-core/src/providers.rs | 7 +- crates/kuncode-core/src/providers/any_chat.rs | 77 +++++ .../src/providers/deepseek/protocol.rs | 10 +- .../src/providers/openai_compatible.rs | 281 ++++++++++++++++++ 8 files changed, 603 insertions(+), 27 deletions(-) create mode 100644 crates/kuncode-core/src/providers/any_chat.rs create mode 100644 crates/kuncode-core/src/providers/openai_compatible.rs diff --git a/README.md b/README.md index b24e12b..503606c 100644 --- a/README.md +++ b/README.md @@ -4,26 +4,27 @@ [`learn-claude-code`](https://github.com/shareAI-lab/learn-claude-code) 的 Harness Engineering 思路:模型负责判断下一步做什么,Harness 负责提供工具、上下文、权限边界、持久化和用户界面。 -当前版本为 `0.1.0`,使用 DeepSeek 模型,提供一次性命令行执行和交互式 TUI 两种使用方式。 +当前版本为 `0.1.0`,默认使用 DeepSeek,也支持 OpenAI Chat Completions +兼容协议,提供一次性命令行执行和交互式 TUI 两种使用方式。 ## 工作区结构 ```text -kuncode-cli ──▶ kuncode-agent ──▶ kuncode-core ──▶ DeepSeek API +kuncode-cli ──▶ kuncode-agent ──▶ kuncode-core ──▶ LLM API │ │ │ │ │ └─ 消息、Completion、流式协议、Provider │ └─ Agent Loop、工具、权限、会话、压缩与编排 └─ 参数、配置、审批、一次性输出与 TUI ``` -- `kuncode-core`:Provider-neutral 的消息与 Completion 抽象,以及 DeepSeek Provider。 +- `kuncode-core`:Provider-neutral 的消息与 Completion 抽象,以及 DeepSeek、OpenAI-compatible Provider。 - `kuncode-agent`:Agent 运行时、工具调度、权限、Hook、Todo、会话持久化和上下文压缩。 - `kuncode-cli`:命令行参数、项目配置、终端审批、普通输出和交互式 TUI。 ## 环境要求 - Rust stable,项目使用 Rust 2024 edition。 -- DeepSeek API Key。 +- 所选模型服务的 API Key;本地无鉴权服务可不配置。 - 支持 ANSI 终端;交互模式需要 stdin 和 stdout 都连接到真实终端。 ## 快速开始 @@ -40,6 +41,42 @@ export DEEPSEEK_API_KEY="your-api-key" DEEPSEEK_API_KEY=your-api-key ``` +使用 OpenAI 官方接口时,在 `.kuncode/settings.json` 配置: + +```json +{ + "model": { + "provider": "openai-compatible", + "name": "your-openai-model", + "apiKeyEnv": "OPENAI_API_KEY", + "maxTokens": 16384 + } +} +``` + +并设置对应环境变量: + +```bash +export OPENAI_API_KEY="your-api-key" +``` + +Kimi、Qwen、智谱、vLLM 等提供 OpenAI Chat Completions 兼容协议的服务,增加 +`baseUrl` 即可。它既可以是服务根地址,也可以是完整 endpoint: + +```json +{ + "model": { + "provider": "openai-compatible", + "name": "your-model", + "baseUrl": "http://localhost:8000/v1", + "apiKeyEnv": "", + "maxTokens": 8192 + } +} +``` + +`apiKeyEnv` 为空表示不发送 `Authorization`,适合无鉴权的本地服务。 + 一次性执行任务: ```bash @@ -66,13 +103,15 @@ cargo build --release -p kuncode-cli ```json { "permissions": { - "allow": ["Read", "Bash(cargo *)"], + "allow": ["Read(**)", "Bash(cargo *)"], "ask": ["Edit(.env)"], "deny": ["Bash(curl *)"], "defaultMode": "default" }, "model": { + "provider": "deepseek", "name": "deepseek-v4-pro", + "apiKeyEnv": "DEEPSEEK_API_KEY", "maxTokens": 65536 }, "agent": { @@ -90,8 +129,12 @@ cargo build --release -p kuncode-cli 补充说明: -- `DEEPSEEK_MODEL` 可以覆盖配置文件中的模型名称。 +- `KUNCODE_MODEL` 可以覆盖配置文件中的模型名称;`DEEPSEEK_MODEL` 作为兼容别名保留。 +- `model.provider` 支持 `deepseek` 和 `openai-compatible`。 +- `openai-compatible` 未设置 `baseUrl` 时默认使用 `https://api.openai.com/v1`。 +- `baseUrl` 会自动补全 `/chat/completions`;完整 endpoint 不会重复追加。 - 内置模型配置包括 `deepseek-v4-pro` 和 `deepseek-v4-flash`。 +- 非内置模型启用上下文压缩时,需要显式设置 `compaction.contextLimit`。 - `compaction.mode` 支持 `disabled`、`shadow` 和 `enabled`,默认是 `disabled`。 - `shadow` 只计算和报告压缩候选,不替换当前上下文。 - `enabled` 会在达到预算阈值时执行压缩,并要求会话持久化状态保持健康。 diff --git a/crates/kuncode-cli/src/runtime.rs b/crates/kuncode-cli/src/runtime.rs index bb04f12..8f32659 100644 --- a/crates/kuncode-cli/src/runtime.rs +++ b/crates/kuncode-cli/src/runtime.rs @@ -27,10 +27,14 @@ use kuncode_agent::system_prompt::{ }; use kuncode_agent::workspace::Workspace; use kuncode_core::completion::{CompletionModel, RetryModel, RetryPolicy}; -use kuncode_core::providers::deepseek::{DeepSeekClient, DeepSeekCompletionModel}; +use kuncode_core::providers::{ + any_chat::{AnyChatClient, AnyChatCompletionModel}, + deepseek::DeepSeekClient, + openai_compatible::OpenAiCompatibleClient, +}; use crate::config::{PermissionFlags, resolve_permissions}; -use crate::settings::{ProjectSettings, ProjectTrust, load_project_settings}; +use crate::settings::{ProjectSettings, ProjectTrust, ProviderKind, load_project_settings}; use crate::{Cli, logging::LoggingObserver}; /// Identity and behavioral instructions rendered as the first system-prompt @@ -46,7 +50,7 @@ Keep working until the task is done, then give a short, direct final answer."; /// observer + approver, plus the bits a frontend renders directly /// ([`model_name`](Self::model_name), [`mode`](Self::mode)). Generic over the /// model so a test or a future provider can supply its own `M`; [`assemble`] -/// pins it to the CLI's [`DeepSeekCompletionModel`] wrapped in a +/// pins it to the configured [`AnyChatCompletionModel`] wrapped in a /// [`RetryModel`] so transient provider failures are retried transparently. /// /// [`assemble`]: Self::assemble @@ -64,21 +68,21 @@ pub struct CliRuntime { persistence_error: Option, } -impl CliRuntime> { +impl CliRuntime> { /// Builds the runtime from parsed CLI args and the project settings file. /// /// Resolves permissions from built-in ∪ project file ∪ CLI flags (mode /// precedence CLI > project > Default), assembles the system prompt from its - /// identity/environment/tools sections, and wires the DeepSeek model + the - /// default workspace tool registry. + /// identity/environment/tools sections, and wires the configured model + + /// the default workspace tool registry. /// /// # Errors /// /// Fails if the current directory is not a usable workspace, the project /// settings or resolved permissions are invalid, active compaction cannot - /// be bound to the selected model, or the DeepSeek client cannot be built - /// from the environment. Failure to open the optional session store is - /// retained as degraded persistence state rather than failing assembly. + /// be bound to the selected model, or the provider client cannot be built + /// from its endpoint and environment. Failure to open the optional session + /// store is retained as degraded persistence state rather than failing assembly. pub async fn assemble(cli: &Cli) -> Result> { let workspace = Workspace::from_current_dir().await?; tracing::debug!( @@ -97,6 +101,7 @@ impl CliRuntime> { let project = load_project_settings(workspace.root(), project_trust)?; let model_name = project.model_name.clone(); let config = agent_config(&project)?; + let client = provider_client(&project)?; let flags = PermissionFlags { allow: &cli.allow, ask: &cli.ask, @@ -160,11 +165,10 @@ impl CliRuntime> { (None, Some("home directory unavailable".to_string())) } }; - let client = DeepSeekClient::from_env()?; // Normal turns inherit the default retry budget. Semantic summaries use // a separate one-retry wrapper so their fallback latency is bounded // independently of ordinary model calls. - let provider = DeepSeekCompletionModel::make(&client, model_name.clone()); + let provider = AnyChatCompletionModel::make(&client, model_name.clone()); let model = RetryModel::with_policy(provider.clone(), RetryPolicy::default()); let summary_model = RetryModel::with_policy(provider, summary_retry_policy()); let registry = ToolRegistry::with_default_workspace_tools(workspace)?; @@ -185,6 +189,32 @@ impl CliRuntime> { } } +fn provider_client(project: &ProjectSettings) -> Result> { + let api_key = if project.api_key_env.trim().is_empty() { + String::new() + } else { + std::env::var(&project.api_key_env).map_err(|error| { + std::io::Error::new( + std::io::ErrorKind::NotFound, + format!( + "provider API key environment variable `{}` is unavailable: {error}", + project.api_key_env + ), + ) + })? + }; + match project.provider { + ProviderKind::DeepSeek => Ok(AnyChatClient::DeepSeek(DeepSeekClient::new(api_key)?)), + ProviderKind::OpenAiCompatible => { + let client = match project.base_url.as_deref() { + Some(base_url) => OpenAiCompatibleClient::new(api_key, base_url)?, + None => OpenAiCompatibleClient::openai(api_key)?, + }; + Ok(AnyChatClient::OpenAiCompatible(client)) + } + } +} + fn agent_config(project: &ProjectSettings) -> Result { let compaction = project .compaction diff --git a/crates/kuncode-cli/src/settings.rs b/crates/kuncode-cli/src/settings.rs index 843ebb7..657dcb1 100644 --- a/crates/kuncode-cli/src/settings.rs +++ b/crates/kuncode-cli/src/settings.rs @@ -76,19 +76,37 @@ struct PermissionsSection { #[derive(Debug, Deserialize)] #[serde(default, rename_all = "camelCase", deny_unknown_fields)] struct ModelSection { + provider: ProviderKind, name: String, + base_url: Option, + api_key_env: Option, max_tokens: Option, } impl Default for ModelSection { fn default() -> Self { Self { + provider: ProviderKind::DeepSeek, name: DEEPSEEK_V4_PRO_MODEL_ID.to_string(), + base_url: None, + api_key_env: None, max_tokens: None, } } } +/// Wire protocol selected for model requests. +#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq)] +pub(crate) enum ProviderKind { + /// Native DeepSeek behavior and environment defaults. + #[default] + #[serde(rename = "deepseek")] + DeepSeek, + /// OpenAI Chat Completions and compatible services such as vLLM or Kimi. + #[serde(rename = "openai-compatible")] + OpenAiCompatible, +} + #[derive(Debug, Deserialize)] #[serde(default, rename_all = "camelCase", deny_unknown_fields)] struct AgentSection { @@ -166,6 +184,12 @@ pub struct ProjectSettings { pub default_mode: Option, /// Trust comes from CLI/user state, never from the project file itself. pub(crate) trust: ProjectTrust, + /// Effective provider protocol. + pub(crate) provider: ProviderKind, + /// Custom OpenAI-compatible service root or full completion endpoint. + pub(crate) base_url: Option, + /// Environment variable holding the provider API key. + pub(crate) api_key_env: String, /// Effective model identifier after file and environment precedence. pub(crate) model_name: String, /// Effective provider output budget for an ordinary turn. @@ -188,6 +212,9 @@ impl Default for ProjectSettings { policy: None, default_mode: None, trust: ProjectTrust::Untrusted, + provider: ProviderKind::DeepSeek, + base_url: None, + api_key_env: "DEEPSEEK_API_KEY".to_string(), model_name: DEEPSEEK_V4_PRO_MODEL_ID.to_string(), max_tokens, max_iterations: DEFAULT_MAX_ITERATIONS, @@ -199,9 +226,10 @@ impl Default for ProjectSettings { /// Loads `.kuncode/settings.json` under `root`. /// -/// A missing file returns defaults. `DEEPSEEK_MODEL` overrides the file's model -/// name when present. Every section forms a closed schema, so misspelled fields -/// fail instead of silently selecting defaults. +/// A missing file returns defaults. `KUNCODE_MODEL` overrides the file's model +/// name; `DEEPSEEK_MODEL` remains a backward-compatible fallback. Every section +/// forms a closed schema, so misspelled fields fail instead of silently selecting +/// defaults. /// /// # Errors /// @@ -213,7 +241,9 @@ pub(crate) fn load_project_settings( root: &Path, trust: ProjectTrust, ) -> Result { - let model_override = std::env::var("DEEPSEEK_MODEL").ok(); + let model_override = std::env::var("KUNCODE_MODEL") + .ok() + .or_else(|| std::env::var("DEEPSEEK_MODEL").ok()); load_project_settings_from(root, model_override.as_deref(), trust) } @@ -266,7 +296,11 @@ fn resolve_settings( "model name must not be blank".to_string(), )); } - let profile = model_profile(&model_name); + let profile = if file.model.provider == ProviderKind::DeepSeek { + model_profile(&model_name) + } else { + None + }; let max_tokens = file.model.max_tokens.unwrap_or_else(|| { profile.map_or_else( default_agent_max_tokens, @@ -276,6 +310,32 @@ fn resolve_settings( validate_model_max_tokens(max_tokens, profile)?; validate_agent(&file.agent)?; validate_log_level(&file.logging.level)?; + let base_url = match file.model.base_url { + Some(url) if url.trim().is_empty() => { + return Err(SettingsError::Model( + "baseUrl must not be blank when provided".to_string(), + )); + } + Some(url) => Some(url.trim().to_string()), + None => None, + }; + if file.model.provider == ProviderKind::DeepSeek && base_url.is_some() { + return Err(SettingsError::Model( + "baseUrl requires provider `openai-compatible`".to_string(), + )); + } + let api_key_env = file + .model + .api_key_env + .unwrap_or_else(|| match file.model.provider { + ProviderKind::DeepSeek => "DEEPSEEK_API_KEY".to_string(), + ProviderKind::OpenAiCompatible => "OPENAI_API_KEY".to_string(), + }); + if file.model.provider == ProviderKind::DeepSeek && api_key_env.trim().is_empty() { + return Err(SettingsError::Model( + "apiKeyEnv must not be blank for the DeepSeek provider".to_string(), + )); + } let canonical_root = std::fs::canonicalize(root).map_err(|error| SettingsError::Workspace(error.to_string()))?; @@ -300,6 +360,9 @@ fn resolve_settings( policy: Some(policy), default_mode, trust, + provider: file.model.provider, + base_url, + api_key_env, model_name, max_tokens, max_iterations: file.agent.max_iterations, @@ -555,6 +618,9 @@ mod tests { .expect("a missing file is fine"); let _ = fs::remove_dir_all(&dir); + assert_eq!(loaded.provider, ProviderKind::DeepSeek); + assert_eq!(loaded.base_url, None); + assert_eq!(loaded.api_key_env, "DEEPSEEK_API_KEY"); assert!(loaded.policy.expect("resolved policy").rules().is_empty()); assert!(loaded.default_mode.is_none()); assert!(loaded.compaction.is_none()); @@ -564,6 +630,67 @@ mod tests { assert_eq!(loaded.todo_reminder_interval, Some(3)); } + #[test] + fn loads_openai_compatible_provider_settings() { + let loaded = load_json( + "openai-compatible", + r#"{ "model": { + "provider": "openai-compatible", + "name": "gpt-test", + "baseUrl": "https://api.example.com/v1/", + "apiKeyEnv": "EXAMPLE_API_KEY", + "maxTokens": 8192 + } }"#, + ) + .expect("loads"); + + assert_eq!(loaded.provider, ProviderKind::OpenAiCompatible); + assert_eq!( + loaded.base_url.as_deref(), + Some("https://api.example.com/v1/") + ); + assert_eq!(loaded.api_key_env, "EXAMPLE_API_KEY"); + assert_eq!(loaded.model_name, "gpt-test"); + assert_eq!(loaded.max_tokens, 8_192); + } + + #[test] + fn loads_explicit_deepseek_provider_name() { + let loaded = load_json( + "deepseek-provider", + r#"{ "model": { "provider": "deepseek" } }"#, + ) + .expect("loads"); + + assert_eq!(loaded.provider, ProviderKind::DeepSeek); + } + + #[test] + fn openai_compatible_provider_allows_unauthenticated_local_endpoint() { + let loaded = load_json( + "openai-compatible-local", + r#"{ "model": { + "provider": "openai-compatible", + "name": "local-model", + "baseUrl": "http://localhost:8000/v1", + "apiKeyEnv": "" + } }"#, + ) + .expect("loads"); + + assert_eq!(loaded.api_key_env, ""); + } + + #[test] + fn deepseek_rejects_openai_only_base_url() { + let result = load_json( + "deepseek-base-url", + r#"{ "model": { "baseUrl": "https://example.com/v1" } }"#, + ); + + assert!(matches!(result, Err(SettingsError::Model(_)))); + } + #[test] fn logging_level_defaults_to_info() { let dir = std::env::temp_dir().join(format!("kuncode-log-absent-{}", std::process::id())); diff --git a/crates/kuncode-core/src/json_utils.rs b/crates/kuncode-core/src/json_utils.rs index 4115cbb..17bea53 100644 --- a/crates/kuncode-core/src/json_utils.rs +++ b/crates/kuncode-core/src/json_utils.rs @@ -65,3 +65,14 @@ where let opt = > as serde::Deserialize>::deserialize(deserializer)?; Ok(opt.unwrap_or_default()) } + +/// Deserializes a defaultable value from a field that providers may send as +/// JSON `null` even though the non-null wire value is required by the schema. +pub fn null_or_default<'de, D, T>(deserializer: D) -> Result +where + D: serde::Deserializer<'de>, + T: serde::Deserialize<'de> + Default, +{ + let opt = as serde::Deserialize>::deserialize(deserializer)?; + Ok(opt.unwrap_or_default()) +} diff --git a/crates/kuncode-core/src/providers.rs b/crates/kuncode-core/src/providers.rs index 990aff0..c722cc2 100644 --- a/crates/kuncode-core/src/providers.rs +++ b/crates/kuncode-core/src/providers.rs @@ -1,7 +1,10 @@ //! Concrete LLM provider integrations. //! //! Each provider owns the mapping from the provider-agnostic -//! [`crate::completion`] types to that provider's HTTP API. Currently this -//! crate ships the [`deepseek`] integration. +//! [`crate::completion`] types to that provider's HTTP API. DeepSeek remains +//! the default, while [`openai_compatible`] supports OpenAI and compatible +//! `/chat/completions` endpoints. +pub mod any_chat; pub mod deepseek; +pub mod openai_compatible; diff --git a/crates/kuncode-core/src/providers/any_chat.rs b/crates/kuncode-core/src/providers/any_chat.rs new file mode 100644 index 0000000..a186485 --- /dev/null +++ b/crates/kuncode-core/src/providers/any_chat.rs @@ -0,0 +1,77 @@ +//! Runtime-selected chat model used when provider choice comes from configuration. + +use serde_json::Value; + +use crate::{ + completion::{ + CompletionError, CompletionModel, CompletionRequest, CompletionResponse, CompletionStream, + }, + providers::{ + deepseek::{DeepSeekClient, DeepSeekCompletionModel}, + openai_compatible::{OpenAiCompatibleClient, OpenAiCompatibleCompletionModel}, + }, +}; + +/// Provider client selected by project configuration. +#[derive(Clone)] +pub enum AnyChatClient { + /// Native DeepSeek protocol behavior. + DeepSeek(DeepSeekClient), + /// OpenAI-compatible Chat Completions behavior. + OpenAiCompatible(OpenAiCompatibleClient), +} + +/// Model handle that keeps the agent runtime independent of provider choice. +#[derive(Clone)] +pub enum AnyChatCompletionModel { + /// Native DeepSeek model. + DeepSeek(DeepSeekCompletionModel), + /// OpenAI-compatible model. + OpenAiCompatible(OpenAiCompatibleCompletionModel), +} + +impl CompletionModel for AnyChatCompletionModel { + type Response = Value; + type Client = AnyChatClient; + + fn make(client: &Self::Client, model: impl Into) -> Self { + let model = model.into(); + match client { + AnyChatClient::DeepSeek(client) => { + Self::DeepSeek(DeepSeekCompletionModel::make(client, model)) + } + AnyChatClient::OpenAiCompatible(client) => { + Self::OpenAiCompatible(OpenAiCompatibleCompletionModel::make(client, model)) + } + } + } + + async fn completion( + &self, + request: CompletionRequest, + ) -> Result, CompletionError> { + match self { + Self::DeepSeek(model) => { + let response = model.completion(request).await?; + let raw_response = serde_json::to_value(response.raw_response)?; + Ok(CompletionResponse { + choice: response.choice, + usage: response.usage, + raw_response, + message_id: response.message_id, + }) + } + Self::OpenAiCompatible(model) => model.completion(request).await, + } + } + + async fn stream( + &self, + request: CompletionRequest, + ) -> Result { + match self { + Self::DeepSeek(model) => model.stream(request).await, + Self::OpenAiCompatible(model) => model.stream(request).await, + } + } +} diff --git a/crates/kuncode-core/src/providers/deepseek/protocol.rs b/crates/kuncode-core/src/providers/deepseek/protocol.rs index c97df9d..f38a071 100644 --- a/crates/kuncode-core/src/providers/deepseek/protocol.rs +++ b/crates/kuncode-core/src/providers/deepseek/protocol.rs @@ -51,6 +51,7 @@ pub enum Message { /// Assistant output: visible text plus optional tool calls and reasoning. Assistant { /// Visible assistant text. + #[serde(default, deserialize_with = "json_utils::null_or_default")] content: String, /// Optional speaker name accepted by OpenAI-compatible APIs. #[serde(skip_serializing_if = "Option::is_none")] @@ -207,6 +208,7 @@ pub struct ToolCall { // Position within parallel calls. DeepSeek includes it in responses and // streaming chunks; request replay uses array order instead, so outbound // conversion fills 0 and inbound projection does not read it. + #[serde(default)] pub index: usize, /// Tool-call kind; currently always [`ToolType::Function`]. pub r#type: ToolType, @@ -377,11 +379,13 @@ pub struct DeepSeekCompletionResponse { pub created: u64, /// Model id that served the request. pub model: String, - /// Provider backend fingerprint. - pub system_fingerprint: String, + /// Provider backend fingerprint; OpenAI-compatible endpoints may omit it. + #[serde(default)] + pub system_fingerprint: Option, /// OpenAI-compatible object type. pub object: String, - /// Token accounting for the call. + /// Token accounting for the call; some compatible endpoints omit it. + #[serde(default)] pub usage: Usage, } diff --git a/crates/kuncode-core/src/providers/openai_compatible.rs b/crates/kuncode-core/src/providers/openai_compatible.rs new file mode 100644 index 0000000..d2e1d35 --- /dev/null +++ b/crates/kuncode-core/src/providers/openai_compatible.rs @@ -0,0 +1,281 @@ +//! OpenAI-compatible `/chat/completions` provider with configurable endpoint. + +use std::time::Duration; + +use serde_json::Value; +use thiserror::Error; + +use crate::{ + completion::{ + CompletionError, CompletionModel, CompletionRequest, CompletionResponse, ReasoningEffort, + }, + json_utils, + providers::deepseek::protocol::{DeepSeekCompletionRequest, DeepSeekCompletionResponse}, +}; + +const DEFAULT_BASE_URL: &str = "https://api.openai.com/v1"; +const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); +const READ_TIMEOUT: Duration = Duration::from_secs(360); +const REQUEST_TIMEOUT: Duration = Duration::from_secs(360); + +/// Errors produced while constructing an OpenAI-compatible client. +#[derive(Debug, Error)] +pub enum Error { + /// The endpoint is empty and cannot be normalized. + #[error("OpenAI-compatible base URL must not be blank")] + BlankBaseUrl, + /// The underlying HTTP client could not be built. + #[error("HTTP client error: {0}")] + Client(#[from] reqwest::Error), +} + +/// Authenticated client for OpenAI and compatible Chat Completions endpoints. +#[derive(Clone)] +pub struct OpenAiCompatibleClient { + http_client: reqwest::Client, + api_key: String, + endpoint: String, +} + +impl OpenAiCompatibleClient { + /// Builds a client, appending `/chat/completions` when `base_url` names a + /// service root such as `https://api.openai.com/v1`. + /// + /// An empty API key is accepted for local endpoints that do not authenticate. + /// + /// # Errors + /// Returns [`enum@Error`] when the base URL is blank or the HTTP client fails. + pub fn new(api_key: impl Into, base_url: impl AsRef) -> Result { + let endpoint = completion_endpoint(base_url.as_ref())?; + let http_client = reqwest::Client::builder() + .read_timeout(READ_TIMEOUT) + .connect_timeout(CONNECT_TIMEOUT) + .build()?; + Ok(Self { + http_client, + api_key: api_key.into(), + endpoint, + }) + } + + /// Builds a client for the official OpenAI endpoint. + /// + /// # Errors + /// Returns [`enum@Error`] when the HTTP client cannot be built. + pub fn openai(api_key: impl Into) -> Result { + Self::new(api_key, DEFAULT_BASE_URL) + } + + fn post(&self) -> reqwest::RequestBuilder { + let request = self.http_client.post(&self.endpoint); + if self.api_key.is_empty() { + request + } else { + request.bearer_auth(&self.api_key) + } + } +} + +fn completion_endpoint(base_url: &str) -> Result { + let base_url = base_url.trim().trim_end_matches('/'); + if base_url.is_empty() { + return Err(Error::BlankBaseUrl); + } + if base_url.ends_with("/chat/completions") { + Ok(base_url.to_string()) + } else { + Ok(format!("{base_url}/chat/completions")) + } +} + +/// Completion model for OpenAI-compatible Chat Completions APIs. +#[derive(Clone)] +pub struct OpenAiCompatibleCompletionModel { + client: OpenAiCompatibleClient, + model: String, +} + +impl CompletionModel for OpenAiCompatibleCompletionModel { + type Response = Value; + type Client = OpenAiCompatibleClient; + + fn make(client: &Self::Client, model: impl Into) -> Self { + Self { + client: client.clone(), + model: model.into(), + } + } + + async fn completion( + &self, + request: CompletionRequest, + ) -> Result, CompletionError> { + let body = request_body(request, &self.model, false)?; + let response = self + .client + .post() + .timeout(REQUEST_TIMEOUT) + .json(&body) + .send() + .await?; + let status = response.status(); + if !status.is_success() { + return Err(CompletionError::ApiError { + status: status.as_u16(), + message: response.text().await.unwrap_or_default(), + }); + } + + let raw: Value = serde_json::from_slice(&response.bytes().await?)?; + normalize_response(raw) + } + + async fn stream( + &self, + request: CompletionRequest, + ) -> Result { + let body = request_body(request, &self.model, true)?; + let response = self.client.post().json(&body).send().await?; + let status = response.status(); + if !status.is_success() { + return Err(CompletionError::ApiError { + status: status.as_u16(), + message: response.text().await.unwrap_or_default(), + }); + } + Ok(crate::providers::deepseek::protocol::streaming::stream_events(response)) + } +} + +fn request_body( + mut request: CompletionRequest, + model: &str, + streaming: bool, +) -> Result { + request.model.get_or_insert_with(|| model.to_string()); + let extra = request.additional_params.take(); + let reasoning = request.reasoning.take(); + let request = DeepSeekCompletionRequest::try_from(request)?; + let mut body = if streaming { + serde_json::to_value(request.into_streaming())? + } else { + serde_json::to_value(request)? + }; + + // The shared DTO carries DeepSeek's `thinking` object. OpenAI-compatible + // endpoints use the flat `reasoning_effort` field instead. + if let Value::Object(fields) = &mut body { + fields.remove("thinking"); + fields.remove("reasoning_effort"); + if let Some(effort) = openai_reasoning_effort(reasoning) { + fields.insert( + "reasoning_effort".to_string(), + Value::String(effort.to_string()), + ); + } + } + + match extra { + Some(extra) => Ok(json_utils::merge(body, extra)), + None => Ok(body), + } +} + +fn normalize_response(raw: Value) -> Result, CompletionError> { + let provider: DeepSeekCompletionResponse = serde_json::from_value(raw.clone())?; + let normalized: CompletionResponse = provider.try_into()?; + Ok(CompletionResponse { + choice: normalized.choice, + usage: normalized.usage, + raw_response: raw, + message_id: normalized.message_id, + }) +} + +fn openai_reasoning_effort(effort: Option) -> Option<&'static str> { + match effort { + None | Some(ReasoningEffort::Off) => None, + Some(ReasoningEffort::Minimal) => Some("minimal"), + Some(ReasoningEffort::Low) => Some("low"), + Some(ReasoningEffort::Medium) => Some("medium"), + Some(ReasoningEffort::High) => Some("high"), + Some(ReasoningEffort::Xhigh) => Some("xhigh"), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::completion::{AssistantContent, CompletionRequestBuilder, Message}; + + #[test] + fn appends_chat_completions_to_service_root() { + assert_eq!( + completion_endpoint("https://api.openai.com/v1/").expect("valid endpoint"), + "https://api.openai.com/v1/chat/completions" + ); + } + + #[test] + fn preserves_full_chat_completions_endpoint() { + assert_eq!( + completion_endpoint("http://localhost:8000/v1/chat/completions") + .expect("valid endpoint"), + "http://localhost:8000/v1/chat/completions" + ); + } + + #[test] + fn request_uses_openai_reasoning_field_without_deepseek_thinking() { + let request = CompletionRequestBuilder::new(Message::user("test")) + .reasoning(Some(ReasoningEffort::Low)) + .build(); + let body = request_body(request, "gpt-test", true).expect("request body"); + + assert_eq!(body["reasoning_effort"], "low"); + assert!(body.get("thinking").is_none()); + assert_eq!(body["stream_options"]["include_usage"], true); + } + + #[test] + fn accepts_openai_tool_call_with_null_content_and_no_fingerprint() { + let response = normalize_response(serde_json::json!({ + "id": "chatcmpl-test", + "choices": [{ + "finish_reason": "tool_calls", + "index": 0, + "message": { + "role": "assistant", + "content": null, + "tool_calls": [{ + "id": "call-1", + "type": "function", + "function": { + "name": "bash", + "arguments": "{\"cmd\":\"pwd\"}" + } + }] + }, + "logprobs": null + }], + "created": 1, + "model": "gpt-test", + "object": "chat.completion", + "usage": { + "prompt_tokens": 10, + "completion_tokens": 4, + "total_tokens": 14, + "prompt_tokens_details": { "cached_tokens": 3 } + } + })) + .expect("OpenAI response normalizes"); + + assert_eq!(response.usage.cached_input_tokens, 3); + assert!(matches!( + response.choice.first(), + AssistantContent::ToolCall(call) + if call.function.name == "bash" + && call.function.arguments.get("cmd").and_then(Value::as_str) == Some("pwd") + )); + } +} From 6e6b49407bf34cb72ebcba5a0fe0257c05670d06 Mon Sep 17 00:00:00 2001 From: wenzr <282277167@qq.com> Date: Tue, 21 Jul 2026 23:29:20 +0800 Subject: [PATCH 2/4] =?UTF-8?q?Provider:=20=E6=94=B9=E8=BF=9B=E9=9D=9ESSE?= =?UTF-8?q?=E5=93=8D=E5=BA=94=E9=94=99=E8=AF=AF=E6=8F=90=E7=A4=BA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../src/providers/openai_compatible.rs | 44 +++++++++++++++++++ 1 file changed, 44 insertions(+) diff --git a/crates/kuncode-core/src/providers/openai_compatible.rs b/crates/kuncode-core/src/providers/openai_compatible.rs index d2e1d35..f958d49 100644 --- a/crates/kuncode-core/src/providers/openai_compatible.rs +++ b/crates/kuncode-core/src/providers/openai_compatible.rs @@ -2,6 +2,7 @@ use std::time::Duration; +use reqwest::header::CONTENT_TYPE; use serde_json::Value; use thiserror::Error; @@ -143,10 +144,30 @@ impl CompletionModel for OpenAiCompatibleCompletionModel { message: response.text().await.unwrap_or_default(), }); } + validate_stream_content_type( + response + .headers() + .get(CONTENT_TYPE) + .and_then(|value| value.to_str().ok()), + )?; Ok(crate::providers::deepseek::protocol::streaming::stream_events(response)) } } +fn validate_stream_content_type(content_type: Option<&str>) -> Result<(), CompletionError> { + let Some(content_type) = content_type else { + return Ok(()); + }; + let media_type = content_type.split(';').next().unwrap_or_default().trim(); + if media_type.eq_ignore_ascii_case("text/event-stream") { + return Ok(()); + } + + Err(CompletionError::ResponseError(format!( + "expected an SSE response with content type `text/event-stream`, but the provider returned `{media_type}`; verify `baseUrl` points to the API root (usually ending in `/v1`) or the full `/chat/completions` endpoint" + ))) +} + fn request_body( mut request: CompletionRequest, model: &str, @@ -225,6 +246,29 @@ mod tests { ); } + #[test] + fn accepts_sse_content_type_with_parameters() { + validate_stream_content_type(Some("text/event-stream; charset=utf-8")) + .expect("SSE content type should be accepted"); + } + + #[test] + fn accepts_missing_content_type_for_compatible_gateways() { + validate_stream_content_type(None) + .expect("a missing content type should fall through to the SSE decoder"); + } + + #[test] + fn rejects_html_with_base_url_guidance() { + let error = validate_stream_content_type(Some("text/html; charset=utf-8")) + .expect_err("HTML is not an SSE response") + .to_string(); + + assert!(error.contains("`text/html`")); + assert!(error.contains("`baseUrl`")); + assert!(error.contains("/v1")); + } + #[test] fn request_uses_openai_reasoning_field_without_deepseek_thinking() { let request = CompletionRequestBuilder::new(Message::user("test")) From b9d1e3dbcce0ccfacb866a47051da41b0c309f78 Mon Sep 17 00:00:00 2001 From: wenzr <282277167@qq.com> Date: Wed, 22 Jul 2026 19:58:36 +0800 Subject: [PATCH 3/4] =?UTF-8?q?Provider:=20=E6=94=B6=E7=AA=84=E5=AE=98?= =?UTF-8?q?=E6=96=B9=20OpenAI=20=E5=8D=8F=E8=AE=AE=E8=BE=B9=E7=95=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 36 +- .../src/compaction/artifact/spill/runtime.rs | 4 +- .../src/compaction/protocol/grouping.rs | 4 +- .../src/compaction/slimming/marker.rs | 1 + crates/kuncode-agent/src/runner/iteration.rs | 4 + crates/kuncode-agent/src/runner/request.rs | 1 + crates/kuncode-agent/src/session_store/dto.rs | 3 + .../src/session_store/dto/conversion.rs | 4 + crates/kuncode-cli/src/runtime.rs | 27 +- crates/kuncode-cli/src/settings.rs | 88 +--- crates/kuncode-core/src/completion.rs | 2 +- crates/kuncode-core/src/completion/message.rs | 23 + .../kuncode-core/src/completion/streaming.rs | 2 + crates/kuncode-core/src/providers.rs | 6 +- crates/kuncode-core/src/providers/any_chat.rs | 14 +- .../src/providers/chat_completions.rs | 3 + .../streaming.rs | 78 ++- crates/kuncode-core/src/providers/deepseek.rs | 6 +- .../src/providers/deepseek/protocol.rs | 5 +- crates/kuncode-core/src/providers/openai.rs | 175 +++++++ .../src/providers/openai/protocol.rs | 472 ++++++++++++++++++ .../src/providers/openai_compatible.rs | 325 ------------ 22 files changed, 792 insertions(+), 491 deletions(-) create mode 100644 crates/kuncode-core/src/providers/chat_completions.rs rename crates/kuncode-core/src/providers/{deepseek/protocol => chat_completions}/streaming.rs (89%) create mode 100644 crates/kuncode-core/src/providers/openai.rs create mode 100644 crates/kuncode-core/src/providers/openai/protocol.rs delete mode 100644 crates/kuncode-core/src/providers/openai_compatible.rs diff --git a/README.md b/README.md index 503606c..547a0b4 100644 --- a/README.md +++ b/README.md @@ -4,8 +4,8 @@ [`learn-claude-code`](https://github.com/shareAI-lab/learn-claude-code) 的 Harness Engineering 思路:模型负责判断下一步做什么,Harness 负责提供工具、上下文、权限边界、持久化和用户界面。 -当前版本为 `0.1.0`,默认使用 DeepSeek,也支持 OpenAI Chat Completions -兼容协议,提供一次性命令行执行和交互式 TUI 两种使用方式。 +当前版本为 `0.1.0`,默认使用 DeepSeek,也支持 OpenAI Chat Completions, +提供一次性命令行执行和交互式 TUI 两种使用方式。 ## 工作区结构 @@ -17,14 +17,14 @@ kuncode-cli ──▶ kuncode-agent ──▶ kuncode-core ──▶ LLM API └─ 参数、配置、审批、一次性输出与 TUI ``` -- `kuncode-core`:Provider-neutral 的消息与 Completion 抽象,以及 DeepSeek、OpenAI-compatible Provider。 +- `kuncode-core`:Provider-neutral 的消息与 Completion 抽象,以及 DeepSeek、OpenAI Provider。 - `kuncode-agent`:Agent 运行时、工具调度、权限、Hook、Todo、会话持久化和上下文压缩。 - `kuncode-cli`:命令行参数、项目配置、终端审批、普通输出和交互式 TUI。 ## 环境要求 - Rust stable,项目使用 Rust 2024 edition。 -- 所选模型服务的 API Key;本地无鉴权服务可不配置。 +- DeepSeek 或 OpenAI API Key。 - 支持 ANSI 终端;交互模式需要 stdin 和 stdout 都连接到真实终端。 ## 快速开始 @@ -46,9 +46,8 @@ DEEPSEEK_API_KEY=your-api-key ```json { "model": { - "provider": "openai-compatible", + "provider": "openai", "name": "your-openai-model", - "apiKeyEnv": "OPENAI_API_KEY", "maxTokens": 16384 } } @@ -60,23 +59,6 @@ DEEPSEEK_API_KEY=your-api-key export OPENAI_API_KEY="your-api-key" ``` -Kimi、Qwen、智谱、vLLM 等提供 OpenAI Chat Completions 兼容协议的服务,增加 -`baseUrl` 即可。它既可以是服务根地址,也可以是完整 endpoint: - -```json -{ - "model": { - "provider": "openai-compatible", - "name": "your-model", - "baseUrl": "http://localhost:8000/v1", - "apiKeyEnv": "", - "maxTokens": 8192 - } -} -``` - -`apiKeyEnv` 为空表示不发送 `Authorization`,适合无鉴权的本地服务。 - 一次性执行任务: ```bash @@ -103,7 +85,7 @@ cargo build --release -p kuncode-cli ```json { "permissions": { - "allow": ["Read(**)", "Bash(cargo *)"], + "allow": ["Read", "Bash(cargo *)"], "ask": ["Edit(.env)"], "deny": ["Bash(curl *)"], "defaultMode": "default" @@ -111,7 +93,6 @@ cargo build --release -p kuncode-cli "model": { "provider": "deepseek", "name": "deepseek-v4-pro", - "apiKeyEnv": "DEEPSEEK_API_KEY", "maxTokens": 65536 }, "agent": { @@ -130,9 +111,8 @@ cargo build --release -p kuncode-cli 补充说明: - `KUNCODE_MODEL` 可以覆盖配置文件中的模型名称;`DEEPSEEK_MODEL` 作为兼容别名保留。 -- `model.provider` 支持 `deepseek` 和 `openai-compatible`。 -- `openai-compatible` 未设置 `baseUrl` 时默认使用 `https://api.openai.com/v1`。 -- `baseUrl` 会自动补全 `/chat/completions`;完整 endpoint 不会重复追加。 +- `model.provider` 支持 `deepseek` 和 `openai`;两者分别使用固定官方 endpoint, + 并读取 `DEEPSEEK_API_KEY` 或 `OPENAI_API_KEY`。 - 内置模型配置包括 `deepseek-v4-pro` 和 `deepseek-v4-flash`。 - 非内置模型启用上下文压缩时,需要显式设置 `compaction.contextLimit`。 - `compaction.mode` 支持 `disabled`、`shadow` 和 `enabled`,默认是 `disabled`。 diff --git a/crates/kuncode-agent/src/compaction/artifact/spill/runtime.rs b/crates/kuncode-agent/src/compaction/artifact/spill/runtime.rs index fd1214d..fe4337c 100644 --- a/crates/kuncode-agent/src/compaction/artifact/spill/runtime.rs +++ b/crates/kuncode-agent/src/compaction/artifact/spill/runtime.rs @@ -255,7 +255,9 @@ fn call_names(message: &Message) -> BTreeMap { if let Message::Assistant { content, .. } = message { for call in content.iter().filter_map(|block| match block { AssistantContent::ToolCall(call) => Some(call), - AssistantContent::Text(_) | AssistantContent::Reasoning(_) => None, + AssistantContent::Text(_) + | AssistantContent::Reasoning(_) + | AssistantContent::Refusal(_) => None, }) { names.insert(call.id.clone(), call.function.name.clone()); } diff --git a/crates/kuncode-agent/src/compaction/protocol/grouping.rs b/crates/kuncode-agent/src/compaction/protocol/grouping.rs index 7dd4d97..dbd3929 100644 --- a/crates/kuncode-agent/src/compaction/protocol/grouping.rs +++ b/crates/kuncode-agent/src/compaction/protocol/grouping.rs @@ -127,7 +127,9 @@ pub fn group_messages(messages: &[Message]) -> Result, Protoc .iter() .filter_map(|block| match block { AssistantContent::ToolCall(call) => Some(call), - AssistantContent::Text(_) | AssistantContent::Reasoning(_) => None, + AssistantContent::Text(_) + | AssistantContent::Reasoning(_) + | AssistantContent::Refusal(_) => None, }) .collect::>(); if calls.is_empty() { diff --git a/crates/kuncode-agent/src/compaction/slimming/marker.rs b/crates/kuncode-agent/src/compaction/slimming/marker.rs index 0ad7770..2d40be1 100644 --- a/crates/kuncode-agent/src/compaction/slimming/marker.rs +++ b/crates/kuncode-agent/src/compaction/slimming/marker.rs @@ -139,6 +139,7 @@ fn assistant_call<'a>( AssistantContent::ToolCall(call) if call.id == result_id => Some(call), AssistantContent::Text(_) | AssistantContent::Reasoning(_) + | AssistantContent::Refusal(_) | AssistantContent::ToolCall(_) => None, }) } diff --git a/crates/kuncode-agent/src/runner/iteration.rs b/crates/kuncode-agent/src/runner/iteration.rs index ab89bd9..7841ace 100644 --- a/crates/kuncode-agent/src/runner/iteration.rs +++ b/crates/kuncode-agent/src/runner/iteration.rs @@ -144,6 +144,9 @@ where StreamEvent::ReasoningDelta(text) => { self.emit(session, Some(iteration), EventKind::ReasoningDelta { text }); } + StreamEvent::RefusalDelta(text) => { + self.emit(session, Some(iteration), EventKind::TextDelta { text }); + } // The "calling X" hint is surfaced by `ToolStart` after the turn // completes and the call is gated; ignore the earlier signal. StreamEvent::ToolCallStart { .. } => {} @@ -188,6 +191,7 @@ fn stream_event_kind(event: &StreamEvent) -> &'static str { match event { StreamEvent::TextDelta(_) => "text_delta", StreamEvent::ReasoningDelta(_) => "reasoning_delta", + StreamEvent::RefusalDelta(_) => "refusal_delta", StreamEvent::ToolCallStart { .. } => "tool_call_start", StreamEvent::Completed { .. } => "completed", } diff --git a/crates/kuncode-agent/src/runner/request.rs b/crates/kuncode-agent/src/runner/request.rs index 93fc5d5..eaa611b 100644 --- a/crates/kuncode-agent/src/runner/request.rs +++ b/crates/kuncode-agent/src/runner/request.rs @@ -170,6 +170,7 @@ pub(super) fn assistant_text(content: &NonEmptyVec) -> String .iter() .filter_map(|content| match content { AssistantContent::Text(text) => Some(text.text_ref()), + AssistantContent::Refusal(refusal) => Some(refusal.text_ref()), _ => None, }) .collect::>() diff --git a/crates/kuncode-agent/src/session_store/dto.rs b/crates/kuncode-agent/src/session_store/dto.rs index 1b77c81..2d57809 100644 --- a/crates/kuncode-agent/src/session_store/dto.rs +++ b/crates/kuncode-agent/src/session_store/dto.rs @@ -72,6 +72,9 @@ enum StoredAssistantContent { Text { text: String, }, + Refusal { + text: String, + }, ToolCall { id: String, #[serde(skip_serializing_if = "Option::is_none")] diff --git a/crates/kuncode-agent/src/session_store/dto/conversion.rs b/crates/kuncode-agent/src/session_store/dto/conversion.rs index 1a430d3..7069c8d 100644 --- a/crates/kuncode-agent/src/session_store/dto/conversion.rs +++ b/crates/kuncode-agent/src/session_store/dto/conversion.rs @@ -129,6 +129,9 @@ impl StoredAssistantContent { AssistantContent::Text(text) => Self::Text { text: text.text_ref().to_string(), }, + AssistantContent::Refusal(refusal) => Self::Refusal { + text: refusal.text_ref().to_string(), + }, AssistantContent::ToolCall(call) => Self::ToolCall { id: call.id.clone(), call_id: call.call_id.clone(), @@ -151,6 +154,7 @@ impl StoredAssistantContent { fn into_content(self) -> Result { match self { Self::Text { text } => Ok(AssistantContent::Text(Text::from(text))), + Self::Refusal { text } => Ok(AssistantContent::refusal(text)), Self::ToolCall { id, call_id, diff --git a/crates/kuncode-cli/src/runtime.rs b/crates/kuncode-cli/src/runtime.rs index 8f32659..605616b 100644 --- a/crates/kuncode-cli/src/runtime.rs +++ b/crates/kuncode-cli/src/runtime.rs @@ -30,7 +30,7 @@ use kuncode_core::completion::{CompletionModel, RetryModel, RetryPolicy}; use kuncode_core::providers::{ any_chat::{AnyChatClient, AnyChatCompletionModel}, deepseek::DeepSeekClient, - openai_compatible::OpenAiCompatibleClient, + openai::OpenAiClient, }; use crate::config::{PermissionFlags, resolve_permissions}; @@ -81,7 +81,7 @@ impl CliRuntime> { /// Fails if the current directory is not a usable workspace, the project /// settings or resolved permissions are invalid, active compaction cannot /// be bound to the selected model, or the provider client cannot be built - /// from its endpoint and environment. Failure to open the optional session + /// from its fixed credential environment. Failure to open the optional session /// store is retained as degraded persistence state rather than failing assembly. pub async fn assemble(cli: &Cli) -> Result> { let workspace = Workspace::from_current_dir().await?; @@ -190,28 +190,9 @@ impl CliRuntime> { } fn provider_client(project: &ProjectSettings) -> Result> { - let api_key = if project.api_key_env.trim().is_empty() { - String::new() - } else { - std::env::var(&project.api_key_env).map_err(|error| { - std::io::Error::new( - std::io::ErrorKind::NotFound, - format!( - "provider API key environment variable `{}` is unavailable: {error}", - project.api_key_env - ), - ) - })? - }; match project.provider { - ProviderKind::DeepSeek => Ok(AnyChatClient::DeepSeek(DeepSeekClient::new(api_key)?)), - ProviderKind::OpenAiCompatible => { - let client = match project.base_url.as_deref() { - Some(base_url) => OpenAiCompatibleClient::new(api_key, base_url)?, - None => OpenAiCompatibleClient::openai(api_key)?, - }; - Ok(AnyChatClient::OpenAiCompatible(client)) - } + ProviderKind::DeepSeek => Ok(AnyChatClient::DeepSeek(DeepSeekClient::from_env()?)), + ProviderKind::OpenAi => Ok(AnyChatClient::OpenAi(OpenAiClient::from_env()?)), } } diff --git a/crates/kuncode-cli/src/settings.rs b/crates/kuncode-cli/src/settings.rs index 657dcb1..be13563 100644 --- a/crates/kuncode-cli/src/settings.rs +++ b/crates/kuncode-cli/src/settings.rs @@ -78,8 +78,6 @@ struct PermissionsSection { struct ModelSection { provider: ProviderKind, name: String, - base_url: Option, - api_key_env: Option, max_tokens: Option, } @@ -88,8 +86,6 @@ impl Default for ModelSection { Self { provider: ProviderKind::DeepSeek, name: DEEPSEEK_V4_PRO_MODEL_ID.to_string(), - base_url: None, - api_key_env: None, max_tokens: None, } } @@ -102,9 +98,9 @@ pub(crate) enum ProviderKind { #[default] #[serde(rename = "deepseek")] DeepSeek, - /// OpenAI Chat Completions and compatible services such as vLLM or Kimi. - #[serde(rename = "openai-compatible")] - OpenAiCompatible, + /// Official OpenAI Chat Completions protocol and endpoint. + #[serde(rename = "openai")] + OpenAi, } #[derive(Debug, Deserialize)] @@ -186,10 +182,6 @@ pub struct ProjectSettings { pub(crate) trust: ProjectTrust, /// Effective provider protocol. pub(crate) provider: ProviderKind, - /// Custom OpenAI-compatible service root or full completion endpoint. - pub(crate) base_url: Option, - /// Environment variable holding the provider API key. - pub(crate) api_key_env: String, /// Effective model identifier after file and environment precedence. pub(crate) model_name: String, /// Effective provider output budget for an ordinary turn. @@ -213,8 +205,6 @@ impl Default for ProjectSettings { default_mode: None, trust: ProjectTrust::Untrusted, provider: ProviderKind::DeepSeek, - base_url: None, - api_key_env: "DEEPSEEK_API_KEY".to_string(), model_name: DEEPSEEK_V4_PRO_MODEL_ID.to_string(), max_tokens, max_iterations: DEFAULT_MAX_ITERATIONS, @@ -310,33 +300,6 @@ fn resolve_settings( validate_model_max_tokens(max_tokens, profile)?; validate_agent(&file.agent)?; validate_log_level(&file.logging.level)?; - let base_url = match file.model.base_url { - Some(url) if url.trim().is_empty() => { - return Err(SettingsError::Model( - "baseUrl must not be blank when provided".to_string(), - )); - } - Some(url) => Some(url.trim().to_string()), - None => None, - }; - if file.model.provider == ProviderKind::DeepSeek && base_url.is_some() { - return Err(SettingsError::Model( - "baseUrl requires provider `openai-compatible`".to_string(), - )); - } - let api_key_env = file - .model - .api_key_env - .unwrap_or_else(|| match file.model.provider { - ProviderKind::DeepSeek => "DEEPSEEK_API_KEY".to_string(), - ProviderKind::OpenAiCompatible => "OPENAI_API_KEY".to_string(), - }); - if file.model.provider == ProviderKind::DeepSeek && api_key_env.trim().is_empty() { - return Err(SettingsError::Model( - "apiKeyEnv must not be blank for the DeepSeek provider".to_string(), - )); - } - let canonical_root = std::fs::canonicalize(root).map_err(|error| SettingsError::Workspace(error.to_string()))?; let canonical_root = CanonicalPath::from_absolute(&canonical_root) @@ -361,8 +324,6 @@ fn resolve_settings( default_mode, trust, provider: file.model.provider, - base_url, - api_key_env, model_name, max_tokens, max_iterations: file.agent.max_iterations, @@ -619,8 +580,6 @@ mod tests { let _ = fs::remove_dir_all(&dir); assert_eq!(loaded.provider, ProviderKind::DeepSeek); - assert_eq!(loaded.base_url, None); - assert_eq!(loaded.api_key_env, "DEEPSEEK_API_KEY"); assert!(loaded.policy.expect("resolved policy").rules().is_empty()); assert!(loaded.default_mode.is_none()); assert!(loaded.compaction.is_none()); @@ -631,25 +590,18 @@ mod tests { } #[test] - fn loads_openai_compatible_provider_settings() { + fn loads_official_openai_provider_settings() { let loaded = load_json( - "openai-compatible", + "openai", r#"{ "model": { - "provider": "openai-compatible", + "provider": "openai", "name": "gpt-test", - "baseUrl": "https://api.example.com/v1/", - "apiKeyEnv": "EXAMPLE_API_KEY", "maxTokens": 8192 } }"#, ) .expect("loads"); - assert_eq!(loaded.provider, ProviderKind::OpenAiCompatible); - assert_eq!( - loaded.base_url.as_deref(), - Some("https://api.example.com/v1/") - ); - assert_eq!(loaded.api_key_env, "EXAMPLE_API_KEY"); + assert_eq!(loaded.provider, ProviderKind::OpenAi); assert_eq!(loaded.model_name, "gpt-test"); assert_eq!(loaded.max_tokens, 8_192); } @@ -665,32 +617,6 @@ mod tests { assert_eq!(loaded.provider, ProviderKind::DeepSeek); } - #[test] - fn openai_compatible_provider_allows_unauthenticated_local_endpoint() { - let loaded = load_json( - "openai-compatible-local", - r#"{ "model": { - "provider": "openai-compatible", - "name": "local-model", - "baseUrl": "http://localhost:8000/v1", - "apiKeyEnv": "" - } }"#, - ) - .expect("loads"); - - assert_eq!(loaded.api_key_env, ""); - } - - #[test] - fn deepseek_rejects_openai_only_base_url() { - let result = load_json( - "deepseek-base-url", - r#"{ "model": { "baseUrl": "https://example.com/v1" } }"#, - ); - - assert!(matches!(result, Err(SettingsError::Model(_)))); - } - #[test] fn logging_level_defaults_to_info() { let dir = std::env::temp_dir().join(format!("kuncode-log-absent-{}", std::process::id())); diff --git a/crates/kuncode-core/src/completion.rs b/crates/kuncode-core/src/completion.rs index 2a80ea7..75094d2 100644 --- a/crates/kuncode-core/src/completion.rs +++ b/crates/kuncode-core/src/completion.rs @@ -11,7 +11,7 @@ pub mod retry; pub mod streaming; pub use message::{ - AssistantContent, Message, Reasoning, ReasoningContent, Text, ToolCall, ToolChoice, + AssistantContent, Message, Reasoning, ReasoningContent, Refusal, Text, ToolCall, ToolChoice, ToolFunction, ToolResult, ToolResultContent, UserContent, }; diff --git a/crates/kuncode-core/src/completion/message.rs b/crates/kuncode-core/src/completion/message.rs index 6139725..82bbb17 100644 --- a/crates/kuncode-core/src/completion/message.rs +++ b/crates/kuncode-core/src/completion/message.rs @@ -150,6 +150,8 @@ pub enum AssistantContent { ToolCall(ToolCall), /// Reasoning/thinking content returned by reasoning-capable models. Reasoning(Reasoning), + /// Safety refusal returned instead of ordinary assistant text. + Refusal(Refusal), } impl AssistantContent { @@ -198,6 +200,26 @@ impl AssistantContent { pub fn reasoning(reasoning: impl AsRef) -> Self { Self::Reasoning(Reasoning::new(reasoning.as_ref())) } + + /// Preserves a provider refusal separately from ordinary assistant text. + pub fn refusal(refusal: impl Into) -> Self { + Self::Refusal(Refusal { + refusal: refusal.into(), + }) + } +} + +/// A safety refusal emitted instead of ordinary assistant content. +#[derive(Clone, Debug, Deserialize, Serialize, PartialEq)] +pub struct Refusal { + refusal: String, +} + +impl Refusal { + /// Returns the refusal text. + pub fn text_ref(&self) -> &str { + &self.refusal + } } /// A plain-text content block. @@ -428,6 +450,7 @@ mod tests { Reasoning::summaries(vec!["s1".into(), "s2".into()]) .with_id("r_1".to_string()), ), + AssistantContent::refusal("cannot comply"), ], ), }, diff --git a/crates/kuncode-core/src/completion/streaming.rs b/crates/kuncode-core/src/completion/streaming.rs index 706e0bd..1a5d5a8 100644 --- a/crates/kuncode-core/src/completion/streaming.rs +++ b/crates/kuncode-core/src/completion/streaming.rs @@ -55,6 +55,8 @@ pub enum StreamEvent { /// A chunk of reasoning/thinking text, kept separate from the answer /// (e.g. DeepSeek's `reasoning_content`). Render in a distinct channel. ReasoningDelta(String), + /// A chunk of a safety refusal, kept distinct from ordinary answer text. + RefusalDelta(String), /// A tool call has started: its `id` and `name` are known before the /// arguments finish streaming. Useful for an immediate "calling X" hint; /// the complete call (with assembled arguments) arrives in diff --git a/crates/kuncode-core/src/providers.rs b/crates/kuncode-core/src/providers.rs index c722cc2..5f50135 100644 --- a/crates/kuncode-core/src/providers.rs +++ b/crates/kuncode-core/src/providers.rs @@ -2,9 +2,9 @@ //! //! Each provider owns the mapping from the provider-agnostic //! [`crate::completion`] types to that provider's HTTP API. DeepSeek remains -//! the default, while [`openai_compatible`] supports OpenAI and compatible -//! `/chat/completions` endpoints. +//! the default, while [`openai`] implements the official OpenAI protocol. pub mod any_chat; +pub(crate) mod chat_completions; pub mod deepseek; -pub mod openai_compatible; +pub mod openai; diff --git a/crates/kuncode-core/src/providers/any_chat.rs b/crates/kuncode-core/src/providers/any_chat.rs index a186485..306c22f 100644 --- a/crates/kuncode-core/src/providers/any_chat.rs +++ b/crates/kuncode-core/src/providers/any_chat.rs @@ -8,7 +8,7 @@ use crate::{ }, providers::{ deepseek::{DeepSeekClient, DeepSeekCompletionModel}, - openai_compatible::{OpenAiCompatibleClient, OpenAiCompatibleCompletionModel}, + openai::{OpenAiClient, OpenAiCompletionModel}, }, }; @@ -18,7 +18,7 @@ pub enum AnyChatClient { /// Native DeepSeek protocol behavior. DeepSeek(DeepSeekClient), /// OpenAI-compatible Chat Completions behavior. - OpenAiCompatible(OpenAiCompatibleClient), + OpenAi(OpenAiClient), } /// Model handle that keeps the agent runtime independent of provider choice. @@ -27,7 +27,7 @@ pub enum AnyChatCompletionModel { /// Native DeepSeek model. DeepSeek(DeepSeekCompletionModel), /// OpenAI-compatible model. - OpenAiCompatible(OpenAiCompatibleCompletionModel), + OpenAi(OpenAiCompletionModel), } impl CompletionModel for AnyChatCompletionModel { @@ -40,8 +40,8 @@ impl CompletionModel for AnyChatCompletionModel { AnyChatClient::DeepSeek(client) => { Self::DeepSeek(DeepSeekCompletionModel::make(client, model)) } - AnyChatClient::OpenAiCompatible(client) => { - Self::OpenAiCompatible(OpenAiCompatibleCompletionModel::make(client, model)) + AnyChatClient::OpenAi(client) => { + Self::OpenAi(OpenAiCompletionModel::make(client, model)) } } } @@ -61,7 +61,7 @@ impl CompletionModel for AnyChatCompletionModel { message_id: response.message_id, }) } - Self::OpenAiCompatible(model) => model.completion(request).await, + Self::OpenAi(model) => model.completion(request).await, } } @@ -71,7 +71,7 @@ impl CompletionModel for AnyChatCompletionModel { ) -> Result { match self { Self::DeepSeek(model) => model.stream(request).await, - Self::OpenAiCompatible(model) => model.stream(request).await, + Self::OpenAi(model) => model.stream(request).await, } } } diff --git a/crates/kuncode-core/src/providers/chat_completions.rs b/crates/kuncode-core/src/providers/chat_completions.rs new file mode 100644 index 0000000..c3d92e0 --- /dev/null +++ b/crates/kuncode-core/src/providers/chat_completions.rs @@ -0,0 +1,3 @@ +//! Shared transport primitives for Chat Completions protocol providers. + +pub(crate) mod streaming; diff --git a/crates/kuncode-core/src/providers/deepseek/protocol/streaming.rs b/crates/kuncode-core/src/providers/chat_completions/streaming.rs similarity index 89% rename from crates/kuncode-core/src/providers/deepseek/protocol/streaming.rs rename to crates/kuncode-core/src/providers/chat_completions/streaming.rs index 40cbe30..3be0df8 100644 --- a/crates/kuncode-core/src/providers/deepseek/protocol/streaming.rs +++ b/crates/kuncode-core/src/providers/chat_completions/streaming.rs @@ -1,4 +1,4 @@ -//! DeepSeek Server-Sent Events (SSE) streaming: wire chunk DTOs, an incremental +//! Chat Completions Server-Sent Events (SSE) streaming: wire chunk DTOs, an incremental //! SSE frame decoder, and an assembler that folds chunks into [`StreamEvent`]s. //! //! Split into pure pieces — [`SseDecoder`] (bytes → `data:` payloads) and @@ -7,22 +7,21 @@ //! that drives a live [`reqwest::Response`] body through them. use async_stream::try_stream; -use serde::Deserialize; +use serde::{Deserialize, de::DeserializeOwned}; -use super::Usage; use crate::completion::{ - AssistantContent, CompletionError, CompletionStream, FinishReason, StreamEvent, + AssistantContent, CompletionError, CompletionStream, FinishReason, StreamEvent, Usage, }; use crate::non_empty_vec::NonEmptyVec; /// One `chat.completion.chunk` frame. The terminal usage-only frame carries an /// empty `choices`, so it defaults rather than failing to deserialize. #[derive(Debug, Deserialize)] -struct StreamChunk { +struct StreamChunk { #[serde(default)] choices: Vec, /// Present only on the final frame when `stream_options.include_usage` is set. - usage: Option, + usage: Option, /// Set when the endpoint reports a failure *mid-stream* as a data frame /// instead of via HTTP status; see [`StreamErrorBody`]. error: Option, @@ -54,6 +53,7 @@ struct ChunkChoice { struct ChunkDelta { content: Option, reasoning_content: Option, + refusal: Option, tool_calls: Option>, } @@ -139,6 +139,7 @@ struct PartialToolCall { struct StreamAssembler { text: String, reasoning: String, + refusal: String, tool_calls: Vec, finish_reason: Option, usage: Option, @@ -148,9 +149,12 @@ impl StreamAssembler { /// Folds one chunk in, returning the render deltas it produced (text / /// reasoning / tool-call-start). The terminal [`StreamEvent::Completed`] is /// produced separately by [`finish`](Self::finish). - fn ingest(&mut self, chunk: StreamChunk) -> Vec { - if chunk.usage.is_some() { - self.usage = chunk.usage; + fn ingest(&mut self, chunk: StreamChunk) -> Vec + where + U: Into, + { + if let Some(usage) = chunk.usage { + self.usage = Some(usage.into()); } let mut events = Vec::new(); for choice in chunk.choices { @@ -162,6 +166,10 @@ impl StreamAssembler { self.reasoning.push_str(&reasoning); events.push(StreamEvent::ReasoningDelta(reasoning)); } + if let Some(refusal) = choice.delta.refusal.filter(|s| !s.is_empty()) { + self.refusal.push_str(&refusal); + events.push(StreamEvent::RefusalDelta(refusal)); + } for delta in choice.delta.tool_calls.into_iter().flatten() { self.ingest_tool_call(delta, &mut events); } @@ -244,6 +252,9 @@ impl StreamAssembler { if !self.reasoning.is_empty() { content.push(AssistantContent::reasoning(self.reasoning)); } + if !self.refusal.is_empty() { + content.push(AssistantContent::refusal(self.refusal)); + } let content = NonEmptyVec::try_from(content).map_err(|err| { CompletionError::ResponseError(format!("stream produced no assistant content: {err}")) @@ -251,7 +262,7 @@ impl StreamAssembler { Ok(StreamEvent::Completed { content, - usage: self.usage.unwrap_or_default().into(), + usage: self.usage.unwrap_or_default(), finish_reason: map_finish_reason(self.finish_reason.as_deref()), }) } @@ -266,7 +277,7 @@ fn stream_error(error: StreamErrorBody) -> CompletionError { CompletionError::ResponseError(format!("provider reported a mid-stream error: {detail}")) } -/// Maps DeepSeek's stop-reason string onto the neutral [`FinishReason`]. A +/// Maps a Chat Completions stop-reason string onto the neutral [`FinishReason`]. A /// missing reason (stream ended without one) reads as a natural stop. fn map_finish_reason(reason: Option<&str>) -> FinishReason { match reason { @@ -286,7 +297,10 @@ fn map_finish_reason(reason: Option<&str>) -> FinishReason { /// [`StreamAssembler`], yielding render deltas as they arrive and a final /// [`StreamEvent::Completed`] once the body ends or `[DONE]` is seen. Dropping /// the returned stream closes the HTTP response and halts generation. -pub(crate) fn stream_events(mut response: reqwest::Response) -> CompletionStream { +pub(crate) fn stream_events(mut response: reqwest::Response) -> CompletionStream +where + U: DeserializeOwned + Into + Send + 'static, +{ Box::pin(try_stream! { let mut decoder = SseDecoder::new(); let mut assembler = StreamAssembler::default(); @@ -295,7 +309,7 @@ pub(crate) fn stream_events(mut response: reqwest::Response) -> CompletionStream match event { SseEvent::Done => break 'body, SseEvent::Data(payload) => { - let mut chunk: StreamChunk = serde_json::from_str(&payload)?; + let mut chunk: StreamChunk = serde_json::from_str(&payload)?; // A mid-stream error frame ends the stream with an error; // otherwise the partial answer would be assembled and // reported as a clean completion. @@ -316,6 +330,7 @@ pub(crate) fn stream_events(mut response: reqwest::Response) -> CompletionStream #[cfg(test)] mod tests { use super::*; + use crate::providers::deepseek::protocol::Usage as TestUsage; /// Runs `sse` through a fresh decoder + assembler, splitting the input into /// `chunk_size`-byte network chunks to exercise cross-chunk buffering. @@ -329,7 +344,7 @@ mod tests { match ev { SseEvent::Done => done = true, SseEvent::Data(payload) => { - let chunk: StreamChunk = + let chunk: StreamChunk = serde_json::from_str(&payload).expect("chunk json"); events.extend(assembler.ingest(chunk)); } @@ -430,6 +445,33 @@ data: [DONE] ); } + #[test] + fn refusal_streams_and_is_preserved_in_completed_content() { + let sse = "\ +data: {\"choices\":[{\"delta\":{\"refusal\":\"Cannot \"}}]} + +data: {\"choices\":[{\"delta\":{\"refusal\":\"comply\"},\"finish_reason\":\"stop\"}]} + +data: [DONE] + +"; + let events = run(sse, 4096); + let refusal: String = events + .iter() + .filter_map(|event| match event { + StreamEvent::RefusalDelta(text) => Some(text.clone()), + _ => None, + }) + .collect(); + assert_eq!(refusal, "Cannot comply"); + + let (content, _) = completed(events); + assert!(matches!( + content.first(), + AssistantContent::Refusal(value) if value.text_ref() == "Cannot comply" + )); + } + #[test] fn tool_call_arguments_assemble_across_fragments() { let sse = "\ @@ -480,7 +522,7 @@ data: [DONE] Some(SseEvent::Data(p)) => p, other => panic!("expected a data line, got {other:?}"), }; - assert!(serde_json::from_str::(payload).is_err()); + assert!(serde_json::from_str::>(payload).is_err()); } #[test] @@ -510,7 +552,7 @@ data: {\"error\":{\"message\":\"rate limited\",\"type\":\"server_error\"}} 'outer: for piece in sse.as_bytes().chunks(4096) { for ev in decoder.push(piece) { if let SseEvent::Data(payload) = ev { - let mut chunk: StreamChunk = + let mut chunk: StreamChunk = serde_json::from_str(&payload).expect("chunk json"); if let Some(err) = chunk.error.take() { error = Some(stream_error(err)); @@ -534,10 +576,10 @@ data: {\"error\":{\"message\":\"rate limited\",\"type\":\"server_error\"}} // The standard trio is required: a usage object missing one is malformed // and must fail the parse, not silently read as zero. let bad = r#"{"choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2}}"#; - assert!(serde_json::from_str::(bad).is_err()); + assert!(serde_json::from_str::>(bad).is_err()); // The DeepSeek cache extensions, by contrast, may be omitted. let ok = r#"{"choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}"#; - assert!(serde_json::from_str::(ok).is_ok()); + assert!(serde_json::from_str::>(ok).is_ok()); } } diff --git a/crates/kuncode-core/src/providers/deepseek.rs b/crates/kuncode-core/src/providers/deepseek.rs index 161b85f..30ef00b 100644 --- a/crates/kuncode-core/src/providers/deepseek.rs +++ b/crates/kuncode-core/src/providers/deepseek.rs @@ -205,7 +205,11 @@ impl CompletionModel for DeepSeekCompletionModel { }); } - Ok(protocol::streaming::stream_events(response)) + Ok( + crate::providers::chat_completions::streaming::stream_events::( + response, + ), + ) } } diff --git a/crates/kuncode-core/src/providers/deepseek/protocol.rs b/crates/kuncode-core/src/providers/deepseek/protocol.rs index f38a071..25bcc9d 100644 --- a/crates/kuncode-core/src/providers/deepseek/protocol.rs +++ b/crates/kuncode-core/src/providers/deepseek/protocol.rs @@ -19,8 +19,6 @@ use crate::{ non_empty_vec::NonEmptyVec, }; -pub(crate) mod streaming; - /// DeepSeek wire message serialized by `role`. /// /// This is flatter than the domain-side [`message::Message`]: `content` is a @@ -160,6 +158,9 @@ impl From for Vec { .as_str(), ); } + message::AssistantContent::Refusal(refusal) => { + text_content.push_str(refusal.text_ref()); + } } } diff --git a/crates/kuncode-core/src/providers/openai.rs b/crates/kuncode-core/src/providers/openai.rs new file mode 100644 index 0000000..3ce23d5 --- /dev/null +++ b/crates/kuncode-core/src/providers/openai.rs @@ -0,0 +1,175 @@ +//! Official OpenAI Chat Completions provider. + +use std::{env::VarError, time::Duration}; + +use reqwest::header::CONTENT_TYPE; +use serde_json::Value; +use thiserror::Error; + +use crate::{ + completion::{CompletionError, CompletionModel, CompletionRequest, CompletionResponse}, + json_utils, +}; + +use self::protocol::{OpenAiCompletionRequest, OpenAiCompletionResponse, Usage}; + +mod protocol; + +const OPENAI_COMPLETIONS_URL: &str = "https://api.openai.com/v1/chat/completions"; +const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); +const READ_TIMEOUT: Duration = Duration::from_secs(360); +const REQUEST_TIMEOUT: Duration = Duration::from_secs(360); + +/// Errors produced while constructing an OpenAI client. +#[derive(Debug, Error)] +pub enum Error { + /// The underlying HTTP client could not be built. + #[error("HTTP client error: {0}")] + Client(#[from] reqwest::Error), + /// `OPENAI_API_KEY` was missing or invalid Unicode. + #[error("environment variable `OPENAI_API_KEY` is not set or is invalid")] + EnvironmentVariable(#[source] VarError), +} + +/// Authenticated client for the official OpenAI API. +#[derive(Clone)] +pub struct OpenAiClient { + http_client: reqwest::Client, + api_key: String, +} + +impl OpenAiClient { + /// Builds a client for the fixed official OpenAI endpoint. + /// + /// # Errors + /// + /// Returns [`enum@Error`] when the HTTP client cannot be configured. + pub fn new(api_key: impl Into) -> Result { + let http_client = reqwest::Client::builder() + .read_timeout(READ_TIMEOUT) + .connect_timeout(CONNECT_TIMEOUT) + .build()?; + Ok(Self { + http_client, + api_key: api_key.into(), + }) + } + + /// Reads `OPENAI_API_KEY` and builds an official OpenAI client. + /// + /// # Errors + /// + /// Returns [`Error::EnvironmentVariable`] when the credential is unavailable, + /// or [`Error::Client`] when the HTTP client cannot be configured. + pub fn from_env() -> Result { + let api_key = std::env::var("OPENAI_API_KEY").map_err(Error::EnvironmentVariable)?; + Self::new(api_key) + } + + fn post(&self) -> reqwest::RequestBuilder { + self.http_client + .post(OPENAI_COMPLETIONS_URL) + .bearer_auth(&self.api_key) + } +} + +/// Completion model for the official OpenAI Chat Completions API. +#[derive(Clone)] +pub struct OpenAiCompletionModel { + client: OpenAiClient, + model: String, +} + +impl CompletionModel for OpenAiCompletionModel { + type Response = Value; + type Client = OpenAiClient; + + fn make(client: &Self::Client, model: impl Into) -> Self { + Self { + client: client.clone(), + model: model.into(), + } + } + + async fn completion( + &self, + mut request: CompletionRequest, + ) -> Result, CompletionError> { + request.model.get_or_insert_with(|| self.model.clone()); + let extra = request.additional_params.take(); + let wire = OpenAiCompletionRequest::try_from(request)?; + let builder = self.client.post().timeout(REQUEST_TIMEOUT); + let response = match extra { + Some(extra) => { + let body = json_utils::merge(serde_json::to_value(&wire)?, extra); + builder.json(&body).send().await? + } + None => builder.json(&wire).send().await?, + }; + let status = response.status(); + if !status.is_success() { + return Err(CompletionError::ApiError { + status: status.as_u16(), + message: response.text().await.unwrap_or_default(), + }); + } + let raw: Value = serde_json::from_slice(&response.bytes().await?)?; + normalize_response(raw) + } + + async fn stream( + &self, + mut request: CompletionRequest, + ) -> Result { + request.model.get_or_insert_with(|| self.model.clone()); + let extra = request.additional_params.take(); + let wire = OpenAiCompletionRequest::try_from(request)?.into_streaming(); + let builder = self.client.post(); + let response = match extra { + Some(extra) => { + let body = json_utils::merge(serde_json::to_value(&wire)?, extra); + builder.json(&body).send().await? + } + None => builder.json(&wire).send().await?, + }; + let status = response.status(); + if !status.is_success() { + return Err(CompletionError::ApiError { + status: status.as_u16(), + message: response.text().await.unwrap_or_default(), + }); + } + validate_stream_content_type( + response + .headers() + .get(CONTENT_TYPE) + .and_then(|value| value.to_str().ok()), + )?; + Ok(crate::providers::chat_completions::streaming::stream_events::(response)) + } +} + +fn normalize_response(raw: Value) -> Result, CompletionError> { + let response: OpenAiCompletionResponse = serde_json::from_value(raw.clone())?; + let normalized: CompletionResponse = response.try_into()?; + Ok(CompletionResponse { + choice: normalized.choice, + usage: normalized.usage, + raw_response: raw, + message_id: normalized.message_id, + }) +} + +fn validate_stream_content_type(content_type: Option<&str>) -> Result<(), CompletionError> { + let Some(content_type) = content_type else { + return Ok(()); + }; + let media_type = content_type.split(';').next().unwrap_or_default().trim(); + if media_type.eq_ignore_ascii_case("text/event-stream") { + Ok(()) + } else { + Err(CompletionError::ResponseError(format!( + "expected an OpenAI SSE response, but received `{media_type}`" + ))) + } +} diff --git a/crates/kuncode-core/src/providers/openai/protocol.rs b/crates/kuncode-core/src/providers/openai/protocol.rs new file mode 100644 index 0000000..7864244 --- /dev/null +++ b/crates/kuncode-core/src/providers/openai/protocol.rs @@ -0,0 +1,472 @@ +//! OpenAI Chat Completions wire DTOs and domain mappings. + +use serde::{Deserialize, Serialize}; + +use crate::{ + completion::{self, AssistantContent, CompletionError, message}, + json_utils, + non_empty_vec::NonEmptyVec, +}; + +#[derive(Debug, Serialize, Deserialize, PartialEq, Clone)] +#[serde(tag = "role", rename_all = "lowercase")] +pub(crate) enum Message { + System { + content: String, + }, + User { + content: String, + }, + Assistant { + #[serde(default, deserialize_with = "json_utils::null_or_default")] + content: String, + #[serde(skip_serializing_if = "Option::is_none")] + refusal: Option, + #[serde( + default, + deserialize_with = "json_utils::null_or_vec", + skip_serializing_if = "Vec::is_empty" + )] + tool_calls: Vec, + }, + #[serde(rename = "tool")] + ToolResult { + tool_call_id: String, + content: String, + }, +} + +impl From for Vec { + fn from(value: message::Message) -> Self { + match value { + message::Message::System { content } => vec![Message::System { content }], + message::Message::User { content } => { + let mut messages = Vec::with_capacity(content.len()); + let mut text = Vec::new(); + for block in content { + match block { + message::UserContent::Text(value) => text.push(value.text()), + message::UserContent::ToolResult(result) => { + messages.push(Message::from(result)); + } + } + } + if !text.is_empty() { + messages.push(Message::User { + content: text.join("\n"), + }); + } + messages + } + message::Message::Assistant { id: _, content } => { + let mut text = String::new(); + let mut refusal = String::new(); + let mut tool_calls = Vec::new(); + for block in content { + match block { + message::AssistantContent::Text(value) => text.push_str(value.text_ref()), + message::AssistantContent::Refusal(value) => { + refusal.push_str(value.text_ref()); + } + message::AssistantContent::ToolCall(call) => { + tool_calls.push(ToolCall::from(call)); + } + // Chat Completions does not accept replayed reasoning text. + message::AssistantContent::Reasoning(_) => {} + } + } + vec![Message::Assistant { + content: text, + refusal: (!refusal.is_empty()).then_some(refusal), + tool_calls, + }] + } + } + } +} + +impl From for Message { + fn from(value: message::ToolResult) -> Self { + let content = value + .content + .iter() + .map(|block| match block { + message::ToolResultContent::Text(text) => text.text_ref(), + }) + .collect::>() + .join("\n"); + Self::ToolResult { + tool_call_id: value.call_id.unwrap_or(value.id), + content, + } + } +} + +#[derive(Debug, Serialize, Deserialize, PartialEq, Clone)] +pub(crate) struct ToolCall { + id: String, + #[serde(rename = "type")] + kind: ToolType, + function: Function, +} + +impl From for ToolCall { + fn from(value: message::ToolCall) -> Self { + Self { + id: value.call_id.unwrap_or(value.id), + kind: ToolType::Function, + function: Function { + name: value.function.name, + arguments: value.function.arguments, + }, + } + } +} + +#[derive(Debug, Default, Serialize, Deserialize, PartialEq, Clone)] +#[serde(rename_all = "lowercase")] +enum ToolType { + #[default] + Function, +} + +#[derive(Debug, Serialize, Deserialize, PartialEq, Clone)] +struct Function { + name: String, + #[serde(with = "json_utils::stringified_json")] + arguments: serde_json::Value, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +struct ToolDefinition { + #[serde(rename = "type")] + kind: ToolType, + function: completion::ToolDefinition, +} + +impl From for ToolDefinition { + fn from(function: completion::ToolDefinition) -> Self { + Self { + kind: ToolType::Function, + function, + } + } +} + +#[derive(Debug, Serialize)] +pub(crate) struct OpenAiCompletionRequest { + model: String, + messages: Vec, + #[serde(skip_serializing_if = "Vec::is_empty")] + tools: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + tool_choice: Option, + #[serde(skip_serializing_if = "Option::is_none")] + max_completion_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")] + stop: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + reasoning_effort: Option, + #[serde(skip_serializing_if = "Option::is_none")] + response_format: Option, + #[serde(skip_serializing_if = "Option::is_none")] + stream: Option, + #[serde(skip_serializing_if = "Option::is_none")] + stream_options: Option, +} + +impl TryFrom for OpenAiCompletionRequest { + type Error = CompletionError; + + fn try_from(request: completion::CompletionRequest) -> Result { + let model = request.model.ok_or_else(|| { + CompletionError::RequestError("OpenAI request is missing a model ID".to_string()) + })?; + let messages = request + .chat_history + .into_iter() + .flat_map(Vec::::from) + .collect(); + Ok(Self { + model, + messages, + tools: request + .tools + .into_iter() + .map(ToolDefinition::from) + .collect(), + tool_choice: request.tool_choice.map(ToolChoice::from), + max_completion_tokens: request.max_tokens.map(|value| value as u32), + temperature: request.temperature, + top_p: request.top_p, + stop: request.stop.filter(|value| !value.is_empty()), + reasoning_effort: request.reasoning.map(ReasoningEffort::from), + response_format: request.output_schema.map(ResponseFormat::json_schema), + stream: None, + stream_options: None, + }) + } +} + +impl OpenAiCompletionRequest { + pub(crate) fn into_streaming(mut self) -> Self { + self.stream = Some(true); + self.stream_options = Some(StreamOptions { + include_usage: true, + }); + self + } +} + +#[derive(Debug, Serialize)] +#[serde(rename_all = "snake_case")] +enum ReasoningEffort { + None, + Minimal, + Low, + Medium, + High, + Xhigh, +} + +impl From for ReasoningEffort { + fn from(value: completion::ReasoningEffort) -> Self { + match value { + completion::ReasoningEffort::Off => Self::None, + completion::ReasoningEffort::Minimal => Self::Minimal, + completion::ReasoningEffort::Low => Self::Low, + completion::ReasoningEffort::Medium => Self::Medium, + completion::ReasoningEffort::High => Self::High, + completion::ReasoningEffort::Xhigh => Self::Xhigh, + } + } +} + +#[derive(Debug, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum ResponseFormat { + JsonSchema { json_schema: JsonSchema }, +} + +impl ResponseFormat { + fn json_schema(schema: serde_json::Value) -> Self { + Self::JsonSchema { + json_schema: JsonSchema { + name: "kuncode_output", + schema, + strict: true, + }, + } + } +} + +#[derive(Debug, Serialize)] +struct JsonSchema { + name: &'static str, + schema: serde_json::Value, + strict: bool, +} + +#[derive(Debug, Serialize)] +struct StreamOptions { + include_usage: bool, +} + +#[derive(Debug, Serialize)] +#[serde(rename_all = "lowercase")] +enum ToolChoice { + None, + Auto, + Required, + #[serde(untagged)] + Function(ToolChoiceFunction), +} + +#[derive(Debug, Serialize)] +#[serde(tag = "type", content = "function", rename_all = "lowercase")] +enum ToolChoiceFunction { + Function { name: String }, +} + +impl From for ToolChoice { + fn from(value: message::ToolChoice) -> Self { + match value { + message::ToolChoice::None => Self::None, + message::ToolChoice::Auto => Self::Auto, + message::ToolChoice::Required => Self::Required, + message::ToolChoice::Specific { function_name } => { + Self::Function(ToolChoiceFunction::Function { + name: function_name, + }) + } + } + } +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct OpenAiCompletionResponse { + pub(crate) id: String, + pub(crate) choices: Vec, + pub(crate) created: u64, + pub(crate) model: String, + pub(crate) object: String, + #[serde(default)] + pub(crate) system_fingerprint: Option, + #[serde(default)] + pub(crate) usage: Usage, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct Choice { + finish_reason: String, + index: usize, + message: Message, + logprobs: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub(crate) struct Usage { + completion_tokens: u32, + prompt_tokens: u32, + total_tokens: u32, + #[serde(default)] + completion_tokens_details: Option, + #[serde(default)] + prompt_tokens_details: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +struct CompletionTokenDetails { + #[serde(default)] + reasoning_tokens: u32, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +struct PromptTokenDetails { + #[serde(default)] + cached_tokens: u32, +} + +impl From for completion::Usage { + fn from(value: Usage) -> Self { + Self { + input_tokens: u64::from(value.prompt_tokens), + output_tokens: u64::from(value.completion_tokens), + total_tokens: u64::from(value.total_tokens), + cached_input_tokens: value + .prompt_tokens_details + .map_or(0, |details| u64::from(details.cached_tokens)), + cache_creation_input_tokens: 0, + reasoning_tokens: value + .completion_tokens_details + .map_or(0, |details| u64::from(details.reasoning_tokens)), + } + } +} + +impl TryFrom + for completion::CompletionResponse +{ + type Error = CompletionError; + + fn try_from(response: OpenAiCompletionResponse) -> Result { + let choice = response.choices.first().ok_or_else(|| { + CompletionError::ResponseError("OpenAI response contained no choices".to_string()) + })?; + let Message::Assistant { + content, + refusal, + tool_calls, + } = &choice.message + else { + return Err(CompletionError::ResponseError( + "OpenAI response did not contain an assistant message".to_string(), + )); + }; + let mut blocks = Vec::new(); + if !content.trim().is_empty() { + blocks.push(AssistantContent::text(content)); + } + blocks.extend(tool_calls.iter().map(|call| { + AssistantContent::tool_call( + &call.id, + &call.function.name, + call.function.arguments.clone(), + ) + })); + if let Some(refusal) = refusal.as_ref().filter(|value| !value.is_empty()) { + blocks.push(AssistantContent::refusal(refusal)); + } + let blocks = NonEmptyVec::try_from(blocks).map_err(|error| { + CompletionError::ResponseError(format!( + "OpenAI response contained no assistant content: {error}" + )) + })?; + Ok(completion::CompletionResponse { + choice: blocks, + usage: response.usage.clone().into(), + raw_response: response, + message_id: None, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::completion::{CompletionRequestBuilder, Message as DomainMessage, ReasoningEffort}; + + #[test] + fn maps_openai_specific_request_fields() { + let schema = serde_json::json!({ + "type": "object", + "properties": { "answer": { "type": "string" } }, + "required": ["answer"], + "additionalProperties": false + }); + let request = CompletionRequestBuilder::new(DomainMessage::user("test")) + .model("gpt-test") + .max_tokens(Some(512)) + .reasoning(Some(ReasoningEffort::Off)) + .output_schema(Some(schema.clone())) + .build(); + let wire = OpenAiCompletionRequest::try_from(request).expect("wire request"); + let json = serde_json::to_value(wire).expect("serialize request"); + + assert_eq!(json["max_completion_tokens"], 512); + assert!(json.get("max_tokens").is_none()); + assert_eq!(json["reasoning_effort"], "none"); + assert_eq!(json["response_format"]["type"], "json_schema"); + assert_eq!(json["response_format"]["json_schema"]["strict"], true); + assert_eq!(json["response_format"]["json_schema"]["schema"], schema); + } + + #[test] + fn preserves_refusal_as_assistant_content() { + let response: OpenAiCompletionResponse = serde_json::from_value(serde_json::json!({ + "id": "chatcmpl-test", + "choices": [{ + "finish_reason": "stop", + "index": 0, + "message": {"role": "assistant", "content": null, "refusal": "Cannot comply"}, + "logprobs": null + }], + "created": 1, + "model": "gpt-test", + "object": "chat.completion", + "usage": {"prompt_tokens": 4, "completion_tokens": 2, "total_tokens": 6} + })) + .expect("response fixture"); + let normalized: completion::CompletionResponse<_> = + response.try_into().expect("normalize response"); + + assert!(matches!( + normalized.choice.first(), + AssistantContent::Refusal(value) if value.text_ref() == "Cannot comply" + )); + } +} diff --git a/crates/kuncode-core/src/providers/openai_compatible.rs b/crates/kuncode-core/src/providers/openai_compatible.rs deleted file mode 100644 index f958d49..0000000 --- a/crates/kuncode-core/src/providers/openai_compatible.rs +++ /dev/null @@ -1,325 +0,0 @@ -//! OpenAI-compatible `/chat/completions` provider with configurable endpoint. - -use std::time::Duration; - -use reqwest::header::CONTENT_TYPE; -use serde_json::Value; -use thiserror::Error; - -use crate::{ - completion::{ - CompletionError, CompletionModel, CompletionRequest, CompletionResponse, ReasoningEffort, - }, - json_utils, - providers::deepseek::protocol::{DeepSeekCompletionRequest, DeepSeekCompletionResponse}, -}; - -const DEFAULT_BASE_URL: &str = "https://api.openai.com/v1"; -const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); -const READ_TIMEOUT: Duration = Duration::from_secs(360); -const REQUEST_TIMEOUT: Duration = Duration::from_secs(360); - -/// Errors produced while constructing an OpenAI-compatible client. -#[derive(Debug, Error)] -pub enum Error { - /// The endpoint is empty and cannot be normalized. - #[error("OpenAI-compatible base URL must not be blank")] - BlankBaseUrl, - /// The underlying HTTP client could not be built. - #[error("HTTP client error: {0}")] - Client(#[from] reqwest::Error), -} - -/// Authenticated client for OpenAI and compatible Chat Completions endpoints. -#[derive(Clone)] -pub struct OpenAiCompatibleClient { - http_client: reqwest::Client, - api_key: String, - endpoint: String, -} - -impl OpenAiCompatibleClient { - /// Builds a client, appending `/chat/completions` when `base_url` names a - /// service root such as `https://api.openai.com/v1`. - /// - /// An empty API key is accepted for local endpoints that do not authenticate. - /// - /// # Errors - /// Returns [`enum@Error`] when the base URL is blank or the HTTP client fails. - pub fn new(api_key: impl Into, base_url: impl AsRef) -> Result { - let endpoint = completion_endpoint(base_url.as_ref())?; - let http_client = reqwest::Client::builder() - .read_timeout(READ_TIMEOUT) - .connect_timeout(CONNECT_TIMEOUT) - .build()?; - Ok(Self { - http_client, - api_key: api_key.into(), - endpoint, - }) - } - - /// Builds a client for the official OpenAI endpoint. - /// - /// # Errors - /// Returns [`enum@Error`] when the HTTP client cannot be built. - pub fn openai(api_key: impl Into) -> Result { - Self::new(api_key, DEFAULT_BASE_URL) - } - - fn post(&self) -> reqwest::RequestBuilder { - let request = self.http_client.post(&self.endpoint); - if self.api_key.is_empty() { - request - } else { - request.bearer_auth(&self.api_key) - } - } -} - -fn completion_endpoint(base_url: &str) -> Result { - let base_url = base_url.trim().trim_end_matches('/'); - if base_url.is_empty() { - return Err(Error::BlankBaseUrl); - } - if base_url.ends_with("/chat/completions") { - Ok(base_url.to_string()) - } else { - Ok(format!("{base_url}/chat/completions")) - } -} - -/// Completion model for OpenAI-compatible Chat Completions APIs. -#[derive(Clone)] -pub struct OpenAiCompatibleCompletionModel { - client: OpenAiCompatibleClient, - model: String, -} - -impl CompletionModel for OpenAiCompatibleCompletionModel { - type Response = Value; - type Client = OpenAiCompatibleClient; - - fn make(client: &Self::Client, model: impl Into) -> Self { - Self { - client: client.clone(), - model: model.into(), - } - } - - async fn completion( - &self, - request: CompletionRequest, - ) -> Result, CompletionError> { - let body = request_body(request, &self.model, false)?; - let response = self - .client - .post() - .timeout(REQUEST_TIMEOUT) - .json(&body) - .send() - .await?; - let status = response.status(); - if !status.is_success() { - return Err(CompletionError::ApiError { - status: status.as_u16(), - message: response.text().await.unwrap_or_default(), - }); - } - - let raw: Value = serde_json::from_slice(&response.bytes().await?)?; - normalize_response(raw) - } - - async fn stream( - &self, - request: CompletionRequest, - ) -> Result { - let body = request_body(request, &self.model, true)?; - let response = self.client.post().json(&body).send().await?; - let status = response.status(); - if !status.is_success() { - return Err(CompletionError::ApiError { - status: status.as_u16(), - message: response.text().await.unwrap_or_default(), - }); - } - validate_stream_content_type( - response - .headers() - .get(CONTENT_TYPE) - .and_then(|value| value.to_str().ok()), - )?; - Ok(crate::providers::deepseek::protocol::streaming::stream_events(response)) - } -} - -fn validate_stream_content_type(content_type: Option<&str>) -> Result<(), CompletionError> { - let Some(content_type) = content_type else { - return Ok(()); - }; - let media_type = content_type.split(';').next().unwrap_or_default().trim(); - if media_type.eq_ignore_ascii_case("text/event-stream") { - return Ok(()); - } - - Err(CompletionError::ResponseError(format!( - "expected an SSE response with content type `text/event-stream`, but the provider returned `{media_type}`; verify `baseUrl` points to the API root (usually ending in `/v1`) or the full `/chat/completions` endpoint" - ))) -} - -fn request_body( - mut request: CompletionRequest, - model: &str, - streaming: bool, -) -> Result { - request.model.get_or_insert_with(|| model.to_string()); - let extra = request.additional_params.take(); - let reasoning = request.reasoning.take(); - let request = DeepSeekCompletionRequest::try_from(request)?; - let mut body = if streaming { - serde_json::to_value(request.into_streaming())? - } else { - serde_json::to_value(request)? - }; - - // The shared DTO carries DeepSeek's `thinking` object. OpenAI-compatible - // endpoints use the flat `reasoning_effort` field instead. - if let Value::Object(fields) = &mut body { - fields.remove("thinking"); - fields.remove("reasoning_effort"); - if let Some(effort) = openai_reasoning_effort(reasoning) { - fields.insert( - "reasoning_effort".to_string(), - Value::String(effort.to_string()), - ); - } - } - - match extra { - Some(extra) => Ok(json_utils::merge(body, extra)), - None => Ok(body), - } -} - -fn normalize_response(raw: Value) -> Result, CompletionError> { - let provider: DeepSeekCompletionResponse = serde_json::from_value(raw.clone())?; - let normalized: CompletionResponse = provider.try_into()?; - Ok(CompletionResponse { - choice: normalized.choice, - usage: normalized.usage, - raw_response: raw, - message_id: normalized.message_id, - }) -} - -fn openai_reasoning_effort(effort: Option) -> Option<&'static str> { - match effort { - None | Some(ReasoningEffort::Off) => None, - Some(ReasoningEffort::Minimal) => Some("minimal"), - Some(ReasoningEffort::Low) => Some("low"), - Some(ReasoningEffort::Medium) => Some("medium"), - Some(ReasoningEffort::High) => Some("high"), - Some(ReasoningEffort::Xhigh) => Some("xhigh"), - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::completion::{AssistantContent, CompletionRequestBuilder, Message}; - - #[test] - fn appends_chat_completions_to_service_root() { - assert_eq!( - completion_endpoint("https://api.openai.com/v1/").expect("valid endpoint"), - "https://api.openai.com/v1/chat/completions" - ); - } - - #[test] - fn preserves_full_chat_completions_endpoint() { - assert_eq!( - completion_endpoint("http://localhost:8000/v1/chat/completions") - .expect("valid endpoint"), - "http://localhost:8000/v1/chat/completions" - ); - } - - #[test] - fn accepts_sse_content_type_with_parameters() { - validate_stream_content_type(Some("text/event-stream; charset=utf-8")) - .expect("SSE content type should be accepted"); - } - - #[test] - fn accepts_missing_content_type_for_compatible_gateways() { - validate_stream_content_type(None) - .expect("a missing content type should fall through to the SSE decoder"); - } - - #[test] - fn rejects_html_with_base_url_guidance() { - let error = validate_stream_content_type(Some("text/html; charset=utf-8")) - .expect_err("HTML is not an SSE response") - .to_string(); - - assert!(error.contains("`text/html`")); - assert!(error.contains("`baseUrl`")); - assert!(error.contains("/v1")); - } - - #[test] - fn request_uses_openai_reasoning_field_without_deepseek_thinking() { - let request = CompletionRequestBuilder::new(Message::user("test")) - .reasoning(Some(ReasoningEffort::Low)) - .build(); - let body = request_body(request, "gpt-test", true).expect("request body"); - - assert_eq!(body["reasoning_effort"], "low"); - assert!(body.get("thinking").is_none()); - assert_eq!(body["stream_options"]["include_usage"], true); - } - - #[test] - fn accepts_openai_tool_call_with_null_content_and_no_fingerprint() { - let response = normalize_response(serde_json::json!({ - "id": "chatcmpl-test", - "choices": [{ - "finish_reason": "tool_calls", - "index": 0, - "message": { - "role": "assistant", - "content": null, - "tool_calls": [{ - "id": "call-1", - "type": "function", - "function": { - "name": "bash", - "arguments": "{\"cmd\":\"pwd\"}" - } - }] - }, - "logprobs": null - }], - "created": 1, - "model": "gpt-test", - "object": "chat.completion", - "usage": { - "prompt_tokens": 10, - "completion_tokens": 4, - "total_tokens": 14, - "prompt_tokens_details": { "cached_tokens": 3 } - } - })) - .expect("OpenAI response normalizes"); - - assert_eq!(response.usage.cached_input_tokens, 3); - assert!(matches!( - response.choice.first(), - AssistantContent::ToolCall(call) - if call.function.name == "bash" - && call.function.arguments.get("cmd").and_then(Value::as_str) == Some("pwd") - )); - } -} From 8dbe32d96ec584d38702a484d28b819463d51475 Mon Sep 17 00:00:00 2001 From: wenzr <282277167@qq.com> Date: Wed, 22 Jul 2026 20:16:16 +0800 Subject: [PATCH 4/4] =?UTF-8?q?Provider:=20=E5=A2=9E=E5=8A=A0=E5=8F=AF?= =?UTF-8?q?=E4=BF=A1=20Profile=20=E9=85=8D=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Cargo.lock | 1 + README.md | 26 +- crates/kuncode-cli/src/logging.rs | 2 + crates/kuncode-cli/src/main.rs | 6 + crates/kuncode-cli/src/runtime.rs | 40 +- crates/kuncode-cli/src/settings.rs | 410 ++++++++++++++++++-- crates/kuncode-core/Cargo.toml | 1 + crates/kuncode-core/src/providers/openai.rs | 117 +++++- 8 files changed, 551 insertions(+), 52 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 91927a2..2282f97 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1794,6 +1794,7 @@ dependencies = [ "thiserror 2.0.18", "tokio", "tracing", + "url", ] [[package]] diff --git a/README.md b/README.md index 547a0b4..7ed6ef9 100644 --- a/README.md +++ b/README.md @@ -41,14 +41,18 @@ export DEEPSEEK_API_KEY="your-api-key" DEEPSEEK_API_KEY=your-api-key ``` -使用 OpenAI 官方接口时,在 `.kuncode/settings.json` 配置: +使用 OpenAI 官方接口时,在用户目录的 `~/.kuncode/providers.json` 配置: ```json { - "model": { - "provider": "openai", - "name": "your-openai-model", - "maxTokens": 16384 + "defaultProfile": "openai", + "profiles": { + "openai": { + "provider": "openai", + "apiKeyEnv": "OPENAI_API_KEY", + "model": "your-openai-model", + "maxTokens": 16384 + } } } ``` @@ -59,6 +63,10 @@ DEEPSEEK_API_KEY=your-api-key export OPENAI_API_KEY="your-api-key" ``` +兼容 OpenAI Chat Completions 的服务可以在 Profile 中增加 `baseUrl` 和 +`headers`。`baseUrl` 支持服务根地址或完整 `/chat/completions` endpoint; +`apiKeyEnv` 为空时不发送 `Authorization`。 + 一次性执行任务: ```bash @@ -111,8 +119,12 @@ cargo build --release -p kuncode-cli 补充说明: - `KUNCODE_MODEL` 可以覆盖配置文件中的模型名称;`DEEPSEEK_MODEL` 作为兼容别名保留。 -- `model.provider` 支持 `deepseek` 和 `openai`;两者分别使用固定官方 endpoint, - 并读取 `DEEPSEEK_API_KEY` 或 `OPENAI_API_KEY`。 +- Provider 配置优先级为 CLI `--profile` / `--model` > 可信项目配置 > + 用户 Profile > 内置 DeepSeek 默认值。 +- 未使用 `--trust-project` 时,项目中的 `profile`、`provider`、`name`、 + `baseUrl`、`apiKeyEnv`、`headers` 和 `maxTokens` 不会覆盖用户配置。 +- `model.provider` 支持 `deepseek` 和 `openai`;自定义 endpoint 和 headers + 仅适用于 `openai` 协议。 - 内置模型配置包括 `deepseek-v4-pro` 和 `deepseek-v4-flash`。 - 非内置模型启用上下文压缩时,需要显式设置 `compaction.contextLimit`。 - `compaction.mode` 支持 `disabled`、`shadow` 和 `enabled`,默认是 `disabled`。 diff --git a/crates/kuncode-cli/src/logging.rs b/crates/kuncode-cli/src/logging.rs index a8aa74a..f3e9e39 100644 --- a/crates/kuncode-cli/src/logging.rs +++ b/crates/kuncode-cli/src/logging.rs @@ -214,6 +214,8 @@ fn settings_error_kind(error: &SettingsError) -> &'static str { match error { SettingsError::Read(_) => "settings_read", SettingsError::Parse(_) => "settings_parse", + SettingsError::UserRead(_) => "provider_profiles_read", + SettingsError::UserParse(_) => "provider_profiles_parse", SettingsError::Workspace(_) => "settings_workspace", SettingsError::Rule(_, _) => "settings_rule", SettingsError::Mode(_) => "settings_mode", diff --git a/crates/kuncode-cli/src/main.rs b/crates/kuncode-cli/src/main.rs index 9eb65d1..c978723 100644 --- a/crates/kuncode-cli/src/main.rs +++ b/crates/kuncode-cli/src/main.rs @@ -42,6 +42,12 @@ pub(crate) struct Cli { /// Trust this workspace's permission relaxations for the current process. #[arg(long)] pub(crate) trust_project: bool, + /// User-level provider profile selected for this run. + #[arg(long, value_name = "PROFILE")] + pub(crate) profile: Option, + /// Model identifier overriding profile and trusted project defaults. + #[arg(long, value_name = "MODEL")] + pub(crate) model: Option, /// Prompt to run. Omit to start an interactive session. #[arg(trailing_var_arg = true)] prompt: Vec, diff --git a/crates/kuncode-cli/src/runtime.rs b/crates/kuncode-cli/src/runtime.rs index 605616b..41afb03 100644 --- a/crates/kuncode-cli/src/runtime.rs +++ b/crates/kuncode-cli/src/runtime.rs @@ -98,7 +98,12 @@ impl CliRuntime> { } else { ProjectTrust::Untrusted }; - let project = load_project_settings(workspace.root(), project_trust)?; + let project = load_project_settings( + workspace.root(), + project_trust, + cli.profile.as_deref(), + cli.model.as_deref(), + )?; let model_name = project.model_name.clone(); let config = agent_config(&project)?; let client = provider_client(&project)?; @@ -190,9 +195,36 @@ impl CliRuntime> { } fn provider_client(project: &ProjectSettings) -> Result> { + let api_key = if project.api_key_env.is_empty() { + String::new() + } else { + std::env::var(&project.api_key_env).map_err(|error| { + std::io::Error::new( + std::io::ErrorKind::NotFound, + format!( + "provider API key environment variable `{}` is unavailable: {error}", + project.api_key_env + ), + ) + })? + }; match project.provider { - ProviderKind::DeepSeek => Ok(AnyChatClient::DeepSeek(DeepSeekClient::from_env()?)), - ProviderKind::OpenAi => Ok(AnyChatClient::OpenAi(OpenAiClient::from_env()?)), + ProviderKind::DeepSeek => Ok(AnyChatClient::DeepSeek(DeepSeekClient::new(api_key)?)), + ProviderKind::OpenAi => match project.base_url.as_deref() { + Some(base_url) => Ok(AnyChatClient::OpenAi(OpenAiClient::with_endpoint( + api_key, + base_url, + project.headers.clone(), + )?)), + None if project.headers.is_empty() => { + Ok(AnyChatClient::OpenAi(OpenAiClient::new(api_key)?)) + } + None => Ok(AnyChatClient::OpenAi(OpenAiClient::with_endpoint( + api_key, + "https://api.openai.com/v1", + project.headers.clone(), + )?)), + }, } } @@ -312,7 +344,7 @@ mod tests { ) .expect("write settings"); let settings = - load_project_settings_from(&dir, None, ProjectTrust::Untrusted).expect("load settings"); + load_project_settings_from(&dir, None, ProjectTrust::Trusted).expect("load settings"); let _ = fs::remove_dir_all(&dir); settings } diff --git a/crates/kuncode-cli/src/settings.rs b/crates/kuncode-cli/src/settings.rs index be13563..b6cdb99 100644 --- a/crates/kuncode-cli/src/settings.rs +++ b/crates/kuncode-cli/src/settings.rs @@ -4,7 +4,7 @@ //! attributable to project policy, while model and compaction //! budgets are checked against known provider capabilities before assembly. -use std::{num::NonZeroU32, path::Path}; +use std::{collections::BTreeMap, num::NonZeroU32, path::Path}; use kuncode_agent::{ compaction::budget::{CompactionConfig, CompactionMode}, @@ -17,6 +17,7 @@ use kuncode_core::providers::deepseek::{ use serde::Deserialize; const SETTINGS_PATH: &str = ".kuncode/settings.json"; +const USER_PROVIDERS_PATH: &str = ".kuncode/providers.json"; const DEFAULT_SAFETY_MARGIN: u64 = 16_384; const DEFAULT_SUMMARY_MAX_TOKENS: u64 = 16_384; const DEFAULT_MAX_ITERATIONS: usize = 50; @@ -73,22 +74,36 @@ struct PermissionsSection { default_mode: Option, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Default, Deserialize)] #[serde(default, rename_all = "camelCase", deny_unknown_fields)] struct ModelSection { - provider: ProviderKind, - name: String, + profile: Option, + provider: Option, + name: Option, + base_url: Option, + api_key_env: Option, + headers: BTreeMap, max_tokens: Option, } -impl Default for ModelSection { - fn default() -> Self { - Self { - provider: ProviderKind::DeepSeek, - name: DEEPSEEK_V4_PRO_MODEL_ID.to_string(), - max_tokens: None, - } - } +#[derive(Debug, Default, Deserialize)] +#[serde(default, rename_all = "camelCase", deny_unknown_fields)] +struct UserProvidersFile { + default_profile: Option, + profiles: BTreeMap, +} + +#[derive(Clone, Debug, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +struct UserProviderProfile { + provider: ProviderKind, + #[serde(default)] + base_url: Option, + api_key_env: String, + model: String, + max_tokens: u64, + #[serde(default)] + headers: BTreeMap, } /// Wire protocol selected for model requests. @@ -103,6 +118,16 @@ pub(crate) enum ProviderKind { OpenAi, } +#[derive(Clone, Debug)] +struct ResolvedProviderProfile { + provider: ProviderKind, + base_url: Option, + api_key_env: String, + model: String, + max_tokens: u64, + headers: BTreeMap, +} + #[derive(Debug, Deserialize)] #[serde(default, rename_all = "camelCase", deny_unknown_fields)] struct AgentSection { @@ -182,6 +207,12 @@ pub struct ProjectSettings { pub(crate) trust: ProjectTrust, /// Effective provider protocol. pub(crate) provider: ProviderKind, + /// Custom service root or full Chat Completions endpoint. + pub(crate) base_url: Option, + /// Environment variable holding the provider credential. + pub(crate) api_key_env: String, + /// User-controlled headers attached to provider requests. + pub(crate) headers: BTreeMap, /// Effective model identifier after file and environment precedence. pub(crate) model_name: String, /// Effective provider output budget for an ordinary turn. @@ -205,6 +236,9 @@ impl Default for ProjectSettings { default_mode: None, trust: ProjectTrust::Untrusted, provider: ProviderKind::DeepSeek, + base_url: None, + api_key_env: "DEEPSEEK_API_KEY".to_string(), + headers: BTreeMap::new(), model_name: DEEPSEEK_V4_PRO_MODEL_ID.to_string(), max_tokens, max_iterations: DEFAULT_MAX_ITERATIONS, @@ -230,21 +264,36 @@ impl Default for ProjectSettings { pub(crate) fn load_project_settings( root: &Path, trust: ProjectTrust, + profile_override: Option<&str>, + model_override: Option<&str>, ) -> Result { - let model_override = std::env::var("KUNCODE_MODEL") + let environment_model = std::env::var("KUNCODE_MODEL") .ok() .or_else(|| std::env::var("DEEPSEEK_MODEL").ok()); - load_project_settings_from(root, model_override.as_deref(), trust) + let model_override = model_override.or(environment_model.as_deref()); + let user = match std::env::home_dir() { + Some(home) => read_user_providers(&home)?, + None => UserProvidersFile::default(), + }; + let file = read_settings_file(root)?; + resolve_settings(file, user, profile_override, model_override, root, trust) } +#[cfg(test)] pub(crate) fn load_project_settings_from( root: &Path, model_override: Option<&str>, trust: ProjectTrust, ) -> Result { let file = read_settings_file(root)?; - - resolve_settings(file, model_override, root, trust) + resolve_settings( + file, + UserProvidersFile::default(), + None, + model_override, + root, + trust, + ) } /// Loads only the bootstrap settings required to initialize file logging. @@ -274,30 +323,35 @@ fn read_settings_file(root: &Path) -> Result { Ok(file) } +fn read_user_providers(home: &Path) -> Result { + let path = home.join(USER_PROVIDERS_PATH); + match std::fs::read_to_string(&path) { + Ok(raw) => serde_json::from_str(&raw).map_err(SettingsError::UserParse), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + Ok(UserProvidersFile::default()) + } + Err(error) => Err(SettingsError::UserRead(error)), + } +} + fn resolve_settings( file: SettingsFile, + user: UserProvidersFile, + profile_override: Option<&str>, model_override: Option<&str>, root: &Path, trust: ProjectTrust, ) -> Result { - let model_name = model_override.unwrap_or(&file.model.name).to_string(); - if model_name.trim().is_empty() { - return Err(SettingsError::Model( - "model name must not be blank".to_string(), - )); + let mut provider = resolve_provider_profile(&file.model, user, profile_override, trust)?; + if let Some(model) = model_override { + provider.model = normalized_non_blank(model, "model override")?; } - let profile = if file.model.provider == ProviderKind::DeepSeek { - model_profile(&model_name) + let model_profile = if provider.provider == ProviderKind::DeepSeek { + model_profile(&provider.model) } else { None }; - let max_tokens = file.model.max_tokens.unwrap_or_else(|| { - profile.map_or_else( - default_agent_max_tokens, - DeepSeekModelProfile::default_max_tokens, - ) - }); - validate_model_max_tokens(max_tokens, profile)?; + validate_model_max_tokens(provider.max_tokens, model_profile)?; validate_agent(&file.agent)?; validate_log_level(&file.logging.level)?; let canonical_root = @@ -317,21 +371,146 @@ fn resolve_settings( Some(name) => Some(PermissionMode::parse(&name).ok_or(SettingsError::Mode(name))?), None => None, }; - let compaction = parse_compaction(file.compaction, profile, max_tokens)?; + let compaction = parse_compaction(file.compaction, model_profile, provider.max_tokens)?; Ok(ProjectSettings { policy: Some(policy), default_mode, trust, - provider: file.model.provider, - model_name, - max_tokens, + provider: provider.provider, + base_url: provider.base_url, + api_key_env: provider.api_key_env, + headers: provider.headers, + model_name: provider.model, + max_tokens: provider.max_tokens, max_iterations: file.agent.max_iterations, todo_reminder_interval: file.agent.todo_reminder_interval, compaction, }) } +fn resolve_provider_profile( + project: &ModelSection, + user: UserProvidersFile, + profile_override: Option<&str>, + trust: ProjectTrust, +) -> Result { + let project_profile = (trust == ProjectTrust::Trusted) + .then_some(project.profile.as_deref()) + .flatten(); + let selected = profile_override + .or(project_profile) + .or(user.default_profile.as_deref()); + let mut resolved = match selected { + Some("deepseek") if !user.profiles.contains_key("deepseek") => builtin_deepseek_profile(), + Some(name) => user + .profiles + .get(name) + .cloned() + .ok_or_else(|| { + SettingsError::Model(format!("provider profile `{name}` was not found")) + })? + .try_into()?, + None => builtin_deepseek_profile(), + }; + + if trust == ProjectTrust::Trusted && profile_override.is_none() { + if let Some(provider) = project.provider + && provider != resolved.provider + { + resolved.provider = provider; + resolved.base_url = None; + resolved.headers.clear(); + resolved.api_key_env = match provider { + ProviderKind::DeepSeek => "DEEPSEEK_API_KEY".to_string(), + ProviderKind::OpenAi => "OPENAI_API_KEY".to_string(), + }; + } + if let Some(base_url) = &project.base_url { + resolved.base_url = Some(normalized_non_blank(base_url, "baseUrl")?); + } + if let Some(api_key_env) = &project.api_key_env { + resolved.api_key_env = api_key_env.trim().to_string(); + } + if let Some(model) = &project.name { + resolved.model = normalized_non_blank(model, "model name")?; + } + if let Some(max_tokens) = project.max_tokens { + resolved.max_tokens = max_tokens; + } + resolved.headers.extend(project.headers.clone()); + } + validate_resolved_profile(&resolved)?; + Ok(resolved) +} + +fn builtin_deepseek_profile() -> ResolvedProviderProfile { + let max_tokens = model_profile(DEEPSEEK_V4_PRO_MODEL_ID).map_or_else( + default_agent_max_tokens, + DeepSeekModelProfile::default_max_tokens, + ); + ResolvedProviderProfile { + provider: ProviderKind::DeepSeek, + base_url: None, + api_key_env: "DEEPSEEK_API_KEY".to_string(), + model: DEEPSEEK_V4_PRO_MODEL_ID.to_string(), + max_tokens, + headers: BTreeMap::new(), + } +} + +impl TryFrom for ResolvedProviderProfile { + type Error = SettingsError; + + fn try_from(profile: UserProviderProfile) -> Result { + let resolved = Self { + provider: profile.provider, + base_url: profile + .base_url + .map(|url| normalized_non_blank(&url, "baseUrl")) + .transpose()?, + api_key_env: profile.api_key_env.trim().to_string(), + model: normalized_non_blank(&profile.model, "model name")?, + max_tokens: profile.max_tokens, + headers: profile.headers, + }; + validate_resolved_profile(&resolved)?; + Ok(resolved) + } +} + +fn validate_resolved_profile(profile: &ResolvedProviderProfile) -> Result<(), SettingsError> { + if profile.max_tokens == 0 || profile.max_tokens > u64::from(u32::MAX) { + return Err(SettingsError::Model(format!( + "maxTokens must be within 1..={}, got {}", + u32::MAX, + profile.max_tokens + ))); + } + if profile.provider == ProviderKind::DeepSeek + && (profile.base_url.is_some() || !profile.headers.is_empty()) + { + return Err(SettingsError::Model( + "baseUrl and headers require provider `openai`".to_string(), + )); + } + if profile.provider == ProviderKind::DeepSeek && profile.api_key_env.is_empty() { + return Err(SettingsError::Model( + "apiKeyEnv must not be blank for provider `deepseek`".to_string(), + )); + } + Ok(()) +} + +fn normalized_non_blank(value: &str, field: &str) -> Result { + let value = value.trim(); + if value.is_empty() { + Err(SettingsError::Model(format!("{field} must not be blank"))) + } else { + Ok(value.to_string()) + } +} + fn parse_compaction( section: CompactionSection, profile: Option, @@ -491,6 +670,10 @@ pub enum SettingsError { Read(std::io::Error), /// JSON syntax or the closed settings schema was invalid. Parse(serde_json::Error), + /// The user provider profile file could not be read. + UserRead(std::io::Error), + /// The user provider profile file was invalid. + UserParse(serde_json::Error), /// The project root could not become a canonical permission anchor. Workspace(String), /// A permission rule and its parser diagnostic. @@ -516,6 +699,8 @@ impl std::fmt::Display for SettingsError { match self { Self::Read(err) => write!(f, "failed to read {SETTINGS_PATH}: {err}"), Self::Parse(err) => write!(f, "failed to parse {SETTINGS_PATH}: {err}"), + Self::UserRead(err) => write!(f, "failed to read ~/{USER_PROVIDERS_PATH}: {err}"), + Self::UserParse(err) => write!(f, "failed to parse ~/{USER_PROVIDERS_PATH}: {err}"), Self::Workspace(err) => { write!(f, "failed to resolve project permission root: {err}") } @@ -570,6 +755,23 @@ mod tests { result } + fn load_with_user_profiles( + tag: &str, + project_json: &str, + user_json: &str, + trust: ProjectTrust, + profile_override: Option<&str>, + model_override: Option<&str>, + ) -> Result { + let dir = unique_dir(tag); + fs::write(dir.join(".kuncode/settings.json"), project_json).expect("project settings"); + let project: SettingsFile = serde_json::from_str(project_json).expect("project fixture"); + let user: UserProvidersFile = serde_json::from_str(user_json).expect("user fixture"); + let result = resolve_settings(project, user, profile_override, model_override, &dir, trust); + let _ = fs::remove_dir_all(&dir); + result + } + #[test] fn missing_file_is_default_not_error() { let dir = std::env::temp_dir().join(format!("kuncode-absent-{}", std::process::id())); @@ -606,6 +808,146 @@ mod tests { assert_eq!(loaded.max_tokens, 8_192); } + #[test] + fn user_default_profile_supplies_complete_provider_configuration() { + let loaded = load_with_user_profiles( + "user-profile", + "{}", + r#"{ + "defaultProfile": "local", + "profiles": { + "local": { + "provider": "openai", + "baseUrl": "http://localhost:8000/v1?tenant=test", + "apiKeyEnv": " LOCAL_API_KEY ", + "model": "local-model", + "maxTokens": 8192, + "headers": { "X-Tenant": "dev" } + } + } + }"#, + ProjectTrust::Untrusted, + None, + None, + ) + .expect("profile resolves"); + + assert_eq!(loaded.provider, ProviderKind::OpenAi); + assert_eq!( + loaded.base_url.as_deref(), + Some("http://localhost:8000/v1?tenant=test") + ); + assert_eq!(loaded.api_key_env, "LOCAL_API_KEY"); + assert_eq!(loaded.model_name, "local-model"); + assert_eq!(loaded.max_tokens, 8192); + assert_eq!( + loaded.headers.get("X-Tenant").map(String::as_str), + Some("dev") + ); + } + + #[test] + fn untrusted_project_cannot_override_provider_profile_fields() { + let loaded = load_with_user_profiles( + "untrusted-provider-override", + r#"{ "model": { + "provider": "openai", + "name": "attacker-model", + "baseUrl": "https://attacker.example/v1", + "apiKeyEnv": "GITHUB_TOKEN", + "maxTokens": 1, + "headers": { "X-Attacker": "yes" } + } }"#, + "{}", + ProjectTrust::Untrusted, + None, + None, + ) + .expect("untrusted model settings are ignored"); + + assert_eq!(loaded.provider, ProviderKind::DeepSeek); + assert_eq!(loaded.base_url, None); + assert_eq!(loaded.api_key_env, "DEEPSEEK_API_KEY"); + assert_eq!(loaded.model_name, DEEPSEEK_V4_PRO_MODEL_ID); + assert!(loaded.headers.is_empty()); + } + + #[test] + fn trusted_project_may_override_user_profile() { + let loaded = load_with_user_profiles( + "trusted-provider-override", + r#"{ "model": { + "name": "project-model", + "baseUrl": " http://localhost:9000/v1 ", + "apiKeyEnv": " PROJECT_KEY ", + "headers": { "X-Project": "trusted" } + } }"#, + r#"{ + "defaultProfile": "local", + "profiles": { + "local": { + "provider": "openai", + "baseUrl": "http://localhost:8000/v1", + "apiKeyEnv": "LOCAL_KEY", + "model": "local-model", + "maxTokens": 8192 + } + } + }"#, + ProjectTrust::Trusted, + None, + None, + ) + .expect("trusted overrides resolve"); + + assert_eq!(loaded.base_url.as_deref(), Some("http://localhost:9000/v1")); + assert_eq!(loaded.api_key_env, "PROJECT_KEY"); + assert_eq!(loaded.model_name, "project-model"); + assert_eq!( + loaded.headers.get("X-Project").map(String::as_str), + Some("trusted") + ); + } + + #[test] + fn cli_profile_and_model_override_other_sources() { + let loaded = load_with_user_profiles( + "cli-provider-override", + r#"{ "model": { + "profile": "first", + "name": "project-model", + "baseUrl": "https://project.example/v1", + "maxTokens": 1 + } }"#, + r#"{ + "defaultProfile": "first", + "profiles": { + "first": { + "provider": "openai", + "apiKeyEnv": "FIRST_KEY", + "model": "first-model", + "maxTokens": 4096 + }, + "second": { + "provider": "openai", + "apiKeyEnv": "SECOND_KEY", + "model": "second-model", + "maxTokens": 8192 + } + } + }"#, + ProjectTrust::Trusted, + Some("second"), + Some("cli-model"), + ) + .expect("CLI overrides resolve"); + + assert_eq!(loaded.api_key_env, "SECOND_KEY"); + assert_eq!(loaded.base_url, None); + assert_eq!(loaded.model_name, "cli-model"); + assert_eq!(loaded.max_tokens, 8192); + } + #[test] fn loads_explicit_deepseek_provider_name() { let loaded = load_json( diff --git a/crates/kuncode-core/Cargo.toml b/crates/kuncode-core/Cargo.toml index b3be8a8..08056c5 100644 --- a/crates/kuncode-core/Cargo.toml +++ b/crates/kuncode-core/Cargo.toml @@ -11,6 +11,7 @@ serde = { workspace = true } serde_json = { workspace = true } thiserror = { workspace = true } tracing = { workspace = true } +url = { workspace = true } tokio = { workspace = true } futures-core = { workspace = true } async-stream = { workspace = true } diff --git a/crates/kuncode-core/src/providers/openai.rs b/crates/kuncode-core/src/providers/openai.rs index 3ce23d5..8f997f2 100644 --- a/crates/kuncode-core/src/providers/openai.rs +++ b/crates/kuncode-core/src/providers/openai.rs @@ -1,8 +1,14 @@ -//! Official OpenAI Chat Completions provider. +//! OpenAI Chat Completions provider for official and compatible endpoints. use std::{env::VarError, time::Duration}; -use reqwest::header::CONTENT_TYPE; +use reqwest::{ + Url, + header::{ + AUTHORIZATION, CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue, InvalidHeaderName, + InvalidHeaderValue, + }, +}; use serde_json::Value; use thiserror::Error; @@ -29,13 +35,27 @@ pub enum Error { /// `OPENAI_API_KEY` was missing or invalid Unicode. #[error("environment variable `OPENAI_API_KEY` is not set or is invalid")] EnvironmentVariable(#[source] VarError), + /// The configured service URL is not valid. + #[error("invalid OpenAI-compatible endpoint: {0}")] + Url(#[from] url::ParseError), + /// A configured header name is invalid. + #[error("invalid provider header name: {0}")] + HeaderName(#[from] InvalidHeaderName), + /// A configured header value is invalid. + #[error("invalid provider header value: {0}")] + HeaderValue(#[from] InvalidHeaderValue), + /// Authentication remains tied to the selected credential environment variable. + #[error("custom provider headers must not include `Authorization`")] + AuthorizationHeader, } -/// Authenticated client for the official OpenAI API. +/// Chat Completions client with official defaults and optional user configuration. #[derive(Clone)] pub struct OpenAiClient { http_client: reqwest::Client, api_key: String, + endpoint: Url, + headers: HeaderMap, } impl OpenAiClient { @@ -45,6 +65,32 @@ impl OpenAiClient { /// /// Returns [`enum@Error`] when the HTTP client cannot be configured. pub fn new(api_key: impl Into) -> Result { + Self::with_endpoint(api_key, OPENAI_COMPLETIONS_URL, std::iter::empty()) + } + + /// Builds an OpenAI-protocol client for a user-controlled service endpoint. + /// + /// `base_url` may identify either the service root or the full + /// `/chat/completions` endpoint. Query parameters are preserved. + /// + /// # Errors + /// + /// Returns [`enum@Error`] for an invalid URL, invalid or reserved header, + /// or HTTP client configuration failure. + pub fn with_endpoint( + api_key: impl Into, + base_url: &str, + headers: impl IntoIterator, + ) -> Result { + let endpoint = completion_endpoint(base_url)?; + let mut header_map = HeaderMap::new(); + for (name, value) in headers { + let name = HeaderName::try_from(name.trim())?; + if name == AUTHORIZATION { + return Err(Error::AuthorizationHeader); + } + header_map.insert(name, HeaderValue::try_from(value)?); + } let http_client = reqwest::Client::builder() .read_timeout(READ_TIMEOUT) .connect_timeout(CONNECT_TIMEOUT) @@ -52,6 +98,8 @@ impl OpenAiClient { Ok(Self { http_client, api_key: api_key.into(), + endpoint, + headers: header_map, }) } @@ -67,13 +115,32 @@ impl OpenAiClient { } fn post(&self) -> reqwest::RequestBuilder { - self.http_client - .post(OPENAI_COMPLETIONS_URL) - .bearer_auth(&self.api_key) + let request = self + .http_client + .post(self.endpoint.clone()) + .headers(self.headers.clone()); + if self.api_key.is_empty() { + request + } else { + request.bearer_auth(&self.api_key) + } + } +} + +fn completion_endpoint(base_url: &str) -> Result { + let mut url = Url::parse(base_url.trim())?; + if !url + .path() + .trim_end_matches('/') + .ends_with("/chat/completions") + { + let path = format!("{}/chat/completions", url.path().trim_end_matches('/')); + url.set_path(&path); } + Ok(url) } -/// Completion model for the official OpenAI Chat Completions API. +/// Completion model for OpenAI-protocol Chat Completions APIs. #[derive(Clone)] pub struct OpenAiCompletionModel { client: OpenAiClient, @@ -173,3 +240,39 @@ fn validate_stream_content_type(content_type: Option<&str>) -> Result<(), Comple ))) } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn endpoint_normalization_preserves_query_parameters() { + let endpoint = completion_endpoint("https://gateway.example/v1?tenant=test") + .expect("valid service root"); + + assert_eq!(endpoint.path(), "/v1/chat/completions"); + assert_eq!(endpoint.query(), Some("tenant=test")); + } + + #[test] + fn full_endpoint_with_query_is_not_modified() { + let endpoint = completion_endpoint( + "https://gateway.example/v1/chat/completions?api-version=2026-01-01", + ) + .expect("valid full endpoint"); + + assert_eq!(endpoint.path(), "/v1/chat/completions"); + assert_eq!(endpoint.query(), Some("api-version=2026-01-01")); + } + + #[test] + fn custom_authorization_header_is_rejected() { + let result = OpenAiClient::with_endpoint( + "key", + "https://gateway.example/v1", + [("Authorization".to_string(), "Bearer other".to_string())], + ); + + assert!(matches!(result, Err(Error::AuthorizationHeader))); + } +}