feat(proxy): 重构 Proxy 节点管理与隧道系统

- 重构 proxy_nodes 管理端,支持节点注册、心跳、隧道生命周期管理
- 增强 tunnel 嵌入式 hub 和隧道协议
- 重构 aether-proxy 配置、隧道客户端、心跳和调度机制
- 调整 admin OAuth/配额/导入等处理器的参数传递
- 扩展数据迁移模块
- 补充 proxy nodes、OAuth、配额、系统导入等测试
- 更新前端 proxy nodes 视图和 API
This commit is contained in:
fawney19
2026-04-14 22:51:02 +08:00
parent fb31928e44
commit a4e7ac1df6
55 changed files with 3477 additions and 892 deletions

View File

@@ -5,7 +5,7 @@ use super::{classified, ClassifiedRoute};
pub(super) fn classify_admin_operations_family_route(
method: &http::Method,
normalized_path: &str,
_normalized_path_no_trailing: &str,
normalized_path_no_trailing: &str,
) -> Option<ClassifiedRoute> {
if method == http::Method::GET
&& matches!(
@@ -113,6 +113,24 @@ pub(super) fn classify_admin_operations_family_route(
"admin:proxy_nodes",
false,
))
} else if method == http::Method::GET
&& normalized_path_no_trailing.starts_with("/api/admin/proxy-nodes/")
&& normalized_path_no_trailing["/api/admin/proxy-nodes/".len()..]
.split('/')
.count()
== 1
&& !matches!(
&normalized_path_no_trailing["/api/admin/proxy-nodes/".len()..],
"register" | "heartbeat" | "unregister" | "manual" | "upgrade" | "test-url"
)
{
Some(classified(
"admin_proxy",
"proxy_nodes_manage",
"get_node",
"admin:proxy_nodes",
false,
))
} else if method == http::Method::POST
&& matches!(
normalized_path,

View File

@@ -91,6 +91,15 @@ fn classifies_admin_proxy_nodes_manual_update_as_admin_proxy_route() {
);
}
#[test]
fn classifies_admin_proxy_nodes_detail_as_admin_proxy_route() {
assert_proxy_nodes_admin_route(
http::Method::GET,
"/api/admin/proxy-nodes/node-1",
"get_node",
);
}
#[test]
fn classifies_admin_proxy_nodes_delete_as_admin_proxy_route() {
assert_proxy_nodes_admin_route(

View File

@@ -793,7 +793,7 @@ mod tests {
use axum::{Json, Router};
use serde_json::json;
use super::DirectSyncExecutionRuntime;
use super::{build_client, DirectSyncExecutionRuntime, ExecutionTransportControls};
use crate::frontdoor_loop_guard::{
frontdoor_self_loop_public_ai_path, gateway_frontdoor_self_loop_guard_error_with_port,
gateway_frontdoor_self_loop_guard_matches_with_port,
@@ -842,6 +842,32 @@ mod tests {
);
}
#[test]
fn direct_sync_execution_runtime_builds_clients_for_socks_proxy_urls() {
let timeouts = ExecutionTimeouts {
connect_ms: Some(5_000),
total_ms: Some(5_000),
..ExecutionTimeouts::default()
};
for proxy_url in ["socks5://127.0.0.1:1080", "socks5h://127.0.0.1:1080"] {
build_client(
Some(&timeouts),
Some(&aether_contracts::ProxySnapshot {
enabled: Some(true),
mode: Some("socks".into()),
node_id: None,
label: Some("manual-proxy".into()),
url: Some(proxy_url.to_string()),
extra: None,
}),
None,
ExecutionTransportControls::default(),
)
.unwrap_or_else(|err| panic!("client should build for {proxy_url}: {err}"));
}
}
fn tunnel_proxy_snapshot(base_url: String) -> aether_contracts::ProxySnapshot {
aether_contracts::ProxySnapshot {
enabled: Some(true),

View File

@@ -31,3 +31,5 @@ pub(crate) use self::request::{
AdminAppState, AdminRequestContext, AdminRouteRequest, AdminRouteResponse, AdminRouteResult,
};
pub(crate) use self::routes::maybe_build_local_admin_response;
#[cfg(test)]
pub(crate) use self::system::override_proxy_connectivity_probe_url_for_tests;

View File

@@ -8,9 +8,11 @@ use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider
use crate::handlers::admin::provider::oauth::provisioning::build_provider_oauth_auth_config_from_token_payload;
use crate::handlers::admin::provider::oauth::provisioning::{
create_provider_oauth_catalog_key, provider_oauth_active_api_formats,
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
update_existing_provider_oauth_catalog_key,
};
use crate::handlers::admin::provider::oauth::runtime::{
provider_oauth_runtime_endpoint_for_provider, refresh_provider_oauth_account_state_after_update,
};
use crate::handlers::admin::provider::oauth::runtime::refresh_provider_oauth_account_state_after_update;
use crate::handlers::admin::provider::oauth::state::{
admin_provider_oauth_template, exchange_admin_provider_oauth_refresh_token,
};
@@ -116,7 +118,18 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
.list_provider_catalog_endpoints_by_provider_ids(&[provider_id.to_string()])
.await?;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id);
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider(provider_type, &endpoints);
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
proxy_node_id,
&[
runtime_endpoint
.as_ref()
.and_then(|endpoint| endpoint.proxy.as_ref()),
provider.proxy.as_ref(),
],
)
.await;
let mut results = Vec::with_capacity(entries.len());
let mut success = 0usize;
let mut failed = 0usize;
@@ -126,7 +139,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
state,
template,
entry.refresh_token.as_str(),
proxy_node_id,
request_proxy.clone(),
)
.await
{
@@ -189,7 +202,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
None,
expires_at,
)
.await?
@@ -232,7 +245,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
None,
expires_at,
)
.await?

View File

@@ -2,9 +2,11 @@ use super::parse::{AdminProviderOAuthBatchImportEntry, AdminProviderOAuthBatchIm
use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider_oauth_key;
use crate::handlers::admin::provider::oauth::provisioning::{
create_provider_oauth_catalog_key, provider_oauth_active_api_formats,
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
update_existing_provider_oauth_catalog_key,
};
use crate::handlers::admin::provider::oauth::runtime::{
provider_oauth_runtime_endpoint_for_provider, refresh_provider_oauth_account_state_after_update,
};
use crate::handlers::admin::provider::oauth::runtime::refresh_provider_oauth_account_state_after_update;
use crate::handlers::admin::provider::oauth::state::decode_jwt_claims;
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
use crate::handlers::admin::request::{AdminAppState, AdminKiroAuthConfig};
@@ -14,6 +16,7 @@ use aether_admin::provider::oauth::{
build_kiro_batch_import_key_name, coerce_admin_provider_oauth_import_str,
parse_admin_provider_oauth_kiro_batch_import_entries,
};
use aether_contracts::ProxySnapshot;
use serde_json::{json, Map, Value};
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
@@ -102,7 +105,7 @@ fn admin_provider_oauth_kiro_refresh_error_detail(
async fn refresh_admin_provider_oauth_kiro_auth_config(
state: &AdminAppState<'_>,
auth_config: &AdminKiroAuthConfig,
proxy_node_id: Option<&str>,
proxy: Option<ProxySnapshot>,
social_refresh_base_url: Option<&str>,
idc_refresh_base_url: Option<&str>,
) -> Result<AdminKiroAuthConfig, String> {
@@ -164,7 +167,7 @@ async fn refresh_admin_provider_oauth_kiro_auth_config(
"grantType": "refresh_token",
})),
None,
proxy_node_id,
proxy.clone(),
)
.await
.map_err(|err| format!("IDC refresh 请求失败: {err}"))?;
@@ -268,7 +271,7 @@ async fn refresh_admin_provider_oauth_kiro_auth_config(
.unwrap_or_default(),
})),
None,
proxy_node_id,
proxy,
)
.await
.map_err(|err| format!("social refresh 请求失败: {err}"))?;
@@ -352,7 +355,18 @@ pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(&[provider_id.to_string()])
.await?;
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id);
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider("kiro", &endpoints);
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
proxy_node_id,
&[
runtime_endpoint
.as_ref()
.and_then(|endpoint| endpoint.proxy.as_ref()),
provider.proxy.as_ref(),
],
)
.await;
let social_refresh_base_url =
admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_social_refresh");
let idc_refresh_base_url =
@@ -392,7 +406,7 @@ pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
refreshed_auth_config = match refresh_admin_provider_oauth_kiro_auth_config(
state,
&refreshed_auth_config,
proxy_node_id,
request_proxy.clone(),
social_refresh_base_url.as_deref(),
idc_refresh_base_url.as_deref(),
)
@@ -477,7 +491,7 @@ pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
None,
refreshed_auth_config.expires_at,
)
.await?
@@ -511,7 +525,7 @@ pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
&access_token,
&auth_config,
&provider_oauth_active_api_formats(&endpoints),
key_proxy.clone(),
None,
refreshed_auth_config.expires_at,
)
.await?

View File

@@ -1,5 +1,6 @@
use super::super::super::errors::build_internal_control_error_response;
use super::super::super::quota::codex::refresh_codex_provider_quota_locally;
use super::super::super::runtime::provider_oauth_runtime_endpoint_for_provider;
use super::super::super::state::{
admin_provider_oauth_template, enrich_admin_provider_oauth_auth_config,
is_fixed_provider_type_for_provider_oauth, json_non_empty_string, json_u64_value,
@@ -123,6 +124,22 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
"该 Provider 不支持 OAuth 授权",
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider(&provider_type, &endpoints);
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
payload.proxy_node_id.as_deref(),
&[
key.proxy.as_ref(),
runtime_endpoint
.as_ref()
.and_then(|endpoint| endpoint.proxy.as_ref()),
provider.proxy.as_ref(),
],
)
.await;
let token_payload = match state
.exchange_admin_provider_oauth_code(
@@ -130,7 +147,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
&callback.code,
&callback.state_nonce,
state_data.pkce_verifier.as_deref(),
payload.proxy_node_id.as_deref(),
request_proxy,
)
.await
{

View File

@@ -2,10 +2,11 @@ use super::super::super::duplicates::find_duplicate_provider_oauth_key;
use super::super::super::errors::build_internal_control_error_response;
use super::super::super::provisioning::{
build_provider_oauth_auth_config_from_token_payload, create_provider_oauth_catalog_key,
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
update_existing_provider_oauth_catalog_key,
provider_oauth_active_api_formats, update_existing_provider_oauth_catalog_key,
};
use super::super::super::runtime::{
provider_oauth_runtime_endpoint_for_provider, refresh_provider_oauth_account_state_after_update,
};
use super::super::super::runtime::refresh_provider_oauth_account_state_after_update;
use super::super::super::state::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
is_fixed_provider_type_for_provider_oauth,
@@ -114,6 +115,21 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
"该 Provider 不支持 OAuth 授权",
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider(&provider_type, &endpoints);
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
payload.proxy_node_id.as_deref(),
&[
runtime_endpoint
.as_ref()
.and_then(|endpoint| endpoint.proxy.as_ref()),
provider.proxy.as_ref(),
],
)
.await;
let token_payload = match state
.exchange_admin_provider_oauth_code(
@@ -121,7 +137,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
&callback.code,
&callback.state_nonce,
state_data.pkce_verifier.as_deref(),
payload.proxy_node_id.as_deref(),
request_proxy,
)
.await
{
@@ -138,11 +154,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let key_proxy = provider_oauth_key_proxy_value(payload.proxy_node_id.as_deref());
let duplicate = match state
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
.await
@@ -163,7 +175,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
None,
expires_at,
)
.await?
@@ -204,7 +216,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
None,
expires_at,
)
.await?

View File

@@ -1,5 +1,6 @@
use super::session::AdminProviderOAuthDeviceAuthorizePayload;
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::runtime::provider_oauth_runtime_endpoint_for_provider;
use crate::handlers::admin::provider::oauth::state::{
build_admin_provider_oauth_backend_unavailable_response, current_unix_secs,
default_kiro_device_start_url, generate_provider_oauth_nonce, json_non_empty_string,
@@ -69,6 +70,21 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
"设备授权仅支持 Kiro provider",
));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider("kiro", &endpoints);
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
payload.proxy_node_id.as_deref(),
&[
runtime_endpoint
.as_ref()
.and_then(|endpoint| endpoint.proxy.as_ref()),
provider.proxy.as_ref(),
],
)
.await;
let region = normalize_kiro_device_region(Some(payload.region.as_str())).ok_or_else(|| {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, "region 格式无效")
@@ -85,11 +101,7 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
};
let client_registration = match state
.register_admin_kiro_device_oidc_client(
&region,
&start_url,
payload.proxy_node_id.as_deref(),
)
.register_admin_kiro_device_oidc_client(&region, &start_url, request_proxy.clone())
.await
{
Ok(payload) => payload,
@@ -114,7 +126,7 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
&client_id,
&client_secret,
&start_url,
payload.proxy_node_id.as_deref(),
request_proxy,
)
.await
{

View File

@@ -5,9 +5,11 @@ use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::provisioning::{
create_provider_oauth_catalog_key, provider_oauth_active_api_formats,
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
update_existing_provider_oauth_catalog_key,
};
use crate::handlers::admin::provider::oauth::runtime::{
provider_oauth_runtime_endpoint_for_provider, refresh_provider_oauth_account_state_after_update,
};
use crate::handlers::admin::provider::oauth::runtime::refresh_provider_oauth_account_state_after_update;
use crate::handlers::admin::provider::oauth::state::{
build_admin_provider_oauth_backend_unavailable_response, build_kiro_device_key_name,
current_unix_secs, decode_jwt_claims, json_non_empty_string, json_u64_value,
@@ -125,6 +127,21 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
"Provider 不存在",
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider("kiro", &endpoints);
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
session.proxy_node_id.as_deref(),
&[
runtime_endpoint
.as_ref()
.and_then(|endpoint| endpoint.proxy.as_ref()),
provider.proxy.as_ref(),
],
)
.await;
let token_result = match state
.poll_admin_kiro_device_token(
@@ -132,7 +149,7 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
&session.client_id,
&session.client_secret,
&session.device_code,
session.proxy_node_id.as_deref(),
request_proxy,
)
.await
{
@@ -252,12 +269,7 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
}
};
let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref());
let api_formats = provider_oauth_active_api_formats(
&state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?,
);
let api_formats = provider_oauth_active_api_formats(&endpoints);
let mut replaced = false;
let persisted_key = if let Some(existing_key) = duplicate {
replaced = true;
@@ -266,7 +278,7 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
None,
Some(expires_at),
)
.await?
@@ -288,7 +300,7 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
None,
Some(expires_at),
)
.await?

