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:
fawney19
2026-04-05 20:23:16 +08:00
parent cbc811f6ce
commit 763ff03a7b
777 changed files with 42659 additions and 21469 deletions

View File

@@ -1,5 +1,7 @@
mod runtime;
#[cfg(test)]
mod tests;
pub(crate) use runtime::{
perform_model_fetch_once, spawn_model_fetch_worker, ModelFetchRunSummary,
};
pub(crate) use aether_model_fetch::ModelFetchRunSummary;
pub(crate) use runtime::state::ModelFetchRuntimeState;
pub(crate) use runtime::{perform_model_fetch_once, spawn_model_fetch_worker};

View File

@@ -1,42 +1,27 @@
use std::collections::{BTreeMap, BTreeSet, HashMap};
use std::collections::HashMap;
use std::time::Duration;
use std::time::{SystemTime, UNIX_EPOCH};
use aether_contracts::{ExecutionPlan, ExecutionResult, RequestBody};
use aether_contracts::ExecutionResult;
use aether_data::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use http::HeaderMap;
use regex::Regex;
use aether_model_fetch::{
apply_model_filters, build_models_fetch_execution_plan, extract_error_message,
json_string_list, model_fetch_interval_minutes, model_fetch_startup_delay_seconds,
model_fetch_startup_enabled, parse_models_response, select_models_fetch_endpoint,
sync_provider_model_whitelist_associations, ModelFetchAssociationStore, ModelFetchRunSummary,
ModelsFetchSuccess,
};
use serde_json::{json, Value};
use tracing::{debug, info, warn};
use crate::gateway::provider_transport::{
apply_local_header_rules, build_passthrough_path_url, ensure_upstream_auth_header,
resolve_local_gemini_auth, resolve_local_openai_chat_auth, resolve_local_standard_auth,
resolve_local_vertex_api_key_query_auth, resolve_transport_execution_timeouts,
resolve_transport_proxy_snapshot_with_tunnel_affinity, resolve_transport_tls_profile,
LocalResolvedOAuthRequestAuth,
};
use crate::gateway::{AppState, GatewayError};
use crate::provider_transport::GatewayProviderTransportSnapshot;
use crate::{AppState, GatewayError};
mod association_sync;
pub(crate) mod state;
use self::association_sync::sync_provider_model_whitelist_associations;
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;
const MODEL_FETCH_CACHE_KEY_PREFIX: &str = "upstream_models";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct ModelFetchRunSummary {
pub(crate) attempted: usize,
pub(crate) succeeded: usize,
pub(crate) failed: usize,
pub(crate) skipped: usize,
}
use self::state::ModelFetchRuntimeState;
#[derive(Debug, Clone)]
struct SelectedFetchTarget {
@@ -45,12 +30,6 @@ struct SelectedFetchTarget {
key: StoredProviderCatalogKey,
}
#[derive(Debug, Clone)]
struct ModelsFetchSuccess {
fetched_model_ids: Vec<String>,
cached_models: Vec<Value>,
}
pub(crate) fn spawn_model_fetch_worker(state: AppState) -> Option<tokio::task::JoinHandle<()>> {
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
return None;
@@ -86,7 +65,16 @@ pub(crate) fn spawn_model_fetch_worker(state: AppState) -> Option<tokio::task::J
pub(crate) async fn perform_model_fetch_once(
state: &AppState,
) -> Result<ModelFetchRunSummary, GatewayError> {
if !state.data.has_provider_catalog_reader() || !state.data.has_provider_catalog_writer() {
perform_model_fetch_once_with_state(state).await
}
async fn perform_model_fetch_once_with_state<S>(
state: &S,
) -> Result<ModelFetchRunSummary, GatewayError>
where
S: ModelFetchRuntimeState + ?Sized,
{
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
return Ok(ModelFetchRunSummary {
attempted: 0,
succeeded: 0,
@@ -120,9 +108,12 @@ pub(crate) async fn perform_model_fetch_once(
.push(endpoint);
}
let mut keys_by_provider = HashMap::<String, Vec<StoredProviderCatalogKey>>::new();
for key in state
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
.await?
for key in <S as ModelFetchAssociationStore>::list_provider_catalog_keys_by_provider_ids(
state,
&provider_ids,
)
.await
.map_err(GatewayError::Internal)?
{
keys_by_provider
.entry(key.provider_id.clone())
@@ -191,8 +182,11 @@ pub(crate) async fn perform_model_fetch_once(
Ok(summary)
}
async fn run_model_fetch_cycle(state: &AppState, phase: &'static str) -> Result<(), GatewayError> {
let summary = perform_model_fetch_once(state).await?;
async fn run_model_fetch_cycle<S>(state: &S, phase: &'static str) -> Result<(), GatewayError>
where
S: ModelFetchRuntimeState + ?Sized,
{
let summary = perform_model_fetch_once_with_state(state).await?;
if summary.attempted == 0 {
debug!(phase, "gateway model fetch found no eligible keys");
return Ok(());
@@ -217,7 +211,7 @@ enum KeyFetchDisposition {
}
async fn fetch_and_persist_key_models(
state: &AppState,
state: &(impl ModelFetchRuntimeState + ?Sized),
target: &SelectedFetchTarget,
) -> Result<KeyFetchDisposition, GatewayError> {
let now_unix_secs = now_unix_secs();
@@ -268,69 +262,23 @@ async fn fetch_and_persist_key_models(
);
persist_key_fetch_success(state, &target.key, now_unix_secs, &filtered_models).await?;
write_upstream_models_cache(
state,
&target.provider.id,
&target.key.id,
&result.cached_models,
)
.await;
state
.write_upstream_models_cache(&target.provider.id, &target.key.id, &result.cached_models)
.await;
sync_provider_model_whitelist_associations(state, &target.provider.id, &filtered_models)
.await?;
.await
.map_err(GatewayError::Internal)?;
Ok(KeyFetchDisposition::Succeeded)
}
async fn execute_models_fetch_request(
state: &AppState,
transport: &crate::gateway::provider_transport::GatewayProviderTransportSnapshot,
state: &(impl ModelFetchRuntimeState + ?Sized),
transport: &GatewayProviderTransportSnapshot,
) -> Result<ModelsFetchSuccess, String> {
let (upstream_url, provider_api_format) = build_models_fetch_url(transport)
.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(state, transport)
.await?
.ok_or_else(|| {
"Rust models fetch auth resolution is not supported for this key".to_string()
})?;
let plan = build_models_fetch_execution_plan(state, transport).await?;
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);
let plan = 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: resolve_transport_proxy_snapshot_with_tunnel_affinity(state, transport).await,
tls_profile: resolve_transport_tls_profile(transport),
timeouts: resolve_transport_execution_timeouts(transport),
};
let result = crate::gateway::execute_execution_runtime_sync_plan(state, None, &plan)
let result = state
.execute_execution_runtime_sync_plan(&plan)
.await
.map_err(|err| format!("{err:?}"))?;
@@ -355,211 +303,11 @@ async fn execute_models_fetch_request(
.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, body_json)
}
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)
})
}
fn build_models_fetch_url(
transport: &crate::gateway::provider_transport::GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
let api_format = normalize_api_format(&transport.endpoint.api_format);
if !crate::gateway::provider_transport::provider_type_supports_model_fetch(
&transport.provider.provider_type,
) {
return None;
}
let url = if api_format.starts_with("openai:") || api_format.starts_with("claude:") {
build_v1_models_url(&transport.endpoint.base_url)
} else if api_format.starts_with("gemini:") {
build_gemini_models_url(&transport.endpoint.base_url)
} else {
return None;
}?;
Some((url, api_format))
}
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))
}
async fn resolve_models_fetch_auth(
state: &AppState,
transport: &crate::gateway::provider_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 state.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(format!("{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)
}
fn parse_models_response(
transport: &crate::gateway::provider_transport::GatewayProviderTransportSnapshot,
body: &Value,
) -> Result<ModelsFetchSuccess, String> {
let api_format = normalize_api_format(&transport.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,
})
}
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)
parse_models_response(&transport.endpoint.api_format, body_json)
}
async fn persist_key_fetch_failure(
state: &AppState,
state: &(impl ModelFetchRuntimeState + ?Sized),
key: &StoredProviderCatalogKey,
now_unix_secs: u64,
error: String,
@@ -573,7 +321,7 @@ async fn persist_key_fetch_failure(
}
async fn persist_key_fetch_success(
state: &AppState,
state: &(impl ModelFetchRuntimeState + ?Sized),
key: &StoredProviderCatalogKey,
now_unix_secs: u64,
allowed_models: &[String],
@@ -591,291 +339,9 @@ async fn persist_key_fetch_success(
Ok(())
}
async fn write_upstream_models_cache(
state: &AppState,
provider_id: &str,
key_id: &str,
cached_models: &[Value],
) {
let Some(runner) = state.redis_kv_runner() else {
return;
};
let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(cached_models)) else {
return;
};
let cache_key = format!("{MODEL_FETCH_CACHE_KEY_PREFIX}:{provider_id}:{key_id}");
if let Err(err) = runner
.setex(
&cache_key,
&serialized,
Some(model_fetch_interval_minutes().saturating_mul(60)),
)
.await
{
debug!(
provider_id = %provider_id,
key_id = %key_id,
error = %err,
"gateway model fetch cache write failed"
);
}
}
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 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()
}
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"
)
}
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()
}
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(&regex::escape(&other.to_string())),
}
}
regex.push('$');
Regex::new(&regex)
.ok()
.is_some_and(|compiled| compiled.is_match(model_id))
}
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()
}
fn normalize_api_format(value: &str) -> String {
value.trim().to_ascii_lowercase()
}
fn now_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
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)
}
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)
}
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::{aggregate_models_for_cache, apply_model_filters, build_gemini_models_url};
use serde_json::json;
#[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"
);
}
}

