mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor: 拆分 gateway 单体为独立 crate,新增 systemd 部署方案
将 gateway 内部的 model-fetch、provider-transport、scheduler-core、 usage-runtime、video-tasks-core 模块提取为独立 crate;重构 gateway 内部模块结构(state/router/cache/data/query 等);移除大量遗留模块 文件;新增 systemd 二进制部署骨架及相关文档;更新前端 usage 相关 API 和组件。
This commit is contained in:
226
crates/aether-model-fetch/src/association_sync.rs
Normal file
226
crates/aether-model-fetch/src/association_sync.rs
Normal file
@@ -0,0 +1,226 @@
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use aether_data::repository::global_models::{
|
||||
AdminGlobalModelListQuery, AdminProviderModelListQuery, StoredAdminGlobalModelPage,
|
||||
StoredAdminProviderModel, UpsertAdminProviderModelRecord,
|
||||
};
|
||||
use aether_data::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
use aether_scheduler_core::matches_model_mapping;
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::json_string_list;
|
||||
|
||||
#[async_trait]
|
||||
pub trait ModelFetchAssociationStore {
|
||||
type Error: Send;
|
||||
|
||||
fn has_global_model_reader(&self) -> bool;
|
||||
fn has_global_model_writer(&self) -> bool;
|
||||
fn model_fetch_internal_error(&self, message: String) -> Self::Error;
|
||||
|
||||
async fn list_admin_provider_models(
|
||||
&self,
|
||||
query: &AdminProviderModelListQuery,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, Self::Error>;
|
||||
|
||||
async fn list_admin_global_models(
|
||||
&self,
|
||||
query: &AdminGlobalModelListQuery,
|
||||
) -> Result<StoredAdminGlobalModelPage, Self::Error>;
|
||||
|
||||
async fn create_admin_provider_model(
|
||||
&self,
|
||||
record: &UpsertAdminProviderModelRecord,
|
||||
) -> Result<Option<StoredAdminProviderModel>, Self::Error>;
|
||||
|
||||
async fn list_provider_catalog_keys_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, Self::Error>;
|
||||
|
||||
async fn delete_admin_provider_model(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
model_id: &str,
|
||||
) -> Result<bool, Self::Error>;
|
||||
}
|
||||
|
||||
pub async fn sync_provider_model_whitelist_associations<S>(
|
||||
state: &S,
|
||||
provider_id: &str,
|
||||
current_allowed_models: &[String],
|
||||
) -> Result<(), S::Error>
|
||||
where
|
||||
S: ModelFetchAssociationStore + Sync + ?Sized,
|
||||
{
|
||||
if !state.has_global_model_reader() || !state.has_global_model_writer() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
auto_associate_provider_by_key_whitelist(state, provider_id, current_allowed_models).await?;
|
||||
auto_disassociate_provider_by_key_whitelist(state, provider_id).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn auto_associate_provider_by_key_whitelist<S>(
|
||||
state: &S,
|
||||
provider_id: &str,
|
||||
allowed_models: &[String],
|
||||
) -> Result<(), S::Error>
|
||||
where
|
||||
S: ModelFetchAssociationStore + Sync + ?Sized,
|
||||
{
|
||||
if allowed_models.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let provider_models = state
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider_id.to_string(),
|
||||
is_active: None,
|
||||
offset: 0,
|
||||
limit: 10_000,
|
||||
})
|
||||
.await?;
|
||||
let linked_global_model_ids = provider_models
|
||||
.iter()
|
||||
.map(|model| model.global_model_id.clone())
|
||||
.collect::<BTreeSet<_>>();
|
||||
let existing_provider_model_names = provider_models
|
||||
.iter()
|
||||
.map(|model| model.provider_model_name.clone())
|
||||
.collect::<BTreeSet<_>>();
|
||||
let global_models = state
|
||||
.list_admin_global_models(&AdminGlobalModelListQuery {
|
||||
offset: 0,
|
||||
limit: 10_000,
|
||||
is_active: Some(true),
|
||||
search: None,
|
||||
})
|
||||
.await?
|
||||
.items;
|
||||
|
||||
for global_model in global_models {
|
||||
if linked_global_model_ids.contains(&global_model.id)
|
||||
|| existing_provider_model_names.contains(&global_model.name)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
|
||||
let mappings = global_model_mapping_patterns(global_model.config.as_ref());
|
||||
if mappings.is_empty() {
|
||||
continue;
|
||||
}
|
||||
if !allowed_models.iter().any(|allowed_model| {
|
||||
mappings
|
||||
.iter()
|
||||
.any(|pattern| matches_model_mapping(pattern, allowed_model))
|
||||
}) {
|
||||
continue;
|
||||
}
|
||||
|
||||
let record = UpsertAdminProviderModelRecord::new(
|
||||
Uuid::new_v4().to_string(),
|
||||
provider_id.to_string(),
|
||||
global_model.id.clone(),
|
||||
global_model.name.clone(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
true,
|
||||
None,
|
||||
)
|
||||
.map_err(|err| state.model_fetch_internal_error(err.to_string()))?;
|
||||
state.create_admin_provider_model(&record).await?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn auto_disassociate_provider_by_key_whitelist<S>(
|
||||
state: &S,
|
||||
provider_id: &str,
|
||||
) -> Result<(), S::Error>
|
||||
where
|
||||
S: ModelFetchAssociationStore + Sync + ?Sized,
|
||||
{
|
||||
let keys = state
|
||||
.list_provider_catalog_keys_by_provider_ids(&[provider_id.to_string()])
|
||||
.await?;
|
||||
let active_non_oauth_keys = keys
|
||||
.into_iter()
|
||||
.filter(|key| key.is_active)
|
||||
.filter(|key| !is_oauth_auth_type(&key.auth_type))
|
||||
.collect::<Vec<_>>();
|
||||
if active_non_oauth_keys.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
if active_non_oauth_keys
|
||||
.iter()
|
||||
.any(|key| key.allowed_models.is_none())
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let all_allowed_models = active_non_oauth_keys
|
||||
.iter()
|
||||
.flat_map(|key| json_string_list(key.allowed_models.as_ref()))
|
||||
.collect::<BTreeSet<_>>();
|
||||
let provider_models = state
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider_id.to_string(),
|
||||
is_active: None,
|
||||
offset: 0,
|
||||
limit: 10_000,
|
||||
})
|
||||
.await?;
|
||||
|
||||
for model in provider_models {
|
||||
let mappings = global_model_mapping_patterns(model.global_model_config.as_ref());
|
||||
if mappings.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let matched = all_allowed_models.iter().any(|allowed_model| {
|
||||
mappings
|
||||
.iter()
|
||||
.any(|pattern| matches_model_mapping(pattern, allowed_model))
|
||||
});
|
||||
if matched {
|
||||
continue;
|
||||
}
|
||||
state
|
||||
.delete_admin_provider_model(provider_id, &model.id)
|
||||
.await?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn global_model_mapping_patterns(config: Option<&Value>) -> Vec<String> {
|
||||
config
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|object| object.get("model_mappings"))
|
||||
.and_then(Value::as_array)
|
||||
.map(|items| {
|
||||
items
|
||||
.iter()
|
||||
.filter_map(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn is_oauth_auth_type(value: &str) -> bool {
|
||||
matches!(value.trim().to_ascii_lowercase().as_str(), "oauth" | "kiro")
|
||||
}
|
||||
77
crates/aether-model-fetch/src/config.rs
Normal file
77
crates/aether-model-fetch/src/config.rs
Normal file
@@ -0,0 +1,77 @@
|
||||
const MODEL_FETCH_INTERVAL_MINUTES_DEFAULT: u64 = 1440;
|
||||
const MODEL_FETCH_INTERVAL_MINUTES_MIN: u64 = 60;
|
||||
const MODEL_FETCH_INTERVAL_MINUTES_MAX: u64 = 10080;
|
||||
const MODEL_FETCH_STARTUP_DELAY_SECONDS_DEFAULT: u64 = 10;
|
||||
|
||||
pub fn model_fetch_interval_minutes() -> u64 {
|
||||
std::env::var("MODEL_FETCH_INTERVAL_MINUTES")
|
||||
.ok()
|
||||
.and_then(|value| value.parse::<u64>().ok())
|
||||
.map(|value| {
|
||||
value.clamp(
|
||||
MODEL_FETCH_INTERVAL_MINUTES_MIN,
|
||||
MODEL_FETCH_INTERVAL_MINUTES_MAX,
|
||||
)
|
||||
})
|
||||
.unwrap_or(MODEL_FETCH_INTERVAL_MINUTES_DEFAULT)
|
||||
}
|
||||
|
||||
pub fn model_fetch_startup_enabled() -> bool {
|
||||
std::env::var("MODEL_FETCH_STARTUP_ENABLED")
|
||||
.ok()
|
||||
.map(|value| value.trim().eq_ignore_ascii_case("true"))
|
||||
.unwrap_or(true)
|
||||
}
|
||||
|
||||
pub fn model_fetch_startup_delay_seconds() -> u64 {
|
||||
std::env::var("MODEL_FETCH_STARTUP_DELAY_SECONDS")
|
||||
.ok()
|
||||
.and_then(|value| value.parse::<u64>().ok())
|
||||
.unwrap_or(MODEL_FETCH_STARTUP_DELAY_SECONDS_DEFAULT)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
model_fetch_interval_minutes, model_fetch_startup_delay_seconds,
|
||||
model_fetch_startup_enabled,
|
||||
};
|
||||
|
||||
struct TestEnvVarGuard {
|
||||
key: &'static str,
|
||||
previous: Option<String>,
|
||||
}
|
||||
|
||||
impl Drop for TestEnvVarGuard {
|
||||
fn drop(&mut self) {
|
||||
if let Some(previous) = self.previous.as_deref() {
|
||||
std::env::set_var(self.key, previous);
|
||||
} else {
|
||||
std::env::remove_var(self.key);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn set_test_env_var(key: &'static str, value: &str) -> TestEnvVarGuard {
|
||||
let previous = std::env::var(key).ok();
|
||||
std::env::set_var(key, value);
|
||||
TestEnvVarGuard { key, previous }
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn interval_minutes_clamps_to_supported_bounds() {
|
||||
let _interval = set_test_env_var("MODEL_FETCH_INTERVAL_MINUTES", "5");
|
||||
assert_eq!(model_fetch_interval_minutes(), 60);
|
||||
|
||||
let _interval = set_test_env_var("MODEL_FETCH_INTERVAL_MINUTES", "20000");
|
||||
assert_eq!(model_fetch_interval_minutes(), 10080);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn startup_flags_read_from_environment() {
|
||||
let _enabled = set_test_env_var("MODEL_FETCH_STARTUP_ENABLED", "false");
|
||||
let _delay = set_test_env_var("MODEL_FETCH_STARTUP_DELAY_SECONDS", "3");
|
||||
assert!(!model_fetch_startup_enabled());
|
||||
assert_eq!(model_fetch_startup_delay_seconds(), 3);
|
||||
}
|
||||
}
|
||||
17
crates/aether-model-fetch/src/lib.rs
Normal file
17
crates/aether-model-fetch/src/lib.rs
Normal file
@@ -0,0 +1,17 @@
|
||||
mod association_sync;
|
||||
mod config;
|
||||
mod logic;
|
||||
mod transport;
|
||||
|
||||
pub use association_sync::{
|
||||
sync_provider_model_whitelist_associations, ModelFetchAssociationStore,
|
||||
};
|
||||
pub use config::{
|
||||
model_fetch_interval_minutes, model_fetch_startup_delay_seconds, model_fetch_startup_enabled,
|
||||
};
|
||||
pub use logic::{
|
||||
aggregate_models_for_cache, apply_model_filters, build_models_fetch_url,
|
||||
endpoint_supports_rust_models_fetch, extract_error_message, json_string_list,
|
||||
parse_models_response, select_models_fetch_endpoint, ModelFetchRunSummary, ModelsFetchSuccess,
|
||||
};
|
||||
pub use transport::{build_models_fetch_execution_plan, ModelFetchTransportRuntime};
|
||||
510
crates/aether-model-fetch/src/logic.rs
Normal file
510
crates/aether-model-fetch/src/logic.rs
Normal file
@@ -0,0 +1,510 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use aether_data::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
use aether_provider_transport::provider_types::provider_type_supports_model_fetch;
|
||||
use regex::Regex;
|
||||
use serde_json::Value;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct ModelFetchRunSummary {
|
||||
pub attempted: usize,
|
||||
pub succeeded: usize,
|
||||
pub failed: usize,
|
||||
pub skipped: usize,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct ModelsFetchSuccess {
|
||||
pub fetched_model_ids: Vec<String>,
|
||||
pub cached_models: Vec<Value>,
|
||||
}
|
||||
|
||||
pub fn extract_error_message(value: &Value) -> Option<String> {
|
||||
value
|
||||
.get("error")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|error| error.get("message"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.or_else(|| {
|
||||
value
|
||||
.get("message")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
}
|
||||
|
||||
pub fn build_models_fetch_url(
|
||||
provider_type: &str,
|
||||
endpoint_api_format: &str,
|
||||
base_url: &str,
|
||||
) -> Option<(String, String)> {
|
||||
let api_format = normalize_api_format(endpoint_api_format);
|
||||
if !provider_type_supports_model_fetch(provider_type) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let url = if api_format.starts_with("openai:") || api_format.starts_with("claude:") {
|
||||
build_v1_models_url(base_url)
|
||||
} else if api_format.starts_with("gemini:") {
|
||||
build_gemini_models_url(base_url)
|
||||
} else {
|
||||
return None;
|
||||
}?;
|
||||
Some((url, api_format))
|
||||
}
|
||||
|
||||
pub fn parse_models_response(
|
||||
endpoint_api_format: &str,
|
||||
body: &Value,
|
||||
) -> Result<ModelsFetchSuccess, String> {
|
||||
let api_format = normalize_api_format(endpoint_api_format);
|
||||
let mut cached_models = Vec::new();
|
||||
let mut fetched_model_ids = Vec::new();
|
||||
let mut seen = BTreeSet::new();
|
||||
|
||||
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) {
|
||||
items
|
||||
} else if let Some(items) = body.as_array() {
|
||||
items
|
||||
} else {
|
||||
return Err("models response is missing data array".to_string());
|
||||
};
|
||||
for item in items {
|
||||
let Some(model_id) = item
|
||||
.get("id")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if !seen.insert(model_id.to_string()) {
|
||||
continue;
|
||||
}
|
||||
fetched_model_ids.push(model_id.to_string());
|
||||
cached_models.push(normalize_cached_model(item, model_id, &api_format));
|
||||
}
|
||||
} else if api_format.starts_with("gemini:") {
|
||||
let items = body
|
||||
.get("models")
|
||||
.and_then(Value::as_array)
|
||||
.ok_or_else(|| "gemini models response is missing models array".to_string())?;
|
||||
for item in items {
|
||||
let Some(name) = item
|
||||
.get("name")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let model_id = name.strip_prefix("models/").unwrap_or(name).trim();
|
||||
if model_id.is_empty() || !seen.insert(model_id.to_string()) {
|
||||
continue;
|
||||
}
|
||||
fetched_model_ids.push(model_id.to_string());
|
||||
cached_models.push(normalize_cached_model(item, model_id, &api_format));
|
||||
}
|
||||
} else {
|
||||
return Err("models response parser does not support this provider format".to_string());
|
||||
}
|
||||
|
||||
Ok(ModelsFetchSuccess {
|
||||
fetched_model_ids,
|
||||
cached_models,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn select_models_fetch_endpoint(
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> Option<StoredProviderCatalogEndpoint> {
|
||||
let key_formats = json_string_list(key.api_formats.as_ref())
|
||||
.into_iter()
|
||||
.map(|value| normalize_api_format(&value))
|
||||
.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 {
|
||||
let api_format = normalize_api_format(api_format);
|
||||
matches!(
|
||||
api_format.as_str(),
|
||||
"openai:chat"
|
||||
| "openai:cli"
|
||||
| "openai:responses"
|
||||
| "openai:compact"
|
||||
| "claude:chat"
|
||||
| "claude:cli"
|
||||
| "gemini:chat"
|
||||
| "gemini:cli"
|
||||
)
|
||||
}
|
||||
|
||||
pub fn apply_model_filters(
|
||||
fetched_model_ids: &[String],
|
||||
locked_models: Vec<String>,
|
||||
include_patterns: Vec<String>,
|
||||
exclude_patterns: Vec<String>,
|
||||
) -> Vec<String> {
|
||||
let mut filtered = BTreeSet::new();
|
||||
for model_id in fetched_model_ids {
|
||||
if model_id.trim().is_empty() {
|
||||
continue;
|
||||
}
|
||||
let included = if include_patterns.is_empty() {
|
||||
true
|
||||
} else {
|
||||
include_patterns
|
||||
.iter()
|
||||
.any(|pattern| wildcard_matches(pattern, model_id))
|
||||
};
|
||||
if !included {
|
||||
continue;
|
||||
}
|
||||
let excluded = exclude_patterns
|
||||
.iter()
|
||||
.any(|pattern| wildcard_matches(pattern, model_id));
|
||||
if !excluded {
|
||||
filtered.insert(model_id.trim().to_string());
|
||||
}
|
||||
}
|
||||
for model in locked_models {
|
||||
let trimmed = model.trim();
|
||||
if !trimmed.is_empty() {
|
||||
filtered.insert(trimmed.to_string());
|
||||
}
|
||||
}
|
||||
filtered.into_iter().collect()
|
||||
}
|
||||
|
||||
pub fn json_string_list(value: Option<&Value>) -> Vec<String> {
|
||||
value
|
||||
.and_then(Value::as_array)
|
||||
.map(|items| {
|
||||
items
|
||||
.iter()
|
||||
.filter_map(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
pub fn aggregate_models_for_cache(models: &[Value]) -> Vec<Value> {
|
||||
let mut aggregated = BTreeMap::<String, serde_json::Map<String, Value>>::new();
|
||||
let mut order = Vec::<String>::new();
|
||||
|
||||
for model in models {
|
||||
let Some(object) = model.as_object() else {
|
||||
continue;
|
||||
};
|
||||
let Some(model_id) = object
|
||||
.get("id")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let entry = aggregated.entry(model_id.to_string()).or_insert_with(|| {
|
||||
order.push(model_id.to_string());
|
||||
let mut cloned = object.clone();
|
||||
cloned.remove("api_format");
|
||||
cloned
|
||||
});
|
||||
|
||||
let api_formats = object
|
||||
.get("api_formats")
|
||||
.and_then(Value::as_array)
|
||||
.map(|items| {
|
||||
items
|
||||
.iter()
|
||||
.filter_map(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<BTreeSet<_>>()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
let existing_formats = entry
|
||||
.get("api_formats")
|
||||
.and_then(Value::as_array)
|
||||
.map(|items| {
|
||||
items
|
||||
.iter()
|
||||
.filter_map(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<BTreeSet<_>>()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
let merged_formats = existing_formats
|
||||
.union(&api_formats)
|
||||
.cloned()
|
||||
.map(Value::String)
|
||||
.collect::<Vec<_>>();
|
||||
entry.insert("api_formats".to_string(), Value::Array(merged_formats));
|
||||
|
||||
for (key, value) in object {
|
||||
if key == "api_format" || entry.contains_key(key) {
|
||||
continue;
|
||||
}
|
||||
entry.insert(key.clone(), value.clone());
|
||||
}
|
||||
}
|
||||
|
||||
order
|
||||
.into_iter()
|
||||
.filter_map(|model_id| aggregated.remove(&model_id))
|
||||
.map(Value::Object)
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn build_v1_models_url(base_url: &str) -> Option<String> {
|
||||
let (trimmed_base_url, query) = split_url_query(base_url);
|
||||
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
|
||||
if trimmed_base_url.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let mut url = if trimmed_base_url.ends_with("/v1") {
|
||||
format!("{trimmed_base_url}/models")
|
||||
} else {
|
||||
format!("{trimmed_base_url}/v1/models")
|
||||
};
|
||||
if let Some(query) = query.filter(|value| !value.trim().is_empty()) {
|
||||
url.push('?');
|
||||
url.push_str(query);
|
||||
}
|
||||
Some(url)
|
||||
}
|
||||
|
||||
fn build_gemini_models_url(base_url: &str) -> Option<String> {
|
||||
let (trimmed_base_url, base_query) = split_url_query(base_url);
|
||||
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
|
||||
if trimmed_base_url.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut url = if trimmed_base_url.ends_with("/v1beta") {
|
||||
format!("{trimmed_base_url}/models")
|
||||
} else if trimmed_base_url.contains("/v1beta/models") {
|
||||
trimmed_base_url.to_string()
|
||||
} else {
|
||||
format!("{trimmed_base_url}/v1beta/models")
|
||||
};
|
||||
if let Some(query) = base_query.filter(|value| !value.trim().is_empty()) {
|
||||
url.push('?');
|
||||
url.push_str(query);
|
||||
}
|
||||
Some(url)
|
||||
}
|
||||
|
||||
fn split_url_query(base_url: &str) -> (&str, Option<&str>) {
|
||||
let trimmed = base_url.trim();
|
||||
trimmed
|
||||
.split_once('?')
|
||||
.map(|(base, query)| (base, Some(query)))
|
||||
.unwrap_or((trimmed, None))
|
||||
}
|
||||
|
||||
fn normalize_cached_model(item: &Value, model_id: &str, api_format: &str) -> Value {
|
||||
let mut object = item.as_object().cloned().unwrap_or_default();
|
||||
object.insert("id".to_string(), Value::String(model_id.to_string()));
|
||||
object.insert(
|
||||
"api_formats".to_string(),
|
||||
Value::Array(vec![Value::String(api_format.to_string())]),
|
||||
);
|
||||
object.remove("api_format");
|
||||
Value::Object(object)
|
||||
}
|
||||
|
||||
fn wildcard_matches(pattern: &str, model_id: &str) -> bool {
|
||||
let mut regex = String::from("^");
|
||||
for ch in pattern.chars() {
|
||||
match ch {
|
||||
'*' => regex.push_str(".*"),
|
||||
'?' => regex.push('.'),
|
||||
other => regex.push_str(®ex::escape(&other.to_string())),
|
||||
}
|
||||
}
|
||||
regex.push('$');
|
||||
Regex::new(®ex)
|
||||
.ok()
|
||||
.is_some_and(|compiled| compiled.is_match(model_id))
|
||||
}
|
||||
|
||||
fn normalize_api_format(value: &str) -> String {
|
||||
value.trim().to_ascii_lowercase()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use aether_data::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
aggregate_models_for_cache, apply_model_filters, build_gemini_models_url,
|
||||
build_models_fetch_url, parse_models_response, select_models_fetch_endpoint,
|
||||
};
|
||||
|
||||
fn sample_endpoint(
|
||||
provider_id: &str,
|
||||
endpoint_id: &str,
|
||||
api_format: &str,
|
||||
base_url: &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(
|
||||
base_url.to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("endpoint transport should build")
|
||||
}
|
||||
|
||||
fn sample_key(provider_id: &str, key_id: &str) -> StoredProviderCatalogKey {
|
||||
StoredProviderCatalogKey::new(
|
||||
key_id.to_string(),
|
||||
provider_id.to_string(),
|
||||
"primary".to_string(),
|
||||
"api_key".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build")
|
||||
.with_transport_fields(
|
||||
Some(json!(["openai:chat"])),
|
||||
"encrypted".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("key transport should build")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn apply_model_filters_respects_include_exclude_and_locked_models() {
|
||||
let filtered = apply_model_filters(
|
||||
&[
|
||||
"gpt-5".to_string(),
|
||||
"gpt-beta".to_string(),
|
||||
"claude-4".to_string(),
|
||||
],
|
||||
vec!["locked-model".to_string()],
|
||||
vec!["gpt-*".to_string()],
|
||||
vec!["gpt-beta".to_string()],
|
||||
);
|
||||
assert_eq!(
|
||||
filtered,
|
||||
vec!["gpt-5".to_string(), "locked-model".to_string()]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn aggregate_models_for_cache_merges_api_formats_by_model_id() {
|
||||
let aggregated = aggregate_models_for_cache(&[
|
||||
json!({"id":"gpt-5","api_formats":["openai:chat"]}),
|
||||
json!({"id":"gpt-5","api_formats":["openai:cli"]}),
|
||||
]);
|
||||
assert_eq!(aggregated.len(), 1);
|
||||
assert_eq!(
|
||||
aggregated[0]["api_formats"],
|
||||
json!(["openai:chat", "openai:cli"])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_gemini_models_url_preserves_base_query() {
|
||||
let url =
|
||||
build_gemini_models_url("https://generativelanguage.googleapis.com/v1beta?key=abc")
|
||||
.expect("gemini models url should build");
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://generativelanguage.googleapis.com/v1beta/models?key=abc"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_models_fetch_url_rejects_provider_types_without_fetch_support() {
|
||||
assert_eq!(
|
||||
build_models_fetch_url("vertex_ai", "gemini:chat", "https://example.com"),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_models_response_normalizes_openai_payload() {
|
||||
let parsed = parse_models_response(
|
||||
"openai:chat",
|
||||
&json!({"data": [{"id": "gpt-5"}, {"id": "gpt-5"}]}),
|
||||
)
|
||||
.expect("response should parse");
|
||||
assert_eq!(parsed.fetched_model_ids, vec!["gpt-5".to_string()]);
|
||||
assert_eq!(
|
||||
parsed.cached_models[0]["api_formats"],
|
||||
json!(["openai:chat"])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn select_models_fetch_endpoint_respects_key_api_formats() {
|
||||
let key = sample_key("provider-1", "key-1");
|
||||
let endpoints = vec![
|
||||
sample_endpoint(
|
||||
"provider-1",
|
||||
"endpoint-cli",
|
||||
"openai:cli",
|
||||
"https://example.com",
|
||||
),
|
||||
sample_endpoint(
|
||||
"provider-1",
|
||||
"endpoint-chat",
|
||||
"openai:chat",
|
||||
"https://example.com",
|
||||
),
|
||||
];
|
||||
let selected =
|
||||
select_models_fetch_endpoint(&endpoints, &key).expect("endpoint should be selected");
|
||||
assert_eq!(selected.id, "endpoint-chat");
|
||||
}
|
||||
}
|
||||
267
crates/aether-model-fetch/src/transport.rs
Normal file
267
crates/aether-model-fetch/src/transport.rs
Normal file
@@ -0,0 +1,267 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_contracts::{ExecutionPlan, ProxySnapshot, RequestBody};
|
||||
use aether_provider_transport::auth::{
|
||||
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::{
|
||||
apply_local_header_rules, ensure_upstream_auth_header, resolve_transport_execution_timeouts,
|
||||
resolve_transport_tls_profile, GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::json;
|
||||
|
||||
use crate::build_models_fetch_url;
|
||||
|
||||
#[async_trait]
|
||||
pub trait ModelFetchTransportRuntime: Send + Sync {
|
||||
async fn resolve_local_oauth_request_auth(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Result<Option<LocalResolvedOAuthRequestAuth>, String>;
|
||||
|
||||
async fn resolve_model_fetch_proxy(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<ProxySnapshot>;
|
||||
}
|
||||
|
||||
pub async fn build_models_fetch_execution_plan(
|
||||
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Result<ExecutionPlan, String> {
|
||||
let (upstream_url, provider_api_format) = 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())?;
|
||||
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())]);
|
||||
if !apply_local_header_rules(
|
||||
&mut headers,
|
||||
transport.endpoint.header_rules.as_ref(),
|
||||
&[auth_header_name.as_str()],
|
||||
&json!({}),
|
||||
None,
|
||||
) {
|
||||
return Err("Endpoint header_rules application failed".to_string());
|
||||
}
|
||||
ensure_upstream_auth_header(&mut headers, &auth_header_name, &auth_header_value);
|
||||
|
||||
Ok(ExecutionPlan {
|
||||
request_id: format!("req-model-fetch-{}", transport.key.id),
|
||||
candidate_id: None,
|
||||
provider_name: Some(transport.provider.name.clone()),
|
||||
provider_id: transport.provider.id.clone(),
|
||||
endpoint_id: transport.endpoint.id.clone(),
|
||||
key_id: transport.key.id.clone(),
|
||||
method: "GET".to_string(),
|
||||
url: upstream_url,
|
||||
headers,
|
||||
content_type: None,
|
||||
content_encoding: None,
|
||||
body: RequestBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: None,
|
||||
body_ref: None,
|
||||
},
|
||||
stream: false,
|
||||
client_api_format: provider_api_format.clone(),
|
||||
provider_api_format,
|
||||
model_name: None,
|
||||
proxy: runtime.resolve_model_fetch_proxy(transport).await,
|
||||
tls_profile: resolve_transport_tls_profile(transport),
|
||||
timeouts: resolve_transport_execution_timeouts(transport),
|
||||
})
|
||||
}
|
||||
|
||||
async fn resolve_models_fetch_auth(
|
||||
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Result<Option<(String, String)>, String> {
|
||||
if transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
||||
|| transport.key.auth_type.trim().eq_ignore_ascii_case("kiro")
|
||||
{
|
||||
return 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),
|
||||
};
|
||||
}
|
||||
|
||||
if let Some(auth) = resolve_local_openai_chat_auth(transport) {
|
||||
return Ok(Some(auth));
|
||||
}
|
||||
if let Some(auth) = resolve_local_standard_auth(transport) {
|
||||
return Ok(Some(auth));
|
||||
}
|
||||
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)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_provider_transport::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::{build_models_fetch_execution_plan, ModelFetchTransportRuntime};
|
||||
|
||||
struct TestRuntime {
|
||||
oauth_auth: Option<aether_provider_transport::LocalResolvedOAuthRequestAuth>,
|
||||
proxy: Option<ProxySnapshot>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ModelFetchTransportRuntime for TestRuntime {
|
||||
async fn resolve_local_oauth_request_auth(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Result<Option<aether_provider_transport::LocalResolvedOAuthRequestAuth>, String>
|
||||
{
|
||||
Ok(self.oauth_auth.clone())
|
||||
}
|
||||
|
||||
async fn resolve_model_fetch_proxy(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<ProxySnapshot> {
|
||||
self.proxy.clone()
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_transport(api_format: &str, auth_type: &str) -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Provider One".to_string(),
|
||||
provider_type: "openai".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: Some(30.0),
|
||||
stream_first_byte_timeout_secs: Some(5.0),
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: api_format.to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: None,
|
||||
is_active: true,
|
||||
base_url: "https://example.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-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: auth_type.to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
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: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn builds_openai_models_fetch_plan_from_transport_snapshot() {
|
||||
let runtime = TestRuntime {
|
||||
oauth_auth: None,
|
||||
proxy: None,
|
||||
};
|
||||
let plan = build_models_fetch_execution_plan(
|
||||
&runtime,
|
||||
&sample_transport("openai:chat", "api_key"),
|
||||
)
|
||||
.await
|
||||
.expect("plan");
|
||||
|
||||
assert_eq!(plan.method, "GET");
|
||||
assert_eq!(plan.url, "https://example.com/v1/models");
|
||||
assert_eq!(
|
||||
plan.headers.get("authorization").map(String::as_str),
|
||||
Some("Bearer secret")
|
||||
);
|
||||
assert_eq!(plan.provider_api_format, "openai:chat");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn builds_oauth_models_fetch_plan_from_runtime_auth() {
|
||||
let runtime = TestRuntime {
|
||||
oauth_auth: Some(
|
||||
aether_provider_transport::LocalResolvedOAuthRequestAuth::Header {
|
||||
name: "authorization".to_string(),
|
||||
value: "Bearer oauth-token".to_string(),
|
||||
},
|
||||
),
|
||||
proxy: Some(ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
mode: Some("fixed".to_string()),
|
||||
node_id: None,
|
||||
label: None,
|
||||
url: Some("http://proxy.internal".to_string()),
|
||||
extra: None,
|
||||
}),
|
||||
};
|
||||
let plan =
|
||||
build_models_fetch_execution_plan(&runtime, &sample_transport("openai:chat", "oauth"))
|
||||
.await
|
||||
.expect("plan");
|
||||
|
||||
assert_eq!(
|
||||
plan.headers.get("authorization").map(String::as_str),
|
||||
Some("Bearer oauth-token")
|
||||
);
|
||||
assert_eq!(
|
||||
plan.proxy.as_ref().and_then(|proxy| proxy.url.as_deref()),
|
||||
Some("http://proxy.internal")
|
||||
);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user