mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
feat(model-fetch): 对齐 Rust 上游模型抓取行为到 Python 语义
将 Rust 版上游可用模型抓取逻辑收敛到 Python 版行为,统一后台自动抓模 与管理员 provider-query 的模型发现路径,消除标准 /models、固定模型目录、 Antigravity、Vertex AI 等 provider 在两端实现上的分叉。 核心变更: - 在 aether-model-fetch 中引入统一抓模策略层 - 覆盖标准 /models、Vertex API Key、Vertex Service Account、 Antigravity fetchAvailableModels、固定模型目录五类抓模路径 - 将 provider-query 与后台自动抓模都切换到共享抓模入口,避免重复拼接 URL、headers 和 provider 特判逻辑 标准模型抓取对齐: - 按 Python 语义调整抓模优先级: openai:chat > openai:cli > openai:compact claude:chat > claude:cli gemini:chat > gemini:cli - 从抓模候选中移除 openai:responses - 为 openai:cli/openai:compact、claude:cli、gemini:* 补齐 Python 同款 User-Agent / 浏览器指纹请求头 - Claude 抓模保留 after_id 分页语义 - Gemini 抓模统一为 v1beta/models?key=... 语义 provider-query 对齐: - 返回结果改为按 model id 聚合,并合并/排序 api_formats - 最终模型列表按 model id 排序,行为与 Python 保持一致 - 固定目录 provider(codex/kiro/claude_code/gemini_cli)不再依赖活跃 endpoint,即使无 endpoint 也能返回预设模型目录 - Antigravity 多 key 查询改为按账户可用性 + tier 排序,首个成功结果即 停止,并接入 provider 级缓存 - 仅配置 openai:responses 的 provider 不再被视为抓模成功路径 自动抓模对齐: - 自动抓模成功时写入 allowed_models、upstream_models cache,并同步 upstream_metadata - upstream_metadata 合并逻辑对齐 Python,对 quota_by_model 做模型级合并, 并保留已有 reset_time - 自动抓模失败时不覆盖已有 allowed_models - 固定目录 provider 在无 endpoint 场景下也可成功更新 allowed_models Antigravity 对齐: - 使用 POST /v1internal:fetchAvailableModels 抓取可用模型 - 按 Python 规则处理 URL fallback 和 429/404/408/5xx fallback 状态 - 强制要求 auth_config.project_id - 过滤 Python 黑名单模型 - 解析并持久化 upstream_metadata.antigravity.quota_by_model Vertex AI 对齐: - API Key 模式仅抓取 publishers/google/models - Service Account 模式新增 JWT token exchange,并按 Python region 顺序 抓取 google + anthropic publishers - 模型 owned_by / display_name / api_format 推断与 Python 对齐 - 软 404 处理行为与 Python 收敛 Gemini CLI / 固定目录对齐: - Gemini CLI 改为返回 Python 预设模型目录 - 在可用时通过 loadCodeAssist 补充 plan_type/project_id 元数据 - Codex/Kiro/Claude Code 改为共享固定模型目录实现 测试: - 扩展 aether-model-fetch 单元测试,覆盖格式优先级、openai:responses 排除、 请求头、Claude 分页、Gemini query auth、固定目录与 metadata 合并 - 调整 provider-query 控制面测试到 Python 语义 - 新增自动抓模运行时测试,覆盖固定目录成功、Antigravity metadata 合并、 失败保留旧 allowed_models 验证: - cargo nextest run -p aether-model-fetch --lib - cargo nextest run -p aether-gateway control::admin::provider_query model_fetch::runtime::tests
This commit is contained in:
3
Cargo.lock
generated
3
Cargo.lock
generated
@@ -206,8 +206,11 @@ dependencies = [
|
|||||||
"aether-provider-transport",
|
"aether-provider-transport",
|
||||||
"aether-scheduler-core",
|
"aether-scheduler-core",
|
||||||
"async-trait",
|
"async-trait",
|
||||||
|
"base64 0.22.1",
|
||||||
"regex",
|
"regex",
|
||||||
|
"rsa",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
|
"sha2",
|
||||||
"tokio",
|
"tokio",
|
||||||
"uuid",
|
"uuid",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -16,12 +16,12 @@ use aether_data_contracts::repository::provider_catalog::{
|
|||||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||||
};
|
};
|
||||||
use aether_model_fetch::{
|
use aether_model_fetch::{
|
||||||
aggregate_models_for_cache, build_models_fetch_execution_plan,
|
aggregate_models_for_cache, fetch_models_from_transports, json_string_list,
|
||||||
endpoint_supports_rust_models_fetch, extract_error_message, parse_models_response,
|
preset_models_for_provider, selected_models_fetch_endpoints,
|
||||||
};
|
};
|
||||||
use axum::{body::Body, http::Response, response::IntoResponse, Json};
|
use axum::{body::Body, http::Response, response::IntoResponse, Json};
|
||||||
use serde_json::{json, Value};
|
use serde_json::{json, Value};
|
||||||
use std::collections::{BTreeMap, BTreeSet};
|
use std::collections::BTreeSet;
|
||||||
|
|
||||||
pub(crate) const ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_MESSAGE: &str =
|
pub(crate) const ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_MESSAGE: &str =
|
||||||
"Rust local provider-query model test is not configured";
|
"Rust local provider-query model test is not configured";
|
||||||
@@ -31,22 +31,15 @@ const ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL: &str =
|
|||||||
"No active endpoints found for this provider";
|
"No active endpoints found for this provider";
|
||||||
const ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_ENDPOINT_DETAIL: &str =
|
const ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_ENDPOINT_DETAIL: &str =
|
||||||
"No models returned from any endpoint";
|
"No models returned from any endpoint";
|
||||||
const PROVIDER_QUERY_FETCH_FORMAT_PRIORITY: &[&[&str]] = &[
|
const ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL: &str = "No models returned from any key";
|
||||||
&[
|
const ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX: &str = "upstream_models_provider:";
|
||||||
"openai:chat",
|
|
||||||
"openai:responses",
|
|
||||||
"openai:cli",
|
|
||||||
"openai:compact",
|
|
||||||
],
|
|
||||||
&["claude:chat", "claude:cli"],
|
|
||||||
&["gemini:chat", "gemini:cli"],
|
|
||||||
];
|
|
||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
struct ProviderQueryKeyFetchResult {
|
struct ProviderQueryKeyFetchResult {
|
||||||
models: Vec<Value>,
|
models: Vec<Value>,
|
||||||
error: Option<String>,
|
error: Option<String>,
|
||||||
from_cache: bool,
|
from_cache: bool,
|
||||||
|
has_success: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
fn provider_query_provider_payload(provider: &StoredProviderCatalogProvider) -> Value {
|
fn provider_query_provider_payload(provider: &StoredProviderCatalogProvider) -> Value {
|
||||||
@@ -66,64 +59,6 @@ fn provider_query_key_display_name(key: &StoredProviderCatalogKey) -> String {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn provider_query_normalize_api_format(value: &str) -> String {
|
|
||||||
value.trim().to_ascii_lowercase()
|
|
||||||
}
|
|
||||||
|
|
||||||
fn provider_query_selected_fetch_endpoints(
|
|
||||||
endpoints: &[StoredProviderCatalogEndpoint],
|
|
||||||
key: &StoredProviderCatalogKey,
|
|
||||||
) -> Vec<StoredProviderCatalogEndpoint> {
|
|
||||||
let allowed_api_formats = key
|
|
||||||
.api_formats
|
|
||||||
.as_ref()
|
|
||||||
.and_then(Value::as_array)
|
|
||||||
.map(|items| {
|
|
||||||
items
|
|
||||||
.iter()
|
|
||||||
.filter_map(Value::as_str)
|
|
||||||
.map(provider_query_normalize_api_format)
|
|
||||||
.filter(|value| !value.is_empty())
|
|
||||||
.collect::<BTreeSet<_>>()
|
|
||||||
})
|
|
||||||
.filter(|items| !items.is_empty());
|
|
||||||
let mut by_format = BTreeMap::<String, StoredProviderCatalogEndpoint>::new();
|
|
||||||
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
|
|
||||||
let api_format = provider_query_normalize_api_format(&endpoint.api_format);
|
|
||||||
if api_format.is_empty() || !endpoint_supports_rust_models_fetch(&api_format) {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
if allowed_api_formats
|
|
||||||
.as_ref()
|
|
||||||
.is_some_and(|formats| !formats.contains(&api_format))
|
|
||||||
{
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
by_format.insert(api_format, endpoint.clone());
|
|
||||||
}
|
|
||||||
|
|
||||||
// 与 Python 版本保持一致:同族优先使用 chat 端点,其次才回退到其他抓取格式。
|
|
||||||
let covered_formats = PROVIDER_QUERY_FETCH_FORMAT_PRIORITY
|
|
||||||
.iter()
|
|
||||||
.flat_map(|items| items.iter().copied())
|
|
||||||
.collect::<BTreeSet<_>>();
|
|
||||||
let mut selected = PROVIDER_QUERY_FETCH_FORMAT_PRIORITY
|
|
||||||
.iter()
|
|
||||||
.filter_map(|candidates| {
|
|
||||||
candidates
|
|
||||||
.iter()
|
|
||||||
.find_map(|api_format| by_format.remove(*api_format))
|
|
||||||
})
|
|
||||||
.collect::<Vec<_>>();
|
|
||||||
selected.extend(
|
|
||||||
by_format
|
|
||||||
.into_iter()
|
|
||||||
.filter(|(api_format, _)| !covered_formats.contains(api_format.as_str()))
|
|
||||||
.map(|(_, endpoint)| endpoint),
|
|
||||||
);
|
|
||||||
selected
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn provider_query_read_cached_models(
|
async fn provider_query_read_cached_models(
|
||||||
state: &AdminAppState<'_>,
|
state: &AdminAppState<'_>,
|
||||||
provider_id: &str,
|
provider_id: &str,
|
||||||
@@ -147,38 +82,97 @@ async fn provider_query_read_cached_models(
|
|||||||
Some(aggregate_models_for_cache(&parsed))
|
Some(aggregate_models_for_cache(&parsed))
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn provider_query_fetch_models_from_transport(
|
async fn provider_query_read_provider_cached_models(
|
||||||
state: &AdminAppState<'_>,
|
state: &AdminAppState<'_>,
|
||||||
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
|
provider_id: &str,
|
||||||
) -> Result<Vec<Value>, String> {
|
) -> Option<Vec<Value>> {
|
||||||
let plan = build_models_fetch_execution_plan(state.app(), transport).await?;
|
let runner = state.app().redis_kv_runner()?;
|
||||||
let result = execution_runtime::execute_execution_runtime_sync_plan(state.app(), None, &plan)
|
let cache_key = runner.keyspace().key(&format!(
|
||||||
|
"{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}"
|
||||||
|
));
|
||||||
|
let mut connection = runner
|
||||||
|
.client()
|
||||||
|
.get_multiplexed_async_connection()
|
||||||
.await
|
.await
|
||||||
.map_err(|err| format!("{err:?}"))?;
|
.ok()?;
|
||||||
|
let raw = redis::cmd("GET")
|
||||||
|
.arg(&cache_key)
|
||||||
|
.query_async::<Option<String>>(&mut connection)
|
||||||
|
.await
|
||||||
|
.ok()??;
|
||||||
|
let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
|
||||||
|
Some(aggregate_models_for_cache(&parsed))
|
||||||
|
}
|
||||||
|
|
||||||
if result.status_code != 200 {
|
async fn provider_query_write_provider_cached_models(
|
||||||
let message = result
|
state: &AdminAppState<'_>,
|
||||||
.body
|
provider_id: &str,
|
||||||
.as_ref()
|
models: &[Value],
|
||||||
.and_then(|body| body.json_body.as_ref())
|
) {
|
||||||
.and_then(extract_error_message)
|
let Some(runner) = state.app().redis_kv_runner() else {
|
||||||
.or_else(|| {
|
return;
|
||||||
result.error.as_ref().and_then(|error| {
|
};
|
||||||
let message = error.message.trim();
|
let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(models)) else {
|
||||||
(!message.is_empty()).then_some(message.to_string())
|
return;
|
||||||
|
};
|
||||||
|
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
|
||||||
|
let _ = runner
|
||||||
|
.setex(
|
||||||
|
&cache_key,
|
||||||
|
&serialized,
|
||||||
|
Some(aether_model_fetch::model_fetch_interval_minutes().saturating_mul(60)),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
|
||||||
|
fn provider_query_antigravity_tier_weight(raw_auth_config: Option<&str>) -> i32 {
|
||||||
|
raw_auth_config
|
||||||
|
.and_then(|value| serde_json::from_str::<Value>(value).ok())
|
||||||
|
.and_then(|value| value.get("tier").cloned())
|
||||||
|
.and_then(|value| value.as_str().map(ToOwned::to_owned))
|
||||||
|
.map(|tier| match tier.trim().to_ascii_lowercase().as_str() {
|
||||||
|
"ultra" => 3,
|
||||||
|
"pro" => 2,
|
||||||
|
"free" => 1,
|
||||||
|
_ => 0,
|
||||||
|
})
|
||||||
|
.unwrap_or(0)
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn provider_query_sort_antigravity_keys(
|
||||||
|
state: &AdminAppState<'_>,
|
||||||
|
provider: &StoredProviderCatalogProvider,
|
||||||
|
endpoints: &[StoredProviderCatalogEndpoint],
|
||||||
|
keys: Vec<StoredProviderCatalogKey>,
|
||||||
|
) -> Result<Vec<StoredProviderCatalogKey>, GatewayError> {
|
||||||
|
let mut ranked = Vec::new();
|
||||||
|
for key in keys {
|
||||||
|
let availability = if key.oauth_invalid_at_unix_secs.is_some() {
|
||||||
|
0
|
||||||
|
} else {
|
||||||
|
1
|
||||||
|
};
|
||||||
|
let tier_weight = if let Some(endpoint) = selected_models_fetch_endpoints(endpoints, &key)
|
||||||
|
.into_iter()
|
||||||
|
.next()
|
||||||
|
{
|
||||||
|
state
|
||||||
|
.app()
|
||||||
|
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||||
|
.await?
|
||||||
|
.map(|transport| {
|
||||||
|
provider_query_antigravity_tier_weight(
|
||||||
|
transport.key.decrypted_auth_config.as_deref(),
|
||||||
|
)
|
||||||
})
|
})
|
||||||
})
|
.unwrap_or(0)
|
||||||
.unwrap_or_else(|| format!("upstream returned status {}", result.status_code));
|
} else {
|
||||||
return Err(message);
|
0
|
||||||
|
};
|
||||||
|
ranked.push(((availability, tier_weight), key));
|
||||||
}
|
}
|
||||||
|
ranked.sort_by(|left, right| right.0.cmp(&left.0));
|
||||||
let body_json = result
|
Ok(ranked.into_iter().map(|(_, key)| key).collect())
|
||||||
.body
|
|
||||||
.as_ref()
|
|
||||||
.and_then(|body| body.json_body.as_ref())
|
|
||||||
.ok_or_else(|| "models fetch response body is missing JSON payload".to_string())?;
|
|
||||||
let parsed = parse_models_response(&transport.endpoint.api_format, body_json)?;
|
|
||||||
Ok(parsed.cached_models)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn provider_query_fetch_models_for_key(
|
async fn provider_query_fetch_models_for_key(
|
||||||
@@ -196,20 +190,30 @@ async fn provider_query_fetch_models_for_key(
|
|||||||
models: cached_models,
|
models: cached_models,
|
||||||
error: None,
|
error: None,
|
||||||
from_cache: true,
|
from_cache: true,
|
||||||
|
has_success: true,
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
let selected_endpoints = provider_query_selected_fetch_endpoints(endpoints, key);
|
let selected_endpoints = selected_models_fetch_endpoints(endpoints, key);
|
||||||
if selected_endpoints.is_empty() {
|
if selected_endpoints.is_empty() {
|
||||||
|
if let Some(models) = preset_models_for_provider(&provider.provider_type) {
|
||||||
|
return Ok(ProviderQueryKeyFetchResult {
|
||||||
|
models: aggregate_models_for_cache(&models),
|
||||||
|
error: None,
|
||||||
|
from_cache: false,
|
||||||
|
has_success: true,
|
||||||
|
});
|
||||||
|
}
|
||||||
return Ok(ProviderQueryKeyFetchResult {
|
return Ok(ProviderQueryKeyFetchResult {
|
||||||
models: Vec::new(),
|
models: Vec::new(),
|
||||||
error: Some(ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL.to_string()),
|
error: Some(ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL.to_string()),
|
||||||
from_cache: false,
|
from_cache: false,
|
||||||
|
has_success: false,
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut all_models = Vec::new();
|
let mut transports = Vec::new();
|
||||||
let mut all_errors = Vec::new();
|
let mut all_errors = Vec::new();
|
||||||
for endpoint in selected_endpoints {
|
for endpoint in selected_endpoints {
|
||||||
let Some(transport) = state
|
let Some(transport) = state
|
||||||
@@ -223,14 +227,34 @@ async fn provider_query_fetch_models_for_key(
|
|||||||
));
|
));
|
||||||
continue;
|
continue;
|
||||||
};
|
};
|
||||||
match provider_query_fetch_models_from_transport(state, &transport).await {
|
transports.push(transport);
|
||||||
Ok(models) => all_models.extend(models),
|
|
||||||
Err(err) => all_errors.push(err),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
let unique_models = aggregate_models_for_cache(&all_models);
|
if transports.is_empty() {
|
||||||
if !unique_models.is_empty() {
|
return Ok(ProviderQueryKeyFetchResult {
|
||||||
|
models: Vec::new(),
|
||||||
|
error: Some(all_errors.join("; ")),
|
||||||
|
from_cache: false,
|
||||||
|
has_success: false,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
let outcome = match fetch_models_from_transports(state.app(), &transports).await {
|
||||||
|
Ok(outcome) => outcome,
|
||||||
|
Err(err) => {
|
||||||
|
all_errors.push(err);
|
||||||
|
return Ok(ProviderQueryKeyFetchResult {
|
||||||
|
models: Vec::new(),
|
||||||
|
error: Some(all_errors.join("; ")),
|
||||||
|
from_cache: false,
|
||||||
|
has_success: false,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
all_errors.extend(outcome.errors);
|
||||||
|
let unique_models = aggregate_models_for_cache(&outcome.cached_models);
|
||||||
|
if outcome.has_success && !unique_models.is_empty() {
|
||||||
<AppState as ModelFetchRuntimeState>::write_upstream_models_cache(
|
<AppState as ModelFetchRuntimeState>::write_upstream_models_cache(
|
||||||
state.app(),
|
state.app(),
|
||||||
&provider.id,
|
&provider.id,
|
||||||
@@ -253,6 +277,7 @@ async fn provider_query_fetch_models_for_key(
|
|||||||
models: unique_models,
|
models: unique_models,
|
||||||
error,
|
error,
|
||||||
from_cache: false,
|
from_cache: false,
|
||||||
|
has_success: outcome.has_success,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -317,7 +342,10 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
|||||||
.into_response());
|
.into_response());
|
||||||
}
|
}
|
||||||
|
|
||||||
let active_keys = keys.iter().filter(|key| key.is_active).collect::<Vec<_>>();
|
let active_keys = keys
|
||||||
|
.into_iter()
|
||||||
|
.filter(|key| key.is_active)
|
||||||
|
.collect::<Vec<_>>();
|
||||||
if active_keys.is_empty() {
|
if active_keys.is_empty() {
|
||||||
return Ok(build_admin_provider_query_bad_request_response(
|
return Ok(build_admin_provider_query_bad_request_response(
|
||||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||||
@@ -325,11 +353,45 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
|||||||
}
|
}
|
||||||
let active_key_count = active_keys.len();
|
let active_key_count = active_keys.len();
|
||||||
|
|
||||||
|
if provider
|
||||||
|
.provider_type
|
||||||
|
.trim()
|
||||||
|
.eq_ignore_ascii_case("antigravity")
|
||||||
|
&& !force_refresh
|
||||||
|
{
|
||||||
|
if let Some(models) = provider_query_read_provider_cached_models(state, &provider.id).await
|
||||||
|
{
|
||||||
|
return Ok(Json(json!({
|
||||||
|
"success": !models.is_empty(),
|
||||||
|
"data": {
|
||||||
|
"models": models,
|
||||||
|
"error": serde_json::Value::Null,
|
||||||
|
"from_cache": true,
|
||||||
|
"keys_total": active_key_count,
|
||||||
|
"keys_cached": active_key_count,
|
||||||
|
"keys_fetched": 0,
|
||||||
|
},
|
||||||
|
"provider": provider_query_provider_payload(&provider),
|
||||||
|
}))
|
||||||
|
.into_response());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let ordered_keys = if provider
|
||||||
|
.provider_type
|
||||||
|
.trim()
|
||||||
|
.eq_ignore_ascii_case("antigravity")
|
||||||
|
{
|
||||||
|
provider_query_sort_antigravity_keys(state, &provider, &endpoints, active_keys).await?
|
||||||
|
} else {
|
||||||
|
active_keys
|
||||||
|
};
|
||||||
|
|
||||||
let mut all_models = Vec::new();
|
let mut all_models = Vec::new();
|
||||||
let mut all_errors = Vec::new();
|
let mut all_errors = Vec::new();
|
||||||
let mut cache_hit_count = 0usize;
|
let mut cache_hit_count = 0usize;
|
||||||
let mut fetch_count = 0usize;
|
let mut fetch_count = 0usize;
|
||||||
for key in active_keys {
|
for key in &ordered_keys {
|
||||||
let result =
|
let result =
|
||||||
provider_query_fetch_models_for_key(state, &provider, &endpoints, key, force_refresh)
|
provider_query_fetch_models_for_key(state, &provider, &endpoints, key, force_refresh)
|
||||||
.await?;
|
.await?;
|
||||||
@@ -346,9 +408,25 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
|||||||
} else {
|
} else {
|
||||||
fetch_count += 1;
|
fetch_count += 1;
|
||||||
}
|
}
|
||||||
|
if provider
|
||||||
|
.provider_type
|
||||||
|
.trim()
|
||||||
|
.eq_ignore_ascii_case("antigravity")
|
||||||
|
&& result.has_success
|
||||||
|
{
|
||||||
|
break;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
let models = aggregate_models_for_cache(&all_models);
|
let models = aggregate_models_for_cache(&all_models);
|
||||||
|
if provider
|
||||||
|
.provider_type
|
||||||
|
.trim()
|
||||||
|
.eq_ignore_ascii_case("antigravity")
|
||||||
|
&& !models.is_empty()
|
||||||
|
{
|
||||||
|
provider_query_write_provider_cached_models(state, &provider.id, &models).await;
|
||||||
|
}
|
||||||
let success = !models.is_empty();
|
let success = !models.is_empty();
|
||||||
let mut error = if all_errors.is_empty() {
|
let mut error = if all_errors.is_empty() {
|
||||||
None
|
None
|
||||||
@@ -356,7 +434,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
|||||||
Some(all_errors.join("; "))
|
Some(all_errors.join("; "))
|
||||||
};
|
};
|
||||||
if !success && error.is_none() {
|
if !success && error.is_none() {
|
||||||
error = Some("No models returned from any key".to_string());
|
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL.to_string());
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(Json(json!({
|
Ok(Json(json!({
|
||||||
|
|||||||
@@ -2,21 +2,18 @@ use std::collections::HashMap;
|
|||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
use std::time::{SystemTime, UNIX_EPOCH};
|
use std::time::{SystemTime, UNIX_EPOCH};
|
||||||
|
|
||||||
use aether_contracts::ExecutionResult;
|
|
||||||
use aether_data_contracts::repository::provider_catalog::{
|
use aether_data_contracts::repository::provider_catalog::{
|
||||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||||
};
|
};
|
||||||
use aether_model_fetch::{
|
use aether_model_fetch::{
|
||||||
apply_model_filters, build_models_fetch_execution_plan, extract_error_message,
|
apply_model_filters, fetch_models_from_transports, json_string_list, merge_upstream_metadata,
|
||||||
json_string_list, model_fetch_interval_minutes, model_fetch_startup_delay_seconds,
|
model_fetch_interval_minutes, model_fetch_startup_delay_seconds, model_fetch_startup_enabled,
|
||||||
model_fetch_startup_enabled, parse_models_response, select_models_fetch_endpoint,
|
preset_models_for_provider, selected_models_fetch_endpoints,
|
||||||
sync_provider_model_whitelist_associations, ModelFetchAssociationStore, ModelFetchRunSummary,
|
sync_provider_model_whitelist_associations, ModelFetchAssociationStore, ModelFetchRunSummary,
|
||||||
ModelsFetchSuccess,
|
|
||||||
};
|
};
|
||||||
use serde_json::{json, Value};
|
use serde_json::{json, Value};
|
||||||
use tracing::{debug, info, warn};
|
use tracing::{debug, info, warn};
|
||||||
|
|
||||||
use crate::provider_transport::GatewayProviderTransportSnapshot;
|
|
||||||
use crate::{AppState, GatewayError};
|
use crate::{AppState, GatewayError};
|
||||||
|
|
||||||
pub(crate) mod state;
|
pub(crate) mod state;
|
||||||
@@ -26,8 +23,8 @@ use self::state::ModelFetchRuntimeState;
|
|||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
struct SelectedFetchTarget {
|
struct SelectedFetchTarget {
|
||||||
provider: StoredProviderCatalogProvider,
|
provider: StoredProviderCatalogProvider,
|
||||||
endpoint: StoredProviderCatalogEndpoint,
|
|
||||||
key: StoredProviderCatalogKey,
|
key: StoredProviderCatalogKey,
|
||||||
|
endpoints: Vec<StoredProviderCatalogEndpoint>,
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn spawn_model_fetch_worker(state: AppState) -> Option<tokio::task::JoinHandle<()>> {
|
pub(crate) fn spawn_model_fetch_worker(state: AppState) -> Option<tokio::task::JoinHandle<()>> {
|
||||||
@@ -131,38 +128,12 @@ where
|
|||||||
if !key.is_active || !key.auto_fetch_models {
|
if !key.is_active || !key.auto_fetch_models {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
if let Some(endpoint) = select_models_fetch_endpoint(&endpoints, &key) {
|
let selected_endpoints = selected_models_fetch_endpoints(&endpoints, &key);
|
||||||
targets.push(SelectedFetchTarget {
|
targets.push(SelectedFetchTarget {
|
||||||
provider: provider.clone(),
|
provider: provider.clone(),
|
||||||
endpoint,
|
key,
|
||||||
key,
|
endpoints: selected_endpoints,
|
||||||
});
|
});
|
||||||
} else {
|
|
||||||
targets.push(SelectedFetchTarget {
|
|
||||||
provider: provider.clone(),
|
|
||||||
endpoint: StoredProviderCatalogEndpoint::new(
|
|
||||||
"__unsupported__".to_string(),
|
|
||||||
provider.id.clone(),
|
|
||||||
"__unsupported__".to_string(),
|
|
||||||
None,
|
|
||||||
None,
|
|
||||||
false,
|
|
||||||
)
|
|
||||||
.expect("unsupported sentinel endpoint should build")
|
|
||||||
.with_transport_fields(
|
|
||||||
"https://unsupported.invalid".to_string(),
|
|
||||||
None,
|
|
||||||
None,
|
|
||||||
None,
|
|
||||||
None,
|
|
||||||
None,
|
|
||||||
None,
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
.expect("unsupported sentinel endpoint transport should build"),
|
|
||||||
key,
|
|
||||||
});
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -215,7 +186,34 @@ async fn fetch_and_persist_key_models(
|
|||||||
target: &SelectedFetchTarget,
|
target: &SelectedFetchTarget,
|
||||||
) -> Result<KeyFetchDisposition, GatewayError> {
|
) -> Result<KeyFetchDisposition, GatewayError> {
|
||||||
let now_unix_secs = now_unix_secs();
|
let now_unix_secs = now_unix_secs();
|
||||||
if target.endpoint.api_format == "__unsupported__" {
|
if target.endpoints.is_empty() {
|
||||||
|
if let Some(models) = preset_models_for_provider(&target.provider.provider_type) {
|
||||||
|
let fetched_model_ids = models
|
||||||
|
.iter()
|
||||||
|
.filter_map(|model| model.get("id"))
|
||||||
|
.filter_map(Value::as_str)
|
||||||
|
.map(ToOwned::to_owned)
|
||||||
|
.collect::<Vec<_>>();
|
||||||
|
let filtered_models = apply_model_filters(
|
||||||
|
&fetched_model_ids,
|
||||||
|
json_string_list(target.key.locked_models.as_ref()),
|
||||||
|
json_string_list(target.key.model_include_patterns.as_ref()),
|
||||||
|
json_string_list(target.key.model_exclude_patterns.as_ref()),
|
||||||
|
);
|
||||||
|
persist_key_fetch_success(state, &target.key, now_unix_secs, &filtered_models, None)
|
||||||
|
.await?;
|
||||||
|
state
|
||||||
|
.write_upstream_models_cache(&target.provider.id, &target.key.id, &models)
|
||||||
|
.await;
|
||||||
|
sync_provider_model_whitelist_associations(
|
||||||
|
state,
|
||||||
|
&target.provider.id,
|
||||||
|
&filtered_models,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.map_err(GatewayError::Internal)?;
|
||||||
|
return Ok(KeyFetchDisposition::Succeeded);
|
||||||
|
}
|
||||||
persist_key_fetch_failure(
|
persist_key_fetch_failure(
|
||||||
state,
|
state,
|
||||||
&target.key,
|
&target.key,
|
||||||
@@ -226,10 +224,25 @@ async fn fetch_and_persist_key_models(
|
|||||||
return Ok(KeyFetchDisposition::Skipped);
|
return Ok(KeyFetchDisposition::Skipped);
|
||||||
}
|
}
|
||||||
|
|
||||||
let Some(transport) = state
|
let mut transports = Vec::new();
|
||||||
.read_provider_transport_snapshot(&target.provider.id, &target.endpoint.id, &target.key.id)
|
for endpoint in &target.endpoints {
|
||||||
.await?
|
match state
|
||||||
else {
|
.read_provider_transport_snapshot(&target.provider.id, &endpoint.id, &target.key.id)
|
||||||
|
.await?
|
||||||
|
{
|
||||||
|
Some(transport) => transports.push(transport),
|
||||||
|
None => {
|
||||||
|
warn!(
|
||||||
|
provider_id = %target.provider.id,
|
||||||
|
endpoint_id = %endpoint.id,
|
||||||
|
key_id = %target.key.id,
|
||||||
|
"gateway model fetch transport snapshot unavailable"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if transports.is_empty() {
|
||||||
persist_key_fetch_failure(
|
persist_key_fetch_failure(
|
||||||
state,
|
state,
|
||||||
&target.key,
|
&target.key,
|
||||||
@@ -238,9 +251,9 @@ async fn fetch_and_persist_key_models(
|
|||||||
)
|
)
|
||||||
.await?;
|
.await?;
|
||||||
return Ok(KeyFetchDisposition::Skipped);
|
return Ok(KeyFetchDisposition::Skipped);
|
||||||
};
|
}
|
||||||
|
|
||||||
let result = match execute_models_fetch_request(state, &transport).await {
|
let result = match fetch_models_from_transports(state, &transports).await {
|
||||||
Ok(result) => result,
|
Ok(result) => result,
|
||||||
Err(err) => {
|
Err(err) => {
|
||||||
persist_key_fetch_failure(state, &target.key, now_unix_secs, err.clone()).await?;
|
persist_key_fetch_failure(state, &target.key, now_unix_secs, err.clone()).await?;
|
||||||
@@ -254,6 +267,22 @@ async fn fetch_and_persist_key_models(
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
if !result.has_success {
|
||||||
|
let error = if result.errors.is_empty() {
|
||||||
|
"Upstream models fetch failed".to_string()
|
||||||
|
} else {
|
||||||
|
result.errors.join("; ")
|
||||||
|
};
|
||||||
|
persist_key_fetch_failure(state, &target.key, now_unix_secs, error.clone()).await?;
|
||||||
|
warn!(
|
||||||
|
provider_id = %target.provider.id,
|
||||||
|
key_id = %target.key.id,
|
||||||
|
message = %error,
|
||||||
|
"gateway model fetch failed"
|
||||||
|
);
|
||||||
|
return Ok(KeyFetchDisposition::Failed);
|
||||||
|
}
|
||||||
|
|
||||||
let filtered_models = apply_model_filters(
|
let filtered_models = apply_model_filters(
|
||||||
&result.fetched_model_ids,
|
&result.fetched_model_ids,
|
||||||
json_string_list(target.key.locked_models.as_ref()),
|
json_string_list(target.key.locked_models.as_ref()),
|
||||||
@@ -261,7 +290,14 @@ async fn fetch_and_persist_key_models(
|
|||||||
json_string_list(target.key.model_exclude_patterns.as_ref()),
|
json_string_list(target.key.model_exclude_patterns.as_ref()),
|
||||||
);
|
);
|
||||||
|
|
||||||
persist_key_fetch_success(state, &target.key, now_unix_secs, &filtered_models).await?;
|
persist_key_fetch_success(
|
||||||
|
state,
|
||||||
|
&target.key,
|
||||||
|
now_unix_secs,
|
||||||
|
&filtered_models,
|
||||||
|
result.upstream_metadata.as_ref(),
|
||||||
|
)
|
||||||
|
.await?;
|
||||||
state
|
state
|
||||||
.write_upstream_models_cache(&target.provider.id, &target.key.id, &result.cached_models)
|
.write_upstream_models_cache(&target.provider.id, &target.key.id, &result.cached_models)
|
||||||
.await;
|
.await;
|
||||||
@@ -271,41 +307,6 @@ async fn fetch_and_persist_key_models(
|
|||||||
Ok(KeyFetchDisposition::Succeeded)
|
Ok(KeyFetchDisposition::Succeeded)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn execute_models_fetch_request(
|
|
||||||
state: &(impl ModelFetchRuntimeState + ?Sized),
|
|
||||||
transport: &GatewayProviderTransportSnapshot,
|
|
||||||
) -> Result<ModelsFetchSuccess, String> {
|
|
||||||
let plan = build_models_fetch_execution_plan(state, transport).await?;
|
|
||||||
|
|
||||||
let result = state
|
|
||||||
.execute_execution_runtime_sync_plan(&plan)
|
|
||||||
.await
|
|
||||||
.map_err(|err| format!("{err:?}"))?;
|
|
||||||
|
|
||||||
if result.status_code != 200 {
|
|
||||||
let message = result
|
|
||||||
.body
|
|
||||||
.as_ref()
|
|
||||||
.and_then(|body| body.json_body.as_ref())
|
|
||||||
.and_then(extract_error_message)
|
|
||||||
.or_else(|| {
|
|
||||||
result.error.as_ref().and_then(|error| {
|
|
||||||
let message = error.message.trim();
|
|
||||||
(!message.is_empty()).then_some(message.to_string())
|
|
||||||
})
|
|
||||||
})
|
|
||||||
.unwrap_or_else(|| format!("upstream returned status {}", result.status_code));
|
|
||||||
return Err(message);
|
|
||||||
}
|
|
||||||
|
|
||||||
let body_json = result
|
|
||||||
.body
|
|
||||||
.as_ref()
|
|
||||||
.and_then(|body| body.json_body.as_ref())
|
|
||||||
.ok_or_else(|| "models fetch response body is missing JSON payload".to_string())?;
|
|
||||||
parse_models_response(&transport.endpoint.api_format, body_json)
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn persist_key_fetch_failure(
|
async fn persist_key_fetch_failure(
|
||||||
state: &(impl ModelFetchRuntimeState + ?Sized),
|
state: &(impl ModelFetchRuntimeState + ?Sized),
|
||||||
key: &StoredProviderCatalogKey,
|
key: &StoredProviderCatalogKey,
|
||||||
@@ -325,6 +326,7 @@ async fn persist_key_fetch_success(
|
|||||||
key: &StoredProviderCatalogKey,
|
key: &StoredProviderCatalogKey,
|
||||||
now_unix_secs: u64,
|
now_unix_secs: u64,
|
||||||
allowed_models: &[String],
|
allowed_models: &[String],
|
||||||
|
upstream_metadata: Option<&Value>,
|
||||||
) -> Result<(), GatewayError> {
|
) -> Result<(), GatewayError> {
|
||||||
let mut updated = key.clone();
|
let mut updated = key.clone();
|
||||||
updated.allowed_models = if allowed_models.is_empty() {
|
updated.allowed_models = if allowed_models.is_empty() {
|
||||||
@@ -332,6 +334,12 @@ async fn persist_key_fetch_success(
|
|||||||
} else {
|
} else {
|
||||||
Some(json!(allowed_models))
|
Some(json!(allowed_models))
|
||||||
};
|
};
|
||||||
|
if let Some(upstream_metadata) = upstream_metadata {
|
||||||
|
updated.upstream_metadata = Some(merge_upstream_metadata(
|
||||||
|
updated.upstream_metadata.as_ref(),
|
||||||
|
upstream_metadata,
|
||||||
|
));
|
||||||
|
}
|
||||||
updated.last_models_fetch_at_unix_secs = Some(now_unix_secs);
|
updated.last_models_fetch_at_unix_secs = Some(now_unix_secs);
|
||||||
updated.last_models_fetch_error = None;
|
updated.last_models_fetch_error = None;
|
||||||
updated.updated_at_unix_secs = Some(now_unix_secs);
|
updated.updated_at_unix_secs = Some(now_unix_secs);
|
||||||
@@ -345,3 +353,533 @@ fn now_unix_secs() -> u64 {
|
|||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
.as_secs()
|
.as_secs()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::{perform_model_fetch_once_with_state, state::ModelFetchRuntimeState};
|
||||||
|
use aether_contracts::{ExecutionPlan, ExecutionResult, ProxySnapshot};
|
||||||
|
use aether_data_contracts::repository::global_models::{
|
||||||
|
AdminGlobalModelListQuery, AdminProviderModelListQuery, StoredAdminGlobalModelPage,
|
||||||
|
StoredAdminProviderModel, UpsertAdminProviderModelRecord,
|
||||||
|
};
|
||||||
|
use aether_data_contracts::repository::provider_catalog::{
|
||||||
|
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||||
|
};
|
||||||
|
use aether_model_fetch::{ModelFetchAssociationStore, ModelFetchTransportRuntime};
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use serde_json::{json, Value};
|
||||||
|
use std::collections::{HashMap, VecDeque};
|
||||||
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
|
use crate::provider_transport::LocalResolvedOAuthRequestAuth;
|
||||||
|
use crate::GatewayError;
|
||||||
|
use aether_provider_transport::snapshot::{
|
||||||
|
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||||
|
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||||
|
};
|
||||||
|
|
||||||
|
#[derive(Clone, Default)]
|
||||||
|
struct TestState {
|
||||||
|
providers: Arc<Vec<StoredProviderCatalogProvider>>,
|
||||||
|
endpoints: Arc<Vec<StoredProviderCatalogEndpoint>>,
|
||||||
|
keys: Arc<Mutex<Vec<StoredProviderCatalogKey>>>,
|
||||||
|
transports: Arc<HashMap<(String, String, String), GatewayProviderTransportSnapshot>>,
|
||||||
|
execution_results: Arc<Mutex<VecDeque<ExecutionResult>>>,
|
||||||
|
cached_models: Arc<Mutex<HashMap<(String, String), Vec<Value>>>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl TestState {
|
||||||
|
fn new(
|
||||||
|
providers: Vec<StoredProviderCatalogProvider>,
|
||||||
|
endpoints: Vec<StoredProviderCatalogEndpoint>,
|
||||||
|
keys: Vec<StoredProviderCatalogKey>,
|
||||||
|
transports: HashMap<(String, String, String), GatewayProviderTransportSnapshot>,
|
||||||
|
execution_results: Vec<ExecutionResult>,
|
||||||
|
) -> Self {
|
||||||
|
Self {
|
||||||
|
providers: Arc::new(providers),
|
||||||
|
endpoints: Arc::new(endpoints),
|
||||||
|
keys: Arc::new(Mutex::new(keys)),
|
||||||
|
transports: Arc::new(transports),
|
||||||
|
execution_results: Arc::new(Mutex::new(VecDeque::from(execution_results))),
|
||||||
|
cached_models: Arc::new(Mutex::new(HashMap::new())),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn key(&self, key_id: &str) -> StoredProviderCatalogKey {
|
||||||
|
self.keys
|
||||||
|
.lock()
|
||||||
|
.expect("keys mutex")
|
||||||
|
.iter()
|
||||||
|
.find(|key| key.id == key_id)
|
||||||
|
.cloned()
|
||||||
|
.expect("key should exist")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl ModelFetchTransportRuntime for TestState {
|
||||||
|
async fn resolve_local_oauth_request_auth(
|
||||||
|
&self,
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
) -> Result<Option<LocalResolvedOAuthRequestAuth>, String> {
|
||||||
|
if transport.key.auth_type.trim().eq_ignore_ascii_case("oauth") {
|
||||||
|
return Ok(Some(LocalResolvedOAuthRequestAuth::Header {
|
||||||
|
name: "authorization".to_string(),
|
||||||
|
value: "Bearer oauth-token".to_string(),
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
Ok(None)
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn resolve_model_fetch_proxy(
|
||||||
|
&self,
|
||||||
|
_transport: &GatewayProviderTransportSnapshot,
|
||||||
|
) -> Option<ProxySnapshot> {
|
||||||
|
None
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn execute_model_fetch_execution_plan(
|
||||||
|
&self,
|
||||||
|
_plan: &ExecutionPlan,
|
||||||
|
) -> Result<ExecutionResult, String> {
|
||||||
|
self.execution_results
|
||||||
|
.lock()
|
||||||
|
.expect("execution result mutex")
|
||||||
|
.pop_front()
|
||||||
|
.ok_or_else(|| "missing execution result".to_string())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl ModelFetchAssociationStore for TestState {
|
||||||
|
type Error = String;
|
||||||
|
|
||||||
|
fn has_global_model_reader(&self) -> bool {
|
||||||
|
false
|
||||||
|
}
|
||||||
|
|
||||||
|
fn has_global_model_writer(&self) -> bool {
|
||||||
|
false
|
||||||
|
}
|
||||||
|
|
||||||
|
fn model_fetch_internal_error(&self, message: String) -> Self::Error {
|
||||||
|
message
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn list_admin_provider_models(
|
||||||
|
&self,
|
||||||
|
_query: &AdminProviderModelListQuery,
|
||||||
|
) -> Result<Vec<StoredAdminProviderModel>, Self::Error> {
|
||||||
|
Ok(Vec::new())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn list_admin_global_models(
|
||||||
|
&self,
|
||||||
|
_query: &AdminGlobalModelListQuery,
|
||||||
|
) -> Result<StoredAdminGlobalModelPage, Self::Error> {
|
||||||
|
Ok(StoredAdminGlobalModelPage {
|
||||||
|
items: Vec::new(),
|
||||||
|
total: 0,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn create_admin_provider_model(
|
||||||
|
&self,
|
||||||
|
_record: &UpsertAdminProviderModelRecord,
|
||||||
|
) -> Result<Option<StoredAdminProviderModel>, Self::Error> {
|
||||||
|
Ok(None)
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn list_provider_catalog_keys_by_provider_ids(
|
||||||
|
&self,
|
||||||
|
provider_ids: &[String],
|
||||||
|
) -> Result<Vec<StoredProviderCatalogKey>, Self::Error> {
|
||||||
|
Ok(self
|
||||||
|
.keys
|
||||||
|
.lock()
|
||||||
|
.expect("keys mutex")
|
||||||
|
.iter()
|
||||||
|
.filter(|key| {
|
||||||
|
provider_ids
|
||||||
|
.iter()
|
||||||
|
.any(|provider_id| provider_id == &key.provider_id)
|
||||||
|
})
|
||||||
|
.cloned()
|
||||||
|
.collect())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn delete_admin_provider_model(
|
||||||
|
&self,
|
||||||
|
_provider_id: &str,
|
||||||
|
_model_id: &str,
|
||||||
|
) -> Result<bool, Self::Error> {
|
||||||
|
Ok(false)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl ModelFetchRuntimeState for TestState {
|
||||||
|
fn has_provider_catalog_data_reader(&self) -> bool {
|
||||||
|
true
|
||||||
|
}
|
||||||
|
|
||||||
|
fn has_provider_catalog_data_writer(&self) -> bool {
|
||||||
|
true
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn list_provider_catalog_providers(
|
||||||
|
&self,
|
||||||
|
_active_only: bool,
|
||||||
|
) -> Result<Vec<StoredProviderCatalogProvider>, GatewayError> {
|
||||||
|
Ok(self.providers.as_ref().clone())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn list_provider_catalog_endpoints_by_provider_ids(
|
||||||
|
&self,
|
||||||
|
provider_ids: &[String],
|
||||||
|
) -> Result<Vec<StoredProviderCatalogEndpoint>, GatewayError> {
|
||||||
|
Ok(self
|
||||||
|
.endpoints
|
||||||
|
.iter()
|
||||||
|
.filter(|endpoint| {
|
||||||
|
provider_ids
|
||||||
|
.iter()
|
||||||
|
.any(|provider_id| provider_id == &endpoint.provider_id)
|
||||||
|
})
|
||||||
|
.cloned()
|
||||||
|
.collect())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn read_provider_transport_snapshot(
|
||||||
|
&self,
|
||||||
|
provider_id: &str,
|
||||||
|
endpoint_id: &str,
|
||||||
|
key_id: &str,
|
||||||
|
) -> Result<Option<GatewayProviderTransportSnapshot>, GatewayError> {
|
||||||
|
Ok(self
|
||||||
|
.transports
|
||||||
|
.get(&(
|
||||||
|
provider_id.to_string(),
|
||||||
|
endpoint_id.to_string(),
|
||||||
|
key_id.to_string(),
|
||||||
|
))
|
||||||
|
.cloned())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn execute_execution_runtime_sync_plan(
|
||||||
|
&self,
|
||||||
|
_plan: &ExecutionPlan,
|
||||||
|
) -> Result<ExecutionResult, GatewayError> {
|
||||||
|
Err(GatewayError::Internal(
|
||||||
|
"execute_execution_runtime_sync_plan should not be called".to_string(),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn update_provider_catalog_key(
|
||||||
|
&self,
|
||||||
|
key: &StoredProviderCatalogKey,
|
||||||
|
) -> Result<(), GatewayError> {
|
||||||
|
let mut keys = self.keys.lock().expect("keys mutex");
|
||||||
|
let Some(slot) = keys.iter_mut().find(|item| item.id == key.id) else {
|
||||||
|
return Err(GatewayError::Internal("key not found".to_string()));
|
||||||
|
};
|
||||||
|
*slot = key.clone();
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn write_upstream_models_cache(
|
||||||
|
&self,
|
||||||
|
provider_id: &str,
|
||||||
|
key_id: &str,
|
||||||
|
cached_models: &[Value],
|
||||||
|
) {
|
||||||
|
self.cached_models.lock().expect("cache mutex").insert(
|
||||||
|
(provider_id.to_string(), key_id.to_string()),
|
||||||
|
cached_models.to_vec(),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sample_provider(provider_id: &str, provider_type: &str) -> StoredProviderCatalogProvider {
|
||||||
|
StoredProviderCatalogProvider::new(
|
||||||
|
provider_id.to_string(),
|
||||||
|
provider_id.to_string(),
|
||||||
|
None,
|
||||||
|
provider_type.to_string(),
|
||||||
|
)
|
||||||
|
.expect("provider should build")
|
||||||
|
.with_transport_fields(true, false, false, None, None, None, None, None, None)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sample_endpoint(
|
||||||
|
endpoint_id: &str,
|
||||||
|
provider_id: &str,
|
||||||
|
api_format: &str,
|
||||||
|
) -> StoredProviderCatalogEndpoint {
|
||||||
|
StoredProviderCatalogEndpoint::new(
|
||||||
|
endpoint_id.to_string(),
|
||||||
|
provider_id.to_string(),
|
||||||
|
api_format.to_string(),
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
true,
|
||||||
|
)
|
||||||
|
.expect("endpoint should build")
|
||||||
|
.with_transport_fields(
|
||||||
|
"https://cloudcode-pa.googleapis.com".to_string(),
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.expect("endpoint transport should build")
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sample_key(
|
||||||
|
key_id: &str,
|
||||||
|
provider_id: &str,
|
||||||
|
auth_type: &str,
|
||||||
|
api_formats: &[&str],
|
||||||
|
) -> StoredProviderCatalogKey {
|
||||||
|
let mut key = StoredProviderCatalogKey::new(
|
||||||
|
key_id.to_string(),
|
||||||
|
provider_id.to_string(),
|
||||||
|
"primary".to_string(),
|
||||||
|
auth_type.to_string(),
|
||||||
|
None,
|
||||||
|
true,
|
||||||
|
)
|
||||||
|
.expect("key should build")
|
||||||
|
.with_transport_fields(
|
||||||
|
Some(json!(api_formats)),
|
||||||
|
"encrypted".to_string(),
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.expect("key transport should build");
|
||||||
|
key.auto_fetch_models = true;
|
||||||
|
key
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sample_transport(
|
||||||
|
provider_type: &str,
|
||||||
|
provider_id: &str,
|
||||||
|
endpoint_id: &str,
|
||||||
|
key_id: &str,
|
||||||
|
api_format: &str,
|
||||||
|
auth_type: &str,
|
||||||
|
decrypted_auth_config: Option<&str>,
|
||||||
|
) -> GatewayProviderTransportSnapshot {
|
||||||
|
GatewayProviderTransportSnapshot {
|
||||||
|
provider: GatewayProviderTransportProvider {
|
||||||
|
id: provider_id.to_string(),
|
||||||
|
name: provider_id.to_string(),
|
||||||
|
provider_type: provider_type.to_string(),
|
||||||
|
website: None,
|
||||||
|
is_active: true,
|
||||||
|
keep_priority_on_conversion: false,
|
||||||
|
enable_format_conversion: false,
|
||||||
|
concurrent_limit: None,
|
||||||
|
max_retries: None,
|
||||||
|
proxy: None,
|
||||||
|
request_timeout_secs: None,
|
||||||
|
stream_first_byte_timeout_secs: None,
|
||||||
|
config: None,
|
||||||
|
},
|
||||||
|
endpoint: GatewayProviderTransportEndpoint {
|
||||||
|
id: endpoint_id.to_string(),
|
||||||
|
provider_id: provider_id.to_string(),
|
||||||
|
api_format: api_format.to_string(),
|
||||||
|
api_family: None,
|
||||||
|
endpoint_kind: None,
|
||||||
|
is_active: true,
|
||||||
|
base_url: "https://cloudcode-pa.googleapis.com".to_string(),
|
||||||
|
header_rules: None,
|
||||||
|
body_rules: None,
|
||||||
|
max_retries: None,
|
||||||
|
custom_path: None,
|
||||||
|
config: None,
|
||||||
|
format_acceptance_config: None,
|
||||||
|
proxy: None,
|
||||||
|
},
|
||||||
|
key: GatewayProviderTransportKey {
|
||||||
|
id: key_id.to_string(),
|
||||||
|
provider_id: provider_id.to_string(),
|
||||||
|
name: "primary".to_string(),
|
||||||
|
auth_type: auth_type.to_string(),
|
||||||
|
is_active: true,
|
||||||
|
api_formats: Some(vec![api_format.to_string()]),
|
||||||
|
allowed_models: None,
|
||||||
|
capabilities: None,
|
||||||
|
rate_multipliers: None,
|
||||||
|
global_priority_by_format: None,
|
||||||
|
expires_at_unix_secs: None,
|
||||||
|
proxy: None,
|
||||||
|
fingerprint: None,
|
||||||
|
decrypted_api_key: "secret".to_string(),
|
||||||
|
decrypted_auth_config: decrypted_auth_config.map(ToOwned::to_owned),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn execution_result(body: Value) -> ExecutionResult {
|
||||||
|
ExecutionResult {
|
||||||
|
request_id: "req-1".to_string(),
|
||||||
|
candidate_id: None,
|
||||||
|
status_code: 200,
|
||||||
|
headers: Default::default(),
|
||||||
|
body: Some(aether_contracts::ResponseBody {
|
||||||
|
json_body: Some(body),
|
||||||
|
body_bytes_b64: None,
|
||||||
|
}),
|
||||||
|
telemetry: None,
|
||||||
|
error: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn model_fetch_uses_preset_models_without_endpoint() {
|
||||||
|
let provider = sample_provider("provider-codex", "codex");
|
||||||
|
let key = sample_key("key-codex", "provider-codex", "api_key", &["openai:cli"]);
|
||||||
|
let state = TestState::new(vec![provider], vec![], vec![key], HashMap::new(), vec![]);
|
||||||
|
|
||||||
|
let summary = perform_model_fetch_once_with_state(&state)
|
||||||
|
.await
|
||||||
|
.expect("fetch should succeed");
|
||||||
|
|
||||||
|
assert_eq!(summary.attempted, 1);
|
||||||
|
assert_eq!(summary.succeeded, 1);
|
||||||
|
let updated = state.key("key-codex");
|
||||||
|
let allowed_models = updated
|
||||||
|
.allowed_models
|
||||||
|
.and_then(|value| value.as_array().cloned())
|
||||||
|
.expect("allowed_models should be set");
|
||||||
|
assert!(allowed_models.iter().any(|model| model == "gpt-5.4"));
|
||||||
|
assert!(state
|
||||||
|
.cached_models
|
||||||
|
.lock()
|
||||||
|
.expect("cache mutex")
|
||||||
|
.contains_key(&("provider-codex".to_string(), "key-codex".to_string())));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn model_fetch_merges_antigravity_metadata_and_preserves_reset_time() {
|
||||||
|
let provider = sample_provider("provider-antigravity", "antigravity");
|
||||||
|
let endpoint = sample_endpoint(
|
||||||
|
"endpoint-antigravity",
|
||||||
|
"provider-antigravity",
|
||||||
|
"gemini:chat",
|
||||||
|
);
|
||||||
|
let mut key = sample_key(
|
||||||
|
"key-antigravity",
|
||||||
|
"provider-antigravity",
|
||||||
|
"oauth",
|
||||||
|
&["gemini:chat"],
|
||||||
|
);
|
||||||
|
key.upstream_metadata = Some(json!({
|
||||||
|
"antigravity": {
|
||||||
|
"quota_by_model": {
|
||||||
|
"gemini-2.5-pro": {
|
||||||
|
"reset_time": "2026-04-12T00:00:00Z"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
let transport = sample_transport(
|
||||||
|
"antigravity",
|
||||||
|
"provider-antigravity",
|
||||||
|
"endpoint-antigravity",
|
||||||
|
"key-antigravity",
|
||||||
|
"gemini:chat",
|
||||||
|
"oauth",
|
||||||
|
Some(r#"{"project_id":"project-1","client_version":"1.2.3","session_id":"sess-1"}"#),
|
||||||
|
);
|
||||||
|
let state = TestState::new(
|
||||||
|
vec![provider],
|
||||||
|
vec![endpoint],
|
||||||
|
vec![key],
|
||||||
|
HashMap::from([(
|
||||||
|
(
|
||||||
|
"provider-antigravity".to_string(),
|
||||||
|
"endpoint-antigravity".to_string(),
|
||||||
|
"key-antigravity".to_string(),
|
||||||
|
),
|
||||||
|
transport,
|
||||||
|
)]),
|
||||||
|
vec![execution_result(json!({
|
||||||
|
"models": {
|
||||||
|
"gemini-2.5-pro": {
|
||||||
|
"displayName": "Gemini 2.5 Pro",
|
||||||
|
"quotaInfo": {
|
||||||
|
"remainingFraction": 0.25
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}))],
|
||||||
|
);
|
||||||
|
|
||||||
|
let summary = perform_model_fetch_once_with_state(&state)
|
||||||
|
.await
|
||||||
|
.expect("fetch should succeed");
|
||||||
|
|
||||||
|
assert_eq!(summary.succeeded, 1);
|
||||||
|
let updated = state.key("key-antigravity");
|
||||||
|
assert_eq!(updated.allowed_models, Some(json!(["gemini-2.5-pro"])));
|
||||||
|
assert_eq!(
|
||||||
|
updated
|
||||||
|
.upstream_metadata
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|value| value.get("antigravity"))
|
||||||
|
.and_then(|value| value.get("quota_by_model"))
|
||||||
|
.and_then(|value| value.get("gemini-2.5-pro"))
|
||||||
|
.and_then(|value| value.get("reset_time")),
|
||||||
|
Some(&json!("2026-04-12T00:00:00Z"))
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn model_fetch_failure_keeps_existing_allowed_models() {
|
||||||
|
let provider = sample_provider("provider-openai", "openai");
|
||||||
|
let endpoint = sample_endpoint(
|
||||||
|
"endpoint-openai-responses",
|
||||||
|
"provider-openai",
|
||||||
|
"openai:responses",
|
||||||
|
);
|
||||||
|
let mut key = sample_key(
|
||||||
|
"key-openai-responses",
|
||||||
|
"provider-openai",
|
||||||
|
"api_key",
|
||||||
|
&["openai:responses"],
|
||||||
|
);
|
||||||
|
key.allowed_models = Some(json!(["gpt-old"]));
|
||||||
|
let state = TestState::new(
|
||||||
|
vec![provider],
|
||||||
|
vec![endpoint],
|
||||||
|
vec![key],
|
||||||
|
HashMap::new(),
|
||||||
|
vec![],
|
||||||
|
);
|
||||||
|
|
||||||
|
let summary = perform_model_fetch_once_with_state(&state)
|
||||||
|
.await
|
||||||
|
.expect("fetch should finish");
|
||||||
|
|
||||||
|
assert_eq!(summary.succeeded, 0);
|
||||||
|
assert_eq!(summary.skipped, 1);
|
||||||
|
let updated = state.key("key-openai-responses");
|
||||||
|
assert_eq!(updated.allowed_models, Some(json!(["gpt-old"])));
|
||||||
|
assert_eq!(
|
||||||
|
updated.last_models_fetch_error.as_deref(),
|
||||||
|
Some("No supported endpoint for Rust models fetch")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -92,6 +92,19 @@ impl ModelFetchTransportRuntime for AppState {
|
|||||||
) -> Option<ProxySnapshot> {
|
) -> Option<ProxySnapshot> {
|
||||||
resolve_transport_proxy_snapshot_with_tunnel_affinity(self, transport).await
|
resolve_transport_proxy_snapshot_with_tunnel_affinity(self, transport).await
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn execute_model_fetch_execution_plan(
|
||||||
|
&self,
|
||||||
|
plan: &ExecutionPlan,
|
||||||
|
) -> Result<ExecutionResult, String> {
|
||||||
|
execution_runtime::execute_execution_runtime_sync_plan(self, None, plan)
|
||||||
|
.await
|
||||||
|
.map_err(|err| match err {
|
||||||
|
GatewayError::UpstreamUnavailable { message, .. }
|
||||||
|
| GatewayError::ControlUnavailable { message, .. }
|
||||||
|
| GatewayError::Internal(message) => message,
|
||||||
|
})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
|
|||||||
@@ -278,16 +278,16 @@ async fn gateway_handles_admin_provider_query_models_with_openai_responses_endpo
|
|||||||
|
|
||||||
assert_eq!(response.status(), StatusCode::OK);
|
assert_eq!(response.status(), StatusCode::OK);
|
||||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||||
assert_eq!(payload["success"], json!(true));
|
assert_eq!(payload["success"], json!(false));
|
||||||
assert_eq!(payload["data"]["error"], serde_json::Value::Null);
|
|
||||||
assert_eq!(payload["data"]["from_cache"], json!(false));
|
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
payload["data"]["models"][0]["api_formats"],
|
payload["data"]["error"],
|
||||||
json!(["openai:responses"])
|
json!("No active endpoints found for this provider")
|
||||||
);
|
);
|
||||||
|
assert_eq!(payload["data"]["from_cache"], json!(false));
|
||||||
|
assert_eq!(payload["data"]["models"], json!([]));
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
*execution_runtime_hits.lock().expect("mutex should lock"),
|
*execution_runtime_hits.lock().expect("mutex should lock"),
|
||||||
1
|
0
|
||||||
);
|
);
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
@@ -557,7 +557,7 @@ async fn gateway_handles_admin_provider_query_models_aggregating_active_keys() {
|
|||||||
.iter()
|
.iter()
|
||||||
.map(|model| model["id"].as_str().expect("id should exist"))
|
.map(|model| model["id"].as_str().expect("id should exist"))
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
assert_eq!(model_ids, vec!["gpt-5", "gpt-4.1"]);
|
assert_eq!(model_ids, vec!["gpt-4.1", "gpt-5"]);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
*execution_runtime_hits.lock().expect("mutex should lock"),
|
*execution_runtime_hits.lock().expect("mutex should lock"),
|
||||||
2
|
2
|
||||||
@@ -567,6 +567,81 @@ async fn gateway_handles_admin_provider_query_models_aggregating_active_keys() {
|
|||||||
execution_runtime_handle.abort();
|
execution_runtime_handle.abort();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn gateway_handles_admin_provider_query_models_for_fixed_provider_without_endpoint() {
|
||||||
|
let execution_runtime_hits = Arc::new(Mutex::new(0usize));
|
||||||
|
let execution_runtime_hits_clone = Arc::clone(&execution_runtime_hits);
|
||||||
|
let execution_runtime = Router::new().route(
|
||||||
|
"/v1/execute/sync",
|
||||||
|
any(move |_request: Request| {
|
||||||
|
let execution_runtime_hits_inner = Arc::clone(&execution_runtime_hits_clone);
|
||||||
|
async move {
|
||||||
|
*execution_runtime_hits_inner
|
||||||
|
.lock()
|
||||||
|
.expect("mutex should lock") += 1;
|
||||||
|
Json(json!({
|
||||||
|
"request_id": "unexpected",
|
||||||
|
"status_code": 500
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
|
||||||
|
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
|
||||||
|
let mut provider = sample_provider("provider-codex", "Codex", 10);
|
||||||
|
provider.provider_type = "codex".to_string();
|
||||||
|
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||||
|
vec![provider],
|
||||||
|
vec![],
|
||||||
|
vec![sample_key(
|
||||||
|
"key-codex-oauth",
|
||||||
|
"provider-codex",
|
||||||
|
"openai:cli",
|
||||||
|
"sk-test-codex",
|
||||||
|
)],
|
||||||
|
));
|
||||||
|
|
||||||
|
let gateway = build_router_with_state(
|
||||||
|
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||||
|
.with_data_state_for_tests(GatewayDataState::with_provider_transport_reader_for_tests(
|
||||||
|
provider_catalog_repository,
|
||||||
|
DEVELOPMENT_ENCRYPTION_KEY.to_string(),
|
||||||
|
)),
|
||||||
|
);
|
||||||
|
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||||
|
|
||||||
|
let response = reqwest::Client::new()
|
||||||
|
.post(format!("{gateway_url}/api/admin/provider-query/models"))
|
||||||
|
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||||
|
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||||
|
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||||
|
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||||
|
.json(&json!({
|
||||||
|
"provider_id": "provider-codex",
|
||||||
|
"api_key_id": "key-codex-oauth"
|
||||||
|
}))
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.expect("request should succeed");
|
||||||
|
|
||||||
|
assert_eq!(response.status(), StatusCode::OK);
|
||||||
|
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||||
|
assert_eq!(payload["success"], json!(true));
|
||||||
|
assert_eq!(payload["data"]["error"], serde_json::Value::Null);
|
||||||
|
assert_eq!(payload["data"]["from_cache"], json!(false));
|
||||||
|
let models = payload["data"]["models"]
|
||||||
|
.as_array()
|
||||||
|
.expect("models should be an array");
|
||||||
|
assert!(models.iter().any(|model| model["id"] == "gpt-5.4"));
|
||||||
|
assert_eq!(
|
||||||
|
*execution_runtime_hits.lock().expect("mutex should lock"),
|
||||||
|
0
|
||||||
|
);
|
||||||
|
|
||||||
|
gateway_handle.abort();
|
||||||
|
execution_runtime_handle.abort();
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn gateway_handles_admin_provider_query_test_model_locally_with_trusted_admin_principal() {
|
async fn gateway_handles_admin_provider_query_test_model_locally_with_trusted_admin_principal() {
|
||||||
assert_admin_provider_query_route(
|
assert_admin_provider_query_route(
|
||||||
|
|||||||
@@ -11,8 +11,11 @@ aether-data-contracts.workspace = true
|
|||||||
aether-provider-transport.workspace = true
|
aether-provider-transport.workspace = true
|
||||||
aether-scheduler-core.workspace = true
|
aether-scheduler-core.workspace = true
|
||||||
async-trait.workspace = true
|
async-trait.workspace = true
|
||||||
|
base64.workspace = true
|
||||||
regex.workspace = true
|
regex.workspace = true
|
||||||
|
rsa = "0.9.10"
|
||||||
serde_json.workspace = true
|
serde_json.workspace = true
|
||||||
|
sha2 = { workspace = true, features = ["oid"] }
|
||||||
uuid.workspace = true
|
uuid.workspace = true
|
||||||
|
|
||||||
[dev-dependencies]
|
[dev-dependencies]
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
mod association_sync;
|
mod association_sync;
|
||||||
mod config;
|
mod config;
|
||||||
mod logic;
|
mod logic;
|
||||||
|
mod strategy;
|
||||||
mod transport;
|
mod transport;
|
||||||
|
|
||||||
pub use association_sync::{
|
pub use association_sync::{
|
||||||
@@ -12,6 +13,13 @@ pub use config::{
|
|||||||
pub use logic::{
|
pub use logic::{
|
||||||
aggregate_models_for_cache, apply_model_filters, build_models_fetch_url,
|
aggregate_models_for_cache, apply_model_filters, build_models_fetch_url,
|
||||||
endpoint_supports_rust_models_fetch, extract_error_message, json_string_list,
|
endpoint_supports_rust_models_fetch, extract_error_message, json_string_list,
|
||||||
parse_models_response, select_models_fetch_endpoint, ModelFetchRunSummary, ModelsFetchSuccess,
|
merge_upstream_metadata, parse_models_response, parse_models_response_page,
|
||||||
|
preset_models_for_provider, provider_type_uses_preset_models, select_models_fetch_endpoint,
|
||||||
|
selected_models_fetch_endpoints, ModelFetchRunSummary, ModelsFetchPage, ModelsFetchSuccess,
|
||||||
|
};
|
||||||
|
pub use strategy::{fetch_models_from_transports, ModelsFetchOutcome};
|
||||||
|
pub use transport::{
|
||||||
|
build_antigravity_fetch_available_models_plan, build_gemini_cli_load_code_assist_plan,
|
||||||
|
build_models_fetch_execution_plan, build_standard_models_fetch_execution_plan,
|
||||||
|
build_vertex_models_fetch_execution_plan, ModelFetchTransportRuntime,
|
||||||
};
|
};
|
||||||
pub use transport::{build_models_fetch_execution_plan, ModelFetchTransportRuntime};
|
|
||||||
|
|||||||
@@ -3,9 +3,14 @@ use std::collections::{BTreeMap, BTreeSet};
|
|||||||
use aether_data_contracts::repository::provider_catalog::{
|
use aether_data_contracts::repository::provider_catalog::{
|
||||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||||
};
|
};
|
||||||
use aether_provider_transport::provider_types::provider_type_supports_model_fetch;
|
|
||||||
use regex::Regex;
|
use regex::Regex;
|
||||||
use serde_json::Value;
|
use serde_json::{json, Value};
|
||||||
|
|
||||||
|
const MODEL_FETCH_FORMAT_PRIORITY: &[&[&str]] = &[
|
||||||
|
&["openai:chat", "openai:cli", "openai:compact"],
|
||||||
|
&["claude:chat", "claude:cli"],
|
||||||
|
&["gemini:chat", "gemini:cli"],
|
||||||
|
];
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
pub struct ModelFetchRunSummary {
|
pub struct ModelFetchRunSummary {
|
||||||
@@ -21,6 +26,14 @@ pub struct ModelsFetchSuccess {
|
|||||||
pub cached_models: Vec<Value>,
|
pub cached_models: Vec<Value>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, PartialEq)]
|
||||||
|
pub struct ModelsFetchPage {
|
||||||
|
pub fetched_model_ids: Vec<String>,
|
||||||
|
pub cached_models: Vec<Value>,
|
||||||
|
pub has_more: bool,
|
||||||
|
pub next_after_id: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
pub fn extract_error_message(value: &Value) -> Option<String> {
|
pub fn extract_error_message(value: &Value) -> Option<String> {
|
||||||
value
|
value
|
||||||
.get("error")
|
.get("error")
|
||||||
@@ -41,21 +54,20 @@ pub fn extract_error_message(value: &Value) -> Option<String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn build_models_fetch_url(
|
pub fn build_models_fetch_url(
|
||||||
provider_type: &str,
|
_provider_type: &str,
|
||||||
endpoint_api_format: &str,
|
endpoint_api_format: &str,
|
||||||
base_url: &str,
|
base_url: &str,
|
||||||
) -> Option<(String, String)> {
|
) -> Option<(String, String)> {
|
||||||
let api_format = normalize_api_format(endpoint_api_format);
|
let api_format = normalize_api_format(endpoint_api_format);
|
||||||
if !provider_type_supports_model_fetch(provider_type) {
|
if !endpoint_supports_rust_models_fetch(&api_format) {
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
|
|
||||||
let url = if api_format.starts_with("openai:") || api_format.starts_with("claude:") {
|
let url = if api_format.starts_with("openai:") || api_format.starts_with("claude:") {
|
||||||
build_v1_models_url(base_url)
|
build_v1_models_url(base_url)
|
||||||
} else if api_format.starts_with("gemini:") {
|
} else if api_format.starts_with("gemini:") {
|
||||||
build_gemini_models_url(base_url)
|
build_gemini_models_url(base_url)
|
||||||
} else {
|
} else {
|
||||||
return None;
|
None
|
||||||
}?;
|
}?;
|
||||||
Some((url, api_format))
|
Some((url, api_format))
|
||||||
}
|
}
|
||||||
@@ -64,13 +76,38 @@ pub fn parse_models_response(
|
|||||||
endpoint_api_format: &str,
|
endpoint_api_format: &str,
|
||||||
body: &Value,
|
body: &Value,
|
||||||
) -> Result<ModelsFetchSuccess, String> {
|
) -> Result<ModelsFetchSuccess, String> {
|
||||||
|
let parsed = parse_models_response_page(endpoint_api_format, body)?;
|
||||||
|
Ok(ModelsFetchSuccess {
|
||||||
|
fetched_model_ids: parsed.fetched_model_ids,
|
||||||
|
cached_models: parsed.cached_models,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn parse_models_response_page(
|
||||||
|
endpoint_api_format: &str,
|
||||||
|
body: &Value,
|
||||||
|
) -> Result<ModelsFetchPage, String> {
|
||||||
let api_format = normalize_api_format(endpoint_api_format);
|
let api_format = normalize_api_format(endpoint_api_format);
|
||||||
let mut cached_models = Vec::new();
|
let mut cached_models = Vec::new();
|
||||||
let mut fetched_model_ids = Vec::new();
|
let mut fetched_model_ids = Vec::new();
|
||||||
let mut seen = BTreeSet::new();
|
let mut seen = BTreeSet::new();
|
||||||
|
let mut has_more = false;
|
||||||
|
let mut next_after_id = None;
|
||||||
|
|
||||||
if api_format.starts_with("openai:") || api_format.starts_with("claude:") {
|
if api_format.starts_with("openai:") || api_format.starts_with("claude:") {
|
||||||
let items = if let Some(items) = body.get("data").and_then(Value::as_array) {
|
let items = if let Some(items) = body.get("data").and_then(Value::as_array) {
|
||||||
|
has_more = body
|
||||||
|
.get("has_more")
|
||||||
|
.and_then(Value::as_bool)
|
||||||
|
.unwrap_or(false);
|
||||||
|
if api_format.starts_with("claude:") && has_more {
|
||||||
|
next_after_id = body
|
||||||
|
.get("last_id")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned);
|
||||||
|
}
|
||||||
items
|
items
|
||||||
} else if let Some(items) = body.as_array() {
|
} else if let Some(items) = body.as_array() {
|
||||||
items
|
items
|
||||||
@@ -117,29 +154,54 @@ pub fn parse_models_response(
|
|||||||
return Err("models response parser does not support this provider format".to_string());
|
return Err("models response parser does not support this provider format".to_string());
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(ModelsFetchSuccess {
|
Ok(ModelsFetchPage {
|
||||||
fetched_model_ids,
|
fetched_model_ids,
|
||||||
cached_models,
|
cached_models,
|
||||||
|
has_more,
|
||||||
|
next_after_id,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn selected_models_fetch_endpoints(
|
||||||
|
endpoints: &[StoredProviderCatalogEndpoint],
|
||||||
|
key: &StoredProviderCatalogKey,
|
||||||
|
) -> Vec<StoredProviderCatalogEndpoint> {
|
||||||
|
let key_formats = json_string_list(key.api_formats.as_ref())
|
||||||
|
.into_iter()
|
||||||
|
.map(|value| normalize_api_format(&value))
|
||||||
|
.collect::<BTreeSet<_>>();
|
||||||
|
let mut by_format = BTreeMap::<String, StoredProviderCatalogEndpoint>::new();
|
||||||
|
|
||||||
|
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
|
||||||
|
let api_format = normalize_api_format(&endpoint.api_format);
|
||||||
|
if api_format.is_empty() || !endpoint_supports_rust_models_fetch(&api_format) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
if !key_formats.is_empty() && !key_formats.contains(&api_format) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
by_format
|
||||||
|
.entry(api_format)
|
||||||
|
.or_insert_with(|| endpoint.clone());
|
||||||
|
}
|
||||||
|
|
||||||
|
MODEL_FETCH_FORMAT_PRIORITY
|
||||||
|
.iter()
|
||||||
|
.filter_map(|candidates| {
|
||||||
|
candidates
|
||||||
|
.iter()
|
||||||
|
.find_map(|api_format| by_format.remove(*api_format))
|
||||||
|
})
|
||||||
|
.collect()
|
||||||
|
}
|
||||||
|
|
||||||
pub fn select_models_fetch_endpoint(
|
pub fn select_models_fetch_endpoint(
|
||||||
endpoints: &[StoredProviderCatalogEndpoint],
|
endpoints: &[StoredProviderCatalogEndpoint],
|
||||||
key: &StoredProviderCatalogKey,
|
key: &StoredProviderCatalogKey,
|
||||||
) -> Option<StoredProviderCatalogEndpoint> {
|
) -> Option<StoredProviderCatalogEndpoint> {
|
||||||
let key_formats = json_string_list(key.api_formats.as_ref())
|
selected_models_fetch_endpoints(endpoints, key)
|
||||||
.into_iter()
|
.into_iter()
|
||||||
.map(|value| normalize_api_format(&value))
|
.next()
|
||||||
.collect::<BTreeSet<_>>();
|
|
||||||
endpoints
|
|
||||||
.iter()
|
|
||||||
.filter(|endpoint| endpoint.is_active)
|
|
||||||
.find(|endpoint| {
|
|
||||||
let api_format = normalize_api_format(&endpoint.api_format);
|
|
||||||
(key_formats.is_empty() || key_formats.contains(&api_format))
|
|
||||||
&& endpoint_supports_rust_models_fetch(&endpoint.api_format)
|
|
||||||
})
|
|
||||||
.cloned()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool {
|
pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool {
|
||||||
@@ -148,7 +210,6 @@ pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool {
|
|||||||
api_format.as_str(),
|
api_format.as_str(),
|
||||||
"openai:chat"
|
"openai:chat"
|
||||||
| "openai:cli"
|
| "openai:cli"
|
||||||
| "openai:responses"
|
|
||||||
| "openai:compact"
|
| "openai:compact"
|
||||||
| "claude:chat"
|
| "claude:chat"
|
||||||
| "claude:cli"
|
| "claude:cli"
|
||||||
@@ -157,6 +218,100 @@ pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn provider_type_uses_preset_models(provider_type: &str) -> bool {
|
||||||
|
matches!(
|
||||||
|
provider_type.trim().to_ascii_lowercase().as_str(),
|
||||||
|
"codex" | "kiro" | "claude_code" | "gemini_cli"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[rustfmt::skip]
|
||||||
|
pub fn preset_models_for_provider(provider_type: &str) -> Option<Vec<Value>> {
|
||||||
|
let models = match provider_type.trim().to_ascii_lowercase().as_str() {
|
||||||
|
"gemini_cli" => vec![
|
||||||
|
preset_model("gemini-2.5-pro", "google", "Gemini 2.5 Pro", "gemini:cli"),
|
||||||
|
preset_model("gemini-2.5-flash", "google", "Gemini 2.5 Flash", "gemini:cli"),
|
||||||
|
preset_model("gemini-3-pro-preview", "google", "Gemini 3 Pro Preview", "gemini:cli"),
|
||||||
|
preset_model("gemini-3-flash-preview", "google", "Gemini 3 Flash Preview", "gemini:cli"),
|
||||||
|
preset_model("gemini-3.1-pro-preview", "google", "Gemini 3.1 Pro Preview", "gemini:cli"),
|
||||||
|
],
|
||||||
|
"kiro" => vec![
|
||||||
|
preset_model("claude-sonnet-4.5", "anthropic", "Claude Sonnet 4.5", "claude:cli"),
|
||||||
|
preset_model("claude-sonnet-4.6", "anthropic", "Claude Sonnet 4.6", "claude:cli"),
|
||||||
|
preset_model("claude-opus-4.5", "anthropic", "Claude Opus 4.5", "claude:cli"),
|
||||||
|
preset_model("claude-opus-4.6", "anthropic", "Claude Opus 4.6", "claude:cli"),
|
||||||
|
preset_model("claude-haiku-4.5", "anthropic", "Claude Haiku 4.5", "claude:cli"),
|
||||||
|
],
|
||||||
|
"claude_code" => vec![
|
||||||
|
preset_model("claude-opus-4-5-20251101", "anthropic", "Claude Opus 4.5", "claude:cli"),
|
||||||
|
preset_model("claude-opus-4-6", "anthropic", "Claude Opus 4.6", "claude:cli"),
|
||||||
|
preset_model("claude-sonnet-4-6", "anthropic", "Claude Sonnet 4.6", "claude:cli"),
|
||||||
|
preset_model("claude-sonnet-4-5-20250929", "anthropic", "Claude Sonnet 4.5", "claude:cli"),
|
||||||
|
preset_model("claude-haiku-4-5-20251001", "anthropic", "Claude Haiku 4.5", "claude:cli"),
|
||||||
|
],
|
||||||
|
"codex" => vec![
|
||||||
|
preset_model("gpt-5", "openai", "GPT-5", "openai:cli"),
|
||||||
|
preset_model("gpt-5-codex", "openai", "GPT-5 Codex", "openai:cli"),
|
||||||
|
preset_model("gpt-5-codex-mini", "openai", "GPT-5 Codex Mini", "openai:cli"),
|
||||||
|
preset_model("gpt-5.1", "openai", "GPT-5.1", "openai:cli"),
|
||||||
|
preset_model("gpt-5.1-codex", "openai", "GPT-5.1 Codex", "openai:cli"),
|
||||||
|
preset_model("gpt-5.1-codex-mini", "openai", "GPT-5.1 Codex Mini", "openai:cli"),
|
||||||
|
preset_model("gpt-5.1-codex-max", "openai", "GPT-5.1 Codex Max", "openai:cli"),
|
||||||
|
preset_model("gpt-5.2", "openai", "GPT-5.2", "openai:cli"),
|
||||||
|
preset_model("gpt-5.2-codex", "openai", "GPT-5.2 Codex", "openai:cli"),
|
||||||
|
preset_model("gpt-5.3-codex", "openai", "GPT-5.3 Codex", "openai:cli"),
|
||||||
|
preset_model("gpt-5.4", "openai", "GPT-5.4", "openai:cli"),
|
||||||
|
],
|
||||||
|
_ => return None,
|
||||||
|
};
|
||||||
|
Some(models)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn merge_upstream_metadata(current: Option<&Value>, incoming: &Value) -> Value {
|
||||||
|
let mut merged = current
|
||||||
|
.and_then(Value::as_object)
|
||||||
|
.cloned()
|
||||||
|
.unwrap_or_default();
|
||||||
|
let Some(incoming_object) = incoming.as_object() else {
|
||||||
|
return Value::Object(merged);
|
||||||
|
};
|
||||||
|
|
||||||
|
for (namespace, value) in incoming_object {
|
||||||
|
let mut next_value = value.clone();
|
||||||
|
if let (Some(next_namespace), Some(old_namespace)) = (
|
||||||
|
next_value.as_object_mut(),
|
||||||
|
merged.get(namespace).and_then(Value::as_object),
|
||||||
|
) {
|
||||||
|
if let (Some(new_quota), Some(old_quota)) = (
|
||||||
|
next_namespace
|
||||||
|
.get_mut("quota_by_model")
|
||||||
|
.and_then(Value::as_object_mut),
|
||||||
|
old_namespace
|
||||||
|
.get("quota_by_model")
|
||||||
|
.and_then(Value::as_object),
|
||||||
|
) {
|
||||||
|
for (model_id, new_info) in new_quota.iter_mut() {
|
||||||
|
let Some(new_info_object) = new_info.as_object_mut() else {
|
||||||
|
continue;
|
||||||
|
};
|
||||||
|
let Some(old_info_object) = old_quota.get(model_id).and_then(Value::as_object)
|
||||||
|
else {
|
||||||
|
continue;
|
||||||
|
};
|
||||||
|
if !new_info_object.contains_key("reset_time") {
|
||||||
|
if let Some(reset_time) = old_info_object.get("reset_time") {
|
||||||
|
new_info_object.insert("reset_time".to_string(), reset_time.clone());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
merged.insert(namespace.clone(), next_value);
|
||||||
|
}
|
||||||
|
|
||||||
|
Value::Object(merged)
|
||||||
|
}
|
||||||
|
|
||||||
pub fn apply_model_filters(
|
pub fn apply_model_filters(
|
||||||
fetched_model_ids: &[String],
|
fetched_model_ids: &[String],
|
||||||
locked_models: Vec<String>,
|
locked_models: Vec<String>,
|
||||||
@@ -211,7 +366,6 @@ pub fn json_string_list(value: Option<&Value>) -> Vec<String> {
|
|||||||
|
|
||||||
pub fn aggregate_models_for_cache(models: &[Value]) -> Vec<Value> {
|
pub fn aggregate_models_for_cache(models: &[Value]) -> Vec<Value> {
|
||||||
let mut aggregated = BTreeMap::<String, serde_json::Map<String, Value>>::new();
|
let mut aggregated = BTreeMap::<String, serde_json::Map<String, Value>>::new();
|
||||||
let mut order = Vec::<String>::new();
|
|
||||||
|
|
||||||
for model in models {
|
for model in models {
|
||||||
let Some(object) = model.as_object() else {
|
let Some(object) = model.as_object() else {
|
||||||
@@ -227,7 +381,6 @@ pub fn aggregate_models_for_cache(models: &[Value]) -> Vec<Value> {
|
|||||||
};
|
};
|
||||||
|
|
||||||
let entry = aggregated.entry(model_id.to_string()).or_insert_with(|| {
|
let entry = aggregated.entry(model_id.to_string()).or_insert_with(|| {
|
||||||
order.push(model_id.to_string());
|
|
||||||
let mut cloned = object.clone();
|
let mut cloned = object.clone();
|
||||||
cloned.remove("api_format");
|
cloned.remove("api_format");
|
||||||
cloned
|
cloned
|
||||||
@@ -286,11 +439,7 @@ pub fn aggregate_models_for_cache(models: &[Value]) -> Vec<Value> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
order
|
aggregated.into_values().map(Value::Object).collect()
|
||||||
.into_iter()
|
|
||||||
.filter_map(|model_id| aggregated.remove(&model_id))
|
|
||||||
.map(Value::Object)
|
|
||||||
.collect()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn build_v1_models_url(base_url: &str) -> Option<String> {
|
fn build_v1_models_url(base_url: &str) -> Option<String> {
|
||||||
@@ -347,10 +496,37 @@ fn normalize_cached_model(item: &Value, model_id: &str, api_format: &str) -> Val
|
|||||||
"api_formats".to_string(),
|
"api_formats".to_string(),
|
||||||
Value::Array(vec![Value::String(api_format.to_string())]),
|
Value::Array(vec![Value::String(api_format.to_string())]),
|
||||||
);
|
);
|
||||||
|
if api_format.starts_with("gemini:") {
|
||||||
|
object
|
||||||
|
.entry("owned_by".to_string())
|
||||||
|
.or_insert_with(|| Value::String("google".to_string()));
|
||||||
|
if !object.contains_key("display_name") {
|
||||||
|
let display_name = item
|
||||||
|
.get("displayName")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.unwrap_or(model_id);
|
||||||
|
object.insert(
|
||||||
|
"display_name".to_string(),
|
||||||
|
Value::String(display_name.to_string()),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
object.remove("api_format");
|
object.remove("api_format");
|
||||||
Value::Object(object)
|
Value::Object(object)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn preset_model(model_id: &str, owned_by: &str, display_name: &str, api_format: &str) -> Value {
|
||||||
|
json!({
|
||||||
|
"id": model_id,
|
||||||
|
"object": "model",
|
||||||
|
"owned_by": owned_by,
|
||||||
|
"display_name": display_name,
|
||||||
|
"api_formats": [api_format],
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
fn wildcard_matches(pattern: &str, model_id: &str) -> bool {
|
fn wildcard_matches(pattern: &str, model_id: &str) -> bool {
|
||||||
let mut regex = String::from("^");
|
let mut regex = String::from("^");
|
||||||
for ch in pattern.chars() {
|
for ch in pattern.chars() {
|
||||||
@@ -379,7 +555,8 @@ mod tests {
|
|||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
aggregate_models_for_cache, apply_model_filters, build_gemini_models_url,
|
aggregate_models_for_cache, apply_model_filters, build_gemini_models_url,
|
||||||
build_models_fetch_url, parse_models_response, select_models_fetch_endpoint,
|
build_models_fetch_url, merge_upstream_metadata, parse_models_response,
|
||||||
|
parse_models_response_page, preset_models_for_provider, selected_models_fetch_endpoints,
|
||||||
};
|
};
|
||||||
|
|
||||||
fn sample_endpoint(
|
fn sample_endpoint(
|
||||||
@@ -410,7 +587,11 @@ mod tests {
|
|||||||
.expect("endpoint transport should build")
|
.expect("endpoint transport should build")
|
||||||
}
|
}
|
||||||
|
|
||||||
fn sample_key(provider_id: &str, key_id: &str) -> StoredProviderCatalogKey {
|
fn sample_key(
|
||||||
|
provider_id: &str,
|
||||||
|
key_id: &str,
|
||||||
|
api_formats: &[&str],
|
||||||
|
) -> StoredProviderCatalogKey {
|
||||||
StoredProviderCatalogKey::new(
|
StoredProviderCatalogKey::new(
|
||||||
key_id.to_string(),
|
key_id.to_string(),
|
||||||
provider_id.to_string(),
|
provider_id.to_string(),
|
||||||
@@ -421,7 +602,7 @@ mod tests {
|
|||||||
)
|
)
|
||||||
.expect("key should build")
|
.expect("key should build")
|
||||||
.with_transport_fields(
|
.with_transport_fields(
|
||||||
Some(json!(["openai:chat"])),
|
Some(json!(api_formats)),
|
||||||
"encrypted".to_string(),
|
"encrypted".to_string(),
|
||||||
None,
|
None,
|
||||||
None,
|
None,
|
||||||
@@ -453,12 +634,15 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn aggregate_models_for_cache_merges_api_formats_by_model_id() {
|
fn aggregate_models_for_cache_merges_api_formats_and_sorts_by_model_id() {
|
||||||
let aggregated = aggregate_models_for_cache(&[
|
let aggregated = aggregate_models_for_cache(&[
|
||||||
json!({"id":"gpt-5","api_formats":["openai:chat"]}),
|
json!({"id":"zeta","api_formats":["openai:chat"]}),
|
||||||
json!({"id":"gpt-5","api_formats":["openai:cli"]}),
|
json!({"id":"alpha","api_formats":["openai:cli"]}),
|
||||||
|
json!({"id":"alpha","api_formats":["openai:chat"]}),
|
||||||
]);
|
]);
|
||||||
assert_eq!(aggregated.len(), 1);
|
assert_eq!(aggregated.len(), 2);
|
||||||
|
assert_eq!(aggregated[0]["id"], "alpha");
|
||||||
|
assert_eq!(aggregated[1]["id"], "zeta");
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
aggregated[0]["api_formats"],
|
aggregated[0]["api_formats"],
|
||||||
json!(["openai:chat", "openai:cli"])
|
json!(["openai:chat", "openai:cli"])
|
||||||
@@ -488,9 +672,9 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn build_models_fetch_url_rejects_provider_types_without_fetch_support() {
|
fn build_models_fetch_url_excludes_openai_responses() {
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
build_models_fetch_url("vertex_ai", "gemini:chat", "https://example.com"),
|
build_models_fetch_url("openai", "openai:responses", "https://example.com"),
|
||||||
None
|
None
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
@@ -510,9 +694,30 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn select_models_fetch_endpoint_respects_key_api_formats() {
|
fn parse_models_response_page_reads_claude_pagination_state() {
|
||||||
let key = sample_key("provider-1", "key-1");
|
let parsed = parse_models_response_page(
|
||||||
|
"claude:chat",
|
||||||
|
&json!({
|
||||||
|
"data": [{"id": "claude-sonnet-4"}],
|
||||||
|
"has_more": true,
|
||||||
|
"last_id": "cursor-2"
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
.expect("response should parse");
|
||||||
|
assert!(parsed.has_more);
|
||||||
|
assert_eq!(parsed.next_after_id.as_deref(), Some("cursor-2"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn selected_models_fetch_endpoints_prefers_chat_and_excludes_responses() {
|
||||||
|
let key = sample_key("provider-1", "key-1", &["openai:chat", "openai:responses"]);
|
||||||
let endpoints = vec![
|
let endpoints = vec![
|
||||||
|
sample_endpoint(
|
||||||
|
"provider-1",
|
||||||
|
"endpoint-responses",
|
||||||
|
"openai:responses",
|
||||||
|
"https://example.com",
|
||||||
|
),
|
||||||
sample_endpoint(
|
sample_endpoint(
|
||||||
"provider-1",
|
"provider-1",
|
||||||
"endpoint-cli",
|
"endpoint-cli",
|
||||||
@@ -526,8 +731,50 @@ mod tests {
|
|||||||
"https://example.com",
|
"https://example.com",
|
||||||
),
|
),
|
||||||
];
|
];
|
||||||
let selected =
|
let selected = selected_models_fetch_endpoints(&endpoints, &key);
|
||||||
select_models_fetch_endpoint(&endpoints, &key).expect("endpoint should be selected");
|
assert_eq!(selected.len(), 1);
|
||||||
assert_eq!(selected.id, "endpoint-chat");
|
assert_eq!(selected[0].id, "endpoint-chat");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn merge_upstream_metadata_keeps_existing_reset_time_for_returned_models() {
|
||||||
|
let merged = merge_upstream_metadata(
|
||||||
|
Some(&json!({
|
||||||
|
"antigravity": {
|
||||||
|
"quota_by_model": {
|
||||||
|
"gemini-2.5-pro": {
|
||||||
|
"remaining_fraction": 0.3,
|
||||||
|
"reset_time": "2026-04-12T00:00:00Z"
|
||||||
|
},
|
||||||
|
"stale-model": {
|
||||||
|
"remaining_fraction": 0.1,
|
||||||
|
"reset_time": "old"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})),
|
||||||
|
&json!({
|
||||||
|
"antigravity": {
|
||||||
|
"quota_by_model": {
|
||||||
|
"gemini-2.5-pro": {
|
||||||
|
"remaining_fraction": 0.6
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
merged["antigravity"]["quota_by_model"]["gemini-2.5-pro"]["reset_time"],
|
||||||
|
"2026-04-12T00:00:00Z"
|
||||||
|
);
|
||||||
|
assert!(merged["antigravity"]["quota_by_model"]
|
||||||
|
.get("stale-model")
|
||||||
|
.is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn preset_models_cover_codex_catalog() {
|
||||||
|
let models = preset_models_for_provider("codex").expect("preset models should exist");
|
||||||
|
assert!(models.iter().any(|model| model["id"] == "gpt-5.4"));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
1063
crates/aether-model-fetch/src/strategy.rs
Normal file
1063
crates/aether-model-fetch/src/strategy.rs
Normal file
File diff suppressed because it is too large
Load Diff
@@ -1,20 +1,47 @@
|
|||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
|
|
||||||
use aether_contracts::{ExecutionPlan, ProxySnapshot, RequestBody};
|
use aether_contracts::{ExecutionPlan, ExecutionResult, ProxySnapshot, RequestBody};
|
||||||
use aether_provider_transport::auth::{
|
use aether_provider_transport::antigravity::{
|
||||||
resolve_local_gemini_auth, resolve_local_openai_chat_auth, resolve_local_standard_auth,
|
build_antigravity_static_identity_headers, resolve_local_antigravity_request_auth,
|
||||||
|
AntigravityRequestAuthSupport, ANTIGRAVITY_REQUEST_USER_AGENT,
|
||||||
|
};
|
||||||
|
use aether_provider_transport::auth::{
|
||||||
|
ensure_upstream_auth_header, resolve_local_gemini_auth, resolve_local_openai_chat_auth,
|
||||||
|
resolve_local_standard_auth,
|
||||||
};
|
};
|
||||||
use aether_provider_transport::url::build_passthrough_path_url;
|
|
||||||
use aether_provider_transport::vertex::resolve_local_vertex_api_key_query_auth;
|
use aether_provider_transport::vertex::resolve_local_vertex_api_key_query_auth;
|
||||||
use aether_provider_transport::{
|
use aether_provider_transport::{
|
||||||
apply_local_header_rules, ensure_upstream_auth_header, resolve_transport_execution_timeouts,
|
apply_local_header_rules, resolve_transport_execution_timeouts, resolve_transport_tls_profile,
|
||||||
resolve_transport_tls_profile, GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
|
GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
|
||||||
};
|
};
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
|
|
||||||
use crate::build_models_fetch_url;
|
use crate::build_models_fetch_url;
|
||||||
|
|
||||||
|
const OPENAI_CLI_USER_AGENT: &str = "openai-codex/1.0";
|
||||||
|
const CLAUDE_CLI_USER_AGENT: &str = "claude-code/1.0.1";
|
||||||
|
const GEMINI_CLI_USER_AGENT: &str = "GeminiCLI/0.1.5 (Windows; AMD64)";
|
||||||
|
const CLAUDE_VERSION_HEADER: &str = "2023-06-01";
|
||||||
|
const ANTIGRAVITY_FETCH_PROVIDER_API_FORMAT: &str = "antigravity:fetch_available_models";
|
||||||
|
const GEMINI_CLI_LOAD_CODE_ASSIST_PROVIDER_API_FORMAT: &str = "gemini_cli:load_code_assist";
|
||||||
|
|
||||||
|
const BROWSER_FINGERPRINT_HEADERS: &[(&str, &str)] = &[
|
||||||
|
(
|
||||||
|
"user-agent",
|
||||||
|
"Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/140.0.7339.249 Electron/38.7.0 Safari/537.36",
|
||||||
|
),
|
||||||
|
("accept", "application/json"),
|
||||||
|
("accept-encoding", "gzip, deflate, br"),
|
||||||
|
("accept-language", "zh-CN"),
|
||||||
|
("sec-ch-ua", "\"Not=A?Brand\";v=\"24\", \"Chromium\";v=\"140\""),
|
||||||
|
("sec-ch-ua-mobile", "?0"),
|
||||||
|
("sec-ch-ua-platform", "\"macOS\""),
|
||||||
|
("sec-fetch-site", "cross-site"),
|
||||||
|
("sec-fetch-mode", "cors"),
|
||||||
|
("sec-fetch-dest", "empty"),
|
||||||
|
];
|
||||||
|
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
pub trait ModelFetchTransportRuntime: Send + Sync {
|
pub trait ModelFetchTransportRuntime: Send + Sync {
|
||||||
async fn resolve_local_oauth_request_auth(
|
async fn resolve_local_oauth_request_auth(
|
||||||
@@ -26,113 +53,415 @@ pub trait ModelFetchTransportRuntime: Send + Sync {
|
|||||||
&self,
|
&self,
|
||||||
transport: &GatewayProviderTransportSnapshot,
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
) -> Option<ProxySnapshot>;
|
) -> Option<ProxySnapshot>;
|
||||||
|
|
||||||
|
async fn execute_model_fetch_execution_plan(
|
||||||
|
&self,
|
||||||
|
plan: &ExecutionPlan,
|
||||||
|
) -> Result<ExecutionResult, String>;
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn build_models_fetch_execution_plan(
|
pub async fn build_models_fetch_execution_plan(
|
||||||
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||||
transport: &GatewayProviderTransportSnapshot,
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
) -> Result<ExecutionPlan, String> {
|
) -> Result<ExecutionPlan, String> {
|
||||||
let (upstream_url, provider_api_format) = build_models_fetch_url(
|
build_standard_models_fetch_execution_plan(runtime, transport, None).await
|
||||||
&transport.provider.provider_type,
|
}
|
||||||
&transport.endpoint.api_format,
|
|
||||||
&transport.endpoint.base_url,
|
|
||||||
)
|
|
||||||
.ok_or_else(|| "Rust models fetch does not support this provider format yet".to_string())?;
|
|
||||||
let (auth_header_name, auth_header_value) = resolve_models_fetch_auth(runtime, transport)
|
|
||||||
.await?
|
|
||||||
.ok_or_else(|| {
|
|
||||||
"Rust models fetch auth resolution is not supported for this key".to_string()
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let mut headers = BTreeMap::from([(auth_header_name.clone(), auth_header_value.clone())]);
|
pub async fn build_standard_models_fetch_execution_plan(
|
||||||
if !apply_local_header_rules(
|
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||||
&mut headers,
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
transport.endpoint.header_rules.as_ref(),
|
after_id: Option<&str>,
|
||||||
&[auth_header_name.as_str()],
|
) -> Result<ExecutionPlan, String> {
|
||||||
&json!({}),
|
let api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
|
||||||
None,
|
let provider_api_format = api_format.clone();
|
||||||
) {
|
let mut headers = standard_models_fetch_headers(&api_format);
|
||||||
return Err("Endpoint header_rules application failed".to_string());
|
let mut protected_headers = Vec::<String>::new();
|
||||||
|
|
||||||
|
if api_format.starts_with("openai:") || api_format.starts_with("claude:") {
|
||||||
|
let (auth_header_name, auth_header_value) =
|
||||||
|
resolve_standard_header_auth(runtime, transport)
|
||||||
|
.await?
|
||||||
|
.ok_or_else(|| {
|
||||||
|
"Rust models fetch auth resolution is not supported for this key".to_string()
|
||||||
|
})?;
|
||||||
|
protected_headers.push(auth_header_name.clone());
|
||||||
|
headers.insert(auth_header_name.clone(), auth_header_value.clone());
|
||||||
|
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
|
||||||
|
ensure_upstream_auth_header(&mut headers, &auth_header_name, &auth_header_value);
|
||||||
|
} else {
|
||||||
|
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
|
||||||
}
|
}
|
||||||
ensure_upstream_auth_header(&mut headers, &auth_header_name, &auth_header_value);
|
|
||||||
|
|
||||||
|
let upstream_url = build_standard_models_fetch_url(transport, after_id)?;
|
||||||
|
build_execution_plan(
|
||||||
|
runtime,
|
||||||
|
transport,
|
||||||
|
"GET",
|
||||||
|
upstream_url,
|
||||||
|
headers,
|
||||||
|
None,
|
||||||
|
RequestBody {
|
||||||
|
json_body: None,
|
||||||
|
body_bytes_b64: None,
|
||||||
|
body_ref: None,
|
||||||
|
},
|
||||||
|
provider_api_format.clone(),
|
||||||
|
provider_api_format,
|
||||||
|
Some("models".to_string()),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn build_antigravity_fetch_available_models_plan(
|
||||||
|
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
base_url: &str,
|
||||||
|
project_id: &str,
|
||||||
|
) -> Result<ExecutionPlan, String> {
|
||||||
|
let authorization = resolve_oauth_header_auth(runtime, transport)
|
||||||
|
.await?
|
||||||
|
.ok_or_else(|| "Antigravity fetch requires OAuth authorization header".to_string())?;
|
||||||
|
let identity_auth = match resolve_local_antigravity_request_auth(transport) {
|
||||||
|
AntigravityRequestAuthSupport::Supported(auth) => auth,
|
||||||
|
AntigravityRequestAuthSupport::Unsupported(reason) => {
|
||||||
|
return Err(format!(
|
||||||
|
"Antigravity fetch auth resolution is not supported: {reason:?}"
|
||||||
|
))
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
let mut headers = build_antigravity_static_identity_headers(&identity_auth);
|
||||||
|
headers.insert(authorization.0.clone(), authorization.1.clone());
|
||||||
|
headers.insert("content-type".to_string(), "application/json".to_string());
|
||||||
|
headers.insert("accept".to_string(), "application/json".to_string());
|
||||||
|
headers
|
||||||
|
.entry("user-agent".to_string())
|
||||||
|
.or_insert_with(|| ANTIGRAVITY_REQUEST_USER_AGENT.to_string());
|
||||||
|
let protected_headers = vec![authorization.0];
|
||||||
|
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
|
||||||
|
|
||||||
|
let url = format!(
|
||||||
|
"{}{}",
|
||||||
|
base_url.trim_end_matches('/'),
|
||||||
|
"/v1internal:fetchAvailableModels"
|
||||||
|
);
|
||||||
|
build_execution_plan(
|
||||||
|
runtime,
|
||||||
|
transport,
|
||||||
|
"POST",
|
||||||
|
url,
|
||||||
|
headers,
|
||||||
|
Some("application/json".to_string()),
|
||||||
|
RequestBody::from_json(json!({ "project": project_id })),
|
||||||
|
"gemini:chat".to_string(),
|
||||||
|
ANTIGRAVITY_FETCH_PROVIDER_API_FORMAT.to_string(),
|
||||||
|
Some("fetchAvailableModels".to_string()),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn build_gemini_cli_load_code_assist_plan(
|
||||||
|
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
) -> Result<ExecutionPlan, String> {
|
||||||
|
let authorization = resolve_bearer_or_oauth_header_auth(runtime, transport)
|
||||||
|
.await?
|
||||||
|
.ok_or_else(|| "GeminiCLI loadCodeAssist requires bearer or OAuth auth".to_string())?;
|
||||||
|
|
||||||
|
let mut headers = BTreeMap::from([
|
||||||
|
("user-agent".to_string(), GEMINI_CLI_USER_AGENT.to_string()),
|
||||||
|
("accept-encoding".to_string(), "identity".to_string()),
|
||||||
|
("content-type".to_string(), "application/json".to_string()),
|
||||||
|
]);
|
||||||
|
headers.insert(authorization.0.clone(), authorization.1.clone());
|
||||||
|
let protected_headers = vec![authorization.0];
|
||||||
|
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
|
||||||
|
|
||||||
|
build_execution_plan(
|
||||||
|
runtime,
|
||||||
|
transport,
|
||||||
|
"POST",
|
||||||
|
"https://cloudcode-pa.googleapis.com/v1internal:loadCodeAssist".to_string(),
|
||||||
|
headers,
|
||||||
|
Some("application/json".to_string()),
|
||||||
|
RequestBody::from_json(json!({
|
||||||
|
"metadata": {
|
||||||
|
"ideType": "ANTIGRAVITY",
|
||||||
|
"platform": "PLATFORM_UNSPECIFIED",
|
||||||
|
"pluginType": "GEMINI",
|
||||||
|
}
|
||||||
|
})),
|
||||||
|
"gemini:cli".to_string(),
|
||||||
|
GEMINI_CLI_LOAD_CODE_ASSIST_PROVIDER_API_FORMAT.to_string(),
|
||||||
|
Some("loadCodeAssist".to_string()),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn build_vertex_models_fetch_execution_plan(
|
||||||
|
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
url: &str,
|
||||||
|
api_format: &str,
|
||||||
|
auth_header: Option<(String, String)>,
|
||||||
|
) -> Result<ExecutionPlan, String> {
|
||||||
|
let mut headers = standard_models_fetch_headers(api_format);
|
||||||
|
let mut protected_headers = Vec::<String>::new();
|
||||||
|
if let Some((name, value)) = auth_header {
|
||||||
|
protected_headers.push(name.clone());
|
||||||
|
headers.insert(name.clone(), value.clone());
|
||||||
|
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
|
||||||
|
ensure_upstream_auth_header(&mut headers, &name, &value);
|
||||||
|
} else {
|
||||||
|
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
|
||||||
|
}
|
||||||
|
|
||||||
|
build_execution_plan(
|
||||||
|
runtime,
|
||||||
|
transport,
|
||||||
|
"GET",
|
||||||
|
url.trim().to_string(),
|
||||||
|
headers,
|
||||||
|
None,
|
||||||
|
RequestBody {
|
||||||
|
json_body: None,
|
||||||
|
body_bytes_b64: None,
|
||||||
|
body_ref: None,
|
||||||
|
},
|
||||||
|
api_format.to_string(),
|
||||||
|
api_format.to_string(),
|
||||||
|
Some("models".to_string()),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn build_execution_plan(
|
||||||
|
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
method: &str,
|
||||||
|
url: String,
|
||||||
|
headers: BTreeMap<String, String>,
|
||||||
|
content_type: Option<String>,
|
||||||
|
body: RequestBody,
|
||||||
|
client_api_format: String,
|
||||||
|
provider_api_format: String,
|
||||||
|
model_name: Option<String>,
|
||||||
|
) -> Result<ExecutionPlan, String> {
|
||||||
Ok(ExecutionPlan {
|
Ok(ExecutionPlan {
|
||||||
request_id: format!("req-model-fetch-{}", transport.key.id),
|
request_id: format!(
|
||||||
|
"req-model-fetch-{}-{}",
|
||||||
|
transport.key.id,
|
||||||
|
provider_api_format.replace(':', "-")
|
||||||
|
),
|
||||||
candidate_id: None,
|
candidate_id: None,
|
||||||
provider_name: Some(transport.provider.name.clone()),
|
provider_name: Some(transport.provider.name.clone()),
|
||||||
provider_id: transport.provider.id.clone(),
|
provider_id: transport.provider.id.clone(),
|
||||||
endpoint_id: transport.endpoint.id.clone(),
|
endpoint_id: transport.endpoint.id.clone(),
|
||||||
key_id: transport.key.id.clone(),
|
key_id: transport.key.id.clone(),
|
||||||
method: "GET".to_string(),
|
method: method.to_string(),
|
||||||
url: upstream_url,
|
url,
|
||||||
headers,
|
headers,
|
||||||
content_type: None,
|
content_type,
|
||||||
content_encoding: None,
|
content_encoding: None,
|
||||||
body: RequestBody {
|
body,
|
||||||
json_body: None,
|
|
||||||
body_bytes_b64: None,
|
|
||||||
body_ref: None,
|
|
||||||
},
|
|
||||||
stream: false,
|
stream: false,
|
||||||
client_api_format: provider_api_format.clone(),
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: None,
|
model_name,
|
||||||
proxy: runtime.resolve_model_fetch_proxy(transport).await,
|
proxy: runtime.resolve_model_fetch_proxy(transport).await,
|
||||||
tls_profile: resolve_transport_tls_profile(transport),
|
tls_profile: resolve_transport_tls_profile(transport),
|
||||||
timeouts: resolve_transport_execution_timeouts(transport),
|
timeouts: resolve_transport_execution_timeouts(transport),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn resolve_models_fetch_auth(
|
async fn resolve_standard_header_auth(
|
||||||
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||||
transport: &GatewayProviderTransportSnapshot,
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
) -> Result<Option<(String, String)>, String> {
|
) -> Result<Option<(String, String)>, String> {
|
||||||
if transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
if transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
||||||
|| transport.key.auth_type.trim().eq_ignore_ascii_case("kiro")
|
|| transport.key.auth_type.trim().eq_ignore_ascii_case("kiro")
|
||||||
{
|
{
|
||||||
return match runtime.resolve_local_oauth_request_auth(transport).await {
|
return resolve_oauth_header_auth(runtime, transport).await;
|
||||||
Ok(Some(LocalResolvedOAuthRequestAuth::Header { name, value })) => {
|
|
||||||
Ok(Some((name, value)))
|
|
||||||
}
|
|
||||||
Ok(Some(LocalResolvedOAuthRequestAuth::Kiro(_))) => Ok(None),
|
|
||||||
Ok(None) => Ok(None),
|
|
||||||
Err(err) => Err(err),
|
|
||||||
};
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some(auth) = resolve_local_openai_chat_auth(transport) {
|
let api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
|
||||||
return Ok(Some(auth));
|
if api_format.starts_with("openai:") {
|
||||||
|
return Ok(resolve_local_openai_chat_auth(transport));
|
||||||
}
|
}
|
||||||
if let Some(auth) = resolve_local_standard_auth(transport) {
|
if api_format.starts_with("claude:") {
|
||||||
return Ok(Some(auth));
|
return Ok(resolve_local_standard_auth(transport));
|
||||||
}
|
|
||||||
if let Some(auth) = resolve_local_gemini_auth(transport) {
|
|
||||||
return Ok(Some(auth));
|
|
||||||
}
|
|
||||||
if let Some(query_auth) = resolve_local_vertex_api_key_query_auth(transport) {
|
|
||||||
let url = build_passthrough_path_url(
|
|
||||||
&transport.endpoint.base_url,
|
|
||||||
"/v1/publishers/google/models",
|
|
||||||
Some(&format!("{}={}", query_auth.name, query_auth.value)),
|
|
||||||
&[],
|
|
||||||
);
|
|
||||||
if url.is_some() {
|
|
||||||
return Ok(None);
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
Ok(None)
|
Ok(None)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn resolve_oauth_header_auth(
|
||||||
|
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
) -> Result<Option<(String, String)>, String> {
|
||||||
|
match runtime.resolve_local_oauth_request_auth(transport).await {
|
||||||
|
Ok(Some(LocalResolvedOAuthRequestAuth::Header { name, value })) => Ok(Some((name, value))),
|
||||||
|
Ok(Some(LocalResolvedOAuthRequestAuth::Kiro(_))) => Ok(None),
|
||||||
|
Ok(None) => Ok(None),
|
||||||
|
Err(err) => Err(err),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn resolve_bearer_or_oauth_header_auth(
|
||||||
|
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
) -> Result<Option<(String, String)>, String> {
|
||||||
|
if let Some(auth) = resolve_oauth_header_auth(runtime, transport).await? {
|
||||||
|
return Ok(Some(auth));
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Some((name, value)) = resolve_local_openai_chat_auth(transport) {
|
||||||
|
return Ok(Some((name, value)));
|
||||||
|
}
|
||||||
|
|
||||||
|
if transport
|
||||||
|
.key
|
||||||
|
.auth_type
|
||||||
|
.trim()
|
||||||
|
.eq_ignore_ascii_case("bearer")
|
||||||
|
{
|
||||||
|
let secret = transport.key.decrypted_api_key.trim();
|
||||||
|
if !secret.is_empty() {
|
||||||
|
return Ok(Some((
|
||||||
|
"authorization".to_string(),
|
||||||
|
format!("Bearer {secret}"),
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(None)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn apply_fetch_header_rules(
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
mut headers: BTreeMap<String, String>,
|
||||||
|
protected_headers: &[String],
|
||||||
|
) -> Result<BTreeMap<String, String>, String> {
|
||||||
|
let protected = protected_headers
|
||||||
|
.iter()
|
||||||
|
.map(String::as_str)
|
||||||
|
.collect::<Vec<_>>();
|
||||||
|
if !apply_local_header_rules(
|
||||||
|
&mut headers,
|
||||||
|
transport.endpoint.header_rules.as_ref(),
|
||||||
|
&protected,
|
||||||
|
&json!({}),
|
||||||
|
None,
|
||||||
|
) {
|
||||||
|
return Err("Endpoint header_rules application failed".to_string());
|
||||||
|
}
|
||||||
|
Ok(headers)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn standard_models_fetch_headers(api_format: &str) -> BTreeMap<String, String> {
|
||||||
|
let api_format = api_format.trim().to_ascii_lowercase();
|
||||||
|
match api_format.as_str() {
|
||||||
|
"openai:cli" | "openai:compact" => {
|
||||||
|
BTreeMap::from([("user-agent".to_string(), OPENAI_CLI_USER_AGENT.to_string())])
|
||||||
|
}
|
||||||
|
"claude:chat" => BTreeMap::from([(
|
||||||
|
"anthropic-version".to_string(),
|
||||||
|
CLAUDE_VERSION_HEADER.to_string(),
|
||||||
|
)]),
|
||||||
|
"claude:cli" => BTreeMap::from([
|
||||||
|
("user-agent".to_string(), CLAUDE_CLI_USER_AGENT.to_string()),
|
||||||
|
(
|
||||||
|
"anthropic-version".to_string(),
|
||||||
|
CLAUDE_VERSION_HEADER.to_string(),
|
||||||
|
),
|
||||||
|
]),
|
||||||
|
"gemini:chat" => BROWSER_FINGERPRINT_HEADERS
|
||||||
|
.iter()
|
||||||
|
.map(|(key, value)| (key.to_string(), value.to_string()))
|
||||||
|
.collect(),
|
||||||
|
"gemini:cli" => {
|
||||||
|
let mut headers = BROWSER_FINGERPRINT_HEADERS
|
||||||
|
.iter()
|
||||||
|
.map(|(key, value)| (key.to_string(), value.to_string()))
|
||||||
|
.collect::<BTreeMap<_, _>>();
|
||||||
|
headers.insert("user-agent".to_string(), GEMINI_CLI_USER_AGENT.to_string());
|
||||||
|
headers
|
||||||
|
}
|
||||||
|
_ => BTreeMap::new(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn build_standard_models_fetch_url(
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
after_id: Option<&str>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
let api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
|
||||||
|
if api_format.starts_with("gemini:") {
|
||||||
|
let secret = resolve_local_vertex_api_key_query_auth(transport)
|
||||||
|
.map(|auth| auth.value)
|
||||||
|
.or_else(|| {
|
||||||
|
resolve_local_gemini_auth(transport).and_then(|(name, value)| {
|
||||||
|
name.eq_ignore_ascii_case("x-goog-api-key").then_some(value)
|
||||||
|
})
|
||||||
|
})
|
||||||
|
.or_else(|| {
|
||||||
|
let secret = transport.key.decrypted_api_key.trim();
|
||||||
|
(!secret.is_empty()).then_some(secret.to_string())
|
||||||
|
})
|
||||||
|
.ok_or_else(|| "Gemini models fetch requires an API key".to_string())?;
|
||||||
|
|
||||||
|
let (url, _) = build_models_fetch_url(
|
||||||
|
&transport.provider.provider_type,
|
||||||
|
&transport.endpoint.api_format,
|
||||||
|
&transport.endpoint.base_url,
|
||||||
|
)
|
||||||
|
.ok_or_else(|| "Rust models fetch does not support this provider format yet".to_string())?;
|
||||||
|
return Ok(append_query_param(url, "key", &secret));
|
||||||
|
}
|
||||||
|
|
||||||
|
let (mut url, _) = build_models_fetch_url(
|
||||||
|
&transport.provider.provider_type,
|
||||||
|
&transport.endpoint.api_format,
|
||||||
|
&transport.endpoint.base_url,
|
||||||
|
)
|
||||||
|
.ok_or_else(|| "Rust models fetch does not support this provider format yet".to_string())?;
|
||||||
|
|
||||||
|
if api_format.starts_with("claude:") {
|
||||||
|
url = append_query_param(url, "limit", "100");
|
||||||
|
if let Some(after_id) = after_id.map(str::trim).filter(|value| !value.is_empty()) {
|
||||||
|
url = append_query_param(url, "after_id", after_id);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(url)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn append_query_param(mut url: String, key: &str, value: &str) -> String {
|
||||||
|
if key.trim().is_empty() || value.trim().is_empty() {
|
||||||
|
return url;
|
||||||
|
}
|
||||||
|
let separator = if url.contains('?') { '&' } else { '?' };
|
||||||
|
url.push(separator);
|
||||||
|
url.push_str(key.trim());
|
||||||
|
url.push('=');
|
||||||
|
url.push_str(value.trim());
|
||||||
|
url
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use aether_contracts::ProxySnapshot;
|
use aether_contracts::{ExecutionPlan, ExecutionResult, ProxySnapshot};
|
||||||
use aether_provider_transport::snapshot::{
|
use aether_provider_transport::snapshot::{
|
||||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||||
};
|
};
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
|
use serde_json::json;
|
||||||
|
|
||||||
use super::{build_models_fetch_execution_plan, ModelFetchTransportRuntime};
|
use super::{
|
||||||
|
build_antigravity_fetch_available_models_plan, build_gemini_cli_load_code_assist_plan,
|
||||||
|
build_models_fetch_execution_plan, build_standard_models_fetch_execution_plan,
|
||||||
|
build_vertex_models_fetch_execution_plan, ModelFetchTransportRuntime,
|
||||||
|
};
|
||||||
|
|
||||||
struct TestRuntime {
|
struct TestRuntime {
|
||||||
oauth_auth: Option<aether_provider_transport::LocalResolvedOAuthRequestAuth>,
|
oauth_auth: Option<aether_provider_transport::LocalResolvedOAuthRequestAuth>,
|
||||||
@@ -155,14 +484,25 @@ mod tests {
|
|||||||
) -> Option<ProxySnapshot> {
|
) -> Option<ProxySnapshot> {
|
||||||
self.proxy.clone()
|
self.proxy.clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn execute_model_fetch_execution_plan(
|
||||||
|
&self,
|
||||||
|
_plan: &ExecutionPlan,
|
||||||
|
) -> Result<ExecutionResult, String> {
|
||||||
|
unreachable!("tests only validate plan construction")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn sample_transport(api_format: &str, auth_type: &str) -> GatewayProviderTransportSnapshot {
|
fn sample_transport(
|
||||||
|
provider_type: &str,
|
||||||
|
api_format: &str,
|
||||||
|
auth_type: &str,
|
||||||
|
) -> GatewayProviderTransportSnapshot {
|
||||||
GatewayProviderTransportSnapshot {
|
GatewayProviderTransportSnapshot {
|
||||||
provider: GatewayProviderTransportProvider {
|
provider: GatewayProviderTransportProvider {
|
||||||
id: "provider-1".to_string(),
|
id: "provider-1".to_string(),
|
||||||
name: "Provider One".to_string(),
|
name: "Provider One".to_string(),
|
||||||
provider_type: "openai".to_string(),
|
provider_type: provider_type.to_string(),
|
||||||
website: None,
|
website: None,
|
||||||
is_active: true,
|
is_active: true,
|
||||||
keep_priority_on_conversion: false,
|
keep_priority_on_conversion: false,
|
||||||
@@ -196,7 +536,7 @@ mod tests {
|
|||||||
name: "key".to_string(),
|
name: "key".to_string(),
|
||||||
auth_type: auth_type.to_string(),
|
auth_type: auth_type.to_string(),
|
||||||
is_active: true,
|
is_active: true,
|
||||||
api_formats: None,
|
api_formats: Some(vec![api_format.to_string()]),
|
||||||
allowed_models: None,
|
allowed_models: None,
|
||||||
capabilities: None,
|
capabilities: None,
|
||||||
rate_multipliers: None,
|
rate_multipliers: None,
|
||||||
@@ -205,35 +545,86 @@ mod tests {
|
|||||||
proxy: None,
|
proxy: None,
|
||||||
fingerprint: None,
|
fingerprint: None,
|
||||||
decrypted_api_key: "secret".to_string(),
|
decrypted_api_key: "secret".to_string(),
|
||||||
decrypted_auth_config: None,
|
decrypted_auth_config: Some(
|
||||||
|
r#"{"project_id":"project-1","client_version":"1.2.3","session_id":"sess-1"}"#
|
||||||
|
.to_string(),
|
||||||
|
),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn builds_openai_models_fetch_plan_from_transport_snapshot() {
|
async fn builds_openai_cli_models_fetch_plan_with_cli_user_agent() {
|
||||||
let runtime = TestRuntime {
|
let runtime = TestRuntime {
|
||||||
oauth_auth: None,
|
oauth_auth: None,
|
||||||
proxy: None,
|
proxy: None,
|
||||||
};
|
};
|
||||||
let plan = build_models_fetch_execution_plan(
|
let mut transport = sample_transport("openai", "openai:cli", "api_key");
|
||||||
&runtime,
|
transport.key.decrypted_auth_config = None;
|
||||||
&sample_transport("openai:chat", "api_key"),
|
let plan = build_models_fetch_execution_plan(&runtime, &transport)
|
||||||
)
|
.await
|
||||||
.await
|
.expect("plan");
|
||||||
.expect("plan");
|
|
||||||
|
|
||||||
assert_eq!(plan.method, "GET");
|
|
||||||
assert_eq!(plan.url, "https://example.com/v1/models");
|
assert_eq!(plan.url, "https://example.com/v1/models");
|
||||||
|
assert_eq!(
|
||||||
|
plan.headers.get("user-agent").map(String::as_str),
|
||||||
|
Some("openai-codex/1.0")
|
||||||
|
);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
plan.headers.get("authorization").map(String::as_str),
|
plan.headers.get("authorization").map(String::as_str),
|
||||||
Some("Bearer secret")
|
Some("Bearer secret")
|
||||||
);
|
);
|
||||||
assert_eq!(plan.provider_api_format, "openai:chat");
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn builds_oauth_models_fetch_plan_from_runtime_auth() {
|
async fn builds_claude_models_fetch_plan_with_pagination() {
|
||||||
|
let runtime = TestRuntime {
|
||||||
|
oauth_auth: None,
|
||||||
|
proxy: None,
|
||||||
|
};
|
||||||
|
let mut transport = sample_transport("custom", "claude:chat", "api_key");
|
||||||
|
transport.key.decrypted_auth_config = None;
|
||||||
|
let plan =
|
||||||
|
build_standard_models_fetch_execution_plan(&runtime, &transport, Some("cursor-1"))
|
||||||
|
.await
|
||||||
|
.expect("plan");
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
plan.url,
|
||||||
|
"https://example.com/v1/models?limit=100&after_id=cursor-1"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
plan.headers.get("anthropic-version").map(String::as_str),
|
||||||
|
Some("2023-06-01")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
plan.headers.get("x-api-key").map(String::as_str),
|
||||||
|
Some("secret")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn builds_gemini_models_fetch_plan_with_browser_headers_and_query_auth() {
|
||||||
|
let runtime = TestRuntime {
|
||||||
|
oauth_auth: None,
|
||||||
|
proxy: None,
|
||||||
|
};
|
||||||
|
let mut transport = sample_transport("custom", "gemini:chat", "api_key");
|
||||||
|
transport.key.decrypted_auth_config = None;
|
||||||
|
let plan = build_models_fetch_execution_plan(&runtime, &transport)
|
||||||
|
.await
|
||||||
|
.expect("plan");
|
||||||
|
|
||||||
|
assert_eq!(plan.url, "https://example.com/v1beta/models?key=secret");
|
||||||
|
assert!(plan.headers.contains_key("sec-ch-ua"));
|
||||||
|
assert_eq!(
|
||||||
|
plan.headers.get("accept-language").map(String::as_str),
|
||||||
|
Some("zh-CN")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn builds_antigravity_fetch_available_models_plan() {
|
||||||
let runtime = TestRuntime {
|
let runtime = TestRuntime {
|
||||||
oauth_auth: Some(
|
oauth_auth: Some(
|
||||||
aether_provider_transport::LocalResolvedOAuthRequestAuth::Header {
|
aether_provider_transport::LocalResolvedOAuthRequestAuth::Header {
|
||||||
@@ -241,27 +632,85 @@ mod tests {
|
|||||||
value: "Bearer oauth-token".to_string(),
|
value: "Bearer oauth-token".to_string(),
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
proxy: Some(ProxySnapshot {
|
proxy: None,
|
||||||
enabled: Some(true),
|
|
||||||
mode: Some("fixed".to_string()),
|
|
||||||
node_id: None,
|
|
||||||
label: None,
|
|
||||||
url: Some("http://proxy.internal".to_string()),
|
|
||||||
extra: None,
|
|
||||||
}),
|
|
||||||
};
|
};
|
||||||
let plan =
|
let transport = sample_transport("antigravity", "gemini:chat", "oauth");
|
||||||
build_models_fetch_execution_plan(&runtime, &sample_transport("openai:chat", "oauth"))
|
let plan = build_antigravity_fetch_available_models_plan(
|
||||||
.await
|
&runtime,
|
||||||
.expect("plan");
|
&transport,
|
||||||
|
"https://daily-cloudcode-pa.sandbox.googleapis.com",
|
||||||
|
"project-1",
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect("plan");
|
||||||
|
|
||||||
|
assert_eq!(plan.method, "POST");
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
plan.headers.get("authorization").map(String::as_str),
|
plan.url,
|
||||||
Some("Bearer oauth-token")
|
"https://daily-cloudcode-pa.sandbox.googleapis.com/v1internal:fetchAvailableModels"
|
||||||
);
|
);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
plan.proxy.as_ref().and_then(|proxy| proxy.url.as_deref()),
|
plan.provider_api_format,
|
||||||
Some("http://proxy.internal")
|
"antigravity:fetch_available_models"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
plan.body
|
||||||
|
.json_body
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|value| value.get("project")),
|
||||||
|
Some(&json!("project-1"))
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn builds_gemini_cli_load_code_assist_plan() {
|
||||||
|
let runtime = TestRuntime {
|
||||||
|
oauth_auth: Some(
|
||||||
|
aether_provider_transport::LocalResolvedOAuthRequestAuth::Header {
|
||||||
|
name: "authorization".to_string(),
|
||||||
|
value: "Bearer oauth-token".to_string(),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
proxy: None,
|
||||||
|
};
|
||||||
|
let transport = sample_transport("gemini_cli", "gemini:cli", "oauth");
|
||||||
|
let plan = build_gemini_cli_load_code_assist_plan(&runtime, &transport)
|
||||||
|
.await
|
||||||
|
.expect("plan");
|
||||||
|
|
||||||
|
assert_eq!(plan.method, "POST");
|
||||||
|
assert_eq!(
|
||||||
|
plan.url,
|
||||||
|
"https://cloudcode-pa.googleapis.com/v1internal:loadCodeAssist"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
plan.headers.get("user-agent").map(String::as_str),
|
||||||
|
Some("GeminiCLI/0.1.5 (Windows; AMD64)")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn builds_vertex_models_fetch_plan_with_auth_override() {
|
||||||
|
let runtime = TestRuntime {
|
||||||
|
oauth_auth: None,
|
||||||
|
proxy: None,
|
||||||
|
};
|
||||||
|
let mut transport = sample_transport("vertex_ai", "claude:chat", "api_key");
|
||||||
|
transport.key.decrypted_auth_config = None;
|
||||||
|
let plan = build_vertex_models_fetch_execution_plan(
|
||||||
|
&runtime,
|
||||||
|
&transport,
|
||||||
|
"https://aiplatform.googleapis.com/v1/publishers/google/models?key=secret",
|
||||||
|
"gemini:chat",
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect("plan");
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
plan.url,
|
||||||
|
"https://aiplatform.googleapis.com/v1/publishers/google/models?key=secret"
|
||||||
|
);
|
||||||
|
assert!(plan.headers.contains_key("sec-ch-ua"));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user