View File

@@ -1,184 +0,0 @@
use std::collections::BTreeSet;
use aether_data::repository::global_models::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, UpsertAdminProviderModelRecord,
};
use serde_json::Value;
use uuid::Uuid;
use crate::gateway::{AppState, GatewayError};
pub(super) async fn sync_provider_model_whitelist_associations(
state: &AppState,
provider_id: &str,
current_allowed_models: &[String],
) -> Result<(), GatewayError> {
if !state.data.has_global_model_reader() || !state.data.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(
state: &AppState,
provider_id: &str,
allowed_models: &[String],
) -> Result<(), GatewayError> {
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| GatewayError::Internal(err.to_string()))?;
state.create_admin_provider_model(&record).await?;
}
Ok(())
}
async fn auto_disassociate_provider_by_key_whitelist(
state: &AppState,
provider_id: &str,
) -> Result<(), GatewayError> {
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| super::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 matches_model_mapping(pattern: &str, model_name: &str) -> bool {
regex::Regex::new(&format!("^(?:{pattern})$"))
.ok()
.is_some_and(|compiled| compiled.is_match(model_name))
}
fn is_oauth_auth_type(value: &str) -> bool {
matches!(value.trim().to_ascii_lowercase().as_str(), "oauth" | "kiro")
}

View File

@@ -0,0 +1,58 @@
use aether_contracts::{ExecutionPlan, ExecutionResult, ProxySnapshot};
use aether_data::repository::global_models::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, StoredAdminGlobalModelPage,
StoredAdminProviderModel, UpsertAdminProviderModelRecord,
};
use aether_data::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_model_fetch::{ModelFetchAssociationStore, ModelFetchTransportRuntime};
use async_trait::async_trait;
use serde_json::Value;
use crate::provider_transport::{
GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
};
use crate::{AppState, GatewayError};
#[async_trait]
pub(crate) trait ModelFetchRuntimeState:
ModelFetchAssociationStore<Error = String> + ModelFetchTransportRuntime + Sync
{
fn has_provider_catalog_data_reader(&self) -> bool;
fn has_provider_catalog_data_writer(&self) -> bool;
async fn list_provider_catalog_providers(
&self,
active_only: bool,
) -> Result<Vec<StoredProviderCatalogProvider>, GatewayError>;
async fn list_provider_catalog_endpoints_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogEndpoint>, GatewayError>;
async fn read_provider_transport_snapshot(
&self,
provider_id: &str,
endpoint_id: &str,
key_id: &str,
) -> Result<Option<GatewayProviderTransportSnapshot>, GatewayError>;
async fn execute_execution_runtime_sync_plan(
&self,
plan: &ExecutionPlan,
) -> Result<ExecutionResult, GatewayError>;
async fn update_provider_catalog_key(
&self,
key: &StoredProviderCatalogKey,
) -> Result<(), GatewayError>;
async fn write_upstream_models_cache(
&self,
provider_id: &str,
key_id: &str,
cached_models: &[Value],
);
}