View File

@@ -2,10 +2,11 @@ use super::super::duplicates::find_duplicate_provider_oauth_key;
use super::super::errors::build_internal_control_error_response;
use super::super::provisioning::{
build_provider_oauth_auth_config_from_token_payload, create_provider_oauth_catalog_key,
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
update_existing_provider_oauth_catalog_key,
provider_oauth_active_api_formats, update_existing_provider_oauth_catalog_key,
};
use super::super::runtime::{
provider_oauth_runtime_endpoint_for_provider, refresh_provider_oauth_account_state_after_update,
};
use super::super::runtime::refresh_provider_oauth_account_state_after_update;
use super::super::state::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
exchange_admin_provider_oauth_refresh_token, is_fixed_provider_type_for_provider_oauth,
@@ -96,13 +97,24 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider(&provider_type, &endpoints);
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
proxy_node_id.as_deref(),
&[
runtime_endpoint
.as_ref()
.and_then(|endpoint| endpoint.proxy.as_ref()),
provider.proxy.as_ref(),
],
)
.await;
let token_payload = match state
.exchange_admin_provider_oauth_refresh_token(
template,
refresh_token_input,
proxy_node_id.as_deref(),
)
.exchange_admin_provider_oauth_refresh_token(template, refresh_token_input, request_proxy)
.await
{
Ok(payload) => payload,
@@ -124,11 +136,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id.as_deref());
let duplicate = match state
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
.await
@@ -149,7 +157,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
None,
expires_at,
)
.await?
@@ -189,7 +197,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
None,
expires_at,
)
.await?

View File

@@ -79,7 +79,7 @@ pub(super) async fn parse_admin_provider_oauth_refresh_request(
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, endpoints)
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, &endpoints)
else {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,

View File

@@ -1,13 +1,14 @@
use super::shared::{
coerce_json_f64, coerce_json_string, execute_provider_quota_plan,
extract_execution_error_message, persist_provider_quota_refresh_state,
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
coerce_json_f64, coerce_json_string, default_provider_quota_execution_timeouts,
execute_provider_quota_plan, extract_execution_error_message,
persist_provider_quota_refresh_state, quota_refresh_success_invalid_state,
ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::provider::shared::payloads::ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH;
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota::parse_antigravity_usage_response;
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_contracts::{ExecutionPlan, RequestBody};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
@@ -31,6 +32,14 @@ async fn execute_antigravity_quota_plan(
.or_insert_with(|| "antigravity".to_string());
let body = json!({ "project": project_id });
let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await;
let timeouts = state
.resolve_transport_execution_timeouts(transport)
.or(Some(default_provider_quota_execution_timeouts(
proxy.as_ref(),
)));
let plan = ExecutionPlan {
request_id: format!("antigravity-quota:{}", transport.key.id),
candidate_id: None,
@@ -56,20 +65,9 @@ async fn execute_antigravity_quota_plan(
client_api_format: "gemini:chat".to_string(),
provider_api_format: "antigravity:fetch_available_models".to_string(),
model_name: Some("fetchAvailableModels".to_string()),
proxy: state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await,
proxy,
tls_profile: state.resolve_transport_tls_profile(transport),
timeouts: state
.resolve_transport_execution_timeouts(transport)
.or(Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
})),
timeouts,
};
execute_provider_quota_plan(state, transport, plan, "antigravity").await

View File

@@ -1,9 +1,12 @@
use super::super::shared::{execute_provider_quota_plan, ProviderQuotaExecutionOutcome};
use super::super::shared::{
default_provider_quota_execution_timeouts, execute_provider_quota_plan,
ProviderQuotaExecutionOutcome,
};
use super::parse::normalize_codex_plan_type;
use crate::handlers::admin::provider::shared::payloads::CODEX_WHAM_USAGE_URL;
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_contracts::{ExecutionPlan, RequestBody};
use std::collections::BTreeMap;
pub(super) fn build_codex_refresh_headers(
@@ -58,6 +61,14 @@ pub(super) async fn execute_codex_quota_plan(
transport: &AdminGatewayProviderTransportSnapshot,
headers: BTreeMap<String, String>,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await;
let timeouts = state
.resolve_transport_execution_timeouts(transport)
.or(Some(default_provider_quota_execution_timeouts(
proxy.as_ref(),
)));
let plan = ExecutionPlan {
request_id: format!("codex-quota:{}", transport.key.id),
candidate_id: None,
@@ -79,20 +90,9 @@ pub(super) async fn execute_codex_quota_plan(
client_api_format: "openai:cli".to_string(),
provider_api_format: "openai:cli".to_string(),
model_name: Some("codex-wham-usage".to_string()),
proxy: state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await,
proxy,
tls_profile: state.resolve_transport_tls_profile(transport),
timeouts: state
.resolve_transport_execution_timeouts(transport)
.or(Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
})),
timeouts,
};
execute_provider_quota_plan(state, transport, plan, "codex").await
}

View File

@@ -1,3 +1,4 @@
use super::super::shared::default_provider_quota_execution_timeouts;
use super::super::shared::{execute_provider_quota_plan, ProviderQuotaExecutionOutcome};
use crate::handlers::admin::provider::shared::payloads::{
KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION,
@@ -6,7 +7,7 @@ use crate::handlers::admin::request::{
AdminAppState, AdminGatewayProviderTransportSnapshot, AdminKiroRequestAuth,
};
use crate::GatewayError;
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_contracts::{ExecutionPlan, RequestBody};
use std::collections::BTreeMap;
use url::form_urlencoded;
use uuid::Uuid;
@@ -66,6 +67,14 @@ pub(super) async fn execute_kiro_quota_plan(
transport: &AdminGatewayProviderTransportSnapshot,
auth: &AdminKiroRequestAuth,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await;
let timeouts = state
.resolve_transport_execution_timeouts(transport)
.or(Some(default_provider_quota_execution_timeouts(
proxy.as_ref(),
)));
let plan = ExecutionPlan {
request_id: format!("kiro-quota:{}", transport.key.id),
candidate_id: None,
@@ -87,20 +96,9 @@ pub(super) async fn execute_kiro_quota_plan(
client_api_format: "claude:cli".to_string(),
provider_api_format: "kiro:usage".to_string(),
model_name: Some("kiro-usage-limits".to_string()),
proxy: state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await,
proxy,
tls_profile: state.resolve_transport_tls_profile(transport),
timeouts: state
.resolve_transport_execution_timeouts(transport)
.or(Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
})),
timeouts,
};
execute_provider_quota_plan(state, transport, plan, "kiro").await

View File

@@ -4,16 +4,37 @@ use crate::handlers::admin::provider::shared::payloads::{
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota as admin_provider_quota_pure;
use aether_contracts::{ExecutionPlan, ExecutionResult};
use aether_contracts::{ExecutionPlan, ExecutionResult, ExecutionTimeouts, ProxySnapshot};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn;
const PROVIDER_QUOTA_DEFAULT_TIMEOUT_MS: u64 = 30_000;
const PROVIDER_QUOTA_PROXY_TIMEOUT_MS: u64 = 60_000;
pub(super) enum ProviderQuotaExecutionOutcome {
Response(ExecutionResult),
Failure(String),
}
pub(super) fn default_provider_quota_execution_timeouts(
proxy: Option<&ProxySnapshot>,
) -> ExecutionTimeouts {
let timeout_ms = if proxy.is_some() {
PROVIDER_QUOTA_PROXY_TIMEOUT_MS
} else {
PROVIDER_QUOTA_DEFAULT_TIMEOUT_MS
};
ExecutionTimeouts {
connect_ms: Some(timeout_ms),
read_ms: Some(timeout_ms),
write_ms: Some(timeout_ms),
pool_ms: Some(timeout_ms),
total_ms: Some(timeout_ms),
..ExecutionTimeouts::default()
}
}
pub(super) fn provider_auto_remove_banned_keys(config: Option<&serde_json::Value>) -> bool {
admin_provider_quota_pure::provider_auto_remove_banned_keys(config)
}
@@ -127,6 +148,9 @@ pub(super) async fn execute_provider_quota_plan(
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let proxy_source = state
.resolve_transport_proxy_source_with_tunnel_affinity(transport)
.await;
let proxy_url_present = plan
.proxy
.as_ref()
@@ -138,6 +162,7 @@ pub(super) async fn execute_provider_quota_plan(
endpoint_id = %transport.endpoint.id,
url = %plan.url,
tls_profile = ?plan.tls_profile.as_deref(),
proxy_source = ?proxy_source,
proxy_node_id = ?proxy_node_id,
proxy_url_present,
error = %error,

View File

@@ -9,28 +9,34 @@ use aether_data_contracts::repository::provider_catalog::{
pub(crate) fn provider_oauth_runtime_endpoint_for_provider(
provider_type: &str,
endpoints: Vec<StoredProviderCatalogEndpoint>,
endpoints: &[StoredProviderCatalogEndpoint],
) -> Option<StoredProviderCatalogEndpoint> {
let provider_type = provider_type.trim().to_ascii_lowercase();
match provider_type.as_str() {
"codex" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:cli")
}),
"antigravity" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& (endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:chat")
|| endpoint
"codex" => endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:cli"))
}),
.eq_ignore_ascii_case("openai:cli")
})
.cloned(),
"antigravity" => endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& (endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:chat")
|| endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:cli"))
})
.cloned(),
"kiro" => endpoints
.iter()
.find(|endpoint| {
@@ -41,8 +47,16 @@ pub(crate) fn provider_oauth_runtime_endpoint_for_provider(
.eq_ignore_ascii_case("claude:cli")
})
.cloned()
.or_else(|| endpoints.into_iter().find(|endpoint| endpoint.is_active)),
_ => endpoints.into_iter().find(|endpoint| endpoint.is_active),
.or_else(|| {
endpoints
.iter()
.find(|endpoint| endpoint.is_active)
.cloned()
}),
_ => endpoints
.iter()
.find(|endpoint| endpoint.is_active)
.cloned(),
}
}
@@ -59,7 +73,7 @@ pub(crate) async fn refresh_provider_oauth_account_state_after_update(
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, endpoints)
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, &endpoints)
else {
return Ok((false, None));
};

View File

@@ -3,16 +3,25 @@ use super::super::errors::{
};
use super::json_non_empty_string;
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
use aether_contracts::ProxySnapshot;
use axum::{body::Body, http, response::Response};
use url::form_urlencoded;
fn provider_oauth_transport_error_detail(prefix: &str, error: &str) -> String {
let error = error.trim();
if error.is_empty() {
return prefix.to_string();
}
format!("{prefix}: {error}")
}
pub(crate) async fn exchange_admin_provider_oauth_code(
state: &AdminAppState<'_>,
template: AdminProviderOAuthTemplate,
code: &str,
state_nonce: &str,
pkce_verifier: Option<&str>,
proxy_node_id: Option<&str>,
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let response = if template.provider_type == "claude_code" {
@@ -63,7 +72,7 @@ pub(crate) async fn exchange_admin_provider_oauth_code(
Some("application/json"),
Some(serde_json::Value::Object(body)),
None,
proxy_node_id,
proxy.clone(),
)
.await
} else {
@@ -100,12 +109,15 @@ pub(crate) async fn exchange_admin_provider_oauth_code(
Some("application/x-www-form-urlencoded"),
None,
Some(form_body),
proxy_node_id,
proxy.clone(),
)
.await
}
.map_err(|_| {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, "token exchange 失败")
.map_err(|error| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
provider_oauth_transport_error_detail("token exchange 失败", &error),
)
})?;
if !response.status.is_success() {
@@ -134,7 +146,7 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
state: &AdminAppState<'_>,
template: AdminProviderOAuthTemplate,
refresh_token: &str,
proxy_node_id: Option<&str>,
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let scope = template.scopes.join(" ");
@@ -175,7 +187,7 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
Some("application/json"),
Some(serde_json::Value::Object(body)),
None,
proxy_node_id,
proxy.clone(),
)
.await
} else {
@@ -211,14 +223,17 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
Some("application/x-www-form-urlencoded"),
None,
Some(form_body),
proxy_node_id,
proxy.clone(),
)
.await
}
.map_err(|_| {
.map_err(|error| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Refresh Token 验证失败: token exchange 失败",
provider_oauth_transport_error_detail(
"Refresh Token 验证失败: token exchange 失败",
&error,
),
)
})?;

View File

@@ -28,9 +28,6 @@ pub(crate) fn normalize_pool_advanced_config(
serde_json::Value::Null => Ok(None),
// `pool_advanced: {}` still means "enable pool mode with defaults".
serde_json::Value::Object(map) => Ok(Some(serde_json::Value::Object(map))),
// Backward compatibility for older boolean payloads.
serde_json::Value::Bool(true) => Ok(Some(serde_json::json!({}))),
serde_json::Value::Bool(false) => Ok(None),
_ => Err("pool_advanced 必须是 JSON 对象".to_string()),
}
}
@@ -78,14 +75,14 @@ mod tests {
}
#[test]
fn normalize_pool_advanced_accepts_legacy_booleans() {
fn normalize_pool_advanced_rejects_legacy_booleans() {
assert_eq!(
normalize_pool_advanced_config(Some(json!(true))).expect("true should normalize"),
Some(json!({}))
normalize_pool_advanced_config(Some(json!(true))).unwrap_err(),
"pool_advanced 必须是 JSON 对象"
);
assert_eq!(
normalize_pool_advanced_config(Some(json!(false))).expect("false should normalize"),
None
normalize_pool_advanced_config(Some(json!(false))).unwrap_err(),
"pool_advanced 必须是 JSON 对象"
);
}
}

View File

