chore: init main branch
No files matched your search
@@ -0,0 +1,35 @@
|
||||
[package]
|
||||
name = "kimiswitch"
|
||||
version = "0.6.0"
|
||||
description = "Kimi Switch - model config manager"
|
||||
authors = ["you"]
|
||||
edition = "2021"
|
||||
|
||||
[build-dependencies]
|
||||
tauri-build = { version = "2.0.0", features = [] }
|
||||
|
||||
[dependencies]
|
||||
tauri = { version = "2.0.0", features = ["tray-icon"] }
|
||||
tauri-plugin-opener = "2.0.0"
|
||||
tauri-plugin-single-instance = "2.0.0"
|
||||
serde = { version = "1.0", features = ["derive"] }
|
||||
serde_json = "1.0"
|
||||
indexmap = { version = "2.2", features = ["serde"] }
|
||||
chrono = { version = "0.4", features = ["serde"] }
|
||||
anyhow = "1.0"
|
||||
regex = "1"
|
||||
walkdir = "2"
|
||||
fs_extra = "1"
|
||||
reqwest = { version = "0.12", features = ["json", "rustls-tls", "stream"], default-features = false }
|
||||
futures-util = "0.3"
|
||||
dirs = "5.0"
|
||||
toml = "0.8"
|
||||
rusqlite = { version = "0.32", features = ["bundled", "chrono"] }
|
||||
|
||||
[lib]
|
||||
name = "kimiswitch_lib"
|
||||
crate-type = ["staticlib", "cdylib", "rlib"]
|
||||
|
||||
[features]
|
||||
default = ["custom-protocol"]
|
||||
custom-protocol = ["tauri/custom-protocol"]
|
||||
@@ -0,0 +1,3 @@
|
||||
fn main() {
|
||||
tauri_build::build()
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
{
|
||||
"$schema": "../gen/schemas/desktop-schema.json",
|
||||
"identifier": "default",
|
||||
"description": "Default capabilities for KimiSwitch",
|
||||
"windows": ["main"],
|
||||
"permissions": [
|
||||
"core:default",
|
||||
"opener:allow-open-path",
|
||||
"opener:allow-open-url",
|
||||
"opener:allow-reveal-item-in-dir"
|
||||
]
|
||||
}
|
||||
|
After Width: | Height: | Size: 6.5 KiB |
|
After Width: | Height: | Size: 15 KiB |
|
After Width: | Height: | Size: 1.4 KiB |
|
After Width: | Height: | Size: 1.2 KiB |
|
After Width: | Height: | Size: 1.9 KiB |
|
After Width: | Height: | Size: 2.4 KiB |
|
After Width: | Height: | Size: 2.5 KiB |
|
After Width: | Height: | Size: 4.5 KiB |
|
After Width: | Height: | Size: 654 B |
|
After Width: | Height: | Size: 4.9 KiB |
|
After Width: | Height: | Size: 874 B |
|
After Width: | Height: | Size: 1.3 KiB |
|
After Width: | Height: | Size: 1.5 KiB |
|
After Width: | Height: | Size: 980 B |
@@ -0,0 +1,5 @@
|
||||
<?xml version="1.0" encoding="utf-8"?>
|
||||
<adaptive-icon xmlns:android="http://schemas.android.com/apk/res/android">
|
||||
<foreground android:drawable="@mipmap/ic_launcher_foreground"/>
|
||||
<background android:drawable="@color/ic_launcher_background"/>
|
||||
</adaptive-icon>
|
||||
|
After Width: | Height: | Size: 993 B |
|
After Width: | Height: | Size: 2.7 KiB |
|
After Width: | Height: | Size: 1.2 KiB |
|
After Width: | Height: | Size: 987 B |
|
After Width: | Height: | Size: 1.9 KiB |
|
After Width: | Height: | Size: 1.1 KiB |
|
After Width: | Height: | Size: 1.7 KiB |
|
After Width: | Height: | Size: 3.3 KiB |
|
After Width: | Height: | Size: 2.2 KiB |
|
After Width: | Height: | Size: 2.6 KiB |
|
After Width: | Height: | Size: 5.1 KiB |
|
After Width: | Height: | Size: 3.2 KiB |
|
After Width: | Height: | Size: 3.2 KiB |
|
After Width: | Height: | Size: 6.9 KiB |
|
After Width: | Height: | Size: 4.1 KiB |
@@ -0,0 +1,4 @@
|
||||
<?xml version="1.0" encoding="utf-8"?>
|
||||
<resources>
|
||||
<color name="ic_launcher_background">#fff</color>
|
||||
</resources>
|
||||
|
After Width: | Height: | Size: 15 KiB |
|
After Width: | Height: | Size: 1.4 KiB |
|
After Width: | Height: | Size: 15 KiB |
|
After Width: | Height: | Size: 478 B |
|
After Width: | Height: | Size: 810 B |
|
After Width: | Height: | Size: 810 B |
|
After Width: | Height: | Size: 1.1 KiB |
|
After Width: | Height: | Size: 634 B |
|
After Width: | Height: | Size: 1.0 KiB |
|
After Width: | Height: | Size: 1.0 KiB |
|
After Width: | Height: | Size: 1.5 KiB |
|
After Width: | Height: | Size: 810 B |
|
After Width: | Height: | Size: 1.4 KiB |
|
After Width: | Height: | Size: 1.4 KiB |
|
After Width: | Height: | Size: 2.0 KiB |
|
After Width: | Height: | Size: 14 KiB |
|
After Width: | Height: | Size: 2.0 KiB |
|
After Width: | Height: | Size: 2.8 KiB |
|
After Width: | Height: | Size: 1.4 KiB |
|
After Width: | Height: | Size: 2.5 KiB |
|
After Width: | Height: | Size: 2.7 KiB |
@@ -0,0 +1,859 @@
|
||||
use crate::db;
|
||||
use crate::models::{Agent, Config, DiscoveredModel, Model, Provider, ProviderType};
|
||||
use crate::pi_io;
|
||||
use crate::services::{self, UsageKind, UsageResult};
|
||||
use indexmap::IndexMap;
|
||||
use serde::Serialize;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::{Mutex, OnceLock};
|
||||
use std::time::{Duration, Instant};
|
||||
use tauri_plugin_opener::OpenerExt;
|
||||
|
||||
fn fmt_anyhow(err: anyhow::Error) -> String {
|
||||
err.chain()
|
||||
.map(|e| e.to_string())
|
||||
.collect::<Vec<_>>()
|
||||
.join(": ")
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub fn debug_log(message: String) {
|
||||
eprintln!("[frontend] {}", message);
|
||||
}
|
||||
|
||||
/// SQLite settings key for a provider's billing/usage query kinds
|
||||
/// (JSON array of kind strings, e.g. `["balance:deepseek"]`).
|
||||
fn usage_kinds_key(provider_name: &str) -> String {
|
||||
format!("usage_kinds:{provider_name}")
|
||||
}
|
||||
|
||||
/// Merge per-provider `usage_kinds` into a loaded config: explicit SQLite
|
||||
/// settings first, host-based detection as fallback so existing installs get
|
||||
/// billing support automatically. The field never enters config.toml.
|
||||
fn merge_usage_kinds(config: &mut Config) {
|
||||
for p in config.providers.values_mut() {
|
||||
let from_settings = db::get_setting_pub(&usage_kinds_key(&p.name))
|
||||
.ok()
|
||||
.flatten()
|
||||
.and_then(|s| serde_json::from_str::<Vec<String>>(&s).ok())
|
||||
.filter(|v| !v.is_empty());
|
||||
p.usage_kinds = from_settings.or_else(|| {
|
||||
let kinds = services::detect_provider(&resolve_base_url(p));
|
||||
if kinds.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(kinds.iter().map(|k| k.as_str().to_string()).collect())
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
fn load_pi_native_config() -> Result<Config, String> {
|
||||
let file = pi_io::load_pi_models().map_err(fmt_anyhow)?;
|
||||
let mut config = pi_io::pi_file_to_config(&file);
|
||||
if config.default_model.is_none() {
|
||||
if let Ok(settings) = pi_io::load_pi_settings() {
|
||||
if let (Some(provider), Some(model_id)) =
|
||||
(settings.default_provider, settings.default_model)
|
||||
{
|
||||
if let Some(alias) = config
|
||||
.models
|
||||
.values()
|
||||
.find(|m| m.provider == provider && m.model == model_id)
|
||||
{
|
||||
config.default_model = Some(alias.alias.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(config)
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub fn load_agent_config_command(agent: Agent) -> Result<Config, String> {
|
||||
// Load Kimi Switch's own SQLite database (metadata + migration fallback).
|
||||
let db_config = db::load_config(&agent).ok();
|
||||
|
||||
let mut config = match agent {
|
||||
Agent::KimiCode => {
|
||||
// config.toml is the authoritative source for provider/model data
|
||||
// because the user can add or edit providers at any time via the
|
||||
// CLI's /provider command. SQLite only enriches with Kimi
|
||||
// Switch-private metadata (note, official_url, remembered default
|
||||
// model) and fills gaps when config.toml is incomplete.
|
||||
let mut config = crate::kimi_code_io::load_kimi_code_config_as_config()
|
||||
.map_err(fmt_anyhow)?;
|
||||
|
||||
if let Some(db) = &db_config {
|
||||
// Enrich config.toml providers with SQLite metadata.
|
||||
for (name, p) in config.providers.iter_mut() {
|
||||
if let Some(db_p) = db.providers.get(name) {
|
||||
p.note = db_p.note.clone();
|
||||
p.official_url = db_p.official_url.clone();
|
||||
// Restore Kimi Switch metadata that does not live in
|
||||
// the agent's config.toml.
|
||||
p.icon = db_p.icon.clone();
|
||||
p.icon_color = db_p.icon_color.clone();
|
||||
// Restore the remembered per-provider default model
|
||||
// (Kimi-Switch-private, stored in raw_other).
|
||||
if let Some(dm) = db_p.raw_other.get("default_model") {
|
||||
match &mut p.raw_other {
|
||||
serde_json::Value::Object(obj) => {
|
||||
obj.insert("default_model".to_string(), dm.clone());
|
||||
}
|
||||
_ => {
|
||||
let mut obj = serde_json::Map::new();
|
||||
obj.insert("default_model".to_string(), dm.clone());
|
||||
p.raw_other = serde_json::Value::Object(obj);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// Migration safety: include providers/models that exist in
|
||||
// SQLite but not in config.toml (e.g. after upgrading from the
|
||||
// old single-provider-write behaviour).
|
||||
for (name, p) in &db.providers {
|
||||
config.providers.entry(name.clone()).or_insert_with(|| p.clone());
|
||||
}
|
||||
for (alias, m) in &db.models {
|
||||
config.models.entry(alias.clone()).or_insert_with(|| m.clone());
|
||||
}
|
||||
}
|
||||
|
||||
config
|
||||
}
|
||||
Agent::Pi => {
|
||||
// Pi: SQLite first, fall back to native config on first use.
|
||||
match db_config {
|
||||
Some(config) if !config.providers.is_empty() => config,
|
||||
_ => load_pi_native_config()?,
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
merge_usage_kinds(&mut config);
|
||||
Ok(config)
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub fn save_agent_config_command(agent: Agent, config: Config) -> Result<(), String> {
|
||||
// Save the full Kimi Switch configuration to local SQLite.
|
||||
db::save_config(&agent, &config).map_err(fmt_anyhow)?;
|
||||
// For Kimi Code, config.toml is the authoritative provider store, so
|
||||
// persist changes there immediately — not only on activation. This
|
||||
// ensures edits (Ctrl+S) survive a restart even without switching.
|
||||
if matches!(agent, Agent::KimiCode) {
|
||||
crate::kimi_code_io::save_config_as_kimi_code(&config).map_err(fmt_anyhow)?;
|
||||
}
|
||||
// Persist usage_kinds to the SQLite settings table (never config.toml;
|
||||
// the field is skip_serializing and the TOML export is hand-built).
|
||||
// None / empty array → delete the key.
|
||||
for provider in config.providers.values() {
|
||||
let key = usage_kinds_key(&provider.name);
|
||||
match &provider.usage_kinds {
|
||||
Some(kinds) if !kinds.is_empty() => {
|
||||
let json = serde_json::to_string(kinds).map_err(|e| e.to_string())?;
|
||||
db::set_setting_pub(&key, &json).map_err(fmt_anyhow)?;
|
||||
}
|
||||
_ => db::delete_setting_pub(&key).map_err(fmt_anyhow)?,
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub fn activate_agent_config_command(agent: Agent) -> Result<(), String> {
|
||||
match agent {
|
||||
Agent::KimiCode => {
|
||||
// No-op: save_agent_config_command already writes config.toml for
|
||||
// Kimi Code. Avoiding a second write here prevents a redundant disk
|
||||
// write + backup on every switch.
|
||||
Ok(())
|
||||
}
|
||||
Agent::Pi => {
|
||||
// Load the full config from SQLite and write only the active
|
||||
// provider to Pi's native config files.
|
||||
let config = db::load_config(&agent).map_err(fmt_anyhow)?;
|
||||
let active_config = build_active_config(&config);
|
||||
let file = pi_io::config_to_pi_file(&active_config);
|
||||
pi_io::save_pi_models(&file).map_err(fmt_anyhow)?;
|
||||
|
||||
// Keep Pi's own default provider / model in sync so the switch is
|
||||
// actually picked up on the next `pi` run.
|
||||
let mut settings = pi_io::load_pi_settings().map_err(fmt_anyhow)?;
|
||||
if let Some((provider_name, model_id)) = active_provider_and_model(&active_config) {
|
||||
settings.default_provider = Some(provider_name);
|
||||
settings.default_model = Some(model_id);
|
||||
}
|
||||
pi_io::save_pi_settings(&settings).map_err(fmt_anyhow)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn active_provider_and_model(config: &Config) -> Option<(String, String)> {
|
||||
let alias = config.default_model.as_ref()?;
|
||||
let model = config.models.get(alias)?;
|
||||
Some((model.provider.clone(), model.model.clone()))
|
||||
}
|
||||
|
||||
fn build_active_config(config: &Config) -> Config {
|
||||
// Used only by Pi: writes only the provider explicitly marked as active
|
||||
// to Pi's native config so Pi follows Kimi Switch's selection instead of
|
||||
// falling back to another provider. Kimi Code does not use this — it
|
||||
// writes all providers and selects via default_model.
|
||||
let providers: IndexMap<String, Provider> = config
|
||||
.providers
|
||||
.iter()
|
||||
.filter(|(_, p)| p.active)
|
||||
.map(|(k, p)| (k.clone(), p.clone()))
|
||||
.collect();
|
||||
|
||||
let active_provider_names: std::collections::HashSet<&str> = providers
|
||||
.values()
|
||||
.map(|p| p.name.as_str())
|
||||
.collect();
|
||||
|
||||
let models: IndexMap<String, Model> = config
|
||||
.models
|
||||
.iter()
|
||||
.filter(|(_, m)| active_provider_names.contains(m.provider.as_str()))
|
||||
.map(|(k, m)| (k.clone(), m.clone()))
|
||||
.collect();
|
||||
|
||||
Config {
|
||||
default_model: config.default_model.clone(),
|
||||
providers,
|
||||
models,
|
||||
raw_other: config.raw_other.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub fn open_agent_config_dir(app: tauri::AppHandle, agent: Agent) -> Result<(), String> {
|
||||
let path = agent.config_dir();
|
||||
std::fs::create_dir_all(&path).map_err(|e| e.to_string())?;
|
||||
let path_str = path.to_string_lossy().to_string();
|
||||
app.opener()
|
||||
.open_path(&path_str, None::<&str>)
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub fn get_app_version() -> String {
|
||||
env!("CARGO_PKG_VERSION").to_string()
|
||||
}
|
||||
|
||||
/// Fetch available models from a provider's API.
|
||||
#[tauri::command]
|
||||
pub async fn list_provider_models(provider: Provider) -> Result<Vec<DiscoveredModel>, String> {
|
||||
let api_key = resolve_api_key(&provider)
|
||||
.ok_or_else(|| format!("Provider '{}' has no API key configured", provider.name))?;
|
||||
|
||||
let base = resolve_base_url(&provider);
|
||||
|
||||
match provider.provider_type {
|
||||
ProviderType::Kimi | ProviderType::Openai | ProviderType::OpenaiResponses => {
|
||||
fetch_openai_models(&base, &api_key).await
|
||||
}
|
||||
ProviderType::Anthropic => fetch_anthropic_models(&base, &api_key).await,
|
||||
ProviderType::GoogleGenai => fetch_google_genai_models(&base, &api_key).await,
|
||||
ProviderType::Vertexai => Err("Vertex AI model discovery requires GCP project/location configuration and is not yet supported".to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Test reachability of a provider's base URL (cc-switch semantics):
|
||||
/// any HTTP response counts as reachable; only network-layer errors fail.
|
||||
#[derive(Debug, Serialize, Clone)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ConnectivityResult {
|
||||
pub ok: bool,
|
||||
pub latency_ms: u64,
|
||||
pub status_code: Option<u16>,
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub async fn test_connectivity(provider: Provider) -> Result<ConnectivityResult, String> {
|
||||
let base = resolve_base_url(&provider);
|
||||
if base.trim().is_empty() {
|
||||
return Ok(ConnectivityResult {
|
||||
ok: false,
|
||||
latency_ms: 0,
|
||||
status_code: None,
|
||||
error: Some("no base URL configured".to_string()),
|
||||
});
|
||||
}
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(8))
|
||||
.build()
|
||||
.map_err(|e| format!("failed to build HTTP client: {e}"))?;
|
||||
let start = std::time::Instant::now();
|
||||
match client.get(base.as_str()).send().await {
|
||||
Ok(resp) => Ok(ConnectivityResult {
|
||||
ok: true,
|
||||
latency_ms: start.elapsed().as_millis() as u64,
|
||||
status_code: Some(resp.status().as_u16()),
|
||||
error: None,
|
||||
}),
|
||||
Err(e) => Ok(ConnectivityResult {
|
||||
ok: false,
|
||||
latency_ms: start.elapsed().as_millis() as u64,
|
||||
status_code: None,
|
||||
error: Some(if e.is_connect() {
|
||||
"connection refused / DNS failed".to_string()
|
||||
} else if e.is_timeout() {
|
||||
"request timed out".to_string()
|
||||
} else {
|
||||
e.to_string()
|
||||
}),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_api_key(provider: &Provider) -> Option<String> {
|
||||
if provider.managed {
|
||||
return Some("managed".to_string());
|
||||
}
|
||||
if let Some(key) = &provider.api_key {
|
||||
if !key.is_empty() {
|
||||
return Some(key.clone());
|
||||
}
|
||||
}
|
||||
let env_key = expected_api_key_key(&provider.provider_type);
|
||||
provider.env.get(env_key).filter(|s| !s.is_empty()).cloned()
|
||||
}
|
||||
|
||||
fn resolve_base_url(provider: &Provider) -> String {
|
||||
provider
|
||||
.base_url
|
||||
.clone()
|
||||
.filter(|s| !s.is_empty())
|
||||
.unwrap_or_else(|| {
|
||||
provider
|
||||
.provider_type
|
||||
.default_base_url()
|
||||
.unwrap_or("")
|
||||
.to_string()
|
||||
})
|
||||
}
|
||||
|
||||
fn expected_api_key_key(provider_type: &ProviderType) -> &'static str {
|
||||
match provider_type {
|
||||
ProviderType::Kimi => "KIMI_API_KEY",
|
||||
ProviderType::Anthropic => "ANTHROPIC_API_KEY",
|
||||
ProviderType::Openai | ProviderType::OpenaiResponses => "OPENAI_API_KEY",
|
||||
ProviderType::GoogleGenai => "GOOGLE_API_KEY",
|
||||
ProviderType::Vertexai => "VERTEXAI_API_KEY",
|
||||
}
|
||||
}
|
||||
|
||||
// ── OpenAI-compatible /models endpoint (with pagination) ─────────────
|
||||
|
||||
async fn fetch_openai_models(base: &str, api_key: &str) -> Result<Vec<DiscoveredModel>, String> {
|
||||
#[derive(serde::Deserialize)]
|
||||
struct OaiModel {
|
||||
id: String,
|
||||
}
|
||||
#[derive(serde::Deserialize)]
|
||||
struct OaiList {
|
||||
data: Vec<OaiModel>,
|
||||
#[serde(default)]
|
||||
has_more: Option<bool>,
|
||||
#[serde(default)]
|
||||
last_id: Option<String>,
|
||||
#[serde(default)]
|
||||
next_page_token: Option<String>,
|
||||
}
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
let root = base.trim_end_matches('/').to_string();
|
||||
let mut url = format!("{}/models", root);
|
||||
let mut all: Vec<OaiModel> = Vec::new();
|
||||
const MAX_PAGES: usize = 50;
|
||||
|
||||
for _ in 0..MAX_PAGES {
|
||||
let resp = client
|
||||
.get(&url)
|
||||
.bearer_auth(api_key)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("HTTP request to {} failed: {e}", url))?;
|
||||
if !resp.status().is_success() {
|
||||
return Err(format!("{} returned HTTP {}", url, resp.status()));
|
||||
}
|
||||
let body: OaiList = resp
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| format!("failed to parse response from {}: {e}", url))?;
|
||||
all.extend(body.data);
|
||||
|
||||
// OpenAI cursor pagination: has_more + last_id → ?after=<last_id>
|
||||
if body.has_more.unwrap_or(false) {
|
||||
if let Some(last_id) = body.last_id.clone() {
|
||||
url = format!("{}/models?after={}", root, last_id);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
// Token-based pagination: next_page_token → ?page_token=<token>
|
||||
if let Some(token) = body.next_page_token.clone() {
|
||||
if !token.is_empty() {
|
||||
url = format!("{}/models?page_token={}", root, token);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
||||
let mut seen = std::collections::HashSet::new();
|
||||
Ok(all
|
||||
.into_iter()
|
||||
.filter(|m| seen.insert(m.id.clone()))
|
||||
.map(|m| DiscoveredModel {
|
||||
id: m.id,
|
||||
display_name: None,
|
||||
max_context_size: None,
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
// ── Anthropic /v1/models endpoint (with pagination) ──────────────────
|
||||
|
||||
async fn fetch_anthropic_models(base: &str, api_key: &str) -> Result<Vec<DiscoveredModel>, String> {
|
||||
#[derive(serde::Deserialize)]
|
||||
struct AntModel {
|
||||
id: String,
|
||||
display_name: Option<String>,
|
||||
}
|
||||
#[derive(serde::Deserialize)]
|
||||
struct AntList {
|
||||
data: Vec<AntModel>,
|
||||
#[serde(default)]
|
||||
has_more: Option<bool>,
|
||||
#[serde(default)]
|
||||
last_id: Option<String>,
|
||||
}
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
let root = base.trim_end_matches('/').to_string();
|
||||
let mut url = format!("{}/v1/models?limit=1000", root);
|
||||
let mut all: Vec<AntModel> = Vec::new();
|
||||
const MAX_PAGES: usize = 20;
|
||||
|
||||
for _ in 0..MAX_PAGES {
|
||||
let resp = client
|
||||
.get(&url)
|
||||
.header("x-api-key", api_key)
|
||||
.header("anthropic-version", "2023-06-01")
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("HTTP request to {} failed: {e}", url))?;
|
||||
if !resp.status().is_success() {
|
||||
return Err(format!("{} returned HTTP {}", url, resp.status()));
|
||||
}
|
||||
let body: AntList = resp
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| format!("failed to parse response from {}: {e}", url))?;
|
||||
let more = body.has_more.unwrap_or(false);
|
||||
let cursor = body.last_id.clone();
|
||||
all.extend(body.data);
|
||||
if !(more && cursor.is_some()) {
|
||||
break;
|
||||
}
|
||||
url = format!("{}/v1/models?limit=1000&after_id={}", root, cursor.unwrap());
|
||||
}
|
||||
|
||||
let mut seen = std::collections::HashSet::new();
|
||||
Ok(all
|
||||
.into_iter()
|
||||
.filter(|m| seen.insert(m.id.clone()))
|
||||
.map(|m| DiscoveredModel {
|
||||
id: m.id,
|
||||
display_name: m.display_name,
|
||||
max_context_size: None,
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
// ── Google GenAI /v1beta/models endpoint (with pagination) ───────────
|
||||
|
||||
async fn fetch_google_genai_models(
|
||||
base: &str,
|
||||
api_key: &str,
|
||||
) -> Result<Vec<DiscoveredModel>, String> {
|
||||
#[derive(serde::Deserialize)]
|
||||
struct GglModel {
|
||||
name: String,
|
||||
#[serde(rename = "displayName")]
|
||||
display_name: Option<String>,
|
||||
#[serde(rename = "outputTokenLimit")]
|
||||
output_token_limit: Option<u64>,
|
||||
}
|
||||
#[derive(serde::Deserialize)]
|
||||
struct GglList {
|
||||
models: Vec<GglModel>,
|
||||
#[serde(rename = "nextPageToken", default)]
|
||||
next_page_token: Option<String>,
|
||||
}
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
let root = base.trim_end_matches('/').to_string();
|
||||
let mut url = format!("{}/v1beta/models?key={}&pageSize=1000", root, api_key);
|
||||
let mut all: Vec<GglModel> = Vec::new();
|
||||
const MAX_PAGES: usize = 50;
|
||||
|
||||
for _ in 0..MAX_PAGES {
|
||||
let resp = client
|
||||
.get(&url)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("HTTP request to {} failed: {e}", url))?;
|
||||
if !resp.status().is_success() {
|
||||
return Err(format!("{} returned HTTP {}", url, resp.status()));
|
||||
}
|
||||
let body: GglList = resp
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| format!("failed to parse response from {}: {e}", url))?;
|
||||
let cursor = body.next_page_token.clone();
|
||||
all.extend(body.models);
|
||||
match cursor {
|
||||
Some(t) if !t.is_empty() => {
|
||||
url = format!(
|
||||
"{}/v1beta/models?key={}&pageSize=1000&pageToken={}",
|
||||
root, api_key, t
|
||||
);
|
||||
}
|
||||
_ => break,
|
||||
}
|
||||
}
|
||||
|
||||
let mut seen = std::collections::HashSet::new();
|
||||
Ok(all
|
||||
.into_iter()
|
||||
.filter(|m| seen.insert(m.name.clone()))
|
||||
.map(|m| {
|
||||
let id = m
|
||||
.name
|
||||
.strip_prefix("models/")
|
||||
.unwrap_or(&m.name)
|
||||
.to_string();
|
||||
DiscoveredModel {
|
||||
id,
|
||||
display_name: m.display_name,
|
||||
max_context_size: m.output_token_limit,
|
||||
}
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// App settings (generic key/value via SQLite)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tauri::command]
|
||||
pub fn get_app_setting(key: String) -> Option<String> {
|
||||
db::get_setting_pub(&key).ok().flatten()
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub fn set_app_setting(key: String, value: String) -> Result<(), String> {
|
||||
db::set_setting_pub(&key, &value).map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Provider billing / usage query (cc-switch semantics)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// 5-minute in-memory cache keyed by (agent, provider_name).
|
||||
/// Only successful results are cached; failures are always re-queryable.
|
||||
const USAGE_CACHE_TTL: Duration = Duration::from_secs(300);
|
||||
|
||||
type UsageCache = Mutex<HashMap<(String, String), (Instant, UsageResult)>>;
|
||||
|
||||
fn usage_cache() -> &'static UsageCache {
|
||||
static CACHE: OnceLock<UsageCache> = OnceLock::new();
|
||||
CACHE.get_or_init(|| Mutex::new(HashMap::new()))
|
||||
}
|
||||
|
||||
/// Query a provider's balance / plan quota. The frontend passes only the
|
||||
/// provider name — base_url, api_key and usage kinds are all resolved here,
|
||||
/// so the API key never crosses IPC and the host routing cannot be spoofed.
|
||||
///
|
||||
/// Error channel semantics (cc-switch):
|
||||
/// - `Err(_)` = transient failure (network/timeout/body read) → frontend
|
||||
/// retries and keeps the last good value.
|
||||
/// - `Ok(success:false)` = deterministic failure (no key / auth / non-2xx /
|
||||
/// bad JSON / unsupported provider) → show the error text directly.
|
||||
#[tauri::command]
|
||||
pub async fn query_provider_usage(
|
||||
agent: Agent,
|
||||
provider_name: String,
|
||||
force_refresh: Option<bool>,
|
||||
) -> Result<UsageResult, String> {
|
||||
let cache_key = (agent.as_str().to_string(), provider_name.clone());
|
||||
|
||||
if !force_refresh.unwrap_or(false) {
|
||||
let cached = usage_cache()
|
||||
.lock()
|
||||
.unwrap()
|
||||
.get(&cache_key)
|
||||
.and_then(|(ts, result)| (ts.elapsed() < USAGE_CACHE_TTL).then(|| result.clone()));
|
||||
if let Some(result) = cached {
|
||||
return Ok(result);
|
||||
}
|
||||
}
|
||||
|
||||
// Load via the same path as load_agent_config_command so usage_kinds
|
||||
// (SQLite merge + host-detect fallback) is already resolved.
|
||||
let config = load_agent_config_command(agent)?;
|
||||
let Some(provider) = config.providers.get(&provider_name) else {
|
||||
return Ok(UsageResult::failure(format!(
|
||||
"provider '{provider_name}' not found"
|
||||
)));
|
||||
};
|
||||
|
||||
// The api_key only ever goes into request headers — never into logs,
|
||||
// error messages, or the cache key.
|
||||
let api_key = provider
|
||||
.api_key
|
||||
.clone()
|
||||
.filter(|s| !s.trim().is_empty())
|
||||
.or_else(|| {
|
||||
provider
|
||||
.env
|
||||
.get(expected_api_key_key(&provider.provider_type))
|
||||
.cloned()
|
||||
.filter(|s| !s.is_empty())
|
||||
});
|
||||
let Some(api_key) = api_key else {
|
||||
return Ok(UsageResult::failure(if provider.managed {
|
||||
"provider uses managed OAuth; usage query requires an API key".to_string()
|
||||
} else {
|
||||
"no API key configured".to_string()
|
||||
}));
|
||||
};
|
||||
|
||||
let base_url = resolve_base_url(provider);
|
||||
let kinds: Vec<UsageKind> = provider
|
||||
.usage_kinds
|
||||
.as_ref()
|
||||
.filter(|v| !v.is_empty())
|
||||
.map(|v| {
|
||||
v.iter()
|
||||
.filter_map(|s| s.parse::<UsageKind>().ok())
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.filter(|v| !v.is_empty())
|
||||
.unwrap_or_else(|| services::detect_provider(&base_url));
|
||||
if kinds.is_empty() {
|
||||
return Ok(UsageResult::failure(
|
||||
"unsupported provider: no usage query available for this base URL".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
// A failing kind must not take down the others: collect successes,
|
||||
// deterministic failures and transient failures separately.
|
||||
let mut data: Vec<crate::services::UsageData> = Vec::new();
|
||||
let mut errors: Vec<String> = Vec::new();
|
||||
let mut transient: Vec<String> = Vec::new();
|
||||
let mut any_success = false;
|
||||
for kind in kinds {
|
||||
match services::query_kind(kind, &base_url, &api_key).await {
|
||||
Ok(result) if result.success => {
|
||||
any_success = true;
|
||||
if let Some(d) = result.data {
|
||||
data.extend(d);
|
||||
}
|
||||
}
|
||||
Ok(result) => {
|
||||
if let Some(e) = result.error {
|
||||
errors.push(format!("{}: {e}", kind.as_str()));
|
||||
}
|
||||
}
|
||||
Err(e) => transient.push(format!("{}: {e}", kind.as_str())),
|
||||
}
|
||||
}
|
||||
|
||||
if any_success {
|
||||
let result = UsageResult {
|
||||
success: true,
|
||||
data: if data.is_empty() { None } else { Some(data) },
|
||||
error: if errors.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(errors.join("; "))
|
||||
},
|
||||
};
|
||||
usage_cache()
|
||||
.lock()
|
||||
.unwrap()
|
||||
.insert(cache_key, (Instant::now(), result.clone()));
|
||||
Ok(result)
|
||||
} else if !transient.is_empty() {
|
||||
// All kinds failed transiently → propagate Err so the frontend
|
||||
// rejects and retries (keep-last-good).
|
||||
Err(transient.join("; "))
|
||||
} else {
|
||||
Ok(UsageResult::failure(errors.join("; ")))
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Version check (lightweight: GET Gitea releases API, compare tag_name)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[derive(Debug, Serialize, Clone)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct UpdateInfo {
|
||||
pub current: String,
|
||||
pub latest: String,
|
||||
pub update_available: bool,
|
||||
pub release_url: String,
|
||||
pub download_url: Option<String>,
|
||||
}
|
||||
|
||||
fn parse_version(s: &str) -> Vec<u32> {
|
||||
s.trim_start_matches('v')
|
||||
.split('.')
|
||||
.filter_map(|p| p.parse::<u32>().ok())
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn version_lt(current: &str, latest: &str) -> bool {
|
||||
let c = parse_version(current);
|
||||
let l = parse_version(latest);
|
||||
for i in 0..c.len().max(l.len()) {
|
||||
let cv = *c.get(i).unwrap_or(&0);
|
||||
let lv = *l.get(i).unwrap_or(&0);
|
||||
if cv < lv {
|
||||
return true;
|
||||
}
|
||||
if cv > lv {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub async fn check_for_update() -> Result<UpdateInfo, String> {
|
||||
let current = env!("CARGO_PKG_VERSION").to_string();
|
||||
let url = "https://git.codingplan.site/api/v1/repos/admin/KimiCodeSwitch/releases?limit=1";
|
||||
|
||||
let resp = reqwest::get(url)
|
||||
.await
|
||||
.map_err(|e| format!("request failed: {e}"))?;
|
||||
|
||||
let releases: serde_json::Value = resp
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| format!("parse failed: {e}"))?;
|
||||
|
||||
let first = releases
|
||||
.as_array()
|
||||
.and_then(|a| a.first())
|
||||
.ok_or("no releases found")?;
|
||||
|
||||
let latest = first
|
||||
.get("tag_name")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("0.0.0")
|
||||
.to_string();
|
||||
|
||||
let release_url = first
|
||||
.get("html_url")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
|
||||
let download_url = first
|
||||
.get("assets")
|
||||
.and_then(|a| a.as_array())
|
||||
.and_then(|a| a.first())
|
||||
.and_then(|a| a.get("browser_download_url"))
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
|
||||
let update_available = version_lt(¤t, &latest);
|
||||
|
||||
Ok(UpdateInfo {
|
||||
current,
|
||||
latest,
|
||||
update_available,
|
||||
release_url,
|
||||
download_url,
|
||||
})
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Silent download with progress events
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tauri::command]
|
||||
pub async fn download_update(
|
||||
app: tauri::AppHandle,
|
||||
url: String,
|
||||
) -> Result<String, String> {
|
||||
use futures_util::StreamExt;
|
||||
use std::io::Write;
|
||||
use tauri::Emitter;
|
||||
|
||||
let resp = reqwest::get(&url)
|
||||
.await
|
||||
.map_err(|e| format!("download request failed: {e}"))?;
|
||||
|
||||
let total = resp.content_length().unwrap_or(0);
|
||||
|
||||
let temp_dir = std::env::temp_dir();
|
||||
let file_path = temp_dir.join("KimiSwitch_update.msi");
|
||||
|
||||
let mut file = std::fs::File::create(&file_path)
|
||||
.map_err(|e| format!("create temp file failed: {e}"))?;
|
||||
|
||||
let mut downloaded: u64 = 0;
|
||||
let mut stream = resp.bytes_stream();
|
||||
|
||||
while let Some(chunk_result) = stream.next().await {
|
||||
let chunk = chunk_result.map_err(|e| format!("read chunk failed: {e}"))?;
|
||||
file.write_all(&chunk)
|
||||
.map_err(|e| format!("write failed: {e}"))?;
|
||||
downloaded += chunk.len() as u64;
|
||||
|
||||
let progress = if total > 0 {
|
||||
((downloaded as f64 / total as f64) * 100.0).min(100.0) as u32
|
||||
} else {
|
||||
0
|
||||
};
|
||||
|
||||
let _ = app.emit(
|
||||
"download-progress",
|
||||
serde_json::json!({
|
||||
"downloaded": downloaded,
|
||||
"total": total,
|
||||
"progress": progress,
|
||||
}),
|
||||
);
|
||||
}
|
||||
|
||||
drop(file);
|
||||
|
||||
let path_str = file_path.to_string_lossy().to_string();
|
||||
|
||||
let _ = app.emit(
|
||||
"download-complete",
|
||||
serde_json::json!({ "path": &path_str }),
|
||||
);
|
||||
|
||||
Ok(path_str)
|
||||
}
|
||||
|
||||
/// Open the downloaded MSI installer using the system default handler.
|
||||
#[tauri::command]
|
||||
pub fn open_installer(app: tauri::AppHandle, path: String) -> Result<(), String> {
|
||||
use tauri_plugin_opener::OpenerExt;
|
||||
app.opener()
|
||||
.open_path(&path, None::<&str>)
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
//! Configuration file I/O utilities for KimiSwitch.
|
||||
//! Reads and writes the global config file and profile files.
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
use anyhow::{Context, Result};
|
||||
use chrono::{Duration, Local};
|
||||
|
||||
/// Creates a backup of `path` inside a `backups/` subdirectory next to the
|
||||
/// original file, then removes backup files older than `retention_days`.
|
||||
pub fn backup_file(path: &Path, retention_days: i64) -> Result<PathBuf> {
|
||||
if !path.exists() {
|
||||
return Ok(PathBuf::new());
|
||||
}
|
||||
|
||||
let parent = path.parent().context("path has no parent directory")?;
|
||||
let backups_dir = parent.join("backups");
|
||||
std::fs::create_dir_all(&backups_dir)
|
||||
.with_context(|| format!("failed to create backups directory {}", backups_dir.display()))?;
|
||||
|
||||
let file_name = path.file_name().context("path has no file name")?;
|
||||
let timestamp = Local::now().format("%Y%m%d_%H%M%S_%.3f").to_string();
|
||||
let backup_name = format!("{}.bak.{}", file_name.to_string_lossy(), timestamp);
|
||||
let mut backup_path = backups_dir.join(&backup_name);
|
||||
|
||||
if backup_path.exists() {
|
||||
for n in 1..1000 {
|
||||
let candidate = backups_dir.join(format!(
|
||||
"{}.bak.{}.{:03}",
|
||||
file_name.to_string_lossy(),
|
||||
timestamp,
|
||||
n
|
||||
));
|
||||
if !candidate.exists() {
|
||||
backup_path = candidate;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
std::fs::copy(path, &backup_path).with_context(|| {
|
||||
format!(
|
||||
"failed to back up {} to {}",
|
||||
path.display(),
|
||||
backup_path.display()
|
||||
)
|
||||
})?;
|
||||
|
||||
cleanup_old_backups(&backups_dir, retention_days)?;
|
||||
|
||||
Ok(backup_path)
|
||||
}
|
||||
|
||||
/// Removes backup files in `dir` older than `retention_days`, based on file
|
||||
/// modification time. Only files whose names contain `.bak.` are deleted.
|
||||
fn cleanup_old_backups(dir: &Path, retention_days: i64) -> Result<()> {
|
||||
let cutoff = Local::now() - Duration::days(retention_days);
|
||||
let entries = std::fs::read_dir(dir)
|
||||
.with_context(|| format!("failed to read backups directory {}", dir.display()))?;
|
||||
|
||||
for entry in entries {
|
||||
let entry = entry.context("failed to read directory entry")?;
|
||||
let path = entry.path();
|
||||
if !path.is_file() {
|
||||
continue;
|
||||
}
|
||||
let name = path.file_name().and_then(|n| n.to_str()).unwrap_or("");
|
||||
if !name.contains(".bak.") {
|
||||
continue;
|
||||
}
|
||||
let metadata = entry
|
||||
.metadata()
|
||||
.with_context(|| format!("failed to read metadata for {}", path.display()))?;
|
||||
let modified = metadata
|
||||
.modified()
|
||||
.with_context(|| format!("failed to read modified time for {}", path.display()))?;
|
||||
let modified: chrono::DateTime<Local> = modified.into();
|
||||
if modified < cutoff {
|
||||
std::fs::remove_file(&path)
|
||||
.with_context(|| format!("failed to remove old backup {}", path.display()))?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,338 @@
|
||||
//! Kimi Switch local SQLite storage.
|
||||
//!
|
||||
//! This module stores the full Kimi Switch configuration (all providers and models
|
||||
//! for both Kimi Code and Pi agents) in a local SQLite database. It is separate
|
||||
//! from the agent-specific config files that are written when the user activates
|
||||
//! a provider.
|
||||
|
||||
use std::path::PathBuf;
|
||||
|
||||
use anyhow::Context;
|
||||
use indexmap::IndexMap;
|
||||
use rusqlite::{params, Connection};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::models::{Agent, Config, Model, Provider, ProviderType};
|
||||
|
||||
pub type DbResult<T> = anyhow::Result<T>;
|
||||
|
||||
pub fn kimi_switch_data_dir() -> PathBuf {
|
||||
dirs::home_dir()
|
||||
.map(|h| h.join(".kimi-switch"))
|
||||
.expect("failed to resolve home directory")
|
||||
}
|
||||
|
||||
pub fn db_path() -> PathBuf {
|
||||
kimi_switch_data_dir().join("kimi-switch.db")
|
||||
}
|
||||
|
||||
/// One-time migration: if the legacy `~/.pi-switch/` data directory exists and
|
||||
/// `~/.kimi-switch/` does not, move it so existing users keep their saved
|
||||
/// configuration. Safe to call on every startup — it is a no-op once the new
|
||||
/// directory exists.
|
||||
fn migrate_legacy_data_dir() {
|
||||
let new_dir = kimi_switch_data_dir();
|
||||
if new_dir.exists() {
|
||||
return;
|
||||
}
|
||||
let old_dir = match dirs::home_dir() {
|
||||
Some(h) => h.join(".pi-switch"),
|
||||
None => return,
|
||||
};
|
||||
if !old_dir.exists() {
|
||||
return;
|
||||
}
|
||||
// Best-effort move; failures are silently ignored so the app can still start.
|
||||
if let Some(new_parent) = new_dir.parent() {
|
||||
if let Err(_) = std::fs::create_dir_all(new_parent) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
let _ = std::fs::rename(&old_dir, &new_dir);
|
||||
}
|
||||
|
||||
pub fn init_db() -> DbResult<Connection> {
|
||||
migrate_legacy_data_dir();
|
||||
let path = db_path();
|
||||
if let Some(parent) = path.parent() {
|
||||
std::fs::create_dir_all(parent)
|
||||
.with_context(|| format!("failed to create directory {}", parent.display()))?;
|
||||
}
|
||||
let conn = Connection::open(&path)
|
||||
.with_context(|| format!("failed to open database {}", path.display()))?;
|
||||
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS providers (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
agent TEXT NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
provider_type TEXT NOT NULL,
|
||||
base_url TEXT,
|
||||
api_key TEXT,
|
||||
env TEXT,
|
||||
note TEXT,
|
||||
official_url TEXT,
|
||||
managed INTEGER NOT NULL DEFAULT 0,
|
||||
enabled INTEGER NOT NULL DEFAULT 1,
|
||||
active INTEGER NOT NULL DEFAULT 0,
|
||||
raw_other TEXT,
|
||||
UNIQUE(agent, name)
|
||||
)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS models (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
agent TEXT NOT NULL,
|
||||
alias TEXT NOT NULL,
|
||||
provider_name TEXT NOT NULL,
|
||||
model TEXT NOT NULL,
|
||||
max_context_size INTEGER NOT NULL,
|
||||
display_name TEXT,
|
||||
supports_1m INTEGER NOT NULL DEFAULT 0,
|
||||
capabilities TEXT,
|
||||
raw_other TEXT,
|
||||
UNIQUE(agent, alias)
|
||||
)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS settings (
|
||||
key TEXT PRIMARY KEY,
|
||||
value TEXT NOT NULL
|
||||
)",
|
||||
[],
|
||||
)?;
|
||||
|
||||
// Add icon columns for provider icon picker (introduced in v0.3.1).
|
||||
// SQLite has no ADD COLUMN IF NOT EXISTS, so ignore duplicate-column errors.
|
||||
for stmt in [
|
||||
"ALTER TABLE providers ADD COLUMN icon TEXT",
|
||||
"ALTER TABLE providers ADD COLUMN icon_color TEXT",
|
||||
] {
|
||||
if let Err(e) = conn.execute(stmt, []) {
|
||||
let msg = e.to_string();
|
||||
if !msg.contains("duplicate column name") {
|
||||
return Err(e).with_context(|| format!("failed to run migration: {}", stmt));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(conn)
|
||||
}
|
||||
|
||||
pub fn load_config(agent: &Agent) -> DbResult<Config> {
|
||||
let mut conn = init_db()?;
|
||||
let tx = conn.transaction()?;
|
||||
|
||||
let default_model = get_setting_tx(&tx, &default_model_key(agent))?;
|
||||
|
||||
let mut providers = IndexMap::new();
|
||||
{
|
||||
let mut stmt = tx.prepare(
|
||||
"SELECT name, provider_type, base_url, api_key, env, note, official_url, managed, enabled, active, icon, icon_color, raw_other
|
||||
FROM providers WHERE agent = ?1 ORDER BY id",
|
||||
)?;
|
||||
let provider_rows = stmt.query_map(params![agent.as_str()], |row| {
|
||||
let provider_type: String = row.get(1)?;
|
||||
let env_json: Option<String> = row.get(4)?;
|
||||
let raw_json: Option<String> = row.get(12)?;
|
||||
Ok(Provider {
|
||||
name: row.get(0)?,
|
||||
provider_type: provider_type_for_str(&provider_type),
|
||||
base_url: row.get(2)?,
|
||||
api_key: row.get(3)?,
|
||||
env: env_json
|
||||
.and_then(|s| serde_json::from_str(&s).ok())
|
||||
.unwrap_or_default(),
|
||||
note: row.get(5)?,
|
||||
official_url: row.get(6)?,
|
||||
managed: row.get::<_, i32>(7)? != 0,
|
||||
enabled: row.get::<_, i32>(8)? != 0,
|
||||
active: row.get::<_, i32>(9)? != 0,
|
||||
icon: row.get(10)?,
|
||||
icon_color: row.get(11)?,
|
||||
raw_other: raw_json
|
||||
.and_then(|s| serde_json::from_str(&s).ok())
|
||||
.unwrap_or(Value::Null),
|
||||
// Merged from the settings table by load_agent_config_command.
|
||||
usage_kinds: None,
|
||||
})
|
||||
})?;
|
||||
|
||||
for provider in provider_rows {
|
||||
let p = provider?;
|
||||
providers.insert(p.name.clone(), p);
|
||||
}
|
||||
}
|
||||
|
||||
let mut models = IndexMap::new();
|
||||
{
|
||||
let mut stmt = tx.prepare(
|
||||
"SELECT alias, provider_name, model, max_context_size, display_name, supports_1m, capabilities, raw_other
|
||||
FROM models WHERE agent = ?1 ORDER BY id",
|
||||
)?;
|
||||
let model_rows = stmt.query_map(params![agent.as_str()], |row| {
|
||||
let caps_json: Option<String> = row.get(6)?;
|
||||
let raw_json: Option<String> = row.get(7)?;
|
||||
Ok(Model {
|
||||
alias: row.get(0)?,
|
||||
provider: row.get(1)?,
|
||||
model: row.get(2)?,
|
||||
max_context_size: row.get::<_, i64>(3)? as u64,
|
||||
display_name: row.get(4)?,
|
||||
supports_1m: row.get::<_, i32>(5)? != 0,
|
||||
capabilities: caps_json
|
||||
.and_then(|s| serde_json::from_str(&s).ok())
|
||||
.unwrap_or_default(),
|
||||
raw_other: raw_json
|
||||
.and_then(|s| serde_json::from_str(&s).ok())
|
||||
.unwrap_or(Value::Null),
|
||||
})
|
||||
})?;
|
||||
|
||||
for model in model_rows {
|
||||
let m = model?;
|
||||
models.insert(m.alias.clone(), m);
|
||||
}
|
||||
}
|
||||
|
||||
tx.commit()?;
|
||||
|
||||
Ok(Config {
|
||||
default_model,
|
||||
providers,
|
||||
models,
|
||||
raw_other: Value::Null,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn save_config(agent: &Agent, config: &Config) -> DbResult<()> {
|
||||
let mut conn = init_db()?;
|
||||
let tx = conn.transaction()?;
|
||||
|
||||
tx.execute("DELETE FROM providers WHERE agent = ?1", params![agent.as_str()])?;
|
||||
tx.execute("DELETE FROM models WHERE agent = ?1", params![agent.as_str()])?;
|
||||
|
||||
{
|
||||
let mut insert_provider = tx.prepare(
|
||||
"INSERT INTO providers
|
||||
(agent, name, provider_type, base_url, api_key, env, note, official_url, managed, enabled, active, icon, icon_color, raw_other)
|
||||
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)",
|
||||
)?;
|
||||
|
||||
for provider in config.providers.values() {
|
||||
insert_provider.execute(params![
|
||||
agent.as_str(),
|
||||
provider.name,
|
||||
provider.provider_type.as_str(),
|
||||
provider.base_url,
|
||||
provider.api_key,
|
||||
serde_json::to_string(&provider.env).ok(),
|
||||
provider.note,
|
||||
provider.official_url,
|
||||
provider.managed as i32,
|
||||
provider.enabled as i32,
|
||||
provider.active as i32,
|
||||
provider.icon,
|
||||
provider.icon_color,
|
||||
serde_json::to_string(&provider.raw_other).ok(),
|
||||
])?;
|
||||
}
|
||||
}
|
||||
|
||||
{
|
||||
let mut insert_model = tx.prepare(
|
||||
"INSERT INTO models
|
||||
(agent, alias, provider_name, model, max_context_size, display_name, supports_1m, capabilities, raw_other)
|
||||
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)",
|
||||
)?;
|
||||
|
||||
for model in config.models.values() {
|
||||
insert_model.execute(params![
|
||||
agent.as_str(),
|
||||
model.alias,
|
||||
model.provider,
|
||||
model.model,
|
||||
model.max_context_size as i64,
|
||||
model.display_name,
|
||||
model.supports_1m as i32,
|
||||
serde_json::to_string(&model.capabilities).ok(),
|
||||
serde_json::to_string(&model.raw_other).ok(),
|
||||
])?;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(default_model) = &config.default_model {
|
||||
set_setting_tx(&tx, &default_model_key(agent), default_model)?;
|
||||
} else {
|
||||
tx.execute("DELETE FROM settings WHERE key = ?1", params![default_model_key(agent)])?;
|
||||
}
|
||||
|
||||
tx.commit()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn default_model_key(agent: &Agent) -> String {
|
||||
format!("default_model:{}", agent.as_str())
|
||||
}
|
||||
|
||||
fn get_setting_tx(tx: &rusqlite::Transaction, key: &str) -> DbResult<Option<String>> {
|
||||
let mut stmt = tx.prepare("SELECT value FROM settings WHERE key = ?1")?;
|
||||
let mut rows = stmt.query(params![key])?;
|
||||
if let Some(row) = rows.next()? {
|
||||
Ok(Some(row.get(0)?))
|
||||
} else {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
|
||||
fn set_setting_tx(tx: &rusqlite::Transaction, key: &str, value: &str) -> DbResult<()> {
|
||||
tx.execute(
|
||||
"INSERT INTO settings (key, value) VALUES (?1, ?2)
|
||||
ON CONFLICT(key) DO UPDATE SET value = excluded.value",
|
||||
params![key, value],
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Public helper: read a single setting without an explicit transaction.
|
||||
pub fn get_setting_pub(key: &str) -> DbResult<Option<String>> {
|
||||
let conn = init_db()?;
|
||||
let mut stmt = conn.prepare("SELECT value FROM settings WHERE key = ?1")?;
|
||||
let mut rows = stmt.query(params![key])?;
|
||||
if let Some(row) = rows.next()? {
|
||||
Ok(Some(row.get(0)?))
|
||||
} else {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
|
||||
/// Public helper: write a single setting in its own transaction.
|
||||
pub fn set_setting_pub(key: &str, value: &str) -> DbResult<()> {
|
||||
let mut conn = init_db()?;
|
||||
let tx = conn.transaction()?;
|
||||
set_setting_tx(&tx, key, value)?;
|
||||
tx.commit()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Public helper: delete a single setting (no-op if the key does not exist).
|
||||
pub fn delete_setting_pub(key: &str) -> DbResult<()> {
|
||||
let conn = init_db()?;
|
||||
conn.execute("DELETE FROM settings WHERE key = ?1", params![key])?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn provider_type_for_str(s: &str) -> ProviderType {
|
||||
match s {
|
||||
"anthropic" => ProviderType::Anthropic,
|
||||
"openai" => ProviderType::Openai,
|
||||
"openai_responses" => ProviderType::OpenaiResponses,
|
||||
"google-genai" => ProviderType::GoogleGenai,
|
||||
"vertexai" => ProviderType::Vertexai,
|
||||
_ => ProviderType::Kimi,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,758 @@
|
||||
//! Kimi Code configuration I/O.
|
||||
//!
|
||||
//! Kimi Code stores its configuration in `~/.kimi-code/config.toml`.
|
||||
//! This module reads/writes that file and converts between Kimi's TOML
|
||||
//! format and Kimi Switch's internal `Config`/`Provider`/`Model` types.
|
||||
|
||||
use std::path::PathBuf;
|
||||
|
||||
use anyhow::Context;
|
||||
use indexmap::IndexMap;
|
||||
use serde_json::Value;
|
||||
use toml::value::{Table, Value as TomlValue};
|
||||
|
||||
use crate::config_io::backup_file;
|
||||
use crate::models::{Agent, Config, Model, Provider, ProviderType};
|
||||
|
||||
pub type KimiResult<T> = anyhow::Result<T>;
|
||||
|
||||
pub fn kimi_code_config_dir() -> PathBuf {
|
||||
std::env::var_os("KIMI_CODE_HOME")
|
||||
.map(PathBuf::from)
|
||||
.or_else(|| dirs::home_dir().map(|h| h.join(".kimi-code")))
|
||||
.expect("failed to resolve Kimi Code config directory")
|
||||
}
|
||||
|
||||
pub fn kimi_code_config_path() -> PathBuf {
|
||||
kimi_code_config_dir().join("config.toml")
|
||||
}
|
||||
|
||||
/// Read the existing Kimi Code `config.toml`, or return an empty table if it
|
||||
/// does not exist yet.
|
||||
pub fn load_kimi_code_config() -> KimiResult<TomlValue> {
|
||||
let path = kimi_code_config_path();
|
||||
if !path.exists() {
|
||||
return Ok(TomlValue::Table(Table::new()));
|
||||
}
|
||||
let content = std::fs::read_to_string(&path)
|
||||
.with_context(|| format!("failed to read {}", path.display()))?;
|
||||
let value: TomlValue = content
|
||||
.parse()
|
||||
.with_context(|| format!("failed to parse {}", path.display()))?;
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
/// Write the given TOML value to Kimi Code's config file, creating the
|
||||
/// directory if needed and backing up the previous file.
|
||||
pub fn save_kimi_code_config(value: &TomlValue) -> KimiResult<()> {
|
||||
let path = kimi_code_config_path();
|
||||
if let Some(parent) = path.parent() {
|
||||
std::fs::create_dir_all(parent)
|
||||
.with_context(|| format!("failed to create directory {}", parent.display()))?;
|
||||
}
|
||||
|
||||
if path.exists() {
|
||||
backup_file(&path, 7).with_context(|| {
|
||||
format!("failed to back up Kimi Code config {}", path.display())
|
||||
})?;
|
||||
}
|
||||
|
||||
let content = toml::to_string_pretty(value)
|
||||
.with_context(|| "failed to serialize Kimi Code config.toml")?;
|
||||
std::fs::write(&path, content)
|
||||
.with_context(|| format!("failed to write {}", path.display()))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Import a Kimi Code TOML config into Kimi Switch's internal `Config`.
|
||||
pub fn kimi_code_to_config(value: &TomlValue) -> Config {
|
||||
let root = value.as_table().cloned().unwrap_or_default();
|
||||
|
||||
let default_model = root
|
||||
.get("default_model")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
|
||||
let mut providers = IndexMap::new();
|
||||
let mut models = IndexMap::new();
|
||||
|
||||
// Read models first so we can determine which provider the default_model
|
||||
// belongs to. Only that provider should be marked active in the UI.
|
||||
if let Some(models_table) = root.get("models").and_then(|v| v.as_table()) {
|
||||
for (alias, mv) in models_table {
|
||||
let table = mv.as_table().cloned().unwrap_or_default();
|
||||
let provider = table
|
||||
.get("provider")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string())
|
||||
.unwrap_or_default();
|
||||
let model_id = table
|
||||
.get("model")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string())
|
||||
.unwrap_or_default();
|
||||
let max_context_size = table
|
||||
.get("max_context_size")
|
||||
.and_then(|v| v.as_integer())
|
||||
.map(|n| n as u64)
|
||||
.unwrap_or(128_000);
|
||||
let display_name = table
|
||||
.get("display_name")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
let capabilities: Vec<String> = table
|
||||
.get("capabilities")
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| {
|
||||
arr.iter()
|
||||
.filter_map(|v| v.as_str().map(|s| s.to_string()))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
let supports_1m = max_context_size >= 1_000_000;
|
||||
|
||||
let raw_other = {
|
||||
let mut rest = table.clone();
|
||||
rest.remove("provider");
|
||||
rest.remove("model");
|
||||
rest.remove("max_context_size");
|
||||
rest.remove("display_name");
|
||||
rest.remove("capabilities");
|
||||
toml_value_to_json(&TomlValue::Table(rest))
|
||||
};
|
||||
|
||||
models.insert(
|
||||
alias.clone(),
|
||||
Model {
|
||||
alias: alias.clone(),
|
||||
provider,
|
||||
model: model_id,
|
||||
max_context_size,
|
||||
display_name,
|
||||
supports_1m,
|
||||
capabilities,
|
||||
raw_other,
|
||||
},
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let active_provider_name = default_model
|
||||
.as_ref()
|
||||
.and_then(|alias| models.get(alias))
|
||||
.map(|m| m.provider.as_str());
|
||||
|
||||
if let Some(providers_table) = root.get("providers").and_then(|v| v.as_table()) {
|
||||
for (name, pv) in providers_table {
|
||||
let table = pv.as_table().cloned().unwrap_or_default();
|
||||
let provider_type = table
|
||||
.get("type")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(provider_type_for_kimi_type)
|
||||
.unwrap_or(ProviderType::Kimi);
|
||||
|
||||
let base_url = table.get("base_url").and_then(|v| v.as_str()).map(|s| s.to_string());
|
||||
let api_key = table.get("api_key").and_then(|v| v.as_str()).map(|s| s.to_string());
|
||||
let managed = table.contains_key("oauth")
|
||||
|| table.get("managed").and_then(|v| v.as_bool()).unwrap_or(false);
|
||||
let enabled = table.get("enabled").and_then(|v| v.as_bool()).unwrap_or(true);
|
||||
|
||||
let env: IndexMap<String, String> = table
|
||||
.get("env")
|
||||
.and_then(|v| v.as_table())
|
||||
.map(|t| {
|
||||
t.iter()
|
||||
.filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_string())))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
let icon = table
|
||||
.get("icon")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
let icon_color = table
|
||||
.get("icon_color")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
|
||||
let raw_other = {
|
||||
let mut rest = table.clone();
|
||||
rest.remove("type");
|
||||
rest.remove("base_url");
|
||||
rest.remove("api_key");
|
||||
rest.remove("managed");
|
||||
rest.remove("enabled");
|
||||
rest.remove("env");
|
||||
// NOTE: `oauth` is intentionally kept in raw_other so that the
|
||||
// exact `storage`/`key` block round-trips verbatim on export.
|
||||
// Previously it was stripped here and regenerated on export
|
||||
// with a key derived from the provider name, which corrupted
|
||||
// the managed:kimi-code credential reference (e.g.
|
||||
// "oauth/kimi-code" became "oauth/managed-kimi-code").
|
||||
toml_value_to_json(&TomlValue::Table(rest))
|
||||
};
|
||||
|
||||
providers.insert(
|
||||
name.clone(),
|
||||
Provider {
|
||||
name: name.clone(),
|
||||
provider_type,
|
||||
base_url: base_url.filter(|s| !s.is_empty()),
|
||||
api_key: api_key.filter(|s| !s.is_empty()),
|
||||
env,
|
||||
note: None,
|
||||
official_url: None,
|
||||
managed,
|
||||
enabled,
|
||||
active: active_provider_name == Some(name.as_str()),
|
||||
icon,
|
||||
icon_color,
|
||||
raw_other,
|
||||
usage_kinds: None,
|
||||
},
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let raw_other = {
|
||||
let mut rest = root;
|
||||
rest.remove("default_model");
|
||||
rest.remove("providers");
|
||||
rest.remove("models");
|
||||
toml_value_to_json(&TomlValue::Table(rest))
|
||||
};
|
||||
|
||||
Config {
|
||||
default_model,
|
||||
providers,
|
||||
models,
|
||||
raw_other,
|
||||
}
|
||||
}
|
||||
|
||||
/// Export Kimi Switch's internal `Config` to a Kimi Code TOML config value.
|
||||
///
|
||||
/// If `existing` is provided, unknown top-level sections (e.g. `services`) are
|
||||
/// preserved; otherwise a fresh TOML table is used.
|
||||
pub fn config_to_kimi_code(config: &Config, existing: Option<&TomlValue>) -> TomlValue {
|
||||
let mut root = existing
|
||||
.and_then(|v| v.as_table().cloned())
|
||||
.unwrap_or_default();
|
||||
|
||||
root.insert(
|
||||
"default_model".to_string(),
|
||||
config
|
||||
.default_model
|
||||
.clone()
|
||||
.map(TomlValue::String)
|
||||
.unwrap_or(TomlValue::String("".to_string())),
|
||||
);
|
||||
|
||||
let mut providers_table = Table::new();
|
||||
for (name, provider) in &config.providers {
|
||||
// Write ALL providers to Kimi Code config. The active provider is
|
||||
// selected via `default_model`, so keeping the full list matches the
|
||||
// CLI's native multi-provider behavior and prevents data loss when
|
||||
// switching.
|
||||
let mut pt = Table::new();
|
||||
pt.insert("type".to_string(), TomlValue::String(provider.provider_type.as_str().to_string()));
|
||||
if let Some(base_url) = provider.base_url.clone().filter(|s| !s.is_empty()) {
|
||||
pt.insert("base_url".to_string(), TomlValue::String(base_url));
|
||||
}
|
||||
if let Some(api_key) = provider.api_key.clone().filter(|s| !s.is_empty()) {
|
||||
pt.insert("api_key".to_string(), TomlValue::String(api_key));
|
||||
}
|
||||
pt.insert("enabled".to_string(), TomlValue::Boolean(provider.enabled));
|
||||
if let Some(icon) = provider.icon.clone().filter(|s| !s.is_empty()) {
|
||||
pt.insert("icon".to_string(), TomlValue::String(icon));
|
||||
}
|
||||
if let Some(icon_color) = provider.icon_color.clone().filter(|s| !s.is_empty()) {
|
||||
pt.insert("icon_color".to_string(), TomlValue::String(icon_color));
|
||||
}
|
||||
if !provider.env.is_empty() {
|
||||
let mut env_table = Table::new();
|
||||
for (k, v) in &provider.env {
|
||||
env_table.insert(k.clone(), TomlValue::String(v.clone()));
|
||||
}
|
||||
pt.insert("env".to_string(), TomlValue::Table(env_table));
|
||||
}
|
||||
if provider.managed {
|
||||
// Preserve existing oauth config if present, otherwise create a default entry.
|
||||
let oauth = provider
|
||||
.raw_other
|
||||
.get("oauth")
|
||||
.and_then(json_to_toml)
|
||||
.unwrap_or_else(|| {
|
||||
let mut t = Table::new();
|
||||
t.insert("storage".to_string(), TomlValue::String("file".to_string()));
|
||||
t.insert("key".to_string(), TomlValue::String(format!("oauth/{}", name.replace(':', "-"))));
|
||||
TomlValue::Table(t)
|
||||
});
|
||||
pt.insert("oauth".to_string(), oauth);
|
||||
}
|
||||
// Merge remaining raw fields (oauth and env are handled explicitly above).
|
||||
if let TomlValue::Table(mut extra) = json_to_toml(&provider.raw_other).unwrap_or(TomlValue::Table(Table::new())) {
|
||||
extra.remove("oauth");
|
||||
extra.remove("env");
|
||||
// Strip Kimi-Switch-private field: the remembered per-provider
|
||||
// default model is stored in raw_other.default_model and must NOT
|
||||
// leak into the agent's config.toml.
|
||||
extra.remove("default_model");
|
||||
for (k, v) in extra {
|
||||
pt.insert(k, v);
|
||||
}
|
||||
}
|
||||
providers_table.insert(name.clone(), TomlValue::Table(pt));
|
||||
}
|
||||
root.insert("providers".to_string(), TomlValue::Table(providers_table));
|
||||
|
||||
// Write all models — each provider's models are kept so the CLI's
|
||||
// /provider command can list and switch between them.
|
||||
let mut models_table = Table::new();
|
||||
for (alias, model) in &config.models {
|
||||
let mut mt = Table::new();
|
||||
mt.insert("provider".to_string(), TomlValue::String(model.provider.clone()));
|
||||
mt.insert("model".to_string(), TomlValue::String(model.model.clone()));
|
||||
mt.insert(
|
||||
"max_context_size".to_string(),
|
||||
TomlValue::Integer(model.max_context_size as i64),
|
||||
);
|
||||
if let Some(display_name) = model.display_name.clone().filter(|s| !s.is_empty()) {
|
||||
mt.insert("display_name".to_string(), TomlValue::String(display_name));
|
||||
}
|
||||
let capabilities = model.capabilities.clone();
|
||||
if !capabilities.is_empty() {
|
||||
mt.insert(
|
||||
"capabilities".to_string(),
|
||||
TomlValue::Array(capabilities.into_iter().map(TomlValue::String).collect()),
|
||||
);
|
||||
}
|
||||
if let TomlValue::Table(mut extra) = json_to_toml(&model.raw_other).unwrap_or(TomlValue::Table(Table::new())) {
|
||||
extra.remove("provider");
|
||||
extra.remove("model");
|
||||
extra.remove("max_context_size");
|
||||
extra.remove("display_name");
|
||||
extra.remove("capabilities");
|
||||
for (k, v) in extra {
|
||||
mt.insert(k, v);
|
||||
}
|
||||
}
|
||||
models_table.insert(alias.clone(), TomlValue::Table(mt));
|
||||
}
|
||||
root.insert("models".to_string(), TomlValue::Table(models_table));
|
||||
|
||||
TomlValue::Table(root)
|
||||
}
|
||||
|
||||
fn provider_type_for_kimi_type(typ: &str) -> ProviderType {
|
||||
match typ {
|
||||
"anthropic" => ProviderType::Anthropic,
|
||||
"openai" => ProviderType::Openai,
|
||||
"openai_responses" => ProviderType::OpenaiResponses,
|
||||
"google-genai" => ProviderType::GoogleGenai,
|
||||
"vertexai" => ProviderType::Vertexai,
|
||||
_ => ProviderType::Kimi,
|
||||
}
|
||||
}
|
||||
|
||||
fn toml_value_to_json(value: &TomlValue) -> Value {
|
||||
match value {
|
||||
TomlValue::String(s) => Value::String(s.clone()),
|
||||
TomlValue::Integer(n) => Value::Number((*n).into()),
|
||||
TomlValue::Float(n) => Value::Number(serde_json::Number::from_f64(*n).unwrap_or(Value::Null.as_u64().unwrap_or(0).into())),
|
||||
TomlValue::Boolean(b) => Value::Bool(*b),
|
||||
TomlValue::Array(arr) => Value::Array(arr.iter().map(toml_value_to_json).collect()),
|
||||
TomlValue::Table(t) => {
|
||||
let mut map = serde_json::Map::new();
|
||||
for (k, v) in t {
|
||||
map.insert(k.clone(), toml_value_to_json(v));
|
||||
}
|
||||
Value::Object(map)
|
||||
}
|
||||
TomlValue::Datetime(dt) => Value::String(dt.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
fn json_to_toml(value: &Value) -> Option<TomlValue> {
|
||||
match value {
|
||||
Value::Null => None,
|
||||
Value::Bool(b) => Some(TomlValue::Boolean(*b)),
|
||||
Value::Number(n) => {
|
||||
if let Some(i) = n.as_i64() {
|
||||
Some(TomlValue::Integer(i))
|
||||
} else if let Some(f) = n.as_f64() {
|
||||
Some(TomlValue::Float(f))
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
Value::String(s) => Some(TomlValue::String(s.clone())),
|
||||
Value::Array(arr) => Some(TomlValue::Array(
|
||||
arr.iter().filter_map(json_to_toml).collect(),
|
||||
)),
|
||||
Value::Object(obj) => {
|
||||
let mut t = Table::new();
|
||||
for (k, v) in obj {
|
||||
if let Some(tv) = json_to_toml(v) {
|
||||
t.insert(k.clone(), tv);
|
||||
}
|
||||
}
|
||||
Some(TomlValue::Table(t))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Convenience entry point used by commands: load and convert.
|
||||
pub fn load_kimi_code_config_as_config() -> KimiResult<Config> {
|
||||
let value = load_kimi_code_config()?;
|
||||
Ok(kimi_code_to_config(&value))
|
||||
}
|
||||
|
||||
/// Convenience entry point used by commands: convert and save.
|
||||
pub fn save_config_as_kimi_code(config: &Config) -> KimiResult<()> {
|
||||
let existing = load_kimi_code_config().ok();
|
||||
let value = config_to_kimi_code(config, existing.as_ref());
|
||||
save_kimi_code_config(&value)
|
||||
}
|
||||
|
||||
/// Open the Kimi Code config directory in the system file manager.
|
||||
pub fn open_kimi_code_config_dir() -> KimiResult<()> {
|
||||
let path = kimi_code_config_dir();
|
||||
std::fs::create_dir_all(&path)
|
||||
.with_context(|| format!("failed to create directory {}", path.display()))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
pub fn config_dir(&self) -> PathBuf {
|
||||
match self {
|
||||
Agent::KimiCode => kimi_code_config_dir(),
|
||||
Agent::Pi => crate::pi_io::pi_agent_dir(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn kimi_code_import_extracts_providers_and_models() {
|
||||
let toml_str = r#"
|
||||
default_model = "kimi-code/kimi-for-coding"
|
||||
default_thinking = true
|
||||
|
||||
[providers."managed:kimi-code"]
|
||||
type = "kimi"
|
||||
api_key = ""
|
||||
base_url = "https://api.kimi.com/coding/v1"
|
||||
|
||||
[providers."managed:kimi-code".oauth]
|
||||
storage = "file"
|
||||
key = "oauth/kimi-code"
|
||||
|
||||
[models."kimi-code/kimi-for-coding"]
|
||||
provider = "managed:kimi-code"
|
||||
model = "kimi-for-coding"
|
||||
max_context_size = 262144
|
||||
capabilities = ["thinking", "always_thinking", "image_in", "video_in", "tool_use"]
|
||||
display_name = "K2.7 Code"
|
||||
|
||||
[services.moonshot_search]
|
||||
base_url = "https://api.kimi.com/coding/v1/search"
|
||||
api_key = ""
|
||||
"#;
|
||||
let value: TomlValue = toml_str.parse().unwrap();
|
||||
let config = kimi_code_to_config(&value);
|
||||
|
||||
assert_eq!(config.default_model.as_deref(), Some("kimi-code/kimi-for-coding"));
|
||||
assert_eq!(config.providers.len(), 1);
|
||||
assert_eq!(config.models.len(), 1);
|
||||
|
||||
let provider = config.providers.get("managed:kimi-code").unwrap();
|
||||
assert_eq!(provider.provider_type, ProviderType::Kimi);
|
||||
assert!(provider.managed);
|
||||
assert_eq!(provider.base_url.as_deref(), Some("https://api.kimi.com/coding/v1"));
|
||||
|
||||
let model = config.models.get("kimi-code/kimi-for-coding").unwrap();
|
||||
assert_eq!(model.provider, "managed:kimi-code");
|
||||
assert_eq!(model.model, "kimi-for-coding");
|
||||
assert_eq!(model.max_context_size, 262144);
|
||||
assert!(model.capabilities.contains(&"thinking".to_string()));
|
||||
// supports_1m reflects context size only, not thinking capability.
|
||||
assert!(!model.supports_1m);
|
||||
|
||||
// Unknown top-level sections are preserved.
|
||||
let services = config.raw_other.get("services");
|
||||
assert!(services.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kimi_code_export_roundtrip_preserves_model() {
|
||||
let mut providers = IndexMap::new();
|
||||
providers.insert(
|
||||
"glmzhongzhuan".to_string(),
|
||||
Provider {
|
||||
name: "glmzhongzhuan".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
base_url: Some("https://fast.cdks.work".to_string()),
|
||||
api_key: Some("sk-test".to_string()),
|
||||
env: IndexMap::new(),
|
||||
note: None,
|
||||
official_url: None,
|
||||
managed: false,
|
||||
enabled: true,
|
||||
active: true,
|
||||
icon: None,
|
||||
icon_color: None,
|
||||
raw_other: Value::Null,
|
||||
usage_kinds: None,
|
||||
},
|
||||
);
|
||||
let mut models = IndexMap::new();
|
||||
models.insert(
|
||||
"glm-5.2".to_string(),
|
||||
Model {
|
||||
alias: "glm-5.2".to_string(),
|
||||
provider: "glmzhongzhuan".to_string(),
|
||||
model: "glm-5.2".to_string(),
|
||||
max_context_size: 900_000,
|
||||
display_name: None,
|
||||
supports_1m: true,
|
||||
capabilities: vec!["thinking".to_string()],
|
||||
raw_other: Value::Null,
|
||||
},
|
||||
);
|
||||
let config = Config {
|
||||
default_model: Some("glm-5.2".to_string()),
|
||||
providers,
|
||||
models,
|
||||
raw_other: Value::Null,
|
||||
};
|
||||
|
||||
let exported = config_to_kimi_code(&config, None);
|
||||
let root = exported.as_table().unwrap();
|
||||
assert_eq!(
|
||||
root.get("default_model").and_then(|v| v.as_str()),
|
||||
Some("glm-5.2")
|
||||
);
|
||||
|
||||
let providers_table = root.get("providers").unwrap().as_table().unwrap();
|
||||
let provider = providers_table.get("glmzhongzhuan").unwrap().as_table().unwrap();
|
||||
assert_eq!(provider.get("type").and_then(|v| v.as_str()), Some("anthropic"));
|
||||
assert_eq!(provider.get("api_key").and_then(|v| v.as_str()), Some("sk-test"));
|
||||
|
||||
let models_table = root.get("models").unwrap().as_table().unwrap();
|
||||
let model = models_table.get("glm-5.2").unwrap().as_table().unwrap();
|
||||
assert_eq!(model.get("model").and_then(|v| v.as_str()), Some("glm-5.2"));
|
||||
let caps = model.get("capabilities").unwrap().as_array().unwrap();
|
||||
assert!(caps.iter().any(|v| v.as_str() == Some("thinking")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kimi_code_export_does_not_inject_thinking_from_supports_1m() {
|
||||
let mut providers = IndexMap::new();
|
||||
providers.insert(
|
||||
"test-provider".to_string(),
|
||||
Provider {
|
||||
name: "test-provider".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
base_url: Some("https://example.com".to_string()),
|
||||
api_key: Some("sk-test".to_string()),
|
||||
env: IndexMap::new(),
|
||||
note: None,
|
||||
official_url: None,
|
||||
managed: false,
|
||||
enabled: true,
|
||||
active: true,
|
||||
icon: None,
|
||||
icon_color: None,
|
||||
raw_other: Value::Null,
|
||||
usage_kinds: None,
|
||||
},
|
||||
);
|
||||
let mut models = IndexMap::new();
|
||||
models.insert(
|
||||
"big-model".to_string(),
|
||||
Model {
|
||||
alias: "big-model".to_string(),
|
||||
provider: "test-provider".to_string(),
|
||||
model: "big-model".to_string(),
|
||||
max_context_size: 2_000_000,
|
||||
display_name: None,
|
||||
supports_1m: true,
|
||||
capabilities: vec![],
|
||||
raw_other: Value::Null,
|
||||
},
|
||||
);
|
||||
let config = Config {
|
||||
default_model: Some("big-model".to_string()),
|
||||
providers,
|
||||
models,
|
||||
raw_other: Value::Null,
|
||||
};
|
||||
|
||||
let exported = config_to_kimi_code(&config, None);
|
||||
let root = exported.as_table().unwrap();
|
||||
let models_table = root.get("models").unwrap().as_table().unwrap();
|
||||
let model = models_table.get("big-model").unwrap().as_table().unwrap();
|
||||
// supports_1m=true but capabilities empty → must NOT inject "thinking".
|
||||
let caps = model.get("capabilities");
|
||||
match caps {
|
||||
None => {}
|
||||
Some(v) => {
|
||||
let arr = v.as_array().unwrap();
|
||||
assert!(
|
||||
!arr.iter().any(|c| c.as_str() == Some("thinking")),
|
||||
"thinking should not be injected from supports_1m"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kimi_code_oauth_key_roundtrip() {
|
||||
// Regression: importing then exporting managed:kimi-code must preserve
|
||||
// the exact oauth key "oauth/kimi-code". Previously the oauth block
|
||||
// was dropped on import and regenerated on export as
|
||||
// "oauth/managed-kimi-code", breaking the official subscription.
|
||||
let toml_str = r#"
|
||||
default_model = "kimi-code/k3"
|
||||
|
||||
[providers."managed:kimi-code"]
|
||||
type = "kimi"
|
||||
api_key = ""
|
||||
base_url = "https://api.kimi.com/coding/v1"
|
||||
|
||||
[providers."managed:kimi-code".oauth]
|
||||
storage = "file"
|
||||
key = "oauth/kimi-code"
|
||||
|
||||
[models."kimi-code/k3"]
|
||||
provider = "managed:kimi-code"
|
||||
model = "k3"
|
||||
max_context_size = 1048576
|
||||
"#;
|
||||
let value: TomlValue = toml_str.parse().unwrap();
|
||||
let config = kimi_code_to_config(&value);
|
||||
|
||||
// Export back to TOML.
|
||||
let exported = config_to_kimi_code(&config, None);
|
||||
let root = exported.as_table().unwrap();
|
||||
let providers = root.get("providers").unwrap().as_table().unwrap();
|
||||
let managed = providers.get("managed:kimi-code").unwrap().as_table().unwrap();
|
||||
let oauth = managed.get("oauth").unwrap().as_table().unwrap();
|
||||
assert_eq!(
|
||||
oauth.get("key").and_then(|v| v.as_str()),
|
||||
Some("oauth/kimi-code"),
|
||||
"oauth key must round-trip verbatim, not be regenerated"
|
||||
);
|
||||
assert_eq!(
|
||||
oauth.get("storage").and_then(|v| v.as_str()),
|
||||
Some("file")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kimi_code_export_writes_all_providers() {
|
||||
// Regression: inactive providers must still be written to config.toml
|
||||
// so that switching does not wipe the provider list.
|
||||
let mut providers = IndexMap::new();
|
||||
providers.insert(
|
||||
"active-one".to_string(),
|
||||
Provider {
|
||||
name: "active-one".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
base_url: Some("https://a.example.com".to_string()),
|
||||
api_key: Some("sk-a".to_string()),
|
||||
env: IndexMap::new(),
|
||||
note: None,
|
||||
official_url: None,
|
||||
managed: false,
|
||||
enabled: true,
|
||||
active: true,
|
||||
icon: None,
|
||||
icon_color: None,
|
||||
raw_other: Value::Null,
|
||||
usage_kinds: None,
|
||||
},
|
||||
);
|
||||
providers.insert(
|
||||
"inactive-one".to_string(),
|
||||
Provider {
|
||||
name: "inactive-one".to_string(),
|
||||
provider_type: ProviderType::Openai,
|
||||
base_url: Some("https://b.example.com".to_string()),
|
||||
api_key: Some("sk-b".to_string()),
|
||||
env: IndexMap::new(),
|
||||
note: None,
|
||||
official_url: None,
|
||||
managed: false,
|
||||
enabled: true,
|
||||
active: false,
|
||||
icon: None,
|
||||
icon_color: None,
|
||||
raw_other: Value::Null,
|
||||
usage_kinds: None,
|
||||
},
|
||||
);
|
||||
let config = Config {
|
||||
default_model: None,
|
||||
providers,
|
||||
models: IndexMap::new(),
|
||||
raw_other: Value::Null,
|
||||
};
|
||||
|
||||
let exported = config_to_kimi_code(&config, None);
|
||||
let root = exported.as_table().unwrap();
|
||||
let providers_table = root.get("providers").unwrap().as_table().unwrap();
|
||||
assert_eq!(providers_table.len(), 2, "both providers must be written");
|
||||
assert!(providers_table.contains_key("active-one"));
|
||||
assert!(providers_table.contains_key("inactive-one"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kimi_code_export_strips_private_default_model() {
|
||||
// The remembered per-provider default model (raw_other.default_model)
|
||||
// is a Kimi-Switch-private field and must NOT leak into config.toml.
|
||||
let mut providers = IndexMap::new();
|
||||
providers.insert(
|
||||
"p".to_string(),
|
||||
Provider {
|
||||
name: "p".to_string(),
|
||||
provider_type: ProviderType::Anthropic,
|
||||
base_url: None,
|
||||
api_key: Some("sk-x".to_string()),
|
||||
env: IndexMap::new(),
|
||||
note: None,
|
||||
official_url: None,
|
||||
managed: false,
|
||||
enabled: true,
|
||||
active: true,
|
||||
icon: None,
|
||||
icon_color: None,
|
||||
raw_other: serde_json::json!({"default_model": "some-alias"}),
|
||||
usage_kinds: None,
|
||||
},
|
||||
);
|
||||
let config = Config {
|
||||
default_model: None,
|
||||
providers,
|
||||
models: IndexMap::new(),
|
||||
raw_other: Value::Null,
|
||||
};
|
||||
|
||||
let exported = config_to_kimi_code(&config, None);
|
||||
let root = exported.as_table().unwrap();
|
||||
let provider = root
|
||||
.get("providers").unwrap()
|
||||
.as_table().unwrap()
|
||||
.get("p").unwrap()
|
||||
.as_table().unwrap();
|
||||
assert!(
|
||||
!provider.contains_key("default_model"),
|
||||
"private default_model must not leak into config.toml"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,119 @@
|
||||
pub mod commands;
|
||||
pub mod config_io;
|
||||
pub mod dashboard;
|
||||
pub mod db;
|
||||
pub mod kimi_code_io;
|
||||
pub mod models;
|
||||
pub mod pi_io;
|
||||
pub mod services;
|
||||
|
||||
use tauri::menu::{Menu, MenuItem, PredefinedMenuItem};
|
||||
use tauri::tray::{MouseButton, MouseButtonState, TrayIconBuilder};
|
||||
use tauri::Manager;
|
||||
|
||||
pub fn run() {
|
||||
tauri::Builder::default()
|
||||
// Single instance: must be registered before other plugins. When a
|
||||
// second process is launched, this callback runs in the existing
|
||||
// instance and simply brings its window back (including from tray).
|
||||
.plugin(tauri_plugin_single_instance::init(|app, _args, _cwd| {
|
||||
if let Some(window) = app.get_webview_window("main") {
|
||||
let _ = window.unminimize();
|
||||
let _ = window.show();
|
||||
let _ = window.set_focus();
|
||||
}
|
||||
}))
|
||||
.plugin(tauri_plugin_opener::init())
|
||||
.setup(|app| {
|
||||
println!("[Tauri] Setup started");
|
||||
let window = app.get_webview_window("main").unwrap();
|
||||
println!("[Tauri] Window label: {}", window.label());
|
||||
|
||||
// Tray menu: show / separator / quit
|
||||
let show_i = MenuItem::with_id(app, "show", "显示", true, None::<&str>)?;
|
||||
let quit_i = MenuItem::with_id(app, "quit", "退出", true, None::<&str>)?;
|
||||
let menu = Menu::with_items(
|
||||
app,
|
||||
&[&show_i, &PredefinedMenuItem::separator(app)?, &quit_i],
|
||||
)?;
|
||||
|
||||
// Tray icon: reuse the window icon
|
||||
let icon = app.default_window_icon().unwrap().clone();
|
||||
TrayIconBuilder::new()
|
||||
.icon(icon)
|
||||
.menu(&menu)
|
||||
.show_menu_on_left_click(false)
|
||||
.on_menu_event(|app, event| match event.id.as_ref() {
|
||||
"show" => {
|
||||
if let Some(window) = app.get_webview_window("main") {
|
||||
let _ = window.show();
|
||||
let _ = window.unminimize();
|
||||
let _ = window.set_focus();
|
||||
}
|
||||
}
|
||||
"quit" => {
|
||||
app.exit(0);
|
||||
}
|
||||
_ => {}
|
||||
})
|
||||
.on_tray_icon_event(|tray, event| {
|
||||
if let tauri::tray::TrayIconEvent::Click {
|
||||
button,
|
||||
button_state,
|
||||
..
|
||||
} = event
|
||||
{
|
||||
if button == MouseButton::Left && button_state == MouseButtonState::Up {
|
||||
let app = tray.app_handle();
|
||||
if let Some(window) = app.get_webview_window("main") {
|
||||
let _ = window.show();
|
||||
let _ = window.unminimize();
|
||||
let _ = window.set_focus();
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
.build(app)?;
|
||||
|
||||
// Close button -> hide to tray instead of quitting. Minimize keeps
|
||||
// the taskbar button (native behaviour); only an explicit close
|
||||
// (window X) retreats to the tray.
|
||||
let window_clone = window.clone();
|
||||
window.on_window_event(move |event| match event {
|
||||
tauri::WindowEvent::CloseRequested { api, .. } => {
|
||||
api.prevent_close();
|
||||
let _ = window_clone.hide();
|
||||
}
|
||||
_ => {}
|
||||
});
|
||||
|
||||
Ok(())
|
||||
})
|
||||
.invoke_handler(tauri::generate_handler![
|
||||
commands::load_agent_config_command,
|
||||
commands::save_agent_config_command,
|
||||
commands::activate_agent_config_command,
|
||||
commands::open_agent_config_dir,
|
||||
commands::get_app_version,
|
||||
commands::list_provider_models,
|
||||
commands::test_connectivity,
|
||||
commands::query_provider_usage,
|
||||
commands::debug_log,
|
||||
commands::get_app_setting,
|
||||
commands::set_app_setting,
|
||||
commands::check_for_update,
|
||||
commands::download_update,
|
||||
commands::open_installer,
|
||||
dashboard::get_paths,
|
||||
dashboard::get_prices,
|
||||
dashboard::get_summary,
|
||||
dashboard::list_sessions,
|
||||
dashboard::archive_session,
|
||||
dashboard::unarchive_session,
|
||||
dashboard::delete_session,
|
||||
dashboard::delete_workspace,
|
||||
dashboard::get_session_preview,
|
||||
])
|
||||
.run(tauri::generate_context!())
|
||||
.expect("error while running tauri application");
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
#![cfg_attr(not(debug_assertions), windows_subsystem = "windows")]
|
||||
|
||||
fn main() {
|
||||
kimiswitch_lib::run();
|
||||
}
|
||||
@@ -0,0 +1,235 @@
|
||||
use indexmap::IndexMap;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
|
||||
/// Target agent whose provider/model config is being edited.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum Agent {
|
||||
KimiCode,
|
||||
Pi,
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
pub fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
Agent::KimiCode => "kimi_code",
|
||||
Agent::Pi => "pi",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ProviderType {
|
||||
Anthropic,
|
||||
Openai,
|
||||
#[serde(rename = "openai_responses")]
|
||||
OpenaiResponses,
|
||||
#[serde(rename = "google-genai")]
|
||||
GoogleGenai,
|
||||
Vertexai,
|
||||
/// Kept for compatibility; mapped to OpenAI-compatible in Pi.
|
||||
Kimi,
|
||||
}
|
||||
|
||||
const fn default_true() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
impl ProviderType {
|
||||
pub fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
ProviderType::Anthropic => "anthropic",
|
||||
ProviderType::Openai => "openai",
|
||||
ProviderType::OpenaiResponses => "openai_responses",
|
||||
ProviderType::GoogleGenai => "google-genai",
|
||||
ProviderType::Vertexai => "vertexai",
|
||||
ProviderType::Kimi => "kimi",
|
||||
}
|
||||
}
|
||||
|
||||
pub fn default_base_url(&self) -> Option<&'static str> {
|
||||
match self {
|
||||
ProviderType::Openai | ProviderType::OpenaiResponses | ProviderType::Kimi => {
|
||||
Some("https://api.openai.com/v1")
|
||||
}
|
||||
ProviderType::GoogleGenai => Some("https://generativelanguage.googleapis.com"),
|
||||
ProviderType::Anthropic | ProviderType::Vertexai => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_openai_compatible(&self) -> bool {
|
||||
matches!(
|
||||
self,
|
||||
ProviderType::Kimi | ProviderType::Openai | ProviderType::OpenaiResponses
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Serialize, Deserialize)]
|
||||
pub struct Provider {
|
||||
pub name: String,
|
||||
pub provider_type: ProviderType,
|
||||
pub base_url: Option<String>,
|
||||
pub api_key: Option<String>,
|
||||
#[serde(default)]
|
||||
pub env: IndexMap<String, String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub note: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub official_url: Option<String>,
|
||||
#[serde(default)]
|
||||
pub managed: bool,
|
||||
#[serde(default = "default_true")]
|
||||
pub enabled: bool,
|
||||
#[serde(default)]
|
||||
pub active: bool,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub icon: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub icon_color: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Value::is_null")]
|
||||
pub raw_other: Value,
|
||||
/// Billing/usage query kinds for this provider (e.g. "balance:deepseek",
|
||||
/// "plan:kimi_coding"). Persisted in the SQLite settings table under
|
||||
/// `usage_kinds:<provider_name>`, NOT in the agent's config.toml — both
|
||||
/// export paths (kimi_code_io manual TOML, pi_io PiProvider struct) are
|
||||
/// explicit and never serialize this field, while the IPC payload to the
|
||||
/// frontend does carry it.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub usage_kinds: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
impl PartialEq for Provider {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
self.name == other.name
|
||||
&& self.provider_type == other.provider_type
|
||||
&& self.base_url == other.base_url
|
||||
&& self.api_key == other.api_key
|
||||
&& self.env == other.env
|
||||
&& self.note == other.note
|
||||
&& self.official_url == other.official_url
|
||||
&& self.managed == other.managed
|
||||
&& self.enabled == other.enabled
|
||||
&& self.active == other.active
|
||||
&& self.icon == other.icon
|
||||
&& self.icon_color == other.icon_color
|
||||
&& self.raw_other == other.raw_other
|
||||
&& self.usage_kinds == other.usage_kinds
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for Provider {}
|
||||
|
||||
impl std::fmt::Debug for Provider {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("Provider")
|
||||
.field("name", &self.name)
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("base_url", &self.base_url)
|
||||
.field("api_key", &"<redacted>")
|
||||
.field("env", &"<redacted>")
|
||||
.field("managed", &self.managed)
|
||||
.field("enabled", &self.enabled)
|
||||
.field("active", &self.active)
|
||||
.field("icon", &self.icon)
|
||||
.field("icon_color", &self.icon_color)
|
||||
.field("raw_other", &"<json>")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Serialize, Deserialize)]
|
||||
pub struct Model {
|
||||
pub alias: String,
|
||||
pub provider: String,
|
||||
pub model: String,
|
||||
pub max_context_size: u64,
|
||||
pub display_name: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
|
||||
pub supports_1m: bool,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub capabilities: Vec<String>,
|
||||
#[serde(default, skip_serializing_if = "Value::is_null")]
|
||||
pub raw_other: Value,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for Model {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("Model")
|
||||
.field("alias", &self.alias)
|
||||
.field("provider", &self.provider)
|
||||
.field("model", &self.model)
|
||||
.field("max_context_size", &self.max_context_size)
|
||||
.field("display_name", &self.display_name)
|
||||
.field("supports_1m", &self.supports_1m)
|
||||
.field("capabilities", &self.capabilities)
|
||||
.field("raw_other", &"<json>")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialEq for Model {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
self.alias == other.alias
|
||||
&& self.provider == other.provider
|
||||
&& self.model == other.model
|
||||
&& self.max_context_size == other.max_context_size
|
||||
&& self.display_name == other.display_name
|
||||
&& self.supports_1m == other.supports_1m
|
||||
&& self.capabilities == other.capabilities
|
||||
&& self.raw_other == other.raw_other
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for Model {}
|
||||
|
||||
#[derive(Clone, Serialize, Deserialize)]
|
||||
pub struct Config {
|
||||
pub default_model: Option<String>,
|
||||
#[serde(default)]
|
||||
pub providers: IndexMap<String, Provider>,
|
||||
#[serde(default)]
|
||||
pub models: IndexMap<String, Model>,
|
||||
#[serde(default, skip_serializing_if = "Value::is_null")]
|
||||
pub raw_other: Value,
|
||||
}
|
||||
|
||||
impl PartialEq for Config {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
self.default_model == other.default_model
|
||||
&& self.providers == other.providers
|
||||
&& self.models == other.models
|
||||
&& self.raw_other == other.raw_other
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for Config {}
|
||||
|
||||
impl std::fmt::Debug for Config {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("Config")
|
||||
.field("default_model", &self.default_model)
|
||||
.field(
|
||||
"providers",
|
||||
&format!(
|
||||
"{:?} ({} providers)",
|
||||
self.providers.keys().collect::<Vec<_>>(),
|
||||
self.providers.len()
|
||||
),
|
||||
)
|
||||
.field("models", &self.models)
|
||||
.field("raw_other", &"<json>")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// A model discovered from a provider's API endpoint.
|
||||
/// The frontend uses this to populate a `Model` form entry.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct DiscoveredModel {
|
||||
pub id: String,
|
||||
pub display_name: Option<String>,
|
||||
pub max_context_size: Option<u64>,
|
||||
}
|
||||
@@ -0,0 +1,626 @@
|
||||
//! Pi CLI configuration I/O.
|
||||
//!
|
||||
//! Pi stores custom providers and models in `~/.pi/agent/models.json`.
|
||||
//! This module reads and writes that file and converts between Pi's JSON
|
||||
//! format and Kimi Switch's internal `Config`/`Provider`/`Model` types.
|
||||
|
||||
use std::path::PathBuf;
|
||||
|
||||
use anyhow::Context;
|
||||
use indexmap::IndexMap;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::config_io::backup_file;
|
||||
use crate::models::{Config, Model, Provider, ProviderType};
|
||||
|
||||
pub type PiResult<T> = anyhow::Result<T>;
|
||||
|
||||
/// Returns the default Pi agent config directory for the current user.
|
||||
pub fn pi_agent_dir() -> PathBuf {
|
||||
std::env::var_os("PI_CODING_AGENT_DIR")
|
||||
.map(PathBuf::from)
|
||||
.or_else(|| dirs::home_dir().map(|h| h.join(".pi").join("agent")))
|
||||
.expect("failed to resolve Pi agent config directory")
|
||||
}
|
||||
|
||||
/// Returns the path to Pi's `models.json` file.
|
||||
pub fn pi_models_path() -> PathBuf {
|
||||
pi_agent_dir().join("models.json")
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct PiCost {
|
||||
pub input: f64,
|
||||
pub output: f64,
|
||||
#[serde(rename = "cacheRead", default)]
|
||||
pub cache_read: f64,
|
||||
#[serde(rename = "cacheWrite", default)]
|
||||
pub cache_write: f64,
|
||||
}
|
||||
|
||||
/// A single Pi model entry.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct PiModel {
|
||||
pub id: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub name: Option<String>,
|
||||
#[serde(default)]
|
||||
pub reasoning: bool,
|
||||
#[serde(default = "default_input", skip_serializing_if = "is_default_input")]
|
||||
pub input: Vec<String>,
|
||||
#[serde(rename = "contextWindow", default = "default_context_window")]
|
||||
pub context_window: u64,
|
||||
#[serde(rename = "maxTokens", default, skip_serializing_if = "Option::is_none")]
|
||||
pub max_tokens: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cost: Option<PiCost>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub compat: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Value,
|
||||
}
|
||||
|
||||
fn default_input() -> Vec<String> {
|
||||
vec!["text".to_string()]
|
||||
}
|
||||
|
||||
fn is_default_input(input: &[String]) -> bool {
|
||||
input.len() == 1 && input.first().map(String::as_str) == Some("text")
|
||||
}
|
||||
|
||||
fn default_context_window() -> u64 {
|
||||
128_000
|
||||
}
|
||||
|
||||
fn default_true() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
/// A single Pi provider entry.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct PiProvider {
|
||||
#[serde(rename = "baseUrl", default, skip_serializing_if = "Option::is_none")]
|
||||
pub base_url: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub api: Option<String>,
|
||||
#[serde(rename = "apiKey", default, skip_serializing_if = "Option::is_none")]
|
||||
pub api_key: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub headers: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub compat: Option<Value>,
|
||||
#[serde(default = "default_true")]
|
||||
pub enabled: bool,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub models: Vec<PiModel>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Value,
|
||||
}
|
||||
|
||||
/// Root Pi `models.json` structure.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct PiModelsFile {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub default_model: Option<String>,
|
||||
#[serde(default)]
|
||||
pub providers: IndexMap<String, PiProvider>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Value,
|
||||
}
|
||||
|
||||
/// Read the existing Pi `models.json` file, or return a default empty struct
|
||||
/// if the file does not exist yet.
|
||||
pub fn load_pi_models() -> PiResult<PiModelsFile> {
|
||||
let path = pi_models_path();
|
||||
if !path.exists() {
|
||||
return Ok(PiModelsFile::default());
|
||||
}
|
||||
let content = std::fs::read_to_string(&path)
|
||||
.with_context(|| format!("failed to read {}", path.display()))?;
|
||||
let file: PiModelsFile = serde_json::from_str(&content)
|
||||
.with_context(|| format!("failed to parse {}", path.display()))?;
|
||||
Ok(file)
|
||||
}
|
||||
|
||||
/// Write the given Pi models file, creating the parent directory if needed.
|
||||
/// A backup of the previous file is created when it already exists.
|
||||
pub fn save_pi_models(file: &PiModelsFile) -> PiResult<()> {
|
||||
let path = pi_models_path();
|
||||
if let Some(parent) = path.parent() {
|
||||
std::fs::create_dir_all(parent)
|
||||
.with_context(|| format!("failed to create directory {}", parent.display()))?;
|
||||
}
|
||||
|
||||
if path.exists() {
|
||||
backup_file(&path, 7).with_context(|| {
|
||||
format!("failed to back up Pi models {}", path.display())
|
||||
})?;
|
||||
}
|
||||
|
||||
let content = serde_json::to_string_pretty(file)
|
||||
.with_context(|| "failed to serialize Pi models.json")?;
|
||||
std::fs::write(&path, content)
|
||||
.with_context(|| format!("failed to write {}", path.display()))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Returns the path to Pi's `settings.json` file.
|
||||
pub fn pi_settings_path() -> PathBuf {
|
||||
pi_agent_dir().join("settings.json")
|
||||
}
|
||||
|
||||
/// Pi's `settings.json` file (only the fields Kimi Switch manipulates).
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct PiSettingsFile {
|
||||
#[serde(rename = "defaultProvider", default, skip_serializing_if = "Option::is_none")]
|
||||
pub default_provider: Option<String>,
|
||||
#[serde(rename = "defaultModel", default, skip_serializing_if = "Option::is_none")]
|
||||
pub default_model: Option<String>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Value,
|
||||
}
|
||||
|
||||
/// Read the existing Pi `settings.json` file, or return a default empty struct
|
||||
/// if the file does not exist yet.
|
||||
pub fn load_pi_settings() -> PiResult<PiSettingsFile> {
|
||||
let path = pi_settings_path();
|
||||
if !path.exists() {
|
||||
return Ok(PiSettingsFile::default());
|
||||
}
|
||||
let content = std::fs::read_to_string(&path)
|
||||
.with_context(|| format!("failed to read {}", path.display()))?;
|
||||
let file: PiSettingsFile = serde_json::from_str(&content)
|
||||
.with_context(|| format!("failed to parse {}", path.display()))?;
|
||||
Ok(file)
|
||||
}
|
||||
|
||||
/// Write the given Pi settings file, creating the parent directory if needed.
|
||||
/// A backup of the previous file is created when it already exists.
|
||||
pub fn save_pi_settings(file: &PiSettingsFile) -> PiResult<()> {
|
||||
let path = pi_settings_path();
|
||||
if let Some(parent) = path.parent() {
|
||||
std::fs::create_dir_all(parent)
|
||||
.with_context(|| format!("failed to create directory {}", parent.display()))?;
|
||||
}
|
||||
|
||||
if path.exists() {
|
||||
backup_file(&path, 7).with_context(|| {
|
||||
format!("failed to back up Pi settings {}", path.display())
|
||||
})?;
|
||||
}
|
||||
|
||||
let content = serde_json::to_string_pretty(file)
|
||||
.with_context(|| "failed to serialize Pi settings.json")?;
|
||||
std::fs::write(&path, content)
|
||||
.with_context(|| format!("failed to write {}", path.display()))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Convert a Kimi Switch provider type to the Pi `api` field value.
|
||||
pub fn pi_api_for_provider(provider_type: &ProviderType) -> &'static str {
|
||||
match provider_type {
|
||||
ProviderType::Openai => "openai-completions",
|
||||
ProviderType::OpenaiResponses => "openai-responses",
|
||||
ProviderType::Anthropic => "anthropic-messages",
|
||||
ProviderType::GoogleGenai => "google-generative-ai",
|
||||
ProviderType::Vertexai => "google-vertex",
|
||||
// Treat Kimi as OpenAI-compatible since it is not a native Pi API.
|
||||
ProviderType::Kimi => "openai-completions",
|
||||
}
|
||||
}
|
||||
|
||||
/// Convert a Pi `api` field value back to a Kimi Switch provider type.
|
||||
pub fn provider_type_for_pi_api(api: &str) -> ProviderType {
|
||||
match api {
|
||||
"openai-responses" => ProviderType::OpenaiResponses,
|
||||
"anthropic-messages" => ProviderType::Anthropic,
|
||||
"google-generative-ai" => ProviderType::GoogleGenai,
|
||||
"google-vertex" => ProviderType::Vertexai,
|
||||
_ => ProviderType::Openai,
|
||||
}
|
||||
}
|
||||
|
||||
/// Merge provider-level fields that have dedicated `PiProvider` fields into the
|
||||
/// `raw_other` blob so they survive the round-trip through Kimi Switch's internal
|
||||
/// `Config`.
|
||||
fn merge_provider_known_into_raw(pi_provider: &PiProvider, raw: &mut Value) {
|
||||
let has_known = pi_provider.headers.is_some()
|
||||
|| pi_provider.compat.is_some()
|
||||
|| pi_provider.api.is_some();
|
||||
if !has_known {
|
||||
return;
|
||||
}
|
||||
|
||||
if raw.is_null() {
|
||||
*raw = Value::Object(serde_json::Map::new());
|
||||
}
|
||||
if let Some(obj) = raw.as_object_mut() {
|
||||
if let Some(headers) = &pi_provider.headers {
|
||||
obj.insert("headers".to_string(), headers.clone());
|
||||
}
|
||||
if let Some(compat) = &pi_provider.compat {
|
||||
obj.insert("compat".to_string(), compat.clone());
|
||||
}
|
||||
if let Some(api) = &pi_provider.api {
|
||||
obj.insert("api".to_string(), Value::String(api.clone()));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Merge model-level fields that have dedicated `PiModel` fields into the
|
||||
/// `raw_other` blob so they survive the round-trip through Kimi Switch's internal
|
||||
/// `Config`.
|
||||
fn merge_model_known_into_raw(pi_model: &PiModel, raw: &mut Value) {
|
||||
let has_known = pi_model.cost.is_some()
|
||||
|| pi_model.compat.is_some()
|
||||
|| pi_model.max_tokens.is_some();
|
||||
if !has_known {
|
||||
return;
|
||||
}
|
||||
|
||||
if raw.is_null() {
|
||||
*raw = Value::Object(serde_json::Map::new());
|
||||
}
|
||||
if let Some(obj) = raw.as_object_mut() {
|
||||
if let Some(cost) = &pi_model.cost {
|
||||
if let Ok(value) = serde_json::to_value(cost) {
|
||||
obj.insert("cost".to_string(), value);
|
||||
}
|
||||
}
|
||||
if let Some(compat) = &pi_model.compat {
|
||||
obj.insert("compat".to_string(), compat.clone());
|
||||
}
|
||||
if let Some(max_tokens) = pi_model.max_tokens {
|
||||
obj.insert(
|
||||
"maxTokens".to_string(),
|
||||
Value::Number(serde_json::Number::from(max_tokens)),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Extract provider-level fields that PiProvider serializes explicitly from the
|
||||
/// `raw_other` blob. Any extracted keys are removed from the returned extra so
|
||||
/// they are not duplicated by `#[serde(flatten)]`.
|
||||
fn extract_provider_fields(raw: &Value) -> (Option<String>, Option<Value>, Option<Value>, Value) {
|
||||
let mut extra = raw.clone();
|
||||
let api = extra
|
||||
.as_object_mut()
|
||||
.and_then(|o| o.remove("api"))
|
||||
.and_then(|v| v.as_str().map(String::from));
|
||||
let headers = extra.as_object_mut().and_then(|o| o.remove("headers"));
|
||||
let compat = extra.as_object_mut().and_then(|o| o.remove("compat"));
|
||||
(api, headers, compat, extra)
|
||||
}
|
||||
|
||||
/// Extract model-level fields that PiModel serializes explicitly from the
|
||||
/// `raw_other` blob. Any extracted keys are removed from the returned extra.
|
||||
fn extract_model_fields(raw: &Value) -> (Option<PiCost>, Option<Value>, Option<u64>, Value) {
|
||||
let mut extra = raw.clone();
|
||||
let cost = extra
|
||||
.as_object_mut()
|
||||
.and_then(|o| o.remove("cost"))
|
||||
.and_then(|v| serde_json::from_value(v).ok());
|
||||
let compat = extra.as_object_mut().and_then(|o| o.remove("compat"));
|
||||
let max_tokens = extra
|
||||
.as_object_mut()
|
||||
.and_then(|o| o.remove("maxTokens"))
|
||||
.and_then(|v| v.as_u64());
|
||||
(cost, compat, max_tokens, extra)
|
||||
}
|
||||
|
||||
/// Import a Pi `models.json` into a Kimi Switch `Config`.
|
||||
pub fn pi_file_to_config(file: &PiModelsFile) -> Config {
|
||||
let mut providers = IndexMap::new();
|
||||
let mut models = IndexMap::new();
|
||||
|
||||
for (name, pi_provider) in &file.providers {
|
||||
let provider_type = pi_provider
|
||||
.api
|
||||
.as_deref()
|
||||
.map(provider_type_for_pi_api)
|
||||
.unwrap_or(ProviderType::Openai);
|
||||
|
||||
let mut provider_raw = pi_provider.extra.clone();
|
||||
merge_provider_known_into_raw(pi_provider, &mut provider_raw);
|
||||
|
||||
let provider = Provider {
|
||||
name: name.clone(),
|
||||
provider_type,
|
||||
base_url: pi_provider.base_url.clone().filter(|s| !s.is_empty()),
|
||||
api_key: pi_provider.api_key.clone().filter(|s| !s.is_empty()),
|
||||
env: IndexMap::new(),
|
||||
note: None,
|
||||
official_url: None,
|
||||
managed: false,
|
||||
enabled: pi_provider.enabled,
|
||||
active: true,
|
||||
icon: None,
|
||||
icon_color: None,
|
||||
raw_other: provider_raw,
|
||||
usage_kinds: None,
|
||||
};
|
||||
|
||||
for (idx, pi_model) in pi_provider.models.iter().enumerate() {
|
||||
let alias = pi_model
|
||||
.name
|
||||
.clone()
|
||||
.filter(|s| !s.is_empty())
|
||||
.unwrap_or_else(|| {
|
||||
if pi_model.id.trim().is_empty() {
|
||||
format!("{}-model-{}", name, idx + 1)
|
||||
} else {
|
||||
pi_model.id.clone()
|
||||
}
|
||||
});
|
||||
let display_name = pi_model.name.clone().filter(|s| !s.is_empty());
|
||||
|
||||
let mut model_raw = pi_model.extra.clone();
|
||||
merge_model_known_into_raw(pi_model, &mut model_raw);
|
||||
|
||||
models.insert(
|
||||
alias.clone(),
|
||||
Model {
|
||||
alias,
|
||||
provider: name.clone(),
|
||||
model: pi_model.id.clone(),
|
||||
max_context_size: pi_model.context_window,
|
||||
display_name,
|
||||
supports_1m: pi_model.reasoning || pi_model.context_window >= 1_000_000,
|
||||
capabilities: vec![],
|
||||
raw_other: model_raw,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
providers.insert(name.clone(), provider);
|
||||
}
|
||||
|
||||
Config {
|
||||
default_model: file.default_model.clone(),
|
||||
providers,
|
||||
models,
|
||||
raw_other: file.extra.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Export a Kimi Switch `Config` to a Pi `models.json`.
|
||||
pub fn config_to_pi_file(config: &Config) -> PiModelsFile {
|
||||
let mut providers = IndexMap::new();
|
||||
|
||||
for (name, provider) in &config.providers {
|
||||
let (provider_api, provider_headers, provider_compat, provider_extra) =
|
||||
extract_provider_fields(&provider.raw_other);
|
||||
let api = provider_api.unwrap_or_else(|| pi_api_for_provider(&provider.provider_type).to_string());
|
||||
|
||||
let mut pi_models = Vec::new();
|
||||
for model in config.models.values() {
|
||||
if model.provider != *name {
|
||||
continue;
|
||||
}
|
||||
let display_name = model.display_name.clone().filter(|s| !s.is_empty());
|
||||
let name_field = display_name.clone().filter(|s| s != &model.model);
|
||||
let (model_cost, model_compat, model_max_tokens, model_extra) =
|
||||
extract_model_fields(&model.raw_other);
|
||||
|
||||
pi_models.push(PiModel {
|
||||
id: model.model.clone(),
|
||||
name: name_field,
|
||||
reasoning: model.supports_1m,
|
||||
input: default_input(),
|
||||
context_window: model.max_context_size,
|
||||
max_tokens: model_max_tokens,
|
||||
cost: model_cost,
|
||||
compat: model_compat,
|
||||
extra: model_extra,
|
||||
});
|
||||
}
|
||||
|
||||
providers.insert(
|
||||
name.clone(),
|
||||
PiProvider {
|
||||
base_url: provider.base_url.clone().filter(|s| !s.is_empty()),
|
||||
api: Some(api),
|
||||
api_key: provider.api_key.clone().filter(|s| !s.is_empty()),
|
||||
headers: provider_headers,
|
||||
compat: provider_compat,
|
||||
enabled: provider.enabled,
|
||||
models: pi_models,
|
||||
extra: provider_extra,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
PiModelsFile {
|
||||
// Pi does not use a top-level default_model field; defaults live in
|
||||
// ~/.pi/agent/settings.json (defaultProvider / defaultModel).
|
||||
default_model: None,
|
||||
providers,
|
||||
extra: config.raw_other.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn pi_roundtrip_preserves_provider_and_model() {
|
||||
let mut providers = IndexMap::new();
|
||||
providers.insert(
|
||||
"my-openai".to_string(),
|
||||
PiProvider {
|
||||
base_url: Some("https://proxy.example.com/v1".to_string()),
|
||||
api: Some("openai-completions".to_string()),
|
||||
api_key: Some("sk-test".to_string()),
|
||||
headers: None,
|
||||
compat: None,
|
||||
enabled: true,
|
||||
models: vec![PiModel {
|
||||
id: "gpt-test".to_string(),
|
||||
name: Some("GPT Test".to_string()),
|
||||
reasoning: true,
|
||||
input: vec!["text".to_string()],
|
||||
context_window: 200_000,
|
||||
max_tokens: Some(4096),
|
||||
cost: Some(PiCost {
|
||||
input: 1.0,
|
||||
output: 2.0,
|
||||
cache_read: 0.0,
|
||||
cache_write: 0.0,
|
||||
}),
|
||||
compat: None,
|
||||
extra: Value::Null,
|
||||
}],
|
||||
extra: Value::Null,
|
||||
},
|
||||
);
|
||||
let file = PiModelsFile {
|
||||
default_model: None,
|
||||
providers,
|
||||
extra: Value::Null,
|
||||
};
|
||||
|
||||
let config = pi_file_to_config(&file);
|
||||
assert_eq!(config.providers.len(), 1);
|
||||
assert_eq!(config.models.len(), 1);
|
||||
|
||||
let provider = config.providers.get("my-openai").unwrap();
|
||||
assert_eq!(provider.provider_type, ProviderType::Openai);
|
||||
assert_eq!(provider.base_url.as_deref(), Some("https://proxy.example.com/v1"));
|
||||
assert_eq!(provider.api_key.as_deref(), Some("sk-test"));
|
||||
|
||||
let model = config.models.get("GPT Test").unwrap();
|
||||
assert_eq!(model.model, "gpt-test");
|
||||
assert!(model.supports_1m);
|
||||
assert_eq!(model.max_context_size, 200_000);
|
||||
|
||||
let exported = config_to_pi_file(&config);
|
||||
let exported_provider = exported.providers.get("my-openai").unwrap();
|
||||
assert_eq!(exported_provider.api.as_deref(), Some("openai-completions"));
|
||||
assert_eq!(exported_provider.models[0].id, "gpt-test");
|
||||
assert!(exported_provider.models[0].reasoning);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pi_roundtrip_preserves_advanced_fields() {
|
||||
let mut providers = IndexMap::new();
|
||||
providers.insert(
|
||||
"my-proxy".to_string(),
|
||||
PiProvider {
|
||||
base_url: Some("https://proxy.example.com/v1".to_string()),
|
||||
api: Some("anthropic-messages".to_string()),
|
||||
api_key: Some("sk-test".to_string()),
|
||||
headers: Some(serde_json::json!({"X-Custom": "value"})),
|
||||
compat: Some(serde_json::json!({"supportsDeveloperRole": false})),
|
||||
enabled: true,
|
||||
models: vec![PiModel {
|
||||
id: "custom-claude".to_string(),
|
||||
name: Some("Custom Claude".to_string()),
|
||||
reasoning: true,
|
||||
input: vec!["text".to_string(), "image".to_string()],
|
||||
context_window: 200_000,
|
||||
max_tokens: Some(4096),
|
||||
cost: Some(PiCost {
|
||||
input: 1.0,
|
||||
output: 2.0,
|
||||
cache_read: 0.5,
|
||||
cache_write: 0.75,
|
||||
}),
|
||||
compat: Some(serde_json::json!({"forceAdaptiveThinking": true})),
|
||||
extra: serde_json::json!({
|
||||
"thinkingLevelMap": {
|
||||
"off": null,
|
||||
"minimal": null,
|
||||
"low": "low",
|
||||
"medium": "medium",
|
||||
"high": "high",
|
||||
"xhigh": "max"
|
||||
},
|
||||
"headers": {"X-Model-Header": "model-value"},
|
||||
"api": "anthropic-messages"
|
||||
}),
|
||||
}],
|
||||
extra: serde_json::json!({
|
||||
"name": "My Proxy",
|
||||
"authHeader": true,
|
||||
"modelOverrides": {
|
||||
"claude-sonnet-4": {"name": "Overridden"}
|
||||
}
|
||||
}),
|
||||
},
|
||||
);
|
||||
let file = PiModelsFile {
|
||||
default_model: None,
|
||||
providers,
|
||||
extra: Value::Null,
|
||||
};
|
||||
|
||||
let config = pi_file_to_config(&file);
|
||||
let exported = config_to_pi_file(&config);
|
||||
let provider = exported.providers.get("my-proxy").unwrap();
|
||||
|
||||
// Provider-level fields
|
||||
assert_eq!(
|
||||
provider.extra.get("name").and_then(|v| v.as_str()),
|
||||
Some("My Proxy")
|
||||
);
|
||||
assert_eq!(
|
||||
provider.extra.get("authHeader").and_then(|v| v.as_bool()),
|
||||
Some(true)
|
||||
);
|
||||
assert_eq!(
|
||||
provider
|
||||
.extra
|
||||
.get("modelOverrides")
|
||||
.and_then(|v| v.as_object())
|
||||
.and_then(|o| o.get("claude-sonnet-4"))
|
||||
.and_then(|v| v.get("name")),
|
||||
Some(&serde_json::json!("Overridden"))
|
||||
);
|
||||
assert_eq!(
|
||||
provider.headers.as_ref().and_then(|h| h.get("X-Custom")),
|
||||
Some(&serde_json::json!("value"))
|
||||
);
|
||||
assert_eq!(
|
||||
provider.compat.as_ref().and_then(|c| c.get("supportsDeveloperRole")),
|
||||
Some(&serde_json::json!(false))
|
||||
);
|
||||
|
||||
// Model-level fields
|
||||
let model = &provider.models[0];
|
||||
assert_eq!(model.id, "custom-claude");
|
||||
assert_eq!(model.max_tokens, Some(4096));
|
||||
assert_eq!(model.cost.as_ref().map(|c| c.output), Some(2.0));
|
||||
assert_eq!(
|
||||
model.compat.as_ref().and_then(|c| c.get("forceAdaptiveThinking")),
|
||||
Some(&serde_json::json!(true))
|
||||
);
|
||||
assert_eq!(
|
||||
model.extra.get("thinkingLevelMap").and_then(|v| v.get("xhigh")),
|
||||
Some(&serde_json::json!("max"))
|
||||
);
|
||||
assert_eq!(
|
||||
model.extra.get("headers").and_then(|v| v.get("X-Model-Header")),
|
||||
Some(&serde_json::json!("model-value"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pi_settings_roundtrip() {
|
||||
let file = PiSettingsFile {
|
||||
default_provider: Some("my-proxy".to_string()),
|
||||
default_model: Some("glm-5.2".to_string()),
|
||||
extra: serde_json::json!({"theme": "dark"}),
|
||||
};
|
||||
let serialized = serde_json::to_string_pretty(&file).unwrap();
|
||||
assert!(serialized.contains("defaultProvider"));
|
||||
assert!(serialized.contains("defaultModel"));
|
||||
|
||||
let deserialized: PiSettingsFile = serde_json::from_str(&serialized).unwrap();
|
||||
assert_eq!(deserialized.default_provider, file.default_provider);
|
||||
assert_eq!(deserialized.default_model, file.default_model);
|
||||
assert_eq!(deserialized.extra.get("theme"), Some(&serde_json::json!("dark")));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
//! Profile management for KimiSwitch.
|
||||
//! Loads, saves, and switches between configuration profiles.
|
||||
@@ -0,0 +1,326 @@
|
||||
// Adapted from cc-switch (MIT, © Jason Young)
|
||||
// https://github.com/farion1231/cc-switch
|
||||
|
||||
//! 供应商余额查询服务
|
||||
//!
|
||||
//! 支持 DeepSeek、StepFun、SiliconFlow、OpenRouter、Novita AI 的账户余额查询。
|
||||
//!
|
||||
//! 错误通道语义(与 cc-switch 一致):
|
||||
//! - `Err(String)` = 瞬时传输失败(网络不可达/超时/读体中断)。前端 invoke reject,
|
||||
//! 触发 retry 并保留上一次成功的 data(keep-last-good)。
|
||||
//! - `Ok(success:false)` = 确定性失败(空 key/鉴权失败/非 2xx/响应体非法 JSON),
|
||||
//! 直接透出错误文案。
|
||||
//!
|
||||
//! HTTP 调用与 JSON→UsageData 解析拆成纯函数,便于无 mock 单元测试。
|
||||
|
||||
use super::usage_types::{UsageData, UsageResult};
|
||||
use std::time::Duration;
|
||||
|
||||
const REQUEST_TIMEOUT: Duration = Duration::from_secs(8);
|
||||
|
||||
/// 鉴权头形式:绝大多数供应商用 `Bearer <key>`;智谱套餐接口不加前缀(见 coding_plan)。
|
||||
pub(crate) enum AuthStyle {
|
||||
Bearer,
|
||||
Raw,
|
||||
}
|
||||
|
||||
/// GET 请求的归类结果。
|
||||
pub(crate) enum Fetched {
|
||||
/// 2xx 且响应体是合法 JSON。
|
||||
Body(serde_json::Value),
|
||||
/// 确定性失败(401/403/非 2xx/解析失败),调用方原样包成 Ok 返回。
|
||||
Failed(UsageResult),
|
||||
}
|
||||
|
||||
/// 统一的 GET + JSON 读取助手。瞬时失败(网络/超时/读体中断)返回 `Err`;
|
||||
/// 确定性失败收进 `Fetched::Failed`。
|
||||
///
|
||||
/// 先 `bytes()` 再解析:读体失败(超时/连接中断)是瞬时 → Err;拿到完整响应体
|
||||
/// 后解析失败才是确定性。reqwest 的 `.json()` 把读体错误也包成 decode,无法区分。
|
||||
pub(crate) async fn get_json(
|
||||
url: &str,
|
||||
api_key: &str,
|
||||
auth: AuthStyle,
|
||||
) -> Result<Fetched, String> {
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(REQUEST_TIMEOUT)
|
||||
.build()
|
||||
.map_err(|e| format!("Failed to build HTTP client: {e}"))?;
|
||||
|
||||
// 注意:api_key 只允许进请求头,严禁拼进 URL / 日志 / 错误信息。
|
||||
let req = client.get(url).header("Accept", "application/json");
|
||||
let req = match auth {
|
||||
AuthStyle::Bearer => req.header("Authorization", format!("Bearer {api_key}")),
|
||||
AuthStyle::Raw => req.header("Authorization", api_key),
|
||||
};
|
||||
|
||||
let resp = match req.send().await {
|
||||
Ok(r) => r,
|
||||
Err(e) => return Err(format!("Network error: {e}")),
|
||||
};
|
||||
|
||||
let status = resp.status();
|
||||
if status == reqwest::StatusCode::UNAUTHORIZED || status == reqwest::StatusCode::FORBIDDEN {
|
||||
return Ok(Fetched::Failed(UsageResult::failure(format!(
|
||||
"Authentication failed (HTTP {status})"
|
||||
))));
|
||||
}
|
||||
if !status.is_success() {
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
return Ok(Fetched::Failed(UsageResult::failure(format!(
|
||||
"API error (HTTP {status}): {body}"
|
||||
))));
|
||||
}
|
||||
|
||||
let raw = match resp.bytes().await {
|
||||
Ok(b) => b,
|
||||
Err(e) => return Err(format!("Failed to read response: {e}")),
|
||||
};
|
||||
match serde_json::from_slice(&raw) {
|
||||
Ok(v) => Ok(Fetched::Body(v)),
|
||||
Err(e) => Ok(Fetched::Failed(UsageResult::failure(format!(
|
||||
"Failed to parse response: {e}"
|
||||
)))),
|
||||
}
|
||||
}
|
||||
|
||||
/// 解析 JSON 字段为 f64,兼容数字和字符串格式。
|
||||
pub(crate) fn parse_f64_field(obj: &serde_json::Value, field: &str) -> Option<f64> {
|
||||
obj.get(field).and_then(|v| {
|
||||
v.as_f64()
|
||||
.or_else(|| v.as_str().and_then(|s| s.parse().ok()))
|
||||
})
|
||||
}
|
||||
|
||||
// ── DeepSeek ────────────────────────────────────────────────
|
||||
// GET https://api.deepseek.com/user/balance
|
||||
// Response: { balance_infos: [{ currency, total_balance, ... }], is_available }
|
||||
|
||||
pub async fn query_deepseek(api_key: &str) -> Result<UsageResult, String> {
|
||||
match get_json("https://api.deepseek.com/user/balance", api_key, AuthStyle::Bearer).await? {
|
||||
Fetched::Body(body) => Ok(UsageResult::ok(parse_deepseek(&body))),
|
||||
Fetched::Failed(err) => Ok(err),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_deepseek(body: &serde_json::Value) -> Vec<UsageData> {
|
||||
let is_available = body
|
||||
.get("is_available")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(true);
|
||||
let mut data = Vec::new();
|
||||
|
||||
if let Some(infos) = body.get("balance_infos").and_then(|v| v.as_array()) {
|
||||
for info in infos {
|
||||
let currency = info
|
||||
.get("currency")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("CNY");
|
||||
data.push(UsageData {
|
||||
plan_name: Some(currency.to_string()),
|
||||
remaining: parse_f64_field(info, "total_balance"),
|
||||
is_valid: Some(is_available),
|
||||
unit: Some(currency.to_string()),
|
||||
..Default::default()
|
||||
});
|
||||
}
|
||||
}
|
||||
data
|
||||
}
|
||||
|
||||
// ── StepFun ─────────────────────────────────────────────────
|
||||
// GET https://api.stepfun.com/v1/accounts
|
||||
// Response: { object, type, balance, total_cash_balance, total_voucher_balance }
|
||||
|
||||
pub async fn query_stepfun(api_key: &str) -> Result<UsageResult, String> {
|
||||
match get_json("https://api.stepfun.com/v1/accounts", api_key, AuthStyle::Bearer).await? {
|
||||
Fetched::Body(body) => Ok(UsageResult::ok(parse_stepfun(&body))),
|
||||
Fetched::Failed(err) => Ok(err),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_stepfun(body: &serde_json::Value) -> Vec<UsageData> {
|
||||
vec![UsageData {
|
||||
plan_name: Some("StepFun".to_string()),
|
||||
remaining: Some(parse_f64_field(body, "balance").unwrap_or(0.0)),
|
||||
unit: Some("CNY".to_string()),
|
||||
is_valid: Some(true),
|
||||
..Default::default()
|
||||
}]
|
||||
}
|
||||
|
||||
// ── SiliconFlow ─────────────────────────────────────────────
|
||||
// GET https://api.siliconflow.cn/v1/user/info (.cn;海外站 .com 单位 USD)
|
||||
// Response: { code, data: { balance, chargeBalance, totalBalance, status } }
|
||||
|
||||
pub async fn query_siliconflow(api_key: &str, is_cn: bool) -> Result<UsageResult, String> {
|
||||
let domain = if is_cn {
|
||||
"api.siliconflow.cn"
|
||||
} else {
|
||||
"api.siliconflow.com"
|
||||
};
|
||||
let url = format!("https://{domain}/v1/user/info");
|
||||
match get_json(&url, api_key, AuthStyle::Bearer).await? {
|
||||
Fetched::Body(body) => Ok(match parse_siliconflow(&body, is_cn) {
|
||||
Ok(data) => UsageResult::ok(data),
|
||||
Err(err) => err,
|
||||
}),
|
||||
Fetched::Failed(err) => Ok(err),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_siliconflow(body: &serde_json::Value, is_cn: bool) -> Result<Vec<UsageData>, UsageResult> {
|
||||
let data = match body.get("data") {
|
||||
Some(d) => d,
|
||||
None => {
|
||||
return Err(UsageResult::failure(
|
||||
"Missing 'data' field in response".to_string(),
|
||||
))
|
||||
}
|
||||
};
|
||||
let (plan_name, unit) = if is_cn {
|
||||
("SiliconFlow", "CNY")
|
||||
} else {
|
||||
("SiliconFlow (EN)", "USD")
|
||||
};
|
||||
Ok(vec![UsageData {
|
||||
plan_name: Some(plan_name.to_string()),
|
||||
remaining: Some(parse_f64_field(data, "totalBalance").unwrap_or(0.0)),
|
||||
unit: Some(unit.to_string()),
|
||||
is_valid: Some(true),
|
||||
..Default::default()
|
||||
}])
|
||||
}
|
||||
|
||||
// ── OpenRouter ──────────────────────────────────────────────
|
||||
// GET https://openrouter.ai/api/v1/credits
|
||||
// Response: { data: { total_credits, total_usage } }
|
||||
|
||||
pub async fn query_openrouter(api_key: &str) -> Result<UsageResult, String> {
|
||||
match get_json(
|
||||
"https://openrouter.ai/api/v1/credits",
|
||||
api_key,
|
||||
AuthStyle::Bearer,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Fetched::Body(body) => Ok(UsageResult::ok(parse_openrouter(&body))),
|
||||
Fetched::Failed(err) => Ok(err),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_openrouter(body: &serde_json::Value) -> Vec<UsageData> {
|
||||
let data = body.get("data").unwrap_or(body);
|
||||
let total_credits = parse_f64_field(data, "total_credits").unwrap_or(0.0);
|
||||
let total_usage = parse_f64_field(data, "total_usage").unwrap_or(0.0);
|
||||
let remaining = total_credits - total_usage;
|
||||
|
||||
vec![UsageData {
|
||||
plan_name: Some("OpenRouter".to_string()),
|
||||
remaining: Some(remaining),
|
||||
total: Some(total_credits),
|
||||
used: Some(total_usage),
|
||||
unit: Some("USD".to_string()),
|
||||
is_valid: Some(remaining > 0.0),
|
||||
..Default::default()
|
||||
}]
|
||||
}
|
||||
|
||||
// ── Novita AI ───────────────────────────────────────────────
|
||||
// GET https://api.novita.ai/v3/user/balance
|
||||
// Response: { availableBalance, ... };金额单位 0.0001 USD,需 /10000。
|
||||
|
||||
pub async fn query_novita(api_key: &str) -> Result<UsageResult, String> {
|
||||
match get_json(
|
||||
"https://api.novita.ai/v3/user/balance",
|
||||
api_key,
|
||||
AuthStyle::Bearer,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Fetched::Body(body) => Ok(UsageResult::ok(parse_novita(&body))),
|
||||
Fetched::Failed(err) => Ok(err),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_novita(body: &serde_json::Value) -> Vec<UsageData> {
|
||||
let available = parse_f64_field(body, "availableBalance").unwrap_or(0.0) / 10000.0;
|
||||
vec![UsageData {
|
||||
plan_name: Some("Novita AI".to_string()),
|
||||
remaining: Some(available),
|
||||
unit: Some("USD".to_string()),
|
||||
is_valid: Some(available > 0.0),
|
||||
..Default::default()
|
||||
}]
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn deepseek_maps_balance_infos() {
|
||||
let body = json!({
|
||||
"is_available": true,
|
||||
"balance_infos": [
|
||||
{ "currency": "CNY", "total_balance": "12.34" },
|
||||
{ "currency": "USD", "total_balance": 5.0 }
|
||||
]
|
||||
});
|
||||
let data = parse_deepseek(&body);
|
||||
assert_eq!(data.len(), 2);
|
||||
assert_eq!(data[0].plan_name.as_deref(), Some("CNY"));
|
||||
assert_eq!(data[0].remaining, Some(12.34));
|
||||
assert_eq!(data[0].unit.as_deref(), Some("CNY"));
|
||||
assert_eq!(data[0].is_valid, Some(true));
|
||||
assert_eq!(data[1].remaining, Some(5.0));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stepfun_reads_balance() {
|
||||
let body = json!({ "balance": 88.5 });
|
||||
let data = parse_stepfun(&body);
|
||||
assert_eq!(data[0].remaining, Some(88.5));
|
||||
assert_eq!(data[0].unit.as_deref(), Some("CNY"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn siliconflow_missing_data_is_deterministic_failure() {
|
||||
let body = json!({ "code": 500 });
|
||||
let err = parse_siliconflow(&body, true).unwrap_err();
|
||||
assert!(!err.success);
|
||||
assert!(err.error.unwrap().contains("Missing 'data'"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn siliconflow_cn_and_en_units() {
|
||||
let body = json!({ "data": { "totalBalance": "42.0" } });
|
||||
let cn = parse_siliconflow(&body, true).unwrap();
|
||||
let en = parse_siliconflow(&body, false).unwrap();
|
||||
assert_eq!(cn[0].unit.as_deref(), Some("CNY"));
|
||||
assert_eq!(en[0].unit.as_deref(), Some("USD"));
|
||||
assert_eq!(cn[0].remaining, Some(42.0));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openrouter_computes_remaining_credits() {
|
||||
let body = json!({ "data": { "total_credits": 20.0, "total_usage": 7.5 } });
|
||||
let data = parse_openrouter(&body);
|
||||
assert_eq!(data[0].remaining, Some(12.5));
|
||||
assert_eq!(data[0].total, Some(20.0));
|
||||
assert_eq!(data[0].used, Some(7.5));
|
||||
assert_eq!(data[0].is_valid, Some(true));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn novita_converts_ten_thousandth_usd() {
|
||||
let body = json!({ "availableBalance": 123400 });
|
||||
let data = parse_novita(&body);
|
||||
assert_eq!(data[0].remaining, Some(12.34));
|
||||
assert_eq!(data[0].unit.as_deref(), Some("USD"));
|
||||
|
||||
let zero = parse_novita(&json!({ "availableBalance": 0 }));
|
||||
assert_eq!(zero[0].is_valid, Some(false));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,414 @@
|
||||
// Adapted from cc-switch (MIT, © Jason Young)
|
||||
// https://github.com/farion1231/cc-switch
|
||||
|
||||
//! Token Plan 套餐额度查询服务
|
||||
//!
|
||||
//! 支持 Kimi For Coding、智谱 GLM、MiniMax 的套餐额度查询。
|
||||
//! cc-switch 的 SubscriptionQuota/tiers 结构在此展平为 `Vec<UsageData>`:
|
||||
//! 每个窗口(tier)一条 UsageData,`plan_name` = tier 名("five_hour" /
|
||||
//! "weekly_limit"),`used` = 已用百分比(0-100),`total` = 100,
|
||||
//! `remaining` = 剩余百分比,`resets_at` 为 ISO 8601 字符串。
|
||||
//!
|
||||
//! 错误通道语义与 balance.rs 一致(Err = 瞬时,Ok(success:false) = 确定性)。
|
||||
|
||||
use super::balance::{get_json, AuthStyle, Fetched};
|
||||
use super::usage_types::{UsageData, UsageResult};
|
||||
|
||||
const TIER_FIVE_HOUR: &str = "five_hour";
|
||||
const TIER_WEEKLY_LIMIT: &str = "weekly_limit";
|
||||
|
||||
/// 套餐条目的统一构造:按百分比表示用量。
|
||||
fn percent_tier(name: &str, used_percent: f64, resets_at: Option<String>) -> UsageData {
|
||||
UsageData {
|
||||
plan_name: Some(name.to_string()),
|
||||
remaining: Some((100.0 - used_percent).max(0.0)),
|
||||
total: Some(100.0),
|
||||
used: Some(used_percent),
|
||||
unit: Some("%".to_string()),
|
||||
is_valid: Some(true),
|
||||
resets_at,
|
||||
}
|
||||
}
|
||||
|
||||
fn millis_to_iso8601(ms: i64) -> Option<String> {
|
||||
let secs = ms / 1000;
|
||||
let nsecs = ((ms % 1000) * 1_000_000) as u32;
|
||||
chrono::DateTime::from_timestamp(secs, nsecs).map(|dt| dt.to_rfc3339())
|
||||
}
|
||||
|
||||
/// 从 JSON 值提取重置时间,兼容字符串和数字格式:
|
||||
/// - 字符串:直接返回(视为 ISO 8601)
|
||||
/// - 数字:自动判断秒/毫秒并转为 ISO 8601;0/负值视为无重置时间
|
||||
fn extract_reset_time(value: &serde_json::Value) -> Option<String> {
|
||||
if let Some(s) = value.as_str() {
|
||||
return Some(s.to_string());
|
||||
}
|
||||
if let Some(n) = value.as_i64() {
|
||||
if n <= 0 {
|
||||
return None;
|
||||
}
|
||||
// 秒级时间戳 < 1e12,毫秒 >= 1e12
|
||||
let ms = if n < 1_000_000_000_000 { n * 1000 } else { n };
|
||||
return millis_to_iso8601(ms);
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
/// 解析 JSON 值为 f64,兼容数字和字符串格式(如 `100` 和 `"100"`)
|
||||
fn parse_f64(value: &serde_json::Value) -> Option<f64> {
|
||||
value
|
||||
.as_f64()
|
||||
.or_else(|| value.as_str().and_then(|s| s.parse().ok()))
|
||||
}
|
||||
|
||||
// ── Kimi For Coding ─────────────────────────────────────────
|
||||
// GET https://api.kimi.com/coding/v1/usages
|
||||
// Response: { limits: [{ detail: { limit, remaining, resetTime } }],
|
||||
// usage: { limit, remaining, resetTime } }
|
||||
|
||||
pub async fn query_kimi_coding(api_key: &str) -> Result<UsageResult, String> {
|
||||
match get_json(
|
||||
"https://api.kimi.com/coding/v1/usages",
|
||||
api_key,
|
||||
AuthStyle::Bearer,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Fetched::Body(body) => Ok(UsageResult::ok(parse_kimi_coding(&body))),
|
||||
Fetched::Failed(err) => Ok(err),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_kimi_coding(body: &serde_json::Value) -> Vec<UsageData> {
|
||||
let mut tiers = Vec::new();
|
||||
|
||||
// 5 小时窗口限额(优先显示)
|
||||
if let Some(limits) = body.get("limits").and_then(|v| v.as_array()) {
|
||||
for limit_item in limits {
|
||||
if let Some(detail) = limit_item.get("detail") {
|
||||
tiers.push(kimi_limit_tier(TIER_FIVE_HOUR, detail));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 总体用量(周限额)
|
||||
if let Some(usage) = body.get("usage") {
|
||||
tiers.push(kimi_limit_tier(TIER_WEEKLY_LIMIT, usage));
|
||||
}
|
||||
|
||||
tiers
|
||||
}
|
||||
|
||||
fn kimi_limit_tier(name: &str, detail: &serde_json::Value) -> UsageData {
|
||||
let limit = detail.get("limit").and_then(parse_f64).unwrap_or(1.0);
|
||||
let remaining = detail.get("remaining").and_then(parse_f64).unwrap_or(0.0);
|
||||
let resets_at = detail.get("resetTime").and_then(extract_reset_time);
|
||||
|
||||
let used = (limit - remaining).max(0.0);
|
||||
let utilization = if limit > 0.0 {
|
||||
(used / limit) * 100.0
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
percent_tier(name, utilization, resets_at)
|
||||
}
|
||||
|
||||
// ── 智谱 GLM ────────────────────────────────────────────────
|
||||
// GET {open.bigmodel.cn | api.z.ai}/api/monitor/usage/quota/limit
|
||||
// 注意:智谱鉴权不加 Bearer 前缀(cc-switch 实测行为,照搬)。
|
||||
|
||||
/// 智谱 TOKENS_LIMIT 条目按 `unit` 字段的显式窗口分类。
|
||||
/// 实测:`unit: 3` → 5 小时滚动窗口;`unit: 6` → 每周窗口。
|
||||
/// 缺失或不识别时走重置时间启发式兜底。
|
||||
fn parse_zhipu(body: &serde_json::Value) -> Result<Vec<UsageData>, UsageResult> {
|
||||
// 业务级别错误
|
||||
if body.get("success").and_then(|v| v.as_bool()) == Some(false) {
|
||||
let msg = body
|
||||
.get("msg")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("Unknown error");
|
||||
return Err(UsageResult::failure(format!("API error: {msg}")));
|
||||
}
|
||||
|
||||
let data = match body.get("data") {
|
||||
Some(d) => d,
|
||||
None => {
|
||||
return Err(UsageResult::failure(
|
||||
"Missing 'data' field in response".to_string(),
|
||||
))
|
||||
}
|
||||
};
|
||||
|
||||
type Entry = (Option<i64>, f64, Option<String>);
|
||||
let mut five_hour: Option<Entry> = None;
|
||||
let mut weekly: Option<Entry> = None;
|
||||
let mut unclassified: Vec<Entry> = Vec::new();
|
||||
|
||||
if let Some(limits) = data.get("limits").and_then(|v| v.as_array()) {
|
||||
for limit_item in limits {
|
||||
let limit_type = limit_item
|
||||
.get("type")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("");
|
||||
if !limit_type.eq_ignore_ascii_case("TOKENS_LIMIT") {
|
||||
continue;
|
||||
}
|
||||
let percentage = limit_item
|
||||
.get("percentage")
|
||||
.and_then(|v| v.as_f64())
|
||||
.unwrap_or(0.0);
|
||||
let reset_ms = limit_item.get("nextResetTime").and_then(|v| v.as_i64());
|
||||
let reset_iso = reset_ms.and_then(millis_to_iso8601);
|
||||
let entry = (reset_ms, percentage, reset_iso);
|
||||
match limit_item.get("unit").and_then(|v| v.as_i64()) {
|
||||
Some(3) if five_hour.is_none() => five_hour = Some(entry),
|
||||
Some(6) if weekly.is_none() => weekly = Some(entry),
|
||||
_ => unclassified.push(entry),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 兜底:无 nextResetTime 的优先归 five_hour,其余按 reset 升序填空槽。
|
||||
unclassified.sort_by_key(|(reset, _, _)| (reset.is_some(), reset.unwrap_or(i64::MIN)));
|
||||
for entry in unclassified {
|
||||
if five_hour.is_none() {
|
||||
five_hour = Some(entry);
|
||||
} else if weekly.is_none() {
|
||||
weekly = Some(entry);
|
||||
}
|
||||
}
|
||||
|
||||
let mut tiers = Vec::new();
|
||||
for (name, slot) in [(TIER_FIVE_HOUR, five_hour), (TIER_WEEKLY_LIMIT, weekly)] {
|
||||
if let Some((_, percentage, resets_at)) = slot {
|
||||
tiers.push(percent_tier(name, percentage, resets_at));
|
||||
}
|
||||
}
|
||||
Ok(tiers)
|
||||
}
|
||||
|
||||
/// 额度接口与推理接口同 host:bigmodel.cn 与 z.ai 共用同一后端与 JSON shape。
|
||||
fn zhipu_quota_base(base_url: &str) -> &'static str {
|
||||
if base_url.to_lowercase().contains("bigmodel.cn") {
|
||||
"https://open.bigmodel.cn"
|
||||
} else {
|
||||
"https://api.z.ai"
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn query_zhipu(base_url: &str, api_key: &str) -> Result<UsageResult, String> {
|
||||
let url = format!(
|
||||
"{}/api/monitor/usage/quota/limit",
|
||||
zhipu_quota_base(base_url)
|
||||
);
|
||||
match get_json(&url, api_key, AuthStyle::Raw).await? {
|
||||
Fetched::Body(body) => Ok(match parse_zhipu(&body) {
|
||||
Ok(data) => UsageResult::ok(data),
|
||||
Err(err) => err,
|
||||
}),
|
||||
Fetched::Failed(err) => Ok(err),
|
||||
}
|
||||
}
|
||||
|
||||
// ── MiniMax ─────────────────────────────────────────────────
|
||||
// GET https://api.minimaxi.com/v1/api/openplatform/coding_plan/remains
|
||||
// (海外站 api.minimax.io)
|
||||
// 接口直接给"剩余百分比",反转为已用百分比;只取 model_name == "general"。
|
||||
|
||||
pub async fn query_minimax(api_key: &str, is_cn: bool) -> Result<UsageResult, String> {
|
||||
let domain = if is_cn {
|
||||
"api.minimaxi.com"
|
||||
} else {
|
||||
"api.minimax.io"
|
||||
};
|
||||
let url = format!("https://{domain}/v1/api/openplatform/coding_plan/remains");
|
||||
match get_json(&url, api_key, AuthStyle::Bearer).await? {
|
||||
Fetched::Body(body) => Ok(match parse_minimax(&body) {
|
||||
Ok(data) => UsageResult::ok(data),
|
||||
Err(err) => err,
|
||||
}),
|
||||
Fetched::Failed(err) => Ok(err),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_minimax(body: &serde_json::Value) -> Result<Vec<UsageData>, UsageResult> {
|
||||
// 业务级别错误
|
||||
if let Some(base_resp) = body.get("base_resp") {
|
||||
let status_code = base_resp
|
||||
.get("status_code")
|
||||
.and_then(|v| v.as_i64())
|
||||
.unwrap_or(-1);
|
||||
if status_code != 0 {
|
||||
let msg = base_resp
|
||||
.get("status_msg")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("Unknown error");
|
||||
return Err(UsageResult::failure(format!(
|
||||
"API error (code {status_code}): {msg}"
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
let mut tiers = Vec::new();
|
||||
|
||||
let Some(model_remains) = body.get("model_remains").and_then(|v| v.as_array()) else {
|
||||
return Ok(tiers);
|
||||
};
|
||||
// 只取 general(编程套餐),跳过 video 等其他模型
|
||||
let Some(item) = model_remains.iter().find(|item| {
|
||||
item.get("model_name")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s == "general")
|
||||
.unwrap_or(false)
|
||||
}) else {
|
||||
return Ok(tiers);
|
||||
};
|
||||
|
||||
// 5h 桶:剩余百分比 → 已用百分比
|
||||
if let Some(remain_pct) = item
|
||||
.get("current_interval_remaining_percent")
|
||||
.and_then(|v| v.as_f64())
|
||||
{
|
||||
let resets_at = item
|
||||
.get("end_time")
|
||||
.and_then(|v| v.as_i64())
|
||||
.and_then(millis_to_iso8601);
|
||||
tiers.push(percent_tier(TIER_FIVE_HOUR, 100.0 - remain_pct, resets_at));
|
||||
}
|
||||
|
||||
// 周桶:仅 status=1 时激活;status=3 等表示该套餐无周限额,跳过
|
||||
if item.get("current_weekly_status").and_then(|v| v.as_i64()) == Some(1) {
|
||||
if let Some(remain_pct) = item
|
||||
.get("current_weekly_remaining_percent")
|
||||
.and_then(|v| v.as_f64())
|
||||
{
|
||||
let resets_at = item
|
||||
.get("weekly_end_time")
|
||||
.and_then(|v| v.as_i64())
|
||||
.and_then(millis_to_iso8601);
|
||||
tiers.push(percent_tier(
|
||||
TIER_WEEKLY_LIMIT,
|
||||
100.0 - remain_pct,
|
||||
resets_at,
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(tiers)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn kimi_coding_flattens_limits_and_usage() {
|
||||
let body = json!({
|
||||
"limits": [
|
||||
{ "detail": { "limit": 100, "remaining": 40, "resetTime": 1_754_000_000_000i64 } }
|
||||
],
|
||||
"usage": { "limit": 1000, "remaining": 900, "resetTime": "2026-08-01T00:00:00Z" }
|
||||
});
|
||||
let tiers = parse_kimi_coding(&body);
|
||||
assert_eq!(tiers.len(), 2);
|
||||
assert_eq!(tiers[0].plan_name.as_deref(), Some("five_hour"));
|
||||
assert_eq!(tiers[0].used, Some(60.0));
|
||||
assert_eq!(tiers[0].total, Some(100.0));
|
||||
assert_eq!(tiers[0].remaining, Some(40.0));
|
||||
assert!(tiers[0].resets_at.is_some());
|
||||
assert_eq!(tiers[1].plan_name.as_deref(), Some("weekly_limit"));
|
||||
assert_eq!(tiers[1].used, Some(10.0));
|
||||
assert_eq!(
|
||||
tiers[1].resets_at.as_deref(),
|
||||
Some("2026-08-01T00:00:00Z")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kimi_coding_reset_time_seconds_vs_millis() {
|
||||
// 秒级时间戳自动 ×1000
|
||||
let v = json!(1_754_000_000i64);
|
||||
assert!(extract_reset_time(&v).is_some());
|
||||
// 0 / 负值视为无重置时间
|
||||
assert_eq!(extract_reset_time(&json!(0)), None);
|
||||
assert_eq!(extract_reset_time(&json!(-1)), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn zhipu_classifies_by_unit_field() {
|
||||
let body = json!({
|
||||
"success": true,
|
||||
"data": {
|
||||
"limits": [
|
||||
{ "type": "TOKENS_LIMIT", "percentage": 35.0, "unit": 3, "nextResetTime": 1_754_000_000_000i64 },
|
||||
{ "type": "TOKENS_LIMIT", "percentage": 80.0, "unit": 6, "nextResetTime": 1_754_500_000_000i64 }
|
||||
]
|
||||
}
|
||||
});
|
||||
let tiers = parse_zhipu(&body).unwrap();
|
||||
assert_eq!(tiers.len(), 2);
|
||||
assert_eq!(tiers[0].plan_name.as_deref(), Some("five_hour"));
|
||||
assert_eq!(tiers[0].used, Some(35.0));
|
||||
assert_eq!(tiers[1].plan_name.as_deref(), Some("weekly_limit"));
|
||||
assert_eq!(tiers[1].used, Some(80.0));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn zhipu_business_error_is_deterministic_failure() {
|
||||
let body = json!({ "success": false, "msg": "token invalid" });
|
||||
let err = parse_zhipu(&body).unwrap_err();
|
||||
assert!(!err.success);
|
||||
assert!(err.error.unwrap().contains("token invalid"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn minimax_picks_general_and_inverts_remaining() {
|
||||
let body = json!({
|
||||
"base_resp": { "status_code": 0, "status_msg": "success" },
|
||||
"model_remains": [
|
||||
{ "model_name": "video", "current_interval_remaining_percent": 50.0 },
|
||||
{
|
||||
"model_name": "general",
|
||||
"current_interval_remaining_percent": 70.0,
|
||||
"end_time": 1_754_000_000_000i64,
|
||||
"current_weekly_status": 1,
|
||||
"current_weekly_remaining_percent": 95.0,
|
||||
"weekly_end_time": 1_754_500_000_000i64
|
||||
}
|
||||
]
|
||||
});
|
||||
let tiers = parse_minimax(&body).unwrap();
|
||||
assert_eq!(tiers.len(), 2);
|
||||
assert_eq!(tiers[0].plan_name.as_deref(), Some("five_hour"));
|
||||
assert_eq!(tiers[0].used, Some(30.0));
|
||||
assert_eq!(tiers[0].remaining, Some(70.0));
|
||||
assert_eq!(tiers[1].plan_name.as_deref(), Some("weekly_limit"));
|
||||
assert_eq!(tiers[1].used, Some(5.0));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn minimax_skips_inactive_weekly_bucket() {
|
||||
let body = json!({
|
||||
"model_remains": [
|
||||
{
|
||||
"model_name": "general",
|
||||
"current_interval_remaining_percent": 70.0,
|
||||
"current_weekly_status": 3,
|
||||
"current_weekly_remaining_percent": 100.0
|
||||
}
|
||||
]
|
||||
});
|
||||
let tiers = parse_minimax(&body).unwrap();
|
||||
assert_eq!(tiers.len(), 1);
|
||||
assert_eq!(tiers[0].plan_name.as_deref(), Some("five_hour"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn minimax_business_error_is_deterministic_failure() {
|
||||
let body = json!({ "base_resp": { "status_code": 1002, "status_msg": "invalid key" } });
|
||||
let err = parse_minimax(&body).unwrap_err();
|
||||
assert!(!err.success);
|
||||
assert!(err.error.unwrap().contains("invalid key"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,182 @@
|
||||
// Adapted from cc-switch (MIT, © Jason Young)
|
||||
// https://github.com/farion1231/cc-switch
|
||||
|
||||
//! 供应商账单/用量查询统一入口。
|
||||
//!
|
||||
//! - [`UsageKind`]:8 种查询类型,字符串形式与前端 / SQLite settings 约定一致
|
||||
//! (如 `"balance:deepseek"`、`"plan:kimi_coding"`)。
|
||||
//! - [`detect_provider`]:按 base_url host 子串匹配,旧用户无显式配置时自动识别。
|
||||
//! - [`query_kind`]:按 kind 路由到 balance / coding_plan 的具体实现。
|
||||
|
||||
pub mod balance;
|
||||
pub mod coding_plan;
|
||||
pub mod usage_types;
|
||||
|
||||
pub use usage_types::{UsageData, UsageResult};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
|
||||
pub enum UsageKind {
|
||||
BalanceDeepseek,
|
||||
BalanceSiliconflow,
|
||||
BalanceOpenrouter,
|
||||
BalanceStepfun,
|
||||
BalanceNovita,
|
||||
PlanKimiCoding,
|
||||
PlanZhipu,
|
||||
PlanMinimax,
|
||||
}
|
||||
|
||||
impl UsageKind {
|
||||
/// 与前端 / SQLite settings(`usage_kinds:<provider_name>`)约定的字符串形式。
|
||||
pub fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
UsageKind::BalanceDeepseek => "balance:deepseek",
|
||||
UsageKind::BalanceSiliconflow => "balance:siliconflow",
|
||||
UsageKind::BalanceOpenrouter => "balance:openrouter",
|
||||
UsageKind::BalanceStepfun => "balance:stepfun",
|
||||
UsageKind::BalanceNovita => "balance:novita",
|
||||
UsageKind::PlanKimiCoding => "plan:kimi_coding",
|
||||
UsageKind::PlanZhipu => "plan:zhipu",
|
||||
UsageKind::PlanMinimax => "plan:minimax",
|
||||
}
|
||||
}
|
||||
|
||||
pub const ALL: [UsageKind; 8] = [
|
||||
UsageKind::BalanceDeepseek,
|
||||
UsageKind::BalanceSiliconflow,
|
||||
UsageKind::BalanceOpenrouter,
|
||||
UsageKind::BalanceStepfun,
|
||||
UsageKind::BalanceNovita,
|
||||
UsageKind::PlanKimiCoding,
|
||||
UsageKind::PlanZhipu,
|
||||
UsageKind::PlanMinimax,
|
||||
];
|
||||
}
|
||||
|
||||
impl std::str::FromStr for UsageKind {
|
||||
type Err = ();
|
||||
|
||||
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||
Ok(match s {
|
||||
"balance:deepseek" => UsageKind::BalanceDeepseek,
|
||||
"balance:siliconflow" => UsageKind::BalanceSiliconflow,
|
||||
"balance:openrouter" => UsageKind::BalanceOpenrouter,
|
||||
"balance:stepfun" => UsageKind::BalanceStepfun,
|
||||
"balance:novita" => UsageKind::BalanceNovita,
|
||||
"plan:kimi_coding" => UsageKind::PlanKimiCoding,
|
||||
"plan:zhipu" => UsageKind::PlanZhipu,
|
||||
"plan:minimax" => UsageKind::PlanMinimax,
|
||||
_ => return Err(()),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// 按 base_url 子串匹配可支持的查询类型;无匹配返回空 vec。
|
||||
/// 一个 base_url 理论上可同时命中多种(套餐 + 余额),故返回 Vec。
|
||||
pub fn detect_provider(base_url: &str) -> Vec<UsageKind> {
|
||||
let url = base_url.to_lowercase();
|
||||
let mut kinds = Vec::new();
|
||||
if url.contains("api.deepseek.com") {
|
||||
kinds.push(UsageKind::BalanceDeepseek);
|
||||
}
|
||||
if url.contains("api.siliconflow.cn") {
|
||||
kinds.push(UsageKind::BalanceSiliconflow);
|
||||
}
|
||||
if url.contains("openrouter.ai") {
|
||||
kinds.push(UsageKind::BalanceOpenrouter);
|
||||
}
|
||||
if url.contains("api.stepfun.com") {
|
||||
kinds.push(UsageKind::BalanceStepfun);
|
||||
}
|
||||
if url.contains("api.novita.ai") {
|
||||
kinds.push(UsageKind::BalanceNovita);
|
||||
}
|
||||
if url.contains("api.kimi.com") && url.contains("/coding") {
|
||||
kinds.push(UsageKind::PlanKimiCoding);
|
||||
}
|
||||
if url.contains("open.bigmodel.cn") || url.contains("api.z.ai") {
|
||||
kinds.push(UsageKind::PlanZhipu);
|
||||
}
|
||||
if url.contains("api.minimaxi.com") {
|
||||
kinds.push(UsageKind::PlanMinimax);
|
||||
}
|
||||
kinds
|
||||
}
|
||||
|
||||
/// 按 kind 路由到对应查询实现。`base_url` 用于消歧同一家供应商的
|
||||
/// 国内/海外站(SiliconFlow .cn/.com、MiniMax .com/.io、智谱 bigmodel/z.ai)。
|
||||
pub async fn query_kind(
|
||||
kind: UsageKind,
|
||||
base_url: &str,
|
||||
api_key: &str,
|
||||
) -> Result<UsageResult, String> {
|
||||
let lower = base_url.to_lowercase();
|
||||
match kind {
|
||||
UsageKind::BalanceDeepseek => balance::query_deepseek(api_key).await,
|
||||
UsageKind::BalanceSiliconflow => {
|
||||
balance::query_siliconflow(api_key, !lower.contains("siliconflow.com")).await
|
||||
}
|
||||
UsageKind::BalanceOpenrouter => balance::query_openrouter(api_key).await,
|
||||
UsageKind::BalanceStepfun => balance::query_stepfun(api_key).await,
|
||||
UsageKind::BalanceNovita => balance::query_novita(api_key).await,
|
||||
UsageKind::PlanKimiCoding => coding_plan::query_kimi_coding(api_key).await,
|
||||
UsageKind::PlanZhipu => coding_plan::query_zhipu(base_url, api_key).await,
|
||||
UsageKind::PlanMinimax => {
|
||||
coding_plan::query_minimax(api_key, !lower.contains("minimax.io")).await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn detect_provider_maps_known_hosts() {
|
||||
let cases: [(&str, UsageKind); 8] = [
|
||||
("https://api.deepseek.com/v1", UsageKind::BalanceDeepseek),
|
||||
("https://api.siliconflow.cn/v1", UsageKind::BalanceSiliconflow),
|
||||
("https://openrouter.ai/api/v1", UsageKind::BalanceOpenrouter),
|
||||
("https://api.stepfun.com/v1", UsageKind::BalanceStepfun),
|
||||
("https://api.novita.ai/v3", UsageKind::BalanceNovita),
|
||||
("https://api.kimi.com/coding/v1", UsageKind::PlanKimiCoding),
|
||||
(
|
||||
"https://open.bigmodel.cn/api/paas/v4",
|
||||
UsageKind::PlanZhipu,
|
||||
),
|
||||
("https://api.minimaxi.com/v1", UsageKind::PlanMinimax),
|
||||
];
|
||||
for (url, expected) in cases {
|
||||
assert_eq!(
|
||||
detect_provider(url),
|
||||
vec![expected],
|
||||
"url: {url}"
|
||||
);
|
||||
}
|
||||
// z.ai 也命中智谱
|
||||
assert_eq!(
|
||||
detect_provider("https://api.z.ai/api/paas/v4"),
|
||||
vec![UsageKind::PlanZhipu]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detect_provider_unknown_url_returns_empty() {
|
||||
assert!(detect_provider("https://api.openai.com/v1").is_empty());
|
||||
assert!(detect_provider("https://example.com").is_empty());
|
||||
assert!(detect_provider("").is_empty());
|
||||
// api.kimi.com 但无 /coding 路径 → 不命中套餐查询
|
||||
assert!(detect_provider("https://api.kimi.com/v1").is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_kind_string_roundtrip() {
|
||||
use std::str::FromStr;
|
||||
for kind in UsageKind::ALL {
|
||||
let s = kind.as_str();
|
||||
assert_eq!(UsageKind::from_str(s), Ok(kind), "kind: {s}");
|
||||
}
|
||||
assert!(UsageKind::from_str("balance:unknown").is_err());
|
||||
assert!(UsageKind::from_str("").is_err());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
// Adapted from cc-switch (MIT, © Jason Young)
|
||||
// https://github.com/farion1231/cc-switch
|
||||
|
||||
//! Shared return contract for provider billing/usage queries.
|
||||
//!
|
||||
//! MUST stay camelCase: the frontend reads `planName` / `resetsAt` /
|
||||
//! `isValid` — snake_case serialization would silently misalign every field.
|
||||
|
||||
use serde::Serialize;
|
||||
|
||||
#[derive(Serialize, Clone, Debug, Default)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct UsageData {
|
||||
pub plan_name: Option<String>,
|
||||
pub remaining: Option<f64>,
|
||||
pub total: Option<f64>,
|
||||
pub used: Option<f64>,
|
||||
pub unit: Option<String>,
|
||||
pub is_valid: Option<bool>,
|
||||
pub resets_at: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Clone, Debug)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct UsageResult {
|
||||
pub success: bool,
|
||||
pub data: Option<Vec<UsageData>>,
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
impl UsageResult {
|
||||
pub fn ok(data: Vec<UsageData>) -> Self {
|
||||
UsageResult {
|
||||
success: true,
|
||||
data: if data.is_empty() { None } else { Some(data) },
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn failure(msg: String) -> Self {
|
||||
UsageResult {
|
||||
success: false,
|
||||
data: None,
|
||||
error: Some(msg),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn usage_result_serializes_camel_case() {
|
||||
let result = UsageResult {
|
||||
success: true,
|
||||
data: Some(vec![UsageData {
|
||||
plan_name: Some("five_hour".to_string()),
|
||||
remaining: Some(65.0),
|
||||
total: Some(100.0),
|
||||
used: Some(35.0),
|
||||
unit: Some("%".to_string()),
|
||||
is_valid: Some(true),
|
||||
resets_at: Some("2026-07-29T12:00:00+00:00".to_string()),
|
||||
}]),
|
||||
error: None,
|
||||
};
|
||||
let json = serde_json::to_string(&result).unwrap();
|
||||
assert!(json.contains("\"planName\""), "json: {json}");
|
||||
assert!(json.contains("\"resetsAt\""), "json: {json}");
|
||||
assert!(json.contains("\"isValid\""), "json: {json}");
|
||||
assert!(json.contains("\"plan_name\"") == false, "json: {json}");
|
||||
assert!(json.contains("\"resets_at\"") == false, "json: {json}");
|
||||
assert!(json.contains("\"is_valid\"") == false, "json: {json}");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
//! Validation helpers for KimiSwitch configuration and command inputs.
|
||||
@@ -0,0 +1,43 @@
|
||||
{
|
||||
"productName": "Kimi Switch",
|
||||
"version": "0.6.0",
|
||||
"identifier": "com.kimiswitch.app",
|
||||
"build": {
|
||||
"beforeDevCommand": "npm run dev",
|
||||
"beforeBuildCommand": "npm run build",
|
||||
"devUrl": "http://localhost:1420",
|
||||
"frontendDist": "../dist"
|
||||
},
|
||||
"app": {
|
||||
"windows": [
|
||||
{
|
||||
"title": "Kimi Switch",
|
||||
"width": 1200,
|
||||
"height": 800,
|
||||
"minWidth": 1000,
|
||||
"minHeight": 700,
|
||||
"resizable": true,
|
||||
"fullscreen": false
|
||||
}
|
||||
],
|
||||
"security": {
|
||||
"csp": "default-src 'self'; script-src 'self' 'unsafe-inline'; style-src 'self' 'unsafe-inline'; img-src 'self' data:;"
|
||||
}
|
||||
},
|
||||
"bundle": {
|
||||
"active": true,
|
||||
"targets": ["msi"],
|
||||
"windows": {
|
||||
"webviewInstallMode": {
|
||||
"type": "downloadBootstrapper"
|
||||
},
|
||||
"nsis": null
|
||||
},
|
||||
"icon": [
|
||||
"icons/32x32.png",
|
||||
"icons/128x128.png",
|
||||
"icons/128x128@2x.png",
|
||||
"icons/icon.ico"
|
||||
]
|
||||
}
|
||||
}
|
||||