Merge branch 'fawney19:main' into main

This commit is contained in:
ZheFox
2026-05-23 19:54:24 +08:00
committed by GitHub
49 changed files with 2041 additions and 195 deletions
@@ -26,7 +26,6 @@ pub(crate) fn mount_public_support_routes(router: Router<AppState>) -> Router<Ap
.route("/api/capabilities/model/{*model_path}", get(proxy_request))
.route("/install/{*install_path}", get(proxy_request))
.route("/install-tunnel/{*install_path}", get(proxy_request))
.route("/install-proxy/{*install_path}", get(proxy_request))
.route("/i/{*install_path}", get(proxy_request))
.route("/", get(proxy_request))
}
@@ -797,7 +797,6 @@ pub(super) fn classify_public_support_route(
} else if method == http::Method::GET
&& (has_single_segment_after_prefix(normalized_path, "/install/")
|| has_single_segment_after_prefix(normalized_path, "/install-tunnel/")
|| has_single_segment_after_prefix(normalized_path, "/install-proxy/")
|| has_single_segment_after_prefix(normalized_path, "/i/"))
{
Some(classified(
@@ -38,6 +38,9 @@ use crate::{AppState, GatewayError};
const HUB_RELAY_CONTENT_TYPE: &str = "application/vnd.aether.tunnel-envelope";
const HUB_RELAY_ERROR_HEADER: &str = "x-aether-tunnel-error";
const TUNNEL_RELAY_PATH_PREFIX: &str = "/api/internal/tunnel/relay";
const DEFAULT_TUNNEL_TIMEOUT_MS: u64 = 60_000;
const MIN_TUNNEL_TIMEOUT_SECS: u64 = 1;
const MAX_TUNNEL_TIMEOUT_SECS: u64 = 300;
pub(crate) fn format_upstream_request_error(err: &reqwest::Error) -> String {
let mut kinds = Vec::new();
if err.is_connect() {
@@ -176,6 +179,12 @@ struct RelayRequestMeta {
method: String,
url: String,
headers: BTreeMap<String, String>,
#[serde(default, skip_serializing_if = "is_false")]
stream: bool,
#[serde(skip_serializing_if = "Option::is_none")]
request_timeout_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
stream_first_byte_timeout_ms: Option<u64>,
timeout: u64,
#[serde(skip_serializing_if = "Option::is_none")]
follow_redirects: Option<bool>,
@@ -195,6 +204,13 @@ pub(crate) struct ExecutionTransportControls {
accept_invalid_certs: bool,
}
#[derive(Debug, Clone, Copy)]
struct TunnelTimeoutMetadata {
request_timeout_ms: Option<u64>,
stream_first_byte_timeout_ms: Option<u64>,
legacy_timeout_secs: u64,
}
pub(crate) enum DirectUpstreamResponse {
Reqwest(reqwest::Response),
BrowserWreq(wreq::Response),
@@ -579,6 +595,7 @@ fn build_direct_tunnel_request_meta(
headers: &HeaderMap,
transport_controls: ExecutionTransportControls,
) -> tunnel_protocol::RequestMeta {
let timeout_metadata = resolve_tunnel_timeout_metadata(plan);
tunnel_protocol::RequestMeta {
provider_id: Some(plan.provider_id.clone()),
endpoint_id: Some(plan.endpoint_id.clone()),
@@ -586,7 +603,10 @@ fn build_direct_tunnel_request_meta(
method: plan.method.clone(),
url: plan.url.clone(),
headers: header_map_to_string_map(headers).into_iter().collect(),
timeout: resolve_relay_timeout_seconds(plan),
stream: plan.stream,
request_timeout_ms: timeout_metadata.request_timeout_ms,
stream_first_byte_timeout_ms: timeout_metadata.stream_first_byte_timeout_ms,
timeout: timeout_metadata.legacy_timeout_secs,
follow_redirects: transport_controls.follow_redirects,
http1_only: transport_controls.http1_only,
transport_profile: plan.transport_profile.clone(),
@@ -753,7 +773,8 @@ async fn send_via_tunnel_relay(
) -> 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);
let timeout_secs = resolve_relay_timeout_seconds(plan);
let timeout_metadata = resolve_tunnel_timeout_metadata(plan);
let timeout_secs = timeout_metadata.legacy_timeout_secs;
let envelope = build_relay_envelope(
RelayRequestMeta {
provider_id: plan.provider_id.clone(),
@@ -762,6 +783,9 @@ async fn send_via_tunnel_relay(
method: method.as_str().to_string(),
url: plan.url.clone(),
headers: header_map_to_string_map(&headers),
stream: plan.stream,
request_timeout_ms: timeout_metadata.request_timeout_ms,
stream_first_byte_timeout_ms: timeout_metadata.stream_first_byte_timeout_ms,
timeout: timeout_secs,
follow_redirects: transport_controls.follow_redirects,
http1_only: transport_controls.http1_only,
@@ -791,15 +815,22 @@ async fn send_via_tunnel_relay(
.request(reqwest::Method::POST, relay_url)
.header(reqwest::header::CONTENT_TYPE, HUB_RELAY_CONTENT_TYPE)
.body(envelope);
if let Some(timeout) = total_timeout {
request = request.timeout(timeout);
if !plan.stream {
if let Some(timeout) = total_timeout {
request = request.timeout(timeout);
}
}
let first_byte_timeout = if plan.stream {
resolve_tunnel_first_byte_timeout(plan)
} else {
None
};
let started_at = Instant::now();
let response = request
.send()
let response = send_relay_request(request, first_byte_timeout)
.await
.map_err(|err| ExecutionRuntimeTransportError::RelayError(err.to_string()))?;
.map_err(ExecutionRuntimeTransportError::RelayError)?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
let status_code = response.status().as_u16();
let proxy_timing = response
@@ -869,6 +900,21 @@ async fn send_via_tunnel_relay(
Ok(response)
}
async fn send_relay_request(
request: reqwest::RequestBuilder,
first_byte_timeout: Option<Duration>,
) -> Result<reqwest::Response, String> {
if let Some(timeout) = first_byte_timeout {
return match tokio::time::timeout(timeout, request.send()).await {
Ok(Ok(response)) => Ok(response),
Ok(Err(error)) => Err(error.to_string()),
Err(_) => Err("tunnel relay first byte timeout".to_string()),
};
}
request.send().await.map_err(|err| err.to_string())
}
pub(crate) fn build_request_body(
plan: &ExecutionPlan,
) -> Result<Vec<u8>, ExecutionRuntimeTransportError> {
@@ -965,18 +1011,49 @@ fn resolve_tunnel_base_url_from_proxy(proxy: &ProxySnapshot) -> Option<String> {
}
fn resolve_relay_timeout_seconds(plan: &ExecutionPlan) -> u64 {
let ms = plan
.timeouts
resolve_tunnel_timeout_metadata(plan).legacy_timeout_secs
}
fn resolve_tunnel_first_byte_timeout(plan: &ExecutionPlan) -> Option<Duration> {
plan.stream.then(|| {
Duration::from_millis(
resolve_selected_tunnel_timeout_ms(plan).unwrap_or(DEFAULT_TUNNEL_TIMEOUT_MS),
)
})
}
fn resolve_tunnel_timeout_metadata(plan: &ExecutionPlan) -> TunnelTimeoutMetadata {
TunnelTimeoutMetadata {
request_timeout_ms: plan
.timeouts
.as_ref()
.and_then(|timeouts| timeouts.total_ms),
stream_first_byte_timeout_ms: plan
.timeouts
.as_ref()
.and_then(|timeouts| timeouts.first_byte_ms),
legacy_timeout_secs: timeout_ms_to_secs(
resolve_selected_tunnel_timeout_ms(plan).unwrap_or(DEFAULT_TUNNEL_TIMEOUT_MS),
),
}
}
fn resolve_selected_tunnel_timeout_ms(plan: &ExecutionPlan) -> Option<u64> {
plan.timeouts
.as_ref()
.and_then(|timeouts| {
timeouts
.read_ms
.or(timeouts.total_ms)
.or(timeouts.connect_ms)
if plan.stream {
timeouts.first_byte_ms.or(timeouts.total_ms)
} else {
timeouts.total_ms.or(timeouts.first_byte_ms)
}
})
.unwrap_or(60_000);
.map(|value| value.max(1))
}
fn timeout_ms_to_secs(ms: u64) -> u64 {
let secs = ms.div_ceil(1_000);
secs.clamp(1, 300)
secs.clamp(MIN_TUNNEL_TIMEOUT_SECS, MAX_TUNNEL_TIMEOUT_SECS)
}
fn resolve_tunnel_node_id(proxy: Option<&ProxySnapshot>) -> Option<String> {
@@ -1522,12 +1599,12 @@ mod tests {
use tokio::sync::watch;
use super::{
build_browser_wreq_client, build_client, build_execution_response_body,
build_request_headers, execute_sync_plan, record_manual_proxy_request_failure,
record_manual_proxy_request_outcome, record_manual_proxy_request_success,
record_manual_proxy_stream_error, resolve_execution_transport_controls,
response_body_is_json, DirectSyncExecutionRuntime, ExecutionRuntimeTransportError,
ExecutionTransportControls,
build_browser_wreq_client, build_client, build_direct_tunnel_request_meta,
build_execution_response_body, build_request_headers, execute_sync_plan,
record_manual_proxy_request_failure, record_manual_proxy_request_outcome,
record_manual_proxy_request_success, record_manual_proxy_stream_error,
resolve_execution_transport_controls, response_body_is_json, DirectSyncExecutionRuntime,
ExecutionRuntimeTransportError, ExecutionTransportControls,
};
use crate::constants::{
EXECUTION_RUNTIME_LOOP_GUARD_HEADER, EXECUTION_RUNTIME_LOOP_GUARD_VIA_TOKEN,
@@ -1629,6 +1706,64 @@ mod tests {
.is_none());
}
#[test]
fn tunnel_request_meta_uses_total_timeout_for_non_stream_requests() {
let plan = tunnel_timeout_plan(false);
let meta = build_direct_tunnel_request_meta(
&plan,
&reqwest::header::HeaderMap::new(),
ExecutionTransportControls::default(),
);
assert!(!meta.stream);
assert_eq!(meta.request_timeout_ms, Some(90_000));
assert_eq!(meta.stream_first_byte_timeout_ms, Some(12_345));
assert_eq!(meta.timeout, 90);
}
#[test]
fn tunnel_request_meta_uses_first_byte_timeout_for_stream_requests() {
let plan = tunnel_timeout_plan(true);
let meta = build_direct_tunnel_request_meta(
&plan,
&reqwest::header::HeaderMap::new(),
ExecutionTransportControls::default(),
);
assert!(meta.stream);
assert_eq!(meta.request_timeout_ms, Some(90_000));
assert_eq!(meta.stream_first_byte_timeout_ms, Some(12_345));
assert_eq!(meta.timeout, 13);
}
fn tunnel_timeout_plan(stream: bool) -> ExecutionPlan {
ExecutionPlan {
request_id: "req-timeout".into(),
candidate_id: None,
provider_name: Some("provider".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())]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({"model": "gpt-4.1"})),
stream,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
model_name: Some("gpt-4.1".into()),
proxy: None,
transport_profile: None,
timeouts: Some(ExecutionTimeouts {
total_ms: Some(90_000),
first_byte_ms: Some(12_345),
..ExecutionTimeouts::default()
}),
}
}
fn tunnel_proxy_snapshot(base_url: String) -> ProxySnapshot {
ProxySnapshot {
enabled: Some(true),
@@ -81,6 +81,133 @@ async fn admin_monitoring_trace_request_returns_local_payload() {
assert_eq!(payload["candidates"][0]["status_code"], json!(502));
}
#[tokio::test]
async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
sample_candidate(
"cand-used",
"trace-1",
0,
RequestCandidateStatus::Success,
Some(101),
Some(33),
Some(200),
),
]));
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider()],
vec![sample_endpoint()],
vec![sample_key()],
));
let mut usage = sample_usage(
"usage-request-1",
"provider-1",
"OpenAI",
40,
0.02,
"completed",
Some(200),
100,
);
usage.id = "usage-row-1".to_string();
usage.candidate_id = Some("cand-used".to_string());
usage.request_headers = Some(json!({
"x-trace-id": "trace-1"
}));
let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![usage]));
let data_state =
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
request_candidates,
usage_repository,
)
.with_provider_catalog_reader(provider_catalog);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
let context = request_context(
http::Method::GET,
"/api/admin/monitoring/trace/usage-row-1?attempted_only=true",
);
let response = local_monitoring_response(&state, &context)
.await
.expect("handler should not error")
.expect("route should be handled locally");
assert_eq!(response.status(), http::StatusCode::OK);
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!(payload["request_id"], json!("trace-1"));
assert_eq!(payload["candidates"][0]["id"], json!("cand-used"));
assert_eq!(
payload["candidates"][0]["extra_data"]["first_byte_time_ms"],
json!(30)
);
}
#[tokio::test]
async fn admin_monitoring_trace_request_resolves_usage_request_id_to_metadata_trace_id() {
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
sample_candidate(
"cand-used",
"trace-2",
0,
RequestCandidateStatus::Success,
Some(101),
Some(33),
Some(200),
),
]));
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider()],
vec![sample_endpoint()],
vec![sample_key()],
));
let mut usage = sample_usage(
"usage-request-2",
"provider-1",
"OpenAI",
40,
0.02,
"completed",
Some(200),
100,
);
usage.candidate_id = Some("cand-used".to_string());
usage.request_metadata = Some(json!({
"trace_id": "trace-2"
}));
let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![usage]));
let data_state =
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
request_candidates,
usage_repository,
)
.with_provider_catalog_reader(provider_catalog);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
let context = request_context(
http::Method::GET,
"/api/admin/monitoring/trace/usage-request-2",
);
let response = local_monitoring_response(&state, &context)
.await
.expect("handler should not error")
.expect("route should be handled locally");
assert_eq!(response.status(), http::StatusCode::OK);
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!(payload["request_id"], json!("trace-2"));
assert_eq!(payload["candidates"][0]["id"], json!("cand-used"));
}
#[tokio::test]
async fn admin_monitoring_trace_request_returns_oauth_account_label_from_auth_config() {
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
@@ -12,6 +12,7 @@ use aether_admin::observability::monitoring::{
use aether_data_contracts::repository::{
candidates::{DecisionTrace, RequestCandidateStatus},
provider_catalog::StoredProviderCatalogKey,
usage::StoredRequestUsageAudit,
};
use axum::{
body::Body,
@@ -21,12 +22,16 @@ use serde_json::{Map, Value};
use std::collections::BTreeMap;
use tracing::debug;
struct ResolvedAdminMonitoringTrace {
trace: DecisionTrace,
usage: Option<StoredRequestUsageAudit>,
}
pub(super) async fn build_admin_monitoring_trace_request_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
let admin_state = state;
let state = state.as_ref();
let Some(request_id) =
admin_monitoring_trace_request_id_from_path(&request_context.request_path)
else {
@@ -39,11 +44,8 @@ pub(super) async fn build_admin_monitoring_trace_request_response(
Err(detail) => return Ok(admin_monitoring_bad_request_response(detail)),
};
let Some(trace) = state
.data
.read_decision_trace(&request_id, attempted_only)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
let Some(resolved) =
resolve_admin_monitoring_trace(admin_state, &request_id, attempted_only).await?
else {
debug!(
event_name = "admin_monitoring_request_trace_not_found",
@@ -58,22 +60,113 @@ pub(super) async fn build_admin_monitoring_trace_request_response(
attempted_only,
));
};
let usage = state
.data
.read_request_usage_audit(&request_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let key_accounts = build_admin_monitoring_key_account_display_map(admin_state, &trace).await?;
let key_accounts =
build_admin_monitoring_key_account_display_map(admin_state, &resolved.trace).await?;
Ok(
build_admin_monitoring_trace_request_payload_response_with_key_accounts(
&trace,
usage.as_ref(),
&resolved.trace,
resolved.usage.as_ref(),
&key_accounts,
),
)
}
async fn resolve_admin_monitoring_trace(
state: &AdminAppState<'_>,
request_id: &str,
attempted_only: bool,
) -> Result<Option<ResolvedAdminMonitoringTrace>, GatewayError> {
let app = state.as_ref();
if let Some(trace) = app
.data
.read_decision_trace(request_id, attempted_only)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
{
let usage = app
.data
.read_request_usage_audit(request_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(Some(ResolvedAdminMonitoringTrace { trace, usage }));
}
let mut usage_candidates = Vec::new();
if let Some(usage) = app
.data
.read_request_usage_audit(request_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
{
usage_candidates.push(usage);
}
if let Some(usage) = state.find_request_usage_by_id(request_id).await? {
if !usage_candidates.iter().any(|item| item.id == usage.id) {
usage_candidates.push(usage);
}
}
for usage in usage_candidates {
for trace_request_id in admin_monitoring_usage_trace_request_ids(&usage) {
if trace_request_id == request_id {
continue;
}
if let Some(trace) = app
.data
.read_decision_trace(&trace_request_id, attempted_only)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
{
return Ok(Some(ResolvedAdminMonitoringTrace {
trace,
usage: Some(usage),
}));
}
}
}
Ok(None)
}
fn admin_monitoring_usage_trace_request_ids(usage: &StoredRequestUsageAudit) -> Vec<String> {
let mut ids = Vec::new();
push_non_empty_unique(&mut ids, usage.request_id.as_str());
if let Some(trace_id) = usage.trace_id() {
push_non_empty_unique(&mut ids, trace_id);
}
if let Some(trace_id) = usage_trace_id_from_headers(usage.request_headers.as_ref()) {
push_non_empty_unique(&mut ids, trace_id.as_str());
}
if let Some(trace_id) = usage_trace_id_from_headers(usage.provider_request_headers.as_ref()) {
push_non_empty_unique(&mut ids, trace_id.as_str());
}
ids
}
fn usage_trace_id_from_headers(headers: Option<&Value>) -> Option<String> {
let object = headers?.as_object()?;
object.iter().find_map(|(key, value)| {
key.eq_ignore_ascii_case(crate::constants::TRACE_ID_HEADER)
.then(|| {
value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty())
})
.flatten()
.map(ToOwned::to_owned)
})
}
fn push_non_empty_unique(values: &mut Vec<String>, value: &str) {
let value = value.trim();
if value.is_empty() || values.iter().any(|existing| existing == value) {
return;
}
values.push(value.to_string());
}
async fn build_admin_monitoring_key_account_display_map(
state: &AdminAppState<'_>,
trace: &DecisionTrace,
@@ -63,6 +63,10 @@ struct ProxyNodeRegisterRequest {
proxy_version: Option<String>,
#[serde(default)]
tunnel_mode: Option<bool>,
#[serde(default)]
tunnel_security: Option<String>,
#[serde(default)]
tunnel_encryption_key: Option<String>,
}
#[derive(Debug, Deserialize)]
@@ -349,9 +353,21 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response(
Ok(mutation) => mutation,
Err(response) => return Ok(Some(response)),
};
let tunnel_encryption_key = mutation
.proxy_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key"))
.and_then(|value| value.as_str())
.map(str::to_string);
let Some(node) = state.register_proxy_node(&mutation).await? else {
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
};
if let Some(key) = tunnel_encryption_key {
state
.app()
.tunnel
.register_secure_tunnel_key(node.id.clone(), key);
}
return Ok(Some(
Json(json!({
"node_id": node.id,
@@ -1358,6 +1374,9 @@ fn build_tunnel_probe_relay_envelope(
method: "GET".to_string(),
url: probe_url.trim().to_string(),
headers: std::collections::HashMap::new(),
stream: false,
request_timeout_ms: None,
stream_first_byte_timeout_ms: None,
timeout: timeout_secs,
follow_redirects: Some(false),
http1_only: false,
@@ -1420,12 +1439,43 @@ fn validate_register_request(
}
validate_optional_object(input.hardware_info.as_ref(), "hardware_info")?;
validate_optional_object(input.proxy_metadata.as_ref(), "proxy_metadata")?;
let tunnel_security =
normalize_optional_string(input.tunnel_security.as_deref(), "tunnel_security", 64)?;
let tunnel_encryption_key = normalize_optional_string(
input.tunnel_encryption_key.as_deref(),
"tunnel_encryption_key",
128,
)?;
let registered_by = request_context
.decision()
.and_then(|decision| decision.admin_principal.as_ref())
.map(|principal| principal.user_id.clone());
let mut proxy_metadata = input.proxy_metadata;
if tunnel_security.as_deref()
== Some(aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED)
{
let key = tunnel_encryption_key.as_deref().ok_or_else(|| {
bad_request_response(
"tunnel_encryption_key is required when tunnel_security=non_tls_required",
)
})?;
aether_contracts::tunnel_security::decode_psk(key)
.map_err(|err| bad_request_response(err.to_string()))?;
let mut metadata = proxy_metadata
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
metadata.insert(
"tunnel_security".to_string(),
json!({
"mode": aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED,
"encryption_key": key,
}),
);
proxy_metadata = Some(Value::Object(metadata));
}
Ok(
aether_data::repository::proxy_nodes::ProxyNodeRegistrationMutation {
name,
@@ -1438,7 +1488,7 @@ fn validate_register_request(
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_metadata,
proxy_version: normalize_optional_string(
input.proxy_version.as_deref(),
"proxy_version",
@@ -16,7 +16,7 @@ const INSTALL_SESSION_TTL_SECS: u64 = 15 * 60;
const INSTALL_SESSION_KEY_PREFIX: &str = "install:session:";
const TUNNEL_INSTALL_SESSION_KEY_PREFIX: &str = "tunnel-install:session:";
const TUNNEL_INSTALL_UNIX_SCRIPT_URL: &str =
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.sh";
"https://raw.githubusercontent.com/fawney19/Aether/refs/heads/main/apps/aether-tunnel/install.sh";
const TUNNEL_INSTALL_POWERSHELL_SCRIPT_URL: &str =
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.ps1";
@@ -59,6 +59,8 @@ struct StoredTunnelInstallSession {
aether_url: String,
management_token: String,
node_name: String,
tunnel_security: String,
tunnel_encryption_key: String,
expires_at_unix_secs: u64,
}
@@ -93,8 +95,7 @@ fn install_code_from_path(request_path: &str) -> Option<(String, bool)> {
fn tunnel_install_code_from_path(request_path: &str) -> Option<(String, bool)> {
let raw = request_path
.strip_prefix("/install-tunnel/")
.or_else(|| request_path.strip_prefix("/install-proxy/"))?
.strip_prefix("/install-tunnel/")?
.trim()
.trim_matches('/');
if raw.is_empty() || raw.contains('/') {
@@ -122,6 +123,17 @@ fn generate_install_code() -> String {
.collect()
}
fn generate_tunnel_encryption_key() -> String {
use base64::Engine;
let first = uuid::Uuid::new_v4();
let second = uuid::Uuid::new_v4();
let mut key = [0_u8; 32];
key[..16].copy_from_slice(first.as_bytes());
key[16..].copy_from_slice(second.as_bytes());
base64::engine::general_purpose::STANDARD.encode(key)
}
fn unix_secs_now() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
@@ -172,6 +184,8 @@ set -eu
export AETHER_TUNNEL_AETHER_URL={aether_url}
export AETHER_TUNNEL_MANAGEMENT_TOKEN={management_token}
export AETHER_TUNNEL_NODE_NAME={node_name}
export AETHER_TUNNEL_SECURITY={tunnel_security}
export AETHER_TUNNEL_ENCRYPTION_KEY={tunnel_encryption_key}
if command -v curl >/dev/null 2>&1; then
curl -fsSL {script_url} | sh
@@ -185,6 +199,8 @@ fi
aether_url = shell_single_quote(&session.aether_url),
management_token = shell_single_quote(&session.management_token),
node_name = shell_single_quote(&session.node_name),
tunnel_security = shell_single_quote(&session.tunnel_security),
tunnel_encryption_key = shell_single_quote(&session.tunnel_encryption_key),
script_url = shell_single_quote(TUNNEL_INSTALL_UNIX_SCRIPT_URL),
)
}
@@ -195,11 +211,15 @@ fn build_tunnel_powershell_script(session: &StoredTunnelInstallSession) -> Strin
$env:AETHER_TUNNEL_AETHER_URL = {aether_url}
$env:AETHER_TUNNEL_MANAGEMENT_TOKEN = {management_token}
$env:AETHER_TUNNEL_NODE_NAME = {node_name}
$env:AETHER_TUNNEL_SECURITY = {tunnel_security}
$env:AETHER_TUNNEL_ENCRYPTION_KEY = {tunnel_encryption_key}
irm {script_url} | iex
"###,
aether_url = powershell_single_quote(&session.aether_url),
management_token = powershell_single_quote(&session.management_token),
node_name = powershell_single_quote(&session.node_name),
tunnel_security = powershell_single_quote(&session.tunnel_security),
tunnel_encryption_key = powershell_single_quote(&session.tunnel_encryption_key),
script_url = powershell_single_quote(TUNNEL_INSTALL_POWERSHELL_SCRIPT_URL),
)
}
@@ -664,6 +684,8 @@ pub(crate) async fn build_proxy_node_install_session_response(
aether_url: base_url_from_request(headers, request_context),
management_token,
node_name,
tunnel_security: "non_tls_required".to_string(),
tunnel_encryption_key: generate_tunnel_encryption_key(),
expires_at_unix_secs,
};
let serialized = match serde_json::to_string(&session) {
@@ -712,9 +734,7 @@ pub(super) async fn maybe_build_local_install_response(
if decision.route_family.as_deref() != Some("install") {
return None;
}
if request_context.request_path.starts_with("/install-tunnel/")
|| request_context.request_path.starts_with("/install-proxy/")
{
if request_context.request_path.starts_with("/install-tunnel/") {
return Some(maybe_build_local_tunnel_install_response(state, request_context).await);
}
let Some((code, wants_powershell)) = install_code_from_path(&request_context.request_path)
@@ -893,6 +913,8 @@ mod tests {
aether_url: "https://aether.example".to_string(),
management_token: "ae-test-token".to_string(),
node_name: "jp-proxy-01".to_string(),
tunnel_security: "non_tls_required".to_string(),
tunnel_encryption_key: "base64-32-bytes".to_string(),
expires_at_unix_secs: u64::MAX,
}
}
@@ -907,10 +929,6 @@ mod tests {
tunnel_install_code_from_path("/install-tunnel/abc123.ps1"),
Some(("abc123".to_string(), true))
);
assert_eq!(
tunnel_install_code_from_path("/install-proxy/abc123"),
Some(("abc123".to_string(), false))
);
assert_eq!(tunnel_install_code_from_path("/install-tunnel/a/b"), None);
}
@@ -921,8 +939,10 @@ mod tests {
assert!(script.contains("export AETHER_TUNNEL_AETHER_URL='https://aether.example'"));
assert!(script.contains("export AETHER_TUNNEL_MANAGEMENT_TOKEN='ae-test-token'"));
assert!(script.contains("export AETHER_TUNNEL_NODE_NAME='jp-proxy-01'"));
assert!(script.contains("export AETHER_TUNNEL_SECURITY='non_tls_required'"));
assert!(script.contains("export AETHER_TUNNEL_ENCRYPTION_KEY='base64-32-bytes'"));
assert!(script.contains(
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.sh"
"https://raw.githubusercontent.com/fawney19/Aether/refs/heads/main/apps/aether-tunnel/install.sh"
));
assert!(!script.contains("aether-rust-pioneer"));
assert!(!script.contains("[[servers]]"));
@@ -935,6 +955,8 @@ mod tests {
assert!(script.contains("$env:AETHER_TUNNEL_AETHER_URL = 'https://aether.example'"));
assert!(script.contains("$env:AETHER_TUNNEL_MANAGEMENT_TOKEN = 'ae-test-token'"));
assert!(script.contains("$env:AETHER_TUNNEL_NODE_NAME = 'jp-proxy-01'"));
assert!(script.contains("$env:AETHER_TUNNEL_SECURITY = 'non_tls_required'"));
assert!(script.contains("$env:AETHER_TUNNEL_ENCRYPTION_KEY = 'base64-32-bytes'"));
assert!(script.contains(
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.ps1"
));
-1
View File
@@ -90,7 +90,6 @@ fn frontend_path_bypasses_static(path: &str) -> bool {
|| path.starts_with("/.well-known/")
|| path.starts_with("/install/")
|| path.starts_with("/install-tunnel/")
|| path.starts_with("/install-proxy/")
|| path.starts_with("/i/")
}
@@ -1307,6 +1307,9 @@ mod tests {
method: "GET".to_string(),
url: "https://example.com".to_string(),
headers: HashMap::new(),
stream: false,
request_timeout_ms: None,
stream_first_byte_timeout_ms: None,
timeout: 30,
follow_redirects: None,
http1_only: false,
@@ -24,6 +24,8 @@ use super::AppState;
pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error";
const MAX_RELAY_META_LEN: usize = 256 * 1024;
const MIN_RELAY_TIMEOUT_MS: u64 = 1;
const MAX_RELAY_TIMEOUT_MS: u64 = 300_000;
struct StreamGuard {
hub: std::sync::Arc<super::hub::HubRouter>,
@@ -92,7 +94,7 @@ pub(crate) async fn open_direct_relay_stream(
return Err(format!("connect: {error}"));
}
let wait_timeout = Duration::from_secs(meta.timeout.clamp(5, 300));
let wait_timeout = relay_header_timeout(&meta);
let response_head = match stream.wait_headers(wait_timeout).await {
Ok(response) => response,
Err(error) => {
@@ -148,6 +150,19 @@ fn map_request_admission_error(error: super::RequestAdmissionError) -> String {
}
}
fn relay_header_timeout(meta: &protocol::RequestMeta) -> Duration {
let timeout_ms = if meta.stream {
meta.stream_first_byte_timeout_ms
.or(meta.request_timeout_ms)
.unwrap_or_else(|| meta.timeout.saturating_mul(1_000))
} else {
meta.request_timeout_ms
.or(meta.stream_first_byte_timeout_ms)
.unwrap_or_else(|| meta.timeout.saturating_mul(1_000))
};
Duration::from_millis(timeout_ms.clamp(MIN_RELAY_TIMEOUT_MS, MAX_RELAY_TIMEOUT_MS))
}
pub async fn relay_request(
Path(node_id): Path<String>,
State(state): State<AppState>,
@@ -322,7 +337,7 @@ pub async fn relay_request(
finished: false,
};
let wait_timeout = Duration::from_secs(meta.timeout.clamp(5, 300));
let wait_timeout = relay_header_timeout(&meta);
let response_head = match stream.wait_headers(wait_timeout).await {
Ok(response) => response,
Err(error) => {
@@ -639,6 +654,9 @@ mod tests {
method: "GET".to_string(),
url: "https://example.com/health".to_string(),
headers: HashMap::new(),
stream: false,
request_timeout_ms: None,
stream_first_byte_timeout_ms: None,
timeout: 30,
follow_redirects: None,
http1_only: false,
@@ -754,6 +772,9 @@ mod tests {
method: "GET".to_string(),
url: "https://example.com/headers".to_string(),
headers: HashMap::new(),
stream: false,
request_timeout_ms: None,
stream_first_byte_timeout_ms: None,
timeout: 30,
follow_redirects: None,
http1_only: false,
@@ -17,6 +17,7 @@ use axum::http::HeaderMap;
use axum::response::{IntoResponse, Json};
use axum::routing::{get, post};
use axum::Router;
use dashmap::DashMap;
use tracing::warn;
use crate::{data::GatewayDataState, middleware};
@@ -34,6 +35,7 @@ pub struct AppState {
data: Arc<GatewayDataState>,
request_gate: Option<Arc<ConcurrencyGate>>,
distributed_request_gate: Option<Arc<RuntimeSemaphore>>,
secure_tunnel_keys: Arc<DashMap<String, String>>,
}
#[derive(Debug)]
@@ -55,9 +57,48 @@ impl AppState {
data: Arc::new(GatewayDataState::disabled()),
request_gate: None,
distributed_request_gate: None,
secure_tunnel_keys: Arc::new(DashMap::new()),
}
}
pub(crate) fn register_secure_tunnel_key(
&self,
node_id: impl Into<String>,
key: impl Into<String>,
) {
self.secure_tunnel_keys.insert(node_id.into(), key.into());
}
pub(crate) fn secure_tunnel_key(&self, node_id: &str) -> Option<String> {
self.secure_tunnel_keys
.get(node_id)
.map(|entry| entry.value().clone())
}
async fn secure_tunnel_key_for_node(&self, node_id: &str) -> Option<String> {
if let Some(key) = self.secure_tunnel_key(node_id) {
return Some(key);
}
let key = self
.data
.find_proxy_node(node_id)
.await
.ok()
.flatten()
.and_then(|node| {
node.proxy_metadata.and_then(|metadata| {
metadata
.pointer("/tunnel_security/encryption_key")
.and_then(|value| value.as_str())
.map(str::to_string)
})
});
if let Some(key) = key.as_ref() {
self.register_secure_tunnel_key(node_id.to_string(), key.clone());
}
key
}
pub(crate) fn with_data(mut self, data: Arc<GatewayDataState>) -> Self {
self.data = data;
self
@@ -212,11 +253,47 @@ pub async fn ws_proxy(
let max_streams = resolve_proxy_max_streams(&headers, state.max_streams);
let protocol_version = resolve_proxy_protocol_version(&headers);
let tunnel_security = headers
.get(aether_contracts::tunnel_security::TUNNEL_SECURITY_HEADER)
.and_then(|value| value.to_str().ok())
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string);
let security_session = headers
.get(aether_contracts::tunnel_security::TUNNEL_SECURITY_SESSION_HEADER)
.and_then(|value| value.to_str().ok())
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string);
if node_id.is_empty() {
warn!("proxy connection rejected: missing X-Node-ID header");
return axum::http::StatusCode::BAD_REQUEST.into_response();
}
let stored_security_key = state.secure_tunnel_key_for_node(&node_id).await;
let (security_key, security_session) = match tunnel_security.as_deref() {
Some(aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED) => {
match stored_security_key {
Some(key) => {
let Some(session) = security_session else {
warn!(node_id = %node_id, "secure tunnel requested without a security session");
return axum::http::StatusCode::BAD_REQUEST.into_response();
};
(Some(key), session)
}
None => {
warn!(node_id = %node_id, "secure tunnel requested but no PSK is registered");
return axum::http::StatusCode::UNAUTHORIZED.into_response();
}
}
}
Some(_) => return axum::http::StatusCode::BAD_REQUEST.into_response(),
None if stored_security_key.is_some() => {
warn!(node_id = %node_id, "proxy connection rejected: stored secure tunnel key requires encrypted frames");
return axum::http::StatusCode::UNAUTHORIZED.into_response();
}
None => (None, String::new()),
};
let request_permit = match state.try_acquire_request_permit().await {
Ok(permit) => permit,
@@ -253,6 +330,8 @@ pub async fn ws_proxy(
node_name,
max_streams,
protocol_version,
security_key,
security_session,
state.proxy_conn_cfg,
)
.await
@@ -13,6 +13,8 @@ use tracing::{debug, info, warn};
use super::hub::{ConnConfig, HubRouter, ProxyConn, SendStatus};
use super::protocol;
use aether_contracts::tunnel::Frame;
use aether_contracts::tunnel_security::{SecureFrameCodec, TunnelSecurityRole};
/// Maximum single frame size: 64 MB
const MAX_FRAME_SIZE: usize = 64 * 1024 * 1024;
@@ -24,6 +26,8 @@ pub async fn handle_proxy_connection(
node_name: String,
max_streams: usize,
protocol_version: u8,
security_key: Option<String>,
security_session: String,
cfg: ConnConfig,
) {
let conn_id = hub.alloc_conn_id();
@@ -31,6 +35,18 @@ pub async fn handle_proxy_connection(
let (tx, mut rx) = bounded_queue::<Message>(cfg.outbound_queue_capacity);
let (close_tx, mut close_rx) = watch::channel(false);
let security = match security_key.as_deref() {
Some(key) => {
match SecureFrameCodec::new(key, &security_session, TunnelSecurityRole::Server) {
Ok(codec) => Some(Arc::new(codec)),
Err(error) => {
warn!(conn_id, node_id = %node_id, error = %error, "secure tunnel codec initialization failed");
return;
}
}
}
None => None,
};
let conn = Arc::new(ProxyConn::new(
conn_id,
@@ -46,6 +62,7 @@ pub async fn handle_proxy_connection(
let writer_conn_id = conn_id;
let writer_conn = conn.clone();
let writer_security = security.clone();
let writer = tokio::spawn(async move {
let mut frames_sent: u64 = 0;
loop {
@@ -58,6 +75,13 @@ pub async fn handle_proxy_connection(
_ => 0,
};
let send_started_at = std::time::Instant::now();
let msg = match encrypt_message(msg, writer_security.as_deref()) {
Ok(msg) => msg,
Err(error) => {
warn!(conn_id = writer_conn_id, error = %error, "failed to encrypt outbound proxy frame");
break;
}
};
let send_result = tokio::time::timeout(
Duration::from_secs(15),
ws_tx.send(msg),
@@ -194,7 +218,7 @@ pub async fn handle_proxy_connection(
let reader_hub = hub.clone();
let reader_conn = conn.clone();
let reader = tokio::spawn(async move {
run_proxy_reader(ws_rx, reader_hub, reader_conn, cfg.idle_timeout).await;
run_proxy_reader(ws_rx, reader_hub, reader_conn, cfg.idle_timeout, security).await;
});
let _ = reader.await;
@@ -223,6 +247,7 @@ async fn run_proxy_reader(
hub: Arc<HubRouter>,
conn: Arc<ProxyConn>,
idle_timeout: Duration,
security: Option<Arc<SecureFrameCodec>>,
) {
let idle_enabled = !idle_timeout.is_zero();
let mut oversized_count = 0u32;
@@ -245,7 +270,14 @@ async fn run_proxy_reader(
match msg {
Some(Ok(Message::Binary(data))) => {
frames_received += 1;
let mut data = data.to_vec();
let mut data = match decrypt_message(data, security.as_deref()) {
Ok(data) => data,
Err(error) => {
warn!(conn_id = conn.id, error = %error, "failed to decrypt secure proxy frame");
conn.request_close();
break;
}
};
if data.len() > MAX_FRAME_SIZE {
oversized_count += 1;
warn!(
@@ -294,3 +326,33 @@ async fn run_proxy_reader(
}
}
}
fn encrypt_message(
msg: Message,
security: Option<&SecureFrameCodec>,
) -> Result<Message, aether_contracts::tunnel_security::TunnelSecurityError> {
let Some(codec) = security else {
return Ok(msg);
};
match msg {
Message::Binary(data) => {
let frame = Frame::decode(bytes::Bytes::from(data.to_vec()))
.map_err(|_| aether_contracts::tunnel_security::TunnelSecurityError::Encrypt)?;
Ok(Message::Binary(codec.encrypt_frame(frame)?))
}
other => Ok(other),
}
}
fn decrypt_message(
data: bytes::Bytes,
security: Option<&SecureFrameCodec>,
) -> Result<Vec<u8>, aether_contracts::tunnel_security::TunnelSecurityError> {
let Some(codec) = security else {
return Ok(data.to_vec());
};
let frame = Frame::decode(data)
.map_err(|_| aether_contracts::tunnel_security::TunnelSecurityError::Decrypt)?;
let frame = codec.decrypt_frame(frame)?;
Ok(frame.encode().to_vec())
}
+11
View File
@@ -455,6 +455,14 @@ impl EmbeddedTunnelState {
self.inner.clone()
}
pub(crate) fn register_secure_tunnel_key(
&self,
node_id: impl Into<String>,
key: impl Into<String>,
) {
self.inner.register_secure_tunnel_key(node_id, key);
}
pub(crate) fn has_local_proxy(&self, node_id: &str) -> bool {
self.inner.hub.has_local_proxy(node_id)
}
@@ -511,6 +519,9 @@ impl EmbeddedTunnelState {
method: "GET".to_string(),
url: url.trim().to_string(),
headers: HashMap::new(),
stream: false,
request_timeout_ms: None,
stream_first_byte_timeout_ms: None,
timeout: timeout_secs,
follow_redirects: Some(false),
http1_only: false,