Skip to main content

weft_core/api/
providers.rs

1use crate::api::openai_compat::AppState;
2use crate::config::store::save_config;
3use crate::config::{ApiKeyConfig, ProviderApi, ProviderConfig};
4use axum::extract::{Path, State};
5use axum::http::StatusCode;
6use axum::response::IntoResponse;
7use axum::Json;
8use serde::{Deserialize, Serialize};
9
10pub async fn list_providers(State(state): State<AppState>) -> Json<serde_json::Value> {
11    let config = state.config.read().await;
12    let providers: Vec<serde_json::Value> = config
13        .providers
14        .iter()
15        .map(|p| {
16            serde_json::json!({
17                "name": p.name,
18                "base_url": p.base_url,
19                "format": p.format,
20                "models": p.models,
21                "key_count": p.keys.len(),
22            })
23        })
24        .collect();
25    // 同时返回路由信息,让前端区分文本LLM(default_provider)与图像(image_provider)用途。
26    Json(serde_json::json!({
27        "providers": providers,
28        "routing": {
29            "default_provider": config.routing.default_provider,
30            "default_model": config.routing.default_model,
31            "image_provider": config.routing.image_provider,
32        }
33    }))
34}
35
36pub async fn get_provider(
37    State(state): State<AppState>,
38    Path(name): Path<String>,
39) -> impl IntoResponse {
40    let config = state.config.read().await;
41    match config.providers.iter().find(|p| p.name == name) {
42        Some(p) => {
43            // get_provider 用于编辑对话框,需返回完整 keys(本地单用户应用,
44            // key 归用户所有,通过 runtime-token 鉴权后可见)。列表接口仍只给 key_count。
45            let keys: Vec<serde_json::Value> = p
46                .keys
47                .iter()
48                .map(|k| {
49                    serde_json::json!({
50                        "value": k.value,
51                        "label": k.label,
52                        "enabled": k.enabled,
53                    })
54                })
55                .collect();
56            let resp = serde_json::json!({
57                "name": p.name,
58                "base_url": p.base_url,
59                "format": p.format,
60                "models": p.models,
61                "key_count": p.keys.len(),
62                "keys": keys,
63            });
64            Json(resp).into_response()
65        }
66        None => StatusCode::NOT_FOUND.into_response(),
67    }
68}
69
70#[derive(Debug, Deserialize, Serialize)]
71pub struct CreateProviderRequest {
72    pub name: String,
73    pub base_url: String,
74    #[serde(default = "default_format")]
75    pub format: String,
76    #[serde(default)]
77    pub keys: Vec<ApiKeyConfig>,
78    #[serde(default)]
79    pub models: Vec<String>,
80}
81
82fn default_format() -> String {
83    "openai".into()
84}
85
86/// 请求体:手动从 provider 拉取可用模型列表(配置对话框里的「获取模型」按钮)。
87#[derive(Debug, Deserialize)]
88pub struct FetchModelsRequest {
89    pub base_url: String,
90    #[serde(default)]
91    pub api_key: String,
92    #[serde(default = "default_format")]
93    pub format: String,
94}
95
96/// 调用 provider 的 /models 接口拉取可用模型 id 列表。
97/// OpenAI 格式:GET {base_url}/models;Anthropic 没有标准 models 接口,返回提示。
98pub async fn fetch_models(
99    State(_state): State<AppState>,
100    Json(req): Json<FetchModelsRequest>,
101) -> impl IntoResponse {
102    let base = req.base_url.trim().trim_end_matches('/');
103    if base.is_empty() {
104        return (
105            StatusCode::BAD_REQUEST,
106            Json(serde_json::json!({ "error": "base_url required" })),
107        )
108            .into_response();
109    }
110    // 两种常见布局:base 已含 /v1 → 直接 /models;否则补 /v1/models;都试。
111    let url = if base.ends_with("/v1") {
112        format!("{base}/models")
113    } else {
114        format!("{base}/v1/models")
115    };
116
117    let client = reqwest::Client::new();
118    let mut rb = client.get(&url);
119    if !req.api_key.trim().is_empty() {
120        rb = rb.bearer_auth(req.api_key.trim());
121    }
122    match rb.send().await {
123        Ok(resp) if resp.status().is_success() => {
124            let body: serde_json::Value = resp.json().await.unwrap_or_default();
125            // OpenAI: { "data": [ { "id": "..." }, ... ] }
126            let models: Vec<String> = body
127                .get("data")
128                .and_then(|d| d.as_array())
129                .map(|arr| {
130                    arr.iter()
131                        .filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(String::from))
132                        .collect()
133                })
134                .unwrap_or_default();
135            Json(serde_json::json!({ "models": models })).into_response()
136        }
137        Ok(resp) => (
138            StatusCode::BAD_GATEWAY,
139            Json(serde_json::json!({
140                "error": format!("provider returned {}", resp.status()),
141            })),
142        )
143            .into_response(),
144        Err(e) => (
145            StatusCode::BAD_GATEWAY,
146            Json(serde_json::json!({ "error": format!("fetch failed: {e}") })),
147        )
148            .into_response(),
149    }
150}
151
152#[derive(Debug, Deserialize, Serialize)]
153pub struct UpdateProviderRequest {
154    #[serde(skip_serializing_if = "Option::is_none")]
155    pub base_url: Option<String>,
156    #[serde(skip_serializing_if = "Option::is_none")]
157    pub format: Option<String>,
158    #[serde(skip_serializing_if = "Option::is_none")]
159    pub keys: Option<Vec<ApiKeyConfig>>,
160    #[serde(skip_serializing_if = "Option::is_none")]
161    pub models: Option<Vec<String>>,
162}
163
164/// Create a new provider
165pub async fn create_provider(
166    State(state): State<AppState>,
167    Json(req): Json<CreateProviderRequest>,
168) -> impl IntoResponse {
169    // Validate format
170    if req.format != "openai" && req.format != "anthropic" {
171        return (
172            StatusCode::BAD_REQUEST,
173            Json(serde_json::json!({
174                "error": "Invalid format. Must be 'openai' or 'anthropic'"
175            })),
176        )
177            .into_response();
178    }
179
180    // Validate base_url
181    if req.base_url.is_empty() {
182        return (
183            StatusCode::BAD_REQUEST,
184            Json(serde_json::json!({
185                "error": "base_url cannot be empty"
186            })),
187        )
188            .into_response();
189    }
190
191    let mut config = state.config.write().await;
192
193    // Check if provider already exists
194    if config.providers.iter().any(|p| p.name == req.name) {
195        return (
196            StatusCode::CONFLICT,
197            Json(serde_json::json!({
198                "error": format!("Provider '{}' already exists", req.name)
199            })),
200        )
201            .into_response();
202    }
203
204    // Add provider
205    let provider = ProviderConfig {
206        name: req.name.clone(),
207        base_url: req.base_url,
208        format: req.format,
209        api: ProviderApi::ChatCompletions,
210        keys: req.keys,
211        models: req.models,
212    };
213    config.providers.push(provider);
214
215    // Persist to disk
216    if let Err(e) = save_config(&state.config_path, &config) {
217        tracing::error!("Failed to save config: {}", e);
218        return (
219            StatusCode::INTERNAL_SERVER_ERROR,
220            Json(serde_json::json!({
221                "error": "Failed to persist configuration"
222            })),
223        )
224            .into_response();
225    }
226
227    // Return the created provider so the client can parse it as ProviderConfig.
228    let created = config.providers.last().unwrap();
229    (
230        StatusCode::CREATED,
231        Json(serde_json::json!({
232            "name": created.name,
233            "base_url": created.base_url,
234            "format": created.format,
235            "models": created.models,
236            "keys": created.keys.iter().map(|k| serde_json::json!({
237                "value": k.value,
238                "label": k.label,
239            })).collect::<Vec<_>>(),
240        })),
241    )
242        .into_response()
243}
244
245/// Update an existing provider
246pub async fn update_provider(
247    State(state): State<AppState>,
248    Path(name): Path<String>,
249    Json(req): Json<UpdateProviderRequest>,
250) -> impl IntoResponse {
251    let mut config = state.config.write().await;
252
253    // Find provider
254    let provider = match config.providers.iter_mut().find(|p| p.name == name) {
255        Some(p) => p,
256        None => {
257            return (
258                StatusCode::NOT_FOUND,
259                Json(serde_json::json!({
260                    "error": format!("Provider '{}' not found", name)
261                })),
262            )
263                .into_response()
264        }
265    };
266
267    // Apply updates
268    if let Some(base_url) = req.base_url {
269        if base_url.is_empty() {
270            return (
271                StatusCode::BAD_REQUEST,
272                Json(serde_json::json!({
273                    "error": "base_url cannot be empty"
274                })),
275            )
276                .into_response();
277        }
278        provider.base_url = base_url;
279    }
280
281    if let Some(format) = req.format {
282        if format != "openai" && format != "anthropic" {
283            return (
284                StatusCode::BAD_REQUEST,
285                Json(serde_json::json!({
286                    "error": "Invalid format. Must be 'openai' or 'anthropic'"
287                })),
288            )
289                .into_response();
290        }
291        provider.format = format;
292    }
293
294    if let Some(keys) = req.keys {
295        provider.keys = keys;
296    }
297
298    if let Some(models) = req.models {
299        provider.models = models;
300    }
301
302    // Persist to disk
303    if let Err(e) = save_config(&state.config_path, &config) {
304        tracing::error!("Failed to save config: {}", e);
305        return (
306            StatusCode::INTERNAL_SERVER_ERROR,
307            Json(serde_json::json!({
308                "error": "Failed to persist configuration"
309            })),
310        )
311            .into_response();
312    }
313
314    (
315        StatusCode::OK,
316        Json(serde_json::json!({
317            "message": format!("Provider '{}' updated successfully", name)
318        })),
319    )
320        .into_response()
321}
322
323/// Delete a provider
324pub async fn delete_provider(
325    State(state): State<AppState>,
326    Path(name): Path<String>,
327) -> impl IntoResponse {
328    let mut config = state.config.write().await;
329
330    // Check if provider exists
331    let index = match config.providers.iter().position(|p| p.name == name) {
332        Some(i) => i,
333        None => {
334            return (
335                StatusCode::NOT_FOUND,
336                Json(serde_json::json!({
337                    "error": format!("Provider '{}' not found", name)
338                })),
339            )
340                .into_response()
341        }
342    };
343
344    // Check dependencies
345    let mut warnings = Vec::new();
346    if config.routing.default_provider.as_ref() == Some(&name) {
347        warnings.push("This is the default provider in routing config");
348    }
349    if config.fallback.priority.contains(&name) {
350        warnings.push("This provider is in the fallback priority list");
351    }
352
353    // Remove provider
354    config.providers.remove(index);
355
356    // Persist to disk
357    if let Err(e) = save_config(&state.config_path, &config) {
358        tracing::error!("Failed to save config: {}", e);
359        return (
360            StatusCode::INTERNAL_SERVER_ERROR,
361            Json(serde_json::json!({
362                "error": "Failed to persist configuration"
363            })),
364        )
365            .into_response();
366    }
367
368    let mut response = serde_json::json!({
369        "message": format!("Provider '{}' deleted successfully", name)
370    });
371
372    if !warnings.is_empty() {
373        response["warnings"] = serde_json::json!(warnings);
374    }
375
376    (StatusCode::OK, Json(response)).into_response()
377}
378
379/// PUT /api/routing — 更新路由配置(default_provider / default_model / image_provider)。
380#[derive(Debug, Deserialize)]
381pub struct UpdateRoutingRequest {
382    #[serde(default)]
383    pub default_provider: Option<String>,
384    #[serde(default)]
385    pub default_model: Option<String>,
386    #[serde(default)]
387    pub image_provider: Option<String>,
388}
389
390pub async fn update_routing(
391    State(state): State<AppState>,
392    Json(req): Json<UpdateRoutingRequest>,
393) -> impl IntoResponse {
394    let mut config = state.config.write().await;
395
396    if let Some(dp) = req.default_provider {
397        config.routing.default_provider = if dp.is_empty() { None } else { Some(dp) };
398    }
399    if let Some(dm) = req.default_model {
400        config.routing.default_model = if dm.is_empty() { None } else { Some(dm) };
401    }
402    if let Some(ip) = req.image_provider {
403        config.routing.image_provider = if ip.is_empty() { None } else { Some(ip) };
404    }
405
406    // Persist to disk
407    if let Err(e) = save_config(&state.config_path, &config) {
408        tracing::error!("Failed to save config: {}", e);
409        return (
410            StatusCode::INTERNAL_SERVER_ERROR,
411            Json(serde_json::json!({
412                "error": "Failed to persist configuration"
413            })),
414        )
415            .into_response();
416    }
417
418    Json(serde_json::json!({
419        "routing": {
420            "default_provider": config.routing.default_provider,
421            "default_model": config.routing.default_model,
422            "image_provider": config.routing.image_provider,
423        }
424    }))
425    .into_response()
426}
427
428/// 从上游 provider 的 /v1/models 端点拉取实际可用模型列表。
429/// 让 web 前端能展示上游中转站的全部模型,而不只是 config 手写的几个。
430pub async fn list_upstream_models(
431    State(state): State<AppState>,
432    Path(name): Path<String>,
433) -> impl IntoResponse {
434    let config = state.config.read().await;
435    let Some(provider) = config.providers.iter().find(|p| p.name == name) else {
436        return (StatusCode::NOT_FOUND, Json(serde_json::json!({"error": "provider not found"}))).into_response();
437    };
438    let api_key = provider.keys.iter().find(|k| k.enabled && !k.value.trim().is_empty()).map(|k| k.value.trim().to_string());
439    let base_url = provider.base_url.trim_end_matches('/').to_string();
440    drop(config); // 释放读锁
441
442    let url = if base_url.ends_with("/v1") {
443        format!("{base_url}/models")
444    } else {
445        format!("{base_url}/v1/models")
446    };
447    let client = reqwest::Client::new();
448    let mut req = client.get(&url);
449    if let Some(key) = &api_key {
450        req = req.header("Authorization", format!("Bearer {key}"));
451    }
452    match req.timeout(std::time::Duration::from_secs(15)).send().await {
453        Ok(resp) => {
454            if let Ok(body) = resp.json::<serde_json::Value>().await {
455                let models: Vec<String> = body
456                    .get("data")
457                    .and_then(|d| d.as_array())
458                    .map(|arr| arr.iter().filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(|s| s.to_string())).collect())
459                    .unwrap_or_default();
460                Json(serde_json::json!({"models": models, "count": models.len()})).into_response()
461            } else {
462                (StatusCode::BAD_GATEWAY, Json(serde_json::json!({"error": "invalid upstream response"}))).into_response()
463            }
464        }
465        Err(e) => {
466            (StatusCode::BAD_GATEWAY, Json(serde_json::json!({"error": format!("upstream request failed: {e}")}))).into_response()
467        }
468    }
469}