use std::time::Duration;
use crate::endpoint::{anthropic_base_url, ends_with_version_segment, join_api_path};
use crate::error::VerifyError;
use crate::kind::Protocol;
use crate::model_filter::clean_fetched_models;
use crate::preset::preset_by_key;
#[derive(Debug, Clone, Copy)]
#[non_exhaustive]
pub struct ServiceConfig<'a> {
pub preset_key: Option<&'a str>,
pub protocol: Protocol,
pub base_url: &'a str,
pub api_key: &'a str,
pub model: &'a str,
pub extra: &'a [(&'a str, &'a str)],
}
impl<'a> ServiceConfig<'a> {
pub fn new(protocol: Protocol, base_url: &'a str) -> Self {
Self {
preset_key: None,
protocol,
base_url,
api_key: "",
model: "",
extra: &[],
}
}
pub fn with_preset(mut self, key: &'a str) -> Self {
self.preset_key = Some(key);
self
}
pub fn with_api_key(mut self, key: &'a str) -> Self {
self.api_key = key;
self
}
pub fn with_model(mut self, model: &'a str) -> Self {
self.model = model;
self
}
pub fn with_extra(mut self, extra: &'a [(&'a str, &'a str)]) -> Self {
self.extra = extra;
self
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
#[serde(rename_all = "camelCase")]
#[non_exhaustive]
pub struct VerifyOk {
pub latency_ms: u32,
pub models: Vec<String>,
pub dropped: usize,
pub dropped_models: Vec<String>,
pub model_in_list: bool,
pub limits: Option<crate::limits::TokenLimits>,
pub model_limits: Vec<(String, crate::limits::TokenLimits)>,
}
const VERIFY_TIMEOUT_SECS: u64 = 20;
#[derive(Debug, Clone)]
pub struct Verifier {
client: reqwest::Client,
}
impl Verifier {
pub fn new() -> Result<Self, VerifyError> {
Self::from_builder(reqwest::Client::builder())
}
pub fn from_builder(builder: reqwest::ClientBuilder) -> Result<Self, VerifyError> {
let client = builder
.timeout(Duration::from_secs(VERIFY_TIMEOUT_SECS))
.connect_timeout(Duration::from_secs(10))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| VerifyError::Malformed {
detail: format!("构造 HTTP 客户端失败: {e}"),
})?;
Ok(Self { client })
}
pub async fn verify(&self, cfg: ServiceConfig<'_>) -> Result<VerifyOk, VerifyError> {
check_required_fields(cfg.preset_key, cfg.extra)?;
let base = if cfg.base_url.trim().is_empty() {
cfg.preset_key
.and_then(preset_by_key)
.and_then(|p| p.base_url)
.unwrap_or("")
} else {
cfg.base_url.trim()
};
if base.is_empty() {
return Err(VerifyError::MissingExtraField {
key: "base_url".to_string(),
});
}
let effective_base = match cfg.protocol {
Protocol::Anthropic => anthropic_base_url(base),
_ => base.to_string(),
};
let base = effective_base.as_str();
let url = join_api_path(base, "models");
let key = cfg.api_key.trim();
let mut req = self.client.get(&url);
req = match cfg.protocol {
Protocol::Anthropic => {
let r = req.header("anthropic-version", "2023-06-01");
if key.is_empty() {
r
} else {
r.header("x-api-key", key)
}
}
_ if key.is_empty() => req,
_ => req.bearer_auth(key),
};
let started = std::time::Instant::now();
let resp = req.send().await.map_err(|_| VerifyError::Unreachable {
proxy_hint: needs_proxy_hint(&url),
})?;
let latency_ms = started.elapsed().as_millis().min(u128::from(u32::MAX)) as u32;
let status = resp.status().as_u16();
let body = resp.text().await.unwrap_or_default();
if !(200..300).contains(&status) {
return Err(diagnose(status, &body, &url, base));
}
let ids = parse_model_ids(&body);
let cleaned = clean_fetched_models(ids);
let model = cfg.model.trim();
let model_in_list = model.is_empty()
|| cleaned.models.is_empty()
|| cleaned.models.iter().any(|m| m == model);
let model_limits = parse_model_limits(&body);
let limits = model_limits
.iter()
.find(|(id, _)| id == model)
.map(|(_, l)| *l);
Ok(VerifyOk {
latency_ms,
models: cleaned.models,
dropped: cleaned.dropped,
dropped_models: cleaned.dropped_models,
model_in_list,
limits,
model_limits,
})
}
}
static DEFAULT_VERIFIER: std::sync::OnceLock<Result<Verifier, String>> = std::sync::OnceLock::new();
fn default_verifier() -> Result<&'static Verifier, VerifyError> {
DEFAULT_VERIFIER
.get_or_init(|| Verifier::new().map_err(|e| e.to_string()))
.as_ref()
.map_err(|detail| VerifyError::Malformed {
detail: detail.clone(),
})
}
pub fn check_required_fields(
preset_key: Option<&str>,
extra: &[(&str, &str)],
) -> Result<(), VerifyError> {
let Some(p) = preset_key.and_then(preset_by_key) else {
return Ok(());
};
for f in p.extra_fields {
if !f.required {
continue;
}
let given = extra
.iter()
.find(|(k, _)| *k == f.key)
.is_some_and(|(_, v)| !v.trim().is_empty());
if !given {
return Err(VerifyError::MissingExtraField {
key: f.key.to_string(),
});
}
}
Ok(())
}
pub fn suggest_url(base_url: &str) -> Option<String> {
let trimmed = base_url.trim().trim_end_matches('/');
if trimmed.is_empty() {
return None;
}
if ends_with_version_segment(trimmed)
|| trimmed.contains("/v1beta/")
|| trimmed.contains("/v1/")
{
return None;
}
Some(format!("{trimmed}/v1"))
}
pub fn diagnose(status: u16, body: &str, requested_url: &str, base_url: &str) -> VerifyError {
let detail = extract_error_message(body).unwrap_or_else(|| format!("HTTP {status}"));
match status {
401 | 403 => VerifyError::AuthFailed { detail },
404 => VerifyError::NotFound {
requested_url: requested_url.to_string(),
suggested_url: suggest_url(base_url),
},
_ => VerifyError::Malformed { detail },
}
}
fn extract_error_message(body: &str) -> Option<String> {
let v: serde_json::Value = serde_json::from_str(body).ok()?;
let msg = v
.get("error")
.and_then(|e| e.get("message"))
.or_else(|| v.get("error").and_then(|e| e.as_str().map(|_| e)))
.and_then(|m| m.as_str())
.or_else(|| v.get("message").and_then(|m| m.as_str()))?;
let msg = msg.trim();
if msg.is_empty() {
return None;
}
Some(msg.chars().take(300).collect())
}
pub fn parse_model_limits(body: &str) -> Vec<(String, crate::limits::TokenLimits)> {
let Ok(v) = serde_json::from_str::<serde_json::Value>(body) else {
return Vec::new();
};
let arr = v
.get("data")
.and_then(|d| d.as_array())
.or_else(|| v.as_array());
let Some(arr) = arr else { return Vec::new() };
arr.iter()
.filter_map(|item| {
let id = item.get("id").and_then(|i| i.as_str())?;
let limits = crate::limits::parse_model_limits(item)?;
Some((id.to_string(), limits))
})
.collect()
}
pub fn parse_model_ids(body: &str) -> Vec<String> {
let Ok(v) = serde_json::from_str::<serde_json::Value>(body) else {
return Vec::new();
};
let arr = v
.get("data")
.and_then(|d| d.as_array())
.or_else(|| v.as_array());
let Some(arr) = arr else { return Vec::new() };
arr.iter()
.filter_map(|item| {
item.get("id")
.and_then(|i| i.as_str())
.or_else(|| item.as_str())
.map(str::to_string)
})
.collect()
}
fn needs_proxy_hint(url: &str) -> bool {
const BLOCKED: &[&str] = &[
"api.openai.com",
"api.anthropic.com",
"generativelanguage.googleapis.com",
"openrouter.ai",
"api.groq.com",
"api.x.ai",
];
let u = url.to_ascii_lowercase();
BLOCKED.iter().any(|h| u.contains(h))
}
pub async fn verify(cfg: ServiceConfig<'_>) -> Result<VerifyOk, VerifyError> {
default_verifier()?.verify(cfg).await
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn verify_ok_serializes_camel_case() {
let ok = VerifyOk {
latency_ms: 320,
models: vec!["deepseek-flash".into()],
dropped: 2,
dropped_models: vec!["bge-m3".into(), "tts-1".into()],
model_in_list: true,
limits: Some(crate::limits::TokenLimits::from_endpoint(
Some(128_000),
Some(8192),
)),
model_limits: Vec::new(),
};
let j = serde_json::to_string(&ok).unwrap();
assert!(j.contains(r#""latencyMs":320"#), "前端读 latencyMs:{j}");
assert!(
j.contains(r#""modelInList":true"#),
"前端读 modelInList:{j}"
);
assert!(j.contains(r#""contextWindow":128000"#), "{j}");
assert!(j.contains(r#""source":"endpoint""#), "来源必须能分辨:{j}");
assert!(
j.contains(r#""dropped":2"#),
"「已滤掉 N 个」的提示靠它:{j}"
);
assert!(
j.contains(r#""droppedModels":["bge-m3","tts-1"]"#),
"多模态调用方靠它把非对话模型接回清单:{j}"
);
}
#[test]
fn suggest_url_only_when_version_missing() {
assert_eq!(
suggest_url("https://api.deepseek.com").as_deref(),
Some("https://api.deepseek.com/v1")
);
assert_eq!(
suggest_url("https://api.deepseek.com/").as_deref(),
Some("https://api.deepseek.com/v1")
);
assert_eq!(suggest_url("https://api.deepseek.com/v1"), None);
assert_eq!(suggest_url("https://open.bigmodel.cn/api/paas/v4"), None);
assert_eq!(
suggest_url("https://generativelanguage.googleapis.com/v1beta/openai"),
None
);
assert_eq!(suggest_url(""), None);
}
#[test]
fn diagnose_maps_status_to_actionable_errors() {
let e = diagnose(401, r#"{"error":{"message":"invalid key"}}"#, "u", "b");
assert!(matches!(e, VerifyError::AuthFailed { ref detail } if detail == "invalid key"));
let e = diagnose(404, "{}", "https://x.com/models", "https://x.com");
match e {
VerifyError::NotFound {
requested_url,
suggested_url,
} => {
assert_eq!(requested_url, "https://x.com/models");
assert_eq!(suggested_url.as_deref(), Some("https://x.com/v1"));
}
other => panic!("404 应映射为 NotFound,实际 {other:?}"),
}
assert!(matches!(
diagnose(403, "{}", "u", "b"),
VerifyError::AuthFailed { .. }
));
assert!(matches!(
diagnose(500, "{}", "u", "b"),
VerifyError::Malformed { .. }
));
}
#[test]
fn extracts_error_message_from_various_shapes() {
assert_eq!(
extract_error_message(r#"{"error":{"message":"no credit"}}"#).as_deref(),
Some("no credit")
);
assert_eq!(
extract_error_message(r#"{"message":"bad request"}"#).as_deref(),
Some("bad request")
);
assert_eq!(extract_error_message("not json"), None);
assert_eq!(extract_error_message(r#"{"error":{"message":" "}}"#), None);
}
#[test]
fn parses_model_ids_from_common_shapes() {
assert_eq!(
parse_model_ids(r#"{"data":[{"id":"gpt-4o"},{"id":"gpt-4o-mini"}]}"#),
vec!["gpt-4o", "gpt-4o-mini"]
);
assert_eq!(
parse_model_ids(r#"["llama3.1:8b","qwen3:8b"]"#),
vec!["llama3.1:8b", "qwen3:8b"]
);
assert_eq!(parse_model_ids(r#"[{"id":"a"}]"#), vec!["a"]);
assert!(parse_model_ids("garbage").is_empty());
}
#[test]
fn check_required_fields_catches_missing() {
assert!(check_required_fields(Some("deepseek"), &[]).is_ok());
assert!(check_required_fields(None, &[]).is_ok());
assert!(check_required_fields(Some("不存在的预置"), &[]).is_ok());
}
#[test]
fn proxy_hint_only_for_blocked_hosts() {
assert!(needs_proxy_hint("https://api.openai.com/v1/models"));
assert!(needs_proxy_hint(
"https://generativelanguage.googleapis.com/x"
));
assert!(!needs_proxy_hint("https://api.deepseek.com/v1/models"));
assert!(!needs_proxy_hint("http://localhost:11434/v1/models"));
}
}