mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
Merge remote-tracking branch 'pr-434/fix-provider-model-test-compat' into aether-rust-pioneer
This commit is contained in:
@@ -19,6 +19,7 @@ pub(crate) mod submission;
|
||||
pub(crate) mod sync;
|
||||
pub(crate) mod transport;
|
||||
|
||||
pub(crate) use self::chatgpt_web_image::maybe_execute_chatgpt_web_image_sync;
|
||||
pub(crate) use self::constants::{
|
||||
MAX_ERROR_BODY_BYTES, MAX_STREAM_PREFETCH_BYTES, MAX_STREAM_PREFETCH_FRAMES,
|
||||
};
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,521 @@
|
||||
use super::payload::{
|
||||
provider_query_extract_api_key_id, provider_query_extract_force_refresh,
|
||||
provider_query_extract_model, provider_query_extract_provider_id,
|
||||
provider_query_extract_request_id,
|
||||
};
|
||||
use super::response::{
|
||||
build_admin_provider_query_bad_request_response, build_admin_provider_query_not_found_response,
|
||||
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL, ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL,
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL, ADMIN_PROVIDER_QUERY_NO_LOCAL_MODELS_DETAIL,
|
||||
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
|
||||
ADMIN_PROVIDER_QUERY_PROVIDER_NOT_FOUND_DETAIL,
|
||||
};
|
||||
use crate::ai_serving::{
|
||||
maybe_build_sync_finalize_outcome, GatewayControlDecision,
|
||||
ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME, GEMINI_CHAT_SYNC_FINALIZE_REPORT_KIND,
|
||||
OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND,
|
||||
};
|
||||
use crate::clock::current_unix_ms;
|
||||
use crate::execution_runtime;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::handlers::shared::provider_pool::{
|
||||
admin_provider_pool_config_from_config_value, read_admin_provider_pool_runtime_state,
|
||||
AdminProviderPoolConfig, AdminProviderPoolRuntimeState,
|
||||
};
|
||||
use crate::handlers::shared::{
|
||||
parse_catalog_auth_config_json, provider_key_health_summary,
|
||||
provider_key_status_snapshot_payload,
|
||||
};
|
||||
use crate::model_fetch::ModelFetchRuntimeState;
|
||||
use crate::provider_key_auth::{
|
||||
provider_key_auth_semantics, provider_key_configured_api_formats,
|
||||
provider_key_inherits_provider_api_formats,
|
||||
};
|
||||
use crate::provider_transport::antigravity::{
|
||||
build_antigravity_safe_v1internal_request, build_antigravity_static_identity_headers,
|
||||
classify_local_antigravity_request_support, AntigravityEnvelopeRequestType,
|
||||
AntigravityRequestEnvelopeSupport, AntigravityRequestSideSupport,
|
||||
AntigravityRequestSideUnsupportedReason,
|
||||
};
|
||||
use crate::provider_transport::kiro::{
|
||||
build_kiro_generate_assistant_response_url, build_kiro_provider_headers,
|
||||
build_kiro_provider_request_body, supports_local_kiro_request_transport_with_network,
|
||||
KiroProviderHeadersInput, KIRO_ENVELOPE_NAME,
|
||||
};
|
||||
use crate::usage::GatewaySyncReportRequest;
|
||||
use crate::{AppState, GatewayError};
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
use aether_ai_serving::{
|
||||
run_ai_pool_scheduler, AiPoolCandidateFacts, AiPoolCandidateInput, AiPoolCatalogKeyContext,
|
||||
AiPoolRuntimeState, AiPoolSchedulingConfig, AiPoolSchedulingPreset,
|
||||
};
|
||||
use aether_contracts::{ExecutionPlan, RequestBody};
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||
};
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
AdminProviderModelListQuery, StoredAdminProviderModel,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_model_fetch::{
|
||||
aggregate_models_for_cache, fetch_models_from_transports, json_string_list,
|
||||
preset_models_for_provider, selected_models_fetch_endpoints,
|
||||
};
|
||||
use axum::{
|
||||
body::{to_bytes, Body},
|
||||
http::{self, HeaderMap, HeaderName, HeaderValue},
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use base64::Engine as _;
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering};
|
||||
use uuid::Uuid;
|
||||
|
||||
pub(crate) const ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_MESSAGE: &str =
|
||||
"Rust local provider-query model test is not configured";
|
||||
pub(crate) const ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_FAILOVER_MESSAGE: &str =
|
||||
"Rust local provider-query failover simulation is not configured";
|
||||
const ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL: &str =
|
||||
"No active endpoints found for this provider";
|
||||
const ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_ENDPOINT_DETAIL: &str =
|
||||
"No models returned from any endpoint";
|
||||
const ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL: &str = "No models returned from any key";
|
||||
const ADMIN_PROVIDER_QUERY_NO_ACTIVE_TEST_CANDIDATE_DETAIL: &str =
|
||||
"No active endpoint or API key found";
|
||||
const ADMIN_PROVIDER_QUERY_INVALID_MAPPED_MODEL_DETAIL: &str =
|
||||
"mapped_model_name is not valid for the selected model and endpoint";
|
||||
const ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX: &str = "upstream_models_provider:";
|
||||
const DEFAULT_PROVIDER_QUERY_TEST_MESSAGE: &str = "Hello! This is a test message.";
|
||||
static PROVIDER_QUERY_POOL_LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
|
||||
|
||||
#[derive(Debug)]
|
||||
struct ProviderQueryKeyFetchResult {
|
||||
models: Vec<Value>,
|
||||
error: Option<String>,
|
||||
from_cache: bool,
|
||||
has_success: bool,
|
||||
}
|
||||
|
||||
fn provider_query_codex_preset_fallback(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
) -> Option<ProviderQueryKeyFetchResult> {
|
||||
if !provider.provider_type.trim().eq_ignore_ascii_case("codex") {
|
||||
return None;
|
||||
}
|
||||
let models = preset_models_for_provider(&provider.provider_type)?;
|
||||
Some(ProviderQueryKeyFetchResult {
|
||||
models: aggregate_models_for_cache(&models),
|
||||
error: None,
|
||||
from_cache: false,
|
||||
has_success: true,
|
||||
})
|
||||
}
|
||||
|
||||
mod model_test;
|
||||
|
||||
pub(crate) use self::model_test::{
|
||||
build_admin_provider_query_test_model_failover_local_response,
|
||||
build_admin_provider_query_test_model_failover_response,
|
||||
build_admin_provider_query_test_model_local_response,
|
||||
build_admin_provider_query_test_model_response,
|
||||
};
|
||||
|
||||
fn provider_query_provider_payload(provider: &StoredProviderCatalogProvider) -> Value {
|
||||
json!({
|
||||
"id": provider.id.clone(),
|
||||
"name": provider.name.clone(),
|
||||
"display_name": provider.name.clone(),
|
||||
"provider_type": provider.provider_type.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_key_display_name(key: &StoredProviderCatalogKey) -> String {
|
||||
let trimmed = key.name.trim();
|
||||
if trimmed.is_empty() {
|
||||
key.id.clone()
|
||||
} else {
|
||||
trimmed.to_string()
|
||||
}
|
||||
}
|
||||
|
||||
async fn provider_query_read_cached_models(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
key_id: &str,
|
||||
) -> Option<Vec<Value>> {
|
||||
let cache_key = format!("upstream_models:{provider_id}:{key_id}");
|
||||
let raw = state.runtime_state().kv_get(&cache_key).await.ok()??;
|
||||
let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
|
||||
Some(aggregate_models_for_cache(&parsed))
|
||||
}
|
||||
|
||||
async fn provider_query_read_provider_cached_models(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
) -> Option<Vec<Value>> {
|
||||
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
|
||||
let raw = state.runtime_state().kv_get(&cache_key).await.ok()??;
|
||||
let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
|
||||
Some(aggregate_models_for_cache(&parsed))
|
||||
}
|
||||
|
||||
async fn provider_query_write_provider_cached_models(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
models: &[Value],
|
||||
) {
|
||||
let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(models)) else {
|
||||
return;
|
||||
};
|
||||
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
|
||||
let _ = state
|
||||
.runtime_state()
|
||||
.kv_set(
|
||||
&cache_key,
|
||||
serialized,
|
||||
Some(std::time::Duration::from_secs(
|
||||
aether_model_fetch::model_fetch_interval_minutes().saturating_mul(60),
|
||||
)),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
fn provider_query_antigravity_tier_weight(raw_auth_config: Option<&str>) -> i32 {
|
||||
raw_auth_config
|
||||
.and_then(|value| serde_json::from_str::<Value>(value).ok())
|
||||
.and_then(|value| value.get("tier").cloned())
|
||||
.and_then(|value| value.as_str().map(ToOwned::to_owned))
|
||||
.map(|tier| match tier.trim().to_ascii_lowercase().as_str() {
|
||||
"ultra" => 3,
|
||||
"pro" => 2,
|
||||
"free" => 1,
|
||||
_ => 0,
|
||||
})
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
async fn provider_query_sort_antigravity_keys(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, GatewayError> {
|
||||
let mut ranked = Vec::new();
|
||||
for key in keys {
|
||||
let availability = if key.oauth_invalid_at_unix_secs.is_some() {
|
||||
0
|
||||
} else {
|
||||
1
|
||||
};
|
||||
let tier_weight = if let Some(endpoint) = selected_models_fetch_endpoints(endpoints, &key)
|
||||
.into_iter()
|
||||
.next()
|
||||
{
|
||||
state
|
||||
.app()
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
.await?
|
||||
.map(|transport| {
|
||||
provider_query_antigravity_tier_weight(
|
||||
transport.key.decrypted_auth_config.as_deref(),
|
||||
)
|
||||
})
|
||||
.unwrap_or(0)
|
||||
} else {
|
||||
0
|
||||
};
|
||||
ranked.push(((availability, tier_weight), key));
|
||||
}
|
||||
ranked.sort_by_key(|entry| std::cmp::Reverse(entry.0));
|
||||
Ok(ranked.into_iter().map(|(_, key)| key).collect())
|
||||
}
|
||||
|
||||
async fn provider_query_fetch_models_for_key(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
key: &StoredProviderCatalogKey,
|
||||
force_refresh: bool,
|
||||
) -> Result<ProviderQueryKeyFetchResult, GatewayError> {
|
||||
if !force_refresh {
|
||||
if let Some(cached_models) =
|
||||
provider_query_read_cached_models(state, &provider.id, &key.id).await
|
||||
{
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: cached_models,
|
||||
error: None,
|
||||
from_cache: true,
|
||||
has_success: true,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let selected_endpoints = selected_models_fetch_endpoints(endpoints, key);
|
||||
if selected_endpoints.is_empty() {
|
||||
if let Some(models) = preset_models_for_provider(&provider.provider_type) {
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: aggregate_models_for_cache(&models),
|
||||
error: None,
|
||||
from_cache: false,
|
||||
has_success: true,
|
||||
});
|
||||
}
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL.to_string()),
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
}
|
||||
|
||||
let mut transports = Vec::new();
|
||||
let mut all_errors = Vec::new();
|
||||
for endpoint in selected_endpoints {
|
||||
let Some(transport) = state
|
||||
.app()
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
.await?
|
||||
else {
|
||||
all_errors.push(format!(
|
||||
"{} transport snapshot unavailable",
|
||||
endpoint.api_format.trim()
|
||||
));
|
||||
continue;
|
||||
};
|
||||
transports.push(transport);
|
||||
}
|
||||
|
||||
if transports.is_empty() {
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(all_errors.join("; ")),
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
}
|
||||
|
||||
let outcome = match fetch_models_from_transports(state.app(), &transports).await {
|
||||
Ok(outcome) => outcome,
|
||||
Err(err) => {
|
||||
all_errors.push(err);
|
||||
if let Some(fallback) = provider_query_codex_preset_fallback(provider) {
|
||||
return Ok(fallback);
|
||||
}
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(all_errors.join("; ")),
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
all_errors.extend(outcome.errors);
|
||||
let unique_models = aggregate_models_for_cache(&outcome.cached_models);
|
||||
if outcome.has_success && !unique_models.is_empty() {
|
||||
<AppState as ModelFetchRuntimeState>::write_upstream_models_cache(
|
||||
state.app(),
|
||||
&provider.id,
|
||||
&key.id,
|
||||
&unique_models,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
if unique_models.is_empty() && !all_errors.is_empty() {
|
||||
if let Some(fallback) = provider_query_codex_preset_fallback(provider) {
|
||||
return Ok(fallback);
|
||||
}
|
||||
}
|
||||
|
||||
let mut error = if all_errors.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(all_errors.join("; "))
|
||||
};
|
||||
if unique_models.is_empty() && error.is_none() {
|
||||
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_ENDPOINT_DETAIL.to_string());
|
||||
}
|
||||
|
||||
Ok(ProviderQueryKeyFetchResult {
|
||||
models: unique_models,
|
||||
error,
|
||||
from_cache: false,
|
||||
has_success: outcome.has_success,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_provider_query_models_response(
|
||||
state: &AdminAppState<'_>,
|
||||
payload: &serde_json::Value,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let Some(provider_id) = provider_query_extract_provider_id(payload) else {
|
||||
return Ok(build_admin_provider_query_bad_request_response(
|
||||
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
|
||||
));
|
||||
};
|
||||
|
||||
let Some(provider) = state
|
||||
.app()
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.find(|item| item.id == provider_id)
|
||||
else {
|
||||
return Ok(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_PROVIDER_NOT_FOUND_DETAIL,
|
||||
));
|
||||
};
|
||||
|
||||
let provider_ids = vec![provider.id.clone()];
|
||||
let endpoints = state
|
||||
.app()
|
||||
.list_provider_catalog_endpoints_by_provider_ids(&provider_ids)
|
||||
.await?;
|
||||
let keys = state
|
||||
.app()
|
||||
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
|
||||
.await?;
|
||||
let force_refresh = provider_query_extract_force_refresh(payload);
|
||||
|
||||
if let Some(api_key_id) = provider_query_extract_api_key_id(payload) {
|
||||
let Some(selected_key) = keys.iter().find(|key| key.id == api_key_id) else {
|
||||
return Ok(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL,
|
||||
));
|
||||
};
|
||||
|
||||
let result = provider_query_fetch_models_for_key(
|
||||
state,
|
||||
&provider,
|
||||
&endpoints,
|
||||
selected_key,
|
||||
force_refresh,
|
||||
)
|
||||
.await?;
|
||||
let success = !result.models.is_empty();
|
||||
return Ok(Json(json!({
|
||||
"success": success,
|
||||
"data": {
|
||||
"models": result.models,
|
||||
"error": result.error,
|
||||
"from_cache": result.from_cache,
|
||||
},
|
||||
"provider": provider_query_provider_payload(&provider),
|
||||
}))
|
||||
.into_response());
|
||||
}
|
||||
|
||||
let active_keys = keys
|
||||
.into_iter()
|
||||
.filter(|key| key.is_active)
|
||||
.collect::<Vec<_>>();
|
||||
if active_keys.is_empty() {
|
||||
return Ok(build_admin_provider_query_bad_request_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
));
|
||||
}
|
||||
let active_key_count = active_keys.len();
|
||||
|
||||
if provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity")
|
||||
&& !force_refresh
|
||||
{
|
||||
if let Some(models) = provider_query_read_provider_cached_models(state, &provider.id).await
|
||||
{
|
||||
return Ok(Json(json!({
|
||||
"success": !models.is_empty(),
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": serde_json::Value::Null,
|
||||
"from_cache": true,
|
||||
"keys_total": active_key_count,
|
||||
"keys_cached": active_key_count,
|
||||
"keys_fetched": 0,
|
||||
},
|
||||
"provider": provider_query_provider_payload(&provider),
|
||||
}))
|
||||
.into_response());
|
||||
}
|
||||
}
|
||||
|
||||
let ordered_keys = if provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity")
|
||||
{
|
||||
provider_query_sort_antigravity_keys(state, &provider, &endpoints, active_keys).await?
|
||||
} else {
|
||||
active_keys
|
||||
};
|
||||
|
||||
let mut all_models = Vec::new();
|
||||
let mut all_errors = Vec::new();
|
||||
let mut cache_hit_count = 0usize;
|
||||
let mut fetch_count = 0usize;
|
||||
for key in &ordered_keys {
|
||||
let result =
|
||||
provider_query_fetch_models_for_key(state, &provider, &endpoints, key, force_refresh)
|
||||
.await?;
|
||||
all_models.extend(result.models);
|
||||
if let Some(error) = result.error {
|
||||
all_errors.push(format!(
|
||||
"Key {}: {}",
|
||||
provider_query_key_display_name(key),
|
||||
error
|
||||
));
|
||||
}
|
||||
if result.from_cache {
|
||||
cache_hit_count += 1;
|
||||
} else {
|
||||
fetch_count += 1;
|
||||
}
|
||||
if provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity")
|
||||
&& result.has_success
|
||||
{
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
let models = aggregate_models_for_cache(&all_models);
|
||||
if provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity")
|
||||
&& !models.is_empty()
|
||||
{
|
||||
provider_query_write_provider_cached_models(state, &provider.id, &models).await;
|
||||
}
|
||||
let success = !models.is_empty();
|
||||
let mut error = if all_errors.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(all_errors.join("; "))
|
||||
};
|
||||
if !success && error.is_none() {
|
||||
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL.to_string());
|
||||
}
|
||||
|
||||
Ok(Json(json!({
|
||||
"success": success,
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": error,
|
||||
"from_cache": fetch_count == 0 && cache_hit_count > 0,
|
||||
"keys_total": active_key_count,
|
||||
"keys_cached": cache_hit_count,
|
||||
"keys_fetched": fetch_count,
|
||||
},
|
||||
"provider": provider_query_provider_payload(&provider),
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,268 @@
|
||||
use super::DEFAULT_PROVIDER_QUERY_TEST_MESSAGE;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::provider_transport::antigravity::{
|
||||
classify_local_antigravity_request_support, AntigravityEnvelopeRequestType,
|
||||
AntigravityRequestSideSupport, AntigravityRequestSideUnsupportedReason,
|
||||
};
|
||||
use crate::provider_transport::kiro::supports_local_kiro_request_transport_with_network;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(super) enum ProviderQueryTestAdapter {
|
||||
Standard,
|
||||
Kiro,
|
||||
OpenAiImage,
|
||||
Antigravity,
|
||||
}
|
||||
pub(super) fn provider_query_unsupported_test_api_format_message(api_format: &str) -> String {
|
||||
let api_format = api_format.trim();
|
||||
if api_format.is_empty() {
|
||||
"Rust local provider-query model test does not support an empty endpoint format".to_string()
|
||||
} else {
|
||||
format!(
|
||||
"Rust local provider-query model test does not support endpoint format {api_format}"
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_standard_test_client_api_format(
|
||||
provider_api_format: &str,
|
||||
) -> &'static str {
|
||||
let normalized_api_format = crate::ai_serving::normalize_api_format_alias(provider_api_format);
|
||||
if aether_ai_formats::is_embedding_api_format(&normalized_api_format) {
|
||||
"openai:embedding"
|
||||
} else if aether_ai_formats::is_rerank_api_format(&normalized_api_format) {
|
||||
"openai:rerank"
|
||||
} else {
|
||||
"openai:chat"
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_standard_test_unsupported_reason(
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> String {
|
||||
let normalized_api_format = crate::ai_serving::normalize_api_format_alias(api_format);
|
||||
let reason = match normalized_api_format.as_str() {
|
||||
"openai:chat" => {
|
||||
crate::provider_transport::policy::local_openai_chat_transport_unsupported_reason(
|
||||
transport,
|
||||
)
|
||||
}
|
||||
"openai:responses"
|
||||
| "openai:responses:compact"
|
||||
| "claude:messages"
|
||||
| "openai:embedding"
|
||||
| "jina:embedding"
|
||||
| "doubao:embedding"
|
||||
| "openai:rerank"
|
||||
| "jina:rerank" => {
|
||||
crate::provider_transport::policy::local_standard_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
api_format,
|
||||
)
|
||||
}
|
||||
"gemini:generate_content"
|
||||
if crate::provider_transport::is_vertex_api_key_transport_context(transport) =>
|
||||
{
|
||||
aether_provider_transport::vertex::local_vertex_api_key_gemini_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
)
|
||||
}
|
||||
"gemini:generate_content" | "gemini:embedding" => {
|
||||
crate::provider_transport::policy::local_gemini_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
api_format,
|
||||
)
|
||||
}
|
||||
_ => Some("transport_api_format_mismatch"),
|
||||
};
|
||||
|
||||
match reason {
|
||||
Some(reason) => format!(
|
||||
"{} ({reason})",
|
||||
provider_query_unsupported_test_api_format_message(api_format)
|
||||
),
|
||||
None => provider_query_unsupported_test_api_format_message(api_format),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_antigravity_unsupported_reason(
|
||||
reason: AntigravityRequestSideUnsupportedReason,
|
||||
) -> &'static str {
|
||||
match reason {
|
||||
AntigravityRequestSideUnsupportedReason::InactiveTransport => "transport_inactive",
|
||||
AntigravityRequestSideUnsupportedReason::WrongProviderType => {
|
||||
"transport_provider_type_unsupported"
|
||||
}
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedApiFormat => {
|
||||
"transport_api_format_mismatch"
|
||||
}
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedCustomPath => {
|
||||
"transport_custom_path_unsupported"
|
||||
}
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedHeaderRules => {
|
||||
"transport_header_rules_unsupported"
|
||||
}
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedBodyRules => {
|
||||
"transport_body_rules_unsupported"
|
||||
}
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedNetworkConfig => {
|
||||
"transport_network_config_unsupported"
|
||||
}
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedAuth(_) => {
|
||||
"transport_antigravity_auth_unsupported"
|
||||
}
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedEnvelope(_) => {
|
||||
"transport_antigravity_envelope_unsupported"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_antigravity_test_unsupported_reason(
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
request_body: &Value,
|
||||
) -> Option<&'static str> {
|
||||
match classify_local_antigravity_request_support(
|
||||
transport,
|
||||
request_body,
|
||||
AntigravityEnvelopeRequestType::EndpointTest,
|
||||
) {
|
||||
AntigravityRequestSideSupport::Supported(_) => None,
|
||||
AntigravityRequestSideSupport::Unsupported(reason) => {
|
||||
Some(provider_query_antigravity_unsupported_reason(reason))
|
||||
}
|
||||
}
|
||||
}
|
||||
pub(super) fn provider_query_normalize_api_format_alias(value: &str) -> String {
|
||||
crate::ai_serving::normalize_api_format_alias(value)
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_test_adapter_for_provider_api_format(
|
||||
provider_type: &str,
|
||||
api_format: &str,
|
||||
) -> Option<ProviderQueryTestAdapter> {
|
||||
if provider_type.trim().eq_ignore_ascii_case("kiro") {
|
||||
return Some(ProviderQueryTestAdapter::Kiro);
|
||||
}
|
||||
|
||||
let normalized_api_format = provider_query_normalize_api_format_alias(api_format);
|
||||
if normalized_api_format == "openai:image" {
|
||||
return Some(ProviderQueryTestAdapter::OpenAiImage);
|
||||
}
|
||||
if provider_type.trim().eq_ignore_ascii_case("antigravity")
|
||||
&& normalized_api_format == "gemini:generate_content"
|
||||
{
|
||||
return Some(ProviderQueryTestAdapter::Antigravity);
|
||||
}
|
||||
if matches!(
|
||||
normalized_api_format.as_str(),
|
||||
"openai:chat"
|
||||
| "openai:responses"
|
||||
| "openai:responses:compact"
|
||||
| "claude:messages"
|
||||
| "gemini:generate_content"
|
||||
| "openai:embedding"
|
||||
| "gemini:embedding"
|
||||
| "jina:embedding"
|
||||
| "doubao:embedding"
|
||||
| "openai:rerank"
|
||||
| "jina:rerank"
|
||||
) {
|
||||
return Some(ProviderQueryTestAdapter::Standard);
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_model_test_endpoint_priority(
|
||||
provider_type: &str,
|
||||
api_format: &str,
|
||||
) -> Option<u8> {
|
||||
let normalized_api_format = provider_query_normalize_api_format_alias(api_format);
|
||||
match provider_query_test_adapter_for_provider_api_format(provider_type, api_format)? {
|
||||
ProviderQueryTestAdapter::Kiro => Some(0),
|
||||
ProviderQueryTestAdapter::Antigravity => Some(1),
|
||||
ProviderQueryTestAdapter::OpenAiImage => Some(2),
|
||||
ProviderQueryTestAdapter::Standard => {
|
||||
if matches!(
|
||||
normalized_api_format.as_str(),
|
||||
"openai:chat" | "claude:messages" | "gemini:generate_content"
|
||||
) {
|
||||
Some(0)
|
||||
} else {
|
||||
Some(1)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_default_antigravity_endpoint_test_body() -> Value {
|
||||
json!({
|
||||
"contents": [{
|
||||
"role": "user",
|
||||
"parts": [{
|
||||
"text": DEFAULT_PROVIDER_QUERY_TEST_MESSAGE
|
||||
}]
|
||||
}]
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_transport_supports_model_test_execution(
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
match provider_query_test_adapter_for_provider_api_format(
|
||||
transport.provider.provider_type.as_str(),
|
||||
api_format,
|
||||
) {
|
||||
Some(ProviderQueryTestAdapter::Kiro) => {
|
||||
supports_local_kiro_request_transport_with_network(transport)
|
||||
}
|
||||
Some(ProviderQueryTestAdapter::OpenAiImage) => {
|
||||
crate::provider_transport::openai_image_transport_unsupported_reason(
|
||||
transport,
|
||||
"openai:image",
|
||||
)
|
||||
.is_none()
|
||||
}
|
||||
Some(ProviderQueryTestAdapter::Antigravity) => {
|
||||
provider_query_antigravity_test_unsupported_reason(
|
||||
transport,
|
||||
&provider_query_default_antigravity_endpoint_test_body(),
|
||||
)
|
||||
.is_none()
|
||||
}
|
||||
Some(ProviderQueryTestAdapter::Standard) => match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"openai:chat" => {
|
||||
crate::provider_transport::policy::supports_local_openai_chat_transport(transport)
|
||||
}
|
||||
"openai:responses"
|
||||
| "openai:responses:compact"
|
||||
| "openai:embedding"
|
||||
| "jina:embedding"
|
||||
| "doubao:embedding"
|
||||
| "openai:rerank"
|
||||
| "jina:rerank" => {
|
||||
crate::provider_transport::policy::supports_local_standard_transport_with_network(
|
||||
transport, api_format,
|
||||
)
|
||||
}
|
||||
"claude:messages" => {
|
||||
crate::provider_transport::policy::supports_local_standard_transport_with_network(
|
||||
transport, api_format,
|
||||
)
|
||||
}
|
||||
"gemini:generate_content" | "gemini:embedding" => {
|
||||
if crate::provider_transport::is_vertex_transport_context(transport) {
|
||||
aether_provider_transport::vertex::supports_local_vertex_gemini_transport_with_network(transport)
|
||||
} else {
|
||||
state.supports_local_gemini_transport_with_network(transport, api_format)
|
||||
}
|
||||
}
|
||||
_ => false,
|
||||
},
|
||||
None => false,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,369 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||
};
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
AdminProviderModelListQuery, StoredAdminProviderModel,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogEndpoint;
|
||||
use serde_json::Value;
|
||||
|
||||
pub(super) async fn provider_query_resolve_global_effective_model(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
requested_model: &str,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
) -> Result<String, GatewayError> {
|
||||
let models = state
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider_id.to_string(),
|
||||
is_active: Some(true),
|
||||
offset: 0,
|
||||
limit: 1024,
|
||||
})
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
|
||||
for model in models
|
||||
.into_iter()
|
||||
.filter(|model| model.is_available && model.is_active)
|
||||
{
|
||||
let mappings =
|
||||
provider_query_parse_provider_model_mappings(model.provider_model_mappings.as_ref())?;
|
||||
if !provider_query_admin_model_matches_requested_model(
|
||||
&model,
|
||||
mappings.as_deref(),
|
||||
requested_model,
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
|
||||
let row = provider_query_admin_model_selection_row(&model, endpoint, mappings);
|
||||
return Ok(aether_scheduler_core::select_provider_model_name(
|
||||
&row,
|
||||
&endpoint.api_format,
|
||||
));
|
||||
}
|
||||
|
||||
Ok(requested_model.to_string())
|
||||
}
|
||||
|
||||
pub(super) async fn provider_query_resolve_explicit_mapped_effective_model(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
_provider_type: &str,
|
||||
requested_model: &str,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
mapped_model: &str,
|
||||
) -> Result<Option<String>, GatewayError> {
|
||||
let models = state
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider_id.to_string(),
|
||||
is_active: Some(true),
|
||||
offset: 0,
|
||||
limit: 1024,
|
||||
})
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
|
||||
for model in models
|
||||
.into_iter()
|
||||
.filter(|model| model.is_available && model.is_active)
|
||||
{
|
||||
let Some(mappings) =
|
||||
provider_query_parse_provider_model_mappings(model.provider_model_mappings.as_ref())?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if !provider_query_admin_model_matches_requested_model(
|
||||
&model,
|
||||
Some(&mappings),
|
||||
requested_model,
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
|
||||
if mappings.iter().any(|mapping| {
|
||||
mapping.name.eq_ignore_ascii_case(mapped_model)
|
||||
&& provider_query_model_mapping_matches_endpoint(mapping, endpoint)
|
||||
}) {
|
||||
return Ok(Some(mapped_model.to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn provider_query_model_mapping_matches_endpoint(
|
||||
mapping: &StoredProviderModelMapping,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
) -> bool {
|
||||
let api_format_matches = mapping.api_formats.as_ref().is_none_or(|api_formats| {
|
||||
api_formats.iter().any(|value| {
|
||||
aether_scheduler_core::normalize_api_format(value)
|
||||
== aether_scheduler_core::normalize_api_format(&endpoint.api_format)
|
||||
})
|
||||
});
|
||||
if !api_format_matches {
|
||||
return false;
|
||||
}
|
||||
|
||||
mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
|
||||
endpoint_ids
|
||||
.iter()
|
||||
.any(|endpoint_id| endpoint_id == &endpoint.id)
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_admin_model_matches_requested_model(
|
||||
model: &StoredAdminProviderModel,
|
||||
mappings: Option<&[StoredProviderModelMapping]>,
|
||||
requested_model: &str,
|
||||
) -> bool {
|
||||
model
|
||||
.global_model_name
|
||||
.as_deref()
|
||||
.is_some_and(|value| value.eq_ignore_ascii_case(requested_model))
|
||||
|| model
|
||||
.provider_model_name
|
||||
.eq_ignore_ascii_case(requested_model)
|
||||
|| mappings.is_some_and(|mappings| {
|
||||
mappings
|
||||
.iter()
|
||||
.any(|mapping| mapping.name.eq_ignore_ascii_case(requested_model))
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_admin_model_selection_row(
|
||||
model: &StoredAdminProviderModel,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
mappings: Option<Vec<StoredProviderModelMapping>>,
|
||||
) -> StoredMinimalCandidateSelectionRow {
|
||||
StoredMinimalCandidateSelectionRow {
|
||||
provider_id: model.provider_id.clone(),
|
||||
provider_name: String::new(),
|
||||
provider_type: String::new(),
|
||||
provider_priority: 0,
|
||||
provider_is_active: true,
|
||||
endpoint_id: endpoint.id.clone(),
|
||||
endpoint_api_format: endpoint.api_format.clone(),
|
||||
endpoint_api_family: endpoint.api_family.clone(),
|
||||
endpoint_kind: endpoint.endpoint_kind.clone(),
|
||||
endpoint_is_active: endpoint.is_active,
|
||||
key_id: String::new(),
|
||||
key_name: String::new(),
|
||||
key_auth_type: String::new(),
|
||||
key_is_active: true,
|
||||
key_api_formats: None,
|
||||
key_allowed_models: None,
|
||||
key_capabilities: None,
|
||||
key_internal_priority: 0,
|
||||
key_global_priority_by_format: None,
|
||||
model_id: model.id.clone(),
|
||||
global_model_id: model.global_model_id.clone(),
|
||||
global_model_name: model.global_model_name.clone().unwrap_or_default(),
|
||||
global_model_mappings: None,
|
||||
global_model_supports_streaming: None,
|
||||
model_provider_model_name: model.provider_model_name.clone(),
|
||||
model_provider_model_mappings: mappings,
|
||||
model_supports_streaming: model.supports_streaming,
|
||||
model_is_active: model.is_active,
|
||||
model_is_available: model.is_available,
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_parse_provider_model_mappings(
|
||||
value: Option<&Value>,
|
||||
) -> Result<Option<Vec<StoredProviderModelMapping>>, GatewayError> {
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
provider_query_parse_provider_model_mappings_value(value)
|
||||
}
|
||||
|
||||
fn provider_query_parse_provider_model_mappings_value(
|
||||
value: &Value,
|
||||
) -> Result<Option<Vec<StoredProviderModelMapping>>, GatewayError> {
|
||||
match value {
|
||||
Value::Null => Ok(None),
|
||||
Value::Array(items) => provider_query_parse_provider_model_mappings_array(items),
|
||||
Value::Object(object) => provider_query_parse_provider_model_mapping_object(object)
|
||||
.map(|mapping| Some(vec![mapping])),
|
||||
Value::String(raw) => provider_query_parse_embedded_provider_model_mappings(raw),
|
||||
_ => Err(GatewayError::Internal(
|
||||
"models.provider_model_mappings is not a JSON array".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_parse_embedded_provider_model_mappings(
|
||||
raw: &str,
|
||||
) -> Result<Option<Vec<StoredProviderModelMapping>>, GatewayError> {
|
||||
let raw = raw.trim();
|
||||
if raw.is_empty() || raw.eq_ignore_ascii_case("null") {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
if let Ok(decoded) = serde_json::from_str::<Value>(raw) {
|
||||
return provider_query_parse_provider_model_mappings_value(&decoded);
|
||||
}
|
||||
|
||||
Ok(Some(vec![StoredProviderModelMapping {
|
||||
name: raw.to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
endpoint_ids: None,
|
||||
}]))
|
||||
}
|
||||
|
||||
fn provider_query_parse_provider_model_mappings_array(
|
||||
items: &[Value],
|
||||
) -> Result<Option<Vec<StoredProviderModelMapping>>, GatewayError> {
|
||||
let mut mappings = Vec::with_capacity(items.len());
|
||||
for item in items {
|
||||
match item {
|
||||
Value::Object(object) => {
|
||||
if let Some(mapping) =
|
||||
provider_query_parse_provider_model_mapping_object_lenient(object)?
|
||||
{
|
||||
mappings.push(mapping);
|
||||
}
|
||||
}
|
||||
Value::String(raw) => {
|
||||
let raw = raw.trim();
|
||||
if !raw.is_empty() {
|
||||
mappings.push(StoredProviderModelMapping {
|
||||
name: raw.to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
endpoint_ids: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
Value::Null => {}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
if mappings.is_empty() {
|
||||
Ok(None)
|
||||
} else {
|
||||
Ok(Some(mappings))
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_parse_provider_model_mapping_object(
|
||||
object: &serde_json::Map<String, Value>,
|
||||
) -> Result<StoredProviderModelMapping, GatewayError> {
|
||||
provider_query_parse_provider_model_mapping_object_lenient(object)?.ok_or_else(|| {
|
||||
GatewayError::Internal(
|
||||
"models.provider_model_mappings item is missing a valid name".to_string(),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_parse_provider_model_mapping_object_lenient(
|
||||
object: &serde_json::Map<String, Value>,
|
||||
) -> Result<Option<StoredProviderModelMapping>, GatewayError> {
|
||||
let Some(name) = object
|
||||
.get("name")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let priority = object
|
||||
.get("priority")
|
||||
.and_then(Value::as_i64)
|
||||
.unwrap_or(1)
|
||||
.max(1);
|
||||
let api_formats = provider_query_parse_mapping_string_list(
|
||||
object.get("api_formats"),
|
||||
"models.provider_model_mappings.api_formats",
|
||||
)?
|
||||
.map(|formats| {
|
||||
formats
|
||||
.into_iter()
|
||||
.map(|value| aether_ai_formats::normalize_api_format_alias(&value))
|
||||
.collect()
|
||||
});
|
||||
let endpoint_ids = provider_query_parse_mapping_string_list(
|
||||
object.get("endpoint_ids"),
|
||||
"models.provider_model_mappings.endpoint_ids",
|
||||
)?;
|
||||
|
||||
Ok(Some(StoredProviderModelMapping {
|
||||
name: name.to_string(),
|
||||
priority: i32::try_from(priority).map_err(|_| {
|
||||
GatewayError::Internal(format!(
|
||||
"invalid models.provider_model_mappings.priority: {priority}"
|
||||
))
|
||||
})?,
|
||||
api_formats,
|
||||
endpoint_ids,
|
||||
}))
|
||||
}
|
||||
|
||||
fn provider_query_parse_mapping_string_list(
|
||||
value: Option<&Value>,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, GatewayError> {
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
provider_query_parse_mapping_string_list_value(value, field_name)
|
||||
}
|
||||
|
||||
fn provider_query_parse_mapping_string_list_value(
|
||||
value: &Value,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, GatewayError> {
|
||||
match value {
|
||||
Value::Null => Ok(None),
|
||||
Value::Array(items) => {
|
||||
provider_query_parse_mapping_string_list_array(items, field_name).map(Some)
|
||||
}
|
||||
Value::String(raw) => provider_query_parse_embedded_mapping_string_list(raw, field_name),
|
||||
_ => Err(GatewayError::Internal(format!(
|
||||
"{field_name} is not a JSON array"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_parse_embedded_mapping_string_list(
|
||||
raw: &str,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, GatewayError> {
|
||||
let raw = raw.trim();
|
||||
if raw.is_empty() || raw.eq_ignore_ascii_case("null") {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
if let Ok(decoded) = serde_json::from_str::<Value>(raw) {
|
||||
return provider_query_parse_mapping_string_list_value(&decoded, field_name);
|
||||
}
|
||||
|
||||
Ok(Some(vec![raw.to_string()]))
|
||||
}
|
||||
|
||||
fn provider_query_parse_mapping_string_list_array(
|
||||
items: &[Value],
|
||||
field_name: &str,
|
||||
) -> Result<Vec<String>, GatewayError> {
|
||||
let mut parsed = Vec::with_capacity(items.len());
|
||||
for item in items {
|
||||
let Some(item) = item.as_str() else {
|
||||
return Err(GatewayError::Internal(format!(
|
||||
"{field_name} contains a non-string item"
|
||||
)));
|
||||
};
|
||||
let item = item.trim();
|
||||
if !item.is_empty() {
|
||||
parsed.push(item.to_string());
|
||||
}
|
||||
}
|
||||
Ok(parsed)
|
||||
}
|
||||
@@ -0,0 +1,135 @@
|
||||
use super::super::provider_query_key_display_name;
|
||||
use super::{ProviderQueryExecutionOutcome, ProviderQueryTestCandidate};
|
||||
use serde_json::{json, Value};
|
||||
|
||||
pub(super) fn provider_query_test_attempt_payload(
|
||||
candidate_index: usize,
|
||||
candidate: &ProviderQueryTestCandidate,
|
||||
execution: &ProviderQueryExecutionOutcome,
|
||||
) -> Value {
|
||||
json!({
|
||||
"candidate_index": candidate_index,
|
||||
"retry_index": 0,
|
||||
"endpoint_api_format": candidate.endpoint.api_format,
|
||||
"endpoint_base_url": candidate.endpoint.base_url,
|
||||
"key_name": provider_query_key_display_name(&candidate.key),
|
||||
"key_id": candidate.key.id,
|
||||
"auth_type": candidate.key.auth_type,
|
||||
"effective_model": candidate.effective_model,
|
||||
"status": execution.status,
|
||||
"skip_reason": execution.skip_reason,
|
||||
"error_message": execution.error_message,
|
||||
"status_code": execution.status_code,
|
||||
"latency_ms": execution.latency_ms,
|
||||
"request_url": execution.request_url,
|
||||
"request_headers": execution.request_headers,
|
||||
"request_body": execution.request_body,
|
||||
"response_headers": execution.response_headers,
|
||||
"response_body": execution.response_body,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_candidate_summary_payload(
|
||||
total_candidates: usize,
|
||||
total_attempts: usize,
|
||||
attempts: &[Value],
|
||||
) -> Value {
|
||||
let success_count = attempts
|
||||
.iter()
|
||||
.filter(|attempt| attempt.get("status").and_then(Value::as_str) == Some("success"))
|
||||
.count();
|
||||
let failed_count = attempts
|
||||
.iter()
|
||||
.filter(|attempt| {
|
||||
matches!(
|
||||
attempt.get("status").and_then(Value::as_str),
|
||||
Some("failed") | Some("cancelled") | Some("stream_interrupted")
|
||||
)
|
||||
})
|
||||
.count();
|
||||
let skipped_count = attempts
|
||||
.iter()
|
||||
.filter(|attempt| attempt.get("status").and_then(Value::as_str) == Some("skipped"))
|
||||
.count();
|
||||
let pending_count = attempts
|
||||
.iter()
|
||||
.filter(|attempt| {
|
||||
matches!(
|
||||
attempt.get("status").and_then(Value::as_str),
|
||||
Some("pending") | Some("streaming")
|
||||
)
|
||||
})
|
||||
.count();
|
||||
let available_count = attempts
|
||||
.iter()
|
||||
.filter(|attempt| attempt.get("status").and_then(Value::as_str) == Some("available"))
|
||||
.count();
|
||||
let unused_count = if success_count > 0 {
|
||||
total_candidates.saturating_sub(success_count + failed_count + skipped_count)
|
||||
} else {
|
||||
0
|
||||
};
|
||||
let stop_reason = if total_candidates == 0 {
|
||||
"no_candidate"
|
||||
} else if success_count > 0 {
|
||||
"first_success"
|
||||
} else if total_attempts == 0 && skipped_count > 0 {
|
||||
"all_skipped"
|
||||
} else if failed_count > 0 || skipped_count > 0 {
|
||||
"exhausted"
|
||||
} else {
|
||||
"pending"
|
||||
};
|
||||
let winning_attempt = attempts
|
||||
.iter()
|
||||
.find(|attempt| attempt.get("status").and_then(Value::as_str) == Some("success"));
|
||||
|
||||
json!({
|
||||
"total_candidates": total_candidates,
|
||||
"attempted": total_attempts,
|
||||
"success": success_count,
|
||||
"failed": failed_count,
|
||||
"skipped": skipped_count,
|
||||
"unused": unused_count,
|
||||
"pending": pending_count,
|
||||
"available": available_count,
|
||||
"completed": success_count + failed_count + skipped_count + unused_count,
|
||||
"stop_reason": stop_reason,
|
||||
"winning_candidate_index": winning_attempt
|
||||
.and_then(|attempt| attempt.get("candidate_index"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
"winning_key_name": winning_attempt
|
||||
.and_then(|attempt| attempt.get("key_name"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
"winning_key_id": winning_attempt
|
||||
.and_then(|attempt| attempt.get("key_id"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
"winning_auth_type": winning_attempt
|
||||
.and_then(|attempt| attempt.get("auth_type"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
"winning_effective_model": winning_attempt
|
||||
.and_then(|attempt| attempt.get("effective_model"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
"winning_endpoint_api_format": winning_attempt
|
||||
.and_then(|attempt| attempt.get("endpoint_api_format"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
"winning_endpoint_base_url": winning_attempt
|
||||
.and_then(|attempt| attempt.get("endpoint_base_url"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
"winning_latency_ms": winning_attempt
|
||||
.and_then(|attempt| attempt.get("latency_ms"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
"winning_status_code": winning_attempt
|
||||
.and_then(|attempt| attempt.get("status_code"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,362 @@
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn provider_query_test_request_body_preserves_custom_model() {
|
||||
let payload = json!({
|
||||
"request_body": {
|
||||
"model": "custom-upstream-model",
|
||||
"messages": []
|
||||
}
|
||||
});
|
||||
|
||||
let body = provider_query_build_test_request_body(&payload, "fallback-model");
|
||||
|
||||
assert_eq!(body["model"], json!("custom-upstream-model"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_test_request_body_defaults_missing_model() {
|
||||
let payload = json!({
|
||||
"request_body": {
|
||||
"messages": []
|
||||
}
|
||||
});
|
||||
|
||||
let body = provider_query_build_test_request_body(&payload, "fallback-model");
|
||||
|
||||
assert_eq!(body["model"], json!("fallback-model"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_failover_request_body_overrides_custom_model() {
|
||||
let payload = json!({
|
||||
"request_body": {
|
||||
"model": "custom-upstream-model",
|
||||
"messages": []
|
||||
}
|
||||
});
|
||||
|
||||
let body = provider_query_build_test_request_body_for_route(
|
||||
&payload,
|
||||
"failover-model",
|
||||
"/api/admin/provider-query/test-model-failover",
|
||||
);
|
||||
|
||||
assert_eq!(body["model"], json!("failover-model"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_failover_request_body_uses_explicit_mapped_model() {
|
||||
let payload = json!({
|
||||
"mapped_model_name": "upstream-mapped-model",
|
||||
"request_body": {
|
||||
"model": "original-selected-model",
|
||||
"messages": []
|
||||
}
|
||||
});
|
||||
|
||||
let body = provider_query_build_test_request_body_for_route(
|
||||
&payload,
|
||||
"upstream-mapped-model",
|
||||
"/api/admin/provider-query/test-model-failover",
|
||||
);
|
||||
|
||||
assert_eq!(body["model"], json!("upstream-mapped-model"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_request_body_model_uses_non_empty_string_only() {
|
||||
let custom = json!({ "model": " custom-model " });
|
||||
let blank = json!({ "model": " " });
|
||||
let non_string = json!({ "model": 123 });
|
||||
|
||||
assert_eq!(
|
||||
provider_query_request_body_model(&custom, "fallback-model"),
|
||||
"custom-model"
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_request_body_model(&blank, "fallback-model"),
|
||||
"fallback-model"
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_request_body_model(&non_string, "fallback-model"),
|
||||
"fallback-model"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_standard_test_resolves_codex_responses_upstream_streaming() {
|
||||
assert!(provider_query_resolve_standard_test_upstream_is_stream(
|
||||
None,
|
||||
"codex",
|
||||
"openai:responses",
|
||||
));
|
||||
assert!(provider_query_resolve_standard_test_upstream_is_stream(
|
||||
Some(&json!({"upstream_stream_policy": "force_non_stream"})),
|
||||
"codex",
|
||||
"openai:responses",
|
||||
));
|
||||
assert!(!provider_query_resolve_standard_test_upstream_is_stream(
|
||||
None,
|
||||
"codex",
|
||||
"openai:responses:compact",
|
||||
));
|
||||
assert!(!provider_query_resolve_standard_test_upstream_is_stream(
|
||||
None,
|
||||
"custom",
|
||||
"openai:responses",
|
||||
));
|
||||
assert!(provider_query_resolve_standard_test_upstream_is_stream(
|
||||
Some(&json!({"upstream_stream_policy": "force_stream"})),
|
||||
"custom",
|
||||
"openai:responses",
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_standard_test_reenforces_upstream_stream_body_field() {
|
||||
let endpoint_config = json!({"upstream_stream_policy": "force_stream"});
|
||||
let mut body = json!({"model": "gpt-5", "input": "hello", "stream": false});
|
||||
let upstream_is_stream = provider_query_resolve_standard_test_upstream_is_stream(
|
||||
Some(&endpoint_config),
|
||||
"codex",
|
||||
"openai:responses",
|
||||
);
|
||||
let require_body_stream_field =
|
||||
provider_query_request_requires_body_stream_field(&body, Some(&endpoint_config));
|
||||
|
||||
crate::ai_serving::enforce_request_body_stream_field(
|
||||
&mut body,
|
||||
"openai:responses",
|
||||
upstream_is_stream,
|
||||
require_body_stream_field,
|
||||
);
|
||||
|
||||
assert_eq!(body["stream"], json!(true));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_standard_test_aggregates_responses_stream_body() {
|
||||
let stream_body = concat!(
|
||||
"event: response.created\n",
|
||||
"data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_123\",\"object\":\"response\",\"model\":\"gpt-5.4-mini\",\"status\":\"in_progress\",\"output\":[]}}\n\n",
|
||||
"event: response.output_text.delta\n",
|
||||
"data: {\"type\":\"response.output_text.delta\",\"output_index\":0,\"content_index\":0,\"delta\":\"Hello\"}\n\n",
|
||||
"event: response.completed\n",
|
||||
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_123\",\"object\":\"response\",\"model\":\"gpt-5.4-mini\",\"status\":\"completed\",\"output\":[]}}\n\n",
|
||||
);
|
||||
let result = aether_contracts::ExecutionResult {
|
||||
request_id: "provider-test".to_string(),
|
||||
candidate_id: Some("candidate-0".to_string()),
|
||||
status_code: 200,
|
||||
headers: BTreeMap::new(),
|
||||
body: Some(aether_contracts::ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(
|
||||
base64::engine::general_purpose::STANDARD.encode(stream_body.as_bytes()),
|
||||
),
|
||||
}),
|
||||
telemetry: None,
|
||||
error: None,
|
||||
};
|
||||
|
||||
let body = provider_query_standard_execution_response_body("openai:responses", &result)
|
||||
.expect("stream body should aggregate");
|
||||
|
||||
assert_eq!(body["model"], json!("gpt-5.4-mini"));
|
||||
assert_eq!(body["output"][0]["content"][0]["text"], json!("Hello"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_test_adapter_routes_fixed_provider_endpoint_types() {
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("custom", "openai:chat"),
|
||||
Some(ProviderQueryTestAdapter::Standard)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("codex", "openai:responses"),
|
||||
Some(ProviderQueryTestAdapter::Standard)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("codex", "openai:responses:compact"),
|
||||
Some(ProviderQueryTestAdapter::Standard)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("chatgpt_web", "openai:image"),
|
||||
Some(ProviderQueryTestAdapter::OpenAiImage)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("kiro", "claude:messages"),
|
||||
Some(ProviderQueryTestAdapter::Kiro)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format(
|
||||
"gemini_cli",
|
||||
"gemini:generate_content"
|
||||
),
|
||||
Some(ProviderQueryTestAdapter::Standard)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format(
|
||||
"antigravity",
|
||||
"gemini:generate_content"
|
||||
),
|
||||
Some(ProviderQueryTestAdapter::Antigravity)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("custom", "openai:embedding"),
|
||||
Some(ProviderQueryTestAdapter::Standard)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("custom", "gemini:embedding"),
|
||||
Some(ProviderQueryTestAdapter::Standard)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("jina", "jina:rerank"),
|
||||
Some(ProviderQueryTestAdapter::Standard)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("custom", "openai:video"),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("custom", "gemini:video"),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("gemini", "gemini:files"),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_endpoint_priority_prefers_text_before_cli_and_image() {
|
||||
assert_eq!(
|
||||
provider_query_model_test_endpoint_priority("custom", "openai:chat"),
|
||||
Some(0)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_model_test_endpoint_priority("codex", "openai:responses:compact"),
|
||||
Some(1)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_model_test_endpoint_priority("chatgpt_web", "openai:image"),
|
||||
Some(2)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_model_test_endpoint_priority("antigravity", "gemini:generate_content"),
|
||||
Some(1)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_candidate_summary_marks_unused_after_first_success() {
|
||||
let attempts = vec![json!({
|
||||
"candidate_index": 0,
|
||||
"key_name": "winning-key",
|
||||
"key_id": "key-1",
|
||||
"auth_type": "api_key",
|
||||
"effective_model": "claude-haiku",
|
||||
"endpoint_api_format": "claude:messages",
|
||||
"endpoint_base_url": "https://api.example",
|
||||
"status": "success",
|
||||
"latency_ms": 123,
|
||||
"status_code": 200
|
||||
})];
|
||||
|
||||
let summary = provider_query_candidate_summary_payload(3, 1, &attempts);
|
||||
|
||||
assert_eq!(summary["total_candidates"], json!(3));
|
||||
assert_eq!(summary["attempted"], json!(1));
|
||||
assert_eq!(summary["success"], json!(1));
|
||||
assert_eq!(summary["unused"], json!(2));
|
||||
assert_eq!(summary["stop_reason"], json!("first_success"));
|
||||
assert_eq!(summary["winning_key_name"], json!("winning-key"));
|
||||
assert_eq!(
|
||||
summary["winning_endpoint_api_format"],
|
||||
json!("claude:messages")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_candidate_summary_reports_skipped_exhaustion() {
|
||||
let attempts = vec![json!({
|
||||
"candidate_index": 0,
|
||||
"status": "skipped",
|
||||
"skip_reason": "transport_api_format_mismatch"
|
||||
})];
|
||||
|
||||
let summary = provider_query_candidate_summary_payload(1, 0, &attempts);
|
||||
|
||||
assert_eq!(summary["total_candidates"], json!(1));
|
||||
assert_eq!(summary["attempted"], json!(0));
|
||||
assert_eq!(summary["skipped"], json!(1));
|
||||
assert_eq!(summary["unused"], json!(0));
|
||||
assert_eq!(summary["stop_reason"], json!("all_skipped"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_candidate_summary_counts_scheduler_skips_before_success() {
|
||||
let attempts = vec![
|
||||
json!({
|
||||
"candidate_index": 0,
|
||||
"status": "skipped",
|
||||
"skip_reason": "pool_account_exhausted"
|
||||
}),
|
||||
json!({
|
||||
"candidate_index": 1,
|
||||
"key_name": "winning-key",
|
||||
"key_id": "key-1",
|
||||
"auth_type": "oauth",
|
||||
"effective_model": "gpt-5.4-mini",
|
||||
"endpoint_api_format": "openai:responses",
|
||||
"endpoint_base_url": "https://chatgpt.com/backend-api/codex",
|
||||
"status": "success",
|
||||
"latency_ms": 123,
|
||||
"status_code": 200
|
||||
}),
|
||||
];
|
||||
|
||||
let summary = provider_query_candidate_summary_payload(4, 1, &attempts);
|
||||
|
||||
assert_eq!(summary["total_candidates"], json!(4));
|
||||
assert_eq!(summary["attempted"], json!(1));
|
||||
assert_eq!(summary["success"], json!(1));
|
||||
assert_eq!(summary["skipped"], json!(1));
|
||||
assert_eq!(summary["unused"], json!(2));
|
||||
assert_eq!(summary["stop_reason"], json!("first_success"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_image_test_request_body_defaults_generation_prompt() {
|
||||
let payload = json!({"message": "draw a small icon"});
|
||||
|
||||
let body = provider_query_build_openai_image_test_request_body_for_route(
|
||||
&payload,
|
||||
"gpt-image-1",
|
||||
"/api/admin/provider-query/test-model",
|
||||
);
|
||||
|
||||
assert_eq!(body["model"], json!("gpt-image-1"));
|
||||
assert_eq!(body["prompt"], json!("draw a small icon"));
|
||||
assert_eq!(body["stream"], json!(true));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_failover_image_test_request_body_overrides_model() {
|
||||
let payload = json!({
|
||||
"request_body": {
|
||||
"model": "old-image-model",
|
||||
"prompt": "draw"
|
||||
}
|
||||
});
|
||||
|
||||
let body = provider_query_build_openai_image_test_request_body_for_route(
|
||||
&payload,
|
||||
"new-image-model",
|
||||
"/api/admin/provider-query/test-model-failover",
|
||||
);
|
||||
|
||||
assert_eq!(body["model"], json!("new-image-model"));
|
||||
}
|
||||
@@ -148,7 +148,14 @@ fn remap_import_proxy(
|
||||
}
|
||||
|
||||
fn normalize_import_endpoint_format(value: &str) -> Result<String, String> {
|
||||
admin_endpoint_signature_parts(value)
|
||||
let normalized = match value.trim().to_ascii_lowercase().as_str() {
|
||||
"openai:cli" => "openai:responses",
|
||||
"openai:compact" => "openai:responses:compact",
|
||||
"claude:chat" | "claude:cli" => "claude:messages",
|
||||
"gemini:chat" | "gemini:cli" => "gemini:generate_content",
|
||||
_ => value.trim(),
|
||||
};
|
||||
admin_endpoint_signature_parts(normalized)
|
||||
.map(|(signature, _, _)| signature.to_string())
|
||||
.ok_or_else(|| format!("无效的 api_format: {value}"))
|
||||
}
|
||||
@@ -1173,7 +1180,7 @@ impl<'a> AdminAppState<'a> {
|
||||
}
|
||||
AdminImportMergeMode::Overwrite => {
|
||||
let Some((normalized_signature, api_family, endpoint_kind)) =
|
||||
admin_endpoint_signature_parts(&imported_endpoint.api_format)
|
||||
admin_endpoint_signature_parts(&normalized_api_format)
|
||||
else {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"无效的 api_format: {}",
|
||||
@@ -1246,7 +1253,7 @@ impl<'a> AdminAppState<'a> {
|
||||
}
|
||||
|
||||
let Some((normalized_signature, api_family, endpoint_kind)) =
|
||||
admin_endpoint_signature_parts(&imported_endpoint.api_format)
|
||||
admin_endpoint_signature_parts(&normalized_api_format)
|
||||
else {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"无效的 api_format: {}",
|
||||
@@ -2832,7 +2839,9 @@ mod tests {
|
||||
use super::{
|
||||
imported_optional_bool, imported_optional_f64, imported_optional_i32,
|
||||
imported_optional_u64, imported_rfc3339_to_unix_secs, imported_string_list_from_value,
|
||||
normalize_import_endpoint_format, normalize_import_key_formats,
|
||||
normalize_imported_wallet_target, validate_imported_system_users_export_version,
|
||||
ImportedProviderKey,
|
||||
};
|
||||
|
||||
#[test]
|
||||
@@ -2849,6 +2858,59 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn config_import_normalizes_python_cli_api_format_aliases() {
|
||||
for (raw, expected) in [
|
||||
("openai:cli", "openai:responses"),
|
||||
("openai:compact", "openai:responses:compact"),
|
||||
("claude:chat", "claude:messages"),
|
||||
("claude:cli", "claude:messages"),
|
||||
("gemini:chat", "gemini:generate_content"),
|
||||
("gemini:cli", "gemini:generate_content"),
|
||||
] {
|
||||
assert_eq!(normalize_import_endpoint_format(raw).unwrap(), expected);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn config_import_normalizes_key_formats_against_imported_endpoint_aliases() {
|
||||
let endpoint_formats = ["claude:messages", "openai:responses:compact"]
|
||||
.into_iter()
|
||||
.map(ToOwned::to_owned)
|
||||
.collect();
|
||||
let item = ImportedProviderKey {
|
||||
api_key: None,
|
||||
auth_type: None,
|
||||
auth_config: None,
|
||||
name: None,
|
||||
note: None,
|
||||
api_formats: Some(vec!["claude:cli".to_string(), "openai:compact".to_string()]),
|
||||
supported_endpoints: None,
|
||||
rate_multipliers: None,
|
||||
internal_priority: None,
|
||||
global_priority_by_format: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
rpm_limit: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
cache_ttl_minutes: None,
|
||||
max_probe_interval_minutes: None,
|
||||
auto_fetch_models: None,
|
||||
locked_models: None,
|
||||
model_include_patterns: None,
|
||||
model_exclude_patterns: None,
|
||||
is_active: true,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
};
|
||||
|
||||
let (formats, missing) = normalize_import_key_formats(&item, &endpoint_formats);
|
||||
|
||||
assert_eq!(formats, vec!["claude:messages", "openai:responses:compact"]);
|
||||
assert!(missing.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn import_handles_legacy_string_scalars() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -578,10 +578,15 @@ fn admin_provider_query_and_strategy_use_specific_local_owners() {
|
||||
"handlers/admin/provider/query/mod.rs should not retain a generic shared module"
|
||||
);
|
||||
|
||||
for path in [
|
||||
"apps/aether-gateway/src/handlers/admin/provider/query/models.rs",
|
||||
"apps/aether-gateway/src/handlers/admin/provider/query/routes.rs",
|
||||
] {
|
||||
let query_model_owners = read_workspace_module_tree(
|
||||
"apps/aether-gateway/src/handlers/admin/provider/query/models/mod.rs",
|
||||
);
|
||||
assert!(
|
||||
!query_model_owners.contains("super::shared::{"),
|
||||
"handlers/admin/provider/query/models should not depend on a generic query::shared hub"
|
||||
);
|
||||
|
||||
for path in ["apps/aether-gateway/src/handlers/admin/provider/query/routes.rs"] {
|
||||
let contents = read_workspace_file(path);
|
||||
assert!(
|
||||
!contents.contains("super::shared::{"),
|
||||
|
||||
@@ -990,7 +990,7 @@ fn admin_route_adjacent_owners_use_wrapped_state_types() {
|
||||
"apps/aether-gateway/src/handlers/admin/model/global/providers.rs",
|
||||
"apps/aether-gateway/src/handlers/admin/model/write.rs",
|
||||
"apps/aether-gateway/src/handlers/admin/users/lifecycle/support.rs",
|
||||
"apps/aether-gateway/src/handlers/admin/provider/query/models.rs",
|
||||
"apps/aether-gateway/src/handlers/admin/provider/query/models/mod.rs",
|
||||
"apps/aether-gateway/src/handlers/admin/provider/strategy/builders.rs",
|
||||
"apps/aether-gateway/src/handlers/admin/features/video_tasks/builders.rs",
|
||||
"apps/aether-gateway/src/handlers/admin/observability/stats/leaderboard.rs",
|
||||
|
||||
@@ -4959,6 +4959,7 @@ fn retired_api_format_occurrences_are_whitelisted() {
|
||||
|
||||
let allowed_paths = [
|
||||
"apps/aether-gateway/src/handlers/admin/provider/write/normalize.rs",
|
||||
"apps/aether-gateway/src/handlers/admin/request/system/import.rs",
|
||||
"crates/aether-ai-formats/src/formats/id.rs",
|
||||
"crates/aether-ai-formats/src/formats/matrix.rs",
|
||||
"crates/aether-ai-formats/src/formats/registry.rs",
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,5 +1,6 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use aether_contracts::ExecutionPlan;
|
||||
use aether_crypto::{
|
||||
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY,
|
||||
};
|
||||
@@ -22,27 +23,24 @@ use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepo
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::http::HeaderMap;
|
||||
use axum::routing::{any, post};
|
||||
use axum::{extract::Request, Router};
|
||||
use axum::{extract::Request, Json, Router};
|
||||
use http::StatusCode;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use super::super::helpers::{sample_endpoint, sample_key, sample_provider};
|
||||
use super::super::{build_router_with_state, start_server, AppState};
|
||||
use super::super::{
|
||||
build_router_with_state, build_state_with_execution_runtime_override, start_server, AppState,
|
||||
};
|
||||
use crate::constants::{
|
||||
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
|
||||
TRUSTED_ADMIN_USER_ROLE_HEADER,
|
||||
};
|
||||
use crate::data::GatewayDataState;
|
||||
|
||||
fn build_empty_admin_system_data_state() -> GatewayDataState {
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
));
|
||||
let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::seed(Vec::<
|
||||
StoredPublicGlobalModel,
|
||||
>::new()));
|
||||
fn build_admin_system_data_state_with_repositories(
|
||||
provider_catalog_repository: Arc<InMemoryProviderCatalogReadRepository>,
|
||||
global_model_repository: Arc<InMemoryGlobalModelReadRepository>,
|
||||
) -> GatewayDataState {
|
||||
let auth_module_repository = Arc::new(InMemoryAuthModuleReadRepository::seed(
|
||||
Vec::<StoredOAuthProviderModuleConfig>::new(),
|
||||
None,
|
||||
@@ -59,6 +57,21 @@ fn build_empty_admin_system_data_state() -> GatewayDataState {
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
|
||||
}
|
||||
|
||||
fn build_empty_admin_system_data_state() -> GatewayDataState {
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
));
|
||||
let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::seed(Vec::<
|
||||
StoredPublicGlobalModel,
|
||||
>::new()));
|
||||
build_admin_system_data_state_with_repositories(
|
||||
provider_catalog_repository,
|
||||
global_model_repository,
|
||||
)
|
||||
}
|
||||
|
||||
fn sample_system_import_payload() -> Value {
|
||||
json!({
|
||||
"version": "2.2",
|
||||
@@ -503,7 +516,232 @@ async fn gateway_returns_503_for_admin_system_config_import_when_local_data_is_u
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_legacy_admin_system_config_import_versions() {
|
||||
async fn gateway_imports_legacy_admin_system_config_versions_and_model_test_succeeds() {
|
||||
for (
|
||||
fixture_name,
|
||||
expected_provider_name,
|
||||
expected_model_name,
|
||||
expected_api_key,
|
||||
expected_base_url,
|
||||
) in [
|
||||
(
|
||||
"v20",
|
||||
"legacy-provider-v20",
|
||||
"legacy-gpt-5-v20",
|
||||
"sk-legacy-v20",
|
||||
"https://legacy-v20.example.com/v1",
|
||||
),
|
||||
(
|
||||
"v21",
|
||||
"legacy-provider-v21",
|
||||
"legacy-gpt-5-v21",
|
||||
"sk-legacy-v21",
|
||||
"https://legacy-v21.example.com/v1",
|
||||
),
|
||||
] {
|
||||
assert_legacy_admin_system_config_import_model_test_succeeds(
|
||||
fixture_name,
|
||||
expected_provider_name,
|
||||
expected_model_name,
|
||||
expected_api_key,
|
||||
expected_base_url,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
async fn assert_legacy_admin_system_config_import_model_test_succeeds(
|
||||
fixture_name: &str,
|
||||
expected_provider_name: &str,
|
||||
expected_model_name: &str,
|
||||
expected_api_key: &str,
|
||||
expected_base_url: &str,
|
||||
) {
|
||||
let execution_runtime_hits = Arc::new(Mutex::new(0usize));
|
||||
let execution_runtime_hits_clone = Arc::clone(&execution_runtime_hits);
|
||||
let expected_base_url_for_runtime = expected_base_url.to_string();
|
||||
let expected_model_for_runtime = expected_model_name.to_string();
|
||||
let expected_bearer_for_runtime = format!("Bearer {expected_api_key}");
|
||||
let execution_runtime = Router::new().route(
|
||||
"/v1/execute/sync",
|
||||
any(move |Json(plan): Json<ExecutionPlan>| {
|
||||
let execution_runtime_hits_inner = Arc::clone(&execution_runtime_hits_clone);
|
||||
let expected_base_url = expected_base_url_for_runtime.clone();
|
||||
let expected_model = expected_model_for_runtime.clone();
|
||||
let expected_bearer = expected_bearer_for_runtime.clone();
|
||||
async move {
|
||||
*execution_runtime_hits_inner
|
||||
.lock()
|
||||
.expect("mutex should lock") += 1;
|
||||
assert_eq!(plan.provider_api_format, "openai:chat");
|
||||
assert!(
|
||||
plan.url.starts_with(expected_base_url.as_str()),
|
||||
"unexpected execution url: {}",
|
||||
plan.url
|
||||
);
|
||||
assert_eq!(plan.model_name.as_deref(), Some(expected_model.as_str()));
|
||||
assert_eq!(
|
||||
plan.headers.get("authorization").map(String::as_str),
|
||||
Some(expected_bearer.as_str())
|
||||
);
|
||||
assert_eq!(
|
||||
plan.body
|
||||
.json_body
|
||||
.as_ref()
|
||||
.and_then(|body| body.get("model"))
|
||||
.and_then(Value::as_str),
|
||||
Some(expected_model.as_str())
|
||||
);
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"candidate_id": plan.candidate_id,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"body": {
|
||||
"json_body": {
|
||||
"id": "chatcmpl-legacy-import",
|
||||
"object": "chat.completion",
|
||||
"choices": [{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Hello from imported provider"
|
||||
}
|
||||
}]
|
||||
}
|
||||
},
|
||||
"telemetry": {
|
||||
"elapsed_ms": 17
|
||||
}
|
||||
}))
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
));
|
||||
let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::seed(Vec::<
|
||||
StoredPublicGlobalModel,
|
||||
>::new()));
|
||||
let data_state = build_admin_system_data_state_with_repositories(
|
||||
Arc::clone(&provider_catalog_repository),
|
||||
Arc::clone(&global_model_repository),
|
||||
);
|
||||
let gateway = build_router_with_state(
|
||||
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(data_state),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
let import_response = client
|
||||
.post(format!("{gateway_url}/api/admin/system/config/import"))
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.json(&fixture_system_import_payload(fixture_name))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
let import_status = import_response.status();
|
||||
let import_payload: Value = import_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(import_status, StatusCode::OK, "payload={import_payload}");
|
||||
assert_eq!(import_payload["stats"]["providers"]["created"], json!(1));
|
||||
assert_eq!(import_payload["stats"]["endpoints"]["created"], json!(1));
|
||||
assert_eq!(import_payload["stats"]["keys"]["created"], json!(1));
|
||||
assert_eq!(import_payload["stats"]["models"]["created"], json!(1));
|
||||
|
||||
let providers = provider_catalog_repository
|
||||
.list_providers(false)
|
||||
.await
|
||||
.expect("providers should load");
|
||||
assert_eq!(providers.len(), 1);
|
||||
assert_eq!(providers[0].name, expected_provider_name);
|
||||
let provider_id = providers[0].id.clone();
|
||||
let provider_ids = vec![provider_id.clone()];
|
||||
let endpoints = provider_catalog_repository
|
||||
.list_endpoints_by_provider_ids(&provider_ids)
|
||||
.await
|
||||
.expect("endpoints should load");
|
||||
assert_eq!(endpoints.len(), 1);
|
||||
assert_eq!(endpoints[0].base_url, expected_base_url);
|
||||
let endpoint_id = endpoints[0].id.clone();
|
||||
|
||||
let keys = provider_catalog_repository
|
||||
.list_keys_by_provider_ids(&provider_ids)
|
||||
.await
|
||||
.expect("keys should load");
|
||||
assert_eq!(keys.len(), 1);
|
||||
assert_eq!(
|
||||
decrypt_python_fernet_ciphertext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
keys[0]
|
||||
.encrypted_api_key
|
||||
.as_deref()
|
||||
.expect("api key should be present"),
|
||||
)
|
||||
.expect("api key should decrypt"),
|
||||
expected_api_key
|
||||
);
|
||||
|
||||
let provider_models = global_model_repository
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider_id.clone(),
|
||||
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].provider_model_name, expected_model_name);
|
||||
|
||||
let test_response = client
|
||||
.post(format!("{gateway_url}/api/admin/provider-query/test-model"))
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.json(&json!({
|
||||
"provider_id": provider_id,
|
||||
"endpoint_id": endpoint_id,
|
||||
"model": expected_model_name,
|
||||
"api_format": "openai:chat"
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(test_response.status(), StatusCode::OK);
|
||||
let test_payload: Value = test_response.json().await.expect("json body should parse");
|
||||
assert_eq!(test_payload["success"], json!(true));
|
||||
assert_eq!(test_payload["model"], json!(expected_model_name));
|
||||
assert_eq!(test_payload["error"], Value::Null);
|
||||
assert_eq!(
|
||||
test_payload["data"]["response"]["choices"][0]["message"]["content"],
|
||||
json!("Hello from imported provider")
|
||||
);
|
||||
assert_eq!(
|
||||
*execution_runtime_hits.lock().expect("mutex should lock"),
|
||||
1
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_unknown_admin_system_config_import_versions() {
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
@@ -512,7 +750,7 @@ async fn gateway_rejects_legacy_admin_system_config_import_versions() {
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
for version in ["2.0", "2.1"] {
|
||||
for version in ["1.9", "2.3"] {
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/admin/system/config/import"))
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
@@ -535,7 +773,7 @@ async fn gateway_rejects_legacy_admin_system_config_import_versions() {
|
||||
.as_str()
|
||||
.expect("detail should be a string");
|
||||
assert!(detail.contains(&format!("不支持的配置版本: {version}")));
|
||||
assert!(detail.contains("支持的版本: 2.2"));
|
||||
assert!(detail.contains("支持的版本: 2.0, 2.1, 2.2"));
|
||||
}
|
||||
|
||||
gateway_handle.abort();
|
||||
@@ -842,13 +1080,8 @@ async fn gateway_imports_admin_system_config_fixture_v22() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_admin_system_config_fixtures_from_removed_legacy_exports() {
|
||||
async fn gateway_imports_admin_system_config_fixtures_from_legacy_exports() {
|
||||
for fixture in ["v20", "v21"] {
|
||||
let version = match fixture {
|
||||
"v20" => "2.0",
|
||||
"v21" => "2.1",
|
||||
_ => unreachable!("unexpected fixture"),
|
||||
};
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
@@ -869,20 +1102,170 @@ async fn gateway_rejects_admin_system_config_fixtures_from_removed_legacy_export
|
||||
|
||||
assert_eq!(
|
||||
response.status(),
|
||||
StatusCode::BAD_REQUEST,
|
||||
"fixture {fixture} should be rejected"
|
||||
StatusCode::OK,
|
||||
"fixture {fixture} should be accepted for Python migration compatibility"
|
||||
);
|
||||
let payload: Value = response.json().await.expect("json body should parse");
|
||||
let detail = payload["detail"]
|
||||
.as_str()
|
||||
.expect("detail should be a string");
|
||||
assert!(detail.contains(&format!("不支持的配置版本: {version}")));
|
||||
assert!(detail.contains("支持的版本: 2.2"));
|
||||
assert_eq!(payload["message"], "配置导入成功");
|
||||
assert_eq!(payload["stats"]["global_models"]["created"], json!(1));
|
||||
assert_eq!(payload["stats"]["providers"]["created"], json!(1));
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_imports_python_cli_alias_export_and_model_test_smoke() {
|
||||
let seen_plan = Arc::new(Mutex::new(None::<ExecutionPlan>));
|
||||
let seen_plan_clone = Arc::clone(&seen_plan);
|
||||
let execution_runtime = Router::new().route(
|
||||
"/v1/execute/sync",
|
||||
any(move |Json(plan): Json<ExecutionPlan>| {
|
||||
let seen_plan_inner = Arc::clone(&seen_plan_clone);
|
||||
async move {
|
||||
assert_eq!(plan.provider_api_format, "claude:messages");
|
||||
assert_eq!(plan.model_name.as_deref(), Some("claude-sonnet-python"));
|
||||
assert_eq!(
|
||||
plan.body
|
||||
.json_body
|
||||
.as_ref()
|
||||
.and_then(|body| body.get("model")),
|
||||
Some(&json!("claude-sonnet-python"))
|
||||
);
|
||||
*seen_plan_inner.lock().expect("mutex should lock") = Some(plan.clone());
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"candidate_id": plan.candidate_id,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"body": {
|
||||
"json_body": {
|
||||
"id": "msg_python_alias_smoke",
|
||||
"type": "message",
|
||||
"model": "claude-sonnet-python",
|
||||
"content": [{
|
||||
"type": "text",
|
||||
"text": "ok"
|
||||
}]
|
||||
}
|
||||
},
|
||||
"telemetry": {
|
||||
"elapsed_ms": 17
|
||||
}
|
||||
}))
|
||||
}
|
||||
}),
|
||||
);
|
||||
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
|
||||
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
));
|
||||
let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::seed(Vec::<
|
||||
StoredPublicGlobalModel,
|
||||
>::new()));
|
||||
let data_state = build_admin_system_data_state_with_repositories(
|
||||
Arc::clone(&provider_catalog_repository),
|
||||
Arc::clone(&global_model_repository),
|
||||
);
|
||||
let gateway = build_router_with_state(
|
||||
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(data_state),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let mut import_payload = sample_system_import_payload();
|
||||
import_payload["global_models"][0]["name"] = json!("claude-sonnet-python");
|
||||
import_payload["global_models"][0]["display_name"] = json!("Claude Sonnet Python");
|
||||
import_payload["providers"][0]["name"] = json!("python-export-claude");
|
||||
import_payload["providers"][0]["provider_type"] = json!("custom");
|
||||
import_payload["providers"][0]["endpoints"][0]["api_format"] = json!("claude:cli");
|
||||
import_payload["providers"][0]["endpoints"][0]["base_url"] =
|
||||
json!("https://python-export-claude.example.com");
|
||||
import_payload["providers"][0]["api_keys"][0]["name"] = json!("python-alias-key");
|
||||
import_payload["providers"][0]["api_keys"][0]["api_formats"] = json!(["claude:cli"]);
|
||||
import_payload["providers"][0]["api_keys"][0]["api_key"] = json!("sk-python-alias");
|
||||
import_payload["providers"][0]["models"][0]["global_model_name"] =
|
||||
json!("claude-sonnet-python");
|
||||
import_payload["providers"][0]["models"][0]["provider_model_name"] =
|
||||
json!("claude-sonnet-python");
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/admin/system/config/import"))
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.json(&import_payload)
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
let status = response.status();
|
||||
let payload: Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(status, StatusCode::OK, "payload={payload}");
|
||||
|
||||
let providers = provider_catalog_repository
|
||||
.list_providers(false)
|
||||
.await
|
||||
.expect("providers should load");
|
||||
assert_eq!(providers.len(), 1);
|
||||
let provider_id = providers[0].id.clone();
|
||||
let endpoints = provider_catalog_repository
|
||||
.list_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
|
||||
.await
|
||||
.expect("endpoints should load");
|
||||
assert_eq!(endpoints.len(), 1);
|
||||
assert_eq!(endpoints[0].api_format, "claude:messages");
|
||||
let keys = provider_catalog_repository
|
||||
.list_keys_by_provider_ids(std::slice::from_ref(&provider_id))
|
||||
.await
|
||||
.expect("keys should load");
|
||||
assert_eq!(keys.len(), 1);
|
||||
assert_eq!(keys[0].api_formats, Some(json!(["claude:messages"])));
|
||||
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/admin/provider-query/test-model"))
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.json(&json!({
|
||||
"provider_id": provider_id,
|
||||
"model": "claude-sonnet-python",
|
||||
"endpoint_id": endpoints[0].id,
|
||||
"api_format": "claude:messages"
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("model test request should succeed");
|
||||
|
||||
let status = response.status();
|
||||
let payload: Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(status, StatusCode::OK, "payload={payload}");
|
||||
assert_eq!(payload["success"], json!(true));
|
||||
assert_eq!(
|
||||
payload["attempts"][0]["endpoint_api_format"],
|
||||
json!("claude:messages")
|
||||
);
|
||||
assert_eq!(
|
||||
payload["attempts"][0]["request_body"]["model"],
|
||||
json!("claude-sonnet-python")
|
||||
);
|
||||
assert!(
|
||||
seen_plan.lock().expect("mutex should lock").is_some(),
|
||||
"post-import model test should execute through runtime"
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_legacy_user_import_string_bool_field() {
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default());
|
||||
@@ -964,11 +1347,23 @@ async fn gateway_reports_field_path_for_invalid_admin_system_config_import_shape
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_admin_system_config_with_numeric_string_prices() {
|
||||
async fn gateway_imports_admin_system_config_with_numeric_string_prices() {
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
));
|
||||
let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::seed(Vec::<
|
||||
StoredPublicGlobalModel,
|
||||
>::new()));
|
||||
let data_state = build_admin_system_data_state_with_repositories(
|
||||
Arc::clone(&provider_catalog_repository),
|
||||
Arc::clone(&global_model_repository),
|
||||
);
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(build_empty_admin_system_data_state()),
|
||||
.with_data_state_for_tests(data_state),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
@@ -989,11 +1384,37 @@ async fn gateway_rejects_admin_system_config_with_numeric_string_prices() {
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
let status = response.status();
|
||||
let body: Value = response.json().await.expect("json body should parse");
|
||||
let detail = body["detail"].as_str().expect("detail should be a string");
|
||||
assert!(detail.contains("配置文件格式无效"));
|
||||
assert!(detail.contains("default_price_per_request"));
|
||||
assert_eq!(status, StatusCode::OK, "payload={body}");
|
||||
|
||||
let global_models = global_model_repository
|
||||
.list_admin_global_models(&AdminGlobalModelListQuery {
|
||||
offset: 0,
|
||||
limit: 10_000,
|
||||
is_active: None,
|
||||
search: None,
|
||||
})
|
||||
.await
|
||||
.expect("global models should load");
|
||||
assert_eq!(global_models.items[0].default_price_per_request, Some(1.8));
|
||||
|
||||
let providers = provider_catalog_repository
|
||||
.list_providers(false)
|
||||
.await
|
||||
.expect("providers should load");
|
||||
let provider_models = global_model_repository
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: providers[0].id.clone(),
|
||||
is_active: None,
|
||||
offset: 0,
|
||||
limit: 10_000,
|
||||
})
|
||||
.await
|
||||
.expect("provider models should load");
|
||||
assert_eq!(provider_models[0].price_per_request, Some(0.7));
|
||||
assert_eq!(providers[0].request_timeout_secs, Some(30.0));
|
||||
assert_eq!(providers[0].stream_first_byte_timeout_secs, Some(15.0));
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user