mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat(proxy): 实现代理节点批量升级回滚、隧道重定向跟随及远程配置管理
核心功能: - 新增代理节点批量升级回滚工作流,支持分批升级、健康探针、跳过/重试/取消等操作 - proxy 隧道流处理器支持 HTTP 重定向跟随(最多 10 跳),区分 307/308 可重播与不可重播请求体 - proxy 协议新增 follow_redirects / http1_only 字段,网关侧同步支持 - 新增代理节点远端配置变更接口(名称、允许端口、调度状态、升级目标等) - 新增代理节点注册/反注册/心跳的 Admin API,及节点过期清理维护任务 - gateway 隧道 owner-relay 支持流式代理大请求体,新增 5 MiB 默认限制 - 新增 ProxyNodeRegistrationMutation / ProxyNodeRemoteConfigMutation 数据类型 - proxy 配置新增重定向重播预算、心跳间隔等参数,TUI 安装向导同步更新 - 前端 ProxyNodes 页面新增批量升级操作面板及滚动进度展示
This commit is contained in:
@@ -178,6 +178,69 @@ pub(super) fn classify_admin_operations_family_route(
|
||||
"admin:proxy_nodes",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& matches!(
|
||||
normalized_path,
|
||||
"/api/admin/proxy-nodes/upgrade/cancel" | "/api/admin/proxy-nodes/upgrade/cancel/"
|
||||
)
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"proxy_nodes_manage",
|
||||
"cancel_upgrade_rollout",
|
||||
"admin:proxy_nodes",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& matches!(
|
||||
normalized_path,
|
||||
"/api/admin/proxy-nodes/upgrade/clear-conflicts"
|
||||
| "/api/admin/proxy-nodes/upgrade/clear-conflicts/"
|
||||
)
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"proxy_nodes_manage",
|
||||
"clear_upgrade_rollout_conflicts",
|
||||
"admin:proxy_nodes",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& matches!(
|
||||
normalized_path,
|
||||
"/api/admin/proxy-nodes/upgrade/restore-skipped"
|
||||
| "/api/admin/proxy-nodes/upgrade/restore-skipped/"
|
||||
)
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"proxy_nodes_manage",
|
||||
"restore_skipped_upgrade_rollout_nodes",
|
||||
"admin:proxy_nodes",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& normalized_path.starts_with("/api/admin/proxy-nodes/")
|
||||
&& normalized_path.ends_with("/upgrade/skip")
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"proxy_nodes_manage",
|
||||
"skip_upgrade_rollout_node",
|
||||
"admin:proxy_nodes",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& normalized_path.starts_with("/api/admin/proxy-nodes/")
|
||||
&& normalized_path.ends_with("/upgrade/retry")
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"proxy_nodes_manage",
|
||||
"retry_upgrade_rollout_node",
|
||||
"admin:proxy_nodes",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& matches!(
|
||||
normalized_path,
|
||||
|
||||
@@ -58,3 +58,113 @@ fn classifies_admin_proxy_nodes_events_as_admin_proxy_route() {
|
||||
);
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_admin_proxy_nodes_upgrade_cancel_as_admin_proxy_route() {
|
||||
let headers = headers(&[]);
|
||||
let uri: Uri = "/api/admin/proxy-nodes/upgrade/cancel"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
|
||||
assert_eq!(decision.route_family.as_deref(), Some("proxy_nodes_manage"));
|
||||
assert_eq!(
|
||||
decision.route_kind.as_deref(),
|
||||
Some("cancel_upgrade_rollout")
|
||||
);
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("admin:proxy_nodes")
|
||||
);
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_admin_proxy_nodes_upgrade_clear_conflicts_as_admin_proxy_route() {
|
||||
let headers = headers(&[]);
|
||||
let uri: Uri = "/api/admin/proxy-nodes/upgrade/clear-conflicts"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
|
||||
assert_eq!(decision.route_family.as_deref(), Some("proxy_nodes_manage"));
|
||||
assert_eq!(
|
||||
decision.route_kind.as_deref(),
|
||||
Some("clear_upgrade_rollout_conflicts")
|
||||
);
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("admin:proxy_nodes")
|
||||
);
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_admin_proxy_nodes_upgrade_restore_skipped_as_admin_proxy_route() {
|
||||
let headers = headers(&[]);
|
||||
let uri: Uri = "/api/admin/proxy-nodes/upgrade/restore-skipped"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
|
||||
assert_eq!(decision.route_family.as_deref(), Some("proxy_nodes_manage"));
|
||||
assert_eq!(
|
||||
decision.route_kind.as_deref(),
|
||||
Some("restore_skipped_upgrade_rollout_nodes")
|
||||
);
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("admin:proxy_nodes")
|
||||
);
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_admin_proxy_nodes_upgrade_skip_node_as_admin_proxy_route() {
|
||||
let headers = headers(&[]);
|
||||
let uri: Uri = "/api/admin/proxy-nodes/node-1/upgrade/skip"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
|
||||
assert_eq!(decision.route_family.as_deref(), Some("proxy_nodes_manage"));
|
||||
assert_eq!(
|
||||
decision.route_kind.as_deref(),
|
||||
Some("skip_upgrade_rollout_node")
|
||||
);
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("admin:proxy_nodes")
|
||||
);
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_admin_proxy_nodes_upgrade_retry_node_as_admin_proxy_route() {
|
||||
let headers = headers(&[]);
|
||||
let uri: Uri = "/api/admin/proxy-nodes/node-1/upgrade/retry"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
|
||||
assert_eq!(decision.route_family.as_deref(), Some("proxy_nodes_manage"));
|
||||
assert_eq!(
|
||||
decision.route_kind.as_deref(),
|
||||
Some("retry_upgrade_rollout_node")
|
||||
);
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("admin:proxy_nodes")
|
||||
);
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
@@ -1,12 +1,13 @@
|
||||
use super::{
|
||||
AuthApiKeyLookupKey, CreateManagementTokenRecord, DataLayerError, GatewayAuthApiKeySnapshot,
|
||||
GatewayDataState, ManagementTokenListQuery, ProxyNodeHeartbeatMutation,
|
||||
ProxyNodeTunnelStatusMutation, RegenerateManagementTokenSecret, StoredAuthApiKeyExportRecord,
|
||||
StoredAuthApiKeySnapshot, StoredLdapModuleConfig, StoredManagementToken,
|
||||
StoredManagementTokenListPage, StoredManagementTokenWithUser, StoredOAuthProviderConfig,
|
||||
StoredOAuthProviderModuleConfig, StoredProxyNode, StoredProxyNodeEvent, StoredUserAuthRecord,
|
||||
StoredUserPreferenceRecord, StoredUserSessionRecord, StoredWalletSnapshot,
|
||||
UpdateManagementTokenRecord, UpsertOAuthProviderConfigRecord,
|
||||
ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTunnelStatusMutation,
|
||||
RegenerateManagementTokenSecret, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
|
||||
StoredLdapModuleConfig, StoredManagementToken, StoredManagementTokenListPage,
|
||||
StoredManagementTokenWithUser, StoredOAuthProviderConfig, StoredOAuthProviderModuleConfig,
|
||||
StoredProxyNode, StoredProxyNodeEvent, StoredUserAuthRecord, StoredUserPreferenceRecord,
|
||||
StoredUserSessionRecord, StoredWalletSnapshot, UpdateManagementTokenRecord,
|
||||
UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
use crate::LocalMutationOutcome;
|
||||
use aether_data::repository::auth::{
|
||||
@@ -1858,6 +1859,20 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn get_management_token_with_user_by_hash(
|
||||
&self,
|
||||
token_hash: &str,
|
||||
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
|
||||
match &self.management_token_reader {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.get_management_token_with_user_by_hash(token_hash)
|
||||
.await
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn create_management_token(
|
||||
&self,
|
||||
record: &CreateManagementTokenRecord,
|
||||
@@ -1901,6 +1916,21 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn record_management_token_usage(
|
||||
&self,
|
||||
token_id: &str,
|
||||
last_used_ip: Option<&str>,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
match &self.management_token_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.record_management_token_usage(token_id, last_used_ip)
|
||||
.await
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn find_proxy_node(
|
||||
&self,
|
||||
node_id: &str,
|
||||
@@ -1929,6 +1959,25 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn register_proxy_node(
|
||||
&self,
|
||||
mutation: &ProxyNodeRegistrationMutation,
|
||||
) -> Result<Option<StoredProxyNode>, DataLayerError> {
|
||||
match &self.proxy_node_writer {
|
||||
Some(repository) => repository.register_node(mutation).await.map(Some),
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn reset_stale_proxy_node_tunnel_statuses(
|
||||
&self,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
match &self.proxy_node_writer {
|
||||
Some(repository) => repository.reset_stale_tunnel_statuses().await,
|
||||
None => Ok(0),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn apply_proxy_node_heartbeat(
|
||||
&self,
|
||||
mutation: &ProxyNodeHeartbeatMutation,
|
||||
@@ -1949,6 +1998,26 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn unregister_proxy_node(
|
||||
&self,
|
||||
node_id: &str,
|
||||
) -> Result<Option<StoredProxyNode>, DataLayerError> {
|
||||
match &self.proxy_node_writer {
|
||||
Some(repository) => repository.unregister_node(node_id).await,
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn update_proxy_node_remote_config(
|
||||
&self,
|
||||
mutation: &ProxyNodeRemoteConfigMutation,
|
||||
) -> Result<Option<StoredProxyNode>, DataLayerError> {
|
||||
match &self.proxy_node_writer {
|
||||
Some(repository) => repository.update_remote_config(mutation).await,
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn set_management_token_active(
|
||||
&self,
|
||||
token_id: &str,
|
||||
|
||||
@@ -224,6 +224,15 @@ impl GatewayDataState {
|
||||
self.proxy_node_writer.is_some()
|
||||
}
|
||||
|
||||
pub(crate) fn has_system_config_store(&self) -> bool {
|
||||
self.system_config_values.is_some()
|
||||
|| self
|
||||
.backends
|
||||
.as_ref()
|
||||
.and_then(|backends| backends.postgres())
|
||||
.is_some()
|
||||
}
|
||||
|
||||
pub(crate) fn oauth_refresh_lock_runner(&self) -> Option<RedisLockRunner> {
|
||||
self.backends
|
||||
.as_ref()
|
||||
|
||||
@@ -40,8 +40,9 @@ use aether_data::repository::oauth_providers::{
|
||||
UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
use aether_data::repository::proxy_nodes::{
|
||||
ProxyNodeHeartbeatMutation, ProxyNodeReadRepository, ProxyNodeTunnelStatusMutation,
|
||||
ProxyNodeWriteRepository, StoredProxyNode, StoredProxyNodeEvent,
|
||||
ProxyNodeHeartbeatMutation, ProxyNodeReadRepository, ProxyNodeRegistrationMutation,
|
||||
ProxyNodeRemoteConfigMutation, ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository,
|
||||
StoredProxyNode, StoredProxyNodeEvent,
|
||||
};
|
||||
use aether_data::repository::shadow_results::{
|
||||
merge_shadow_result_sample, RecordShadowResultSample, ShadowResultLookupKey,
|
||||
|
||||
@@ -1,16 +1,20 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::error::Error as _;
|
||||
use std::io::Read;
|
||||
use std::io::Write;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use aether_contracts::{
|
||||
ExecutionPlan, ExecutionResult, ExecutionTelemetry, ProxySnapshot, ResponseBody,
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
|
||||
};
|
||||
use aether_http::{apply_http_client_config, HttpClientConfig};
|
||||
use base64::Engine as _;
|
||||
use flate2::read::{DeflateDecoder, GzDecoder};
|
||||
use flate2::write::GzEncoder;
|
||||
use flate2::Compression;
|
||||
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
|
||||
use reqwest::redirect::Policy;
|
||||
use reqwest::tls::Version;
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
@@ -116,11 +120,21 @@ struct RelayRequestMeta {
|
||||
url: String,
|
||||
headers: BTreeMap<String, String>,
|
||||
timeout: u64,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
follow_redirects: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "is_false")]
|
||||
http1_only: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub(crate) struct DirectSyncExecutionRuntime;
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
struct ExecutionTransportControls {
|
||||
follow_redirects: Option<bool>,
|
||||
http1_only: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(crate) struct DirectUpstreamStreamExecution {
|
||||
pub(crate) request_id: String,
|
||||
@@ -150,6 +164,8 @@ impl DirectSyncExecutionRuntime {
|
||||
let body_bytes = response.bytes().await.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
|
||||
})?;
|
||||
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)
|
||||
.unwrap_or_else(|| body_bytes.to_vec());
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
let upstream_bytes = body_bytes.len() as u64;
|
||||
|
||||
@@ -160,8 +176,8 @@ impl DirectSyncExecutionRuntime {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(&body_bytes)),
|
||||
})
|
||||
} else if response_body_is_json(&headers, &body_bytes) {
|
||||
let body_json: Value = serde_json::from_slice(&body_bytes)
|
||||
} else if response_body_is_json(&headers, &decoded_body_bytes) {
|
||||
let body_json: Value = serde_json::from_slice(&decoded_body_bytes)
|
||||
.map_err(ExecutionRuntimeTransportError::InvalidJson)?;
|
||||
Some(ResponseBody {
|
||||
json_body: Some(body_json),
|
||||
@@ -253,6 +269,7 @@ async fn send_request(
|
||||
}
|
||||
|
||||
let method = plan.method.parse::<reqwest::Method>()?;
|
||||
let transport_controls = resolve_execution_transport_controls(&plan.headers);
|
||||
let headers = build_request_headers(
|
||||
&plan.headers,
|
||||
plan.content_encoding.as_deref(),
|
||||
@@ -266,14 +283,23 @@ async fn send_request(
|
||||
.map(Duration::from_millis);
|
||||
|
||||
if let Some(node_id) = resolve_tunnel_node_id(plan.proxy.as_ref()) {
|
||||
return send_via_tunnel_relay(plan, method, headers, body_bytes, &node_id, total_timeout)
|
||||
.await;
|
||||
return send_via_tunnel_relay(
|
||||
plan,
|
||||
method,
|
||||
headers,
|
||||
body_bytes,
|
||||
&node_id,
|
||||
total_timeout,
|
||||
transport_controls,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
let client = build_client(
|
||||
plan.timeouts.as_ref(),
|
||||
plan.proxy.as_ref(),
|
||||
plan.tls_profile.as_deref(),
|
||||
transport_controls,
|
||||
)?;
|
||||
let mut request = client.request(method, &plan.url);
|
||||
request = request.headers(headers).body(body_bytes);
|
||||
@@ -320,6 +346,7 @@ async fn send_via_tunnel_relay(
|
||||
body_bytes: Vec<u8>,
|
||||
node_id: &str,
|
||||
total_timeout: Option<Duration>,
|
||||
transport_controls: ExecutionTransportControls,
|
||||
) -> Result<reqwest::Response, ExecutionRuntimeTransportError> {
|
||||
let client = build_relay_client(plan.timeouts.as_ref())?;
|
||||
let relay_url = build_relay_url(plan.proxy.as_ref(), node_id);
|
||||
@@ -329,6 +356,8 @@ async fn send_via_tunnel_relay(
|
||||
url: plan.url.clone(),
|
||||
headers: header_map_to_string_map(&headers),
|
||||
timeout: resolve_relay_timeout_seconds(plan),
|
||||
follow_redirects: transport_controls.follow_redirects,
|
||||
http1_only: transport_controls.http1_only,
|
||||
},
|
||||
&body_bytes,
|
||||
)?;
|
||||
@@ -501,9 +530,17 @@ fn build_client(
|
||||
timeouts: Option<&aether_contracts::ExecutionTimeouts>,
|
||||
proxy: Option<&ProxySnapshot>,
|
||||
tls_profile: Option<&str>,
|
||||
transport_controls: ExecutionTransportControls,
|
||||
) -> Result<reqwest::Client, ExecutionRuntimeTransportError> {
|
||||
let mut builder = reqwest::Client::builder();
|
||||
if transport_controls.follow_redirects != Some(true) {
|
||||
builder = builder.redirect(Policy::none());
|
||||
}
|
||||
if transport_controls.http1_only {
|
||||
builder = builder.http1_only();
|
||||
}
|
||||
let mut builder = apply_http_client_config(
|
||||
reqwest::Client::builder(),
|
||||
builder,
|
||||
&HttpClientConfig {
|
||||
connect_timeout_ms: timeouts.and_then(|timeouts| timeouts.connect_ms),
|
||||
..HttpClientConfig::default()
|
||||
@@ -608,7 +645,11 @@ fn build_request_headers(
|
||||
}
|
||||
for (key, value) in headers {
|
||||
let normalized_key = key.trim().to_ascii_lowercase();
|
||||
if is_hop_by_hop_header(&normalized_key) || normalized_key == "content-encoding" {
|
||||
if is_hop_by_hop_header(&normalized_key)
|
||||
|| normalized_key == "content-encoding"
|
||||
|| normalized_key == EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER
|
||||
|| normalized_key == EXECUTION_REQUEST_HTTP1_ONLY_HEADER
|
||||
{
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -629,6 +670,43 @@ fn build_request_headers(
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
fn resolve_execution_transport_controls(
|
||||
headers: &BTreeMap<String, String>,
|
||||
) -> ExecutionTransportControls {
|
||||
ExecutionTransportControls {
|
||||
follow_redirects: execution_transport_header_value(
|
||||
headers,
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
|
||||
)
|
||||
.and_then(|value| parse_execution_transport_bool(value)),
|
||||
http1_only: execution_transport_header_value(headers, EXECUTION_REQUEST_HTTP1_ONLY_HEADER)
|
||||
.and_then(|value| parse_execution_transport_bool(value))
|
||||
.unwrap_or(false),
|
||||
}
|
||||
}
|
||||
|
||||
fn execution_transport_header_value<'a>(
|
||||
headers: &'a BTreeMap<String, String>,
|
||||
target: &str,
|
||||
) -> Option<&'a str> {
|
||||
headers
|
||||
.iter()
|
||||
.find(|(name, _)| name.eq_ignore_ascii_case(target))
|
||||
.map(|(_, value)| value.as_str())
|
||||
}
|
||||
|
||||
fn parse_execution_transport_bool(value: &str) -> Option<bool> {
|
||||
match value.trim().to_ascii_lowercase().as_str() {
|
||||
"1" | "true" | "yes" | "on" => Some(true),
|
||||
"0" | "false" | "no" | "off" => Some(false),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn is_false(value: &bool) -> bool {
|
||||
!*value
|
||||
}
|
||||
|
||||
fn header_map_to_string_map(headers: &HeaderMap) -> BTreeMap<String, String> {
|
||||
headers
|
||||
.iter()
|
||||
@@ -661,6 +739,33 @@ fn collect_response_headers(headers: &HeaderMap) -> BTreeMap<String, String> {
|
||||
header_map_to_string_map(headers)
|
||||
}
|
||||
|
||||
fn decode_response_body_bytes(
|
||||
headers: &BTreeMap<String, String>,
|
||||
body_bytes: &[u8],
|
||||
) -> Option<Vec<u8>> {
|
||||
let encoding = headers
|
||||
.get("content-encoding")
|
||||
.map(String::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| value.to_ascii_lowercase());
|
||||
match encoding.as_deref() {
|
||||
Some("gzip") => {
|
||||
let mut decoder = GzDecoder::new(body_bytes);
|
||||
let mut out = Vec::new();
|
||||
decoder.read_to_end(&mut out).ok()?;
|
||||
Some(out)
|
||||
}
|
||||
Some("deflate") => {
|
||||
let mut decoder = DeflateDecoder::new(body_bytes);
|
||||
let mut out = Vec::new();
|
||||
decoder.read_to_end(&mut out).ok()?;
|
||||
Some(out)
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn response_body_is_json(headers: &BTreeMap<String, String>, body_bytes: &[u8]) -> bool {
|
||||
if headers
|
||||
.get("content-type")
|
||||
@@ -678,7 +783,10 @@ mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
use std::io::Read;
|
||||
|
||||
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
|
||||
use aether_contracts::{
|
||||
ExecutionPlan, ExecutionTimeouts, RequestBody, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
|
||||
EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
|
||||
};
|
||||
use axum::body::Bytes;
|
||||
use axum::extract::Path;
|
||||
use axum::routing::post;
|
||||
@@ -884,6 +992,229 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn direct_sync_execution_runtime_disables_redirects_by_default() {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("local addr should resolve");
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/redirect",
|
||||
post(|| async {
|
||||
(
|
||||
axum::http::StatusCode::TEMPORARY_REDIRECT,
|
||||
[(
|
||||
axum::http::header::LOCATION,
|
||||
axum::http::HeaderValue::from_static("/final"),
|
||||
)],
|
||||
)
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/final",
|
||||
post(|| async {
|
||||
(
|
||||
axum::http::StatusCode::OK,
|
||||
Json(json!({"redirected": true})),
|
||||
)
|
||||
}),
|
||||
);
|
||||
let server = tokio::spawn(async move {
|
||||
axum::serve(listener, app)
|
||||
.await
|
||||
.expect("test server should run");
|
||||
});
|
||||
|
||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||
let result = execution_runtime
|
||||
.execute_sync(ExecutionPlan {
|
||||
request_id: "req-redirect-1".into(),
|
||||
candidate_id: None,
|
||||
provider_name: Some("provider_ops".into()),
|
||||
provider_id: "prov-1".into(),
|
||||
endpoint_id: "ep-1".into(),
|
||||
key_id: "key-1".into(),
|
||||
method: "POST".into(),
|
||||
url: format!("http://{addr}/redirect"),
|
||||
headers: BTreeMap::from([("content-type".into(), "application/json".into())]),
|
||||
content_type: Some("application/json".into()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(json!({"model": "gpt-4.1"})),
|
||||
stream: false,
|
||||
client_api_format: "provider_ops:verify".into(),
|
||||
provider_api_format: "provider_ops:verify".into(),
|
||||
model_name: Some("verify-auth".into()),
|
||||
proxy: None,
|
||||
tls_profile: None,
|
||||
timeouts: Some(ExecutionTimeouts {
|
||||
connect_ms: Some(5_000),
|
||||
total_ms: Some(5_000),
|
||||
..ExecutionTimeouts::default()
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.expect("sync execution should succeed");
|
||||
|
||||
server.abort();
|
||||
|
||||
assert_eq!(result.status_code, 307);
|
||||
assert_eq!(
|
||||
result.headers.get("location").map(String::as_str),
|
||||
Some("/final")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn direct_sync_execution_runtime_follows_redirects_when_explicitly_enabled() {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("local addr should resolve");
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/redirect",
|
||||
post(|| async {
|
||||
(
|
||||
axum::http::StatusCode::TEMPORARY_REDIRECT,
|
||||
[(
|
||||
axum::http::header::LOCATION,
|
||||
axum::http::HeaderValue::from_static("/final"),
|
||||
)],
|
||||
)
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/final",
|
||||
post(|| async {
|
||||
(
|
||||
axum::http::StatusCode::OK,
|
||||
Json(json!({"redirected": true})),
|
||||
)
|
||||
}),
|
||||
);
|
||||
let server = tokio::spawn(async move {
|
||||
axum::serve(listener, app)
|
||||
.await
|
||||
.expect("test server should run");
|
||||
});
|
||||
|
||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||
let result = execution_runtime
|
||||
.execute_sync(ExecutionPlan {
|
||||
request_id: "req-redirect-2".into(),
|
||||
candidate_id: None,
|
||||
provider_name: Some("provider_oauth".into()),
|
||||
provider_id: "prov-1".into(),
|
||||
endpoint_id: "ep-1".into(),
|
||||
key_id: "key-1".into(),
|
||||
method: "POST".into(),
|
||||
url: format!("http://{addr}/redirect"),
|
||||
headers: BTreeMap::from([
|
||||
("content-type".into(), "application/json".into()),
|
||||
(
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.into(),
|
||||
"true".into(),
|
||||
),
|
||||
]),
|
||||
content_type: Some("application/json".into()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(json!({"model": "gpt-4.1"})),
|
||||
stream: false,
|
||||
client_api_format: "provider_oauth:exchange".into(),
|
||||
provider_api_format: "provider_oauth:exchange".into(),
|
||||
model_name: Some("oauth-exchange".into()),
|
||||
proxy: None,
|
||||
tls_profile: None,
|
||||
timeouts: Some(ExecutionTimeouts {
|
||||
connect_ms: Some(5_000),
|
||||
total_ms: Some(5_000),
|
||||
..ExecutionTimeouts::default()
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.expect("sync execution should succeed");
|
||||
|
||||
server.abort();
|
||||
|
||||
assert_eq!(result.status_code, 200);
|
||||
assert_eq!(
|
||||
result.body.and_then(|body| body.json_body),
|
||||
Some(json!({"redirected": true}))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn direct_sync_execution_runtime_forwards_http1_only_control_to_tunnel_relay() {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("local addr should resolve");
|
||||
let app = Router::new().route(
|
||||
"/api/internal/tunnel/relay/{node_id}",
|
||||
post(|Path(node_id): Path<String>, body: Bytes| async move {
|
||||
let (meta, request_body) = decode_relay_envelope(&body);
|
||||
assert_eq!(node_id, "node-1");
|
||||
assert_eq!(meta["http1_only"], true);
|
||||
assert_eq!(meta["follow_redirects"], json!(false));
|
||||
let request_json: serde_json::Value =
|
||||
serde_json::from_slice(&request_body).expect("request body should be json");
|
||||
assert_eq!(request_json["model"], "gpt-4.1");
|
||||
(axum::http::StatusCode::OK, Json(json!({"ok": true})))
|
||||
}),
|
||||
);
|
||||
let server = tokio::spawn(async move {
|
||||
axum::serve(listener, app)
|
||||
.await
|
||||
.expect("relay test server should run");
|
||||
});
|
||||
|
||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||
let result = execution_runtime
|
||||
.execute_sync(ExecutionPlan {
|
||||
request_id: "req-relay-http1-1".into(),
|
||||
candidate_id: None,
|
||||
provider_name: Some("provider_ops".into()),
|
||||
provider_id: "prov-1".into(),
|
||||
endpoint_id: "ep-1".into(),
|
||||
key_id: "key-1".into(),
|
||||
method: "POST".into(),
|
||||
url: "https://example.com/chat".into(),
|
||||
headers: BTreeMap::from([
|
||||
("content-type".into(), "application/json".into()),
|
||||
(EXECUTION_REQUEST_HTTP1_ONLY_HEADER.into(), "true".into()),
|
||||
(
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.into(),
|
||||
"false".into(),
|
||||
),
|
||||
]),
|
||||
content_type: Some("application/json".into()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(json!({"model": "gpt-4.1"})),
|
||||
stream: false,
|
||||
client_api_format: "provider_ops:verify".into(),
|
||||
provider_api_format: "provider_ops:verify".into(),
|
||||
model_name: Some("verify-auth".into()),
|
||||
proxy: Some(tunnel_proxy_snapshot(format!("http://{addr}"))),
|
||||
tls_profile: None,
|
||||
timeouts: Some(ExecutionTimeouts {
|
||||
connect_ms: Some(5_000),
|
||||
total_ms: Some(5_000),
|
||||
..ExecutionTimeouts::default()
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.expect("tunnel relay execution should succeed");
|
||||
|
||||
server.abort();
|
||||
|
||||
assert_eq!(result.status_code, 200);
|
||||
assert_eq!(
|
||||
result.body.and_then(|body| body.json_body),
|
||||
Some(json!({"ok": true}))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn direct_sync_execution_runtime_allows_tls_profile_best_effort() {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
|
||||
|
||||
@@ -126,6 +126,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
state,
|
||||
template,
|
||||
entry.refresh_token.as_str(),
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
{
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
use super::parse::{AdminProviderOAuthBatchImportEntry, AdminProviderOAuthBatchImportOutcome};
|
||||
use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider_oauth_key;
|
||||
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,
|
||||
@@ -8,9 +7,8 @@ use crate::handlers::admin::provider::oauth::provisioning::{
|
||||
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, AdminKiroOAuthRefreshAdapter,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminKiroAuthConfig};
|
||||
use crate::provider_transport::kiro::generate_machine_id;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::oauth::{
|
||||
build_kiro_batch_import_key_name, coerce_admin_provider_oauth_import_str,
|
||||
@@ -20,6 +18,9 @@ use serde_json::{json, Map, Value};
|
||||
use std::collections::BTreeSet;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
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";
|
||||
|
||||
fn admin_provider_oauth_kiro_refresh_base_url_override(
|
||||
state: &AdminAppState<'_>,
|
||||
override_key: &str,
|
||||
@@ -29,6 +30,293 @@ fn admin_provider_oauth_kiro_refresh_base_url_override(
|
||||
(!normalized.is_empty()).then(|| normalized.to_string())
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_kiro_build_refresh_url(
|
||||
auth_config: &AdminKiroAuthConfig,
|
||||
override_base_url: Option<&str>,
|
||||
path: &str,
|
||||
default_host: impl FnOnce(&str) -> String,
|
||||
) -> String {
|
||||
if let Some(base_url) = override_base_url
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return format!("{}/{}", base_url.trim_end_matches('/'), path);
|
||||
}
|
||||
let region = auth_config.effective_auth_region();
|
||||
default_host(region)
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_kiro_effective_host(url: &str, fallback_host: String) -> String {
|
||||
reqwest::Url::parse(url)
|
||||
.ok()
|
||||
.and_then(|value| value.host_str().map(ToOwned::to_owned))
|
||||
.unwrap_or(fallback_host)
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_kiro_ide_tag(kiro_version: &str, machine_id: &str) -> String {
|
||||
if machine_id.trim().is_empty() {
|
||||
format!("KiroIDE-{kiro_version}")
|
||||
} else {
|
||||
format!("KiroIDE-{kiro_version}-{machine_id}")
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_kiro_refresh_expires_at(payload: &Value) -> u64 {
|
||||
let expires_in = payload
|
||||
.get("expiresIn")
|
||||
.and_then(|value| {
|
||||
value
|
||||
.as_u64()
|
||||
.or_else(|| value.as_str()?.parse::<u64>().ok())
|
||||
})
|
||||
.unwrap_or(3600);
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|value| value.as_secs())
|
||||
.unwrap_or_default()
|
||||
.saturating_add(expires_in)
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_kiro_refresh_response_json(
|
||||
body_text: &str,
|
||||
json_body: Option<Value>,
|
||||
) -> Result<Value, String> {
|
||||
json_body
|
||||
.or_else(|| serde_json::from_str::<Value>(body_text).ok())
|
||||
.ok_or_else(|| "refresh 接口返回了非 JSON 响应".to_string())
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_kiro_refresh_error_detail(
|
||||
status: http::StatusCode,
|
||||
body_text: &str,
|
||||
) -> String {
|
||||
let detail = body_text.trim();
|
||||
if detail.is_empty() {
|
||||
format!("HTTP {}", status.as_u16())
|
||||
} else {
|
||||
detail.to_string()
|
||||
}
|
||||
}
|
||||
|
||||
async fn refresh_admin_provider_oauth_kiro_auth_config(
|
||||
state: &AdminAppState<'_>,
|
||||
auth_config: &AdminKiroAuthConfig,
|
||||
proxy_node_id: Option<&str>,
|
||||
social_refresh_base_url: Option<&str>,
|
||||
idc_refresh_base_url: Option<&str>,
|
||||
) -> Result<AdminKiroAuthConfig, String> {
|
||||
if auth_config.is_idc_auth() {
|
||||
let fallback_host = format!("oidc.{}.amazonaws.com", auth_config.effective_auth_region());
|
||||
let url = admin_provider_oauth_kiro_build_refresh_url(
|
||||
auth_config,
|
||||
idc_refresh_base_url,
|
||||
"token",
|
||||
|region| format!("https://oidc.{region}.amazonaws.com/token"),
|
||||
);
|
||||
let host = admin_provider_oauth_kiro_effective_host(&url, fallback_host);
|
||||
let headers = reqwest::header::HeaderMap::from_iter([
|
||||
(
|
||||
reqwest::header::CONTENT_TYPE,
|
||||
reqwest::header::HeaderValue::from_static("application/json"),
|
||||
),
|
||||
(
|
||||
reqwest::header::HOST,
|
||||
reqwest::header::HeaderValue::from_str(&host)
|
||||
.map_err(|_| "IDC host 无效".to_string())?,
|
||||
),
|
||||
(
|
||||
reqwest::header::HeaderName::from_static("x-amz-user-agent"),
|
||||
reqwest::header::HeaderValue::from_static(KIRO_IDC_AMZ_USER_AGENT),
|
||||
),
|
||||
(
|
||||
reqwest::header::USER_AGENT,
|
||||
reqwest::header::HeaderValue::from_static("node"),
|
||||
),
|
||||
(
|
||||
reqwest::header::ACCEPT,
|
||||
reqwest::header::HeaderValue::from_static("*/*"),
|
||||
),
|
||||
]);
|
||||
let response = state
|
||||
.execute_admin_provider_oauth_http_request(
|
||||
"kiro_batch_refresh:idc",
|
||||
reqwest::Method::POST,
|
||||
&url,
|
||||
&headers,
|
||||
Some("application/json"),
|
||||
Some(json!({
|
||||
"clientId": auth_config
|
||||
.client_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.unwrap_or_default(),
|
||||
"clientSecret": auth_config
|
||||
.client_secret
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.unwrap_or_default(),
|
||||
"refreshToken": auth_config
|
||||
.refresh_token
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.unwrap_or_default(),
|
||||
"grantType": "refresh_token",
|
||||
})),
|
||||
None,
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| format!("IDC refresh 请求失败: {err}"))?;
|
||||
if !response.status.is_success() {
|
||||
return Err(format!(
|
||||
"IDC refresh 失败: {}",
|
||||
admin_provider_oauth_kiro_refresh_error_detail(
|
||||
response.status,
|
||||
&response.body_text
|
||||
)
|
||||
));
|
||||
}
|
||||
let payload = admin_provider_oauth_kiro_refresh_response_json(
|
||||
&response.body_text,
|
||||
response.json_body,
|
||||
)?;
|
||||
let access_token = payload
|
||||
.get("accessToken")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or_else(|| "IDC refresh 返回了空 accessToken".to_string())?;
|
||||
|
||||
let mut refreshed = auth_config.clone();
|
||||
refreshed.access_token = Some(access_token.to_string());
|
||||
refreshed.expires_at = Some(admin_provider_oauth_kiro_refresh_expires_at(&payload));
|
||||
if refreshed
|
||||
.machine_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_none_or(|value| value.is_empty())
|
||||
{
|
||||
refreshed.machine_id = generate_machine_id(auth_config, None);
|
||||
}
|
||||
if let Some(refresh_token) = payload
|
||||
.get("refreshToken")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
refreshed.refresh_token = Some(refresh_token.to_string());
|
||||
}
|
||||
return Ok(refreshed);
|
||||
}
|
||||
|
||||
let machine_id = generate_machine_id(auth_config, None)
|
||||
.ok_or_else(|| "缺少 machine_id 种子,无法刷新 social token".to_string())?;
|
||||
let fallback_host = format!(
|
||||
"prod.{}.auth.desktop.kiro.dev",
|
||||
auth_config.effective_auth_region()
|
||||
);
|
||||
let url = admin_provider_oauth_kiro_build_refresh_url(
|
||||
auth_config,
|
||||
social_refresh_base_url,
|
||||
"refreshToken",
|
||||
|region| format!("https://prod.{region}.auth.desktop.kiro.dev/refreshToken"),
|
||||
);
|
||||
let host = admin_provider_oauth_kiro_effective_host(&url, fallback_host);
|
||||
let user_agent =
|
||||
admin_provider_oauth_kiro_ide_tag(auth_config.effective_kiro_version(), &machine_id);
|
||||
let headers = reqwest::header::HeaderMap::from_iter([
|
||||
(
|
||||
reqwest::header::USER_AGENT,
|
||||
reqwest::header::HeaderValue::from_str(&user_agent)
|
||||
.map_err(|_| "Kiro User-Agent 无效".to_string())?,
|
||||
),
|
||||
(
|
||||
reqwest::header::HOST,
|
||||
reqwest::header::HeaderValue::from_str(&host)
|
||||
.map_err(|_| "Kiro host 无效".to_string())?,
|
||||
),
|
||||
(
|
||||
reqwest::header::ACCEPT,
|
||||
reqwest::header::HeaderValue::from_static("application/json, text/plain, */*"),
|
||||
),
|
||||
(
|
||||
reqwest::header::CONTENT_TYPE,
|
||||
reqwest::header::HeaderValue::from_static("application/json"),
|
||||
),
|
||||
(
|
||||
reqwest::header::CONNECTION,
|
||||
reqwest::header::HeaderValue::from_static("close"),
|
||||
),
|
||||
(
|
||||
reqwest::header::ACCEPT_ENCODING,
|
||||
reqwest::header::HeaderValue::from_static("gzip, compress, deflate, br"),
|
||||
),
|
||||
]);
|
||||
let response = state
|
||||
.execute_admin_provider_oauth_http_request(
|
||||
"kiro_batch_refresh:social",
|
||||
reqwest::Method::POST,
|
||||
&url,
|
||||
&headers,
|
||||
Some("application/json"),
|
||||
Some(json!({
|
||||
"refreshToken": auth_config
|
||||
.refresh_token
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.unwrap_or_default(),
|
||||
})),
|
||||
None,
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| format!("social refresh 请求失败: {err}"))?;
|
||||
if !response.status.is_success() {
|
||||
return Err(format!(
|
||||
"social refresh 失败: {}",
|
||||
admin_provider_oauth_kiro_refresh_error_detail(response.status, &response.body_text)
|
||||
));
|
||||
}
|
||||
let payload =
|
||||
admin_provider_oauth_kiro_refresh_response_json(&response.body_text, response.json_body)?;
|
||||
let access_token = payload
|
||||
.get("accessToken")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or_else(|| "social refresh 返回了空 accessToken".to_string())?;
|
||||
|
||||
let mut refreshed = auth_config.clone();
|
||||
refreshed.access_token = Some(access_token.to_string());
|
||||
refreshed.expires_at = Some(admin_provider_oauth_kiro_refresh_expires_at(&payload));
|
||||
if refreshed
|
||||
.machine_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_none_or(|value| value.is_empty())
|
||||
{
|
||||
refreshed.machine_id = Some(machine_id);
|
||||
}
|
||||
if let Some(refresh_token) = payload
|
||||
.get("refreshToken")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
refreshed.refresh_token = Some(refresh_token.to_string());
|
||||
}
|
||||
if let Some(profile_arn) = payload
|
||||
.get("profileArn")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
refreshed.profile_arn = Some(profile_arn.to_string());
|
||||
}
|
||||
Ok(refreshed)
|
||||
}
|
||||
|
||||
pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
@@ -65,10 +353,10 @@ pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
|
||||
.list_provider_catalog_endpoints_by_provider_ids(&[provider_id.to_string()])
|
||||
.await?;
|
||||
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id);
|
||||
let adapter = AdminKiroOAuthRefreshAdapter::default().with_refresh_base_urls(
|
||||
admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_social_refresh"),
|
||||
admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_idc_refresh"),
|
||||
);
|
||||
let social_refresh_base_url =
|
||||
admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_social_refresh");
|
||||
let idc_refresh_base_url =
|
||||
admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_idc_refresh");
|
||||
let mut results = Vec::with_capacity(entries.len());
|
||||
let mut success = 0usize;
|
||||
let mut failed = 0usize;
|
||||
@@ -101,9 +389,14 @@ pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
|
||||
continue;
|
||||
}
|
||||
|
||||
refreshed_auth_config = match adapter
|
||||
.refresh_auth_config(state.http_client(), &refreshed_auth_config)
|
||||
.await
|
||||
refreshed_auth_config = match refresh_admin_provider_oauth_kiro_auth_config(
|
||||
state,
|
||||
&refreshed_auth_config,
|
||||
proxy_node_id,
|
||||
social_refresh_base_url.as_deref(),
|
||||
idc_refresh_base_url.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(config) => config,
|
||||
Err(err) => {
|
||||
@@ -111,7 +404,7 @@ pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": format!("Token 验证失败: {err:?}"),
|
||||
"error": format!("Token 验证失败: {err}"),
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
|
||||
@@ -130,6 +130,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(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
|
||||
@@ -121,6 +121,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(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
|
||||
@@ -85,7 +85,11 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
||||
};
|
||||
|
||||
let client_registration = match state
|
||||
.register_admin_kiro_device_oidc_client(®ion, &start_url)
|
||||
.register_admin_kiro_device_oidc_client(
|
||||
®ion,
|
||||
&start_url,
|
||||
payload.proxy_node_id.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => payload,
|
||||
@@ -105,7 +109,13 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
||||
};
|
||||
|
||||
let device_authorization = match state
|
||||
.start_admin_kiro_device_authorization(®ion, &client_id, &client_secret, &start_url)
|
||||
.start_admin_kiro_device_authorization(
|
||||
®ion,
|
||||
&client_id,
|
||||
&client_secret,
|
||||
&start_url,
|
||||
payload.proxy_node_id.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => payload,
|
||||
|
||||
@@ -132,6 +132,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(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
|
||||
@@ -98,7 +98,11 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
};
|
||||
|
||||
let token_payload = match state
|
||||
.exchange_admin_provider_oauth_refresh_token(template, refresh_token_input)
|
||||
.exchange_admin_provider_oauth_refresh_token(
|
||||
template,
|
||||
refresh_token_input,
|
||||
proxy_node_id.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => payload,
|
||||
|
||||
@@ -4,6 +4,7 @@ use super::super::errors::{
|
||||
use super::json_non_empty_string;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
|
||||
use axum::{body::Body, http, response::Response};
|
||||
use url::form_urlencoded;
|
||||
|
||||
pub(crate) async fn exchange_admin_provider_oauth_code(
|
||||
state: &AdminAppState<'_>,
|
||||
@@ -11,9 +12,9 @@ pub(crate) async fn exchange_admin_provider_oauth_code(
|
||||
code: &str,
|
||||
state_nonce: &str,
|
||||
pkce_verifier: Option<&str>,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<serde_json::Value, Response<Body>> {
|
||||
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
|
||||
let request = state.http_client().post(token_url);
|
||||
let response = if template.provider_type == "claude_code" {
|
||||
let mut body = serde_json::Map::from_iter([
|
||||
(
|
||||
@@ -43,44 +44,78 @@ pub(crate) async fn exchange_admin_provider_oauth_code(
|
||||
serde_json::Value::String(verifier.to_string()),
|
||||
);
|
||||
}
|
||||
request
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "application/json")
|
||||
.json(&serde_json::Value::Object(body))
|
||||
.send()
|
||||
let headers = reqwest::header::HeaderMap::from_iter([
|
||||
(
|
||||
reqwest::header::CONTENT_TYPE,
|
||||
reqwest::header::HeaderValue::from_static("application/json"),
|
||||
),
|
||||
(
|
||||
reqwest::header::ACCEPT,
|
||||
reqwest::header::HeaderValue::from_static("application/json"),
|
||||
),
|
||||
]);
|
||||
state
|
||||
.execute_admin_provider_oauth_http_request(
|
||||
"provider-oauth:exchange-code",
|
||||
reqwest::Method::POST,
|
||||
&token_url,
|
||||
&headers,
|
||||
Some("application/json"),
|
||||
Some(serde_json::Value::Object(body)),
|
||||
None,
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
let mut form = vec![
|
||||
("grant_type", "authorization_code".to_string()),
|
||||
("client_id", template.client_id.to_string()),
|
||||
("redirect_uri", template.redirect_uri.to_string()),
|
||||
("code", code.to_string()),
|
||||
];
|
||||
if !template.client_secret.trim().is_empty() {
|
||||
form.push(("client_secret", template.client_secret.to_string()));
|
||||
}
|
||||
if let Some(verifier) = pkce_verifier {
|
||||
form.push(("code_verifier", verifier.to_string()));
|
||||
}
|
||||
request
|
||||
.header("Content-Type", "application/x-www-form-urlencoded")
|
||||
.header("Accept", "application/json")
|
||||
.form(&form)
|
||||
.send()
|
||||
let form_body = {
|
||||
let mut form = form_urlencoded::Serializer::new(String::new());
|
||||
form.append_pair("grant_type", "authorization_code");
|
||||
form.append_pair("client_id", template.client_id);
|
||||
form.append_pair("redirect_uri", template.redirect_uri);
|
||||
form.append_pair("code", code);
|
||||
if !template.client_secret.trim().is_empty() {
|
||||
form.append_pair("client_secret", template.client_secret);
|
||||
}
|
||||
if let Some(verifier) = pkce_verifier {
|
||||
form.append_pair("code_verifier", verifier);
|
||||
}
|
||||
form.finish().into_bytes()
|
||||
};
|
||||
let headers = reqwest::header::HeaderMap::from_iter([
|
||||
(
|
||||
reqwest::header::CONTENT_TYPE,
|
||||
reqwest::header::HeaderValue::from_static("application/x-www-form-urlencoded"),
|
||||
),
|
||||
(
|
||||
reqwest::header::ACCEPT,
|
||||
reqwest::header::HeaderValue::from_static("application/json"),
|
||||
),
|
||||
]);
|
||||
state
|
||||
.execute_admin_provider_oauth_http_request(
|
||||
"provider-oauth:exchange-code",
|
||||
reqwest::Method::POST,
|
||||
&token_url,
|
||||
&headers,
|
||||
Some("application/x-www-form-urlencoded"),
|
||||
None,
|
||||
Some(form_body),
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
}
|
||||
.map_err(|_| {
|
||||
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, "token exchange 失败")
|
||||
})?;
|
||||
|
||||
if !response.status().is_success() {
|
||||
if !response.status.is_success() {
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"token exchange 失败",
|
||||
));
|
||||
}
|
||||
|
||||
let payload = response.json::<serde_json::Value>().await.map_err(|_| {
|
||||
let payload = response.json_body.ok_or_else(|| {
|
||||
build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"token exchange 返回缺少 access_token",
|
||||
@@ -99,9 +134,9 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
|
||||
state: &AdminAppState<'_>,
|
||||
template: AdminProviderOAuthTemplate,
|
||||
refresh_token: &str,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<serde_json::Value, Response<Body>> {
|
||||
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
|
||||
let request = state.http_client().post(token_url);
|
||||
let scope = template.scopes.join(" ");
|
||||
let response = if template.provider_type == "claude_code" {
|
||||
let mut body = serde_json::Map::from_iter([
|
||||
@@ -121,29 +156,63 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
|
||||
if !scope.trim().is_empty() {
|
||||
body.insert("scope".to_string(), serde_json::Value::String(scope));
|
||||
}
|
||||
request
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "application/json")
|
||||
.json(&serde_json::Value::Object(body))
|
||||
.send()
|
||||
let headers = reqwest::header::HeaderMap::from_iter([
|
||||
(
|
||||
reqwest::header::CONTENT_TYPE,
|
||||
reqwest::header::HeaderValue::from_static("application/json"),
|
||||
),
|
||||
(
|
||||
reqwest::header::ACCEPT,
|
||||
reqwest::header::HeaderValue::from_static("application/json"),
|
||||
),
|
||||
]);
|
||||
state
|
||||
.execute_admin_provider_oauth_http_request(
|
||||
"provider-oauth:refresh-token",
|
||||
reqwest::Method::POST,
|
||||
&token_url,
|
||||
&headers,
|
||||
Some("application/json"),
|
||||
Some(serde_json::Value::Object(body)),
|
||||
None,
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
let mut form = vec![
|
||||
("grant_type", "refresh_token".to_string()),
|
||||
("client_id", template.client_id.to_string()),
|
||||
("refresh_token", refresh_token.to_string()),
|
||||
];
|
||||
if !scope.trim().is_empty() {
|
||||
form.push(("scope", scope));
|
||||
}
|
||||
if !template.client_secret.trim().is_empty() {
|
||||
form.push(("client_secret", template.client_secret.to_string()));
|
||||
}
|
||||
request
|
||||
.header("Content-Type", "application/x-www-form-urlencoded")
|
||||
.header("Accept", "application/json")
|
||||
.form(&form)
|
||||
.send()
|
||||
let form_body = {
|
||||
let mut form = form_urlencoded::Serializer::new(String::new());
|
||||
form.append_pair("grant_type", "refresh_token");
|
||||
form.append_pair("client_id", template.client_id);
|
||||
form.append_pair("refresh_token", refresh_token);
|
||||
if !scope.trim().is_empty() {
|
||||
form.append_pair("scope", &scope);
|
||||
}
|
||||
if !template.client_secret.trim().is_empty() {
|
||||
form.append_pair("client_secret", template.client_secret);
|
||||
}
|
||||
form.finish().into_bytes()
|
||||
};
|
||||
let headers = reqwest::header::HeaderMap::from_iter([
|
||||
(
|
||||
reqwest::header::CONTENT_TYPE,
|
||||
reqwest::header::HeaderValue::from_static("application/x-www-form-urlencoded"),
|
||||
),
|
||||
(
|
||||
reqwest::header::ACCEPT,
|
||||
reqwest::header::HeaderValue::from_static("application/json"),
|
||||
),
|
||||
]);
|
||||
state
|
||||
.execute_admin_provider_oauth_http_request(
|
||||
"provider-oauth:refresh-token",
|
||||
reqwest::Method::POST,
|
||||
&token_url,
|
||||
&headers,
|
||||
Some("application/x-www-form-urlencoded"),
|
||||
None,
|
||||
Some(form_body),
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
}
|
||||
.map_err(|_| {
|
||||
@@ -153,13 +222,8 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
|
||||
)
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
let body = response.text().await.map_err(|_| {
|
||||
build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Refresh Token 验证失败: token exchange 失败",
|
||||
)
|
||||
})?;
|
||||
let status = response.status;
|
||||
let body = response.body_text;
|
||||
if !status.is_success() {
|
||||
let reason =
|
||||
normalize_provider_oauth_refresh_error_message(Some(status.as_u16()), Some(&body));
|
||||
@@ -169,12 +233,15 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
|
||||
));
|
||||
}
|
||||
|
||||
let payload = serde_json::from_str::<serde_json::Value>(&body).map_err(|_| {
|
||||
build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"token refresh 返回缺少 access_token",
|
||||
)
|
||||
})?;
|
||||
let payload = response
|
||||
.json_body
|
||||
.or_else(|| serde_json::from_str::<serde_json::Value>(&body).ok())
|
||||
.ok_or_else(|| {
|
||||
build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"token refresh 返回缺少 access_token",
|
||||
)
|
||||
})?;
|
||||
if json_non_empty_string(payload.get("access_token")).is_none() {
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use super::super::super::support::AdminProviderOpsCheckinOutcome;
|
||||
use super::super::super::verify::admin_provider_ops_execute_proxy_json_request;
|
||||
use super::super::super::verify::{
|
||||
admin_provider_ops_execute_json_request, AdminProviderOpsExecuteJsonError,
|
||||
};
|
||||
use super::super::support::{admin_provider_ops_json_object_map, admin_provider_ops_request_url};
|
||||
use super::shared::{
|
||||
admin_provider_ops_checkin_already_done, admin_provider_ops_checkin_auth_failure,
|
||||
@@ -27,40 +29,20 @@ pub(in super::super) async fn admin_provider_ops_probe_new_api_checkin(
|
||||
&admin_provider_ops_json_object_map(json!({ "endpoint": endpoint })),
|
||||
endpoint,
|
||||
);
|
||||
let (status, response_json) = if let Some(proxy_snapshot) = proxy_snapshot {
|
||||
match admin_provider_ops_execute_proxy_json_request(
|
||||
state,
|
||||
"provider-ops-action:probe_checkin",
|
||||
reqwest::Method::POST,
|
||||
&url,
|
||||
headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(_) => return None,
|
||||
}
|
||||
} else {
|
||||
let response = match state
|
||||
.http_client()
|
||||
.request(reqwest::Method::POST, url)
|
||||
.headers(headers.clone())
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(response) => response,
|
||||
Err(_) => return None,
|
||||
};
|
||||
let status = response.status();
|
||||
let response_json = match response.bytes().await {
|
||||
Ok(bytes) => {
|
||||
serde_json::from_slice::<serde_json::Value>(&bytes).unwrap_or_else(|_| json!({}))
|
||||
}
|
||||
Err(_) => json!({}),
|
||||
};
|
||||
(status, response_json)
|
||||
let (status, response_json) = match admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
"provider-ops-action:probe_checkin",
|
||||
reqwest::Method::POST,
|
||||
&url,
|
||||
headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(AdminProviderOpsExecuteJsonError::InvalidJson(_))
|
||||
| Err(AdminProviderOpsExecuteJsonError::Transport(_)) => return None,
|
||||
};
|
||||
|
||||
if status == http::StatusCode::NOT_FOUND {
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use super::super::super::support::ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE;
|
||||
use super::super::super::verify::admin_provider_ops_execute_proxy_json_request;
|
||||
use super::super::super::verify::{
|
||||
admin_provider_ops_execute_json_request, AdminProviderOpsExecuteJsonError,
|
||||
};
|
||||
use super::super::responses::{
|
||||
admin_provider_ops_action_error, admin_provider_ops_action_not_supported,
|
||||
admin_provider_ops_action_response,
|
||||
@@ -32,80 +34,37 @@ pub(in super::super) async fn admin_provider_ops_run_checkin_action(
|
||||
|
||||
let url = admin_provider_ops_request_url(base_url, action_config, "/api/user/checkin");
|
||||
let method = admin_provider_ops_request_method(action_config, "POST");
|
||||
let (status, response_json) = if let Some(proxy_snapshot) = proxy_snapshot {
|
||||
match admin_provider_ops_execute_proxy_json_request(
|
||||
state,
|
||||
&format!(
|
||||
"provider-ops-action:{}:checkin",
|
||||
architecture.architecture_id
|
||||
),
|
||||
method,
|
||||
&url,
|
||||
headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"checkin",
|
||||
admin_provider_ops_network_error_message(&err),
|
||||
None,
|
||||
);
|
||||
}
|
||||
let (status, response_json) = match admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
&format!(
|
||||
"provider-ops-action:{}:checkin",
|
||||
architecture.architecture_id
|
||||
),
|
||||
method,
|
||||
&url,
|
||||
headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(AdminProviderOpsExecuteJsonError::InvalidJson(_)) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"parse_error",
|
||||
"checkin",
|
||||
"响应不是有效的 JSON",
|
||||
Some(start.elapsed().as_millis() as u64),
|
||||
);
|
||||
}
|
||||
Err(AdminProviderOpsExecuteJsonError::Transport(err)) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"checkin",
|
||||
admin_provider_ops_network_error_message(&err),
|
||||
None,
|
||||
);
|
||||
}
|
||||
} else {
|
||||
let response = match state
|
||||
.http_client()
|
||||
.request(method, url)
|
||||
.headers(headers.clone())
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(response) => response,
|
||||
Err(err) if err.is_timeout() => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"checkin",
|
||||
"请求超时",
|
||||
None,
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"checkin",
|
||||
format!("网络错误: {err}"),
|
||||
None,
|
||||
);
|
||||
}
|
||||
};
|
||||
let status = response.status();
|
||||
let response_json = match response.bytes().await {
|
||||
Ok(bytes) => match serde_json::from_slice::<serde_json::Value>(&bytes) {
|
||||
Ok(value) => value,
|
||||
Err(_) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"parse_error",
|
||||
"checkin",
|
||||
"响应不是有效的 JSON",
|
||||
Some(start.elapsed().as_millis() as u64),
|
||||
);
|
||||
}
|
||||
},
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"checkin",
|
||||
format!("网络错误: {err}"),
|
||||
Some(start.elapsed().as_millis() as u64),
|
||||
);
|
||||
}
|
||||
};
|
||||
(status, response_json)
|
||||
};
|
||||
let response_time_ms = Some(start.elapsed().as_millis() as u64);
|
||||
|
||||
|
||||
@@ -2,7 +2,9 @@ mod sub2api;
|
||||
mod yescode;
|
||||
|
||||
use super::super::support::AdminProviderOpsCheckinOutcome;
|
||||
use super::super::verify::admin_provider_ops_execute_proxy_json_request;
|
||||
use super::super::verify::{
|
||||
admin_provider_ops_execute_json_request, AdminProviderOpsExecuteJsonError,
|
||||
};
|
||||
use super::checkin::admin_provider_ops_probe_new_api_checkin;
|
||||
use super::responses::{admin_provider_ops_action_error, admin_provider_ops_action_response};
|
||||
use super::support::{admin_provider_ops_request_method, admin_provider_ops_request_url};
|
||||
@@ -73,80 +75,37 @@ pub(super) async fn admin_provider_ops_run_query_balance_action(
|
||||
let start = std::time::Instant::now();
|
||||
let url = admin_provider_ops_request_url(base_url, action_config, "/api/user/balance");
|
||||
let method = admin_provider_ops_request_method(action_config, "GET");
|
||||
let (status, response_json) = if let Some(proxy_snapshot) = proxy_snapshot {
|
||||
match admin_provider_ops_execute_proxy_json_request(
|
||||
state,
|
||||
&format!(
|
||||
"provider-ops-action:{}:query_balance:{provider_id}",
|
||||
architecture.architecture_id
|
||||
),
|
||||
method,
|
||||
&url,
|
||||
headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
admin_provider_ops_network_error_message(&err),
|
||||
None,
|
||||
);
|
||||
}
|
||||
let (status, response_json) = match admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
&format!(
|
||||
"provider-ops-action:{}:query_balance:{provider_id}",
|
||||
architecture.architecture_id
|
||||
),
|
||||
method,
|
||||
&url,
|
||||
headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(AdminProviderOpsExecuteJsonError::InvalidJson(_)) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"parse_error",
|
||||
"query_balance",
|
||||
"响应不是有效的 JSON",
|
||||
Some(start.elapsed().as_millis() as u64),
|
||||
);
|
||||
}
|
||||
Err(AdminProviderOpsExecuteJsonError::Transport(err)) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
admin_provider_ops_network_error_message(&err),
|
||||
None,
|
||||
);
|
||||
}
|
||||
} else {
|
||||
let response = match state
|
||||
.http_client()
|
||||
.request(method, url)
|
||||
.headers(headers.clone())
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(response) => response,
|
||||
Err(err) if err.is_timeout() => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
"请求超时",
|
||||
None,
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
format!("网络错误: {err}"),
|
||||
None,
|
||||
);
|
||||
}
|
||||
};
|
||||
let status = response.status();
|
||||
let response_json = match response.bytes().await {
|
||||
Ok(bytes) => match serde_json::from_slice::<serde_json::Value>(&bytes) {
|
||||
Ok(value) => value,
|
||||
Err(_) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"parse_error",
|
||||
"query_balance",
|
||||
"响应不是有效的 JSON",
|
||||
Some(start.elapsed().as_millis() as u64),
|
||||
);
|
||||
}
|
||||
},
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
format!("网络错误: {err}"),
|
||||
Some(start.elapsed().as_millis() as u64),
|
||||
);
|
||||
}
|
||||
};
|
||||
(status, response_json)
|
||||
};
|
||||
let response_time_ms = Some(start.elapsed().as_millis() as u64);
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use super::super::super::config::persist_admin_provider_ops_runtime_credentials;
|
||||
use super::super::super::verify::{
|
||||
admin_provider_ops_execute_proxy_json_request, admin_provider_ops_sub2api_exchange_token,
|
||||
admin_provider_ops_sub2api_request_url,
|
||||
admin_provider_ops_execute_json_request, admin_provider_ops_sub2api_exchange_token,
|
||||
admin_provider_ops_sub2api_request_url, AdminProviderOpsExecuteJsonError,
|
||||
};
|
||||
use super::super::responses::{
|
||||
admin_provider_ops_action_error, admin_provider_ops_action_response,
|
||||
@@ -93,97 +93,37 @@ pub(super) async fn admin_provider_ops_sub2api_balance_payload(
|
||||
};
|
||||
let auth_headers =
|
||||
reqwest::header::HeaderMap::from_iter([(reqwest::header::AUTHORIZATION, auth_value)]);
|
||||
let (me_result, subscription_result) = if let Some(proxy_snapshot) = proxy_snapshot {
|
||||
let me_request_id = format!("provider-ops-action:sub2api:me:{provider_id}");
|
||||
let subscription_request_id =
|
||||
format!("provider-ops-action:sub2api:subscriptions:{provider_id}");
|
||||
tokio::join!(
|
||||
admin_provider_ops_execute_proxy_json_request(
|
||||
state,
|
||||
&me_request_id,
|
||||
reqwest::Method::GET,
|
||||
&me_url,
|
||||
&auth_headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
),
|
||||
admin_provider_ops_execute_proxy_json_request(
|
||||
state,
|
||||
&subscription_request_id,
|
||||
reqwest::Method::GET,
|
||||
&subscription_url,
|
||||
&auth_headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
let me_request_id = format!("provider-ops-action:sub2api:me:{provider_id}");
|
||||
let subscription_request_id =
|
||||
format!("provider-ops-action:sub2api:subscriptions:{provider_id}");
|
||||
let (me_result, subscription_result) = tokio::join!(
|
||||
admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
&me_request_id,
|
||||
reqwest::Method::GET,
|
||||
&me_url,
|
||||
&auth_headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
),
|
||||
admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
&subscription_request_id,
|
||||
reqwest::Method::GET,
|
||||
&subscription_url,
|
||||
&auth_headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
} else {
|
||||
let http_client = state.http_client();
|
||||
let (me_response, subscription_response) = tokio::join!(
|
||||
http_client.get(me_url).bearer_auth(&access_token).send(),
|
||||
http_client
|
||||
.get(subscription_url)
|
||||
.bearer_auth(&access_token)
|
||||
.send()
|
||||
);
|
||||
let me_result = match me_response {
|
||||
Ok(response) => {
|
||||
let status = response.status();
|
||||
let value = match response.bytes().await {
|
||||
Ok(bytes) => {
|
||||
serde_json::from_slice::<Value>(&bytes).unwrap_or_else(|_| json!({}))
|
||||
}
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
format!("网络错误: {err}"),
|
||||
Some(start.elapsed().as_millis() as u64),
|
||||
)
|
||||
}
|
||||
};
|
||||
Ok((status, value))
|
||||
}
|
||||
Err(err) if err.is_timeout() => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
"请求超时",
|
||||
Some(start.elapsed().as_millis() as u64),
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
format!("网络错误: {err}"),
|
||||
Some(start.elapsed().as_millis() as u64),
|
||||
);
|
||||
}
|
||||
};
|
||||
let subscription_result = match subscription_response {
|
||||
Ok(response) => {
|
||||
let status = response.status();
|
||||
let value = match response.bytes().await {
|
||||
Ok(bytes) => {
|
||||
serde_json::from_slice::<Value>(&bytes).unwrap_or_else(|_| json!({}))
|
||||
}
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
format!("网络错误: {err}"),
|
||||
Some(start.elapsed().as_millis() as u64),
|
||||
)
|
||||
}
|
||||
};
|
||||
Ok((status, value))
|
||||
}
|
||||
Err(err) if err.is_timeout() => Err("请求超时".to_string()),
|
||||
Err(err) => Err(format!("网络错误: {err}")),
|
||||
};
|
||||
(me_result, subscription_result)
|
||||
};
|
||||
);
|
||||
let me_result = me_result.map_err(|err| match err {
|
||||
AdminProviderOpsExecuteJsonError::InvalidJson(message)
|
||||
| AdminProviderOpsExecuteJsonError::Transport(message) => message,
|
||||
});
|
||||
let subscription_result = subscription_result.map_err(|err| match err {
|
||||
AdminProviderOpsExecuteJsonError::InvalidJson(message)
|
||||
| AdminProviderOpsExecuteJsonError::Transport(message) => message,
|
||||
});
|
||||
let response_time_ms = Some(start.elapsed().as_millis() as u64);
|
||||
|
||||
let (me_status, me_json) = match me_result {
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
use super::super::super::verify::admin_provider_ops_execute_proxy_json_request;
|
||||
use super::super::super::verify::{
|
||||
admin_provider_ops_execute_json_request, AdminProviderOpsExecuteJsonError,
|
||||
};
|
||||
use super::super::responses::{
|
||||
admin_provider_ops_action_error, admin_provider_ops_action_response,
|
||||
};
|
||||
@@ -17,65 +19,34 @@ pub(super) async fn admin_provider_ops_yescode_balance_payload(
|
||||
let start = std::time::Instant::now();
|
||||
let balance_url = format!("{}/api/v1/user/balance", base_url.trim_end_matches('/'));
|
||||
let profile_url = format!("{}/api/v1/auth/profile", base_url.trim_end_matches('/'));
|
||||
let (balance_result, profile_result) = if let Some(proxy_snapshot) = proxy_snapshot {
|
||||
tokio::join!(
|
||||
admin_provider_ops_execute_proxy_json_request(
|
||||
state,
|
||||
"provider-ops-action:yescode:balance",
|
||||
reqwest::Method::GET,
|
||||
&balance_url,
|
||||
headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
),
|
||||
admin_provider_ops_execute_proxy_json_request(
|
||||
state,
|
||||
"provider-ops-action:yescode:profile",
|
||||
reqwest::Method::GET,
|
||||
&profile_url,
|
||||
headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
let (balance_result, profile_result) = tokio::join!(
|
||||
admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
"provider-ops-action:yescode:balance",
|
||||
reqwest::Method::GET,
|
||||
&balance_url,
|
||||
headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
),
|
||||
admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
"provider-ops-action:yescode:profile",
|
||||
reqwest::Method::GET,
|
||||
&profile_url,
|
||||
headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
} else {
|
||||
let balance_future = state
|
||||
.http_client()
|
||||
.request(reqwest::Method::GET, balance_url)
|
||||
.headers(headers.clone())
|
||||
.send();
|
||||
let profile_future = state
|
||||
.http_client()
|
||||
.request(reqwest::Method::GET, profile_url)
|
||||
.headers(headers.clone())
|
||||
.send();
|
||||
let (balance_result, profile_result) = tokio::join!(balance_future, profile_future);
|
||||
let balance_result = match balance_result {
|
||||
Ok(response) => {
|
||||
let status = response.status();
|
||||
let value = match response.bytes().await {
|
||||
Ok(bytes) => serde_json::from_slice::<serde_json::Value>(&bytes)
|
||||
.unwrap_or_else(|_| json!({})),
|
||||
Err(_) => json!({}),
|
||||
};
|
||||
Ok((status, value))
|
||||
}
|
||||
Err(err) => Err(err.to_string()),
|
||||
};
|
||||
let profile_result = match profile_result {
|
||||
Ok(response) => {
|
||||
let status = response.status();
|
||||
let value = match response.bytes().await {
|
||||
Ok(bytes) => serde_json::from_slice::<serde_json::Value>(&bytes)
|
||||
.unwrap_or_else(|_| json!({})),
|
||||
Err(_) => json!({}),
|
||||
};
|
||||
Ok((status, value))
|
||||
}
|
||||
Err(err) => Err(err.to_string()),
|
||||
};
|
||||
(balance_result, profile_result)
|
||||
};
|
||||
);
|
||||
let balance_result = balance_result.map_err(|err| match err {
|
||||
AdminProviderOpsExecuteJsonError::InvalidJson(message)
|
||||
| AdminProviderOpsExecuteJsonError::Transport(message) => message,
|
||||
});
|
||||
let profile_result = profile_result.map_err(|err| match err {
|
||||
AdminProviderOpsExecuteJsonError::InvalidJson(message)
|
||||
| AdminProviderOpsExecuteJsonError::Transport(message) => message,
|
||||
});
|
||||
let response_time_ms = Some(start.elapsed().as_millis() as u64);
|
||||
|
||||
let mut combined = serde_json::Map::new();
|
||||
|
||||
@@ -12,7 +12,10 @@ use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogPr
|
||||
pub(super) use proxy::{
|
||||
admin_provider_ops_anyrouter_acw_cookie, admin_provider_ops_resolve_proxy_snapshot,
|
||||
};
|
||||
pub(super) use request::admin_provider_ops_execute_proxy_json_request;
|
||||
pub(super) use request::{
|
||||
admin_provider_ops_execute_json_request, admin_provider_ops_execute_proxy_json_request,
|
||||
AdminProviderOpsExecuteJsonError,
|
||||
};
|
||||
pub(super) use sub2api::{
|
||||
admin_provider_ops_sub2api_exchange_token, admin_provider_ops_sub2api_request_url,
|
||||
};
|
||||
|
||||
@@ -1,18 +1,9 @@
|
||||
use super::request::{
|
||||
admin_provider_ops_execute_get_text, admin_provider_ops_execute_get_text_no_redirect,
|
||||
};
|
||||
use super::request::admin_provider_ops_execute_get_text_no_redirect;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use aether_admin::provider::ops::admin_provider_ops_anyrouter_compute_acw_sc_v2;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_data::repository::proxy_nodes::StoredProxyNode;
|
||||
use aether_provider_transport::TransportTunnelAffinityLookup;
|
||||
use regex::Regex;
|
||||
use serde_json::{json, Map, Value};
|
||||
use url::Url;
|
||||
|
||||
const TUNNEL_BASE_URL_EXTRA_KEY: &str = "tunnel_base_url";
|
||||
const TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY: &str = "tunnel_owner_instance_id";
|
||||
const TUNNEL_OWNER_OBSERVED_AT_EXTRA_KEY: &str = "tunnel_owner_observed_at_unix_secs";
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub(in super::super) struct AdminProviderOpsAnyrouterChallenge {
|
||||
pub(in super::super) acw_cookie: String,
|
||||
@@ -30,25 +21,15 @@ pub(in super::super) async fn admin_provider_ops_anyrouter_acw_cookie(
|
||||
"Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36",
|
||||
),
|
||||
)]);
|
||||
let response = if admin_provider_ops_proxy_uses_tunnel(proxy_snapshot.as_ref()) {
|
||||
admin_provider_ops_execute_get_text(
|
||||
state,
|
||||
"provider-ops-acw:anyrouter",
|
||||
base_url.trim_end_matches('/'),
|
||||
&headers,
|
||||
proxy_snapshot.as_ref(),
|
||||
)
|
||||
.await
|
||||
.ok()?
|
||||
} else {
|
||||
admin_provider_ops_execute_get_text_no_redirect(
|
||||
base_url.trim_end_matches('/'),
|
||||
&headers,
|
||||
proxy_snapshot.as_ref(),
|
||||
)
|
||||
.await
|
||||
.ok()?
|
||||
};
|
||||
let response = admin_provider_ops_execute_get_text_no_redirect(
|
||||
state,
|
||||
"provider-ops-acw:anyrouter",
|
||||
base_url.trim_end_matches('/'),
|
||||
&headers,
|
||||
proxy_snapshot.as_ref(),
|
||||
)
|
||||
.await
|
||||
.ok()?;
|
||||
let compiled = Regex::new(r"var\s+arg1\s*=\s*'([0-9a-fA-F]{40})'").ok()?;
|
||||
let captures = compiled.captures(&response.body)?;
|
||||
let arg1 = captures.get(1)?.as_str();
|
||||
@@ -63,222 +44,7 @@ pub(in super::super) async fn admin_provider_ops_resolve_proxy_snapshot(
|
||||
state: &AdminAppState<'_>,
|
||||
connector_config: Option<&Map<String, Value>>,
|
||||
) -> Option<ProxySnapshot> {
|
||||
let explicit_node_id = connector_config
|
||||
.and_then(|config| admin_provider_ops_string_field(config, "proxy_node_id"));
|
||||
if let Some(snapshot) =
|
||||
admin_provider_ops_resolve_proxy_node_snapshot(state, explicit_node_id.as_deref()).await
|
||||
{
|
||||
return Some(snapshot);
|
||||
}
|
||||
|
||||
if explicit_node_id.is_none() {
|
||||
let system_node_id = state
|
||||
.read_system_config_json_value("system_proxy_node_id")
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
.and_then(|value| value.as_str().map(str::trim).map(ToOwned::to_owned))
|
||||
.filter(|value| !value.is_empty());
|
||||
if let Some(snapshot) =
|
||||
admin_provider_ops_resolve_proxy_node_snapshot(state, system_node_id.as_deref()).await
|
||||
{
|
||||
return Some(snapshot);
|
||||
}
|
||||
}
|
||||
|
||||
connector_config
|
||||
.and_then(|config| config.get("proxy"))
|
||||
.and_then(admin_provider_ops_legacy_proxy_snapshot)
|
||||
}
|
||||
|
||||
async fn admin_provider_ops_resolve_proxy_node_snapshot(
|
||||
state: &AdminAppState<'_>,
|
||||
node_id: Option<&str>,
|
||||
) -> Option<ProxySnapshot> {
|
||||
let node_id = node_id.map(str::trim).filter(|value| !value.is_empty())?;
|
||||
let node = state.find_proxy_node(node_id).await.ok().flatten()?;
|
||||
if node.status.trim() != "online" {
|
||||
return None;
|
||||
}
|
||||
if node.tunnel_mode && node.tunnel_connected {
|
||||
let mut extra = Map::new();
|
||||
if let Ok(Some(owner)) = state.app().lookup_tunnel_attachment_owner(node_id).await {
|
||||
extra.insert(
|
||||
TUNNEL_BASE_URL_EXTRA_KEY.to_string(),
|
||||
Value::String(owner.relay_base_url),
|
||||
);
|
||||
extra.insert(
|
||||
TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY.to_string(),
|
||||
Value::String(owner.gateway_instance_id),
|
||||
);
|
||||
extra.insert(
|
||||
TUNNEL_OWNER_OBSERVED_AT_EXTRA_KEY.to_string(),
|
||||
json!(owner.observed_at_unix_secs),
|
||||
);
|
||||
}
|
||||
return Some(ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
mode: Some("tunnel".to_string()),
|
||||
node_id: Some(node_id.to_string()),
|
||||
label: Some(node.name),
|
||||
url: None,
|
||||
extra: if extra.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(Value::Object(extra))
|
||||
},
|
||||
});
|
||||
}
|
||||
if !node.is_manual {
|
||||
return None;
|
||||
}
|
||||
let proxy_url = node
|
||||
.proxy_url
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())?;
|
||||
Some(ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
mode: admin_provider_ops_proxy_mode(Some(proxy_url)),
|
||||
node_id: Some(node.id.clone()),
|
||||
label: Some(node.name.clone()),
|
||||
url: admin_provider_ops_proxy_url_with_node_auth(&node),
|
||||
extra: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn admin_provider_ops_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_ops_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_ops_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_ops_inject_proxy_auth(proxy_url, username, password)
|
||||
.or_else(|| Some(proxy_url.to_string())),
|
||||
extra: None,
|
||||
})
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_provider_ops_proxy_url_with_node_auth(node: &StoredProxyNode) -> Option<String> {
|
||||
let proxy_url = node
|
||||
.proxy_url
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())?;
|
||||
let username = node
|
||||
.proxy_username
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
let password = node
|
||||
.proxy_password
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
admin_provider_ops_inject_proxy_auth(proxy_url, username, password)
|
||||
.or_else(|| Some(proxy_url.to_string()))
|
||||
}
|
||||
|
||||
fn admin_provider_ops_inject_proxy_auth(
|
||||
proxy_url: &str,
|
||||
username: Option<&str>,
|
||||
password: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let username = username.filter(|value| !value.is_empty())?;
|
||||
let mut parsed = Url::parse(proxy_url).ok()?;
|
||||
parsed.set_username(username).ok()?;
|
||||
parsed.set_password(password).ok()?;
|
||||
Some(parsed.to_string())
|
||||
}
|
||||
|
||||
fn admin_provider_ops_proxy_mode(proxy_url: Option<&str>) -> Option<String> {
|
||||
proxy_url
|
||||
.and_then(|value| {
|
||||
Url::parse(value)
|
||||
.ok()
|
||||
.map(|parsed| parsed.scheme().to_string())
|
||||
})
|
||||
.or_else(|| {
|
||||
proxy_url.and_then(|value| {
|
||||
value
|
||||
.split_once("://")
|
||||
.map(|(scheme, _)| scheme.trim().to_ascii_lowercase())
|
||||
.filter(|scheme| !scheme.is_empty())
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn admin_provider_ops_proxy_uses_tunnel(proxy_snapshot: Option<&ProxySnapshot>) -> bool {
|
||||
proxy_snapshot.is_some_and(|proxy| {
|
||||
proxy.mode.as_deref().map(str::trim) == Some("tunnel")
|
||||
|| (proxy
|
||||
.url
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.unwrap_or_default()
|
||||
.is_empty()
|
||||
&& proxy
|
||||
.node_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty()))
|
||||
})
|
||||
}
|
||||
|
||||
fn admin_provider_ops_string_field(config: &Map<String, Value>, key: &str) -> Option<String> {
|
||||
config
|
||||
.get(key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
state
|
||||
.resolve_admin_connector_proxy_snapshot(connector_config)
|
||||
.await
|
||||
}
|
||||
|
||||
@@ -2,11 +2,10 @@ use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::{
|
||||
ExecutionPlan, ExecutionResult, ExecutionTimeouts, ProxySnapshot, RequestBody,
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
|
||||
};
|
||||
use aether_http::{apply_http_client_config, HttpClientConfig};
|
||||
use base64::{engine::general_purpose::STANDARD, Engine as _};
|
||||
use flate2::read::{DeflateDecoder, GzDecoder};
|
||||
use reqwest::redirect::Policy;
|
||||
use serde_json::{json, Value};
|
||||
use std::collections::BTreeMap;
|
||||
use std::io::Read;
|
||||
@@ -17,6 +16,11 @@ pub(super) struct AdminProviderOpsTextResponse {
|
||||
pub(super) body: String,
|
||||
}
|
||||
|
||||
pub(in super::super) enum AdminProviderOpsExecuteJsonError {
|
||||
InvalidJson(String),
|
||||
Transport(String),
|
||||
}
|
||||
|
||||
pub(super) async fn admin_provider_ops_execute_get_json(
|
||||
state: &AdminAppState<'_>,
|
||||
request_id: &str,
|
||||
@@ -24,37 +28,7 @@ pub(super) async fn admin_provider_ops_execute_get_json(
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
proxy_snapshot: Option<&ProxySnapshot>,
|
||||
) -> Result<(http::StatusCode, Value), String> {
|
||||
if proxy_snapshot.is_none() {
|
||||
let response = match state
|
||||
.http_client()
|
||||
.get(url)
|
||||
.headers(headers.clone())
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(response) => response,
|
||||
Err(err) if err.is_timeout() => return Err("timeout".to_string()),
|
||||
Err(err) => return Err(err.to_string()),
|
||||
};
|
||||
let status = response.status();
|
||||
let content_encoding = response
|
||||
.headers()
|
||||
.get(reqwest::header::CONTENT_ENCODING)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(ToOwned::to_owned);
|
||||
let bytes = response.bytes().await.map_err(|err| err.to_string())?;
|
||||
let decoded_bytes =
|
||||
admin_provider_ops_decode_response_bytes(bytes.as_ref(), content_encoding.as_deref())
|
||||
.unwrap_or_else(|| bytes.to_vec());
|
||||
let response_json = match serde_json::from_slice::<Value>(&decoded_bytes) {
|
||||
Ok(value) => value,
|
||||
Err(err) if status != http::StatusCode::OK => json!({}),
|
||||
Err(err) => return Err(format!("upstream response is not valid JSON: {err}")),
|
||||
};
|
||||
return Ok((status, response_json));
|
||||
}
|
||||
|
||||
let result = admin_provider_ops_execute_request(
|
||||
match admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
request_id,
|
||||
reqwest::Method::GET,
|
||||
@@ -63,11 +37,35 @@ pub(super) async fn admin_provider_ops_execute_get_json(
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await?;
|
||||
Ok((
|
||||
admin_provider_ops_execution_status_code(&result),
|
||||
admin_provider_ops_execution_json_body(&result),
|
||||
))
|
||||
.await
|
||||
{
|
||||
Ok(result) => Ok(result),
|
||||
Err(AdminProviderOpsExecuteJsonError::InvalidJson(message))
|
||||
| Err(AdminProviderOpsExecuteJsonError::Transport(message)) => Err(message),
|
||||
}
|
||||
}
|
||||
|
||||
pub(in super::super) async fn admin_provider_ops_execute_json_request(
|
||||
state: &AdminAppState<'_>,
|
||||
request_id: &str,
|
||||
method: reqwest::Method,
|
||||
url: &str,
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
json_body: Option<Value>,
|
||||
proxy_snapshot: Option<&ProxySnapshot>,
|
||||
) -> Result<(http::StatusCode, Value), AdminProviderOpsExecuteJsonError> {
|
||||
let result = admin_provider_ops_execute_request(
|
||||
state,
|
||||
request_id,
|
||||
method,
|
||||
url,
|
||||
headers,
|
||||
json_body,
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
.map_err(AdminProviderOpsExecuteJsonError::Transport)?;
|
||||
admin_provider_ops_execution_json_response(&result)
|
||||
}
|
||||
|
||||
pub(in super::super) async fn admin_provider_ops_execute_proxy_json_request(
|
||||
@@ -79,7 +77,7 @@ pub(in super::super) async fn admin_provider_ops_execute_proxy_json_request(
|
||||
json_body: Option<Value>,
|
||||
proxy_snapshot: &ProxySnapshot,
|
||||
) -> Result<(http::StatusCode, Value), String> {
|
||||
let result = admin_provider_ops_execute_request(
|
||||
match admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
request_id,
|
||||
method,
|
||||
@@ -88,11 +86,12 @@ pub(in super::super) async fn admin_provider_ops_execute_proxy_json_request(
|
||||
json_body,
|
||||
Some(proxy_snapshot),
|
||||
)
|
||||
.await?;
|
||||
Ok((
|
||||
admin_provider_ops_execution_status_code(&result),
|
||||
admin_provider_ops_execution_json_body(&result),
|
||||
))
|
||||
.await
|
||||
{
|
||||
Ok(result) => Ok(result),
|
||||
Err(AdminProviderOpsExecuteJsonError::InvalidJson(message))
|
||||
| Err(AdminProviderOpsExecuteJsonError::Transport(message)) => Err(message),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn admin_provider_ops_execute_get_text(
|
||||
@@ -129,50 +128,20 @@ pub(super) async fn admin_provider_ops_execute_get_text(
|
||||
}
|
||||
|
||||
pub(super) async fn admin_provider_ops_execute_get_text_no_redirect(
|
||||
state: &AdminAppState<'_>,
|
||||
request_id: &str,
|
||||
url: &str,
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
proxy_snapshot: Option<&ProxySnapshot>,
|
||||
) -> Result<AdminProviderOpsTextResponse, String> {
|
||||
let mut builder = apply_http_client_config(
|
||||
reqwest::Client::builder().redirect(Policy::none()),
|
||||
&HttpClientConfig {
|
||||
connect_timeout_ms: Some(10_000),
|
||||
request_timeout_ms: Some(ADMIN_PROVIDER_OPS_VERIFY_TIMEOUT_MS),
|
||||
use_rustls_tls: true,
|
||||
http2_adaptive_window: true,
|
||||
..HttpClientConfig::default()
|
||||
},
|
||||
);
|
||||
if let Some(proxy_url) = proxy_snapshot
|
||||
.and_then(|proxy| proxy.url.as_deref())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
let proxy = reqwest::Proxy::all(proxy_url).map_err(|err| format!("连接失败: {err}"))?;
|
||||
builder = builder.proxy(proxy);
|
||||
}
|
||||
let client = builder.build().map_err(|err| format!("验证失败: {err}"))?;
|
||||
let response = match client.get(url).headers(headers.clone()).send().await {
|
||||
Ok(response) => response,
|
||||
Err(err) if err.is_timeout() => return Err("连接超时".to_string()),
|
||||
Err(err) if err.is_connect() => return Err(format!("连接失败: {err}")),
|
||||
Err(err) => return Err(format!("验证失败: {err}")),
|
||||
};
|
||||
let content_encoding = response
|
||||
.headers()
|
||||
.get(reqwest::header::CONTENT_ENCODING)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(ToOwned::to_owned);
|
||||
let body = response
|
||||
.bytes()
|
||||
.await
|
||||
.map_err(|err| format!("验证失败: {err}"))
|
||||
.map(|bytes| {
|
||||
admin_provider_ops_decode_response_bytes(bytes.as_ref(), content_encoding.as_deref())
|
||||
.unwrap_or_else(|| bytes.to_vec())
|
||||
})
|
||||
.map(|bytes| String::from_utf8_lossy(&bytes).to_string())?;
|
||||
Ok(AdminProviderOpsTextResponse { body })
|
||||
admin_provider_ops_execute_get_text(
|
||||
state,
|
||||
request_id,
|
||||
url,
|
||||
&admin_provider_ops_headers_with_transport_controls(headers, Some(false), false),
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn admin_provider_ops_execute_request(
|
||||
@@ -240,23 +209,55 @@ fn admin_provider_ops_execution_headers(
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(in super::super) fn admin_provider_ops_headers_with_transport_controls(
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
follow_redirects: Option<bool>,
|
||||
http1_only: bool,
|
||||
) -> reqwest::header::HeaderMap {
|
||||
let mut headers = headers.clone();
|
||||
if let Some(follow_redirects) = follow_redirects {
|
||||
let value = if follow_redirects { "true" } else { "false" };
|
||||
headers.insert(
|
||||
reqwest::header::HeaderName::from_static(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER),
|
||||
reqwest::header::HeaderValue::from_static(value),
|
||||
);
|
||||
}
|
||||
if http1_only {
|
||||
headers.insert(
|
||||
reqwest::header::HeaderName::from_static(EXECUTION_REQUEST_HTTP1_ONLY_HEADER),
|
||||
reqwest::header::HeaderValue::from_static("true"),
|
||||
);
|
||||
}
|
||||
headers
|
||||
}
|
||||
|
||||
fn admin_provider_ops_execution_status_code(result: &ExecutionResult) -> http::StatusCode {
|
||||
http::StatusCode::from_u16(result.status_code).unwrap_or(http::StatusCode::BAD_GATEWAY)
|
||||
}
|
||||
|
||||
fn admin_provider_ops_execution_json_body(result: &ExecutionResult) -> Value {
|
||||
result
|
||||
fn admin_provider_ops_execution_json_response(
|
||||
result: &ExecutionResult,
|
||||
) -> Result<(http::StatusCode, Value), AdminProviderOpsExecuteJsonError> {
|
||||
let status = admin_provider_ops_execution_status_code(result);
|
||||
if let Some(json_body) = result.body.as_ref().and_then(|body| body.json_body.clone()) {
|
||||
return Ok((status, json_body));
|
||||
}
|
||||
|
||||
let Some(bytes) = result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.json_body.clone())
|
||||
.or_else(|| {
|
||||
result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| admin_provider_ops_execution_body_bytes(&result.headers, body))
|
||||
.and_then(|bytes| serde_json::from_slice::<Value>(&bytes).ok())
|
||||
})
|
||||
.unwrap_or_else(|| json!({}))
|
||||
.and_then(|body| admin_provider_ops_execution_body_bytes(&result.headers, body))
|
||||
else {
|
||||
return Ok((status, json!({})));
|
||||
};
|
||||
|
||||
match serde_json::from_slice::<Value>(&bytes) {
|
||||
Ok(value) => Ok((status, value)),
|
||||
Err(_) if status != http::StatusCode::OK => Ok((status, json!({}))),
|
||||
Err(err) => Err(AdminProviderOpsExecuteJsonError::InvalidJson(format!(
|
||||
"upstream response is not valid JSON: {err}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_provider_ops_execution_body_bytes(
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use super::request::{
|
||||
admin_provider_ops_execute_proxy_json_request,
|
||||
admin_provider_ops_verify_execution_error_message,
|
||||
admin_provider_ops_execute_json_request, admin_provider_ops_headers_with_transport_controls,
|
||||
admin_provider_ops_verify_execution_error_message, AdminProviderOpsExecuteJsonError,
|
||||
};
|
||||
use crate::handlers::admin::provider::ops::providers::config::persist_admin_provider_ops_runtime_credentials;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
@@ -10,7 +10,6 @@ use aether_admin::provider::ops::{
|
||||
};
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
|
||||
use aether_http::{apply_http_client_config, HttpClientConfig};
|
||||
use serde_json::{json, Map, Value};
|
||||
use tracing::warn;
|
||||
|
||||
@@ -64,51 +63,26 @@ pub(super) async fn admin_provider_ops_local_sub2api_verify_response(
|
||||
reqwest::header::HeaderValue::from_static("*/*"),
|
||||
),
|
||||
]);
|
||||
let (status, response_json) = if let Some(proxy_snapshot) = proxy_snapshot {
|
||||
match admin_provider_ops_execute_proxy_json_request(
|
||||
state,
|
||||
"provider-ops-verify:sub2api",
|
||||
reqwest::Method::GET,
|
||||
&verify_url,
|
||||
&auth_headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(error) => {
|
||||
return admin_provider_ops_verify_failure(
|
||||
admin_provider_ops_verify_execution_error_message(&error),
|
||||
);
|
||||
}
|
||||
let auth_headers =
|
||||
admin_provider_ops_headers_with_transport_controls(&auth_headers, None, true);
|
||||
let (status, response_json) = match admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
"provider-ops-verify:sub2api",
|
||||
reqwest::Method::GET,
|
||||
&verify_url,
|
||||
&auth_headers,
|
||||
None,
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(AdminProviderOpsExecuteJsonError::InvalidJson(message))
|
||||
| Err(AdminProviderOpsExecuteJsonError::Transport(message)) => {
|
||||
return admin_provider_ops_verify_failure(
|
||||
admin_provider_ops_verify_execution_error_message(&message),
|
||||
);
|
||||
}
|
||||
} else {
|
||||
let http_client = match admin_provider_ops_sub2api_http_client() {
|
||||
Ok(client) => client,
|
||||
Err(err) => {
|
||||
return admin_provider_ops_verify_failure(format!("验证失败: {err}"));
|
||||
}
|
||||
};
|
||||
let response = match http_client
|
||||
.get(&verify_url)
|
||||
.headers(auth_headers)
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(response) => response,
|
||||
Err(err) if err.is_timeout() => return admin_provider_ops_verify_failure("连接超时"),
|
||||
Err(err) if err.is_connect() => {
|
||||
return admin_provider_ops_verify_failure(format!("连接失败: {err}"));
|
||||
}
|
||||
Err(err) => return admin_provider_ops_verify_failure(format!("验证失败: {err}")),
|
||||
};
|
||||
let status = response.status();
|
||||
let response_json = match response.bytes().await {
|
||||
Ok(bytes) => serde_json::from_slice::<Value>(&bytes).unwrap_or_else(|_| json!({})),
|
||||
Err(_) => json!({}),
|
||||
};
|
||||
(status, response_json)
|
||||
};
|
||||
|
||||
parse_verify_payload(
|
||||
@@ -119,20 +93,6 @@ pub(super) async fn admin_provider_ops_local_sub2api_verify_response(
|
||||
)
|
||||
}
|
||||
|
||||
fn admin_provider_ops_sub2api_http_client() -> Result<reqwest::Client, reqwest::Error> {
|
||||
let builder = apply_http_client_config(
|
||||
reqwest::Client::builder().http1_only(),
|
||||
&HttpClientConfig {
|
||||
connect_timeout_ms: Some(10_000),
|
||||
request_timeout_ms: Some(30_000),
|
||||
use_rustls_tls: true,
|
||||
user_agent: Some(ADMIN_PROVIDER_OPS_USER_AGENT.to_string()),
|
||||
..HttpClientConfig::default()
|
||||
},
|
||||
);
|
||||
builder.build()
|
||||
}
|
||||
|
||||
// 对齐 Python httpx.AsyncClient(base_url=...) 的行为:
|
||||
// 以 "/" 开头的端点始终相对站点根路径解析,而不是简单字符串拼接。
|
||||
pub(in super::super) fn admin_provider_ops_sub2api_request_url(
|
||||
@@ -263,39 +223,24 @@ async fn admin_provider_ops_sub2api_token_request(
|
||||
reqwest::header::HeaderValue::from_static("*/*"),
|
||||
),
|
||||
]);
|
||||
let (status, response_json) = if let Some(proxy_snapshot) = proxy_snapshot {
|
||||
admin_provider_ops_execute_proxy_json_request(
|
||||
state,
|
||||
&format!("provider-ops-sub2api:{path}"),
|
||||
reqwest::Method::POST,
|
||||
&url,
|
||||
&default_headers,
|
||||
Some(body),
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
.map_err(|error| admin_provider_ops_verify_execution_error_message(&error))?
|
||||
} else {
|
||||
let client =
|
||||
admin_provider_ops_sub2api_http_client().map_err(|err| format!("验证失败: {err}"))?;
|
||||
let response = match client
|
||||
.post(url)
|
||||
.headers(default_headers)
|
||||
.json(&body)
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(response) => response,
|
||||
Err(err) if err.is_timeout() => return Err("连接超时".to_string()),
|
||||
Err(err) if err.is_connect() => return Err(format!("连接失败: {err}")),
|
||||
Err(err) => return Err(format!("验证失败: {err}")),
|
||||
};
|
||||
let status = response.status();
|
||||
let response_json = match response.bytes().await {
|
||||
Ok(bytes) => serde_json::from_slice::<Value>(&bytes).unwrap_or_else(|_| json!({})),
|
||||
Err(_) => json!({}),
|
||||
};
|
||||
(status, response_json)
|
||||
let default_headers =
|
||||
admin_provider_ops_headers_with_transport_controls(&default_headers, None, true);
|
||||
let (status, response_json) = match admin_provider_ops_execute_json_request(
|
||||
state,
|
||||
&format!("provider-ops-sub2api:{path}"),
|
||||
reqwest::Method::POST,
|
||||
&url,
|
||||
&default_headers,
|
||||
Some(body),
|
||||
proxy_snapshot,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(AdminProviderOpsExecuteJsonError::InvalidJson(message))
|
||||
| Err(AdminProviderOpsExecuteJsonError::Transport(message)) => {
|
||||
return Err(admin_provider_ops_verify_execution_error_message(&message));
|
||||
}
|
||||
};
|
||||
let payload = response_json.as_object().cloned().unwrap_or_default();
|
||||
if status != http::StatusCode::OK
|
||||
|
||||
@@ -88,6 +88,10 @@ impl<'a> AdminAppState<'a> {
|
||||
self.app.has_proxy_node_reader()
|
||||
}
|
||||
|
||||
pub(crate) fn has_proxy_node_writer(&self) -> bool {
|
||||
self.app.has_proxy_node_writer()
|
||||
}
|
||||
|
||||
pub(crate) fn has_auth_api_key_writer(&self) -> bool {
|
||||
self.app.data.has_auth_api_key_writer()
|
||||
}
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
use super::*;
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use aether_contracts::{
|
||||
ExecutionPlan, ExecutionResult, ExecutionTimeouts, RequestBody,
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
|
||||
};
|
||||
use aether_data::repository::provider_oauth::{
|
||||
build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_storage_key,
|
||||
provider_oauth_device_session_storage_key, provider_oauth_state_storage_key,
|
||||
@@ -7,11 +11,22 @@ use aether_data::repository::provider_oauth::{
|
||||
PROVIDER_OAUTH_BATCH_TASK_TTL_SECS, PROVIDER_OAUTH_STATE_TTL_SECS,
|
||||
};
|
||||
use axum::http;
|
||||
use base64::{engine::general_purpose::STANDARD, Engine as _};
|
||||
use flate2::read::{DeflateDecoder, GzDecoder};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
use std::io::Read;
|
||||
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;
|
||||
|
||||
pub(crate) struct AdminProviderOAuthHttpResponse {
|
||||
pub(crate) status: http::StatusCode,
|
||||
pub(crate) body_text: String,
|
||||
pub(crate) json_body: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
impl<'a> AdminAppState<'a> {
|
||||
pub(crate) async fn update_provider_catalog_key_oauth_credentials(
|
||||
@@ -117,6 +132,7 @@ impl<'a> AdminAppState<'a> {
|
||||
code: &str,
|
||||
state_nonce: &str,
|
||||
pkce_verifier: Option<&str>,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<serde_json::Value, Response<Body>> {
|
||||
crate::handlers::admin::provider::oauth::state::exchange_admin_provider_oauth_code(
|
||||
self,
|
||||
@@ -124,6 +140,7 @@ impl<'a> AdminAppState<'a> {
|
||||
code,
|
||||
state_nonce,
|
||||
pkce_verifier,
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -132,11 +149,13 @@ impl<'a> AdminAppState<'a> {
|
||||
&self,
|
||||
template: AdminProviderOAuthTemplate,
|
||||
refresh_token: &str,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> 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,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -296,6 +315,7 @@ impl<'a> AdminAppState<'a> {
|
||||
&self,
|
||||
region: &str,
|
||||
start_url: &str,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<serde_json::Value, Response<Body>> {
|
||||
let payload = post_kiro_device_oidc_json(
|
||||
self,
|
||||
@@ -317,6 +337,7 @@ impl<'a> AdminAppState<'a> {
|
||||
],
|
||||
"issuerUrl": start_url,
|
||||
}),
|
||||
proxy_node_id,
|
||||
)
|
||||
.await?;
|
||||
if payload
|
||||
@@ -343,6 +364,7 @@ impl<'a> AdminAppState<'a> {
|
||||
client_id: &str,
|
||||
client_secret: &str,
|
||||
start_url: &str,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<serde_json::Value, Response<Body>> {
|
||||
let payload = post_kiro_device_oidc_json(
|
||||
self,
|
||||
@@ -353,6 +375,7 @@ impl<'a> AdminAppState<'a> {
|
||||
"clientSecret": client_secret,
|
||||
"startUrl": start_url,
|
||||
}),
|
||||
proxy_node_id,
|
||||
)
|
||||
.await?;
|
||||
if payload
|
||||
@@ -379,6 +402,7 @@ impl<'a> AdminAppState<'a> {
|
||||
client_id: &str,
|
||||
client_secret: &str,
|
||||
device_code: &str,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<serde_json::Value, Response<Body>> {
|
||||
post_kiro_device_oidc_json(
|
||||
self,
|
||||
@@ -390,6 +414,7 @@ impl<'a> AdminAppState<'a> {
|
||||
"grantType": "urn:ietf:params:oauth:grant-type:device_code",
|
||||
"deviceCode": device_code,
|
||||
}),
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -479,22 +504,43 @@ async fn post_kiro_device_oidc_json(
|
||||
endpoint_key: &str,
|
||||
default_url: String,
|
||||
body: serde_json::Value,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<serde_json::Value, Response<Body>> {
|
||||
let url = state.provider_oauth_token_url(endpoint_key, &default_url);
|
||||
let host = Url::parse(&url)
|
||||
.ok()
|
||||
.and_then(|value| value.host_str().map(ToOwned::to_owned))
|
||||
.unwrap_or_default();
|
||||
let headers = reqwest::header::HeaderMap::from_iter([
|
||||
(
|
||||
reqwest::header::CONTENT_TYPE,
|
||||
reqwest::header::HeaderValue::from_static("application/json"),
|
||||
),
|
||||
(
|
||||
reqwest::header::ACCEPT,
|
||||
reqwest::header::HeaderValue::from_static("*/*"),
|
||||
),
|
||||
(
|
||||
reqwest::header::USER_AGENT,
|
||||
reqwest::header::HeaderValue::from_static("node"),
|
||||
),
|
||||
(
|
||||
reqwest::header::HeaderName::from_static("x-amz-user-agent"),
|
||||
reqwest::header::HeaderValue::from_static(KIRO_IDC_AMZ_USER_AGENT),
|
||||
),
|
||||
]);
|
||||
let headers = maybe_insert_host_header(headers, host.as_str());
|
||||
let response = state
|
||||
.http_client()
|
||||
.post(url)
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "*/*")
|
||||
.header("User-Agent", "node")
|
||||
.header("x-amz-user-agent", KIRO_IDC_AMZ_USER_AGENT)
|
||||
.header("Host", host)
|
||||
.json(&body)
|
||||
.send()
|
||||
.execute_admin_provider_oauth_http_request(
|
||||
endpoint_key,
|
||||
reqwest::Method::POST,
|
||||
&url,
|
||||
&headers,
|
||||
Some("application/json"),
|
||||
Some(body),
|
||||
None,
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
.map_err(|_| {
|
||||
build_internal_control_error_response(
|
||||
@@ -502,13 +548,8 @@ async fn post_kiro_device_oidc_json(
|
||||
"发起设备授权失败: unknown",
|
||||
)
|
||||
})?;
|
||||
let status = response.status();
|
||||
let body_text = response.text().await.map_err(|_| {
|
||||
build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"发起设备授权失败: unknown",
|
||||
)
|
||||
})?;
|
||||
let status = response.status;
|
||||
let body_text = response.body_text;
|
||||
Ok(
|
||||
serde_json::from_str::<serde_json::Value>(&body_text).unwrap_or_else(|_| {
|
||||
json!({
|
||||
@@ -518,3 +559,179 @@ async fn post_kiro_device_oidc_json(
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
impl<'a> AdminAppState<'a> {
|
||||
pub(crate) async fn execute_admin_provider_oauth_http_request(
|
||||
&self,
|
||||
request_id: &str,
|
||||
method: reqwest::Method,
|
||||
url: &str,
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
content_type: Option<&str>,
|
||||
json_body: Option<serde_json::Value>,
|
||||
body_bytes: Option<Vec<u8>>,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<AdminProviderOAuthHttpResponse, String> {
|
||||
let body = if let Some(json_body) = json_body {
|
||||
RequestBody::from_json(json_body)
|
||||
} else {
|
||||
RequestBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: body_bytes.map(|bytes| STANDARD.encode(bytes)),
|
||||
body_ref: None,
|
||||
}
|
||||
};
|
||||
let plan = ExecutionPlan {
|
||||
request_id: request_id.to_string(),
|
||||
candidate_id: None,
|
||||
provider_name: Some("provider_oauth".to_string()),
|
||||
provider_id: String::new(),
|
||||
endpoint_id: String::new(),
|
||||
key_id: String::new(),
|
||||
method: method.as_str().to_string(),
|
||||
url: url.to_string(),
|
||||
headers: admin_provider_oauth_execution_headers(headers),
|
||||
content_type: content_type
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
content_encoding: None,
|
||||
body,
|
||||
stream: false,
|
||||
client_api_format: "provider_oauth:exchange".to_string(),
|
||||
provider_api_format: "provider_oauth:exchange".to_string(),
|
||||
model_name: Some("oauth-exchange".to_string()),
|
||||
proxy: self.resolve_admin_proxy_node_snapshot(proxy_node_id).await,
|
||||
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),
|
||||
..ExecutionTimeouts::default()
|
||||
}),
|
||||
};
|
||||
let result = self
|
||||
.execute_execution_runtime_sync_plan(None, &plan)
|
||||
.await
|
||||
.map_err(admin_provider_oauth_gateway_error_message)?;
|
||||
Ok(AdminProviderOAuthHttpResponse {
|
||||
status: http::StatusCode::from_u16(result.status_code)
|
||||
.unwrap_or(http::StatusCode::BAD_GATEWAY),
|
||||
body_text: admin_provider_oauth_execution_body_text(&result),
|
||||
json_body: admin_provider_oauth_execution_json_body(&result),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn maybe_insert_host_header(
|
||||
mut headers: reqwest::header::HeaderMap,
|
||||
host: &str,
|
||||
) -> reqwest::header::HeaderMap {
|
||||
let host = host.trim();
|
||||
if host.is_empty() {
|
||||
return headers;
|
||||
}
|
||||
if let Ok(value) = reqwest::header::HeaderValue::from_str(host) {
|
||||
headers.insert(reqwest::header::HOST, value);
|
||||
}
|
||||
headers
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_execution_headers(
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut headers: BTreeMap<String, String> = headers
|
||||
.iter()
|
||||
.filter_map(|(name, value)| {
|
||||
value
|
||||
.to_str()
|
||||
.ok()
|
||||
.map(|text| (name.as_str().to_string(), text.to_string()))
|
||||
})
|
||||
.collect();
|
||||
headers.insert(
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(),
|
||||
"true".to_string(),
|
||||
);
|
||||
headers
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_execution_json_body(result: &ExecutionResult) -> Option<serde_json::Value> {
|
||||
result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.json_body.clone())
|
||||
.or_else(|| {
|
||||
result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| admin_provider_oauth_execution_body_bytes(&result.headers, body))
|
||||
.and_then(|bytes| serde_json::from_slice::<serde_json::Value>(&bytes).ok())
|
||||
})
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_execution_body_text(result: &ExecutionResult) -> String {
|
||||
result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| admin_provider_oauth_execution_body_bytes(&result.headers, body))
|
||||
.map(|bytes| String::from_utf8_lossy(&bytes).to_string())
|
||||
.or_else(|| {
|
||||
result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.json_body.as_ref())
|
||||
.and_then(|value| serde_json::to_string(value).ok())
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_execution_body_bytes(
|
||||
headers: &BTreeMap<String, String>,
|
||||
body: &aether_contracts::ResponseBody,
|
||||
) -> Option<Vec<u8>> {
|
||||
let bytes = body
|
||||
.body_bytes_b64
|
||||
.as_deref()
|
||||
.and_then(|value| STANDARD.decode(value).ok())?;
|
||||
admin_provider_oauth_decode_response_bytes(
|
||||
&bytes,
|
||||
headers.get("content-encoding").map(String::as_str),
|
||||
)
|
||||
.or(Some(bytes))
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_decode_response_bytes(
|
||||
bytes: &[u8],
|
||||
content_encoding: Option<&str>,
|
||||
) -> Option<Vec<u8>> {
|
||||
let encoding = content_encoding
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| value.to_ascii_lowercase());
|
||||
match encoding.as_deref() {
|
||||
Some("gzip") => {
|
||||
let mut decoder = GzDecoder::new(bytes);
|
||||
let mut out = Vec::new();
|
||||
decoder.read_to_end(&mut out).ok()?;
|
||||
Some(out)
|
||||
}
|
||||
Some("deflate") => {
|
||||
let mut decoder = DeflateDecoder::new(bytes);
|
||||
let mut out = Vec::new();
|
||||
decoder.read_to_end(&mut out).ok()?;
|
||||
Some(out)
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_gateway_error_message(error: GatewayError) -> String {
|
||||
match error {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,13 @@
|
||||
use super::*;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_data::repository::proxy_nodes::{proxy_node_accepts_new_tunnels, StoredProxyNode};
|
||||
use aether_provider_transport::TransportTunnelAffinityLookup;
|
||||
use serde_json::{json, Map, Value};
|
||||
use url::Url;
|
||||
|
||||
const TUNNEL_BASE_URL_EXTRA_KEY: &str = "tunnel_base_url";
|
||||
const TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY: &str = "tunnel_owner_instance_id";
|
||||
const TUNNEL_OWNER_OBSERVED_AT_EXTRA_KEY: &str = "tunnel_owner_observed_at_unix_secs";
|
||||
|
||||
impl<'a> AdminAppState<'a> {
|
||||
pub(crate) async fn read_provider_transport_snapshot(
|
||||
@@ -100,6 +109,100 @@ impl<'a> AdminAppState<'a> {
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_admin_connector_proxy_snapshot(
|
||||
&self,
|
||||
connector_config: Option<&Map<String, Value>>,
|
||||
) -> Option<ProxySnapshot> {
|
||||
let explicit_node_id = connector_config
|
||||
.and_then(|config| admin_provider_transport_string_field(config, "proxy_node_id"));
|
||||
if let Some(snapshot) = self
|
||||
.resolve_admin_proxy_node_snapshot(explicit_node_id.as_deref())
|
||||
.await
|
||||
{
|
||||
return Some(snapshot);
|
||||
}
|
||||
|
||||
if explicit_node_id.is_none() {
|
||||
let system_node_id = self
|
||||
.read_system_config_json_value("system_proxy_node_id")
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
.and_then(|value| value.as_str().map(str::trim).map(ToOwned::to_owned))
|
||||
.filter(|value| !value.is_empty());
|
||||
if let Some(snapshot) = self
|
||||
.resolve_admin_proxy_node_snapshot(system_node_id.as_deref())
|
||||
.await
|
||||
{
|
||||
return Some(snapshot);
|
||||
}
|
||||
}
|
||||
|
||||
connector_config
|
||||
.and_then(|config| config.get("proxy"))
|
||||
.and_then(admin_provider_transport_legacy_proxy_snapshot)
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_admin_proxy_node_snapshot(
|
||||
&self,
|
||||
node_id: Option<&str>,
|
||||
) -> Option<ProxySnapshot> {
|
||||
let node_id = node_id.map(str::trim).filter(|value| !value.is_empty())?;
|
||||
let node = self.find_proxy_node(node_id).await.ok().flatten()?;
|
||||
if node.status.trim() != "online" {
|
||||
return None;
|
||||
}
|
||||
if !proxy_node_accepts_new_tunnels(&node) {
|
||||
return None;
|
||||
}
|
||||
if node.tunnel_mode && node.tunnel_connected {
|
||||
let mut extra = Map::new();
|
||||
if let Ok(Some(owner)) = self.app().lookup_tunnel_attachment_owner(node_id).await {
|
||||
extra.insert(
|
||||
TUNNEL_BASE_URL_EXTRA_KEY.to_string(),
|
||||
Value::String(owner.relay_base_url),
|
||||
);
|
||||
extra.insert(
|
||||
TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY.to_string(),
|
||||
Value::String(owner.gateway_instance_id),
|
||||
);
|
||||
extra.insert(
|
||||
TUNNEL_OWNER_OBSERVED_AT_EXTRA_KEY.to_string(),
|
||||
json!(owner.observed_at_unix_secs),
|
||||
);
|
||||
}
|
||||
return Some(ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
mode: Some("tunnel".to_string()),
|
||||
node_id: Some(node_id.to_string()),
|
||||
label: Some(node.name),
|
||||
url: None,
|
||||
extra: if extra.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(Value::Object(extra))
|
||||
},
|
||||
});
|
||||
}
|
||||
if !node.is_manual {
|
||||
return None;
|
||||
}
|
||||
let proxy_url = node
|
||||
.proxy_url
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())?;
|
||||
Some(ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
mode: admin_provider_transport_proxy_mode(Some(proxy_url)),
|
||||
node_id: Some(node.id.clone()),
|
||||
label: Some(node.name.clone()),
|
||||
url: admin_provider_transport_proxy_url_with_node_auth(&node)
|
||||
.or_else(|| Some(proxy_url.to_string())),
|
||||
extra: None,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn supports_local_gemini_transport_with_network(
|
||||
&self,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
@@ -222,3 +325,121 @@ impl<'a> AdminAppState<'a> {
|
||||
crate::provider_transport::url::build_openai_chat_url(upstream_base_url, query)
|
||||
}
|
||||
}
|
||||
|
||||
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_url_with_node_auth(node: &StoredProxyNode) -> Option<String> {
|
||||
let proxy_url = node
|
||||
.proxy_url
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())?;
|
||||
let username = node
|
||||
.proxy_username
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
let password = node
|
||||
.proxy_password
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
admin_provider_transport_inject_proxy_auth(proxy_url, username, password)
|
||||
}
|
||||
|
||||
fn admin_provider_transport_inject_proxy_auth(
|
||||
proxy_url: &str,
|
||||
username: Option<&str>,
|
||||
password: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let username = username.filter(|value| !value.is_empty())?;
|
||||
let mut parsed = Url::parse(proxy_url).ok()?;
|
||||
parsed.set_username(username).ok()?;
|
||||
parsed.set_password(password).ok()?;
|
||||
Some(parsed.to_string())
|
||||
}
|
||||
|
||||
fn admin_provider_transport_proxy_mode(proxy_url: Option<&str>) -> Option<String> {
|
||||
proxy_url
|
||||
.and_then(|value| {
|
||||
Url::parse(value)
|
||||
.ok()
|
||||
.map(|parsed| parsed.scheme().to_string())
|
||||
})
|
||||
.or_else(|| {
|
||||
proxy_url.and_then(|value| {
|
||||
value
|
||||
.split_once("://")
|
||||
.map(|(scheme, _)| scheme.trim().to_ascii_lowercase())
|
||||
.filter(|scheme| !scheme.is_empty())
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn admin_provider_transport_string_field(config: &Map<String, Value>, key: &str) -> Option<String> {
|
||||
config
|
||||
.get(key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::shared::unix_secs_to_rfc3339;
|
||||
use crate::maintenance::{inspect_proxy_upgrade_rollout, ProxyUpgradeRolloutStatus};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::system::{
|
||||
build_admin_proxy_node_event_payload, build_admin_proxy_node_events_payload_response,
|
||||
@@ -9,6 +11,20 @@ use aether_admin::system::{
|
||||
use axum::{body::Body, response::Response};
|
||||
|
||||
impl<'a> AdminAppState<'a> {
|
||||
pub(crate) async fn register_proxy_node(
|
||||
&self,
|
||||
mutation: &aether_data::repository::proxy_nodes::ProxyNodeRegistrationMutation,
|
||||
) -> Result<Option<aether_data::repository::proxy_nodes::StoredProxyNode>, GatewayError> {
|
||||
self.app.register_proxy_node(mutation).await
|
||||
}
|
||||
|
||||
pub(crate) async fn apply_proxy_node_heartbeat(
|
||||
&self,
|
||||
mutation: &aether_data::repository::proxy_nodes::ProxyNodeHeartbeatMutation,
|
||||
) -> Result<Option<aether_data::repository::proxy_nodes::StoredProxyNode>, GatewayError> {
|
||||
self.app.apply_proxy_node_heartbeat(mutation).await
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_proxy_nodes_list_response(
|
||||
&self,
|
||||
skip: usize,
|
||||
@@ -42,8 +58,12 @@ impl<'a> AdminAppState<'a> {
|
||||
.take(limit)
|
||||
.map(|node| build_admin_proxy_node_payload(&node))
|
||||
.collect::<Vec<_>>();
|
||||
let rollout = inspect_proxy_upgrade_rollout(self.app().data.as_ref())
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
.map(build_admin_proxy_upgrade_rollout_payload);
|
||||
Ok(build_admin_proxy_nodes_list_payload_response(
|
||||
items, total, skip, limit,
|
||||
items, total, skip, limit, rollout,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -66,4 +86,55 @@ impl<'a> AdminAppState<'a> {
|
||||
.collect::<Vec<_>>();
|
||||
Ok(build_admin_proxy_node_events_payload_response(items))
|
||||
}
|
||||
|
||||
pub(crate) async fn unregister_proxy_node(
|
||||
&self,
|
||||
node_id: &str,
|
||||
) -> Result<Option<aether_data::repository::proxy_nodes::StoredProxyNode>, GatewayError> {
|
||||
self.app.unregister_proxy_node(node_id).await
|
||||
}
|
||||
|
||||
pub(crate) async fn update_proxy_node_remote_config(
|
||||
&self,
|
||||
mutation: &aether_data::repository::proxy_nodes::ProxyNodeRemoteConfigMutation,
|
||||
) -> Result<Option<aether_data::repository::proxy_nodes::StoredProxyNode>, GatewayError> {
|
||||
self.app.update_proxy_node_remote_config(mutation).await
|
||||
}
|
||||
}
|
||||
|
||||
fn build_admin_proxy_upgrade_rollout_payload(
|
||||
rollout: ProxyUpgradeRolloutStatus,
|
||||
) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"version": rollout.version,
|
||||
"batch_size": rollout.batch_size,
|
||||
"cooldown_secs": rollout.cooldown_secs,
|
||||
"started_at": unix_secs_to_rfc3339(rollout.started_at_unix_secs),
|
||||
"last_dispatched_at": rollout
|
||||
.last_dispatched_at_unix_secs
|
||||
.and_then(unix_secs_to_rfc3339),
|
||||
"updated_at": unix_secs_to_rfc3339(rollout.updated_at_unix_secs),
|
||||
"probe": rollout.probe.map(|probe| serde_json::json!({
|
||||
"url": probe.url,
|
||||
"timeout_secs": probe.timeout_secs,
|
||||
})),
|
||||
"blocked": rollout.blocked,
|
||||
"online_eligible_total": rollout.online_eligible_total,
|
||||
"completed_node_ids": rollout.completed_node_ids,
|
||||
"pending_node_ids": rollout.pending_node_ids,
|
||||
"conflict_node_ids": rollout.conflict_node_ids,
|
||||
"skipped_node_ids": rollout.skipped_node_ids,
|
||||
"tracked_nodes": rollout.tracked_nodes.into_iter().map(|tracked| serde_json::json!({
|
||||
"node_id": tracked.node_id,
|
||||
"state": tracked.state,
|
||||
"dispatched_at": unix_secs_to_rfc3339(tracked.dispatched_at_unix_secs),
|
||||
"version_confirmed_at": tracked
|
||||
.version_confirmed_at_unix_secs
|
||||
.and_then(unix_secs_to_rfc3339),
|
||||
"traffic_confirmed_at": tracked
|
||||
.traffic_confirmed_at_unix_secs
|
||||
.and_then(unix_secs_to_rfc3339),
|
||||
"cooldown_remaining_secs": tracked.cooldown_remaining_secs,
|
||||
})).collect::<Vec<_>>(),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,15 +1,103 @@
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::query_param_value;
|
||||
use crate::maintenance::{
|
||||
cancel_proxy_upgrade_rollout, clear_proxy_upgrade_rollout_conflicts,
|
||||
restore_proxy_upgrade_rollout_skipped_nodes, retry_proxy_upgrade_rollout_node,
|
||||
skip_proxy_upgrade_rollout_node, start_proxy_upgrade_rollout, ProxyUpgradeRolloutProbeConfig,
|
||||
};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::system::{
|
||||
admin_proxy_node_event_node_id_from_path, build_admin_proxy_nodes_data_unavailable_response,
|
||||
build_admin_proxy_nodes_not_found_response,
|
||||
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 axum::{body::Body, http, response::Response};
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde::de::DeserializeOwned;
|
||||
use serde::Deserialize;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct ProxyNodeRegisterRequest {
|
||||
name: String,
|
||||
ip: String,
|
||||
#[serde(default)]
|
||||
port: Option<u16>,
|
||||
#[serde(default)]
|
||||
region: Option<String>,
|
||||
#[serde(default)]
|
||||
heartbeat_interval: Option<i32>,
|
||||
#[serde(default)]
|
||||
active_connections: Option<i32>,
|
||||
#[serde(default)]
|
||||
total_requests: Option<i64>,
|
||||
#[serde(default)]
|
||||
avg_latency_ms: Option<f64>,
|
||||
#[serde(default)]
|
||||
hardware_info: Option<Value>,
|
||||
#[serde(default)]
|
||||
estimated_max_concurrency: Option<i32>,
|
||||
#[serde(default)]
|
||||
proxy_metadata: Option<Value>,
|
||||
#[serde(default)]
|
||||
proxy_version: Option<String>,
|
||||
#[serde(default)]
|
||||
tunnel_mode: Option<bool>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct ProxyNodeHeartbeatRequest {
|
||||
node_id: String,
|
||||
#[serde(default)]
|
||||
heartbeat_interval: Option<i32>,
|
||||
#[serde(default)]
|
||||
active_connections: Option<i32>,
|
||||
#[serde(default)]
|
||||
total_requests: Option<i64>,
|
||||
#[serde(default)]
|
||||
avg_latency_ms: Option<f64>,
|
||||
#[serde(default)]
|
||||
failed_requests: Option<i64>,
|
||||
#[serde(default)]
|
||||
dns_failures: Option<i64>,
|
||||
#[serde(default)]
|
||||
stream_errors: Option<i64>,
|
||||
#[serde(default)]
|
||||
proxy_metadata: Option<Value>,
|
||||
#[serde(default)]
|
||||
proxy_version: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct ProxyNodeUnregisterRequest {
|
||||
node_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct ProxyNodeBatchUpgradeRequest {
|
||||
version: String,
|
||||
#[serde(default)]
|
||||
batch_size: Option<usize>,
|
||||
#[serde(default)]
|
||||
cooldown_secs: Option<u64>,
|
||||
#[serde(default)]
|
||||
probe_url: Option<String>,
|
||||
#[serde(default)]
|
||||
probe_timeout_secs: Option<u64>,
|
||||
}
|
||||
|
||||
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;
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_proxy_nodes_response(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
@@ -61,5 +149,806 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response(
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("register_node")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
if !state.has_proxy_node_writer() {
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
}
|
||||
let input = match parse_json_body::<ProxyNodeRegisterRequest>(request_body) {
|
||||
Ok(input) => input,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let mutation = match validate_register_request(input, request_context) {
|
||||
Ok(mutation) => mutation,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let Some(node) = state.register_proxy_node(&mutation).await? else {
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
};
|
||||
return Ok(Some(
|
||||
Json(json!({
|
||||
"node_id": node.id,
|
||||
"node": build_admin_proxy_node_payload(&node),
|
||||
}))
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("heartbeat_node")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
if !state.has_proxy_node_writer() {
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
}
|
||||
let input = match parse_json_body::<ProxyNodeHeartbeatRequest>(request_body) {
|
||||
Ok(input) => input,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let mutation = match validate_heartbeat_request(input) {
|
||||
Ok(mutation) => mutation,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let Some(existing) = state.find_proxy_node(&mutation.node_id).await? else {
|
||||
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
|
||||
};
|
||||
if !existing.tunnel_mode {
|
||||
return Ok(Some(bad_request_response(
|
||||
"non-tunnel mode is no longer supported, please upgrade aether-proxy to use tunnel mode",
|
||||
)));
|
||||
}
|
||||
let Some(node) = state.apply_proxy_node_heartbeat(&mutation).await? else {
|
||||
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
|
||||
};
|
||||
return Ok(Some(
|
||||
Json(json!({
|
||||
"message": "heartbeat ok",
|
||||
"node": build_admin_proxy_node_payload(&node),
|
||||
}))
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("unregister_node")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
if !state.has_proxy_node_writer() {
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
}
|
||||
let input = match parse_json_body::<ProxyNodeUnregisterRequest>(request_body) {
|
||||
Ok(input) => input,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let node_id = match validate_node_id(&input.node_id) {
|
||||
Ok(node_id) => node_id,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let Some(node) = state.unregister_proxy_node(&node_id).await? else {
|
||||
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
|
||||
};
|
||||
return Ok(Some(
|
||||
Json(json!({
|
||||
"message": "unregistered",
|
||||
"node_id": node.id,
|
||||
}))
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("update_node_config")
|
||||
&& request_context.method() == http::Method::PUT
|
||||
{
|
||||
if !state.has_proxy_node_writer() {
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
}
|
||||
let Some(node_id) = admin_proxy_node_config_node_id_from_path(request_context.path())
|
||||
else {
|
||||
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
|
||||
};
|
||||
let raw = match parse_json_object_body(request_body) {
|
||||
Ok(raw) => raw,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let Some(existing) = state.find_proxy_node(&node_id).await? else {
|
||||
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
|
||||
};
|
||||
if existing.is_manual {
|
||||
return Ok(Some(bad_request_response("手动节点不支持远程配置下发")));
|
||||
}
|
||||
let mutation = match validate_remote_config_request(node_id, &raw) {
|
||||
Ok(mutation) => mutation,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let Some(node) = state.update_proxy_node_remote_config(&mutation).await? else {
|
||||
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
|
||||
};
|
||||
return Ok(Some(
|
||||
Json(json!({
|
||||
"node_id": node.id,
|
||||
"config_version": node.config_version,
|
||||
"remote_config": node.remote_config,
|
||||
"node": build_admin_proxy_node_payload(&node),
|
||||
}))
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("batch_upgrade_nodes")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
if !state.has_proxy_node_reader()
|
||||
|| !state.has_proxy_node_writer()
|
||||
|| !state.app().data.has_system_config_store()
|
||||
{
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
}
|
||||
let input = match parse_json_body::<ProxyNodeBatchUpgradeRequest>(request_body) {
|
||||
Ok(input) => input,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let version = match validate_version(&input.version) {
|
||||
Ok(version) => version,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let batch_size = match validate_batch_size(input.batch_size) {
|
||||
Ok(batch_size) => batch_size,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let cooldown_secs = match validate_cooldown_secs(input.cooldown_secs) {
|
||||
Ok(cooldown_secs) => cooldown_secs,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let probe =
|
||||
match validate_probe_config(input.probe_url.as_deref(), input.probe_timeout_secs) {
|
||||
Ok(probe) => probe,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let rollout = start_proxy_upgrade_rollout(
|
||||
&state.app().data,
|
||||
version.clone(),
|
||||
batch_size,
|
||||
cooldown_secs,
|
||||
probe,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
|
||||
return Ok(Some(
|
||||
Json(json!({
|
||||
"version": version,
|
||||
"batch_size": rollout.batch_size,
|
||||
"cooldown_secs": rollout.cooldown_secs,
|
||||
"updated": rollout.updated,
|
||||
"skipped": rollout.skipped,
|
||||
"node_ids": rollout.node_ids,
|
||||
"blocked": rollout.blocked,
|
||||
"pending_node_ids": rollout.pending_node_ids,
|
||||
"rollout_active": rollout.rollout_active,
|
||||
"completed": rollout.completed,
|
||||
"remaining": rollout.remaining,
|
||||
}))
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("cancel_upgrade_rollout")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
if !state.has_proxy_node_reader()
|
||||
|| !state.has_proxy_node_writer()
|
||||
|| !state.app().data.has_system_config_store()
|
||||
{
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
}
|
||||
|
||||
let summary = cancel_proxy_upgrade_rollout(&state.app().data)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
return Ok(Some(
|
||||
Json(match summary {
|
||||
Some(summary) => json!({
|
||||
"cancelled": true,
|
||||
"version": summary.version,
|
||||
"pending_node_ids": summary.pending_node_ids,
|
||||
"conflict_node_ids": summary.conflict_node_ids,
|
||||
"completed": summary.completed,
|
||||
"remaining": summary.remaining,
|
||||
}),
|
||||
None => json!({
|
||||
"cancelled": false,
|
||||
"rollout_active": false,
|
||||
}),
|
||||
})
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("clear_upgrade_rollout_conflicts")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
if !state.has_proxy_node_reader()
|
||||
|| !state.has_proxy_node_writer()
|
||||
|| !state.app().data.has_system_config_store()
|
||||
{
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
}
|
||||
|
||||
let summary = clear_proxy_upgrade_rollout_conflicts(&state.app().data)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
return Ok(Some(
|
||||
Json(match summary {
|
||||
Some(summary) => json!({
|
||||
"version": summary.version,
|
||||
"cleared": summary.cleared_node_ids.len(),
|
||||
"node_ids": summary.cleared_node_ids,
|
||||
"updated": summary.updated,
|
||||
"blocked": summary.blocked,
|
||||
"pending_node_ids": summary.pending_node_ids,
|
||||
"rollout_active": summary.rollout_active,
|
||||
"completed": summary.completed,
|
||||
"remaining": summary.remaining,
|
||||
}),
|
||||
None => json!({
|
||||
"version": null,
|
||||
"cleared": 0,
|
||||
"node_ids": [],
|
||||
"updated": 0,
|
||||
"blocked": false,
|
||||
"pending_node_ids": [],
|
||||
"rollout_active": false,
|
||||
"completed": 0,
|
||||
"remaining": 0,
|
||||
}),
|
||||
})
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("restore_skipped_upgrade_rollout_nodes")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
if !state.has_proxy_node_reader()
|
||||
|| !state.has_proxy_node_writer()
|
||||
|| !state.app().data.has_system_config_store()
|
||||
{
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
}
|
||||
|
||||
let summary = restore_proxy_upgrade_rollout_skipped_nodes(&state.app().data)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
return Ok(Some(
|
||||
Json(match summary {
|
||||
Some(summary) => json!({
|
||||
"version": summary.version,
|
||||
"restored": summary.restored_node_ids.len(),
|
||||
"node_ids": summary.restored_node_ids,
|
||||
"skipped_node_ids": summary.skipped_node_ids,
|
||||
"updated": summary.updated,
|
||||
"blocked": summary.blocked,
|
||||
"pending_node_ids": summary.pending_node_ids,
|
||||
"rollout_active": summary.rollout_active,
|
||||
"completed": summary.completed,
|
||||
"remaining": summary.remaining,
|
||||
}),
|
||||
None => json!({
|
||||
"version": null,
|
||||
"restored": 0,
|
||||
"node_ids": [],
|
||||
"skipped_node_ids": [],
|
||||
"updated": 0,
|
||||
"blocked": false,
|
||||
"pending_node_ids": [],
|
||||
"rollout_active": false,
|
||||
"completed": 0,
|
||||
"remaining": 0,
|
||||
}),
|
||||
})
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("skip_upgrade_rollout_node")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
if !state.has_proxy_node_reader()
|
||||
|| !state.has_proxy_node_writer()
|
||||
|| !state.app().data.has_system_config_store()
|
||||
{
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
}
|
||||
let Some(node_id) = admin_proxy_node_upgrade_action_node_id_from_path(
|
||||
request_context.path(),
|
||||
"/upgrade/skip",
|
||||
) else {
|
||||
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
|
||||
};
|
||||
|
||||
let summary = skip_proxy_upgrade_rollout_node(&state.app().data, &node_id)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
return Ok(Some(
|
||||
Json(match summary {
|
||||
Some(summary) => json!({
|
||||
"version": summary.version,
|
||||
"node_id": summary.node_id,
|
||||
"skipped_node_ids": summary.skipped_node_ids,
|
||||
"updated": summary.updated,
|
||||
"blocked": summary.blocked,
|
||||
"pending_node_ids": summary.pending_node_ids,
|
||||
"rollout_active": summary.rollout_active,
|
||||
"completed": summary.completed,
|
||||
"remaining": summary.remaining,
|
||||
}),
|
||||
None => json!({
|
||||
"version": null,
|
||||
"node_id": node_id,
|
||||
"skipped_node_ids": [],
|
||||
"updated": 0,
|
||||
"blocked": false,
|
||||
"pending_node_ids": [],
|
||||
"rollout_active": false,
|
||||
"completed": 0,
|
||||
"remaining": 0,
|
||||
}),
|
||||
})
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("retry_upgrade_rollout_node")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
if !state.has_proxy_node_reader()
|
||||
|| !state.has_proxy_node_writer()
|
||||
|| !state.app().data.has_system_config_store()
|
||||
{
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
}
|
||||
let Some(node_id) = admin_proxy_node_upgrade_action_node_id_from_path(
|
||||
request_context.path(),
|
||||
"/upgrade/retry",
|
||||
) else {
|
||||
return Ok(Some(build_admin_proxy_nodes_not_found_response()));
|
||||
};
|
||||
|
||||
let summary = retry_proxy_upgrade_rollout_node(&state.app().data, &node_id)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
return Ok(Some(
|
||||
Json(match summary {
|
||||
Some(summary) => json!({
|
||||
"version": summary.version,
|
||||
"node_id": summary.node_id,
|
||||
"skipped_node_ids": summary.skipped_node_ids,
|
||||
"updated": summary.updated,
|
||||
"blocked": summary.blocked,
|
||||
"pending_node_ids": summary.pending_node_ids,
|
||||
"rollout_active": summary.rollout_active,
|
||||
"completed": summary.completed,
|
||||
"remaining": summary.remaining,
|
||||
}),
|
||||
None => json!({
|
||||
"version": null,
|
||||
"node_id": node_id,
|
||||
"skipped_node_ids": [],
|
||||
"updated": 0,
|
||||
"blocked": false,
|
||||
"pending_node_ids": [],
|
||||
"rollout_active": false,
|
||||
"completed": 0,
|
||||
"remaining": 0,
|
||||
}),
|
||||
})
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(Some(build_admin_proxy_nodes_data_unavailable_response()))
|
||||
}
|
||||
|
||||
fn validate_register_request(
|
||||
input: ProxyNodeRegisterRequest,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> Result<aether_data::repository::proxy_nodes::ProxyNodeRegistrationMutation, Response<Body>> {
|
||||
let name = normalize_required_string(&input.name, "name", 100)?;
|
||||
let ip = normalize_ip_address(&input.ip)?;
|
||||
let heartbeat_interval = validate_optional_i32_range(
|
||||
input.heartbeat_interval.unwrap_or(30),
|
||||
"heartbeat_interval",
|
||||
5,
|
||||
600,
|
||||
)?;
|
||||
if !input.tunnel_mode.unwrap_or(true) {
|
||||
return Err(bad_request_response("仅支持 tunnel_mode=true"));
|
||||
}
|
||||
validate_optional_counter(
|
||||
input.active_connections.map(i64::from),
|
||||
"active_connections",
|
||||
)?;
|
||||
validate_optional_counter(input.total_requests, "total_requests")?;
|
||||
validate_optional_counter(
|
||||
input.estimated_max_concurrency.map(i64::from),
|
||||
"estimated_max_concurrency",
|
||||
)?;
|
||||
if input
|
||||
.avg_latency_ms
|
||||
.is_some_and(|value| !value.is_finite() || value < 0.0)
|
||||
{
|
||||
return Err(bad_request_response("avg_latency_ms 必须是非负有限数值"));
|
||||
}
|
||||
validate_optional_object(input.hardware_info.as_ref(), "hardware_info")?;
|
||||
validate_optional_object(input.proxy_metadata.as_ref(), "proxy_metadata")?;
|
||||
|
||||
let registered_by = request_context
|
||||
.decision()
|
||||
.and_then(|decision| decision.admin_principal.as_ref())
|
||||
.map(|principal| principal.user_id.clone());
|
||||
|
||||
Ok(
|
||||
aether_data::repository::proxy_nodes::ProxyNodeRegistrationMutation {
|
||||
name,
|
||||
ip,
|
||||
port: i32::from(input.port.unwrap_or_default()),
|
||||
region: normalize_optional_string(input.region.as_deref(), "region", 100)?,
|
||||
heartbeat_interval,
|
||||
active_connections: input.active_connections,
|
||||
total_requests: input.total_requests,
|
||||
avg_latency_ms: input.avg_latency_ms,
|
||||
hardware_info: input.hardware_info,
|
||||
estimated_max_concurrency: input.estimated_max_concurrency,
|
||||
proxy_metadata: input.proxy_metadata,
|
||||
proxy_version: normalize_optional_string(
|
||||
input.proxy_version.as_deref(),
|
||||
"proxy_version",
|
||||
20,
|
||||
)?,
|
||||
registered_by,
|
||||
tunnel_mode: true,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn admin_proxy_node_upgrade_action_node_id_from_path(path: &str, suffix: &str) -> Option<String> {
|
||||
let normalized = path.trim_end_matches('/');
|
||||
let node_id = normalized.strip_prefix("/api/admin/proxy-nodes/")?;
|
||||
let node_id = node_id.strip_suffix(suffix)?;
|
||||
if node_id.is_empty() || node_id.contains('/') {
|
||||
None
|
||||
} else {
|
||||
Some(node_id.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_batch_size(batch_size: Option<usize>) -> Result<usize, Response<Body>> {
|
||||
let batch_size = batch_size.unwrap_or(DEFAULT_PROXY_UPGRADE_BATCH_SIZE);
|
||||
if (1..=100).contains(&batch_size) {
|
||||
Ok(batch_size)
|
||||
} else {
|
||||
Err(bad_request_response("batch_size 必须在 1 到 100 之间"))
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_cooldown_secs(cooldown_secs: Option<u64>) -> Result<u64, Response<Body>> {
|
||||
let cooldown_secs = cooldown_secs.unwrap_or(DEFAULT_PROXY_UPGRADE_COOLDOWN_SECS);
|
||||
if cooldown_secs <= 3600 {
|
||||
Ok(cooldown_secs)
|
||||
} else {
|
||||
Err(bad_request_response("cooldown_secs 不能超过 3600"))
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_probe_config(
|
||||
probe_url: Option<&str>,
|
||||
probe_timeout_secs: Option<u64>,
|
||||
) -> Result<Option<ProxyUpgradeRolloutProbeConfig>, Response<Body>> {
|
||||
let Some(probe_url) = probe_url.map(str::trim).filter(|value| !value.is_empty()) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let parsed = reqwest::Url::parse(probe_url)
|
||||
.map_err(|_| bad_request_response("probe_url 必须是合法的 http/https URL"))?;
|
||||
if !matches!(parsed.scheme(), "http" | "https") {
|
||||
return Err(bad_request_response("probe_url 仅支持 http 或 https"));
|
||||
}
|
||||
if parsed.as_str().len() > 2048 {
|
||||
return Err(bad_request_response("probe_url 长度不能超过 2048"));
|
||||
}
|
||||
let timeout_secs = probe_timeout_secs.unwrap_or(DEFAULT_PROXY_UPGRADE_PROBE_TIMEOUT_SECS);
|
||||
if !(5..=60).contains(&timeout_secs) {
|
||||
return Err(bad_request_response(
|
||||
"probe_timeout_secs 必须在 5 到 60 秒之间",
|
||||
));
|
||||
}
|
||||
Ok(Some(ProxyUpgradeRolloutProbeConfig {
|
||||
url: parsed.to_string(),
|
||||
timeout_secs,
|
||||
}))
|
||||
}
|
||||
|
||||
fn validate_heartbeat_request(
|
||||
input: ProxyNodeHeartbeatRequest,
|
||||
) -> Result<aether_data::repository::proxy_nodes::ProxyNodeHeartbeatMutation, Response<Body>> {
|
||||
let node_id = validate_node_id(&input.node_id)?;
|
||||
if let Some(interval) = input.heartbeat_interval {
|
||||
validate_optional_i32_range(interval, "heartbeat_interval", 5, 600)?;
|
||||
}
|
||||
validate_optional_counter(
|
||||
input.active_connections.map(i64::from),
|
||||
"active_connections",
|
||||
)?;
|
||||
validate_optional_counter(input.total_requests, "total_requests")?;
|
||||
validate_optional_counter(input.failed_requests, "failed_requests")?;
|
||||
validate_optional_counter(input.dns_failures, "dns_failures")?;
|
||||
validate_optional_counter(input.stream_errors, "stream_errors")?;
|
||||
if input
|
||||
.avg_latency_ms
|
||||
.is_some_and(|value| !value.is_finite() || value < 0.0)
|
||||
{
|
||||
return Err(bad_request_response("avg_latency_ms 必须是非负有限数值"));
|
||||
}
|
||||
validate_optional_object(input.proxy_metadata.as_ref(), "proxy_metadata")?;
|
||||
|
||||
Ok(
|
||||
aether_data::repository::proxy_nodes::ProxyNodeHeartbeatMutation {
|
||||
node_id,
|
||||
heartbeat_interval: input.heartbeat_interval,
|
||||
active_connections: input.active_connections,
|
||||
total_requests_delta: input.total_requests,
|
||||
avg_latency_ms: input.avg_latency_ms,
|
||||
failed_requests_delta: input.failed_requests,
|
||||
dns_failures_delta: input.dns_failures,
|
||||
stream_errors_delta: input.stream_errors,
|
||||
proxy_metadata: input.proxy_metadata,
|
||||
proxy_version: normalize_optional_string(
|
||||
input.proxy_version.as_deref(),
|
||||
"proxy_version",
|
||||
20,
|
||||
)?,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn validate_remote_config_request(
|
||||
node_id: String,
|
||||
raw: &serde_json::Map<String, Value>,
|
||||
) -> Result<aether_data::repository::proxy_nodes::ProxyNodeRemoteConfigMutation, Response<Body>> {
|
||||
let node_name = match raw.get("node_name") {
|
||||
Some(Value::Null) | None => None,
|
||||
Some(Value::String(value)) => Some(normalize_required_string(value, "node_name", 100)?),
|
||||
Some(_) => return Err(bad_request_response("node_name 必须是字符串")),
|
||||
};
|
||||
|
||||
let allowed_ports = match raw.get("allowed_ports") {
|
||||
Some(Value::Null) | None => None,
|
||||
Some(Value::Array(items)) => {
|
||||
let mut ports = Vec::with_capacity(items.len());
|
||||
for item in items {
|
||||
let Some(port) = item.as_u64() else {
|
||||
return Err(bad_request_response("allowed_ports 必须是端口数字数组"));
|
||||
};
|
||||
if !(1..=65535).contains(&port) {
|
||||
return Err(bad_request_response("allowed_ports 仅支持 1-65535"));
|
||||
}
|
||||
ports.push(port as u16);
|
||||
}
|
||||
Some(ports)
|
||||
}
|
||||
Some(_) => return Err(bad_request_response("allowed_ports 必须是端口数字数组")),
|
||||
};
|
||||
|
||||
let log_level = match raw.get("log_level") {
|
||||
Some(Value::Null) | None => None,
|
||||
Some(Value::String(value)) => {
|
||||
let normalized = normalize_required_string(value, "log_level", 16)?;
|
||||
if !matches!(
|
||||
normalized.as_str(),
|
||||
"trace" | "debug" | "info" | "warn" | "error"
|
||||
) {
|
||||
return Err(bad_request_response(
|
||||
"log_level 必须是 trace/debug/info/warn/error 之一",
|
||||
));
|
||||
}
|
||||
Some(normalized)
|
||||
}
|
||||
Some(_) => return Err(bad_request_response("log_level 必须是字符串")),
|
||||
};
|
||||
|
||||
let heartbeat_interval = match raw.get("heartbeat_interval") {
|
||||
Some(Value::Null) | None => None,
|
||||
Some(value) => Some(validate_json_i32_range(
|
||||
value,
|
||||
"heartbeat_interval",
|
||||
5,
|
||||
600,
|
||||
)?),
|
||||
};
|
||||
|
||||
let scheduling_state = if raw.contains_key("scheduling_state") {
|
||||
match raw.get("scheduling_state") {
|
||||
Some(Value::Null) | None => Some(None),
|
||||
Some(Value::String(value)) => {
|
||||
let normalized = normalize_required_string(value, "scheduling_state", 16)?;
|
||||
match normalized.as_str() {
|
||||
"active" => Some(None),
|
||||
"draining" | "cordoned" => Some(Some(normalized)),
|
||||
_ => {
|
||||
return Err(bad_request_response(
|
||||
"scheduling_state 必须是 active/draining/cordoned 之一",
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
Some(_) => return Err(bad_request_response("scheduling_state 必须是字符串或 null")),
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let upgrade_to = if raw.contains_key("upgrade_to") {
|
||||
match raw.get("upgrade_to") {
|
||||
Some(Value::Null) | None => Some(None),
|
||||
Some(Value::String(value)) => {
|
||||
let normalized = value.trim();
|
||||
if normalized.is_empty() {
|
||||
Some(None)
|
||||
} else {
|
||||
Some(Some(validate_version(normalized)?))
|
||||
}
|
||||
}
|
||||
Some(_) => return Err(bad_request_response("upgrade_to 必须是字符串或 null")),
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
Ok(
|
||||
aether_data::repository::proxy_nodes::ProxyNodeRemoteConfigMutation {
|
||||
node_id,
|
||||
node_name,
|
||||
allowed_ports,
|
||||
log_level,
|
||||
heartbeat_interval,
|
||||
scheduling_state,
|
||||
upgrade_to,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn admin_proxy_node_config_node_id_from_path(path: &str) -> Option<String> {
|
||||
let value = path
|
||||
.strip_prefix("/api/admin/proxy-nodes/")?
|
||||
.strip_suffix("/config")?;
|
||||
if value.is_empty() || value.contains('/') {
|
||||
None
|
||||
} else {
|
||||
Some(value.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_json_body<T: DeserializeOwned>(request_body: Option<&Bytes>) -> Result<T, Response<Body>> {
|
||||
let Some(request_body) = request_body else {
|
||||
return Err(bad_request_response("请求体不能为空"));
|
||||
};
|
||||
let raw_value = serde_json::from_slice::<Value>(request_body)
|
||||
.map_err(|_| bad_request_response(JSON_OBJECT_REQUIRED_DETAIL))?;
|
||||
serde_json::from_value::<T>(raw_value)
|
||||
.map_err(|_| bad_request_response(JSON_OBJECT_REQUIRED_DETAIL))
|
||||
}
|
||||
|
||||
fn parse_json_object_body(
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<serde_json::Map<String, Value>, Response<Body>> {
|
||||
let Some(request_body) = request_body else {
|
||||
return Err(bad_request_response("请求体不能为空"));
|
||||
};
|
||||
let raw_value = serde_json::from_slice::<Value>(request_body)
|
||||
.map_err(|_| bad_request_response(JSON_OBJECT_REQUIRED_DETAIL))?;
|
||||
raw_value
|
||||
.as_object()
|
||||
.cloned()
|
||||
.ok_or_else(|| bad_request_response(JSON_OBJECT_REQUIRED_DETAIL))
|
||||
}
|
||||
|
||||
fn validate_node_id(value: &str) -> Result<String, Response<Body>> {
|
||||
normalize_required_string(value, "node_id", 36)
|
||||
}
|
||||
|
||||
fn validate_version(value: &str) -> Result<String, Response<Body>> {
|
||||
normalize_required_string(value, "version", 50)
|
||||
}
|
||||
|
||||
fn normalize_required_string(
|
||||
value: &str,
|
||||
field: &str,
|
||||
max_len: usize,
|
||||
) -> Result<String, Response<Body>> {
|
||||
let normalized = value.trim();
|
||||
if normalized.is_empty() {
|
||||
return Err(bad_request_response(format!("{field} 不能为空")));
|
||||
}
|
||||
if normalized.chars().count() > max_len {
|
||||
return Err(bad_request_response(format!(
|
||||
"{field} 长度不能超过 {max_len}"
|
||||
)));
|
||||
}
|
||||
Ok(normalized.to_string())
|
||||
}
|
||||
|
||||
fn normalize_optional_string(
|
||||
value: Option<&str>,
|
||||
field: &str,
|
||||
max_len: usize,
|
||||
) -> Result<Option<String>, Response<Body>> {
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
let normalized = value.trim();
|
||||
if normalized.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
if normalized.chars().count() > max_len {
|
||||
return Err(bad_request_response(format!(
|
||||
"{field} 长度不能超过 {max_len}"
|
||||
)));
|
||||
}
|
||||
Ok(Some(normalized.to_string()))
|
||||
}
|
||||
|
||||
fn normalize_ip_address(value: &str) -> Result<String, Response<Body>> {
|
||||
let normalized = value.trim();
|
||||
normalized
|
||||
.parse::<std::net::IpAddr>()
|
||||
.map(|ip| ip.to_string())
|
||||
.map_err(|_| bad_request_response("ip 必须是合法的 IPv4/IPv6 地址"))
|
||||
}
|
||||
|
||||
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} 必须是非负整数")));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_optional_i32_range(
|
||||
value: i32,
|
||||
field: &str,
|
||||
min: i32,
|
||||
max: i32,
|
||||
) -> Result<i32, Response<Body>> {
|
||||
if !(min..=max).contains(&value) {
|
||||
return Err(bad_request_response(format!(
|
||||
"{field} 必须在 {min}-{max} 范围内"
|
||||
)));
|
||||
}
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
fn validate_json_i32_range(
|
||||
value: &Value,
|
||||
field: &str,
|
||||
min: i32,
|
||||
max: i32,
|
||||
) -> Result<i32, Response<Body>> {
|
||||
let Some(raw) = value.as_i64() else {
|
||||
return Err(bad_request_response(format!("{field} 必须是整数")));
|
||||
};
|
||||
let parsed =
|
||||
i32::try_from(raw).map_err(|_| bad_request_response(format!("{field} 超出范围")))?;
|
||||
validate_optional_i32_range(parsed, field, min, max)
|
||||
}
|
||||
|
||||
fn validate_optional_object(value: Option<&Value>, field: &str) -> Result<(), Response<Body>> {
|
||||
if value.is_some_and(|value| !value.is_object()) {
|
||||
return Err(bad_request_response(format!("{field} 必须是 JSON 对象")));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn bad_request_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
@@ -38,6 +38,7 @@ pub(crate) async fn maybe_build_local_admin_system_response(
|
||||
if let Some(response) = proxy_nodes::maybe_build_local_admin_proxy_nodes_response(
|
||||
&request.state(),
|
||||
&request.request_context(),
|
||||
request.request_body(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
|
||||
@@ -820,7 +820,11 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
||||
};
|
||||
|
||||
let response = match state.apply_proxy_node_heartbeat(&mutation).await {
|
||||
Ok(Some(node)) => Json(build_internal_tunnel_heartbeat_ack(&node)).into_response(),
|
||||
Ok(Some(node)) => Json(build_internal_tunnel_heartbeat_ack(
|
||||
&node,
|
||||
payload.heartbeat_id,
|
||||
))
|
||||
.into_response(),
|
||||
Ok(None) => build_internal_control_error_response(
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("heartbeat sync failed: ProxyNode {node_id} 不存在"),
|
||||
@@ -850,7 +854,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
||||
connected: payload.connected,
|
||||
conn_count: payload.conn_count,
|
||||
detail: None,
|
||||
observed_at_unix_secs: None,
|
||||
observed_at_unix_secs: payload.observed_at_unix_secs,
|
||||
};
|
||||
|
||||
let response = match state.update_proxy_node_tunnel_status(&mutation).await {
|
||||
|
||||
@@ -349,22 +349,26 @@ pub(crate) fn gateway_error_message(error: GatewayError) -> String {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn build_internal_tunnel_heartbeat_ack(node: &StoredProxyNode) -> serde_json::Value {
|
||||
let Some(remote_config) = node.remote_config.as_ref() else {
|
||||
return json!({});
|
||||
};
|
||||
|
||||
pub(crate) fn build_internal_tunnel_heartbeat_ack(
|
||||
node: &StoredProxyNode,
|
||||
heartbeat_id: Option<u64>,
|
||||
) -> serde_json::Value {
|
||||
let mut payload = serde_json::Map::new();
|
||||
payload.insert("remote_config".to_string(), remote_config.clone());
|
||||
payload.insert("config_version".to_string(), json!(node.config_version));
|
||||
if let Some(upgrade_to) = remote_config
|
||||
.as_object()
|
||||
.and_then(|value| value.get("upgrade_to"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
payload.insert("upgrade_to".to_string(), json!(upgrade_to));
|
||||
if let Some(heartbeat_id) = 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));
|
||||
if let Some(upgrade_to) = remote_config
|
||||
.as_object()
|
||||
.and_then(|value| value.get("upgrade_to"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
payload.insert("upgrade_to".to_string(), json!(upgrade_to));
|
||||
}
|
||||
}
|
||||
serde_json::Value::Object(payload)
|
||||
}
|
||||
|
||||
@@ -51,6 +51,7 @@ use crate::{
|
||||
use axum::body::{to_bytes, Body, Bytes};
|
||||
use axum::extract::{ConnectInfo, Request, State};
|
||||
use axum::http::{self, header::HeaderName, header::HeaderValue, Response};
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::time::Instant;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
@@ -96,6 +97,141 @@ fn execution_runtime_candidate_header_value(decision: &GatewayControlDecision) -
|
||||
}
|
||||
}
|
||||
|
||||
fn extract_management_token_bearer(headers: &http::HeaderMap) -> Option<String> {
|
||||
let header = crate::headers::header_value_str(headers, http::header::AUTHORIZATION.as_str())?;
|
||||
let token = header
|
||||
.strip_prefix("Bearer ")
|
||||
.or_else(|| header.strip_prefix("bearer "))?
|
||||
.trim()
|
||||
.to_string();
|
||||
(!token.is_empty() && token.starts_with("ae_")).then_some(token)
|
||||
}
|
||||
|
||||
fn hash_management_token(value: &str) -> String {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(value.as_bytes());
|
||||
format!("{:x}", hasher.finalize())
|
||||
}
|
||||
|
||||
fn remote_ip_allowed(allowed_ips: Option<&serde_json::Value>, remote_ip: std::net::IpAddr) -> bool {
|
||||
let Some(allowed_ips) = allowed_ips else {
|
||||
return true;
|
||||
};
|
||||
let Some(items) = allowed_ips.as_array() else {
|
||||
return false;
|
||||
};
|
||||
if items.is_empty() {
|
||||
return false;
|
||||
}
|
||||
items
|
||||
.iter()
|
||||
.filter_map(serde_json::Value::as_str)
|
||||
.any(|value| ip_or_cidr_matches(value, remote_ip))
|
||||
}
|
||||
|
||||
fn ip_or_cidr_matches(pattern: &str, remote_ip: std::net::IpAddr) -> bool {
|
||||
let pattern = pattern.trim();
|
||||
if pattern.is_empty() {
|
||||
return false;
|
||||
}
|
||||
if let Ok(ip) = pattern.parse::<std::net::IpAddr>() {
|
||||
return ip == remote_ip;
|
||||
}
|
||||
let Some((network, prefix)) = pattern.split_once('/') else {
|
||||
return false;
|
||||
};
|
||||
let Ok(prefix) = prefix.trim().parse::<u8>() else {
|
||||
return false;
|
||||
};
|
||||
match (network.trim().parse::<std::net::IpAddr>(), remote_ip) {
|
||||
(Ok(std::net::IpAddr::V4(network)), std::net::IpAddr::V4(remote)) if prefix <= 32 => {
|
||||
let mask = if prefix == 0 {
|
||||
0
|
||||
} else {
|
||||
u32::MAX << (32 - prefix)
|
||||
};
|
||||
(u32::from(network) & mask) == (u32::from(remote) & mask)
|
||||
}
|
||||
(Ok(std::net::IpAddr::V6(network)), std::net::IpAddr::V6(remote)) if prefix <= 128 => {
|
||||
let mask = if prefix == 0 {
|
||||
0
|
||||
} else {
|
||||
u128::MAX << (128 - prefix)
|
||||
};
|
||||
(u128::from(network) & mask) == (u128::from(remote) & mask)
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
async fn maybe_promote_management_token_admin_principal(
|
||||
state: &AppState,
|
||||
remote_addr: &std::net::SocketAddr,
|
||||
headers: &http::HeaderMap,
|
||||
trace_id: &str,
|
||||
request_context: &mut GatewayPublicRequestContext,
|
||||
) -> Result<(), GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_mut() else {
|
||||
return Ok(());
|
||||
};
|
||||
if decision.route_class.as_deref() != Some("admin_proxy") || decision.admin_principal.is_some()
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let Some(token) = extract_management_token_bearer(headers) else {
|
||||
return Ok(());
|
||||
};
|
||||
let token_hash = hash_management_token(&token);
|
||||
let Some(token_with_user) = state
|
||||
.get_management_token_with_user_by_hash(&token_hash)
|
||||
.await?
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
if !token_with_user.token.is_active {
|
||||
return Ok(());
|
||||
}
|
||||
if token_with_user
|
||||
.token
|
||||
.expires_at_unix_secs
|
||||
.is_some_and(|value| value <= chrono::Utc::now().timestamp().max(0) as u64)
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
if !remote_ip_allowed(token_with_user.token.allowed_ips.as_ref(), remote_addr.ip()) {
|
||||
return Ok(());
|
||||
}
|
||||
let Some(user) = state.find_user_auth_by_id(&token_with_user.user.id).await? else {
|
||||
return Ok(());
|
||||
};
|
||||
if !user.is_active || user.is_deleted || !user.role.eq_ignore_ascii_case("admin") {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
decision.admin_principal = Some(crate::control::GatewayAdminPrincipalContext {
|
||||
user_id: user.id.clone(),
|
||||
user_role: user.role.clone(),
|
||||
session_id: None,
|
||||
management_token_id: Some(token_with_user.token.id.clone()),
|
||||
});
|
||||
|
||||
let remote_ip = remote_addr.ip().to_string();
|
||||
if let Err(err) = state
|
||||
.record_management_token_usage(&token_with_user.token.id, Some(remote_ip.as_str()))
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
token_id = %token_with_user.token.id,
|
||||
error = ?err,
|
||||
"gateway failed to record management token usage"
|
||||
);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn maybe_forward_public_request_to_tunnel_owner(
|
||||
state: &AppState,
|
||||
remote_addr: &std::net::SocketAddr,
|
||||
@@ -379,7 +515,7 @@ pub(crate) async fn proxy_request(
|
||||
));
|
||||
}
|
||||
let request_context_started_at = Instant::now();
|
||||
let request_context = crate::control::resolve_public_request_context(
|
||||
let mut request_context = crate::control::resolve_public_request_context(
|
||||
&state,
|
||||
&parts.method,
|
||||
&parts.uri,
|
||||
@@ -387,6 +523,14 @@ pub(crate) async fn proxy_request(
|
||||
&trace_id,
|
||||
)
|
||||
.await?;
|
||||
maybe_promote_management_token_admin_principal(
|
||||
&state,
|
||||
&remote_addr,
|
||||
&parts.headers,
|
||||
&trace_id,
|
||||
&mut request_context,
|
||||
)
|
||||
.await?;
|
||||
let request_context_ms = request_context_started_at.elapsed().as_millis() as u64;
|
||||
if request_context
|
||||
.control_decision
|
||||
|
||||
@@ -5,6 +5,8 @@ use std::collections::BTreeMap;
|
||||
pub(crate) struct InternalTunnelHeartbeatRequest {
|
||||
pub(crate) node_id: String,
|
||||
#[serde(default)]
|
||||
pub(crate) heartbeat_id: Option<u64>,
|
||||
#[serde(default)]
|
||||
pub(crate) heartbeat_interval: Option<i32>,
|
||||
#[serde(default)]
|
||||
pub(crate) active_connections: Option<i32>,
|
||||
@@ -30,6 +32,8 @@ pub(crate) struct InternalTunnelNodeStatusRequest {
|
||||
pub(crate) connected: bool,
|
||||
#[serde(default)]
|
||||
pub(crate) conn_count: i32,
|
||||
#[serde(default)]
|
||||
pub(crate) observed_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
|
||||
@@ -305,6 +305,11 @@ pub(crate) fn admin_proxy_local_requires_buffered_body(
|
||||
| (Some("api_keys_manage"), http::Method::PUT, Some("update_api_key"))
|
||||
| (Some("api_keys_manage"), http::Method::PATCH, Some("toggle_api_key"))
|
||||
| (Some("adaptive_manage"), http::Method::PATCH, Some("toggle_mode"))
|
||||
| (Some("proxy_nodes_manage"), http::Method::POST, Some("register_node"))
|
||||
| (Some("proxy_nodes_manage"), http::Method::POST, Some("heartbeat_node"))
|
||||
| (Some("proxy_nodes_manage"), http::Method::POST, Some("unregister_node"))
|
||||
| (Some("proxy_nodes_manage"), http::Method::POST, Some("batch_upgrade_nodes"))
|
||||
| (Some("proxy_nodes_manage"), http::Method::PUT, Some("update_node_config"))
|
||||
| (Some("security_manage"), http::Method::POST, Some("blacklist_add"))
|
||||
| (Some("security_manage"), http::Method::POST, Some("whitelist_add"))
|
||||
| (Some("users_manage"), http::Method::POST, Some("create_user"))
|
||||
|
||||
@@ -904,6 +904,13 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
if state.run_postgres_migrations().await? {
|
||||
info!("database migrations complete");
|
||||
}
|
||||
let reset_stale_proxy_nodes = state.reset_stale_proxy_node_tunnel_statuses().await?;
|
||||
if reset_stale_proxy_nodes > 0 {
|
||||
info!(
|
||||
reset_stale_proxy_nodes,
|
||||
"reset stale tunnel-connected proxy nodes on startup"
|
||||
);
|
||||
}
|
||||
state.bootstrap_admin_from_env().await?;
|
||||
|
||||
let background_tasks = if args.node_role.spawns_background_tasks() {
|
||||
|
||||
@@ -3,10 +3,18 @@ mod runtime;
|
||||
mod tests;
|
||||
|
||||
pub(crate) use runtime::{
|
||||
perform_provider_checkin_once, spawn_audit_cleanup_worker, spawn_db_maintenance_worker,
|
||||
spawn_gemini_file_mapping_cleanup_worker, spawn_pending_cleanup_worker,
|
||||
spawn_pool_monitor_worker, spawn_provider_checkin_worker,
|
||||
cancel_proxy_upgrade_rollout, clear_proxy_upgrade_rollout_conflicts,
|
||||
inspect_proxy_upgrade_rollout, perform_provider_checkin_once,
|
||||
record_proxy_upgrade_traffic_success, restore_proxy_upgrade_rollout_skipped_nodes,
|
||||
retry_proxy_upgrade_rollout_node, skip_proxy_upgrade_rollout_node, spawn_audit_cleanup_worker,
|
||||
spawn_db_maintenance_worker, spawn_gemini_file_mapping_cleanup_worker,
|
||||
spawn_pending_cleanup_worker, spawn_pool_monitor_worker, spawn_provider_checkin_worker,
|
||||
spawn_proxy_node_stale_cleanup_worker, spawn_proxy_upgrade_rollout_worker,
|
||||
spawn_request_candidate_cleanup_worker, spawn_stats_aggregation_worker,
|
||||
spawn_stats_hourly_aggregation_worker, spawn_usage_cleanup_worker,
|
||||
spawn_wallet_daily_usage_aggregation_worker, ProviderCheckinRunSummary,
|
||||
spawn_wallet_daily_usage_aggregation_worker, start_proxy_upgrade_rollout,
|
||||
ProviderCheckinRunSummary, ProxyUpgradeRolloutCancelSummary,
|
||||
ProxyUpgradeRolloutConflictClearSummary, ProxyUpgradeRolloutNodeActionSummary,
|
||||
ProxyUpgradeRolloutProbeConfig, ProxyUpgradeRolloutSkippedRestoreSummary,
|
||||
ProxyUpgradeRolloutStatus, ProxyUpgradeRolloutTrackedNodeState,
|
||||
};
|
||||
|
||||
@@ -29,6 +29,10 @@ mod db_maintenance;
|
||||
mod pending_cleanup;
|
||||
#[path = "runtime/provider_checkin.rs"]
|
||||
mod provider_checkin;
|
||||
#[path = "runtime/proxy_node_staleness.rs"]
|
||||
mod proxy_node_staleness;
|
||||
#[path = "runtime/proxy_upgrade_rollout.rs"]
|
||||
mod proxy_upgrade_rollout;
|
||||
#[path = "runtime/request_candidate_cleanup.rs"]
|
||||
mod request_candidate_cleanup;
|
||||
#[path = "runtime/runners.rs"]
|
||||
@@ -53,6 +57,18 @@ use config::*;
|
||||
use db_maintenance::*;
|
||||
use pending_cleanup::*;
|
||||
pub(crate) use provider_checkin::{perform_provider_checkin_once, ProviderCheckinRunSummary};
|
||||
use proxy_node_staleness::*;
|
||||
use proxy_upgrade_rollout::*;
|
||||
pub(crate) use proxy_upgrade_rollout::{
|
||||
cancel_proxy_upgrade_rollout, clear_proxy_upgrade_rollout_conflicts,
|
||||
collect_proxy_upgrade_rollout_probes, inspect_proxy_upgrade_rollout,
|
||||
record_proxy_upgrade_traffic_success, restore_proxy_upgrade_rollout_skipped_nodes,
|
||||
retry_proxy_upgrade_rollout_node, skip_proxy_upgrade_rollout_node, start_proxy_upgrade_rollout,
|
||||
ProxyUpgradeRolloutCancelSummary, ProxyUpgradeRolloutConflictClearSummary,
|
||||
ProxyUpgradeRolloutNodeActionSummary, ProxyUpgradeRolloutPendingProbe,
|
||||
ProxyUpgradeRolloutProbeConfig, ProxyUpgradeRolloutSkippedRestoreSummary,
|
||||
ProxyUpgradeRolloutStatus, ProxyUpgradeRolloutSummary, ProxyUpgradeRolloutTrackedNodeState,
|
||||
};
|
||||
use request_candidate_cleanup::*;
|
||||
use runners::*;
|
||||
use schedule::*;
|
||||
@@ -71,6 +87,10 @@ 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_UPGRADE_ROLLOUT_INTERVAL: Duration = Duration::from_secs(15);
|
||||
const PROXY_NODE_STALE_MIN_GRACE_SECS: u64 = 90;
|
||||
const PROXY_NODE_STALE_MISSED_HEARTBEATS: u64 = 3;
|
||||
const POOL_MONITOR_INTERVAL: Duration = Duration::from_secs(5 * 60);
|
||||
const PROVIDER_CHECKIN_CONCURRENCY: usize = 3;
|
||||
const PROVIDER_CHECKIN_DEFAULT_TIME: &str = "01:05";
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use aether_data::repository::proxy_nodes::ProxyNodeTunnelStatusMutation;
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use crate::data::GatewayDataState;
|
||||
|
||||
use super::{PROXY_NODE_STALE_MIN_GRACE_SECS, PROXY_NODE_STALE_MISSED_HEARTBEATS};
|
||||
|
||||
fn current_unix_secs() -> u64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs()
|
||||
}
|
||||
|
||||
fn stale_proxy_node_grace_secs(heartbeat_interval: i32) -> u64 {
|
||||
let interval = u64::try_from(heartbeat_interval.max(5)).unwrap_or(5);
|
||||
interval
|
||||
.saturating_mul(PROXY_NODE_STALE_MISSED_HEARTBEATS)
|
||||
.max(PROXY_NODE_STALE_MIN_GRACE_SECS)
|
||||
}
|
||||
|
||||
pub(super) async fn cleanup_stale_proxy_nodes_once(
|
||||
data: &GatewayDataState,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
if !data.has_proxy_node_reader() || !data.has_proxy_node_writer() {
|
||||
return Ok(0);
|
||||
}
|
||||
|
||||
let now_unix_secs = current_unix_secs();
|
||||
let nodes = data.list_proxy_nodes().await?;
|
||||
let mut updated = 0usize;
|
||||
|
||||
for node in nodes {
|
||||
if node.is_manual || !node.tunnel_connected {
|
||||
continue;
|
||||
}
|
||||
|
||||
let grace_secs = stale_proxy_node_grace_secs(node.heartbeat_interval);
|
||||
let is_stale = node
|
||||
.last_heartbeat_at_unix_secs
|
||||
.map(|last_seen| last_seen.saturating_add(grace_secs) < now_unix_secs)
|
||||
.unwrap_or(true);
|
||||
if !is_stale {
|
||||
continue;
|
||||
}
|
||||
|
||||
let detail = format!(
|
||||
"[heartbeat_timeout] last_heartbeat_at={} grace_secs={}",
|
||||
node.last_heartbeat_at_unix_secs.unwrap_or(0),
|
||||
grace_secs
|
||||
);
|
||||
let mutation = ProxyNodeTunnelStatusMutation {
|
||||
node_id: node.id.clone(),
|
||||
connected: false,
|
||||
conn_count: 0,
|
||||
detail: Some(detail),
|
||||
observed_at_unix_secs: Some(now_unix_secs),
|
||||
};
|
||||
if data
|
||||
.update_proxy_node_tunnel_status(&mutation)
|
||||
.await?
|
||||
.is_some()
|
||||
{
|
||||
updated = updated.saturating_add(1);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(updated)
|
||||
}
|
||||
1070
apps/aether-gateway/src/maintenance/runtime/proxy_upgrade_rollout.rs
Normal file
1070
apps/aether-gateway/src/maintenance/runtime/proxy_upgrade_rollout.rs
Normal file
File diff suppressed because it is too large
Load Diff
@@ -1,15 +1,18 @@
|
||||
use aether_data_contracts::DataLayerError;
|
||||
use tracing::info;
|
||||
use tracing::{info, warn};
|
||||
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
use super::{
|
||||
cleanup_audit_logs_once, cleanup_expired_gemini_file_mappings_once,
|
||||
cleanup_request_candidates_once, cleanup_stale_pending_requests_once,
|
||||
perform_db_maintenance_once, perform_provider_checkin_once, perform_stats_aggregation_once,
|
||||
advance_proxy_upgrade_rollout_once, cleanup_audit_logs_once,
|
||||
cleanup_expired_gemini_file_mappings_once, cleanup_request_candidates_once,
|
||||
cleanup_stale_pending_requests_once, cleanup_stale_proxy_nodes_once,
|
||||
collect_proxy_upgrade_rollout_probes, perform_db_maintenance_once,
|
||||
perform_provider_checkin_once, perform_stats_aggregation_once,
|
||||
perform_stats_hourly_aggregation_once, perform_usage_cleanup_once,
|
||||
perform_wallet_daily_usage_aggregation_once, summarize_postgres_pool,
|
||||
perform_wallet_daily_usage_aggregation_once, record_proxy_upgrade_traffic_success,
|
||||
summarize_postgres_pool,
|
||||
};
|
||||
|
||||
pub(super) async fn run_audit_cleanup_once(data: &GatewayDataState) -> Result<(), DataLayerError> {
|
||||
@@ -42,6 +45,91 @@ pub(super) async fn run_gemini_file_mapping_cleanup_once(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) async fn run_proxy_node_stale_cleanup_once(
|
||||
data: &GatewayDataState,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let stale_marked = cleanup_stale_proxy_nodes_once(data).await?;
|
||||
if stale_marked > 0 {
|
||||
info!(
|
||||
event_name = "proxy_node_stale_cleanup_completed",
|
||||
log_type = "ops",
|
||||
worker = "proxy_node_stale_cleanup",
|
||||
stale_marked,
|
||||
"gateway marked stale proxy nodes offline"
|
||||
);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) async fn run_proxy_upgrade_rollout_once(state: &AppState) -> Result<(), DataLayerError> {
|
||||
let mut summary = advance_proxy_upgrade_rollout_once(&state.data).await?;
|
||||
let probes = collect_proxy_upgrade_rollout_probes(&state.data).await?;
|
||||
let mut probe_recorded = false;
|
||||
for probe in probes {
|
||||
match state
|
||||
.tunnel
|
||||
.probe_node_url(&probe.node_id, &probe.url, probe.timeout_secs)
|
||||
.await
|
||||
{
|
||||
Ok(status) if (200..300).contains(&status) => {
|
||||
let _ = record_proxy_upgrade_traffic_success(&state.data, &probe.node_id).await?;
|
||||
probe_recorded = true;
|
||||
info!(
|
||||
event_name = "proxy_upgrade_rollout_probe_succeeded",
|
||||
log_type = "ops",
|
||||
worker = "proxy_upgrade_rollout",
|
||||
node_id = %probe.node_id,
|
||||
url = %probe.url,
|
||||
status,
|
||||
"gateway confirmed proxy upgrade health probe"
|
||||
);
|
||||
}
|
||||
Ok(status) => {
|
||||
warn!(
|
||||
event_name = "proxy_upgrade_rollout_probe_unhealthy",
|
||||
log_type = "ops",
|
||||
worker = "proxy_upgrade_rollout",
|
||||
node_id = %probe.node_id,
|
||||
url = %probe.url,
|
||||
status,
|
||||
"gateway proxy upgrade health probe returned non-success status"
|
||||
);
|
||||
}
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "proxy_upgrade_rollout_probe_failed",
|
||||
log_type = "ops",
|
||||
worker = "proxy_upgrade_rollout",
|
||||
node_id = %probe.node_id,
|
||||
url = %probe.url,
|
||||
error = %error,
|
||||
"gateway proxy upgrade health probe failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if probe_recorded {
|
||||
summary = advance_proxy_upgrade_rollout_once(&state.data).await?;
|
||||
}
|
||||
if summary.updated > 0 || !summary.pending_node_ids.is_empty() || !summary.version.is_empty() {
|
||||
info!(
|
||||
event_name = "proxy_upgrade_rollout_checked",
|
||||
log_type = "ops",
|
||||
worker = "proxy_upgrade_rollout",
|
||||
version = %summary.version,
|
||||
updated = summary.updated,
|
||||
blocked = summary.blocked,
|
||||
pending = summary.pending_node_ids.len(),
|
||||
completed = summary.completed,
|
||||
remaining = summary.remaining,
|
||||
rollout_active = summary.rollout_active,
|
||||
"gateway checked proxy upgrade rollout"
|
||||
);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) async fn run_db_maintenance_once(data: &GatewayDataState) -> Result<(), DataLayerError> {
|
||||
let summary = perform_db_maintenance_once(data).await?;
|
||||
if summary.attempted > 0 {
|
||||
|
||||
@@ -1,24 +1,35 @@
|
||||
use std::collections::{HashSet, VecDeque};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Duration;
|
||||
|
||||
use aether_data::repository::proxy_nodes::{
|
||||
InMemoryProxyNodeRepository, ProxyNodeHeartbeatMutation, ProxyNodeReadRepository,
|
||||
ProxyNodeWriteRepository, StoredProxyNode,
|
||||
};
|
||||
use aether_runtime::bounded_queue;
|
||||
use axum::extract::ws::Message;
|
||||
use chrono::{DateTime, Utc};
|
||||
use chrono_tz::Tz;
|
||||
use serde_json::json;
|
||||
use tokio::sync::watch;
|
||||
|
||||
use super::{
|
||||
cleanup_audit_logs_with, next_daily_run_after, next_db_maintenance_run_after,
|
||||
advance_proxy_upgrade_rollout_once, cleanup_audit_logs_with, cleanup_stale_proxy_nodes_once,
|
||||
inspect_proxy_upgrade_rollout, next_daily_run_after, next_db_maintenance_run_after,
|
||||
next_stats_aggregation_run_after, next_stats_hourly_aggregation_run_after,
|
||||
pending_cleanup_batch_size, pending_cleanup_timeout_minutes, plan_pending_cleanup_batch,
|
||||
provider_checkin_schedule, run_db_maintenance_with, spawn_audit_cleanup_worker,
|
||||
spawn_db_maintenance_worker, spawn_pending_cleanup_worker, spawn_pool_monitor_worker,
|
||||
spawn_provider_checkin_worker, spawn_stats_aggregation_worker,
|
||||
spawn_stats_hourly_aggregation_worker, spawn_usage_cleanup_worker,
|
||||
spawn_wallet_daily_usage_aggregation_worker, stats_aggregation_target_day,
|
||||
provider_checkin_schedule, record_proxy_upgrade_traffic_success, run_db_maintenance_with,
|
||||
run_proxy_upgrade_rollout_once, spawn_audit_cleanup_worker, spawn_db_maintenance_worker,
|
||||
spawn_pending_cleanup_worker, spawn_pool_monitor_worker, spawn_provider_checkin_worker,
|
||||
spawn_proxy_node_stale_cleanup_worker, spawn_proxy_upgrade_rollout_worker,
|
||||
spawn_stats_aggregation_worker, spawn_stats_hourly_aggregation_worker,
|
||||
spawn_usage_cleanup_worker, spawn_wallet_daily_usage_aggregation_worker,
|
||||
start_proxy_upgrade_rollout, stats_aggregation_target_day,
|
||||
stats_hourly_aggregation_target_hour, summarize_postgres_pool, usage_cleanup_settings,
|
||||
usage_cleanup_window, wallet_daily_usage_aggregation_target, AppState, DbMaintenanceRunSummary,
|
||||
FailedPendingUsageRow, GatewayDataState, StalePendingUsageRow, UsageCleanupSettings,
|
||||
USAGE_CLEANUP_HOUR, USAGE_CLEANUP_MINUTE, WALLET_DAILY_USAGE_AGGREGATION_HOUR,
|
||||
WALLET_DAILY_USAGE_AGGREGATION_MINUTE,
|
||||
FailedPendingUsageRow, GatewayDataState, ProxyUpgradeRolloutProbeConfig, StalePendingUsageRow,
|
||||
UsageCleanupSettings, USAGE_CLEANUP_HOUR, USAGE_CLEANUP_MINUTE,
|
||||
WALLET_DAILY_USAGE_AGGREGATION_HOUR, WALLET_DAILY_USAGE_AGGREGATION_MINUTE,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
@@ -36,11 +47,476 @@ async fn spawn_pending_cleanup_worker_skips_when_postgres_unavailable() {
|
||||
assert!(spawn_pending_cleanup_worker(Arc::new(GatewayDataState::disabled())).is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn spawn_proxy_node_stale_cleanup_worker_skips_when_proxy_nodes_unavailable() {
|
||||
assert!(
|
||||
spawn_proxy_node_stale_cleanup_worker(Arc::new(GatewayDataState::disabled())).is_none()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn spawn_proxy_upgrade_rollout_worker_skips_when_proxy_nodes_unavailable() {
|
||||
let state = AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(GatewayDataState::disabled());
|
||||
assert!(spawn_proxy_upgrade_rollout_worker(state).is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn spawn_proxy_upgrade_rollout_worker_skips_when_system_config_unavailable() {
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![]));
|
||||
let data = GatewayDataState::with_proxy_node_repository_for_tests(repository);
|
||||
let state = AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(data);
|
||||
assert!(spawn_proxy_upgrade_rollout_worker(state).is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn spawn_pool_monitor_worker_skips_when_postgres_unavailable() {
|
||||
assert!(spawn_pool_monitor_worker(Arc::new(GatewayDataState::disabled())).is_none());
|
||||
}
|
||||
|
||||
fn sample_connected_proxy_node(
|
||||
node_id: &str,
|
||||
heartbeat_interval: i32,
|
||||
last_heartbeat_at_unix_secs: u64,
|
||||
) -> StoredProxyNode {
|
||||
StoredProxyNode::new(
|
||||
node_id.to_string(),
|
||||
format!("proxy-{node_id}"),
|
||||
"127.0.0.1".to_string(),
|
||||
0,
|
||||
false,
|
||||
"online".to_string(),
|
||||
heartbeat_interval,
|
||||
3,
|
||||
10,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
true,
|
||||
true,
|
||||
0,
|
||||
)
|
||||
.expect("node should build")
|
||||
.with_runtime_fields(
|
||||
Some("test".to_string()),
|
||||
None,
|
||||
Some(last_heartbeat_at_unix_secs),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
Some(last_heartbeat_at_unix_secs),
|
||||
None,
|
||||
Some(last_heartbeat_at_unix_secs),
|
||||
Some(last_heartbeat_at_unix_secs),
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stale_proxy_node_cleanup_marks_timed_out_tunnel_offline() {
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
|
||||
sample_connected_proxy_node("node-stale", 30, 1),
|
||||
]));
|
||||
let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository));
|
||||
|
||||
let updated = cleanup_stale_proxy_nodes_once(&data)
|
||||
.await
|
||||
.expect("cleanup should succeed");
|
||||
|
||||
assert_eq!(updated, 1);
|
||||
let node = repository
|
||||
.find_proxy_node("node-stale")
|
||||
.await
|
||||
.expect("lookup should succeed")
|
||||
.expect("node should exist");
|
||||
assert_eq!(node.status, "offline");
|
||||
assert_eq!(node.tunnel_connected, false);
|
||||
assert_eq!(node.active_connections, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn proxy_upgrade_rollout_advances_next_wave_after_version_health_confirmation() {
|
||||
let mut alpha = sample_connected_proxy_node("node-alpha", 30, 1_800_000_000);
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.proxy_metadata = Some(json!({"version": "1.0.0"}));
|
||||
alpha.remote_config = None;
|
||||
alpha.config_version = 0;
|
||||
|
||||
let mut zeta = sample_connected_proxy_node("node-zeta", 30, 1_800_000_000);
|
||||
zeta.name = "zeta".to_string();
|
||||
zeta.proxy_metadata = Some(json!({"version": "1.0.0"}));
|
||||
zeta.remote_config = None;
|
||||
zeta.config_version = 0;
|
||||
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![zeta, alpha]));
|
||||
let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository))
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
|
||||
let first_wave = start_proxy_upgrade_rollout(&data, "2.0.0".to_string(), 1, 0, None)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
assert_eq!(first_wave.updated, 1);
|
||||
assert_eq!(first_wave.node_ids, vec!["node-alpha".to_string()]);
|
||||
assert!(first_wave.rollout_active);
|
||||
|
||||
let alpha_after_first = repository
|
||||
.find_proxy_node("node-alpha")
|
||||
.await
|
||||
.expect("lookup should succeed")
|
||||
.expect("alpha should exist");
|
||||
assert_eq!(
|
||||
alpha_after_first
|
||||
.remote_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("upgrade_to")),
|
||||
Some(&json!("2.0.0"))
|
||||
);
|
||||
|
||||
repository
|
||||
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: "node-alpha".to_string(),
|
||||
heartbeat_interval: None,
|
||||
active_connections: Some(2),
|
||||
total_requests_delta: Some(1),
|
||||
avg_latency_ms: Some(2.0),
|
||||
failed_requests_delta: Some(0),
|
||||
dns_failures_delta: Some(0),
|
||||
stream_errors_delta: Some(0),
|
||||
proxy_metadata: Some(json!({"version": "2.0.0"})),
|
||||
proxy_version: Some("2.0.0".to_string()),
|
||||
})
|
||||
.await
|
||||
.expect("heartbeat should succeed");
|
||||
|
||||
let observed = advance_proxy_upgrade_rollout_once(&data)
|
||||
.await
|
||||
.expect("rollout should observe confirmed version");
|
||||
assert_eq!(observed.updated, 0);
|
||||
assert!(observed.blocked);
|
||||
assert_eq!(observed.pending_node_ids, vec!["node-alpha".to_string()]);
|
||||
|
||||
assert!(!record_proxy_upgrade_traffic_success(&data, "node-zeta")
|
||||
.await
|
||||
.expect("untracked node traffic should be ignored"));
|
||||
assert!(record_proxy_upgrade_traffic_success(&data, "node-alpha")
|
||||
.await
|
||||
.expect("traffic confirmation should be recorded"));
|
||||
|
||||
let second_wave = advance_proxy_upgrade_rollout_once(&data)
|
||||
.await
|
||||
.expect("rollout should advance after a healthy observation cycle");
|
||||
assert_eq!(second_wave.updated, 1);
|
||||
assert_eq!(second_wave.node_ids, vec!["node-zeta".to_string()]);
|
||||
assert!(second_wave.rollout_active);
|
||||
|
||||
repository
|
||||
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: "node-zeta".to_string(),
|
||||
heartbeat_interval: None,
|
||||
active_connections: Some(2),
|
||||
total_requests_delta: Some(1),
|
||||
avg_latency_ms: Some(2.0),
|
||||
failed_requests_delta: Some(0),
|
||||
dns_failures_delta: Some(0),
|
||||
stream_errors_delta: Some(0),
|
||||
proxy_metadata: Some(json!({"version": "2.0.0"})),
|
||||
proxy_version: Some("proxy-v2.0.0".to_string()),
|
||||
})
|
||||
.await
|
||||
.expect("heartbeat should succeed");
|
||||
|
||||
let zeta_observed = advance_proxy_upgrade_rollout_once(&data)
|
||||
.await
|
||||
.expect("rollout should observe second wave");
|
||||
assert_eq!(zeta_observed.updated, 0);
|
||||
assert!(zeta_observed.blocked);
|
||||
assert_eq!(
|
||||
zeta_observed.pending_node_ids,
|
||||
vec!["node-zeta".to_string()]
|
||||
);
|
||||
|
||||
assert!(record_proxy_upgrade_traffic_success(&data, "node-zeta")
|
||||
.await
|
||||
.expect("traffic confirmation should be recorded"));
|
||||
|
||||
let finished = advance_proxy_upgrade_rollout_once(&data)
|
||||
.await
|
||||
.expect("rollout should finish after the second healthy observation cycle");
|
||||
assert!(!finished.rollout_active);
|
||||
assert_eq!(finished.completed, 2);
|
||||
assert_eq!(finished.remaining, 0);
|
||||
assert!(data
|
||||
.list_system_config_entries()
|
||||
.await
|
||||
.expect("system config list should succeed")
|
||||
.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn proxy_upgrade_rollout_excludes_draining_nodes_from_online_eligible_pool() {
|
||||
let mut alpha = sample_connected_proxy_node("node-alpha", 30, 1_800_000_000);
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.proxy_metadata = Some(json!({"version": "1.0.0"}));
|
||||
alpha.remote_config = Some(json!({"scheduling_state": "draining"}));
|
||||
alpha.config_version = 1;
|
||||
|
||||
let mut zeta = sample_connected_proxy_node("node-zeta", 30, 1_800_000_000);
|
||||
zeta.name = "zeta".to_string();
|
||||
zeta.proxy_metadata = Some(json!({"version": "1.0.0"}));
|
||||
zeta.remote_config = None;
|
||||
zeta.config_version = 0;
|
||||
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![zeta, alpha]));
|
||||
let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository))
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
|
||||
let rollout = start_proxy_upgrade_rollout(&data, "2.0.0".to_string(), 2, 0, None)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
assert_eq!(rollout.updated, 1);
|
||||
assert_eq!(rollout.node_ids, vec!["node-zeta".to_string()]);
|
||||
|
||||
let alpha_after = repository
|
||||
.find_proxy_node("node-alpha")
|
||||
.await
|
||||
.expect("lookup should succeed")
|
||||
.expect("alpha should exist");
|
||||
assert_eq!(
|
||||
alpha_after
|
||||
.remote_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("upgrade_to")),
|
||||
None
|
||||
);
|
||||
|
||||
let rollout_status = inspect_proxy_upgrade_rollout(&data)
|
||||
.await
|
||||
.expect("inspect should succeed")
|
||||
.expect("rollout should exist");
|
||||
assert_eq!(rollout_status.online_eligible_total, 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn proxy_upgrade_rollout_blocks_next_wave_after_post_upgrade_transport_errors() {
|
||||
let mut alpha = sample_connected_proxy_node("node-alpha", 30, 1_800_000_000);
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.proxy_metadata = Some(json!({"version": "1.0.0"}));
|
||||
alpha.remote_config = None;
|
||||
alpha.config_version = 0;
|
||||
|
||||
let mut zeta = sample_connected_proxy_node("node-zeta", 30, 1_800_000_000);
|
||||
zeta.name = "zeta".to_string();
|
||||
zeta.proxy_metadata = Some(json!({"version": "1.0.0"}));
|
||||
zeta.remote_config = None;
|
||||
zeta.config_version = 0;
|
||||
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![zeta, alpha]));
|
||||
let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository))
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
|
||||
let first_wave = start_proxy_upgrade_rollout(&data, "2.0.0".to_string(), 1, 0, None)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
assert_eq!(first_wave.updated, 1);
|
||||
assert_eq!(first_wave.node_ids, vec!["node-alpha".to_string()]);
|
||||
|
||||
repository
|
||||
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: "node-alpha".to_string(),
|
||||
heartbeat_interval: None,
|
||||
active_connections: Some(2),
|
||||
total_requests_delta: Some(1),
|
||||
avg_latency_ms: Some(2.0),
|
||||
failed_requests_delta: Some(0),
|
||||
dns_failures_delta: Some(0),
|
||||
stream_errors_delta: Some(0),
|
||||
proxy_metadata: Some(json!({"version": "2.0.0"})),
|
||||
proxy_version: Some("2.0.0".to_string()),
|
||||
})
|
||||
.await
|
||||
.expect("heartbeat should succeed");
|
||||
|
||||
let observed = advance_proxy_upgrade_rollout_once(&data)
|
||||
.await
|
||||
.expect("rollout should observe the first upgraded node");
|
||||
assert!(observed.blocked);
|
||||
assert_eq!(observed.pending_node_ids, vec!["node-alpha".to_string()]);
|
||||
|
||||
assert!(record_proxy_upgrade_traffic_success(&data, "node-alpha")
|
||||
.await
|
||||
.expect("traffic confirmation should be recorded"));
|
||||
|
||||
tokio::time::sleep(Duration::from_millis(5)).await;
|
||||
|
||||
repository
|
||||
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: "node-alpha".to_string(),
|
||||
heartbeat_interval: None,
|
||||
active_connections: Some(2),
|
||||
total_requests_delta: Some(1),
|
||||
avg_latency_ms: Some(2.0),
|
||||
failed_requests_delta: Some(1),
|
||||
dns_failures_delta: Some(0),
|
||||
stream_errors_delta: Some(0),
|
||||
proxy_metadata: Some(json!({"version": "2.0.0"})),
|
||||
proxy_version: Some("2.0.0".to_string()),
|
||||
})
|
||||
.await
|
||||
.expect("heartbeat should succeed");
|
||||
|
||||
let blocked = advance_proxy_upgrade_rollout_once(&data)
|
||||
.await
|
||||
.expect("rollout should stay blocked after post-upgrade transport errors");
|
||||
assert_eq!(blocked.updated, 0);
|
||||
assert!(blocked.blocked);
|
||||
assert_eq!(blocked.pending_node_ids, vec!["node-alpha".to_string()]);
|
||||
|
||||
let zeta_after = repository
|
||||
.find_proxy_node("node-zeta")
|
||||
.await
|
||||
.expect("lookup should succeed")
|
||||
.expect("zeta should exist");
|
||||
assert!(zeta_after.remote_config.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn proxy_upgrade_rollout_active_probe_advances_next_wave_after_version_confirmation() {
|
||||
let mut alpha = sample_connected_proxy_node("node-alpha", 30, 1_800_000_000);
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.proxy_metadata = Some(json!({"version": "1.0.0"}));
|
||||
alpha.remote_config = None;
|
||||
alpha.config_version = 0;
|
||||
|
||||
let mut zeta = sample_connected_proxy_node("node-zeta", 30, 1_800_000_000);
|
||||
zeta.name = "zeta".to_string();
|
||||
zeta.proxy_metadata = Some(json!({"version": "1.0.0"}));
|
||||
zeta.remote_config = None;
|
||||
zeta.config_version = 0;
|
||||
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![zeta, alpha]));
|
||||
let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository))
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
let state = AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(data.clone());
|
||||
|
||||
let first_wave = start_proxy_upgrade_rollout(
|
||||
&data,
|
||||
"2.0.0".to_string(),
|
||||
1,
|
||||
0,
|
||||
Some(ProxyUpgradeRolloutProbeConfig {
|
||||
url: "https://probe.example/health".to_string(),
|
||||
timeout_secs: 5,
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
assert_eq!(first_wave.node_ids, vec!["node-alpha".to_string()]);
|
||||
|
||||
repository
|
||||
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: "node-alpha".to_string(),
|
||||
heartbeat_interval: None,
|
||||
active_connections: Some(2),
|
||||
total_requests_delta: Some(1),
|
||||
avg_latency_ms: Some(2.0),
|
||||
failed_requests_delta: Some(0),
|
||||
dns_failures_delta: Some(0),
|
||||
stream_errors_delta: Some(0),
|
||||
proxy_metadata: Some(json!({"version": "2.0.0"})),
|
||||
proxy_version: Some("2.0.0".to_string()),
|
||||
})
|
||||
.await
|
||||
.expect("heartbeat should succeed");
|
||||
|
||||
let tunnel_state = state.tunnel.app_state();
|
||||
let (proxy_tx, mut proxy_rx) = bounded_queue(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
tunnel_state
|
||||
.hub
|
||||
.register_proxy(Arc::new(crate::tunnel::TunnelProxyConn::new(
|
||||
700,
|
||||
"node-alpha".to_string(),
|
||||
"Node Alpha".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
)));
|
||||
|
||||
let responder_hub = tunnel_state.hub.clone();
|
||||
let responder = tokio::spawn(async move {
|
||||
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 = crate::tunnel::tunnel_protocol::FrameHeader::parse(&request_headers)
|
||||
.expect("probe request headers should parse");
|
||||
assert_eq!(
|
||||
request_header.msg_type,
|
||||
crate::tunnel::tunnel_protocol::REQUEST_HEADERS
|
||||
);
|
||||
|
||||
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 = crate::tunnel::tunnel_protocol::FrameHeader::parse(&request_body)
|
||||
.expect("probe request body should parse");
|
||||
assert_eq!(
|
||||
request_body_header.msg_type,
|
||||
crate::tunnel::tunnel_protocol::REQUEST_BODY
|
||||
);
|
||||
|
||||
let response_meta = crate::tunnel::tunnel_protocol::ResponseMeta {
|
||||
status: 204,
|
||||
headers: vec![],
|
||||
};
|
||||
let response_payload =
|
||||
serde_json::to_vec(&response_meta).expect("response meta should serialize");
|
||||
let mut response_headers_frame = crate::tunnel::tunnel_protocol::encode_frame(
|
||||
request_header.stream_id,
|
||||
crate::tunnel::tunnel_protocol::RESPONSE_HEADERS,
|
||||
0,
|
||||
&response_payload,
|
||||
);
|
||||
responder_hub
|
||||
.handle_proxy_frame(700, &mut response_headers_frame)
|
||||
.await;
|
||||
|
||||
let mut response_end_frame = crate::tunnel::tunnel_protocol::encode_frame(
|
||||
request_header.stream_id,
|
||||
crate::tunnel::tunnel_protocol::STREAM_END,
|
||||
0,
|
||||
&[],
|
||||
);
|
||||
responder_hub
|
||||
.handle_proxy_frame(700, &mut response_end_frame)
|
||||
.await;
|
||||
});
|
||||
|
||||
run_proxy_upgrade_rollout_once(&state)
|
||||
.await
|
||||
.expect("rollout worker should succeed");
|
||||
responder.await.expect("probe responder should complete");
|
||||
|
||||
let zeta_after = repository
|
||||
.find_proxy_node("node-zeta")
|
||||
.await
|
||||
.expect("lookup should succeed")
|
||||
.expect("zeta should exist");
|
||||
assert_eq!(
|
||||
zeta_after
|
||||
.remote_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("upgrade_to")),
|
||||
Some(&json!("2.0.0"))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn spawn_stats_aggregation_worker_skips_when_postgres_unavailable() {
|
||||
assert!(spawn_stats_aggregation_worker(Arc::new(GatewayDataState::disabled())).is_none());
|
||||
|
||||
@@ -11,13 +11,14 @@ use super::{
|
||||
duration_until_next_stats_aggregation_run, duration_until_next_stats_hourly_aggregation_run,
|
||||
maintenance_timezone, parse_hhmm_time, provider_checkin_schedule, run_audit_cleanup_once,
|
||||
run_db_maintenance_once, run_gemini_file_mapping_cleanup_once, run_pending_cleanup_once,
|
||||
run_pool_monitor_once, run_provider_checkin_once, run_request_candidate_cleanup_once,
|
||||
run_stats_aggregation_once, run_stats_hourly_aggregation_once, run_usage_cleanup_once,
|
||||
run_pool_monitor_once, run_provider_checkin_once, run_proxy_node_stale_cleanup_once,
|
||||
run_proxy_upgrade_rollout_once, run_request_candidate_cleanup_once, run_stats_aggregation_once,
|
||||
run_stats_hourly_aggregation_once, run_usage_cleanup_once,
|
||||
run_wallet_daily_usage_aggregation_once, AUDIT_LOG_CLEANUP_INTERVAL,
|
||||
GEMINI_FILE_MAPPING_CLEANUP_INTERVAL, PENDING_CLEANUP_INTERVAL, POOL_MONITOR_INTERVAL,
|
||||
PROVIDER_CHECKIN_DEFAULT_TIME, REQUEST_CANDIDATE_CLEANUP_INTERVAL, USAGE_CLEANUP_HOUR,
|
||||
USAGE_CLEANUP_MINUTE, WALLET_DAILY_USAGE_AGGREGATION_HOUR,
|
||||
WALLET_DAILY_USAGE_AGGREGATION_MINUTE,
|
||||
PROVIDER_CHECKIN_DEFAULT_TIME, PROXY_NODE_STALE_SWEEP_INTERVAL, PROXY_UPGRADE_ROLLOUT_INTERVAL,
|
||||
REQUEST_CANDIDATE_CLEANUP_INTERVAL, USAGE_CLEANUP_HOUR, USAGE_CLEANUP_MINUTE,
|
||||
WALLET_DAILY_USAGE_AGGREGATION_HOUR, WALLET_DAILY_USAGE_AGGREGATION_MINUTE,
|
||||
};
|
||||
|
||||
fn log_maintenance_worker_failure(
|
||||
@@ -227,6 +228,55 @@ pub(crate) fn spawn_pending_cleanup_worker(
|
||||
}))
|
||||
}
|
||||
|
||||
pub(crate) fn spawn_proxy_node_stale_cleanup_worker(
|
||||
data: Arc<GatewayDataState>,
|
||||
) -> Option<tokio::task::JoinHandle<()>> {
|
||||
if !data.has_proxy_node_reader() || !data.has_proxy_node_writer() {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(tokio::spawn(async move {
|
||||
if let Err(err) = run_proxy_node_stale_cleanup_once(&data).await {
|
||||
log_maintenance_worker_failure("proxy_node_stale_cleanup", "startup", &err);
|
||||
}
|
||||
let mut interval = tokio::time::interval(PROXY_NODE_STALE_SWEEP_INTERVAL);
|
||||
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
|
||||
interval.tick().await;
|
||||
loop {
|
||||
interval.tick().await;
|
||||
if let Err(err) = run_proxy_node_stale_cleanup_once(&data).await {
|
||||
log_maintenance_worker_failure("proxy_node_stale_cleanup", "tick", &err);
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
|
||||
pub(crate) fn spawn_proxy_upgrade_rollout_worker(
|
||||
state: AppState,
|
||||
) -> Option<tokio::task::JoinHandle<()>> {
|
||||
if !state.data.has_proxy_node_reader()
|
||||
|| !state.data.has_proxy_node_writer()
|
||||
|| !state.data.has_system_config_store()
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(tokio::spawn(async move {
|
||||
if let Err(err) = run_proxy_upgrade_rollout_once(&state).await {
|
||||
log_maintenance_worker_failure("proxy_upgrade_rollout", "startup", &err);
|
||||
}
|
||||
let mut interval = tokio::time::interval(PROXY_UPGRADE_ROLLOUT_INTERVAL);
|
||||
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
|
||||
interval.tick().await;
|
||||
loop {
|
||||
interval.tick().await;
|
||||
if let Err(err) = run_proxy_upgrade_rollout_once(&state).await {
|
||||
log_maintenance_worker_failure("proxy_upgrade_rollout", "tick", &err);
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
|
||||
pub(crate) fn spawn_pool_monitor_worker(
|
||||
data: Arc<GatewayDataState>,
|
||||
) -> Option<tokio::task::JoinHandle<()>> {
|
||||
|
||||
@@ -86,6 +86,19 @@ impl AppState {
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn get_management_token_with_user_by_hash(
|
||||
&self,
|
||||
token_hash: &str,
|
||||
) -> Result<
|
||||
Option<aether_data::repository::management_tokens::StoredManagementTokenWithUser>,
|
||||
GatewayError,
|
||||
> {
|
||||
self.data
|
||||
.get_management_token_with_user_by_hash(token_hash)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn create_management_token(
|
||||
&self,
|
||||
record: &aether_data::repository::management_tokens::CreateManagementTokenRecord,
|
||||
@@ -122,6 +135,20 @@ impl AppState {
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn record_management_token_usage(
|
||||
&self,
|
||||
token_id: &str,
|
||||
last_used_ip: Option<&str>,
|
||||
) -> Result<
|
||||
Option<aether_data::repository::management_tokens::StoredManagementToken>,
|
||||
GatewayError,
|
||||
> {
|
||||
self.data
|
||||
.record_management_token_usage(token_id, last_used_ip)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn set_management_token_active(
|
||||
&self,
|
||||
token_id: &str,
|
||||
|
||||
@@ -40,6 +40,8 @@ use crate::maintenance::spawn_gemini_file_mapping_cleanup_worker;
|
||||
use crate::maintenance::spawn_pending_cleanup_worker;
|
||||
use crate::maintenance::spawn_pool_monitor_worker;
|
||||
use crate::maintenance::spawn_provider_checkin_worker;
|
||||
use crate::maintenance::spawn_proxy_node_stale_cleanup_worker;
|
||||
use crate::maintenance::spawn_proxy_upgrade_rollout_worker;
|
||||
use crate::maintenance::spawn_request_candidate_cleanup_worker;
|
||||
use crate::maintenance::spawn_stats_aggregation_worker;
|
||||
use crate::maintenance::spawn_stats_hourly_aggregation_worker;
|
||||
@@ -296,6 +298,10 @@ impl AppState {
|
||||
self.data.has_proxy_node_reader()
|
||||
}
|
||||
|
||||
pub(crate) fn has_proxy_node_writer(&self) -> bool {
|
||||
self.data.has_proxy_node_writer()
|
||||
}
|
||||
|
||||
pub(crate) fn frontdoor_cors(&self) -> Option<Arc<FrontdoorCorsConfig>> {
|
||||
self.frontdoor_cors.clone()
|
||||
}
|
||||
@@ -430,6 +436,23 @@ impl AppState {
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn register_proxy_node(
|
||||
&self,
|
||||
mutation: &aether_data::repository::proxy_nodes::ProxyNodeRegistrationMutation,
|
||||
) -> Result<Option<StoredProxyNode>, GatewayError> {
|
||||
self.data
|
||||
.register_proxy_node(mutation)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub async fn reset_stale_proxy_node_tunnel_statuses(&self) -> std::io::Result<usize> {
|
||||
self.data
|
||||
.reset_stale_proxy_node_tunnel_statuses()
|
||||
.await
|
||||
.map_err(|err| std::io::Error::other(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn apply_proxy_node_heartbeat(
|
||||
&self,
|
||||
mutation: &ProxyNodeHeartbeatMutation,
|
||||
@@ -440,6 +463,26 @@ impl AppState {
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn unregister_proxy_node(
|
||||
&self,
|
||||
node_id: &str,
|
||||
) -> Result<Option<StoredProxyNode>, GatewayError> {
|
||||
self.data
|
||||
.unregister_proxy_node(node_id)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn update_proxy_node_remote_config(
|
||||
&self,
|
||||
mutation: &aether_data::repository::proxy_nodes::ProxyNodeRemoteConfigMutation,
|
||||
) -> Result<Option<StoredProxyNode>, GatewayError> {
|
||||
self.data
|
||||
.update_proxy_node_remote_config(mutation)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn update_proxy_node_tunnel_status(
|
||||
&self,
|
||||
mutation: &ProxyNodeTunnelStatusMutation,
|
||||
@@ -668,6 +711,12 @@ impl AppState {
|
||||
if let Some(handle) = spawn_pending_cleanup_worker(self.data.clone()) {
|
||||
tasks.push(handle);
|
||||
}
|
||||
if let Some(handle) = spawn_proxy_node_stale_cleanup_worker(self.data.clone()) {
|
||||
tasks.push(handle);
|
||||
}
|
||||
if let Some(handle) = spawn_proxy_upgrade_rollout_worker(self.clone()) {
|
||||
tasks.push(handle);
|
||||
}
|
||||
if let Some(handle) = spawn_provider_checkin_worker(self.clone()) {
|
||||
tasks.push(handle);
|
||||
}
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use aether_contracts::{ExecutionPlan, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER};
|
||||
use aether_crypto::{
|
||||
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY,
|
||||
};
|
||||
@@ -10,6 +11,7 @@ use aether_data::repository::oauth_providers::{
|
||||
InMemoryOAuthProviderRepository, OAuthProviderReadRepository,
|
||||
};
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository;
|
||||
use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository;
|
||||
use axum::body::{to_bytes, Body, Bytes};
|
||||
use axum::response::{IntoResponse, Response};
|
||||
@@ -19,8 +21,9 @@ use http::{HeaderMap, HeaderValue, StatusCode};
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::{
|
||||
build_router_with_state, sample_endpoint, sample_key, sample_management_token,
|
||||
sample_oauth_provider_config, sample_provider, start_server, AppState,
|
||||
build_router_with_state, build_state_with_execution_runtime_override, sample_endpoint,
|
||||
sample_key, sample_management_token, sample_oauth_provider_config, sample_provider,
|
||||
sample_proxy_node, start_server, AppState,
|
||||
};
|
||||
use crate::admin_api::{
|
||||
maybe_build_local_admin_provider_oauth_response, AdminAppState, AdminRequestContext,
|
||||
@@ -549,6 +552,152 @@ async fn local_admin_provider_oauth_device_poll_attaches_audit_only_when_transit
|
||||
token_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_admin_provider_oauth_device_authorize_via_execution_runtime_proxy_node() {
|
||||
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-kiro"));
|
||||
assert_eq!(
|
||||
plan.headers.get("host").map(String::as_str),
|
||||
Some("oidc.us-east-1.amazonaws.com")
|
||||
);
|
||||
assert_eq!(
|
||||
plan.headers
|
||||
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)
|
||||
.map(String::as_str),
|
||||
Some("true")
|
||||
);
|
||||
if plan.request_id == "kiro_device_register" {
|
||||
assert_eq!(plan.url, "https://oidc.us-east-1.amazonaws.com/client/register");
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"body": {
|
||||
"json_body": {
|
||||
"clientId": "kiro-device-client",
|
||||
"clientSecret": "kiro-device-secret"
|
||||
}
|
||||
}
|
||||
}))
|
||||
} else {
|
||||
assert_eq!(plan.request_id, "kiro_device_authorize");
|
||||
assert_eq!(
|
||||
plan.url,
|
||||
"https://oidc.us-east-1.amazonaws.com/device_authorization"
|
||||
);
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"body": {
|
||||
"json_body": {
|
||||
"deviceCode": "device-code-123",
|
||||
"userCode": "USER-CODE",
|
||||
"verificationUri": "https://device.example.com/verify",
|
||||
"verificationUriComplete": "https://device.example.com/verify?user_code=USER-CODE",
|
||||
"expiresIn": 600,
|
||||
"interval": 5
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let mut provider = sample_provider("provider-kiro", "kiro", 10);
|
||||
provider.provider_type = "kiro".to_string();
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![],
|
||||
vec![],
|
||||
));
|
||||
let mut manual_node = sample_proxy_node("proxy-node-kiro");
|
||||
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 (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
|
||||
let state = build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_reader_for_tests(provider_catalog_repository)
|
||||
.attach_proxy_node_repository_for_tests(proxy_node_repository),
|
||||
)
|
||||
.with_provider_oauth_device_session_entry_for_tests(
|
||||
"seed-session",
|
||||
json!({"status":"seed"}),
|
||||
)
|
||||
.with_provider_oauth_token_url_for_tests(
|
||||
"kiro_device_register",
|
||||
"https://oidc.us-east-1.amazonaws.com/client/register",
|
||||
)
|
||||
.with_provider_oauth_token_url_for_tests(
|
||||
"kiro_device_authorize",
|
||||
"https://oidc.us-east-1.amazonaws.com/device_authorization",
|
||||
);
|
||||
let gateway = build_router_with_state(state.clone());
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/provider-oauth/providers/provider-kiro/device-authorize"
|
||||
))
|
||||
.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!({
|
||||
"start_url": "https://view.awsapps.com/start",
|
||||
"region": "us-east-1",
|
||||
"proxy_node_id": "proxy-node-kiro",
|
||||
}))
|
||||
.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}");
|
||||
let session_id = payload["session_id"]
|
||||
.as_str()
|
||||
.expect("session_id should exist")
|
||||
.to_string();
|
||||
assert_eq!(payload["user_code"], "USER-CODE");
|
||||
assert_eq!(payload["expires_in"], 600);
|
||||
assert_eq!(payload["interval"], 5);
|
||||
|
||||
let stored = state
|
||||
.load_provider_oauth_device_session_for_tests(&format!("device_auth_session:{session_id}"))
|
||||
.expect("device session should be stored");
|
||||
let stored: serde_json::Value =
|
||||
serde_json::from_str(&stored).expect("device session json should parse");
|
||||
assert_eq!(stored["proxy_node_id"], "proxy-node-kiro");
|
||||
|
||||
let plans = execution_plans.lock().expect("mutex should lock");
|
||||
assert_eq!(plans.len(), 2);
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_admin_provider_oauth_start_key_locally_with_trusted_admin_principal() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -1655,6 +1804,134 @@ async fn gateway_imports_admin_provider_oauth_refresh_token_locally_with_trusted
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_imports_admin_provider_oauth_refresh_token_via_execution_runtime_proxy_node() {
|
||||
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-import"));
|
||||
assert_eq!(plan.request_id, "provider-oauth:refresh-token");
|
||||
assert_eq!(plan.method, "POST");
|
||||
assert_eq!(plan.url, "https://oauth.example/oauth/token");
|
||||
assert_eq!(
|
||||
plan.headers.get("content-type").map(String::as_str),
|
||||
Some("application/x-www-form-urlencoded")
|
||||
);
|
||||
assert_eq!(
|
||||
plan.headers
|
||||
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)
|
||||
.map(String::as_str),
|
||||
Some("true")
|
||||
);
|
||||
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();
|
||||
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 manual_node = sample_proxy_node("proxy-node-codex-import");
|
||||
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 (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_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",
|
||||
"proxy_node_id": "proxy-node-codex-import",
|
||||
"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");
|
||||
assert_eq!(payload["has_refresh_token"], true);
|
||||
assert_eq!(payload["email"], "alice@example.com");
|
||||
assert_eq!(payload["replaced"], false);
|
||||
|
||||
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,
|
||||
Some(json!({"node_id":"proxy-node-codex-import","enabled":true}))
|
||||
);
|
||||
|
||||
let plans = execution_plans.lock().expect("mutex should lock");
|
||||
assert_eq!(plans.len(), 1);
|
||||
|
||||
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));
|
||||
@@ -1797,6 +2074,167 @@ async fn gateway_batch_imports_admin_provider_oauth_kiro_locally_with_trusted_ad
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_batch_imports_admin_provider_oauth_kiro_via_execution_runtime_proxy_node() {
|
||||
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());
|
||||
assert_eq!(plan.request_id, "kiro_batch_refresh:social");
|
||||
assert_eq!(plan.url, "https://oauth.example/refreshToken");
|
||||
assert_eq!(
|
||||
plan.proxy
|
||||
.as_ref()
|
||||
.and_then(|proxy| proxy.node_id.as_deref()),
|
||||
Some("proxy-node-kiro-batch-runtime")
|
||||
);
|
||||
assert_eq!(
|
||||
plan.headers
|
||||
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)
|
||||
.map(String::as_str),
|
||||
Some("true")
|
||||
);
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"body": {
|
||||
"json_body": {
|
||||
"accessToken": sample_kiro_device_access_token("kiro-runtime@example.com"),
|
||||
"refreshToken": "kiro-runtime-refresh-token-new",
|
||||
"expiresIn": 1800,
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let mut provider = sample_provider("provider-kiro", "kiro", 10);
|
||||
provider.provider_type = "kiro".to_string();
|
||||
let endpoint = sample_endpoint(
|
||||
"endpoint-kiro-chat",
|
||||
"provider-kiro",
|
||||
"kiro:generateAssistantResponse",
|
||||
"https://service.kiro.dev",
|
||||
);
|
||||
|
||||
let mut existing_key = sample_key(
|
||||
"key-kiro-batch-runtime",
|
||||
"provider-kiro",
|
||||
"kiro:generateAssistantResponse",
|
||||
"stale-kiro-runtime-access-token",
|
||||
);
|
||||
existing_key.auth_type = "oauth".to_string();
|
||||
existing_key.is_active = false;
|
||||
existing_key.encrypted_auth_config = Some(
|
||||
encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
r#"{"provider_type":"kiro","auth_method":"social","email":"kiro-runtime@example.com","refresh_token":"kiro-runtime-refresh-token-old"}"#,
|
||||
)
|
||||
.expect("auth config ciphertext should build"),
|
||||
);
|
||||
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![endpoint],
|
||||
vec![existing_key],
|
||||
));
|
||||
let mut manual_node = sample_proxy_node("proxy-node-kiro-batch-runtime");
|
||||
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 (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_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
)
|
||||
.with_provider_oauth_token_url_for_tests(
|
||||
"kiro_social_refresh",
|
||||
"https://oauth.example",
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/provider-oauth/providers/provider-kiro/batch-import"
|
||||
))
|
||||
.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!({
|
||||
"credentials": "kiro-runtime-refresh-token-old",
|
||||
"proxy_node_id": "proxy-node-kiro-batch-runtime"
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("payload should parse");
|
||||
assert_eq!(payload["total"], 1);
|
||||
assert_eq!(payload["success"], 1);
|
||||
assert_eq!(payload["failed"], 0);
|
||||
assert_eq!(payload["results"][0]["status"], "success");
|
||||
assert_eq!(payload["results"][0]["key_id"], "key-kiro-batch-runtime");
|
||||
assert_eq!(payload["results"][0]["replaced"], true);
|
||||
|
||||
let plans = execution_plans.lock().expect("mutex should lock");
|
||||
assert_eq!(plans.len(), 1);
|
||||
drop(plans);
|
||||
|
||||
let stored_key = provider_catalog_repository
|
||||
.list_keys_by_ids(&["key-kiro-batch-runtime".to_string()])
|
||||
.await
|
||||
.expect("keys should load")
|
||||
.into_iter()
|
||||
.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}))
|
||||
);
|
||||
let decrypted_auth_config = decrypt_python_fernet_ciphertext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
stored_key
|
||||
.encrypted_auth_config
|
||||
.as_deref()
|
||||
.expect("auth config should exist"),
|
||||
)
|
||||
.expect("auth config should decrypt");
|
||||
let auth_config: serde_json::Value =
|
||||
serde_json::from_str(&decrypted_auth_config).expect("auth config should parse");
|
||||
assert_eq!(auth_config["email"], "kiro-runtime@example.com");
|
||||
assert_eq!(
|
||||
auth_config["refresh_token"],
|
||||
"kiro-runtime-refresh-token-new"
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_starts_admin_provider_oauth_kiro_batch_import_task_locally_with_trusted_admin_principal(
|
||||
) {
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use aether_contracts::ExecutionPlan;
|
||||
use aether_contracts::{
|
||||
ExecutionPlan, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
|
||||
};
|
||||
use aether_crypto::{
|
||||
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY,
|
||||
};
|
||||
@@ -466,7 +468,7 @@ async fn gateway_saves_admin_provider_ops_config_locally_with_trusted_admin_prin
|
||||
|
||||
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_eq!(payload["success"], true, "payload={payload}");
|
||||
assert_eq!(payload["message"], "配置保存成功");
|
||||
|
||||
let stored_provider = provider_catalog_repository
|
||||
@@ -1142,6 +1144,12 @@ async fn gateway_verifies_admin_provider_ops_locally_for_anyrouter_proxy_mode()
|
||||
assert_eq!(parsed_proxy.password(), Some("supersecret"));
|
||||
|
||||
if plan.url.ends_with("/api/user/self") {
|
||||
assert_eq!(
|
||||
plan.headers
|
||||
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)
|
||||
.map(String::as_str),
|
||||
None
|
||||
);
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"status_code": 200,
|
||||
@@ -1160,6 +1168,14 @@ async fn gateway_verifies_admin_provider_ops_locally_for_anyrouter_proxy_mode()
|
||||
}
|
||||
}))
|
||||
} else {
|
||||
assert_eq!(plan.request_id, "provider-ops-acw:anyrouter");
|
||||
assert_eq!(plan.url, "https://ops.example");
|
||||
assert_eq!(
|
||||
plan.headers
|
||||
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)
|
||||
.map(String::as_str),
|
||||
Some("false")
|
||||
);
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"status_code": 200,
|
||||
@@ -1242,9 +1258,159 @@ async fn gateway_verifies_admin_provider_ops_locally_for_anyrouter_proxy_mode()
|
||||
assert_eq!(payload["data"]["request_count"], 8);
|
||||
|
||||
let plans = execution_plans.lock().expect("mutex should lock");
|
||||
assert_eq!(plans.len(), 1);
|
||||
assert_eq!(plans[0].request_id, "provider-ops-verify:anyrouter");
|
||||
assert_eq!(plans[0].url, "https://ops.example/api/user/self");
|
||||
assert_eq!(plans.len(), 2);
|
||||
assert_eq!(plans[0].request_id, "provider-ops-acw:anyrouter");
|
||||
assert_eq!(plans[0].url, "https://ops.example");
|
||||
assert_eq!(plans[1].request_id, "provider-ops-verify:anyrouter");
|
||||
assert_eq!(plans[1].url, "https://ops.example/api/user/self");
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_verifies_admin_provider_ops_sub2api_proxy_mode_via_execution_runtime_http1_only() {
|
||||
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-sub2api"));
|
||||
assert_eq!(
|
||||
plan.headers
|
||||
.get(EXECUTION_REQUEST_HTTP1_ONLY_HEADER)
|
||||
.map(String::as_str),
|
||||
Some("true")
|
||||
);
|
||||
|
||||
if plan.url.ends_with("/api/v1/auth/refresh") {
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"body": {
|
||||
"json_body": {
|
||||
"code": 0,
|
||||
"data": {
|
||||
"access_token": "sub2api-access-token",
|
||||
"refresh_token": "sub2api-refresh-token-new"
|
||||
}
|
||||
}
|
||||
}
|
||||
}))
|
||||
} else {
|
||||
assert_eq!(
|
||||
plan.headers.get("authorization").map(String::as_str),
|
||||
Some("Bearer sub2api-access-token")
|
||||
);
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"body": {
|
||||
"json_body": {
|
||||
"code": 0,
|
||||
"data": {
|
||||
"username": "sub2api-user",
|
||||
"email": "sub2api@example.com",
|
||||
"balance": 8.5,
|
||||
"points": 1.5,
|
||||
"status": "active",
|
||||
"concurrency": 3
|
||||
}
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider("provider-openai", "openai", 10)],
|
||||
vec![],
|
||||
vec![],
|
||||
));
|
||||
let mut manual_node = sample_proxy_node("proxy-node-sub2api");
|
||||
manual_node.name = "sub2api-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());
|
||||
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![manual_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,
|
||||
)
|
||||
.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-sub2api"),
|
||||
)]),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/provider-ops/providers/provider-openai/verify"
|
||||
))
|
||||
.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!({
|
||||
"architecture_id": "sub2api",
|
||||
"base_url": "https://sub2api.example",
|
||||
"connector": {
|
||||
"auth_type": "session_login",
|
||||
"config": {},
|
||||
"credentials": {
|
||||
"refresh_token": "refresh-token-old",
|
||||
}
|
||||
},
|
||||
"actions": {},
|
||||
"schedule": {},
|
||||
}))
|
||||
.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["success"], true);
|
||||
assert_eq!(payload["data"]["username"], "sub2api-user");
|
||||
assert_eq!(
|
||||
payload["updated_credentials"]["refresh_token"],
|
||||
"sub2api-refresh-token-new"
|
||||
);
|
||||
|
||||
let plans = execution_plans.lock().expect("mutex should lock");
|
||||
assert_eq!(plans.len(), 2);
|
||||
assert_eq!(plans[0].url, "https://sub2api.example/api/v1/auth/refresh");
|
||||
assert!(
|
||||
plans[1]
|
||||
.url
|
||||
.starts_with("https://sub2api.example/api/v1/auth/me"),
|
||||
"url={}",
|
||||
plans[1].url
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
@@ -1516,6 +1682,109 @@ async fn gateway_verifies_admin_provider_ops_locally_for_new_api_proxy_node_mode
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_verifies_admin_provider_ops_locally_for_new_api_without_proxy_via_execution_runtime(
|
||||
) {
|
||||
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());
|
||||
assert_eq!(plan.request_id, "provider-ops-verify:new_api");
|
||||
assert_eq!(plan.url, "https://ops.example/api/user/self");
|
||||
assert_eq!(
|
||||
plan.headers.get("authorization").map(String::as_str),
|
||||
Some("Bearer live-secret-api-key")
|
||||
);
|
||||
assert_eq!(
|
||||
plan.headers.get("new-api-user").map(String::as_str),
|
||||
Some("42")
|
||||
);
|
||||
assert!(plan.proxy.is_none());
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"body": {
|
||||
"json_body": {
|
||||
"success": true,
|
||||
"data": {
|
||||
"username": "alice",
|
||||
"display_name": "Alice",
|
||||
"quota": 42.5,
|
||||
"used_quota": 12.5,
|
||||
"request_count": 9
|
||||
}
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider("provider-openai", "openai", 10)],
|
||||
vec![],
|
||||
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,
|
||||
),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/provider-ops/providers/provider-openai/verify"
|
||||
))
|
||||
.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!({
|
||||
"architecture_id": "new_api",
|
||||
"base_url": "https://ops.example",
|
||||
"connector": {
|
||||
"auth_type": "api_key",
|
||||
"config": {},
|
||||
"credentials": {
|
||||
"api_key": "live-secret-api-key",
|
||||
"user_id": "42",
|
||||
"cookie": "session=foo"
|
||||
}
|
||||
},
|
||||
"actions": {},
|
||||
"schedule": {},
|
||||
}))
|
||||
.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["success"], true);
|
||||
assert_eq!(payload["data"]["username"], "alice");
|
||||
assert_eq!(payload["data"]["used_quota"], 12.5);
|
||||
assert_eq!(payload["data"]["request_count"], 9);
|
||||
assert_eq!(execution_plans.lock().expect("mutex should lock").len(), 1);
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_verifies_admin_provider_ops_locally_for_sub2api_with_trusted_admin_principal() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -2446,6 +2715,145 @@ async fn gateway_handles_admin_provider_ops_balance_locally_for_generic_api_prox
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_admin_provider_ops_balance_locally_without_proxy_via_execution_runtime() {
|
||||
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());
|
||||
assert!(plan.proxy.is_none());
|
||||
assert_eq!(
|
||||
plan.headers.get("authorization").map(String::as_str),
|
||||
Some("Bearer live-secret-api-key")
|
||||
);
|
||||
|
||||
if plan.url.ends_with("/api/user/checkin") {
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"body": {
|
||||
"json_body": {
|
||||
"success": true,
|
||||
"message": "执行层签到成功"
|
||||
}
|
||||
}
|
||||
}))
|
||||
} else {
|
||||
assert!(plan.url.ends_with("/api/user/balance"));
|
||||
Json(json!({
|
||||
"request_id": plan.request_id,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"body": {
|
||||
"json_body": {
|
||||
"success": true,
|
||||
"data": {
|
||||
"quota": 2500000,
|
||||
"used_quota": 500000
|
||||
}
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![
|
||||
sample_provider("provider-openai", "openai", 10).with_transport_fields(
|
||||
true,
|
||||
false,
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
Some(json!({
|
||||
"provider_ops": {
|
||||
"architecture_id": "generic_api",
|
||||
"base_url": "https://ops.example",
|
||||
"connector": {
|
||||
"auth_type": "api_key",
|
||||
"config": {
|
||||
"auth_method": "bearer"
|
||||
},
|
||||
"credentials": {
|
||||
"api_key": encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
"live-secret-api-key",
|
||||
).expect("api key should encrypt"),
|
||||
}
|
||||
}
|
||||
}
|
||||
})),
|
||||
),
|
||||
],
|
||||
vec![],
|
||||
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,
|
||||
),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/provider-ops/providers/provider-openai/balance?refresh=false"
|
||||
))
|
||||
.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 payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["status"], "success");
|
||||
assert_eq!(payload["data"]["total_available"], 5.0);
|
||||
assert_eq!(payload["data"]["total_used"], 1.0);
|
||||
assert_eq!(payload["data"]["extra"]["checkin_success"], true);
|
||||
assert_eq!(
|
||||
payload["data"]["extra"]["checkin_message"],
|
||||
"执行层签到成功"
|
||||
);
|
||||
|
||||
let plans = execution_plans.lock().expect("mutex should lock");
|
||||
assert_eq!(plans.len(), 2);
|
||||
assert!(plans.iter().all(|plan| plan.proxy.is_none()));
|
||||
assert!(plans
|
||||
.iter()
|
||||
.any(|plan| plan.request_id == "provider-ops-action:probe_checkin"));
|
||||
assert!(plans.iter().any(|plan| {
|
||||
plan.request_id == "provider-ops-action:generic_api:query_balance:provider-openai"
|
||||
}));
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_admin_provider_ops_checkin_locally_with_trusted_admin_principal() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
|
||||
@@ -1,18 +1,28 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNodeEvent};
|
||||
use aether_data::repository::management_tokens::InMemoryManagementTokenRepository;
|
||||
use aether_data::repository::proxy_nodes::{
|
||||
InMemoryProxyNodeRepository, ProxyNodeHeartbeatMutation, StoredProxyNodeEvent,
|
||||
};
|
||||
use axum::body::Body;
|
||||
use axum::routing::any;
|
||||
use axum::{extract::Request, Router};
|
||||
use http::StatusCode;
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::{build_router_with_state, sample_proxy_node, start_server, AppState};
|
||||
use super::super::{
|
||||
build_router_with_state, hash_management_token, sample_management_token, sample_proxy_node,
|
||||
start_server, AppState,
|
||||
};
|
||||
use crate::constants::{
|
||||
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
|
||||
TRUSTED_ADMIN_USER_ROLE_HEADER,
|
||||
};
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::maintenance::{
|
||||
record_proxy_upgrade_traffic_success, skip_proxy_upgrade_rollout_node,
|
||||
start_proxy_upgrade_rollout,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_admin_proxy_nodes_locally_with_trusted_admin_principal() {
|
||||
@@ -77,6 +87,7 @@ async fn gateway_handles_admin_proxy_nodes_locally_with_trusted_admin_principal(
|
||||
assert_eq!(payload["total"], 1);
|
||||
assert_eq!(payload["skip"], 0);
|
||||
assert_eq!(payload["limit"], 10);
|
||||
assert!(payload["rollout"].is_null());
|
||||
|
||||
let items = payload["items"].as_array().expect("items should be array");
|
||||
assert_eq!(items.len(), 1);
|
||||
@@ -95,7 +106,502 @@ async fn gateway_handles_admin_proxy_nodes_locally_with_trusted_admin_principal(
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_admin_proxy_nodes_unavailable_routes_locally() {
|
||||
async fn gateway_reports_active_proxy_upgrade_rollout_in_proxy_node_list() {
|
||||
let mut alpha = sample_proxy_node("node-alpha");
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.status = "online".to_string();
|
||||
alpha.tunnel_connected = true;
|
||||
alpha.remote_config = None;
|
||||
alpha.proxy_metadata = Some(json!({ "version": "1.9.0" }));
|
||||
|
||||
let mut beta = sample_proxy_node("node-beta");
|
||||
beta.name = "beta".to_string();
|
||||
beta.status = "online".to_string();
|
||||
beta.tunnel_connected = true;
|
||||
beta.remote_config = None;
|
||||
beta.proxy_metadata = Some(json!({ "version": "1.9.0" }));
|
||||
|
||||
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![beta, alpha]));
|
||||
let data_state = GatewayDataState::with_proxy_node_repository_for_tests(proxy_node_repository)
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
|
||||
let rollout = start_proxy_upgrade_rollout(
|
||||
&data_state,
|
||||
"2.0.0".to_string(),
|
||||
2,
|
||||
120,
|
||||
Some(crate::maintenance::ProxyUpgradeRolloutProbeConfig {
|
||||
url: "https://probe.example/health".to_string(),
|
||||
timeout_secs: 15,
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
assert_eq!(rollout.updated, 2);
|
||||
|
||||
data_state
|
||||
.apply_proxy_node_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: "node-alpha".to_string(),
|
||||
heartbeat_interval: None,
|
||||
active_connections: None,
|
||||
total_requests_delta: None,
|
||||
avg_latency_ms: None,
|
||||
failed_requests_delta: None,
|
||||
dns_failures_delta: None,
|
||||
stream_errors_delta: None,
|
||||
proxy_metadata: None,
|
||||
proxy_version: Some("2.0.0".to_string()),
|
||||
})
|
||||
.await
|
||||
.expect("heartbeat should apply");
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes?skip=0&limit=10"
|
||||
))
|
||||
.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["rollout"]["version"], "2.0.0");
|
||||
assert_eq!(payload["rollout"]["batch_size"], 2);
|
||||
assert_eq!(payload["rollout"]["cooldown_secs"], 120);
|
||||
assert_eq!(
|
||||
payload["rollout"]["probe"]["url"],
|
||||
"https://probe.example/health"
|
||||
);
|
||||
assert_eq!(payload["rollout"]["probe"]["timeout_secs"], 15);
|
||||
assert_eq!(
|
||||
payload["rollout"]["pending_node_ids"],
|
||||
json!(["node-alpha", "node-beta"])
|
||||
);
|
||||
assert_eq!(payload["rollout"]["completed_node_ids"], json!([]));
|
||||
assert_eq!(payload["rollout"]["conflict_node_ids"], json!([]));
|
||||
assert_eq!(payload["rollout"]["blocked"], true);
|
||||
assert!(payload["rollout"]["started_at"].is_string());
|
||||
assert!(payload["rollout"]["last_dispatched_at"].is_string());
|
||||
assert!(payload["rollout"]["updated_at"].is_string());
|
||||
|
||||
let tracked_nodes = payload["rollout"]["tracked_nodes"]
|
||||
.as_array()
|
||||
.expect("tracked_nodes should be array");
|
||||
assert_eq!(tracked_nodes.len(), 2);
|
||||
let alpha_status = tracked_nodes
|
||||
.iter()
|
||||
.find(|tracked| tracked["node_id"] == "node-alpha")
|
||||
.expect("alpha status should exist");
|
||||
assert_eq!(alpha_status["state"], "awaiting_traffic");
|
||||
assert!(alpha_status["version_confirmed_at"].is_string());
|
||||
assert!(alpha_status["traffic_confirmed_at"].is_null());
|
||||
|
||||
let beta_status = tracked_nodes
|
||||
.iter()
|
||||
.find(|tracked| tracked["node_id"] == "node-beta")
|
||||
.expect("beta status should exist");
|
||||
assert_eq!(beta_status["state"], "awaiting_version");
|
||||
assert!(beta_status["version_confirmed_at"].is_null());
|
||||
assert!(beta_status["traffic_confirmed_at"].is_null());
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_cancels_active_proxy_upgrade_rollout_locally() {
|
||||
let mut alpha = sample_proxy_node("node-alpha");
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.status = "online".to_string();
|
||||
alpha.tunnel_connected = true;
|
||||
alpha.remote_config = None;
|
||||
|
||||
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![alpha]));
|
||||
let data_state = GatewayDataState::with_proxy_node_repository_for_tests(proxy_node_repository)
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
|
||||
let rollout = start_proxy_upgrade_rollout(&data_state, "2.0.0".to_string(), 1, 120, None)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
assert!(rollout.rollout_active);
|
||||
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state.clone()),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes/upgrade/cancel"
|
||||
))
|
||||
.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["cancelled"], true);
|
||||
assert_eq!(payload["version"], "2.0.0");
|
||||
assert_eq!(payload["pending_node_ids"], json!(["node-alpha"]));
|
||||
assert_eq!(payload["conflict_node_ids"], json!([]));
|
||||
|
||||
let list_response = reqwest::Client::new()
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes?skip=0&limit=10"
|
||||
))
|
||||
.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");
|
||||
let list_payload: serde_json::Value =
|
||||
list_response.json().await.expect("json body should parse");
|
||||
assert!(list_payload["rollout"].is_null());
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_clears_proxy_upgrade_rollout_conflicts_locally() {
|
||||
let mut alpha = sample_proxy_node("node-alpha");
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.status = "online".to_string();
|
||||
alpha.tunnel_connected = true;
|
||||
alpha.remote_config = None;
|
||||
|
||||
let mut beta = sample_proxy_node("node-beta");
|
||||
beta.name = "beta".to_string();
|
||||
beta.status = "online".to_string();
|
||||
beta.tunnel_connected = true;
|
||||
beta.remote_config = None;
|
||||
|
||||
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![beta, alpha]));
|
||||
let data_state = GatewayDataState::with_proxy_node_repository_for_tests(proxy_node_repository)
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
|
||||
let rollout = start_proxy_upgrade_rollout(&data_state, "2.0.0".to_string(), 1, 120, None)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
assert_eq!(rollout.updated, 1);
|
||||
|
||||
data_state
|
||||
.update_proxy_node_remote_config(
|
||||
&aether_data::repository::proxy_nodes::ProxyNodeRemoteConfigMutation {
|
||||
node_id: "node-beta".to_string(),
|
||||
node_name: None,
|
||||
allowed_ports: None,
|
||||
log_level: None,
|
||||
heartbeat_interval: None,
|
||||
scheduling_state: None,
|
||||
upgrade_to: Some(Some("3.0.0".to_string())),
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("conflict target should update");
|
||||
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state.clone()),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes/upgrade/clear-conflicts"
|
||||
))
|
||||
.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["cleared"], 1);
|
||||
assert_eq!(payload["node_ids"], json!(["node-beta"]));
|
||||
assert_eq!(payload["blocked"], true);
|
||||
assert_eq!(payload["pending_node_ids"], json!(["node-alpha"]));
|
||||
|
||||
let updated_beta = data_state
|
||||
.find_proxy_node("node-beta")
|
||||
.await
|
||||
.expect("node lookup should succeed")
|
||||
.expect("beta should exist");
|
||||
let beta_upgrade_to = updated_beta
|
||||
.remote_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("upgrade_to"))
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
assert!(beta_upgrade_to.is_null());
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_skips_proxy_upgrade_rollout_node_and_advances_next_wave_locally() {
|
||||
let mut alpha = sample_proxy_node("node-alpha");
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.status = "online".to_string();
|
||||
alpha.tunnel_connected = true;
|
||||
alpha.remote_config = None;
|
||||
|
||||
let mut beta = sample_proxy_node("node-beta");
|
||||
beta.name = "beta".to_string();
|
||||
beta.status = "online".to_string();
|
||||
beta.tunnel_connected = true;
|
||||
beta.remote_config = None;
|
||||
|
||||
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![beta, alpha]));
|
||||
let data_state = GatewayDataState::with_proxy_node_repository_for_tests(proxy_node_repository)
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
|
||||
let rollout = start_proxy_upgrade_rollout(&data_state, "2.0.0".to_string(), 1, 120, None)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
assert_eq!(rollout.node_ids, vec!["node-alpha"]);
|
||||
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state.clone()),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes/node-alpha/upgrade/skip"
|
||||
))
|
||||
.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"], "node-alpha");
|
||||
assert_eq!(payload["skipped_node_ids"], json!(["node-alpha"]));
|
||||
assert_eq!(payload["updated"], 1);
|
||||
assert_eq!(payload["pending_node_ids"], json!(["node-beta"]));
|
||||
|
||||
let alpha_after = data_state
|
||||
.find_proxy_node("node-alpha")
|
||||
.await
|
||||
.expect("node lookup should succeed")
|
||||
.expect("alpha should exist");
|
||||
let alpha_upgrade_to = alpha_after
|
||||
.remote_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("upgrade_to"))
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
assert!(alpha_upgrade_to.is_null());
|
||||
|
||||
let list_response = reqwest::Client::new()
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes?skip=0&limit=10"
|
||||
))
|
||||
.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");
|
||||
let list_payload: serde_json::Value =
|
||||
list_response.json().await.expect("json body should parse");
|
||||
assert_eq!(
|
||||
list_payload["rollout"]["skipped_node_ids"],
|
||||
json!(["node-alpha"])
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_retries_proxy_upgrade_rollout_node_locally() {
|
||||
let mut alpha = sample_proxy_node("node-alpha");
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.status = "online".to_string();
|
||||
alpha.tunnel_connected = true;
|
||||
alpha.remote_config = None;
|
||||
|
||||
let mut beta = sample_proxy_node("node-beta");
|
||||
beta.name = "beta".to_string();
|
||||
beta.status = "online".to_string();
|
||||
beta.tunnel_connected = true;
|
||||
beta.remote_config = None;
|
||||
|
||||
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![beta, alpha]));
|
||||
let data_state = GatewayDataState::with_proxy_node_repository_for_tests(proxy_node_repository)
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
|
||||
start_proxy_upgrade_rollout(&data_state, "2.0.0".to_string(), 1, 120, None)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
let _ = skip_proxy_upgrade_rollout_node(&data_state, "node-alpha")
|
||||
.await
|
||||
.expect("skip should succeed");
|
||||
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state.clone()),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes/node-alpha/upgrade/retry"
|
||||
))
|
||||
.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"], "node-alpha");
|
||||
assert_eq!(payload["skipped_node_ids"], json!([]));
|
||||
assert_eq!(payload["blocked"], true);
|
||||
|
||||
let alpha_after = data_state
|
||||
.find_proxy_node("node-alpha")
|
||||
.await
|
||||
.expect("node lookup should succeed")
|
||||
.expect("alpha should exist");
|
||||
let alpha_upgrade_to = alpha_after
|
||||
.remote_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("upgrade_to"))
|
||||
.and_then(serde_json::Value::as_str);
|
||||
assert_eq!(alpha_upgrade_to, Some("2.0.0"));
|
||||
|
||||
let list_response = reqwest::Client::new()
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes?skip=0&limit=10"
|
||||
))
|
||||
.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");
|
||||
let list_payload: serde_json::Value =
|
||||
list_response.json().await.expect("json body should parse");
|
||||
assert_eq!(list_payload["rollout"]["skipped_node_ids"], json!([]));
|
||||
let tracked_nodes = list_payload["rollout"]["tracked_nodes"]
|
||||
.as_array()
|
||||
.expect("tracked nodes should be array");
|
||||
let alpha_status = tracked_nodes
|
||||
.iter()
|
||||
.find(|tracked| tracked["node_id"] == "node-alpha")
|
||||
.expect("alpha should be tracked again");
|
||||
assert_eq!(alpha_status["state"], "awaiting_version");
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_restores_skipped_proxy_upgrade_rollout_nodes_locally() {
|
||||
let mut alpha = sample_proxy_node("node-alpha");
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.status = "online".to_string();
|
||||
alpha.tunnel_connected = true;
|
||||
alpha.remote_config = None;
|
||||
|
||||
let mut beta = sample_proxy_node("node-beta");
|
||||
beta.name = "beta".to_string();
|
||||
beta.status = "online".to_string();
|
||||
beta.tunnel_connected = true;
|
||||
beta.remote_config = None;
|
||||
|
||||
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![beta, alpha]));
|
||||
let data_state = GatewayDataState::with_proxy_node_repository_for_tests(proxy_node_repository)
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
|
||||
start_proxy_upgrade_rollout(&data_state, "2.0.0".to_string(), 1, 120, None)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
let _ = skip_proxy_upgrade_rollout_node(&data_state, "node-alpha")
|
||||
.await
|
||||
.expect("skip should succeed");
|
||||
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state.clone()),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes/upgrade/restore-skipped"
|
||||
))
|
||||
.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["restored"], 1);
|
||||
assert_eq!(payload["node_ids"], json!(["node-alpha"]));
|
||||
assert_eq!(payload["skipped_node_ids"], json!([]));
|
||||
assert_eq!(payload["updated"], 0);
|
||||
assert_eq!(payload["blocked"], true);
|
||||
assert_eq!(payload["pending_node_ids"], json!(["node-beta"]));
|
||||
|
||||
let list_response = reqwest::Client::new()
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes?skip=0&limit=10"
|
||||
))
|
||||
.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");
|
||||
let list_payload: serde_json::Value =
|
||||
list_response.json().await.expect("json body should parse");
|
||||
assert_eq!(list_payload["rollout"]["skipped_node_ids"], json!([]));
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_registers_and_unregisters_proxy_nodes_locally_with_management_token_principal() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
||||
let upstream = Router::new().route(
|
||||
@@ -109,35 +615,100 @@ async fn gateway_rejects_admin_proxy_nodes_unavailable_routes_locally() {
|
||||
}),
|
||||
);
|
||||
|
||||
let raw_token = "ae_proxy_register_test";
|
||||
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::default());
|
||||
let state = AppState::new().expect("gateway should build");
|
||||
let admin_user = state
|
||||
.create_local_auth_user_with_settings(
|
||||
Some("proxy-admin@example.com".to_string()),
|
||||
true,
|
||||
"admin".to_string(),
|
||||
"hash".to_string(),
|
||||
"admin".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("admin user should be created")
|
||||
.expect("admin user should exist");
|
||||
let mut management_token =
|
||||
sample_management_token("token-proxy-register", &admin_user.id, "proxy-admin", true);
|
||||
management_token.token.allowed_ips = None;
|
||||
let management_token_repository =
|
||||
Arc::new(InMemoryManagementTokenRepository::seed_with_hashes(
|
||||
vec![management_token],
|
||||
vec![(
|
||||
hash_management_token(raw_token),
|
||||
"token-proxy-register".to_string(),
|
||||
)],
|
||||
));
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router_with_state(AppState::new().expect("gateway should build"));
|
||||
let state = state.with_data_state_for_tests(
|
||||
GatewayDataState::with_management_token_repository_for_tests(management_token_repository)
|
||||
.attach_proxy_node_repository_for_tests(proxy_node_repository),
|
||||
);
|
||||
let token_lookup = state
|
||||
.get_management_token_with_user_by_hash(&hash_management_token(raw_token))
|
||||
.await
|
||||
.expect("token lookup should succeed");
|
||||
assert!(token_lookup.is_some());
|
||||
let user_lookup = state
|
||||
.find_user_auth_by_id(&admin_user.id)
|
||||
.await
|
||||
.expect("user lookup should succeed");
|
||||
assert!(user_lookup.is_some());
|
||||
let gateway = build_router_with_state(state);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
let register_response = client
|
||||
.post(format!("{gateway_url}/api/admin/proxy-nodes/register"))
|
||||
.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")
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
.bearer_auth(raw_token)
|
||||
.json(&json!({
|
||||
"name": "proxy-1",
|
||||
"ip": "1.1.1.1",
|
||||
"port": 8080
|
||||
"port": 0,
|
||||
"heartbeat_interval": 30,
|
||||
"tunnel_mode": true
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(register_response.status(), StatusCode::SERVICE_UNAVAILABLE);
|
||||
assert_eq!(register_response.status(), StatusCode::OK);
|
||||
let register_payload: serde_json::Value = register_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
let node_id = register_payload["node_id"]
|
||||
.as_str()
|
||||
.expect("node_id should be present")
|
||||
.to_string();
|
||||
assert_eq!(register_payload["node"]["name"], "proxy-1");
|
||||
assert_eq!(register_payload["node"]["status"], "offline");
|
||||
assert_eq!(
|
||||
register_payload["detail"],
|
||||
"Admin proxy nodes data unavailable"
|
||||
register_payload["node"]["registered_by"],
|
||||
json!(admin_user.id)
|
||||
);
|
||||
|
||||
let unregister_response = client
|
||||
.post(format!("{gateway_url}/api/admin/proxy-nodes/unregister"))
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
.bearer_auth(raw_token)
|
||||
.json(&json!({ "node_id": node_id }))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(unregister_response.status(), StatusCode::OK);
|
||||
let unregister_payload: serde_json::Value = unregister_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(unregister_payload["message"], "unregistered");
|
||||
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
@@ -214,3 +785,336 @@ async fn gateway_handles_admin_proxy_node_events_locally_with_trusted_admin_prin
|
||||
gateway_handle.abort();
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_updates_proxy_node_config_and_batches_upgrade_locally() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
||||
let upstream = Router::new()
|
||||
.route(
|
||||
"/api/admin/proxy-nodes/node-online/config",
|
||||
any(move |_request: Request| {
|
||||
let upstream_hits_inner = Arc::clone(&upstream_hits_clone);
|
||||
async move {
|
||||
*upstream_hits_inner.lock().expect("mutex should lock") += 1;
|
||||
(StatusCode::OK, Body::from("unexpected upstream hit"))
|
||||
}
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/admin/proxy-nodes/upgrade",
|
||||
any(|| async move { (StatusCode::OK, Body::from("unexpected upstream hit")) }),
|
||||
);
|
||||
|
||||
let mut online_node = sample_proxy_node("node-online");
|
||||
online_node.status = "online".to_string();
|
||||
online_node.tunnel_connected = true;
|
||||
let mut online_node_2 = sample_proxy_node("node-zeta");
|
||||
online_node_2.name = "zeta-online".to_string();
|
||||
online_node_2.status = "online".to_string();
|
||||
online_node_2.tunnel_connected = true;
|
||||
online_node_2.remote_config = None;
|
||||
let mut offline_node = sample_proxy_node("node-offline");
|
||||
offline_node.status = "offline".to_string();
|
||||
offline_node.tunnel_connected = false;
|
||||
offline_node.remote_config = None;
|
||||
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
|
||||
online_node,
|
||||
online_node_2,
|
||||
offline_node,
|
||||
]));
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let data_state =
|
||||
GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&proxy_node_repository))
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state.clone()),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
let config_response = client
|
||||
.put(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes/node-online/config"
|
||||
))
|
||||
.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!({
|
||||
"node_name": "edge-online",
|
||||
"allowed_ports": [443, 8443],
|
||||
"log_level": "info",
|
||||
"heartbeat_interval": 45,
|
||||
"upgrade_to": null
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(config_response.status(), StatusCode::OK);
|
||||
let config_payload: serde_json::Value = config_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(config_payload["node_id"], "node-online");
|
||||
assert_eq!(config_payload["config_version"], 8);
|
||||
assert_eq!(config_payload["node"]["name"], "edge-online");
|
||||
assert_eq!(config_payload["remote_config"]["log_level"], "info");
|
||||
assert!(config_payload["remote_config"].get("upgrade_to").is_none());
|
||||
|
||||
let upgrade_response = client
|
||||
.post(format!("{gateway_url}/api/admin/proxy-nodes/upgrade"))
|
||||
.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.0.0", "cooldown_secs": 0 }))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(upgrade_response.status(), StatusCode::OK);
|
||||
let upgrade_payload: serde_json::Value = upgrade_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(upgrade_payload["version"], "2.0.0");
|
||||
assert_eq!(upgrade_payload["batch_size"], 1);
|
||||
assert_eq!(upgrade_payload["updated"], 1);
|
||||
assert_eq!(upgrade_payload["skipped"], 1);
|
||||
assert_eq!(upgrade_payload["blocked"], false);
|
||||
assert_eq!(upgrade_payload["pending_node_ids"], json!(["node-online"]));
|
||||
assert_eq!(upgrade_payload["node_ids"], json!(["node-online"]));
|
||||
assert_eq!(upgrade_payload["completed"], 0);
|
||||
assert_eq!(upgrade_payload["remaining"], 2);
|
||||
assert_eq!(upgrade_payload["rollout_active"], true);
|
||||
|
||||
let blocked_upgrade_response = client
|
||||
.post(format!("{gateway_url}/api/admin/proxy-nodes/upgrade"))
|
||||
.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.0.0", "cooldown_secs": 0 }))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(blocked_upgrade_response.status(), StatusCode::OK);
|
||||
let blocked_upgrade_payload: serde_json::Value = blocked_upgrade_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(blocked_upgrade_payload["updated"], 0);
|
||||
assert_eq!(blocked_upgrade_payload["blocked"], true);
|
||||
assert_eq!(
|
||||
blocked_upgrade_payload["pending_node_ids"],
|
||||
json!(["node-online"])
|
||||
);
|
||||
|
||||
let heartbeat_response = client
|
||||
.post(format!("{gateway_url}/api/internal/tunnel/heartbeat"))
|
||||
.json(&json!({
|
||||
"node_id": "node-online",
|
||||
"heartbeat_interval": 45,
|
||||
"active_connections": 3,
|
||||
"total_requests": 5,
|
||||
"avg_latency_ms": 10.0,
|
||||
"proxy_metadata": { "arch": "arm64" },
|
||||
"proxy_version": "2.0.0"
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(heartbeat_response.status(), StatusCode::OK);
|
||||
let heartbeat_payload: serde_json::Value = heartbeat_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
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);
|
||||
assert_eq!(heartbeat_payload["remote_config"]["log_level"], "info");
|
||||
|
||||
let second_upgrade_response = client
|
||||
.post(format!("{gateway_url}/api/admin/proxy-nodes/upgrade"))
|
||||
.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.0.0", "cooldown_secs": 0 }))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(second_upgrade_response.status(), StatusCode::OK);
|
||||
let second_upgrade_payload: serde_json::Value = second_upgrade_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(second_upgrade_payload["batch_size"], 1);
|
||||
assert_eq!(second_upgrade_payload["updated"], 0);
|
||||
assert_eq!(second_upgrade_payload["skipped"], 2);
|
||||
assert_eq!(second_upgrade_payload["blocked"], true);
|
||||
assert_eq!(
|
||||
second_upgrade_payload["pending_node_ids"],
|
||||
json!(["node-online"])
|
||||
);
|
||||
assert_eq!(second_upgrade_payload["node_ids"], json!([]));
|
||||
assert_eq!(second_upgrade_payload["completed"], 0);
|
||||
assert_eq!(second_upgrade_payload["remaining"], 3);
|
||||
assert_eq!(second_upgrade_payload["rollout_active"], true);
|
||||
|
||||
assert!(
|
||||
record_proxy_upgrade_traffic_success(&data_state, "node-online")
|
||||
.await
|
||||
.expect("traffic confirmation should be recorded")
|
||||
);
|
||||
|
||||
let third_upgrade_response = client
|
||||
.post(format!("{gateway_url}/api/admin/proxy-nodes/upgrade"))
|
||||
.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.0.0", "cooldown_secs": 0 }))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(third_upgrade_response.status(), StatusCode::OK);
|
||||
let third_upgrade_payload: serde_json::Value = third_upgrade_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(third_upgrade_payload["batch_size"], 1);
|
||||
assert_eq!(third_upgrade_payload["updated"], 1);
|
||||
assert_eq!(third_upgrade_payload["skipped"], 1);
|
||||
assert_eq!(third_upgrade_payload["blocked"], false);
|
||||
assert_eq!(
|
||||
third_upgrade_payload["pending_node_ids"],
|
||||
json!(["node-zeta"])
|
||||
);
|
||||
assert_eq!(third_upgrade_payload["node_ids"], json!(["node-zeta"]));
|
||||
assert_eq!(third_upgrade_payload["completed"], 1);
|
||||
assert_eq!(third_upgrade_payload["remaining"], 1);
|
||||
assert_eq!(third_upgrade_payload["rollout_active"], true);
|
||||
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_marks_draining_proxy_nodes_unschedulable_and_rollout_skips_them() {
|
||||
let mut alpha = sample_proxy_node("node-alpha");
|
||||
alpha.name = "alpha".to_string();
|
||||
alpha.status = "online".to_string();
|
||||
alpha.tunnel_connected = true;
|
||||
alpha.remote_config = None;
|
||||
|
||||
let mut zeta = sample_proxy_node("node-zeta");
|
||||
zeta.name = "zeta".to_string();
|
||||
zeta.status = "online".to_string();
|
||||
zeta.tunnel_connected = true;
|
||||
zeta.remote_config = None;
|
||||
|
||||
let proxy_node_repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![zeta, alpha]));
|
||||
let data_state =
|
||||
GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&proxy_node_repository))
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new());
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
let config_response = client
|
||||
.put(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes/node-alpha/config"
|
||||
))
|
||||
.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!({
|
||||
"scheduling_state": "draining"
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(config_response.status(), StatusCode::OK);
|
||||
let config_payload: serde_json::Value = config_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(
|
||||
config_payload["remote_config"]["scheduling_state"],
|
||||
"draining"
|
||||
);
|
||||
|
||||
let list_response = client
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes?skip=0&limit=10"
|
||||
))
|
||||
.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!(list_response.status(), StatusCode::OK);
|
||||
let list_payload: serde_json::Value =
|
||||
list_response.json().await.expect("json body should parse");
|
||||
let alpha_payload = list_payload["items"]
|
||||
.as_array()
|
||||
.expect("items should be array")
|
||||
.iter()
|
||||
.find(|item| item["id"] == "node-alpha")
|
||||
.expect("alpha should exist");
|
||||
assert_eq!(
|
||||
alpha_payload["remote_config"]["scheduling_state"],
|
||||
"draining"
|
||||
);
|
||||
|
||||
let upgrade_response = client
|
||||
.post(format!("{gateway_url}/api/admin/proxy-nodes/upgrade"))
|
||||
.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.0.0", "batch_size": 2, "cooldown_secs": 0 }))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(upgrade_response.status(), StatusCode::OK);
|
||||
let upgrade_payload: serde_json::Value = upgrade_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(upgrade_payload["updated"], 1);
|
||||
assert_eq!(upgrade_payload["node_ids"], json!(["node-zeta"]));
|
||||
assert_eq!(upgrade_payload["pending_node_ids"], json!(["node-zeta"]));
|
||||
|
||||
let rollout_response = client
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/proxy-nodes?skip=0&limit=10"
|
||||
))
|
||||
.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");
|
||||
let rollout_payload: serde_json::Value = rollout_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(rollout_payload["rollout"]["online_eligible_total"], 1);
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
@@ -69,6 +69,10 @@ pub(super) fn hash_api_key(value: &str) -> String {
|
||||
format!("{:x}", hasher.finalize())
|
||||
}
|
||||
|
||||
pub(super) fn hash_management_token(value: &str) -> String {
|
||||
hash_api_key(value)
|
||||
}
|
||||
|
||||
pub(super) fn test_auth_secret() -> String {
|
||||
std::env::var("JWT_SECRET_KEY")
|
||||
.ok()
|
||||
|
||||
@@ -1,8 +1,12 @@
|
||||
use std::io;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use aether_data::repository::proxy_nodes::ProxyNodeReadRepository;
|
||||
use axum::body::Body;
|
||||
use axum::routing::{any, post};
|
||||
use axum::{extract::Request, Json, Router};
|
||||
use bytes::Bytes;
|
||||
use futures_util::stream;
|
||||
use http::header::HeaderValue;
|
||||
use http::StatusCode;
|
||||
use serde_json::json;
|
||||
@@ -45,6 +49,7 @@ async fn gateway_handles_internal_tunnel_heartbeat_locally_with_loopback() {
|
||||
.post(format!("{gateway_url}/api/internal/tunnel/heartbeat"))
|
||||
.json(&json!({
|
||||
"node_id": "node-123",
|
||||
"heartbeat_id": 77,
|
||||
"heartbeat_interval": 45,
|
||||
"active_connections": 5,
|
||||
"total_requests": 9,
|
||||
@@ -61,6 +66,7 @@ async fn gateway_handles_internal_tunnel_heartbeat_locally_with_loopback() {
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["heartbeat_id"], 77);
|
||||
assert_eq!(payload["config_version"], 7);
|
||||
assert_eq!(payload["upgrade_to"], "1.2.3");
|
||||
assert_eq!(payload["remote_config"]["allowed_ports"][0], 443);
|
||||
@@ -105,6 +111,7 @@ async fn gateway_handles_internal_tunnel_node_status_locally_with_loopback() {
|
||||
"node_id": "node-123",
|
||||
"connected": true,
|
||||
"conn_count": 4,
|
||||
"observed_at_unix_secs": 1_800_000_321u64,
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
@@ -114,6 +121,12 @@ async fn gateway_handles_internal_tunnel_node_status_locally_with_loopback() {
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["updated"], json!(true));
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
let node = repository
|
||||
.find_proxy_node("node-123")
|
||||
.await
|
||||
.expect("lookup should succeed")
|
||||
.expect("node should exist");
|
||||
assert_eq!(node.tunnel_connected_at_unix_secs, Some(1_800_000_321));
|
||||
|
||||
gateway_handle.abort();
|
||||
upstream_handle.abort();
|
||||
@@ -261,6 +274,66 @@ async fn gateway_forwards_tunnel_relay_to_attachment_owner() {
|
||||
owner_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_streams_tunnel_relay_body_to_attachment_owner() {
|
||||
let owner_hits = Arc::new(Mutex::new(0usize));
|
||||
let owner_hits_clone = Arc::clone(&owner_hits);
|
||||
let owner = Router::new().route(
|
||||
"/api/internal/tunnel/relay/node-123",
|
||||
post(move |body: Body| {
|
||||
let owner_hits_inner = Arc::clone(&owner_hits_clone);
|
||||
async move {
|
||||
*owner_hits_inner.lock().expect("mutex should lock") += 1;
|
||||
let body = axum::body::to_bytes(body, usize::MAX)
|
||||
.await
|
||||
.expect("body should read");
|
||||
assert_eq!(body, Bytes::from_static(b"relay-stream-envelope"));
|
||||
(StatusCode::OK, Body::from("stream-ok"))
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let (owner_url, owner_handle) = start_server(owner).await;
|
||||
let data_state = GatewayDataState::disabled().with_system_config_values_for_tests(vec![(
|
||||
"tunnel.attachments.node-123".to_string(),
|
||||
json!({
|
||||
"gateway_instance_id": "gateway-b",
|
||||
"relay_base_url": owner_url,
|
||||
"conn_count": 1,
|
||||
"observed_at_unix_secs": 4_102_444_800u64,
|
||||
}),
|
||||
)]);
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state)
|
||||
.with_tunnel_identity_for_tests("gateway-a", Some("http://gateway-a.internal")),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let request_body = reqwest::Body::wrap_stream(stream::iter(vec![
|
||||
Ok::<Bytes, io::Error>(Bytes::from_static(b"relay-")),
|
||||
Ok::<Bytes, io::Error>(Bytes::from_static(b"stream-")),
|
||||
Ok::<Bytes, io::Error>(Bytes::from_static(b"envelope")),
|
||||
]));
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!("{gateway_url}/api/internal/tunnel/relay/node-123"))
|
||||
.body(request_body)
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
response.text().await.expect("body should read"),
|
||||
"stream-ok"
|
||||
);
|
||||
assert_eq!(*owner_hits.lock().expect("mutex should lock"), 1);
|
||||
|
||||
gateway_handle.abort();
|
||||
owner_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_does_not_forward_tunnel_relay_twice() {
|
||||
let owner_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -310,3 +383,53 @@ async fn gateway_does_not_forward_tunnel_relay_twice() {
|
||||
gateway_handle.abort();
|
||||
owner_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_owner_relay_body_above_configured_limit() {
|
||||
let owner_hits = Arc::new(Mutex::new(0usize));
|
||||
let owner_hits_clone = Arc::clone(&owner_hits);
|
||||
let owner = Router::new().route(
|
||||
"/api/internal/tunnel/relay/node-123",
|
||||
post(move |_request: Request| {
|
||||
let owner_hits_inner = Arc::clone(&owner_hits_clone);
|
||||
async move {
|
||||
*owner_hits_inner.lock().expect("mutex should lock") += 1;
|
||||
(StatusCode::OK, Body::from("unexpected owner hit"))
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let (owner_url, owner_handle) = start_server(owner).await;
|
||||
let data_state = GatewayDataState::disabled().with_system_config_values_for_tests(vec![
|
||||
(
|
||||
"tunnel.attachments.node-123".to_string(),
|
||||
json!({
|
||||
"gateway_instance_id": "gateway-b",
|
||||
"relay_base_url": owner_url,
|
||||
"conn_count": 1,
|
||||
"observed_at_unix_secs": 4_102_444_800u64,
|
||||
}),
|
||||
),
|
||||
("max_request_body_size".to_string(), json!(8)),
|
||||
]);
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state)
|
||||
.with_tunnel_identity_for_tests("gateway-a", Some("http://gateway-a.internal")),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!("{gateway_url}/api/internal/tunnel/relay/node-123"))
|
||||
.body("relay-envelope")
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
|
||||
assert_eq!(*owner_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
owner_handle.abort();
|
||||
}
|
||||
|
||||
@@ -6,7 +6,7 @@ use std::sync::Arc;
|
||||
type HeartbeatAckCallback =
|
||||
dyn Fn(Vec<u8>) -> BoxFuture<'static, Result<Vec<u8>, String>> + Send + Sync;
|
||||
type NodeStatusCallback =
|
||||
dyn Fn(String, bool, usize) -> BoxFuture<'static, Result<(), String>> + Send + Sync;
|
||||
dyn Fn(String, bool, usize, u64) -> BoxFuture<'static, Result<(), String>> + Send + Sync;
|
||||
|
||||
enum ControlPlaneMode {
|
||||
Disabled,
|
||||
@@ -51,7 +51,7 @@ impl ControlPlaneClient {
|
||||
where
|
||||
HeartbeatAck:
|
||||
Fn(Vec<u8>) -> BoxFuture<'static, Result<Vec<u8>, String>> + Send + Sync + 'static,
|
||||
PushNodeStatus: Fn(String, bool, usize) -> BoxFuture<'static, Result<(), String>>
|
||||
PushNodeStatus: Fn(String, bool, usize, u64) -> BoxFuture<'static, Result<(), String>>
|
||||
+ Send
|
||||
+ Sync
|
||||
+ 'static,
|
||||
@@ -103,6 +103,7 @@ impl ControlPlaneClient {
|
||||
node_id: &str,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
observed_at_unix_secs: u64,
|
||||
) -> Result<(), String> {
|
||||
match self.inner.as_ref() {
|
||||
ControlPlaneMode::Disabled => Ok(()),
|
||||
@@ -120,6 +121,7 @@ impl ControlPlaneClient {
|
||||
"node_id": node_id,
|
||||
"connected": connected,
|
||||
"conn_count": conn_count,
|
||||
"observed_at_unix_secs": observed_at_unix_secs,
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
@@ -135,7 +137,15 @@ impl ControlPlaneClient {
|
||||
}
|
||||
ControlPlaneMode::Local {
|
||||
push_node_status, ..
|
||||
} => push_node_status(node_id.to_string(), connected, conn_count).await,
|
||||
} => {
|
||||
push_node_status(
|
||||
node_id.to_string(),
|
||||
connected,
|
||||
conn_count,
|
||||
observed_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use std::collections::HashMap;
|
||||
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicUsize, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use aether_runtime::{BoundedQueueSender, MetricKind, MetricSample, QueueSendError};
|
||||
use axum::extract::ws::Message;
|
||||
@@ -324,10 +324,46 @@ pub struct HubRouter {
|
||||
next_conn_id: AtomicU64,
|
||||
next_local_stream_id: AtomicU64,
|
||||
control_plane: ControlPlaneClient,
|
||||
node_status_tx: mpsc::UnboundedSender<NodeStatusEvent>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct NodeStatusEvent {
|
||||
node_id: String,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
observed_at_unix_secs: u64,
|
||||
}
|
||||
|
||||
impl HubRouter {
|
||||
pub fn new(control_plane: ControlPlaneClient) -> Arc<Self> {
|
||||
let (node_status_tx, mut node_status_rx) = mpsc::unbounded_channel::<NodeStatusEvent>();
|
||||
let worker_control_plane = control_plane.clone();
|
||||
if let Ok(handle) = tokio::runtime::Handle::try_current() {
|
||||
handle.spawn(async move {
|
||||
while let Some(event) = node_status_rx.recv().await {
|
||||
if let Err(error) = worker_control_plane
|
||||
.push_node_status(
|
||||
&event.node_id,
|
||||
event.connected,
|
||||
event.conn_count,
|
||||
event.observed_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
node_id = %event.node_id,
|
||||
connected = event.connected,
|
||||
conn_count = event.conn_count,
|
||||
observed_at_unix_secs = event.observed_at_unix_secs,
|
||||
error = %error,
|
||||
"failed to push node status to app control plane"
|
||||
);
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
Arc::new(Self {
|
||||
proxy_conns: RwLock::new(HashMap::new()),
|
||||
proxy_conns_by_id: DashMap::new(),
|
||||
@@ -336,6 +372,7 @@ impl HubRouter {
|
||||
next_conn_id: AtomicU64::new(1),
|
||||
next_local_stream_id: AtomicU64::new(1),
|
||||
control_plane,
|
||||
node_status_tx,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -405,21 +442,21 @@ impl HubRouter {
|
||||
}
|
||||
|
||||
fn notify_node_status(&self, node_id: String, connected: bool, conn_count: usize) {
|
||||
let control_plane = self.control_plane.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(error) = control_plane
|
||||
.push_node_status(&node_id, connected, conn_count)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
node_id = %node_id,
|
||||
connected = connected,
|
||||
conn_count = conn_count,
|
||||
error = %error,
|
||||
"failed to push node status to app control plane"
|
||||
);
|
||||
}
|
||||
});
|
||||
let event = NodeStatusEvent {
|
||||
node_id,
|
||||
connected,
|
||||
conn_count,
|
||||
observed_at_unix_secs: current_unix_secs(),
|
||||
};
|
||||
if let Err(error) = self.node_status_tx.send(event) {
|
||||
warn!(
|
||||
node_id = %error.0.node_id,
|
||||
connected = error.0.connected,
|
||||
conn_count = error.0.conn_count,
|
||||
observed_at_unix_secs = error.0.observed_at_unix_secs,
|
||||
"node status worker unavailable"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
fn get_proxy_conn(&self, node_id: &str) -> Option<Arc<ProxyConn>> {
|
||||
@@ -750,8 +787,12 @@ impl HubRouter {
|
||||
let ack_payload = match self.control_plane.heartbeat_ack(&payload).await {
|
||||
Ok(payload) => payload,
|
||||
Err(error) => {
|
||||
warn!(proxy_conn_id = proxy_conn_id, error = %error, "control-plane heartbeat callback failed");
|
||||
b"{}".to_vec()
|
||||
warn!(
|
||||
proxy_conn_id = proxy_conn_id,
|
||||
error = %error,
|
||||
"control-plane heartbeat callback failed; keeping heartbeat pending"
|
||||
);
|
||||
return;
|
||||
}
|
||||
};
|
||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
@@ -796,6 +837,13 @@ impl HubRouter {
|
||||
}
|
||||
}
|
||||
|
||||
fn current_unix_secs() -> u64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs()
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct HubStats {
|
||||
pub proxy_connections: usize,
|
||||
@@ -845,6 +893,8 @@ mod tests {
|
||||
url: "https://example.com".to_string(),
|
||||
headers: HashMap::new(),
|
||||
timeout: 30,
|
||||
follow_redirects: None,
|
||||
http1_only: false,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -924,4 +974,34 @@ mod tests {
|
||||
assert_eq!(second_header.msg_type, protocol::REQUEST_BODY);
|
||||
assert_ne!(second_header.flags & protocol::FLAG_END_STREAM, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn heartbeat_callback_failure_does_not_send_fake_ack() {
|
||||
let hub = HubRouter::new(ControlPlaneClient::local(
|
||||
|_payload| Box::pin(async { Err("db unavailable".to_string()) }),
|
||||
|_node_id, _connected, _conn_count, _observed_at_unix_secs| Box::pin(async { Ok(()) }),
|
||||
));
|
||||
|
||||
let (proxy_tx, mut proxy_rx) = bounded_queue(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
let proxy = Arc::new(ProxyConn::new(
|
||||
300,
|
||||
"node-3".to_string(),
|
||||
"Node 3".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
));
|
||||
hub.register_proxy(proxy);
|
||||
|
||||
let payload = serde_json::to_vec(&serde_json::json!({
|
||||
"node_id": "node-3",
|
||||
"heartbeat_id": 99u64,
|
||||
}))
|
||||
.expect("payload should serialize");
|
||||
let mut frame = protocol::encode_frame(1, protocol::HEARTBEAT_DATA, 0, &payload);
|
||||
hub.handle_proxy_frame(300, &mut frame).await;
|
||||
|
||||
assert!(proxy_rx.try_recv().is_err());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,6 +14,7 @@ use futures_util::StreamExt;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::api::response::apply_streaming_response_headers;
|
||||
use crate::maintenance::record_proxy_upgrade_traffic_success;
|
||||
|
||||
use super::hub::{LocalBodyEvent, LocalStream};
|
||||
use super::protocol;
|
||||
@@ -222,6 +223,13 @@ pub async fn relay_request(
|
||||
);
|
||||
}
|
||||
};
|
||||
if let Err(error) = record_proxy_upgrade_traffic_success(state.data.as_ref(), &node_id).await {
|
||||
warn!(
|
||||
node_id = %node_id,
|
||||
error = %error,
|
||||
"failed to record proxy upgrade traffic confirmation"
|
||||
);
|
||||
}
|
||||
|
||||
let Some(mut body_rx) = stream.take_body_receiver() else {
|
||||
state
|
||||
@@ -340,12 +348,25 @@ fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Respo
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::{AppState, ConnConfig, ControlPlaneClient};
|
||||
use super::super::hub::ProxyConn;
|
||||
use super::super::{protocol, AppState, ConnConfig, ControlPlaneClient};
|
||||
use super::{relay_request, Body, Request, SocketAddr, StatusCode, TUNNEL_ERROR_HEADER};
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::maintenance::start_proxy_upgrade_rollout;
|
||||
use aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER;
|
||||
use aether_data::repository::proxy_nodes::{
|
||||
InMemoryProxyNodeRepository, ProxyNodeHeartbeatMutation, ProxyNodeWriteRepository,
|
||||
StoredProxyNode,
|
||||
};
|
||||
use axum::extract::ws::Message;
|
||||
use axum::extract::{ConnectInfo, Path, State};
|
||||
use axum::response::IntoResponse;
|
||||
use bytes::Bytes;
|
||||
use serde_json::json;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use tokio::sync::watch;
|
||||
|
||||
fn test_app_state() -> AppState {
|
||||
AppState::new(
|
||||
@@ -359,6 +380,49 @@ mod tests {
|
||||
)
|
||||
}
|
||||
|
||||
fn sample_connected_proxy_node(node_id: &str) -> StoredProxyNode {
|
||||
StoredProxyNode::new(
|
||||
node_id.to_string(),
|
||||
format!("proxy-{node_id}"),
|
||||
"127.0.0.1".to_string(),
|
||||
0,
|
||||
false,
|
||||
"online".to_string(),
|
||||
30,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
true,
|
||||
true,
|
||||
0,
|
||||
)
|
||||
.expect("node should build")
|
||||
.with_runtime_fields(
|
||||
Some("test".to_string()),
|
||||
None,
|
||||
Some(1_800_000_000),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
Some(1_800_000_000),
|
||||
None,
|
||||
Some(1_800_000_000),
|
||||
Some(1_800_000_000),
|
||||
)
|
||||
}
|
||||
|
||||
fn encode_relay_envelope(meta: &protocol::RequestMeta, body: &[u8]) -> Vec<u8> {
|
||||
let meta_bytes = serde_json::to_vec(meta).expect("meta should serialize");
|
||||
let mut payload = Vec::with_capacity(4 + meta_bytes.len() + body.len());
|
||||
payload.extend_from_slice(&(meta_bytes.len() as u32).to_be_bytes());
|
||||
payload.extend_from_slice(&meta_bytes);
|
||||
payload.extend_from_slice(body);
|
||||
payload
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn relay_rejects_non_loopback_without_forwarded_header() {
|
||||
let request = Request::builder()
|
||||
@@ -400,4 +464,149 @@ mod tests {
|
||||
Some("bad_request")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn relay_records_real_traffic_confirmation_for_upgrade_rollout() {
|
||||
let mut node = sample_connected_proxy_node("node-123");
|
||||
node.proxy_metadata = Some(json!({"version": "1.0.0"}));
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![node]));
|
||||
let data = Arc::new(
|
||||
GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository))
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new()),
|
||||
);
|
||||
|
||||
let started = start_proxy_upgrade_rollout(data.as_ref(), "2.0.0".to_string(), 1, 0, None)
|
||||
.await
|
||||
.expect("rollout should start");
|
||||
assert_eq!(started.node_ids, vec!["node-123".to_string()]);
|
||||
|
||||
repository
|
||||
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: "node-123".to_string(),
|
||||
heartbeat_interval: None,
|
||||
active_connections: Some(1),
|
||||
total_requests_delta: Some(1),
|
||||
avg_latency_ms: Some(2.0),
|
||||
failed_requests_delta: Some(0),
|
||||
dns_failures_delta: Some(0),
|
||||
stream_errors_delta: Some(0),
|
||||
proxy_metadata: Some(json!({"version": "2.0.0"})),
|
||||
proxy_version: Some("2.0.0".to_string()),
|
||||
})
|
||||
.await
|
||||
.expect("heartbeat should succeed");
|
||||
|
||||
let observed = start_proxy_upgrade_rollout(data.as_ref(), "2.0.0".to_string(), 1, 0, None)
|
||||
.await
|
||||
.expect("rollout should observe version confirmation");
|
||||
assert!(observed.blocked);
|
||||
assert_eq!(observed.pending_node_ids, vec!["node-123".to_string()]);
|
||||
|
||||
let state = test_app_state().with_data(Arc::clone(&data));
|
||||
let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
state.hub.register_proxy(Arc::new(ProxyConn::new(
|
||||
500,
|
||||
"node-123".to_string(),
|
||||
"Node 123".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
)));
|
||||
|
||||
let meta = protocol::RequestMeta {
|
||||
method: "GET".to_string(),
|
||||
url: "https://example.com/health".to_string(),
|
||||
headers: HashMap::new(),
|
||||
timeout: 30,
|
||||
follow_redirects: None,
|
||||
http1_only: false,
|
||||
};
|
||||
let request = Request::builder()
|
||||
.body(Body::from(encode_relay_envelope(&meta, &[])))
|
||||
.expect("request should build");
|
||||
|
||||
let relay_state = state.clone();
|
||||
let relay_task = tokio::spawn(async move {
|
||||
relay_request(
|
||||
Path("node-123".to_string()),
|
||||
State(relay_state),
|
||||
ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4242))),
|
||||
request,
|
||||
)
|
||||
.await
|
||||
.into_response()
|
||||
});
|
||||
|
||||
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 = protocol::FrameHeader::parse(&request_headers)
|
||||
.expect("request header frame should parse");
|
||||
assert_eq!(request_header.msg_type, protocol::REQUEST_HEADERS);
|
||||
|
||||
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 =
|
||||
protocol::FrameHeader::parse(&request_body).expect("request body frame should parse");
|
||||
assert_eq!(request_body_header.msg_type, protocol::REQUEST_BODY);
|
||||
|
||||
let response_meta = protocol::ResponseMeta {
|
||||
status: 200,
|
||||
headers: vec![("content-type".to_string(), "text/plain".to_string())],
|
||||
};
|
||||
let response_payload =
|
||||
serde_json::to_vec(&response_meta).expect("response meta should serialize");
|
||||
let mut response_headers_frame = protocol::encode_frame(
|
||||
request_header.stream_id,
|
||||
protocol::RESPONSE_HEADERS,
|
||||
0,
|
||||
&response_payload,
|
||||
);
|
||||
state
|
||||
.hub
|
||||
.handle_proxy_frame(500, &mut response_headers_frame)
|
||||
.await;
|
||||
|
||||
let mut response_body_frame = protocol::encode_frame(
|
||||
request_header.stream_id,
|
||||
protocol::RESPONSE_BODY,
|
||||
0,
|
||||
Bytes::new().as_ref(),
|
||||
);
|
||||
state
|
||||
.hub
|
||||
.handle_proxy_frame(500, &mut response_body_frame)
|
||||
.await;
|
||||
let mut response_end_frame =
|
||||
protocol::encode_frame(request_header.stream_id, protocol::STREAM_END, 0, &[]);
|
||||
state
|
||||
.hub
|
||||
.handle_proxy_frame(500, &mut response_end_frame)
|
||||
.await;
|
||||
|
||||
let response = relay_task.await.expect("relay task should complete");
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("response body should read");
|
||||
assert!(body.is_empty());
|
||||
|
||||
let rollout_entry = data
|
||||
.list_system_config_entries()
|
||||
.await
|
||||
.expect("system config list should succeed")
|
||||
.into_iter()
|
||||
.find(|entry| entry.key == "proxy_node_upgrade_rollout")
|
||||
.expect("rollout entry should exist");
|
||||
let tracked_nodes = rollout_entry.value["tracked_nodes"]
|
||||
.as_array()
|
||||
.expect("tracked nodes should be an array");
|
||||
assert_eq!(tracked_nodes.len(), 1);
|
||||
assert!(tracked_nodes[0]["version_confirmed_at_unix_secs"].is_u64());
|
||||
assert!(tracked_nodes[0]["traffic_confirmed_at_unix_secs"].is_u64());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -19,8 +19,10 @@ use axum::routing::{get, post};
|
||||
use axum::Router;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::data::GatewayDataState;
|
||||
|
||||
pub use control_plane::ControlPlaneClient;
|
||||
pub use hub::{ConnConfig, HubRouter};
|
||||
pub use hub::{ConnConfig, HubRouter, ProxyConn};
|
||||
pub use local_relay::relay_request;
|
||||
|
||||
#[derive(Clone)]
|
||||
@@ -28,6 +30,7 @@ pub struct AppState {
|
||||
pub hub: Arc<HubRouter>,
|
||||
pub proxy_conn_cfg: ConnConfig,
|
||||
pub max_streams: usize,
|
||||
data: Arc<GatewayDataState>,
|
||||
request_gate: Option<Arc<ConcurrencyGate>>,
|
||||
distributed_request_gate: Option<Arc<DistributedConcurrencyGate>>,
|
||||
}
|
||||
@@ -48,11 +51,17 @@ impl AppState {
|
||||
hub: HubRouter::new(control_plane),
|
||||
proxy_conn_cfg,
|
||||
max_streams,
|
||||
data: Arc::new(GatewayDataState::disabled()),
|
||||
request_gate: None,
|
||||
distributed_request_gate: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn with_data(mut self, data: Arc<GatewayDataState>) -> Self {
|
||||
self.data = data;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_request_concurrency_limit(mut self, limit: Option<usize>) -> Self {
|
||||
self.request_gate = limit
|
||||
.filter(|limit| *limit > 0)
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
mod embedded;
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::fmt;
|
||||
use std::io;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, SystemTime};
|
||||
|
||||
@@ -11,11 +14,13 @@ use aether_data::repository::proxy_nodes::{
|
||||
ProxyNodeHeartbeatMutation, ProxyNodeTunnelStatusMutation, StoredProxyNode,
|
||||
};
|
||||
use aether_runtime::MetricSample;
|
||||
use axum::body::{to_bytes, Body};
|
||||
use async_stream::stream;
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::extract::ws::WebSocketUpgrade;
|
||||
use axum::extract::{ConnectInfo, Path, Request, State};
|
||||
use axum::http::{HeaderMap, StatusCode};
|
||||
use axum::response::IntoResponse;
|
||||
use futures_util::StreamExt;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::json;
|
||||
use tracing::warn;
|
||||
@@ -28,6 +33,7 @@ use super::error::GatewayError;
|
||||
use super::headers::{extract_or_generate_trace_id, should_skip_request_header};
|
||||
use super::AppState;
|
||||
|
||||
pub(crate) use embedded::ProxyConn as TunnelProxyConn;
|
||||
pub use embedded::{
|
||||
build_router_with_state as build_tunnel_runtime_router_with_state, protocol as tunnel_protocol,
|
||||
AppState as TunnelRuntimeState, ConnConfig as TunnelConnConfig,
|
||||
@@ -45,6 +51,7 @@ const DEFAULT_PING_INTERVAL_SECS: u64 = 15;
|
||||
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 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";
|
||||
@@ -55,6 +62,8 @@ const TUNNEL_ATTACHMENT_TTL_ENV: &str = "AETHER_TUNNEL_ATTACHMENT_TTL_SECS";
|
||||
struct InternalTunnelHeartbeatRequest {
|
||||
node_id: String,
|
||||
#[serde(default)]
|
||||
heartbeat_id: Option<u64>,
|
||||
#[serde(default)]
|
||||
heartbeat_interval: Option<i32>,
|
||||
#[serde(default)]
|
||||
active_connections: Option<i32>,
|
||||
@@ -183,6 +192,7 @@ impl TunnelAttachmentDirectory {
|
||||
node_id: &str,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
observed_at_unix_secs: u64,
|
||||
) -> Result<(), String> {
|
||||
let node_id = node_id.trim();
|
||||
if node_id.is_empty() {
|
||||
@@ -203,7 +213,7 @@ impl TunnelAttachmentDirectory {
|
||||
gateway_instance_id: self.identity.instance_id.clone(),
|
||||
relay_base_url: relay_base_url.clone(),
|
||||
conn_count,
|
||||
observed_at_unix_secs: current_unix_secs(),
|
||||
observed_at_unix_secs,
|
||||
},
|
||||
)
|
||||
.await
|
||||
@@ -404,7 +414,8 @@ impl EmbeddedTunnelState {
|
||||
outbound_queue_capacity: DEFAULT_OUTBOUND_QUEUE_CAPACITY,
|
||||
},
|
||||
DEFAULT_MAX_STREAMS,
|
||||
),
|
||||
)
|
||||
.with_data(data),
|
||||
attachment_directory,
|
||||
}
|
||||
}
|
||||
@@ -434,6 +445,39 @@ impl EmbeddedTunnelState {
|
||||
self.inner.hub.stats().to_metric_samples()
|
||||
}
|
||||
|
||||
pub(crate) async fn probe_node_url(
|
||||
&self,
|
||||
node_id: &str,
|
||||
url: &str,
|
||||
timeout_secs: u64,
|
||||
) -> Result<u16, String> {
|
||||
let timeout_secs = timeout_secs.clamp(5, 60);
|
||||
let meta = tunnel_protocol::RequestMeta {
|
||||
method: "GET".to_string(),
|
||||
url: url.trim().to_string(),
|
||||
headers: HashMap::new(),
|
||||
timeout: timeout_secs,
|
||||
follow_redirects: Some(false),
|
||||
http1_only: false,
|
||||
};
|
||||
let stream = self.inner.hub.open_local_stream(node_id, &meta)?;
|
||||
let stream_id = stream.id;
|
||||
let result = async {
|
||||
self.inner
|
||||
.hub
|
||||
.push_local_request_body(stream_id, Bytes::new(), true)?;
|
||||
let response = stream
|
||||
.wait_headers(Duration::from_secs(timeout_secs))
|
||||
.await?;
|
||||
Ok(response.status)
|
||||
}
|
||||
.await;
|
||||
self.inner
|
||||
.hub
|
||||
.cancel_local_stream(stream_id, "tunnel health probe completed");
|
||||
result
|
||||
}
|
||||
|
||||
pub(crate) fn local_instance_id(&self) -> &str {
|
||||
self.attachment_directory.local_instance_id()
|
||||
}
|
||||
@@ -579,14 +623,26 @@ fn build_embedded_control_plane(
|
||||
Ok(ack)
|
||||
})
|
||||
},
|
||||
move |node_id, connected, conn_count| {
|
||||
move |node_id, connected, conn_count, observed_at_unix_secs| {
|
||||
let data = Arc::clone(&node_status_data);
|
||||
let directory = node_status_directory.clone();
|
||||
Box::pin(async move {
|
||||
apply_embedded_tunnel_node_status(data.as_ref(), &node_id, connected, conn_count)
|
||||
.await?;
|
||||
apply_embedded_tunnel_node_status(
|
||||
data.as_ref(),
|
||||
&node_id,
|
||||
connected,
|
||||
conn_count,
|
||||
Some(observed_at_unix_secs),
|
||||
)
|
||||
.await?;
|
||||
if let Err(error) = directory
|
||||
.sync_node_status(data.as_ref(), &node_id, connected, conn_count)
|
||||
.sync_node_status(
|
||||
data.as_ref(),
|
||||
&node_id,
|
||||
connected,
|
||||
conn_count,
|
||||
observed_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(error = %error, node_id = %node_id, "failed to sync tunnel attachment");
|
||||
@@ -606,9 +662,16 @@ async fn forward_relay_request_to_owner(
|
||||
) -> Result<axum::http::Response<Body>, GatewayError> {
|
||||
let owner_url = build_owner_relay_url(&owner.relay_base_url, node_id)?;
|
||||
let (parts, body) = request.into_parts();
|
||||
let body = to_bytes(body, usize::MAX)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let body_limit = owner_relay_body_limit_bytes(state.data.as_ref()).await;
|
||||
if request_content_length_exceeds_limit(&parts.headers, body_limit) {
|
||||
return build_local_http_error_response(
|
||||
trace_id,
|
||||
None,
|
||||
StatusCode::PAYLOAD_TOO_LARGE,
|
||||
&format!("tunnel relay body exceeds {body_limit} bytes"),
|
||||
);
|
||||
}
|
||||
let limit_exceeded = Arc::new(AtomicBool::new(false));
|
||||
|
||||
let mut upstream_request = state.client.post(owner_url);
|
||||
for (name, value) in &parts.headers {
|
||||
@@ -630,11 +693,30 @@ async fn forward_relay_request_to_owner(
|
||||
upstream_request = upstream_request.header(TRACE_ID_HEADER, trace_id);
|
||||
}
|
||||
|
||||
let upstream_response = upstream_request
|
||||
.body(body)
|
||||
let upstream_response = match upstream_request
|
||||
.body(build_owner_relay_request_body(
|
||||
body,
|
||||
body_limit,
|
||||
Arc::clone(&limit_exceeded),
|
||||
))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(format!("owner tunnel relay failed: {err}")))?;
|
||||
{
|
||||
Ok(response) => response,
|
||||
Err(err) if limit_exceeded.load(Ordering::SeqCst) => {
|
||||
return build_local_http_error_response(
|
||||
trace_id,
|
||||
None,
|
||||
StatusCode::PAYLOAD_TOO_LARGE,
|
||||
&format!("tunnel relay body exceeds {body_limit} bytes"),
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
return Err(GatewayError::Internal(format!(
|
||||
"owner tunnel relay failed: {err}"
|
||||
)));
|
||||
}
|
||||
};
|
||||
|
||||
build_client_response(upstream_response, trace_id, None)
|
||||
}
|
||||
@@ -656,6 +738,57 @@ fn build_owner_relay_url(relay_base_url: &str, node_id: &str) -> Result<String,
|
||||
Ok(url.to_string())
|
||||
}
|
||||
|
||||
async fn owner_relay_body_limit_bytes(data: &GatewayDataState) -> usize {
|
||||
data.find_system_config_value("max_request_body_size")
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
.and_then(|value| value.as_u64())
|
||||
.and_then(|value| usize::try_from(value).ok())
|
||||
.filter(|value| *value > 0)
|
||||
.unwrap_or(DEFAULT_OWNER_RELAY_BODY_LIMIT_BYTES)
|
||||
}
|
||||
|
||||
fn request_content_length_exceeds_limit(headers: &HeaderMap, body_limit: usize) -> bool {
|
||||
headers
|
||||
.get(http::header::CONTENT_LENGTH)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| value.parse::<u64>().ok())
|
||||
.and_then(|value| usize::try_from(value).ok())
|
||||
.is_some_and(|value| value > body_limit)
|
||||
}
|
||||
|
||||
fn build_owner_relay_request_body(
|
||||
body: Body,
|
||||
body_limit: usize,
|
||||
limit_exceeded: Arc<AtomicBool>,
|
||||
) -> reqwest::Body {
|
||||
let mut body_stream = body.into_data_stream();
|
||||
reqwest::Body::wrap_stream(stream! {
|
||||
let mut forwarded = 0usize;
|
||||
while let Some(next_chunk) = body_stream.next().await {
|
||||
match next_chunk {
|
||||
Ok(chunk) => {
|
||||
forwarded = forwarded.saturating_add(chunk.len());
|
||||
if forwarded > body_limit {
|
||||
limit_exceeded.store(true, Ordering::SeqCst);
|
||||
yield Err::<Bytes, io::Error>(io::Error::new(
|
||||
io::ErrorKind::InvalidData,
|
||||
format!("tunnel relay body exceeds {body_limit} bytes"),
|
||||
));
|
||||
break;
|
||||
}
|
||||
yield Ok::<Bytes, io::Error>(chunk);
|
||||
}
|
||||
Err(err) => {
|
||||
yield Err::<Bytes, io::Error>(io::Error::other(err));
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn tunnel_attachment_key(node_id: &str) -> String {
|
||||
format!("{TUNNEL_ATTACHMENT_KEY_PREFIX}{}", node_id.trim())
|
||||
}
|
||||
@@ -719,7 +852,10 @@ async fn apply_embedded_tunnel_heartbeat(
|
||||
.map_err(|err| format!("heartbeat sync failed: {err}"))?
|
||||
.ok_or_else(|| format!("heartbeat sync failed: ProxyNode {node_id} 不存在"))?;
|
||||
|
||||
Ok(build_embedded_tunnel_heartbeat_ack(&node))
|
||||
Ok(build_embedded_tunnel_heartbeat_ack(
|
||||
&node,
|
||||
payload.heartbeat_id,
|
||||
))
|
||||
}
|
||||
|
||||
async fn apply_embedded_tunnel_node_status(
|
||||
@@ -727,13 +863,14 @@ async fn apply_embedded_tunnel_node_status(
|
||||
node_id: &str,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
observed_at_unix_secs: Option<u64>,
|
||||
) -> Result<(), String> {
|
||||
let mutation = ProxyNodeTunnelStatusMutation {
|
||||
node_id: node_id.trim().to_string(),
|
||||
connected,
|
||||
conn_count: conn_count.min(i32::MAX as usize) as i32,
|
||||
detail: None,
|
||||
observed_at_unix_secs: None,
|
||||
observed_at_unix_secs,
|
||||
};
|
||||
|
||||
data.update_proxy_node_tunnel_status(&mutation)
|
||||
@@ -742,22 +879,26 @@ async fn apply_embedded_tunnel_node_status(
|
||||
.map_err(|err| format!("node status sync failed: {err}"))
|
||||
}
|
||||
|
||||
fn build_embedded_tunnel_heartbeat_ack(node: &StoredProxyNode) -> Vec<u8> {
|
||||
let Some(remote_config) = node.remote_config.as_ref() else {
|
||||
return b"{}".to_vec();
|
||||
};
|
||||
|
||||
fn build_embedded_tunnel_heartbeat_ack(
|
||||
node: &StoredProxyNode,
|
||||
heartbeat_id: Option<u64>,
|
||||
) -> Vec<u8> {
|
||||
let mut payload = serde_json::Map::new();
|
||||
payload.insert("remote_config".to_string(), remote_config.clone());
|
||||
payload.insert("config_version".to_string(), json!(node.config_version));
|
||||
if let Some(upgrade_to) = remote_config
|
||||
.as_object()
|
||||
.and_then(|value| value.get("upgrade_to"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
payload.insert("upgrade_to".to_string(), json!(upgrade_to));
|
||||
if let Some(heartbeat_id) = 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));
|
||||
if let Some(upgrade_to) = remote_config
|
||||
.as_object()
|
||||
.and_then(|value| value.get("upgrade_to"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
payload.insert("upgrade_to".to_string(), json!(upgrade_to));
|
||||
}
|
||||
}
|
||||
|
||||
serde_json::to_vec(&serde_json::Value::Object(payload)).unwrap_or_else(|_| b"{}".to_vec())
|
||||
@@ -857,6 +998,7 @@ mod tests {
|
||||
&data,
|
||||
br#"{
|
||||
"node_id": "node-123",
|
||||
"heartbeat_id": 42,
|
||||
"heartbeat_interval": 45,
|
||||
"active_connections": 5,
|
||||
"total_requests": 9,
|
||||
@@ -873,6 +1015,7 @@ mod tests {
|
||||
|
||||
let payload: serde_json::Value =
|
||||
serde_json::from_slice(&ack).expect("ack payload should parse");
|
||||
assert_eq!(payload["heartbeat_id"], 42);
|
||||
assert_eq!(payload["config_version"], 7);
|
||||
assert_eq!(payload["upgrade_to"], "1.2.3");
|
||||
assert_eq!(payload["remote_config"]["allowed_ports"][0], 443);
|
||||
@@ -906,7 +1049,7 @@ mod tests {
|
||||
)]));
|
||||
let data = GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository));
|
||||
|
||||
apply_embedded_tunnel_node_status(&data, "node-123", true, 4)
|
||||
apply_embedded_tunnel_node_status(&data, "node-123", true, 4, Some(1_800_000_123))
|
||||
.await
|
||||
.expect("node status should succeed");
|
||||
|
||||
@@ -917,6 +1060,7 @@ mod tests {
|
||||
.expect("node should exist");
|
||||
assert_eq!(node.status, "online");
|
||||
assert_eq!(node.tunnel_connected, true);
|
||||
assert_eq!(node.tunnel_connected_at_unix_secs, Some(1_800_000_123));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -929,7 +1073,7 @@ mod tests {
|
||||
);
|
||||
|
||||
directory
|
||||
.sync_node_status(&data, "node-123", true, 2)
|
||||
.sync_node_status(&data, "node-123", true, 2, 1_800_000_010)
|
||||
.await
|
||||
.expect("attachment should sync");
|
||||
let record = directory
|
||||
@@ -940,9 +1084,10 @@ mod tests {
|
||||
assert_eq!(record.gateway_instance_id, "gateway-a");
|
||||
assert_eq!(record.relay_base_url, "http://gateway-a.internal");
|
||||
assert_eq!(record.conn_count, 2);
|
||||
assert_eq!(record.observed_at_unix_secs, 1_800_000_010);
|
||||
|
||||
directory
|
||||
.sync_node_status(&data, "node-123", false, 0)
|
||||
.sync_node_status(&data, "node-123", false, 0, 1_800_000_011)
|
||||
.await
|
||||
.expect("attachment should clear");
|
||||
assert!(directory
|
||||
|
||||
Reference in New Issue
Block a user