mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +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:
@@ -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};
|
||||
|
||||
@@ -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(®ex::escape(&other.to_string())),
|
||||
}
|
||||
}
|
||||
regex.push('$');
|
||||
Regex::new(®ex)
|
||||
.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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
58
apps/aether-gateway/src/model_fetch/runtime/state.rs
Normal file
58
apps/aether-gateway/src/model_fetch/runtime/state.rs
Normal 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],
|
||||
);
|
||||
}
|
||||
544
apps/aether-gateway/src/model_fetch/tests.rs
Normal file
544
apps/aether-gateway/src/model_fetch/tests.rs
Normal 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();
|
||||
}
|
||||
Reference in New Issue
Block a user