@@ -1,7 +1,7 @@
use super::*;
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use aether_contracts::{
ExecutionPlan, ExecutionResult, ExecutionTimeouts, RequestBody,
ExecutionPlan, ExecutionResult, ExecutionTimeouts, ProxySnapshot, RequestBody,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
};
use aether_data::repository::provider_oauth::{
@@ -21,6 +21,7 @@ use url::Url;
const KIRO_IDC_AMZ_USER_AGENT: &str =
"aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
const ADMIN_PROVIDER_OAUTH_TIMEOUT_MS: u64 = 30_000;
const ADMIN_PROVIDER_OAUTH_PROXY_TIMEOUT_MS: u64 = 60_000;
pub(crate) struct AdminProviderOAuthHttpResponse {
pub(crate) status: http::StatusCode,
@@ -132,7 +133,7 @@ impl<'a> AdminAppState<'a> {
code: &str,
state_nonce: &str,
pkce_verifier: Option<&str>,
proxy_node_id: Option<&str>,
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
crate::handlers::admin::provider::oauth::state::exchange_admin_provider_oauth_code(
self,
@@ -140,7 +141,7 @@ impl<'a> AdminAppState<'a> {
code,
state_nonce,
pkce_verifier,
proxy_node_id,
proxy,
)
.await
}
@@ -149,13 +150,13 @@ impl<'a> AdminAppState<'a> {
&self,
template: AdminProviderOAuthTemplate,
refresh_token: &str,
proxy_node_id: Option<&str>,
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
crate::handlers::admin::provider::oauth::state::exchange_admin_provider_oauth_refresh_token(
self,
template,
refresh_token,
proxy_node_id,
proxy,
)
.await
}
@@ -315,7 +316,7 @@ impl<'a> AdminAppState<'a> {
&self,
region: &str,
start_url: &str,
proxy_node_id: Option<&str>,
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
let payload = post_kiro_device_oidc_json(
self,
@@ -337,7 +338,7 @@ impl<'a> AdminAppState<'a> {
],
"issuerUrl": start_url,
}),
proxy_node_id,
proxy,
)
.await?;
if payload
@@ -364,7 +365,7 @@ impl<'a> AdminAppState<'a> {
client_id: &str,
client_secret: &str,
start_url: &str,
proxy_node_id: Option<&str>,
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
let payload = post_kiro_device_oidc_json(
self,
@@ -375,7 +376,7 @@ impl<'a> AdminAppState<'a> {
"clientSecret": client_secret,
"startUrl": start_url,
}),
proxy_node_id,
proxy,
)
.await?;
if payload
@@ -402,7 +403,7 @@ impl<'a> AdminAppState<'a> {
client_id: &str,
client_secret: &str,
device_code: &str,
proxy_node_id: Option<&str>,
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
post_kiro_device_oidc_json(
self,
@@ -414,11 +415,35 @@ impl<'a> AdminAppState<'a> {
"grantType": "urn:ietf:params:oauth:grant-type:device_code",
"deviceCode": device_code,
}),
proxy_node_id,
proxy,
)
.await
}
pub(crate) async fn resolve_admin_provider_oauth_operation_proxy_snapshot(
&self,
temporary_proxy_node_id: Option<&str>,
configured_proxies: &[Option<&serde_json::Value>],
) -> Option<ProxySnapshot> {
if let Some(snapshot) = self
.resolve_admin_proxy_node_snapshot(temporary_proxy_node_id)
.await
{
return Some(snapshot);
}
for proxy in configured_proxies {
if let Some(snapshot) = self
.app
.resolve_configured_proxy_snapshot_with_tunnel_affinity(*proxy)
.await
{
return Some(snapshot);
}
}
self.app.resolve_system_proxy_snapshot().await
}
pub(crate) async fn find_duplicate_provider_oauth_key(
&self,
provider_id: &str,
@@ -504,7 +529,7 @@ async fn post_kiro_device_oidc_json(
endpoint_key: &str,
default_url: String,
body: serde_json::Value,
proxy_node_id: Option<&str>,
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
let url = state.provider_oauth_token_url(endpoint_key, &default_url);
let host = Url::parse(&url)
@@ -539,7 +564,7 @@ async fn post_kiro_device_oidc_json(
Some("application/json"),
Some(body),
None,
proxy_node_id,
proxy,
)
.await
.map_err(|_| {
@@ -570,7 +595,7 @@ impl<'a> AdminAppState<'a> {
content_type: Option<&str>,
json_body: Option<serde_json::Value>,
body_bytes: Option<Vec<u8>>,
proxy_node_id: Option<&str>,
proxy: Option<ProxySnapshot>,
) -> Result<AdminProviderOAuthHttpResponse, String> {
let body = if let Some(json_body) = json_body {
RequestBody::from_json(json_body)
@@ -581,6 +606,7 @@ impl<'a> AdminAppState<'a> {
body_ref: None,
}
};
let timeout_ms = admin_provider_oauth_timeout_ms(proxy.as_ref());
let plan = ExecutionPlan {
request_id: request_id.to_string(),
candidate_id: None,
@@ -601,21 +627,14 @@ impl<'a> AdminAppState<'a> {
client_api_format: "provider_oauth:exchange".to_string(),
provider_api_format: "provider_oauth:exchange".to_string(),
model_name: Some("oauth-exchange".to_string()),
proxy: if proxy_node_id
.map(str::trim)
.is_some_and(|value| !value.is_empty())
{
self.resolve_admin_proxy_node_snapshot(proxy_node_id).await
} else {
self.app.resolve_system_proxy_snapshot().await
},
proxy,
tls_profile: None,
timeouts: Some(ExecutionTimeouts {
connect_ms: Some(ADMIN_PROVIDER_OAUTH_TIMEOUT_MS),
read_ms: Some(ADMIN_PROVIDER_OAUTH_TIMEOUT_MS),
write_ms: Some(ADMIN_PROVIDER_OAUTH_TIMEOUT_MS),
pool_ms: Some(ADMIN_PROVIDER_OAUTH_TIMEOUT_MS),
total_ms: Some(ADMIN_PROVIDER_OAUTH_TIMEOUT_MS),
connect_ms: Some(timeout_ms),
read_ms: Some(timeout_ms),
write_ms: Some(timeout_ms),
pool_ms: Some(timeout_ms),
total_ms: Some(timeout_ms),
..ExecutionTimeouts::default()
}),
};
@@ -632,6 +651,14 @@ impl<'a> AdminAppState<'a> {
}
}
fn admin_provider_oauth_timeout_ms(proxy: Option<&ProxySnapshot>) -> u64 {
if proxy.is_some() {
ADMIN_PROVIDER_OAUTH_PROXY_TIMEOUT_MS
} else {
ADMIN_PROVIDER_OAUTH_TIMEOUT_MS
}
}
fn maybe_insert_host_header(
mut headers: reqwest::header::HeaderMap,
host: &str,

View File

@@ -82,6 +82,15 @@ impl<'a> AdminAppState<'a> {
.await
}
pub(crate) async fn resolve_transport_proxy_source_with_tunnel_affinity(
&self,
transport: &AdminGatewayProviderTransportSnapshot,
) -> Option<&'static str> {
self.app
.resolve_transport_proxy_source_with_tunnel_affinity(transport)
.await
}
pub(crate) fn fixed_provider_template(
&self,
provider_type: &str,
@@ -115,15 +124,14 @@ impl<'a> AdminAppState<'a> {
return Some(snapshot);
}
if explicit_node_id.is_none() {
if let Some(snapshot) = self.app.resolve_system_proxy_snapshot().await {
return Some(snapshot);
}
let proxy = connector_config
.and_then(|config| config.get("proxy"))
.and_then(admin_provider_transport_proxy_snapshot);
if proxy.is_some() {
return proxy;
}
connector_config
.and_then(|config| config.get("proxy"))
.and_then(admin_provider_transport_legacy_proxy_snapshot)
self.app.resolve_system_proxy_snapshot().await
}
pub(crate) async fn resolve_admin_proxy_node_snapshot(
@@ -256,65 +264,46 @@ impl<'a> AdminAppState<'a> {
}
}
fn admin_provider_transport_legacy_proxy_snapshot(value: &Value) -> Option<ProxySnapshot> {
match value {
Value::String(proxy_url) => {
let proxy_url = proxy_url.trim();
if proxy_url.is_empty() {
return None;
}
Some(ProxySnapshot {
enabled: Some(true),
mode: admin_provider_transport_proxy_mode(Some(proxy_url)),
node_id: None,
label: None,
url: Some(proxy_url.to_string()),
extra: None,
})
}
Value::Object(object) => {
if object.get("enabled").and_then(Value::as_bool) == Some(false) {
return None;
}
let proxy_url = object
.get("url")
.or_else(|| object.get("proxy_url"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())?;
let username = object
.get("username")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
let password = object
.get("password")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
Some(ProxySnapshot {
enabled: Some(true),
mode: object
.get("mode")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.or_else(|| admin_provider_transport_proxy_mode(Some(proxy_url))),
node_id: None,
label: object
.get("label")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
url: admin_provider_transport_inject_proxy_auth(proxy_url, username, password)
.or_else(|| Some(proxy_url.to_string())),
extra: None,
})
}
_ => None,
fn admin_provider_transport_proxy_snapshot(value: &Value) -> Option<ProxySnapshot> {
let object = value.as_object()?;
if object.get("enabled").and_then(Value::as_bool) == Some(false) {
return None;
}
let proxy_url = object
.get("url")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())?;
let username = object
.get("username")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
let password = object
.get("password")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
Some(ProxySnapshot {
enabled: Some(true),
mode: object
.get("mode")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.or_else(|| admin_provider_transport_proxy_mode(Some(proxy_url))),
node_id: None,
label: object
.get("label")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
url: admin_provider_transport_inject_proxy_auth(proxy_url, username, password)
.or_else(|| Some(proxy_url.to_string())),
extra: None,
})
}
fn admin_provider_transport_inject_proxy_auth(
@@ -354,3 +343,50 @@ fn admin_provider_transport_string_field(config: &Map<String, Value>, key: &str)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
#[cfg(test)]
mod tests {
use aether_contracts::ProxySnapshot;
use serde_json::json;
use super::admin_provider_transport_proxy_snapshot;
#[test]
fn connector_proxy_snapshot_requires_object_value() {
assert_eq!(
admin_provider_transport_proxy_snapshot(&json!("http://proxy.example:8080")),
None
);
}
#[test]
fn connector_proxy_snapshot_requires_url_field() {
assert_eq!(
admin_provider_transport_proxy_snapshot(&json!({
"proxy_url": "http://proxy.example:8080"
})),
None
);
}
#[test]
fn connector_proxy_snapshot_keeps_current_object_shape() {
assert_eq!(
admin_provider_transport_proxy_snapshot(&json!({
"url": "http://proxy.example:8080",
"username": "alice",
"password": "secret",
"mode": "http",
"label": "manual"
})),
Some(ProxySnapshot {
enabled: Some(true),
mode: Some("http".to_string()),
node_id: None,
label: Some("manual".to_string()),
url: Some("http://alice:secret@proxy.example:8080/".to_string()),
extra: None,
})
);
}
}

View File

@@ -46,6 +46,7 @@ use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
const ADMIN_SYSTEM_IMPORT_MAX_SIZE_BYTES: usize = 10 * 1024 * 1024;
const MIN_ADMIN_SYSTEM_IMPORT_VERSION: (u32, u32) = (2, 2);
fn invalid_request(detail: impl Into<String>) -> (http::StatusCode, Value) {
(
@@ -342,19 +343,33 @@ struct ImportedWalletTarget {
updated_at_unix_secs: Option<u64>,
}
fn imported_users_export_is_legacy(version: Option<&Value>) -> bool {
fn imported_system_export_version(version: Option<&Value>) -> Result<(u32, u32), String> {
let Some(Value::String(version)) = version else {
return true;
return Err("version 必须是 x.y 字符串".to_string());
};
let version = version.trim();
if version.is_empty() {
return Err("version 必须是 x.y 字符串".to_string());
}
let mut parts = version.split('.');
let Some(major) = parts.next().and_then(|value| value.parse::<u32>().ok()) else {
return true;
return Err("version 必须是 x.y 字符串".to_string());
};
let Some(minor) = parts.next().and_then(|value| value.parse::<u32>().ok()) else {
return true;
return Err("version 必须是 x.y 字符串".to_string());
};
(major, minor) < (1, 3)
Ok((major, minor))
}
fn validate_imported_system_export_version(version: Option<&Value>) -> Result<(), String> {
let parsed = imported_system_export_version(version)?;
if parsed < MIN_ADMIN_SYSTEM_IMPORT_VERSION {
return Err(format!(
"version {}.{} 已不再支持;仅支持 2.2+ 导出格式",
parsed.0, parsed.1
));
}
Ok(())
}
fn imported_object_field<'a>(
@@ -385,11 +400,6 @@ fn imported_optional_bool(value: Option<&Value>) -> Result<Option<bool>, String>
match value {
None | Some(Value::Null) => Ok(None),
Some(Value::Bool(value)) => Ok(Some(*value)),
Some(Value::String(raw)) => match raw.trim().to_ascii_lowercase().as_str() {
"true" => Ok(Some(true)),
"false" => Ok(Some(false)),
_ => Err("字段必须是布尔值".to_string()),
},
_ => Err("字段必须是布尔值".to_string()),
}
}
@@ -402,11 +412,6 @@ fn imported_optional_i32(value: Option<&Value>, field_name: &str) -> Result<Opti
.ok_or_else(|| format!("{field_name} 必须是整数"))
.and_then(|value| i32::try_from(value).map_err(|_| format!("{field_name} 超出范围")))
.map(Some),
Some(Value::String(raw)) => raw
.trim()
.parse::<i32>()
.map(Some)
.map_err(|_| format!("{field_name} 必须是整数")),
_ => Err(format!("{field_name} 必须是整数")),
}
}
@@ -418,11 +423,6 @@ fn imported_optional_u64(value: Option<&Value>, field_name: &str) -> Result<Opti
.as_u64()
.ok_or_else(|| format!("{field_name} 必须是非负整数"))
.map(Some),
Some(Value::String(raw)) => raw
.trim()
.parse::<u64>()
.map(Some)
.map_err(|_| format!("{field_name} 必须是非负整数")),
_ => Err(format!("{field_name} 必须是非负整数")),
}
}
@@ -435,13 +435,6 @@ fn imported_optional_f64(value: Option<&Value>, field_name: &str) -> Result<Opti
.filter(|value| value.is_finite())
.ok_or_else(|| format!("{field_name} 必须是有限数值"))
.map(Some),
Some(Value::String(raw)) => raw
.trim()
.parse::<f64>()
.ok()
.filter(|value| value.is_finite())
.ok_or_else(|| format!("{field_name} 必须是有限数值"))
.map(Some),
_ => Err(format!("{field_name} 必须是有限数值")),
}
}
@@ -491,16 +484,6 @@ fn imported_string_list_from_value(
.map(ToOwned::to_owned)
.collect(),
)),
Value::String(raw) => {
let trimmed = raw.trim();
if trimmed.is_empty() || trimmed.eq_ignore_ascii_case("null") {
return Ok(None);
}
if let Ok(decoded) = serde_json::from_str::<Value>(trimmed) {
return imported_string_list_from_value(Some(&decoded), field_name);
}
Ok(Some(vec![trimmed.to_string()]))
}
_ => Err(format!("{field_name} 必须是字符串列表")),
}
}
@@ -1584,7 +1567,6 @@ impl<'a> AdminAppState<'a> {
)))
}
};
let legacy_export = imported_users_export_is_legacy(root.get("version"));
let empty = Vec::new();
let users = match root.get("users") {
Some(Value::Array(items)) => items,
@@ -1614,6 +1596,8 @@ impl<'a> AdminAppState<'a> {
};
}
invalid_value!(validate_imported_system_export_version(root.get("version")));
let mut stats = AdminSystemUsersImportStats::default();
for (index, raw_user) in users.iter().enumerate() {
@@ -1883,7 +1867,7 @@ impl<'a> AdminAppState<'a> {
invalid_value!(normalize_imported_user_string_list(key, "allowed_models"));
let rate_limit =
invalid_value!(imported_optional_i32(key.get("rate_limit"), "rate_limit"))
.unwrap_or(if legacy_export { 0 } else { 0 });
.unwrap_or(0);
let concurrent_limit = invalid_value!(imported_optional_i32(
key.get("concurrent_limit"),
"concurrent_limit"
@@ -2063,7 +2047,7 @@ impl<'a> AdminAppState<'a> {
invalid_value!(normalize_imported_user_string_list(key, "allowed_models"));
let rate_limit =
invalid_value!(imported_optional_i32(key.get("rate_limit"), "rate_limit"))
.unwrap_or(if legacy_export { 0 } else { 0 });
.unwrap_or(0);
let concurrent_limit = invalid_value!(imported_optional_i32(
key.get("concurrent_limit"),
"concurrent_limit"
@@ -2321,3 +2305,64 @@ enum WalletOwner<'a> {
User(&'a str),
ApiKey(&'a str),
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{
imported_optional_bool, imported_optional_f64, imported_optional_i32,
imported_optional_u64, imported_string_list_from_value,
validate_imported_system_export_version,
};
#[test]
fn import_requires_supported_export_version() {
assert!(validate_imported_system_export_version(Some(&json!("2.2"))).is_ok());
assert_eq!(
validate_imported_system_export_version(Some(&json!("2.1"))).unwrap_err(),
"version 2.1 已不再支持;仅支持 2.2+ 导出格式"
);
assert_eq!(
validate_imported_system_export_version(Some(&json!(null))).unwrap_err(),
"version 必须是 x.y 字符串"
);
}
#[test]
fn import_rejects_legacy_string_scalars() {
assert_eq!(
imported_optional_bool(Some(&json!("true"))).unwrap_err(),
"字段必须是布尔值"
);
assert_eq!(
imported_optional_i32(Some(&json!("5")), "rate_limit").unwrap_err(),
"rate_limit 必须是整数"
);
assert_eq!(
imported_optional_u64(Some(&json!("5")), "total_requests").unwrap_err(),
"total_requests 必须是非负整数"
);
assert_eq!(
imported_optional_f64(Some(&json!("1.25")), "total_cost_usd").unwrap_err(),
"total_cost_usd 必须是有限数值"
);
}
#[test]
fn import_rejects_legacy_string_lists() {
assert_eq!(
imported_string_list_from_value(Some(&json!("openai")), "allowed_providers")
.unwrap_err(),
"allowed_providers 必须是字符串列表"
);
assert_eq!(
imported_string_list_from_value(
Some(&json!("[\"openai:chat\"]")),
"allowed_api_formats"
)
.unwrap_err(),
"allowed_api_formats 必须是字符串列表"
);
}
}

View File

@@ -6,4 +6,6 @@ mod proxy_nodes;
mod routes;
pub(super) mod shared;
#[cfg(test)]
pub(crate) use self::proxy_nodes::override_proxy_connectivity_probe_url_for_tests;
pub(super) use self::routes::maybe_build_local_admin_system_response;

View File

@@ -1,5 +1,6 @@
use std::time::{Duration, Instant};
use crate::execution_runtime::transport::format_upstream_request_error;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::query_param_value;
use crate::maintenance::{
@@ -12,6 +13,9 @@ use aether_admin::system::{
admin_proxy_node_event_node_id_from_path, build_admin_proxy_node_payload,
build_admin_proxy_nodes_data_unavailable_response, build_admin_proxy_nodes_not_found_response,
};
use aether_contracts::tunnel::{
TUNNEL_RELAY_FORWARDED_BY_HEADER, TUNNEL_RELAY_OWNER_INSTANCE_HEADER,
};
use axum::{
body::{Body, Bytes},
http,
@@ -21,7 +25,6 @@ use axum::{
use serde::de::DeserializeOwned;
use serde::Deserialize;
use serde_json::{json, Value};
use tokio::{net::TcpStream, time::timeout};
#[derive(Debug, Deserialize)]
struct ProxyNodeRegisterRequest {
@@ -131,6 +134,61 @@ const JSON_OBJECT_REQUIRED_DETAIL: &str = "请求体必须是合法的 JSON 对
const DEFAULT_PROXY_UPGRADE_BATCH_SIZE: usize = 1;
const DEFAULT_PROXY_UPGRADE_COOLDOWN_SECS: u64 = 60;
const DEFAULT_PROXY_UPGRADE_PROBE_TIMEOUT_SECS: u64 = 10;
const DEFAULT_PROXY_CONNECTIVITY_PROBE_URL: &str = "https://www.cloudflare.com/cdn-cgi/trace";
const PROXY_CONNECTIVITY_TIMEOUT_SECS: u64 = 10;
const TUNNEL_RELAY_ENVELOPE_CONTENT_TYPE: &str = "application/vnd.aether.tunnel-envelope";
const MAX_PROXY_CONNECTIVITY_RESPONSE_BYTES: usize = 64 * 1024;
#[cfg(test)]
fn manual_proxy_connectivity_probe_url_override() -> &'static std::sync::RwLock<Option<String>> {
static OVERRIDE: std::sync::OnceLock<std::sync::RwLock<Option<String>>> =
std::sync::OnceLock::new();
OVERRIDE.get_or_init(|| std::sync::RwLock::new(None))
}
#[cfg(test)]
fn manual_proxy_connectivity_probe_url_override_lock() -> &'static std::sync::Mutex<()> {
static LOCK: std::sync::OnceLock<std::sync::Mutex<()>> = std::sync::OnceLock::new();
LOCK.get_or_init(|| std::sync::Mutex::new(()))
}
fn proxy_connectivity_probe_url() -> String {
#[cfg(test)]
if let Some(url) = manual_proxy_connectivity_probe_url_override()
.read()
.expect("probe url override lock should read")
.clone()
{
return url;
}
DEFAULT_PROXY_CONNECTIVITY_PROBE_URL.to_string()
}
#[cfg(test)]
pub(crate) struct ProxyConnectivityProbeUrlOverrideGuard(std::sync::MutexGuard<'static, ()>);
#[cfg(test)]
pub(crate) fn override_proxy_connectivity_probe_url_for_tests(
url: impl Into<String>,
) -> ProxyConnectivityProbeUrlOverrideGuard {
let guard = manual_proxy_connectivity_probe_url_override_lock()
.lock()
.expect("probe url override lock should acquire");
*manual_proxy_connectivity_probe_url_override()
.write()
.expect("probe url override lock should write") = Some(url.into());
ProxyConnectivityProbeUrlOverrideGuard(guard)
}
#[cfg(test)]
impl Drop for ProxyConnectivityProbeUrlOverrideGuard {
fn drop(&mut self) {
*manual_proxy_connectivity_probe_url_override()
.write()
.expect("probe url override lock should write") = None;
}
}
pub(crate) async fn maybe_build_local_admin_proxy_nodes_response(
state: &AdminAppState<'_>,
@@ -187,6 +245,26 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response(
));
}
if decision.route_kind.as_deref() == Some("get_node")
&& request_context.method() == http::Method::GET
{
if !state.has_proxy_node_reader() {
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
}
let Some(node_id) = admin_proxy_node_node_id_from_path(request_context.path()) else {
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
};
let Some(node) = state.find_proxy_node(&node_id).await? else {
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
};
return Ok(Some(
Json(json!({
"node": build_admin_proxy_node_detail_payload(&node),
}))
.into_response(),
));
}
if decision.route_kind.as_deref() == Some("register_node")
&& request_context.method() == http::Method::POST
{
@@ -367,7 +445,7 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response(
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
};
return Ok(Some(
Json(test_proxy_node_connectivity(&node).await).into_response(),
Json(test_proxy_node_connectivity(state, &node).await).into_response(),
));
}
@@ -383,15 +461,7 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response(
Err(response) => return Ok(Some(response)),
};
return Ok(Some(
Json(
test_manual_proxy_connectivity(
&normalized.proxy_url,
normalized.host.as_str(),
normalized.port,
)
.await,
)
.into_response(),
Json(test_manual_proxy_connectivity(&normalized).await).into_response(),
));
}
@@ -716,6 +786,18 @@ struct DeletedProxyNodeCleanup {
cleared_keys: usize,
}
fn build_admin_proxy_node_detail_payload(
node: &aether_data::repository::proxy_nodes::StoredProxyNode,
) -> Value {
let mut payload = build_admin_proxy_node_payload(node);
if node.is_manual {
if let Value::Object(object) = &mut payload {
object.insert("proxy_password".to_string(), json!(node.proxy_password));
}
}
payload
}
#[derive(Debug, Clone)]
struct NormalizedManualProxyEndpoint {
proxy_url: String,
@@ -832,6 +914,7 @@ fn proxy_reference_matches_node_id(value: Option<&Value>, node_id: &str) -> bool
}
async fn test_proxy_node_connectivity(
state: &AdminAppState<'_>,
node: &aether_data::repository::proxy_nodes::StoredProxyNode,
) -> Value {
if node.is_manual {
@@ -854,12 +937,13 @@ async fn test_proxy_node_connectivity(
});
}
};
return test_manual_proxy_connectivity(
let proxy_url = proxy_url_with_auth(
&endpoint.proxy_url,
endpoint.host.as_str(),
endpoint.port,
node.proxy_username.as_deref(),
node.proxy_password.as_deref(),
)
.await;
.unwrap_or(endpoint.proxy_url);
return test_manual_proxy_connectivity(&proxy_url).await;
}
if !node.tunnel_mode {
@@ -880,39 +964,269 @@ async fn test_proxy_node_connectivity(
});
}
let probe_url = proxy_connectivity_probe_url();
match probe_tunnel_proxy_connectivity(state.app(), &node.id, &probe_url).await {
Ok(result) => {
if let Ok(status) = reqwest::StatusCode::from_u16(result.status) {
if status.is_success() {
return json!({
"success": true,
"latency_ms": result.latency_ms,
"exit_ip": parse_proxy_probe_exit_ip(&result.body),
"error": null,
});
}
return json!({
"success": false,
"latency_ms": null,
"exit_ip": null,
"error": sanitize_proxy_error(&format_proxy_probe_status_error(status, &result.body)),
});
}
json!({
"success": false,
"latency_ms": null,
"exit_ip": null,
"error": format!("代理探测返回非法状态码: {}", result.status),
})
}
Err(error) => json!({
"success": false,
"latency_ms": null,
"exit_ip": null,
"error": sanitize_proxy_error(&error),
}),
}
}
async fn test_manual_proxy_connectivity(proxy_url: &str) -> Value {
let probe_url = proxy_connectivity_probe_url();
test_manual_proxy_connectivity_with_probe_url(proxy_url, &probe_url).await
}
async fn test_manual_proxy_connectivity_with_probe_url(proxy_url: &str, probe_url: &str) -> Value {
let started_at = Instant::now();
let proxy = match reqwest::Proxy::all(proxy_url) {
Ok(proxy) => proxy,
Err(error) => {
return json!({
"success": false,
"latency_ms": null,
"exit_ip": null,
"error": sanitize_proxy_error(&format_upstream_request_error(&error)),
});
}
};
let mut builder = reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(Duration::from_secs(5))
.timeout(Duration::from_secs(PROXY_CONNECTIVITY_TIMEOUT_SECS))
.proxy(proxy)
.user_agent("aether-gateway/proxy-connectivity");
if proxy_url
.trim()
.to_ascii_lowercase()
.starts_with("https://")
{
builder = builder.danger_accept_invalid_certs(true);
}
let client = match builder.build() {
Ok(client) => client,
Err(error) => {
return json!({
"success": false,
"latency_ms": null,
"exit_ip": null,
"error": sanitize_proxy_error(&format_upstream_request_error(&error)),
});
}
};
let response = match client.get(probe_url).send().await {
Ok(response) => response,
Err(error) => {
return json!({
"success": false,
"latency_ms": null,
"exit_ip": null,
"error": sanitize_proxy_error(&format_upstream_request_error(&error)),
});
}
};
let status = response.status();
let body = match response.text().await {
Ok(body) => body,
Err(error) => {
return json!({
"success": false,
"latency_ms": null,
"exit_ip": null,
"error": sanitize_proxy_error(&format_upstream_request_error(&error)),
});
}
};
if !status.is_success() {
return json!({
"success": false,
"latency_ms": null,
"exit_ip": null,
"error": sanitize_proxy_error(&format_proxy_probe_status_error(status, &body)),
});
}
json!({
"success": true,
"latency_ms": node.avg_latency_ms.map(|value| value.max(0.0).round() as u64),
"exit_ip": null,
"latency_ms": started_at.elapsed().as_millis() as u64,
"exit_ip": parse_proxy_probe_exit_ip(&body),
"error": null,
})
}
async fn test_manual_proxy_connectivity(_proxy_url: &str, host: &str, port: u16) -> Value {
let started_at = Instant::now();
match timeout(Duration::from_secs(5), TcpStream::connect((host, port))).await {
Ok(Ok(stream)) => {
drop(stream);
json!({
"success": true,
"latency_ms": started_at.elapsed().as_millis() as u64,
"exit_ip": null,
"error": null,
})
}
Ok(Err(error)) => json!({
"success": false,
"latency_ms": null,
"exit_ip": null,
"error": sanitize_proxy_error(&error.to_string()),
}),
Err(_) => json!({
"success": false,
"latency_ms": null,
"exit_ip": null,
"error": "连接超时",
}),
struct TunnelConnectivityProbeResult {
status: u16,
body: String,
latency_ms: u64,
}
async fn probe_tunnel_proxy_connectivity(
state: &crate::AppState,
node_id: &str,
probe_url: &str,
) -> Result<TunnelConnectivityProbeResult, String> {
let trimmed_node_id = node_id.trim();
if trimmed_node_id.is_empty() {
return Err("proxy node id is empty".to_string());
}
if state.tunnel.has_local_proxy(trimmed_node_id) {
return probe_tunnel_proxy_connectivity_locally(state, trimmed_node_id, probe_url).await;
}
if let Some(owner) = state
.tunnel
.lookup_attachment_owner(state.data.as_ref(), trimmed_node_id)
.await
.map_err(|err| format!("lookup tunnel attachment owner failed: {err}"))?
{
if owner.gateway_instance_id != state.tunnel.local_instance_id() {
return probe_tunnel_proxy_connectivity_via_owner(
state,
trimmed_node_id,
probe_url,
&owner.relay_base_url,
&owner.gateway_instance_id,
)
.await;
}
state
.tunnel
.clear_local_attachment_if_stale(state.data.as_ref(), trimmed_node_id)
.await
.map_err(|err| format!("clear stale local tunnel attachment failed: {err}"))?;
}
probe_tunnel_proxy_connectivity_locally(state, trimmed_node_id, probe_url).await
}
async fn probe_tunnel_proxy_connectivity_locally(
state: &crate::AppState,
node_id: &str,
probe_url: &str,
) -> Result<TunnelConnectivityProbeResult, String> {
let started_at = Instant::now();
let result = state
.tunnel
.probe_node_url_with_response(node_id, probe_url, PROXY_CONNECTIVITY_TIMEOUT_SECS)
.await?;
Ok(TunnelConnectivityProbeResult {
status: result.status,
body: result.body,
latency_ms: started_at.elapsed().as_millis() as u64,
})
}
async fn probe_tunnel_proxy_connectivity_via_owner(
state: &crate::AppState,
node_id: &str,
probe_url: &str,
relay_base_url: &str,
owner_instance_id: &str,
) -> Result<TunnelConnectivityProbeResult, String> {
let owner_url = build_tunnel_owner_relay_url(relay_base_url, node_id)?;
let started_at = Instant::now();
let response = state
.client
.post(owner_url)
.header(
http::header::CONTENT_TYPE,
TUNNEL_RELAY_ENVELOPE_CONTENT_TYPE,
)
.header(
TUNNEL_RELAY_FORWARDED_BY_HEADER,
state.tunnel.local_instance_id(),
)
.header(TUNNEL_RELAY_OWNER_INSTANCE_HEADER, owner_instance_id)
.timeout(Duration::from_secs(PROXY_CONNECTIVITY_TIMEOUT_SECS))
.body(build_tunnel_probe_relay_envelope(probe_url)?)
.send()
.await
.map_err(|error| format!("owner tunnel relay probe failed: {error}"))?;
let status = response.status();
let body = response
.bytes()
.await
.map_err(|error| format!("failed to read owner tunnel relay probe body: {error}"))?;
if body.len() > MAX_PROXY_CONNECTIVITY_RESPONSE_BYTES {
return Err(format!(
"owner tunnel relay probe body exceeds {} bytes",
MAX_PROXY_CONNECTIVITY_RESPONSE_BYTES
));
}
Ok(TunnelConnectivityProbeResult {
status: status.as_u16(),
body: String::from_utf8_lossy(&body).to_string(),
latency_ms: started_at.elapsed().as_millis() as u64,
})
}
fn build_tunnel_probe_relay_envelope(probe_url: &str) -> Result<Vec<u8>, String> {
let meta = crate::tunnel::tunnel_protocol::RequestMeta {
method: "GET".to_string(),
url: probe_url.trim().to_string(),
headers: std::collections::HashMap::new(),
timeout: PROXY_CONNECTIVITY_TIMEOUT_SECS,
follow_redirects: Some(false),
http1_only: false,
};
let meta_bytes = serde_json::to_vec(&meta)
.map_err(|error| format!("encode tunnel probe metadata failed: {error}"))?;
let mut envelope = Vec::with_capacity(4 + meta_bytes.len());
envelope.extend_from_slice(&(meta_bytes.len() as u32).to_be_bytes());
envelope.extend_from_slice(&meta_bytes);
Ok(envelope)
}
fn build_tunnel_owner_relay_url(relay_base_url: &str, node_id: &str) -> Result<String, String> {
let mut url = url::Url::parse(relay_base_url)
.map_err(|error| format!("invalid owner relay base url: {error}"))?;
{
let mut segments = url
.path_segments_mut()
.map_err(|_| "owner relay base url cannot be a base-less URL".to_string())?;
segments.pop_if_empty();
segments.push("api");
segments.push("internal");
segments.push("tunnel");
segments.push("relay");
segments.push(node_id.trim());
}
Ok(url.to_string())
}
fn validate_register_request(
@@ -1039,10 +1353,16 @@ fn validate_manual_update_request(
fn validate_proxy_test_url_request(
input: ProxyNodeTestUrlRequest,
) -> Result<NormalizedManualProxyEndpoint, Response<Body>> {
let _ = normalize_optional_string(input.username.as_deref(), "username", 255)?;
let _ = normalize_optional_string(input.password.as_deref(), "password", 500)?;
normalize_manual_proxy_endpoint(&input.proxy_url)
) -> Result<String, Response<Body>> {
let username = normalize_optional_string(input.username.as_deref(), "username", 255)?;
let password = normalize_optional_string(input.password.as_deref(), "password", 500)?;
let endpoint = normalize_manual_proxy_endpoint(&input.proxy_url)?;
Ok(proxy_url_with_auth(
&endpoint.proxy_url,
username.as_deref(),
password.as_deref(),
)
.unwrap_or(endpoint.proxy_url))
}
fn admin_proxy_node_upgrade_action_node_id_from_path(path: &str, suffix: &str) -> Option<String> {
@@ -1425,6 +1745,53 @@ fn sanitize_proxy_error(detail: &str) -> String {
}
}
fn proxy_url_with_auth(
proxy_url: &str,
username: Option<&str>,
password: Option<&str>,
) -> Option<String> {
let username = username.map(str::trim).filter(|value| !value.is_empty())?;
let mut parsed = url::Url::parse(proxy_url).ok()?;
if parsed.set_username(username).is_err() {
return None;
}
let password = password.map(str::trim).filter(|value| !value.is_empty());
if parsed.set_password(password).is_err() {
return None;
}
Some(parsed.to_string())
}
fn parse_proxy_probe_exit_ip(body: &str) -> Option<String> {
body.lines().find_map(|line| {
let (key, value) = line.split_once('=')?;
if key.trim() != "ip" {
return None;
}
let value = value.trim();
if value.is_empty() {
return None;
}
Some(value.to_string())
})
}
fn format_proxy_probe_status_error(status: reqwest::StatusCode, body: &str) -> String {
let body = body.trim();
if body.is_empty() {
return format!("代理探测返回 HTTP {}", status.as_u16());
}
let truncated = if body.chars().count() > 200 {
let shortened: String = body.chars().take(200).collect();
format!("{shortened}...")
} else {
body.to_string()
};
format!("代理探测返回 HTTP {}: {truncated}", status.as_u16())
}
fn validate_optional_counter(value: Option<i64>, field: &str) -> Result<(), Response<Body>> {
if value.is_some_and(|value| value < 0) {
return Err(bad_request_response(format!("{field} 必须是非负整数")));

View File

@@ -351,12 +351,10 @@ pub(crate) fn gateway_error_message(error: GatewayError) -> String {
pub(crate) fn build_internal_tunnel_heartbeat_ack(
node: &StoredProxyNode,
heartbeat_id: Option<u64>,
heartbeat_id: u64,
) -> serde_json::Value {
let mut payload = serde_json::Map::new();
if let Some(heartbeat_id) = heartbeat_id {
payload.insert("heartbeat_id".to_string(), json!(heartbeat_id));
}
payload.insert("heartbeat_id".to_string(), json!(heartbeat_id));
if let Some(remote_config) = node.remote_config.as_ref() {
payload.insert("remote_config".to_string(), remote_config.clone());
payload.insert("config_version".to_string(), json!(node.config_version));
@@ -385,7 +383,7 @@ pub(crate) fn parse_internal_tunnel_heartbeat_request(
})?;
let node_id = payload.node_id.trim();
if node_id.is_empty() || node_id.len() > 36 {
if node_id.is_empty() || node_id.len() > 36 || payload.heartbeat_id == 0 {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"invalid heartbeat payload",

View File

@@ -4,8 +4,7 @@ use std::collections::BTreeMap;
#[derive(Debug, Deserialize)]
pub(crate) struct InternalTunnelHeartbeatRequest {
pub(crate) node_id: String,
#[serde(default)]
pub(crate) heartbeat_id: Option<u64>,
pub(crate) heartbeat_id: u64,
#[serde(default)]
pub(crate) heartbeat_interval: Option<i32>,
#[serde(default)]

View File

@@ -87,9 +87,9 @@ pub(super) fn postgres_error(
const AUDIT_LOG_CLEANUP_INTERVAL: Duration = Duration::from_secs(24 * 60 * 60);
const GEMINI_FILE_MAPPING_CLEANUP_INTERVAL: Duration = Duration::from_secs(60 * 60);
const PENDING_CLEANUP_INTERVAL: Duration = Duration::from_secs(5 * 60);
const PROXY_NODE_STALE_SWEEP_INTERVAL: Duration = Duration::from_secs(30);
const PROXY_NODE_STALE_SWEEP_INTERVAL: Duration = Duration::from_secs(5);
const PROXY_UPGRADE_ROLLOUT_INTERVAL: Duration = Duration::from_secs(15);
const PROXY_NODE_STALE_MIN_GRACE_SECS: u64 = 90;
const PROXY_NODE_STALE_MIN_GRACE_SECS: u64 = 15;
const PROXY_NODE_STALE_MISSED_HEARTBEATS: u64 = 3;
const POOL_MONITOR_INTERVAL: Duration = Duration::from_secs(5 * 60);
const PROVIDER_CHECKIN_CONCURRENCY: usize = 3;

View File

@@ -92,26 +92,21 @@ impl AppState {
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<ProxySnapshot> {
if let Some(snapshot) = self
.resolve_proxy_snapshot_from_config(transport.key.proxy.as_ref())
.await
{
return Some(snapshot);
}
if let Some(snapshot) = self
.resolve_proxy_snapshot_from_config(transport.provider.proxy.as_ref())
.await
{
return Some(snapshot);
}
if let Some(snapshot) = self.resolve_system_proxy_snapshot().await {
return Some(snapshot);
}
self.resolve_proxy_snapshot_from_config(transport.endpoint.proxy.as_ref())
self.resolve_transport_proxy_with_source_with_tunnel_affinity(transport)
.await
.map(|(snapshot, _)| snapshot)
}
async fn resolve_proxy_snapshot_from_config(
pub(crate) async fn resolve_transport_proxy_source_with_tunnel_affinity(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<&'static str> {
self.resolve_transport_proxy_with_source_with_tunnel_affinity(transport)
.await
.map(|(_, source)| source)
}
pub(crate) async fn resolve_configured_proxy_snapshot_with_tunnel_affinity(
&self,
raw: Option<&Value>,
) -> Option<ProxySnapshot> {
@@ -127,6 +122,37 @@ impl AppState {
proxy_snapshot_from_object(object)
}
async fn resolve_transport_proxy_with_source_with_tunnel_affinity(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<(ProxySnapshot, &'static str)> {
if let Some(snapshot) = self
.resolve_configured_proxy_snapshot_with_tunnel_affinity(transport.key.proxy.as_ref())
.await
{
return Some((snapshot, "key"));
}
if let Some(snapshot) = self
.resolve_configured_proxy_snapshot_with_tunnel_affinity(
transport.endpoint.proxy.as_ref(),
)
.await
{
return Some((snapshot, "endpoint"));
}
if let Some(snapshot) = self
.resolve_configured_proxy_snapshot_with_tunnel_affinity(
transport.provider.proxy.as_ref(),
)
.await
{
return Some((snapshot, "provider"));
}
self.resolve_system_proxy_snapshot()
.await
.map(|snapshot| (snapshot, "system"))
}
}
fn proxy_enabled(object: &Map<String, Value>) -> bool {
@@ -206,12 +232,22 @@ fn proxy_url_with_node_auth(
if parsed.set_username(username).is_err() {
return None;
}
let password = password
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or_default();
if parsed.set_password(Some(password)).is_err() {
let password = password.map(str::trim).filter(|value| !value.is_empty());
if parsed.set_password(password).is_err() {
return None;
}
Some(parsed.to_string())
}
#[cfg(test)]
mod tests {
use super::proxy_url_with_node_auth;
#[test]
fn proxy_url_with_node_auth_omits_empty_password_separator() {
assert_eq!(
proxy_url_with_node_auth("socks5://proxy.example:1080", Some("alice"), None).as_deref(),
Some("socks5://alice@proxy.example:1080")
);
}
}

View File

@@ -3,6 +3,7 @@ use std::sync::{Arc, Mutex};
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogReadRepository, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
@@ -14,7 +15,7 @@ use serde_json::json;
use super::super::super::{
build_router_with_state, build_state_with_execution_runtime_override, sample_endpoint,
sample_key, start_server,
sample_key, sample_proxy_node, start_server,
};
use crate::constants::{
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
@@ -29,6 +30,7 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_codex_with_trusted_a
url: String,
authorization: String,
provider_api_format: String,
total_ms: Option<u64>,
}
let upstream_hits = Arc::new(Mutex::new(0usize));
@@ -67,6 +69,10 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_codex_with_trusted_a
.cloned()
.unwrap_or_default(),
provider_api_format: plan.provider_api_format.clone(),
total_ms: plan
.timeouts
.as_ref()
.and_then(|timeouts| timeouts.total_ms),
});
let result = aether_contracts::ExecutionResult {
request_id: plan.request_id,
@@ -178,6 +184,7 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_codex_with_trusted_a
seen_execution_runtime_request.provider_api_format,
"openai:cli"
);
assert_eq!(seen_execution_runtime_request.total_ms, Some(30_000));
let reloaded = provider_catalog_repository
.list_keys_by_ids(&["key-codex-a".to_string()])
@@ -215,6 +222,143 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_codex_with_trusted_a
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_refreshes_admin_provider_quota_for_codex_proxy_with_extended_timeout() {
let upstream =
Router::new().route(
"/api/admin/endpoints/providers/provider-codex/refresh-quota",
any(|_request: Request| async move {
(StatusCode::OK, Body::from("unexpected upstream hit"))
}),
);
let seen_execution_runtime = Arc::new(Mutex::new(None::<aether_contracts::ExecutionPlan>));
let seen_execution_runtime_clone = Arc::clone(&seen_execution_runtime);
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any(move |request: Request| {
let seen_execution_runtime_inner = Arc::clone(&seen_execution_runtime_clone);
async move {
let plan: aether_contracts::ExecutionPlan = serde_json::from_slice(
&to_bytes(request.into_body(), usize::MAX)
.await
.expect("body should read"),
)
.expect("plan should parse");
*seen_execution_runtime_inner
.lock()
.expect("mutex should lock") = Some(plan.clone());
let result = aether_contracts::ExecutionResult {
request_id: plan.request_id,
candidate_id: None,
status_code: 200,
headers: BTreeMap::new(),
body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({
"plan_type": "plus",
"rate_limit": {
"primary_window": {
"used_percent": 12.5,
"reset_after_seconds": 18000,
"reset_at": 1_900_000_000u64,
"window_minutes": 300
}
}
})),
body_bytes_b64: None,
}),
telemetry: None,
error: None,
};
(StatusCode::OK, Json(result))
}
}),
);
let mut provider = StoredProviderCatalogProvider::new(
"provider-codex".to_string(),
"codex".to_string(),
Some("https://example.com".to_string()),
"codex".to_string(),
)
.expect("provider should build");
provider.proxy = Some(json!({
"node_id": "proxy-node-codex-quota",
"enabled": true
}));
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![sample_endpoint(
"endpoint-codex-cli",
"provider-codex",
"openai:cli",
"https://chatgpt.com/backend-api",
)],
vec![sample_key(
"key-codex-a",
"provider-codex",
"openai:cli",
"sk-codex-123",
)],
));
let mut manual_node = sample_proxy_node("proxy-node-codex-quota");
manual_node.status = "online".to_string();
manual_node.is_manual = true;
manual_node.tunnel_mode = false;
manual_node.tunnel_connected = false;
manual_node.proxy_url = Some("http://proxy.example:8080".to_string());
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![manual_node]));
let (_upstream_url, upstream_handle) = start_server(upstream).await;
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url.clone())
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository,
)
.attach_proxy_node_repository_for_tests(proxy_node_repository)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/endpoints/providers/provider-codex/refresh-quota"
))
.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")
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let plan = seen_execution_runtime
.lock()
.expect("mutex should lock")
.clone()
.expect("execution runtime request should be captured");
assert_eq!(
plan.proxy
.as_ref()
.and_then(|proxy| proxy.node_id.as_deref()),
Some("proxy-node-codex-quota")
);
let timeouts = plan.timeouts.expect("timeouts should exist");
assert_eq!(timeouts.connect_ms, Some(60_000));
assert_eq!(timeouts.read_ms, Some(60_000));
assert_eq!(timeouts.write_ms, Some(60_000));
assert_eq!(timeouts.pool_ms, Some(60_000));
assert_eq!(timeouts.total_ms, Some(60_000));
gateway_handle.abort();
execution_runtime_handle.abort();
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_refreshes_admin_provider_quota_locally_for_kiro_with_trusted_admin_principal() {
#[derive(Debug, Clone)]

View File

@@ -398,10 +398,7 @@ async fn gateway_handles_admin_provider_oauth_device_poll_locally_with_trusted_a
.next()
.expect("persisted key should exist");
assert_eq!(persisted.auth_type, "oauth");
assert_eq!(
persisted.proxy,
Some(json!({"node_id":"proxy-node-kiro","enabled":true}))
);
assert_eq!(persisted.proxy, None);
let decrypted_api_key =
decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, &persisted.encrypted_api_key)
.expect("api key should decrypt");
@@ -1133,10 +1130,7 @@ async fn gateway_batch_imports_admin_provider_oauth_locally_with_trusted_admin_p
.expect("keys should load");
let persisted = reloaded.first().expect("persisted key should exist");
assert!(persisted.is_active);
assert_eq!(
persisted.proxy,
Some(json!({"node_id":"proxy-node-batch-import","enabled":true}))
);
assert_eq!(persisted.proxy, None);
let decrypted_api_key =
decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, &persisted.encrypted_api_key)
.expect("api key should decrypt");
@@ -1604,10 +1598,7 @@ async fn gateway_completes_admin_provider_oauth_provider_locally_with_trusted_ad
.expect("keys should load");
let persisted = reloaded.first().expect("persisted key should exist");
assert!(persisted.is_active);
assert_eq!(
persisted.proxy,
Some(json!({"node_id":"proxy-node-codex-oauth","enabled":true}))
);
assert_eq!(persisted.proxy, None);
let decrypted_api_key =
decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, &persisted.encrypted_api_key)
.expect("api key should decrypt");
@@ -1775,10 +1766,7 @@ async fn gateway_imports_admin_provider_oauth_refresh_token_locally_with_trusted
.expect("keys should load");
let persisted = reloaded.first().expect("persisted key should exist");
assert!(persisted.is_active);
assert_eq!(
persisted.proxy,
Some(json!({"node_id":"proxy-node-codex-import","enabled":true}))
);
assert_eq!(persisted.proxy, None);
let decrypted_api_key =
decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, &persisted.encrypted_api_key)
.expect("api key should decrypt");
@@ -1819,6 +1807,12 @@ async fn gateway_imports_admin_provider_oauth_refresh_token_via_execution_runtim
.push(plan.clone());
let proxy = plan.proxy.as_ref().expect("proxy snapshot should exist");
assert_eq!(proxy.node_id.as_deref(), Some("proxy-node-codex-import"));
let timeouts = plan.timeouts.as_ref().expect("timeouts should exist");
assert_eq!(timeouts.connect_ms, Some(60_000));
assert_eq!(timeouts.read_ms, Some(60_000));
assert_eq!(timeouts.write_ms, Some(60_000));
assert_eq!(timeouts.pool_ms, Some(60_000));
assert_eq!(timeouts.total_ms, Some(60_000));
assert_eq!(plan.request_id, "provider-oauth:refresh-token");
assert_eq!(plan.method, "POST");
assert_eq!(plan.url, "https://oauth.example/oauth/token");
@@ -1920,10 +1914,7 @@ async fn gateway_imports_admin_provider_oauth_refresh_token_via_execution_runtim
.await
.expect("keys should load");
assert_eq!(keys.len(), 1);
assert_eq!(
keys[0].proxy,
Some(json!({"node_id":"proxy-node-codex-import","enabled":true}))
);
assert_eq!(keys[0].proxy, None);
let plans = execution_plans.lock().expect("mutex should lock");
assert_eq!(plans.len(), 1);
@@ -1932,6 +1923,136 @@ async fn gateway_imports_admin_provider_oauth_refresh_token_via_execution_runtim
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_imports_admin_provider_oauth_refresh_token_via_execution_runtime_provider_proxy_before_system_proxy(
) {
let execution_plans = Arc::new(Mutex::new(Vec::<ExecutionPlan>::new()));
let execution_plans_clone = Arc::clone(&execution_plans);
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any(move |Json(plan): Json<ExecutionPlan>| {
let execution_plans_inner = Arc::clone(&execution_plans_clone);
async move {
execution_plans_inner
.lock()
.expect("mutex should lock")
.push(plan.clone());
let proxy = plan.proxy.as_ref().expect("proxy snapshot should exist");
assert_eq!(proxy.node_id.as_deref(), Some("proxy-node-codex-provider"));
Json(json!({
"request_id": plan.request_id,
"status_code": 200,
"headers": {
"content-type": "application/json"
},
"body": {
"json_body": {
"access_token": "imported-codex-access-token",
"refresh_token": "imported-codex-refresh-token",
"token_type": "Bearer",
"expires_in": 1800,
"scope": "openid email profile offline_access",
"email": "alice@example.com",
"account_id": "acct-codex-123",
"plan_type": "plus"
}
}
}))
}
}),
);
let mut provider = sample_provider("provider-codex", "codex", 10);
provider.provider_type = "codex".to_string();
provider.proxy = Some(json!({"node_id":"proxy-node-codex-provider","enabled":true}));
let endpoint = sample_endpoint(
"endpoint-codex-chat",
"provider-codex",
"openai:chat",
"https://chatgpt.com/backend-api/codex",
);
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![],
));
let mut provider_node = sample_proxy_node("proxy-node-codex-provider");
provider_node.status = "online".to_string();
provider_node.is_manual = true;
provider_node.tunnel_mode = false;
provider_node.tunnel_connected = false;
provider_node.proxy_url = Some("http://proxy-provider.example:8080".to_string());
let mut system_node = sample_proxy_node("proxy-node-codex-system");
system_node.status = "online".to_string();
system_node.is_manual = true;
system_node.tunnel_mode = false;
system_node.tunnel_connected = false;
system_node.proxy_url = Some("http://proxy-system.example:8080".to_string());
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
provider_node,
system_node,
]));
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository.clone(),
)
.attach_proxy_node_repository_for_tests(proxy_node_repository)
.with_system_config_values_for_tests(vec![(
"system_proxy_node_id".to_string(),
json!("proxy-node-codex-system"),
)])
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
.with_provider_oauth_token_url_for_tests("codex", "https://oauth.example/oauth/token"),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/provider-oauth/providers/provider-codex/import-refresh-token"
))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"refresh_token": "provider-import-refresh-token",
"name": "codex-import"
}))
.send()
.await
.expect("request should succeed");
let status = response.status();
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(status, StatusCode::OK, "payload={payload}");
assert_eq!(payload["provider_type"], "codex");
let keys = provider_catalog_repository
.list_keys_by_provider_ids(&["provider-codex".to_string()])
.await
.expect("keys should load");
assert_eq!(keys.len(), 1);
assert_eq!(keys[0].proxy, None);
let plans = execution_plans.lock().expect("mutex should lock");
assert_eq!(plans.len(), 1);
assert_eq!(
plans[0]
.proxy
.as_ref()
.and_then(|proxy| proxy.node_id.as_deref()),
Some("proxy-node-codex-provider")
);
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_imports_admin_provider_oauth_refresh_token_via_execution_runtime_system_proxy() {
let execution_plans = Arc::new(Mutex::new(Vec::<ExecutionPlan>::new()));
@@ -2044,6 +2165,71 @@ async fn gateway_imports_admin_provider_oauth_refresh_token_via_execution_runtim
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_import_refresh_token_surfaces_execution_runtime_error_detail() {
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any(|| async { StatusCode::INTERNAL_SERVER_ERROR }),
);
let mut provider = sample_provider("provider-codex", "codex", 10);
provider.provider_type = "codex".to_string();
let endpoint = sample_endpoint(
"endpoint-codex-chat",
"provider-codex",
"openai:chat",
"https://chatgpt.com/backend-api/codex",
);
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![],
));
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository,
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
.with_provider_oauth_token_url_for_tests("codex", "https://oauth.example/oauth/token"),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/provider-oauth/providers/provider-codex/import-refresh-token"
))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"refresh_token": "provider-import-refresh-token",
"name": "codex-import"
}))
.send()
.await
.expect("request should succeed");
let status = response.status();
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(status, StatusCode::BAD_REQUEST, "payload={payload}");
assert!(
payload["detail"]
.as_str()
.expect("detail should be string")
.contains("execution runtime returned HTTP 500"),
"payload={payload}"
);
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_batch_imports_admin_provider_oauth_kiro_locally_with_trusted_admin_principal() {
let upstream_hits = Arc::new(Mutex::new(0usize));
@@ -2162,10 +2348,7 @@ async fn gateway_batch_imports_admin_provider_oauth_kiro_locally_with_trusted_ad
.next()
.expect("persisted key should exist");
assert!(stored_key.is_active);
assert_eq!(
stored_key.proxy,
Some(json!({"node_id":"proxy-node-kiro-batch","enabled":true}))
);
assert_eq!(stored_key.proxy, None);
let decrypted_auth_config = decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY,
stored_key
@@ -2324,10 +2507,7 @@ async fn gateway_batch_imports_admin_provider_oauth_kiro_via_execution_runtime_p
.next()
.expect("persisted key should exist");
assert!(stored_key.is_active);
assert_eq!(
stored_key.proxy,
Some(json!({"node_id":"proxy-node-kiro-batch-runtime","enabled":true}))
);
assert_eq!(stored_key.proxy, None);
let decrypted_auth_config = decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY,
stored_key
@@ -2653,6 +2833,160 @@ async fn gateway_refreshes_admin_provider_oauth_key_locally_with_trusted_admin_p
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_refreshes_admin_provider_oauth_key_locally_via_execution_runtime_provider_proxy_before_system_proxy(
) {
let execution_plans = Arc::new(Mutex::new(Vec::<ExecutionPlan>::new()));
let execution_plans_clone = Arc::clone(&execution_plans);
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any(move |Json(plan): Json<ExecutionPlan>| {
let execution_plans_inner = Arc::clone(&execution_plans_clone);
async move {
execution_plans_inner
.lock()
.expect("mutex should lock")
.push(plan.clone());
if plan.request_id == "provider-oauth:local-refresh-token" {
Json(json!({
"request_id": plan.request_id,
"status_code": 200,
"headers": {
"content-type": "application/json"
},
"body": {
"json_body": {
"access_token": "refreshed-codex-access-token",
"refresh_token": "refreshed-codex-refresh-token",
"token_type": "Bearer",
"expires_in": 1800,
"scope": "openid email profile offline_access",
"email": "alice@example.com",
"account_id": "acct-codex-123",
"plan_type": "plus"
}
}
}))
} else {
Json(json!({
"request_id": plan.request_id,
"status_code": 200,
"headers": {
"content-type": "application/json"
},
"body": {
"json_body": {}
}
}))
}
}
}),
);
let mut provider = sample_provider("provider-codex", "codex", 10);
provider.provider_type = "codex".to_string();
provider.proxy = Some(json!({"node_id":"proxy-node-provider","enabled":true}));
let endpoint = sample_endpoint(
"endpoint-codex-cli",
"provider-codex",
"openai:cli",
"https://chatgpt.com/backend-api/codex",
);
let mut key = sample_key(
"key-codex-oauth-refresh-provider",
"provider-codex",
"openai:cli",
"stale-codex-access-token",
);
key.auth_type = "oauth".to_string();
key.encrypted_auth_config = Some(
encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"provider_type":"codex","refresh_token":"old-codex-refresh-token","email":"alice@example.com","account_id":"acct-codex-123","plan_type":"plus","expires_at":1}"#,
)
.expect("auth config ciphertext should build"),
);
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
let mut provider_node = sample_proxy_node("proxy-node-provider");
provider_node.status = "online".to_string();
provider_node.is_manual = true;
provider_node.tunnel_mode = false;
provider_node.tunnel_connected = false;
provider_node.proxy_url = Some("http://proxy-provider.example:8080".to_string());
let mut system_node = sample_proxy_node("proxy-node-system");
system_node.status = "online".to_string();
system_node.is_manual = true;
system_node.tunnel_mode = false;
system_node.tunnel_connected = false;
system_node.proxy_url = Some("http://proxy-system.example:8080".to_string());
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
provider_node,
system_node,
]));
let oauth_refresh =
crate::provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![
Arc::new(
crate::provider_transport::oauth_refresh::GenericOAuthRefreshAdapter::default()
.with_token_url_for_tests("codex", "https://oauth.example/oauth/token"),
),
]);
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository.clone(),
)
.attach_proxy_node_repository_for_tests(proxy_node_repository)
.with_system_config_values_for_tests(vec![(
"system_proxy_node_id".to_string(),
json!("proxy-node-system"),
)])
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
.with_oauth_refresh_coordinator_for_tests(oauth_refresh),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/provider-oauth/keys/key-codex-oauth-refresh-provider/refresh"
))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let plans = execution_plans.lock().expect("mutex should lock");
let refresh_plan = plans
.iter()
.find(|plan| plan.request_id == "provider-oauth:local-refresh-token")
.expect("local refresh plan should exist");
assert_eq!(
refresh_plan
.proxy
.as_ref()
.and_then(|proxy| proxy.node_id.as_deref()),
Some("proxy-node-provider")
);
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_refreshes_admin_provider_oauth_key_locally_via_execution_runtime_key_proxy_before_system_proxy(
) {

View File

@@ -6,11 +6,13 @@ use aether_data::repository::proxy_nodes::{
InMemoryProxyNodeRepository, ProxyNodeHeartbeatMutation, StoredProxyNodeEvent,
};
use axum::body::Body;
use axum::extract::ws::Message;
use axum::routing::any;
use axum::{extract::Request, Router};
use base64::Engine as _;
use http::StatusCode;
use serde_json::json;
use tokio::net::TcpListener;
use tokio::sync::watch;
use super::super::{
build_router_with_state, hash_management_token, sample_endpoint, sample_key,
@@ -25,6 +27,7 @@ use crate::maintenance::{
record_proxy_upgrade_traffic_success, skip_proxy_upgrade_rollout_node,
start_proxy_upgrade_rollout,
};
use crate::tunnel::{tunnel_protocol, TunnelProxyConn};
#[tokio::test]
async fn gateway_handles_admin_proxy_nodes_locally_with_trusted_admin_principal() {
@@ -107,6 +110,49 @@ async fn gateway_handles_admin_proxy_nodes_locally_with_trusted_admin_principal(
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_returns_full_manual_proxy_node_detail_locally_with_trusted_admin_principal() {
let mut manual_node = sample_proxy_node("proxy-node-manual");
manual_node.name = "alpha-manual".to_string();
manual_node.status = "online".to_string();
manual_node.is_manual = true;
manual_node.tunnel_mode = false;
manual_node.tunnel_connected = false;
manual_node.proxy_url = Some("http://proxy.example:8080".to_string());
manual_node.proxy_username = Some("alice".to_string());
manual_node.proxy_password = Some("supersecret".to_string());
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![manual_node]));
let gateway = build_router_with_state(
AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(GatewayDataState::with_proxy_node_repository_for_tests(
proxy_node_repository,
)),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.get(format!(
"{gateway_url}/api/admin/proxy-nodes/proxy-node-manual"
))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["node"]["id"], "proxy-node-manual");
assert_eq!(payload["node"]["proxy_username"], "alice");
assert_eq!(payload["node"]["proxy_password"], "supersecret");
gateway_handle.abort();
}
#[tokio::test]
async fn gateway_reports_active_proxy_upgrade_rollout_in_proxy_node_list() {
let mut alpha = sample_proxy_node("node-alpha");
@@ -719,15 +765,28 @@ async fn gateway_registers_and_unregisters_proxy_nodes_locally_with_management_t
#[tokio::test]
async fn gateway_creates_updates_and_tests_manual_proxy_nodes_locally() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let proxy_port = listener
.local_addr()
.expect("listener addr should resolve")
.port();
let accept_handle =
tokio::spawn(async move { while let Ok((_stream, _addr)) = listener.accept().await {} });
let proxy_auths = Arc::new(Mutex::new(Vec::<Option<String>>::new()));
let proxy_auths_clone = Arc::clone(&proxy_auths);
let proxy = Router::new().fallback(any(move |request: Request| {
let proxy_auths_inner = Arc::clone(&proxy_auths_clone);
async move {
proxy_auths_inner.lock().expect("mutex should lock").push(
request
.headers()
.get("proxy-authorization")
.and_then(|value| value.to_str().ok())
.map(str::to_string),
);
(
StatusCode::OK,
Body::from("fl=1234\nip=203.0.113.10\nwarp=off\n"),
)
}
}));
let (proxy_url, proxy_handle) = start_server(proxy).await;
let _probe_url_guard = crate::handlers::admin::override_proxy_connectivity_probe_url_for_tests(
"http://probe.example/cdn-cgi/trace",
);
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::default());
let gateway = build_router_with_state(
@@ -739,7 +798,6 @@ async fn gateway_creates_updates_and_tests_manual_proxy_nodes_locally() {
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let client = reqwest::Client::new();
let proxy_url = format!("http://127.0.0.1:{proxy_port}");
let create_response = client
.post(format!("{gateway_url}/api/admin/proxy-nodes/manual"))
@@ -749,7 +807,9 @@ async fn gateway_creates_updates_and_tests_manual_proxy_nodes_locally() {
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"name": "manual-node",
"proxy_url": proxy_url,
"proxy_url": proxy_url.clone(),
"username": "alice",
"password": "supersecret",
"region": "US-West"
}))
.send()
@@ -770,6 +830,8 @@ async fn gateway_creates_updates_and_tests_manual_proxy_nodes_locally() {
assert_eq!(create_payload["node"]["is_manual"], true);
assert_eq!(create_payload["node"]["status"], "online");
assert_eq!(create_payload["node"]["proxy_url"], proxy_url);
assert_eq!(create_payload["node"]["proxy_username"], "alice");
assert_eq!(create_payload["node"]["proxy_password"], "su****et");
let test_url_response = client
.post(format!("{gateway_url}/api/admin/proxy-nodes/test-url"))
@@ -778,7 +840,9 @@ async fn gateway_creates_updates_and_tests_manual_proxy_nodes_locally() {
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"proxy_url": proxy_url
"proxy_url": proxy_url.clone(),
"username": "alice",
"password": "supersecret"
}))
.send()
.await
@@ -790,6 +854,7 @@ async fn gateway_creates_updates_and_tests_manual_proxy_nodes_locally() {
.expect("json body should parse");
assert_eq!(test_url_payload["success"], true);
assert!(test_url_payload["latency_ms"].is_u64());
assert_eq!(test_url_payload["exit_ip"], "203.0.113.10");
let test_node_response = client
.post(format!(
@@ -808,6 +873,19 @@ async fn gateway_creates_updates_and_tests_manual_proxy_nodes_locally() {
.await
.expect("json body should parse");
assert_eq!(test_node_payload["success"], true);
assert_eq!(test_node_payload["exit_ip"], "203.0.113.10");
let expected_proxy_auth = format!(
"Basic {}",
base64::engine::general_purpose::STANDARD.encode("alice:supersecret")
);
assert_eq!(
proxy_auths.lock().expect("mutex should lock").as_slice(),
[
Some(expected_proxy_auth.clone()),
Some(expected_proxy_auth.clone()),
]
);
let update_response = client
.patch(format!("{gateway_url}/api/admin/proxy-nodes/{node_id}"))
@@ -832,7 +910,7 @@ async fn gateway_creates_updates_and_tests_manual_proxy_nodes_locally() {
assert_eq!(update_payload["node"]["proxy_url"], proxy_url);
gateway_handle.abort();
accept_handle.abort();
proxy_handle.abort();
}
#[tokio::test]
@@ -869,6 +947,135 @@ async fn gateway_tests_disconnected_tunnel_proxy_nodes_locally() {
gateway_handle.abort();
}
#[tokio::test]
async fn gateway_tests_connected_tunnel_proxy_nodes_with_active_probe() {
let _probe_url_guard = crate::handlers::admin::override_proxy_connectivity_probe_url_for_tests(
"https://probe.example/cdn-cgi/trace",
);
let mut node = sample_proxy_node("node-online");
node.status = "online".to_string();
node.tunnel_connected = true;
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![node]));
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(GatewayDataState::with_proxy_node_repository_for_tests(
proxy_node_repository,
));
let tunnel_state = state.tunnel.app_state();
let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8);
let (proxy_close_tx, _) = watch::channel(false);
tunnel_state
.hub
.register_proxy(Arc::new(TunnelProxyConn::new(
500,
"node-online".to_string(),
"Node Online".to_string(),
proxy_tx,
proxy_close_tx,
16,
)));
let gateway = build_router_with_state(state);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let request_task = tokio::spawn({
let gateway_url = gateway_url.clone();
async move {
reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/proxy-nodes/node-online/test"
))
.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")
.send()
.await
}
});
let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") {
Message::Binary(data) => data,
other => panic!("unexpected message: {other:?}"),
};
let request_header =
tunnel_protocol::FrameHeader::parse(&request_headers).expect("request header should parse");
assert_eq!(request_header.msg_type, tunnel_protocol::REQUEST_HEADERS);
let meta_payload = tunnel_protocol::decode_payload(&request_headers, &request_header)
.expect("request header payload should decode");
let meta: tunnel_protocol::RequestMeta =
serde_json::from_slice(&meta_payload).expect("request meta should parse");
assert_eq!(meta.method, "GET");
assert_eq!(meta.url, "https://probe.example/cdn-cgi/trace");
assert_eq!(meta.follow_redirects, Some(false));
let request_body = match proxy_rx.recv().await.expect("body frame should arrive") {
Message::Binary(data) => data,
other => panic!("unexpected message: {other:?}"),
};
let request_body_header =
tunnel_protocol::FrameHeader::parse(&request_body).expect("request body should parse");
assert_eq!(request_body_header.msg_type, tunnel_protocol::REQUEST_BODY);
assert_ne!(
request_body_header.flags & tunnel_protocol::FLAG_END_STREAM,
0,
"probe body frame should close the stream"
);
let response_meta = tunnel_protocol::ResponseMeta {
status: 200,
headers: vec![("content-type".to_string(), "text/plain".to_string())],
};
let response_meta_bytes =
serde_json::to_vec(&response_meta).expect("response meta should serialize");
let mut response_headers_frame = tunnel_protocol::encode_frame(
request_header.stream_id,
tunnel_protocol::RESPONSE_HEADERS,
0,
&response_meta_bytes,
);
tunnel_state
.hub
.handle_proxy_frame(500, &mut response_headers_frame)
.await;
let mut response_body_frame = tunnel_protocol::encode_frame(
request_header.stream_id,
tunnel_protocol::RESPONSE_BODY,
0,
b"fl=1234\nip=203.0.113.10\nwarp=off\n",
);
tunnel_state
.hub
.handle_proxy_frame(500, &mut response_body_frame)
.await;
let mut response_end_frame = tunnel_protocol::encode_frame(
request_header.stream_id,
tunnel_protocol::STREAM_END,
0,
&[],
);
tunnel_state
.hub
.handle_proxy_frame(500, &mut response_end_frame)
.await;
let response = request_task
.await
.expect("request task should complete")
.expect("test-node request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["success"], true);
assert!(payload["latency_ms"].is_u64());
assert_eq!(payload["exit_ip"], "203.0.113.10");
gateway_handle.abort();
}
#[tokio::test]
async fn gateway_deletes_proxy_nodes_and_clears_proxy_refs_locally() {
let mut manual_node = sample_proxy_node("manual-node-1");
@@ -1164,6 +1371,7 @@ async fn gateway_updates_proxy_node_config_and_batches_upgrade_locally() {
.post(format!("{gateway_url}/api/internal/tunnel/heartbeat"))
.json(&json!({
"node_id": "node-online",
"heartbeat_id": 77,
"heartbeat_interval": 45,
"active_connections": 3,
"total_requests": 5,
@@ -1179,6 +1387,7 @@ async fn gateway_updates_proxy_node_config_and_batches_upgrade_locally() {
.json()
.await
.expect("json body should parse");
assert_eq!(heartbeat_payload["heartbeat_id"], 77);
assert_eq!(heartbeat_payload["config_version"], 10);
assert!(heartbeat_payload.get("upgrade_to").is_none());
assert_eq!(heartbeat_payload["remote_config"]["allowed_ports"][1], 8443);

View File

@@ -493,7 +493,7 @@ async fn gateway_returns_503_for_admin_system_config_import_when_local_data_is_u
}
#[tokio::test]
async fn gateway_accepts_legacy_admin_system_config_import_versions() {
async fn gateway_rejects_legacy_admin_system_config_import_versions() {
let gateway = build_router_with_state(
AppState::new()
.expect("gateway should build")
@@ -519,10 +519,13 @@ async fn gateway_accepts_legacy_admin_system_config_import_versions() {
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let payload: Value = response.json().await.expect("json body should parse");
assert_eq!(payload["message"], "配置导入成功");
assert_eq!(payload["stats"]["errors"], json!([]));
let detail = payload["detail"]
.as_str()
.expect("detail should be a string");
assert!(detail.contains(&format!("不支持的配置版本: {version}")));
assert!(detail.contains("支持的版本: 2.2"));
}
gateway_handle.abort();
@@ -562,7 +565,7 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data() {
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"version": "1.3",
"version": "2.2",
"merge_mode": "overwrite",
"users": [{
"email": "alice@example.com",
@@ -761,8 +764,40 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data() {
}
#[tokio::test]
async fn gateway_imports_admin_system_config_fixtures_from_legacy_exports() {
for fixture in ["v20", "v21", "v22"] {
async fn gateway_imports_admin_system_config_fixture_v22() {
let gateway = build_router_with_state(
AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(build_empty_admin_system_data_state()),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.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("v22"))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let payload: Value = response.json().await.expect("json body should parse");
assert_eq!(payload["message"], "配置导入成功");
gateway_handle.abort();
}
#[tokio::test]
async fn gateway_rejects_admin_system_config_fixtures_from_removed_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")
@@ -783,16 +818,59 @@ async fn gateway_imports_admin_system_config_fixtures_from_legacy_exports() {
assert_eq!(
response.status(),
StatusCode::OK,
"fixture {fixture} should import"
StatusCode::BAD_REQUEST,
"fixture {fixture} should be rejected"
);
let payload: Value = response.json().await.expect("json body should parse");
assert_eq!(payload["message"], "配置导入成功");
let detail = payload["detail"]
.as_str()
.expect("detail should be a string");
assert!(detail.contains(&format!("不支持的配置版本: {version}")));
assert!(detail.contains("支持的版本: 2.2"));
gateway_handle.abort();
}
}
#[tokio::test]
async fn gateway_rejects_legacy_user_import_string_bool_field() {
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default());
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(
GatewayDataState::with_auth_api_key_repository_for_tests(Arc::clone(&auth_repository))
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
.with_auth_users_for_tests([sample_import_admin_user("admin-user-123")])
.with_auth_wallets_for_tests(Vec::<StoredWalletSnapshot>::new());
let gateway = build_router_with_state(state);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!("{gateway_url}/api/admin/system/users/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(&json!({
"version": "2.2",
"merge_mode": "overwrite",
"users": [{
"email": "legacy@example.com",
"email_verified": "true"
}]
}))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let payload: Value = response.json().await.expect("json body should parse");
assert_eq!(payload["detail"], "字段必须是布尔值");
gateway_handle.abort();
}
#[tokio::test]
async fn gateway_reports_field_path_for_invalid_admin_system_config_import_shape() {
let gateway = build_router_with_state(
@@ -835,7 +913,7 @@ async fn gateway_reports_field_path_for_invalid_admin_system_config_import_shape
}
#[tokio::test]
async fn gateway_imports_admin_system_config_with_numeric_string_prices() {
async fn gateway_rejects_admin_system_config_with_numeric_string_prices() {
let gateway = build_router_with_state(
AppState::new()
.expect("gateway should build")
@@ -860,12 +938,11 @@ async fn gateway_imports_admin_system_config_with_numeric_string_prices() {
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body: Value = response.json().await.expect("json body should parse");
assert_eq!(body["message"], "配置导入成功");
assert_eq!(body["stats"]["global_models"]["created"], json!(1));
assert_eq!(body["stats"]["providers"]["created"], json!(1));
assert_eq!(body["stats"]["models"]["created"], json!(1));
let detail = body["detail"].as_str().expect("detail should be a string");
assert!(detail.contains("配置文件格式无效"));
assert!(detail.contains("default_price_per_request"));
gateway_handle.abort();
}

View File

@@ -76,6 +76,37 @@ async fn gateway_handles_internal_tunnel_heartbeat_locally_with_loopback() {
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_rejects_internal_tunnel_heartbeat_without_heartbeat_id() {
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node(
"node-123",
)]));
let gateway = build_router_with_state(
AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(GatewayDataState::with_proxy_node_repository_for_tests(
Arc::clone(&repository),
)),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!("{gateway_url}/api/internal/tunnel/heartbeat"))
.json(&json!({
"node_id": "node-123",
"heartbeat_interval": 45,
"active_connections": 5
}))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
gateway_handle.abort();
}
#[tokio::test]
async fn gateway_handles_internal_tunnel_node_status_locally_with_loopback() {
let upstream_hits = Arc::new(Mutex::new(0usize));

View File

@@ -85,6 +85,7 @@ pub struct ProxyConn {
next_stream_id: AtomicU32,
pub stream_count: AtomicUsize,
pub max_streams: usize,
draining: AtomicBool,
}
impl ProxyConn {
@@ -104,6 +105,7 @@ impl ProxyConn {
next_stream_id: AtomicU32::new(2),
stream_count: AtomicUsize::new(0),
max_streams,
draining: AtomicBool::new(false),
}
}
@@ -159,13 +161,21 @@ impl ProxyConn {
}
pub fn is_available(&self) -> bool {
!self.outbound.is_closing()
!self.outbound.is_closing() && !self.is_draining()
}
pub fn request_close(&self) {
self.outbound.mark_closing();
}
pub fn mark_draining(&self) -> bool {
!self.draining.swap(true, Ordering::AcqRel)
}
pub fn is_draining(&self) -> bool {
self.draining.load(Ordering::Acquire)
}
pub fn send(&self, msg: Message) -> SendStatus {
let was_closing = self.outbound.is_closing();
let status = self.outbound.send(msg);
@@ -659,10 +669,16 @@ impl HubRouter {
}
protocol::PONG => {}
protocol::GOAWAY => {
warn!(
proxy_conn_id = proxy_conn_id,
"received GOAWAY from proxy connection"
);
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let first = pc.mark_draining();
if first {
warn!(
proxy_conn_id = proxy_conn_id,
node_id = %pc.node_id,
"received GOAWAY from proxy connection; marking connection draining"
);
}
}
}
_ => {
debug!(
@@ -975,6 +991,63 @@ mod tests {
assert_ne!(second_header.flags & protocol::FLAG_END_STREAM, 0);
}
#[tokio::test]
async fn goaway_marks_connection_draining_and_reroutes_new_streams() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_one_tx, mut proxy_one_rx) = bounded_queue(8);
let (proxy_one_close_tx, _) = watch::channel(false);
let proxy_one = Arc::new(ProxyConn::new(
201,
"node-drain".to_string(),
"Node Drain".to_string(),
proxy_one_tx,
proxy_one_close_tx,
16,
));
hub.register_proxy(Arc::clone(&proxy_one));
let (proxy_two_tx, mut proxy_two_rx) = bounded_queue(8);
let (proxy_two_close_tx, _) = watch::channel(false);
let proxy_two = Arc::new(ProxyConn::new(
202,
"node-drain".to_string(),
"Node Drain".to_string(),
proxy_two_tx,
proxy_two_close_tx,
16,
));
hub.register_proxy(Arc::clone(&proxy_two));
let mut goaway = protocol::encode_goaway();
hub.handle_proxy_frame(201, &mut goaway).await;
assert!(
proxy_one.is_draining(),
"first connection should be draining"
);
assert!(
!proxy_two.is_draining(),
"second connection should remain schedulable"
);
let _stream = hub
.open_local_stream("node-drain", &build_meta())
.expect("open local stream");
assert!(
proxy_one_rx.try_recv().is_err(),
"draining connection should not receive new streams"
);
let routed = proxy_two_rx
.try_recv()
.expect("headers should route to second connection");
let routed_data = match routed {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let header = protocol::FrameHeader::parse(&routed_data).expect("frame header");
assert_eq!(header.msg_type, protocol::REQUEST_HEADERS);
}
#[tokio::test]
async fn heartbeat_callback_failure_does_not_send_fake_ack() {
let hub = HubRouter::new(ControlPlaneClient::local(

View File

@@ -22,7 +22,7 @@ use tracing::warn;
use crate::data::GatewayDataState;
pub use control_plane::ControlPlaneClient;
pub use hub::{ConnConfig, HubRouter, ProxyConn};
pub use hub::{ConnConfig, HubRouter, LocalBodyEvent, ProxyConn};
pub use local_relay::relay_request;
#[derive(Clone)]

View File

@@ -46,12 +46,13 @@ pub(crate) const TUNNEL_NODE_STATUS_PATH: &str = "/api/internal/tunnel/node-stat
pub(crate) const TUNNEL_RELAY_PATH_PATTERN: &str = "/api/internal/tunnel/relay/{node_id}";
pub(crate) const TUNNEL_ROUTE_FAMILY: &str = "tunnel_manage";
const DEFAULT_PROXY_IDLE_TIMEOUT_SECS: u64 = 0;
const DEFAULT_PING_INTERVAL_SECS: u64 = 15;
const DEFAULT_PROXY_IDLE_TIMEOUT_MS: u64 = 900;
const DEFAULT_PING_INTERVAL_MS: u64 = 250;
const DEFAULT_MAX_STREAMS: usize = 2048;
const DEFAULT_OUTBOUND_QUEUE_CAPACITY: usize = 128;
const DEFAULT_ATTACHMENT_TTL_SECS: u64 = 90;
const DEFAULT_OWNER_RELAY_BODY_LIMIT_BYTES: usize = 5_242_880;
const DEFAULT_TUNNEL_PROBE_BODY_LIMIT_BYTES: usize = 64 * 1024;
const TUNNEL_ATTACHMENT_KEY_PREFIX: &str = "tunnel.attachments.";
const TUNNEL_ATTACHMENT_REDIS_KEY_PREFIX: &str = "tunnel:attachments:";
const TUNNEL_INSTANCE_ID_ENV: &str = "AETHER_GATEWAY_INSTANCE_ID";
@@ -61,8 +62,7 @@ const TUNNEL_ATTACHMENT_TTL_ENV: &str = "AETHER_TUNNEL_ATTACHMENT_TTL_SECS";
#[derive(Debug, Deserialize)]
struct InternalTunnelHeartbeatRequest {
node_id: String,
#[serde(default)]
heartbeat_id: Option<u64>,
heartbeat_id: u64,
#[serde(default)]
heartbeat_interval: Option<i32>,
#[serde(default)]
@@ -380,6 +380,12 @@ pub(crate) struct TunnelStatsSnapshot {
pub(crate) active_streams: usize,
}
#[derive(Debug, Clone)]
pub(crate) struct TunnelProbeResponse {
pub(crate) status: u16,
pub(crate) body: String,
}
impl EmbeddedTunnelState {
pub(crate) fn new() -> Self {
Self::with_data(Arc::new(GatewayDataState::disabled()))
@@ -409,8 +415,8 @@ impl EmbeddedTunnelState {
inner: TunnelAppState::new(
build_embedded_control_plane(Arc::clone(&data), attachment_directory.clone()),
ConnConfig {
ping_interval: Duration::from_secs(DEFAULT_PING_INTERVAL_SECS),
idle_timeout: Duration::from_secs(DEFAULT_PROXY_IDLE_TIMEOUT_SECS),
ping_interval: Duration::from_millis(DEFAULT_PING_INTERVAL_MS),
idle_timeout: Duration::from_millis(DEFAULT_PROXY_IDLE_TIMEOUT_MS),
outbound_queue_capacity: DEFAULT_OUTBOUND_QUEUE_CAPACITY,
},
DEFAULT_MAX_STREAMS,
@@ -451,6 +457,18 @@ impl EmbeddedTunnelState {
url: &str,
timeout_secs: u64,
) -> Result<u16, String> {
Ok(self
.probe_node_url_with_response(node_id, url, timeout_secs)
.await?
.status)
}
pub(crate) async fn probe_node_url_with_response(
&self,
node_id: &str,
url: &str,
timeout_secs: u64,
) -> Result<TunnelProbeResponse, String> {
let timeout_secs = timeout_secs.clamp(5, 60);
let meta = tunnel_protocol::RequestMeta {
method: "GET".to_string(),
@@ -469,7 +487,35 @@ impl EmbeddedTunnelState {
let response = stream
.wait_headers(Duration::from_secs(timeout_secs))
.await?;
Ok(response.status)
let Some(mut body_rx) = stream.take_body_receiver() else {
return Err("missing tunnel probe response body receiver".to_string());
};
let body = tokio::time::timeout(Duration::from_secs(timeout_secs), async {
let mut body_bytes = Vec::new();
while let Some(event) = body_rx.recv().await {
match event {
embedded::LocalBodyEvent::Chunk(chunk) => {
let next_len = body_bytes.len().saturating_add(chunk.len());
if next_len > DEFAULT_TUNNEL_PROBE_BODY_LIMIT_BYTES {
return Err(format!(
"tunnel probe body exceeds {} bytes",
DEFAULT_TUNNEL_PROBE_BODY_LIMIT_BYTES
));
}
body_bytes.extend_from_slice(&chunk);
}
embedded::LocalBodyEvent::End => break,
embedded::LocalBodyEvent::Error(error) => return Err(error),
}
}
Ok::<String, String>(String::from_utf8_lossy(&body_bytes).to_string())
})
.await
.map_err(|_| "timed out waiting for tunnel probe response body".to_string())??;
Ok(TunnelProbeResponse {
status: response.status,
body,
})
}
.await;
self.inner
@@ -510,8 +556,8 @@ impl Default for EmbeddedTunnelState {
impl fmt::Debug for EmbeddedTunnelState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("EmbeddedTunnelState")
.field("proxy_idle_timeout_secs", &DEFAULT_PROXY_IDLE_TIMEOUT_SECS)
.field("ping_interval_secs", &DEFAULT_PING_INTERVAL_SECS)
.field("proxy_idle_timeout_ms", &DEFAULT_PROXY_IDLE_TIMEOUT_MS)
.field("ping_interval_ms", &DEFAULT_PING_INTERVAL_MS)
.field("max_streams", &DEFAULT_MAX_STREAMS)
.field("outbound_queue_capacity", &DEFAULT_OUTBOUND_QUEUE_CAPACITY)
.field(
@@ -879,14 +925,9 @@ async fn apply_embedded_tunnel_node_status(
.map_err(|err| format!("node status sync failed: {err}"))
}
fn build_embedded_tunnel_heartbeat_ack(
node: &StoredProxyNode,
heartbeat_id: Option<u64>,
) -> Vec<u8> {
fn build_embedded_tunnel_heartbeat_ack(node: &StoredProxyNode, heartbeat_id: u64) -> Vec<u8> {
let mut payload = serde_json::Map::new();
if let Some(heartbeat_id) = heartbeat_id {
payload.insert("heartbeat_id".to_string(), json!(heartbeat_id));
}
payload.insert("heartbeat_id".to_string(), json!(heartbeat_id));
if let Some(remote_config) = node.remote_config.as_ref() {
payload.insert("remote_config".to_string(), remote_config.clone());
payload.insert("config_version".to_string(), json!(node.config_version));
@@ -911,7 +952,7 @@ fn parse_embedded_tunnel_heartbeat_request(
.map_err(|_| "invalid heartbeat payload".to_string())?;
let node_id = payload.node_id.trim();
if node_id.is_empty() || node_id.len() > 36 {
if node_id.is_empty() || node_id.len() > 36 || payload.heartbeat_id == 0 {
return Err("invalid heartbeat payload".to_string());
}
if payload
@@ -1042,6 +1083,27 @@ mod tests {
);
}
#[tokio::test]
async fn embedded_tunnel_heartbeat_rejects_missing_heartbeat_id() {
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node(
"node-123",
)]));
let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository));
let error = apply_embedded_tunnel_heartbeat(
&data,
br#"{
"node_id": "node-123",
"heartbeat_interval": 45,
"active_connections": 5
}"#,
)
.await
.expect_err("heartbeat without heartbeat_id should fail");
assert_eq!(error, "invalid heartbeat payload");
}
#[tokio::test]
async fn embedded_tunnel_node_status_updates_proxy_node_repository() {
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![sample_proxy_node(