389 lines
18 KiB
Rust
389 lines
18 KiB
Rust
//! 厂商模型列表拉取。
|
|
//!
|
|
//! 按 `provider_type` 分派拉取厂商模型列表,过滤非 chat 模型,返回模型名 Vec。
|
|
//! `fetch_and_probe` 在拉取基础上对每个模型名调 `model_probe::probe` 探测出完整 `ModelConfig`。
|
|
//!
|
|
//! 协议分派:
|
|
//! - `openai_compat` → `GET {base_url}/v1/models`(Bearer 鉴权)
|
|
//! - `anthropic_compat` → `GET {base_url}/v1/models`(x-api-key + anthropic-version 鉴权)
|
|
//! - 其他 → `UnsupportedProvider` 错误
|
|
//!
|
|
//! 本模块仅含网络 IO 入口;URL 拼接 / 噪音过滤 / 响应反序列化等纯逻辑见
|
|
//! [`model_fetch_helpers`]。
|
|
//!
|
|
//! 设计来源:docs/02-架构设计/已编号方案/F-01-模型能力系统与智能路由设计-2026-06-16.md §5
|
|
//! (设计 §5.1 端点表 + §5.3 URL 智能拼接 + §5.4 噪音过滤)。
|
|
|
|
use std::time::Duration;
|
|
|
|
use anyhow::{anyhow, Result};
|
|
use df_ai_core::model::ModelConfig;
|
|
use serde_json::Value;
|
|
|
|
use crate::model_fetch_helpers::{build_models_url, filter_chat_models, ModelsList};
|
|
use crate::model_probe::probe;
|
|
|
|
// ────────────────────────────────────────────────────────────
|
|
// 公共入口
|
|
// ────────────────────────────────────────────────────────────
|
|
|
|
/// HTTP 请求超时(秒)。模型列表接口响应小、无需长超时。
|
|
const FETCH_TIMEOUT_SECS: u64 = 10;
|
|
|
|
/// 拉取厂商模型列表,返回已过滤噪音的模型名 Vec。
|
|
///
|
|
/// 协议分派(见模块文档)。错误信息含 `provider_type` + HTTP 状态/原因,便于前端友好提示。
|
|
///
|
|
/// 不真实调测:本函数本身是网络 IO,单测只覆盖 URL 拼接 / 噪音过滤等纯逻辑;
|
|
/// 集成验证留给阶段 5 的 IPC `ai_fetch_models`(用户在 Settings 点「测试连接」触发)。
|
|
pub async fn fetch_model_names(
|
|
provider_type: &str,
|
|
base_url: &str,
|
|
api_key: &str,
|
|
) -> Result<Vec<String>> {
|
|
match provider_type {
|
|
"openai_compat" => fetch_openai_compat(base_url, api_key).await,
|
|
"anthropic_compat" => fetch_anthropic_compat(base_url, api_key).await,
|
|
other => Err(anyhow!(
|
|
"拉取模型列表失败:不支持的 provider_type={other}(仅支持 openai_compat / anthropic_compat)"
|
|
)),
|
|
}
|
|
}
|
|
|
|
/// 拉取模型列表 + 对每个模型名调 `model_probe::probe` 探测出完整 `ModelConfig`。
|
|
///
|
|
/// 探测为纯 CPU(无网络),不会因单模型探测失败中断整体 — 探测对任意输入都有兜底(见 model_probe)。
|
|
pub async fn fetch_and_probe(
|
|
provider_type: &str,
|
|
base_url: &str,
|
|
api_key: &str,
|
|
) -> Result<Vec<ModelConfig>> {
|
|
let names = fetch_model_names(provider_type, base_url, api_key).await?;
|
|
Ok(names.into_iter().map(|n| probe(&n)).collect())
|
|
}
|
|
|
|
// ────────────────────────────────────────────────────────────
|
|
// 协议分派实现
|
|
// ────────────────────────────────────────────────────────────
|
|
|
|
/// OpenAI 兼容(GLM / DeepSeek / Kimi / 通义 / 中转站):`GET /v1/models`,Bearer 鉴权。
|
|
async fn fetch_openai_compat(base_url: &str, api_key: &str) -> Result<Vec<String>> {
|
|
let url = build_models_url(base_url);
|
|
let client = build_client()?;
|
|
|
|
let resp = client
|
|
.get(&url)
|
|
.bearer_auth(api_key)
|
|
.send()
|
|
.await
|
|
.map_err(|e| map_network_error("openai_compat", &url, e))?;
|
|
|
|
let status = resp.status();
|
|
if !status.is_success() {
|
|
return Err(map_status_error("openai_compat", status, &url));
|
|
}
|
|
|
|
// OpenAI 响应:`{data:[{id, owned_by, ...}]}`。中转站通常同构。
|
|
// 不用 resp.json():reqwest::Error::Decode 的 Display 吞 serde 详情(只给 "error decoding
|
|
// response body"),SenseNova 等厂商解析失败时无法定位根因。改 text() + serde_json::from_str,
|
|
// 解析失败时 serde_json::Error 含具体 field/type/position;再叠加宽松 Value fallback 兜底。
|
|
let body = resp
|
|
.text()
|
|
.await
|
|
.map_err(|e| anyhow!("openai_compat 读取响应体失败({url}):{e}"))?;
|
|
|
|
Ok(filter_chat_models(parse_models_compat("openai_compat", &url, &body)?))
|
|
}
|
|
|
|
/// Anthropic 兼容(Claude 官方 / GLM 订阅端点):`GET /v1/models`,x-api-key + anthropic-version 鉴权。
|
|
async fn fetch_anthropic_compat(base_url: &str, api_key: &str) -> Result<Vec<String>> {
|
|
let url = build_models_url(base_url);
|
|
let client = build_client()?;
|
|
|
|
let resp = client
|
|
.get(&url)
|
|
.header("x-api-key", api_key)
|
|
.header("anthropic-version", "2023-06-01")
|
|
.send()
|
|
.await
|
|
.map_err(|e| map_network_error("anthropic_compat", &url, e))?;
|
|
|
|
let status = resp.status();
|
|
if !status.is_success() {
|
|
return Err(map_status_error("anthropic_compat", status, &url));
|
|
}
|
|
|
|
// Anthropic 响应:`{data:[{id, display_name, type, ...}]}`(has_more 分页字段忽略)。
|
|
// 兼容兜底:`{models:[{name, ...}]}`(Ollama 风格,理论 anthropic_compat 不会命中,
|
|
// 但中转站行为不可控,用 `#[serde(alias)]` 零成本兜底 — 见 issues)。
|
|
// 与 openai_compat 同:text() + 严格 serde + Value 宽松 fallback,见 parse_models_compat。
|
|
let body = resp
|
|
.text()
|
|
.await
|
|
.map_err(|e| anyhow!("anthropic_compat 读取响应体失败({url}):{e}"))?;
|
|
|
|
Ok(filter_chat_models(parse_models_compat("anthropic_compat", &url, &body)?))
|
|
}
|
|
|
|
// ────────────────────────────────────────────────────────────
|
|
// 响应解析(text → 严格 serde → Value 宽松 fallback)
|
|
// ────────────────────────────────────────────────────────────
|
|
|
|
/// 响应体诊断片段最大字符数。完整 body 可能巨大,日志只取前缀定位结构。
|
|
const BODY_DIAGNOSTIC_CHARS: usize = 200;
|
|
|
|
/// 解析厂商 `/v1/models` 响应体,返回模型 id 列表(过滤前)。
|
|
///
|
|
/// 三层解析(诊断优先,兜底保成功):
|
|
/// 1. **严格**:`serde_json::from_str::<ModelsList>` — 标准结构命中,错误信息含具体
|
|
/// field/type/position(serde_json::Error Display 自带 line/column,不丢 detail)。
|
|
/// 2. **宽松 fallback**:`serde_json::Value` 解析 → 取 `data` / `models` 任一数组 →
|
|
/// 遍历项取 `id` / `name` 字符串。容错厂商额外字段、类型变体(如 id 漏成 number)。
|
|
/// 3. **诊断错误**:严格 + 宽松都失败时,返回含 HTTP 标识 + serde detail + body 前缀
|
|
/// 的友好错误,而非 reqwest 默认 "error decoding response body"。
|
|
///
|
|
/// 注:fallback 只取 id/name(模型名),丢弃 ModelEntry 上的其他字段 — 厂商变体下
|
|
/// 我们关心的就是模型名,ModelsList 本身也只消费 id/name,语义对齐。
|
|
fn parse_models_compat(provider_type: &str, url: &str, body: &str) -> Result<Vec<String>> {
|
|
// 1) 严格解析(标准结构,serde 错误 detail 完整)。
|
|
match serde_json::from_str::<ModelsList>(body) {
|
|
Ok(list) => return Ok(list.into_ids()),
|
|
Err(strict_err) => {
|
|
// 2) 宽松 Value fallback — 不依赖 ModelsList 结构,容错厂商变体。
|
|
if let Some(ids) = parse_ids_loose(body) {
|
|
return Ok(ids);
|
|
}
|
|
// 3) 双双失败:叠 HTTP 标识 + serde detail + body 前缀诊断。
|
|
return Err(anyhow!(
|
|
"{provider_type} 响应解析失败({url}):{strict_err} | body 前缀:{}",
|
|
body_preview(body)
|
|
));
|
|
}
|
|
}
|
|
}
|
|
|
|
/// 用 `serde_json::Value` 宽松提取模型 id/name。失败(非 JSON / 无 data / 无 id)返回 None。
|
|
///
|
|
/// 取数组字段优先级:`data`(OpenAI/Anthropic)→ `models`(Ollama 风格 alias)。
|
|
/// 项里取 `id` → 兜底 `name`,只接受字符串值(number/bool 等跳过)。
|
|
fn parse_ids_loose(body: &str) -> Option<Vec<String>> {
|
|
let val: Value = serde_json::from_str(body).ok()?;
|
|
let obj = val.as_object()?;
|
|
// 任一存在即取;data 优先(标准结构)。
|
|
let arr = obj.get("data").or_else(|| obj.get("models"))?;
|
|
let arr = arr.as_array()?;
|
|
let mut ids = Vec::with_capacity(arr.len());
|
|
for item in arr {
|
|
let id = item
|
|
.get("id")
|
|
.or_else(|| item.get("name"))
|
|
.and_then(|v| v.as_str());
|
|
if let Some(id) = id {
|
|
ids.push(id.to_string());
|
|
}
|
|
}
|
|
Some(ids)
|
|
}
|
|
|
|
/// body 前缀诊断(截断 + 控制字符占位,避免换行/制表符污染日志单行)。
|
|
fn body_preview(body: &str) -> String {
|
|
let prefix: String = body.chars().take(BODY_DIAGNOSTIC_CHARS).collect();
|
|
if prefix.chars().all(|c| c.is_control()) && !prefix.is_empty() {
|
|
// 整段控制字符(二进制?)→ 给长度提示而非乱码。
|
|
return format!("<非文本 body,长度 {}>", body.len());
|
|
}
|
|
let truncated = body.chars().count() > BODY_DIAGNOSTIC_CHARS;
|
|
// 把控制字符(换行/制表等)压成空格,保持日志单行可读。
|
|
let cleaned: String = prefix
|
|
.chars()
|
|
.map(|c| if c.is_control() { ' ' } else { c })
|
|
.collect();
|
|
if truncated {
|
|
format!("{cleaned}…")
|
|
} else {
|
|
cleaned
|
|
}
|
|
}
|
|
|
|
// ────────────────────────────────────────────────────────────
|
|
// HTTP 客户端 + 错误映射
|
|
// ────────────────────────────────────────────────────────────
|
|
|
|
fn build_client() -> Result<reqwest::Client> {
|
|
reqwest::Client::builder()
|
|
.timeout(Duration::from_secs(FETCH_TIMEOUT_SECS))
|
|
.build()
|
|
.map_err(|e| anyhow!("构建 HTTP 客户端失败:{e}"))
|
|
}
|
|
|
|
/// 网络错误(连接拒绝 / DNS / 超时)→ 友好提示,标注 provider_type + url 便于排查。
|
|
fn map_network_error(provider_type: &str, url: &str, e: reqwest::Error) -> anyhow::Error {
|
|
if e.is_timeout() {
|
|
anyhow!("{provider_type} 拉取模型列表超时({url},>{FETCH_TIMEOUT_SECS}s)— 检查 base_url 可达性")
|
|
} else {
|
|
anyhow!("{provider_type} 拉取模型列表网络错误({url}):{e}")
|
|
}
|
|
}
|
|
|
|
/// HTTP 状态错误 → 按 401/403/404/5xx 分类友好提示(对齐设计 §5.5 失败降级)。
|
|
fn map_status_error(provider_type: &str, status: reqwest::StatusCode, url: &str) -> anyhow::Error {
|
|
if status.as_u16() == 401 || status.as_u16() == 403 {
|
|
anyhow!("{provider_type} 鉴权失败({status}, {url})— 检查 api_key 是否正确 / 是否有该接口权限")
|
|
} else if status.as_u16() == 404 {
|
|
anyhow!("{provider_type} 端点不存在(404, {url})— 该厂商可能不支持模型列表拉取,请手动输入模型名")
|
|
} else if status.is_server_error() {
|
|
anyhow!("{provider_type} 厂商服务异常({status}, {url})— 稍后重试或手动输入模型名")
|
|
} else {
|
|
anyhow!("{provider_type} 拉取模型列表失败(HTTP {status}, {url})")
|
|
}
|
|
}
|
|
|
|
// ────────────────────────────────────────────────────────────
|
|
// 测试(协议分派:不支持类型直接报错,不触网)
|
|
// ────────────────────────────────────────────────────────────
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
// ── fetch_model_names 协议分派:不支持类型直接报错(不触网) ──
|
|
|
|
#[tokio::test]
|
|
async fn fetch_unsupported_provider_errors_without_network() {
|
|
// ollama / glm 等非 openai_compat / anthropic_compat 类型 → 直接错误,不发请求
|
|
let err = fetch_model_names("ollama", "http://localhost:11434", "k")
|
|
.await
|
|
.unwrap_err();
|
|
let msg = format!("{err}");
|
|
assert!(msg.contains("ollama"), "err={msg}");
|
|
assert!(msg.contains("provider_type"), "err={msg}");
|
|
}
|
|
|
|
// ── parse_models_compat:严格 / fallback / 诊断三层 ──
|
|
|
|
#[test]
|
|
fn parse_strict_openai_format() {
|
|
// 标准 OpenAI 结构 → 严格解析命中,不进 fallback
|
|
let body = r#"{"data":[{"id":"gpt-4o","owned_by":"openai"},{"id":"gpt-4o-mini"}]}"#;
|
|
let ids = parse_models_compat("openai_compat", "http://x/v1/models", body).unwrap();
|
|
assert_eq!(ids, vec!["gpt-4o", "gpt-4o-mini"]);
|
|
}
|
|
|
|
#[test]
|
|
fn parse_loose_fallback_on_unknown_field_type_variant() {
|
|
// 厂商变体:data 项里多了非标准字段、且某项漏 id → 严格可能仍过(serde default),
|
|
// 此用例构造严格失败 + 宽松应成功:id 字段为 number(非字符串)致 ModelEntry serde 失败。
|
|
// 宽松 fallback 应:跳过 number id,保留 string id。
|
|
let body = r#"{"data":[{"id":12345},{"id":"glm-4-flash"}]}"#;
|
|
// 严格 ModelsList 的 id: Option<String>,number 12345 无法反序列化为 String → 失败
|
|
let ids = parse_models_compat("openai_compat", "http://x/v1/models", body).unwrap();
|
|
assert_eq!(ids, vec!["glm-4-flash"]);
|
|
}
|
|
|
|
#[test]
|
|
fn parse_loose_fallback_via_models_alias() {
|
|
// 严格解析缺 data 字段时进 fallback,走 models alias 取 name
|
|
let body = r#"{"models":[{"name":"llama3:8b"},{"name":"qwen2:7b"}]}"#;
|
|
let ids = parse_models_compat("openai_compat", "http://x/v1/models", body).unwrap();
|
|
assert_eq!(ids, vec!["llama3:8b", "qwen2:7b"]);
|
|
}
|
|
|
|
#[test]
|
|
fn parse_loose_fallback_tolerates_extra_top_level_fields() {
|
|
// 宽松 fallback 应容错顶层额外字段、非 id 项(只关心 data[].id/name)
|
|
let body = r#"{"object":"list","data":[{"id":"deepseek-chat","object":"model"},{"id":"deepseek-coder"}],"supported_ids":["x"]}"#;
|
|
let ids = parse_models_compat("openai_compat", "http://x/v1/models", body).unwrap();
|
|
assert_eq!(ids, vec!["deepseek-chat", "deepseek-coder"]);
|
|
}
|
|
|
|
#[test]
|
|
fn parse_diagnostic_error_has_serde_detail_and_body_prefix() {
|
|
// 完全无法解析(非 JSON)→ 严格 + 宽松双失败 → 错误含 serde detail + body 前缀 + HTTP 标识
|
|
let body = "this is not json at all {{{";
|
|
let err = parse_models_compat("openai_compat", "http://x/v1/models", body).unwrap_err();
|
|
let msg = format!("{err}");
|
|
// provider_type 标识
|
|
assert!(msg.contains("openai_compat"), "err={msg}");
|
|
// url 便于定位
|
|
assert!(msg.contains("http://x/v1/models"), "err={msg}");
|
|
// serde detail(serde_json 错误含 line/column 或 expected 字样)
|
|
assert!(
|
|
msg.contains("line") || msg.contains("column") || msg.contains("expected"),
|
|
"err={msg}"
|
|
);
|
|
// body 前缀诊断片段
|
|
assert!(msg.contains("this is not json"), "err={msg}");
|
|
}
|
|
|
|
#[test]
|
|
fn parse_diagnostic_truncates_long_body() {
|
|
// 超长 body → 前缀截断(… 标记),不整段灌进错误信息
|
|
let long_id = "a".repeat(500);
|
|
let body = format!(r#"{{"garbage":"{long_id}""#); // 缺尾 → 非 JSON
|
|
let err = parse_models_compat("openai_compat", "http://x/v1/models", &body).unwrap_err();
|
|
let msg = format!("{err}");
|
|
assert!(msg.contains("…"), "长 body 应截断(err={})\n{}", msg.len(), msg);
|
|
// 诊断片段不应超过 BODY_DIAGNOSTIC_CHARS + 容差
|
|
assert!(
|
|
msg.len() < long_id.len(),
|
|
"错误信息不应含完整 500 字符 body"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn parse_diagnostic_empty_body() {
|
|
// 空 body → 双失败,错误信息不 panic、含 provider 标识
|
|
let err = parse_models_compat("openai_compat", "http://x/v1/models", "").unwrap_err();
|
|
let msg = format!("{err}");
|
|
assert!(msg.contains("openai_compat"), "err={msg}");
|
|
assert!(msg.contains("解析失败"), "err={msg}");
|
|
}
|
|
|
|
// ── parse_ids_loose:边界 ──
|
|
|
|
#[test]
|
|
fn parse_ids_loose_returns_none_on_non_json() {
|
|
assert!(parse_ids_loose("not json").is_none());
|
|
}
|
|
|
|
#[test]
|
|
fn parse_ids_loose_returns_none_on_missing_data_field() {
|
|
// 合法 JSON 但无 data/models → None(parse_models_compat 会进而报诊断错误)
|
|
assert!(parse_ids_loose(r#"{"foo":"bar"}"#).is_none());
|
|
}
|
|
|
|
#[test]
|
|
fn parse_ids_loose_data_not_array_returns_none() {
|
|
// data 存在但非数组 → None
|
|
assert!(parse_ids_loose(r#"{"data":"not-an-array"}"#).is_none());
|
|
}
|
|
|
|
#[test]
|
|
fn parse_ids_loose_skips_non_string_id() {
|
|
// id 为 number/null/object → 跳过,只留字符串 id
|
|
let body = r#"{"data":[{"id":1},{"id":null},{"id":"keep-me"},{"name":"named"}]}"#;
|
|
let ids = parse_ids_loose(body).unwrap();
|
|
assert_eq!(ids, vec!["keep-me", "named"]);
|
|
}
|
|
|
|
// ── body_preview:控制字符 + 截断 ──
|
|
|
|
#[test]
|
|
fn body_preview_replaces_control_chars_with_space() {
|
|
// 换行/制表压成空格,保持日志单行
|
|
let preview = body_preview("line1\nline2\tcol");
|
|
assert!(!preview.contains('\n'), "preview={preview}");
|
|
assert!(!preview.contains('\t'), "preview={preview}");
|
|
assert!(preview.contains("line1"), "preview={preview}");
|
|
}
|
|
|
|
#[test]
|
|
fn body_preview_truncates_with_ellipsis() {
|
|
let body = "abcdefghij".repeat(100); // 1000 chars
|
|
let preview = body_preview(&body);
|
|
assert!(preview.ends_with('…'), "preview should end with ellipsis");
|
|
// 不应含完整 body
|
|
assert!(preview.len() < body.len());
|
|
}
|
|
}
|