View File

@@ -0,0 +1,544 @@
use std::sync::{Arc, Mutex};
use std::time::Duration;
use aether_contracts::ExecutionPlan;
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_data::repository::global_models::{
GlobalModelReadRepository, InMemoryGlobalModelReadRepository, StoredAdminGlobalModel,
StoredAdminProviderModel,
};
use aether_data::repository::provider_catalog::{
InMemoryProviderCatalogReadRepository, ProviderCatalogReadRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use axum::routing::any;
use axum::{extract::Request, Json, Router};
use serde_json::json;
use super::{perform_model_fetch_once, ModelFetchRunSummary};
use crate::AppState;
async fn start_server(app: Router) -> (String, tokio::task::JoinHandle<()>) {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let handle = tokio::spawn(async move {
axum::serve(
listener,
app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
)
.await
.expect("server should run");
});
(format!("http://{addr}"), handle)
}
fn build_state_with_execution_runtime_override(
execution_runtime_override_base_url: impl Into<String>,
) -> AppState {
AppState::new()
.expect("gateway should build")
.with_execution_runtime_override_base_url(execution_runtime_override_base_url)
}
fn sample_provider(provider_id: &str) -> StoredProviderCatalogProvider {
StoredProviderCatalogProvider::new(
provider_id.to_string(),
"openai".to_string(),
Some("https://example.com".to_string()),
"custom".to_string(),
)
.expect("provider should build")
.with_transport_fields(true, false, true, None, None, None, None, None, None)
}
fn sample_endpoint(
provider_id: &str,
endpoint_id: &str,
base_url: &str,
) -> StoredProviderCatalogEndpoint {
StoredProviderCatalogEndpoint::new(
endpoint_id.to_string(),
provider_id.to_string(),
"openai:chat".to_string(),
Some("openai".to_string()),
Some("chat".to_string()),
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 {
let mut key = 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"])),
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "live-secret-api-key")
.expect("api key should encrypt"),
None,
None,
None,
Some(json!(["gpt-4.1"])),
None,
None,
None,
)
.expect("key transport should build");
key.auto_fetch_models = true;
key.locked_models = Some(json!(["locked-model"]));
key.model_include_patterns = Some(json!(["gpt-*"]));
key.model_exclude_patterns = Some(json!(["gpt-beta"]));
key
}
fn sample_global_model(id: &str, name: &str, mappings: &[&str]) -> StoredAdminGlobalModel {
StoredAdminGlobalModel::new(
id.to_string(),
name.to_string(),
name.to_string(),
true,
None,
None,
None,
Some(json!({ "model_mappings": mappings })),
Some(1_711_000_000),
Some(1_711_000_000),
)
.expect("global model should build")
}
fn sample_provider_model(
id: &str,
provider_id: &str,
global_model_id: &str,
provider_model_name: &str,
global_model_name: &str,
mappings: &[&str],
) -> StoredAdminProviderModel {
StoredAdminProviderModel::new(
id.to_string(),
provider_id.to_string(),
global_model_id.to_string(),
provider_model_name.to_string(),
None,
None,
None,
None,
None,
None,
None,
None,
true,
true,
None,
Some(1_711_000_000),
Some(1_711_000_000),
Some(global_model_name.to_string()),
Some(global_model_name.to_string()),
None,
None,
Some(json!({ "model_mappings": mappings })),
)
.expect("provider model should build")
}
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 }
}
#[tokio::test]
async fn gateway_model_fetch_updates_key_and_syncs_provider_model_whitelist_associations() {
let seen_execution_runtime_plan = Arc::new(Mutex::new(None::<ExecutionPlan>));
let seen_execution_runtime_plan_clone = Arc::clone(&seen_execution_runtime_plan);
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any(move |Json(plan): Json<ExecutionPlan>| {
let seen_execution_runtime_plan = Arc::clone(&seen_execution_runtime_plan_clone);
async move {
*seen_execution_runtime_plan
.lock()
.expect("mutex should lock") = Some(plan);
Json(json!({
"request_id": "req-model-fetch-key-1",
"status_code": 200,
"headers": {
"content-type": "application/json"
},
"body": {
"json_body": {
"data": [
{"id": "gpt-5"},
{"id": "gpt-beta"},
{"id": "other-model"}
]
}
}
}))
}
}),
);
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider("provider-openai")],
vec![sample_endpoint(
"provider-openai",
"endpoint-openai",
"https://api.openai.example",
)],
vec![sample_key("provider-openai", "key-openai")],
));
let global_model_repository = Arc::new(
InMemoryGlobalModelReadRepository::seed(Vec::new())
.with_admin_global_models(vec![
sample_global_model("global-model-gpt5", "gpt-5", &["gpt-5"]),
sample_global_model("global-model-gpt4", "gpt-4.1", &["gpt-4\\.1"]),
])
.with_admin_provider_models(vec![sample_provider_model(
"provider-model-gpt4",
"provider-openai",
"global-model-gpt4",
"gpt-4.1",
"gpt-4.1",
&["gpt-4\\.1"],
)]),
);
let data_state = crate::data::GatewayDataState::disabled()
.attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository))
.with_global_model_repository_for_tests(Arc::clone(&global_model_repository))
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
let state = build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(data_state);
let summary = perform_model_fetch_once(&state)
.await
.expect("model fetch should succeed");
assert_eq!(
summary,
ModelFetchRunSummary {
attempted: 1,
succeeded: 1,
failed: 0,
skipped: 0,
}
);
let seen_plan = seen_execution_runtime_plan
.lock()
.expect("mutex should lock")
.clone()
.expect("execution runtime plan should be captured");
assert_eq!(seen_plan.method, "GET");
assert_eq!(seen_plan.url, "https://api.openai.example/v1/models");
assert_eq!(
seen_plan.headers.get("authorization").map(String::as_str),
Some("Bearer live-secret-api-key")
);
let updated_key = provider_catalog_repository
.list_keys_by_ids(&["key-openai".to_string()])
.await
.expect("keys should load")
.into_iter()
.next()
.expect("updated key should exist");
assert_eq!(
updated_key.allowed_models,
Some(json!(["gpt-5", "locked-model"]))
);
assert_eq!(updated_key.last_models_fetch_error, None);
assert!(updated_key.last_models_fetch_at_unix_secs.is_some());
let provider_models = global_model_repository
.list_admin_provider_models(
&aether_data::repository::global_models::AdminProviderModelListQuery {
provider_id: "provider-openai".to_string(),
is_active: None,
offset: 0,
limit: 10_000,
},
)
.await
.expect("provider models should load");
assert_eq!(provider_models.len(), 1);
assert_eq!(
provider_models[0].global_model_name.as_deref(),
Some("gpt-5")
);
assert_eq!(provider_models[0].provider_model_name, "gpt-5");
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_model_fetch_updates_key_and_syncs_provider_model_whitelist_associations_without_execution_runtime_override(
) {
#[derive(Debug, Clone, PartialEq, Eq)]
struct SeenUpstreamRequest {
method: String,
authorization: String,
}
let seen_upstream = Arc::new(Mutex::new(None::<SeenUpstreamRequest>));
let seen_upstream_clone = Arc::clone(&seen_upstream);
let upstream = Router::new().route(
"/v1/models",
any(move |request: Request| {
let seen_upstream_inner = Arc::clone(&seen_upstream_clone);
async move {
let authorization = request
.headers()
.get(http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string();
*seen_upstream_inner.lock().expect("mutex should lock") =
Some(SeenUpstreamRequest {
method: request.method().as_str().to_string(),
authorization,
});
Json(json!({
"data": [
{"id": "gpt-5"},
{"id": "gpt-beta"},
{"id": "other-model"}
]
}))
}
}),
);
let (upstream_url, upstream_handle) = start_server(upstream).await;
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider("provider-openai")],
vec![sample_endpoint(
"provider-openai",
"endpoint-openai",
&upstream_url,
)],
vec![sample_key("provider-openai", "key-openai")],
));
let global_model_repository = Arc::new(
InMemoryGlobalModelReadRepository::seed(Vec::new())
.with_admin_global_models(vec![
sample_global_model("global-model-gpt5", "gpt-5", &["gpt-5"]),
sample_global_model("global-model-gpt4", "gpt-4.1", &["gpt-4\\.1"]),
])
.with_admin_provider_models(vec![sample_provider_model(
"provider-model-gpt4",
"provider-openai",
"global-model-gpt4",
"gpt-4.1",
"gpt-4.1",
&["gpt-4\\.1"],
)]),
);
let data_state = crate::data::GatewayDataState::disabled()
.attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository))
.with_global_model_repository_for_tests(Arc::clone(&global_model_repository))
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
let summary = perform_model_fetch_once(&state)
.await
.expect("model fetch should succeed");
assert_eq!(
summary,
ModelFetchRunSummary {
attempted: 1,
succeeded: 1,
failed: 0,
skipped: 0,
}
);
assert_eq!(
seen_upstream.lock().expect("mutex should lock").clone(),
Some(SeenUpstreamRequest {
method: "GET".to_string(),
authorization: "Bearer live-secret-api-key".to_string(),
})
);
let updated_key = provider_catalog_repository
.list_keys_by_ids(&["key-openai".to_string()])
.await
.expect("keys should load")
.into_iter()
.next()
.expect("updated key should exist");
assert_eq!(
updated_key.allowed_models,
Some(json!(["gpt-5", "locked-model"]))
);
assert_eq!(updated_key.last_models_fetch_error, None);
assert!(updated_key.last_models_fetch_at_unix_secs.is_some());
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_background_model_fetch_updates_key_and_syncs_provider_model_whitelist_associations(
) {
let _startup_enabled = set_test_env_var("MODEL_FETCH_STARTUP_ENABLED", "true");
let _startup_delay = set_test_env_var("MODEL_FETCH_STARTUP_DELAY_SECONDS", "0");
let seen_execution_runtime_plan = Arc::new(Mutex::new(None::<ExecutionPlan>));
let seen_execution_runtime_plan_clone = Arc::clone(&seen_execution_runtime_plan);
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any(move |Json(plan): Json<ExecutionPlan>| {
let seen_execution_runtime_plan = Arc::clone(&seen_execution_runtime_plan_clone);
async move {
*seen_execution_runtime_plan
.lock()
.expect("mutex should lock") = Some(plan);
Json(json!({
"request_id": "req-model-fetch-key-1",
"status_code": 200,
"headers": {
"content-type": "application/json"
},
"body": {
"json_body": {
"data": [
{"id": "gpt-5"},
{"id": "gpt-beta"},
{"id": "other-model"}
]
}
}
}))
}
}),
);
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider("provider-openai")],
vec![sample_endpoint(
"provider-openai",
"endpoint-openai",
"https://api.openai.example",
)],
vec![sample_key("provider-openai", "key-openai")],
));
let global_model_repository = Arc::new(
InMemoryGlobalModelReadRepository::seed(Vec::new())
.with_admin_global_models(vec![
sample_global_model("global-model-gpt5", "gpt-5", &["gpt-5"]),
sample_global_model("global-model-gpt4", "gpt-4.1", &["gpt-4\\.1"]),
])
.with_admin_provider_models(vec![sample_provider_model(
"provider-model-gpt4",
"provider-openai",
"global-model-gpt4",
"gpt-4.1",
"gpt-4.1",
&["gpt-4\\.1"],
)]),
);
let data_state = crate::data::GatewayDataState::disabled()
.attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository))
.with_global_model_repository_for_tests(Arc::clone(&global_model_repository))
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
let gateway_state = build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(data_state);
let background_tasks = gateway_state.spawn_background_tasks();
assert!(
!background_tasks.is_empty(),
"model fetch worker should spawn"
);
let updated_key = tokio::time::timeout(Duration::from_secs(5), async {
loop {
let updated_key = provider_catalog_repository
.list_keys_by_ids(&["key-openai".to_string()])
.await
.expect("keys should load")
.into_iter()
.next()
.expect("updated key should exist");
if updated_key.last_models_fetch_at_unix_secs.is_some() {
break updated_key;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("background worker should fetch models on startup");
assert_eq!(
updated_key.allowed_models,
Some(json!(["gpt-5", "locked-model"]))
);
assert_eq!(updated_key.last_models_fetch_error, None);
let provider_models = global_model_repository
.list_admin_provider_models(
&aether_data::repository::global_models::AdminProviderModelListQuery {
provider_id: "provider-openai".to_string(),
is_active: None,
offset: 0,
limit: 10_000,
},
)
.await
.expect("provider models should load");
assert_eq!(provider_models.len(), 1);
assert_eq!(
provider_models[0].global_model_name.as_deref(),
Some("gpt-5")
);
let seen_plan = seen_execution_runtime_plan
.lock()
.expect("mutex should lock")
.clone()
.expect("execution runtime plan should be captured");
assert_eq!(seen_plan.url, "https://api.openai.example/v1/models");
for handle in background_tasks {
handle.abort();
}
execution_runtime_handle.abort();
}