fix: align oauth tunnel refresh transport

This commit is contained in:
fawney19
2026-04-28 01:10:34 +08:00
parent 6426b5d80e
commit 9a67497d8c
3 changed files with 433 additions and 10 deletions

View File

@@ -345,6 +345,22 @@ async fn execute_sync_plan_via_local_tunnel(
plan.content_encoding.as_deref(),
plan.body.body_bytes_b64.is_some(),
)?;
let timeout_secs = resolve_relay_timeout_seconds(plan);
tracing::info!(
request_id = %plan.request_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
method = %plan.method,
upstream_host = %execution_log_url_host(plan.url.as_str()),
node_id = %node_id,
path = "local_tunnel",
body_bytes_len = body_bytes.len(),
timeout_secs,
follow_redirects = ?transport_controls.follow_redirects,
http1_only = transport_controls.http1_only,
"gateway execution runtime local tunnel request prepared"
);
let started_at = Instant::now();
let mut response = state
.tunnel
@@ -358,6 +374,7 @@ async fn execute_sync_plan_via_local_tunnel(
let ttfb_ms = started_at.elapsed().as_millis() as u64;
let status_code = response.status();
let headers = collect_tunnel_response_headers(response.headers());
let proxy_timing = execution_header_for_log(&headers, "x-proxy-timing").unwrap_or("-");
let mut body_bytes = Vec::new();
while let Some(chunk) = response
.next_chunk()
@@ -370,6 +387,39 @@ async fn execute_sync_plan_via_local_tunnel(
decode_response_body_bytes(&headers, &body_bytes).unwrap_or_else(|| body_bytes.clone());
let elapsed_ms = started_at.elapsed().as_millis() as u64;
let upstream_bytes = body_bytes.len() as u64;
if status_code >= 400 {
tracing::warn!(
request_id = %plan.request_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
method = %plan.method,
upstream_host = %execution_log_url_host(plan.url.as_str()),
node_id = %node_id,
path = "local_tunnel",
status_code,
elapsed_ms,
upstream_bytes,
proxy_timing,
"gateway execution runtime local tunnel response returned error"
);
} else {
tracing::info!(
request_id = %plan.request_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
method = %plan.method,
upstream_host = %execution_log_url_host(plan.url.as_str()),
node_id = %node_id,
path = "local_tunnel",
status_code,
elapsed_ms,
upstream_bytes,
proxy_timing,
"gateway execution runtime local tunnel response received"
);
}
let body = if body_bytes.is_empty() {
None
@@ -483,17 +533,35 @@ 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 envelope = build_relay_envelope(
RelayRequestMeta {
method: method.as_str().to_string(),
url: plan.url.clone(),
headers: header_map_to_string_map(&headers),
timeout: resolve_relay_timeout_seconds(plan),
timeout: timeout_secs,
follow_redirects: transport_controls.follow_redirects,
http1_only: transport_controls.http1_only,
},
&body_bytes,
)?;
tracing::info!(
request_id = %plan.request_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
method = %method,
upstream_host = %execution_log_url_host(plan.url.as_str()),
relay_host = %execution_log_url_host(relay_url.as_str()),
node_id,
path = "tunnel_relay",
body_bytes_len = body_bytes.len(),
envelope_bytes_len = envelope.len(),
timeout_secs,
follow_redirects = ?transport_controls.follow_redirects,
http1_only = transport_controls.http1_only,
"gateway execution runtime tunnel relay request prepared"
);
let mut request = client
.request(reqwest::Method::POST, relay_url)
@@ -503,10 +571,49 @@ async fn send_via_tunnel_relay(
request = request.timeout(timeout);
}
let started_at = Instant::now();
let response = request
.send()
.await
.map_err(|err| ExecutionRuntimeTransportError::RelayError(err.to_string()))?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
let status_code = response.status().as_u16();
let proxy_timing = response
.headers()
.get("x-proxy-timing")
.and_then(|value| value.to_str().ok())
.unwrap_or("-");
if status_code >= 400 {
tracing::warn!(
request_id = %plan.request_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
method = %method,
upstream_host = %execution_log_url_host(plan.url.as_str()),
node_id,
path = "tunnel_relay",
status_code,
elapsed_ms,
proxy_timing,
"gateway execution runtime tunnel relay response returned error"
);
} else {
tracing::info!(
request_id = %plan.request_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
method = %method,
upstream_host = %execution_log_url_host(plan.url.as_str()),
node_id,
path = "tunnel_relay",
status_code,
elapsed_ms,
proxy_timing,
"gateway execution runtime tunnel relay response received"
);
}
if let Some(kind) = response
.headers()
@@ -514,6 +621,20 @@ async fn send_via_tunnel_relay(
.and_then(|value| value.to_str().ok())
.map(str::to_owned)
{
tracing::warn!(
request_id = %plan.request_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
method = %method,
upstream_host = %execution_log_url_host(plan.url.as_str()),
node_id,
path = "tunnel_relay",
status_code,
elapsed_ms,
error_kind = %kind,
"gateway execution runtime tunnel relay returned relay error"
);
let message = response
.text()
.await
@@ -884,6 +1005,23 @@ fn collect_tunnel_response_headers(headers: &[(String, String)]) -> BTreeMap<Str
.collect()
}
fn execution_header_for_log<'a>(
headers: &'a BTreeMap<String, String>,
name: &str,
) -> Option<&'a str> {
headers
.iter()
.find(|(header_name, _)| header_name.eq_ignore_ascii_case(name))
.map(|(_, value)| value.as_str())
}
fn execution_log_url_host(url: &str) -> String {
url::Url::parse(url)
.ok()
.and_then(|url| url.host_str().map(ToOwned::to_owned))
.unwrap_or_else(|| "-".to_string())
}
fn decode_response_body_bytes(
headers: &BTreeMap<String, String>,
body_bytes: &[u8],

View File

@@ -10,7 +10,8 @@ use super::super::provider_transport;
use crate::provider_key_auth::provider_key_is_oauth_managed;
use aether_admin::provider::quota as admin_provider_quota_pure;
use aether_contracts::{
ExecutionPlan, ExecutionTimeouts, RequestBody, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
ExecutionPlan, ExecutionTimeouts, ProxySnapshot, RequestBody,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use base64::{engine::general_purpose::STANDARD, Engine as _};
@@ -70,6 +71,69 @@ fn oauth_metadata_refresh_token_fingerprint(metadata: Option<&Value>) -> Option<
.map(secret_fingerprint)
}
fn local_oauth_request_refresh_token_fingerprint(
request: &provider_transport::LocalOAuthHttpRequest,
) -> (Option<String>, Option<usize>) {
if let Some(json_body) = request.json_body.as_ref() {
return json_body
.as_object()
.and_then(|object| object.get("refresh_token"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| (Some(secret_fingerprint(value)), Some(value.len())))
.unwrap_or((None, None));
}
let Some(body_bytes) = request.body_bytes.as_ref() else {
return (None, None);
};
for (key, value) in url::form_urlencoded::parse(body_bytes) {
if key == "refresh_token" {
let value = value.trim();
if !value.is_empty() {
return (Some(secret_fingerprint(value)), Some(value.len()));
}
}
}
(None, None)
}
fn local_oauth_log_excerpt(body: &str) -> String {
let body = body.trim();
if body.is_empty() {
return "-".to_string();
}
body.chars().take(300).collect()
}
fn local_oauth_proxy_is_tunnel(proxy: Option<&ProxySnapshot>) -> bool {
let Some(proxy) = proxy else {
return false;
};
if proxy.enabled == Some(false) {
return false;
}
proxy
.mode
.as_deref()
.map(str::trim)
.is_some_and(|mode| mode.eq_ignore_ascii_case("tunnel"))
}
fn local_oauth_proxy_extra_string<'a>(
proxy: Option<&'a ProxySnapshot>,
key: &str,
) -> Option<&'a str> {
proxy?
.extra
.as_ref()
.and_then(|extra| extra.get(key))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
}
fn secret_fingerprint(value: &str) -> String {
let digest = Sha256::digest(value.as_bytes());
let mut fingerprint = String::with_capacity(16);
@@ -1188,11 +1252,17 @@ impl AppState {
body_ref: None,
}
};
let proxy_snapshot = self
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await;
let proxy_is_tunnel = local_oauth_proxy_is_tunnel(proxy_snapshot.as_ref());
let mut headers = request.headers.clone();
headers.insert(
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(),
"true".to_string(),
);
if !proxy_is_tunnel {
headers.insert(
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(),
"true".to_string(),
);
}
let plan = ExecutionPlan {
request_id: request.request_id.to_string(),
candidate_id: None,
@@ -1216,9 +1286,7 @@ impl AppState {
client_api_format: "provider_oauth:local_refresh".to_string(),
provider_api_format: "provider_oauth:local_refresh".to_string(),
model_name: Some(provider_type.to_string()),
proxy: self
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await,
proxy: proxy_snapshot,
tls_profile: None,
timeouts: Some(ExecutionTimeouts {
connect_ms: Some(LOCAL_OAUTH_HTTP_TIMEOUT_MS),
@@ -1229,6 +1297,50 @@ impl AppState {
..ExecutionTimeouts::default()
}),
};
let (request_refresh_token_fingerprint, request_refresh_token_len) =
local_oauth_request_refresh_token_fingerprint(request);
tracing::info!(
key_id = %transport.key.id,
provider_id = %transport.provider.id,
endpoint_id = %transport.endpoint.id,
provider_type,
request_id = %request.request_id,
method = %plan.method,
token_url = %plan.url,
content_type = plan.content_type.as_deref().unwrap_or("-"),
body_bytes_len = ?request.body_bytes.as_ref().map(Vec::len),
json_body_present = request.json_body.is_some(),
request_refresh_token_fingerprint = request_refresh_token_fingerprint
.as_deref()
.unwrap_or("-"),
request_refresh_token_len = ?request_refresh_token_len,
proxy_node_id = ?plan.proxy.as_ref().and_then(|proxy| proxy.node_id.as_deref()),
proxy_mode = plan.proxy.as_ref().and_then(|proxy| proxy.mode.as_deref()).unwrap_or("-"),
proxy_enabled = ?plan.proxy.as_ref().and_then(|proxy| proxy.enabled),
proxy_url_present = plan
.proxy
.as_ref()
.and_then(|proxy| proxy.url.as_deref())
.map(str::trim)
.is_some_and(|value| !value.is_empty()),
proxy_is_tunnel,
tunnel_base_url_present = local_oauth_proxy_extra_string(
plan.proxy.as_ref(),
"tunnel_base_url"
)
.is_some(),
tunnel_owner_instance_id = local_oauth_proxy_extra_string(
plan.proxy.as_ref(),
"tunnel_owner_instance_id"
)
.unwrap_or("-"),
follow_redirects = plan
.headers
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)
.map(String::as_str)
.unwrap_or("-"),
"gateway local oauth execution request prepared"
);
let result =
crate::execution_runtime::execute_execution_runtime_sync_plan(self, None, &plan)
.await
@@ -1242,9 +1354,38 @@ impl AppState {
},
},
)?;
let response_body_text = local_oauth_execution_body_text(&result);
if (200..300).contains(&result.status_code) {
tracing::info!(
key_id = %transport.key.id,
provider_id = %transport.provider.id,
endpoint_id = %transport.endpoint.id,
provider_type,
request_id = %request.request_id,
status_code = result.status_code,
request_refresh_token_fingerprint = request_refresh_token_fingerprint
.as_deref()
.unwrap_or("-"),
"gateway local oauth execution response received"
);
} else {
tracing::warn!(
key_id = %transport.key.id,
provider_id = %transport.provider.id,
endpoint_id = %transport.endpoint.id,
provider_type,
request_id = %request.request_id,
status_code = result.status_code,
request_refresh_token_fingerprint = request_refresh_token_fingerprint
.as_deref()
.unwrap_or("-"),
body_excerpt = %local_oauth_log_excerpt(response_body_text.as_str()),
"gateway local oauth execution response returned error"
);
}
Ok(provider_transport::LocalOAuthHttpResponse {
status_code: result.status_code,
body_text: local_oauth_execution_body_text(&result),
body_text: response_body_text,
})
}

View File

@@ -4568,6 +4568,150 @@ async fn gateway_refreshes_admin_provider_oauth_key_locally_via_execution_runtim
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_refreshes_admin_provider_oauth_key_tunnel_proxy_without_follow_redirects_like_master(
) {
let execution_plans = Arc::new(Mutex::new(Vec::<ExecutionPlan>::new()));
let execution_plans_clone = Arc::clone(&execution_plans);
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any(move |Json(plan): Json<ExecutionPlan>| {
let execution_plans_inner = Arc::clone(&execution_plans_clone);
async move {
execution_plans_inner
.lock()
.expect("mutex should lock")
.push(plan.clone());
if plan.request_id == "provider-oauth:local-refresh-token" {
Json(json!({
"request_id": plan.request_id,
"status_code": 200,
"headers": {
"content-type": "application/json"
},
"body": {
"json_body": {
"access_token": "refreshed-codex-access-token",
"refresh_token": "refreshed-codex-refresh-token",
"token_type": "Bearer",
"expires_in": 1800,
"scope": "openid email profile offline_access",
"email": "alice@example.com",
"account_id": "acct-codex-123",
"plan_type": "plus"
}
}
}))
} else {
Json(json!({
"request_id": plan.request_id,
"status_code": 200,
"headers": {
"content-type": "application/json"
},
"body": {
"json_body": {}
}
}))
}
}
}),
);
let mut provider = sample_provider("provider-codex", "codex", 10);
provider.provider_type = "codex".to_string();
provider.proxy = Some(json!({
"mode": "tunnel",
"node_id": "proxy-node-tunnel",
"enabled": true,
"tunnel_base_url": "http://gateway-owner.internal"
}));
let endpoint = sample_endpoint(
"endpoint-codex-cli",
"provider-codex",
"openai:responses",
"https://chatgpt.com/backend-api/codex",
);
let mut key = sample_key(
"key-codex-oauth-refresh-tunnel",
"provider-codex",
"openai:responses",
"stale-codex-access-token",
);
key.auth_type = "oauth".to_string();
key.encrypted_auth_config = Some(
encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"provider_type":"codex","refresh_token":"old-codex-refresh-token","email":"alice@example.com","account_id":"acct-codex-123","plan_type":"plus","expires_at":1}"#,
)
.expect("auth config ciphertext should build"),
);
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
let oauth_refresh =
crate::provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![
Arc::new(
crate::provider_transport::oauth_refresh::GenericOAuthRefreshAdapter::default()
.with_token_url_for_tests("codex", "https://oauth.example/oauth/token"),
),
]);
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository,
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
.with_oauth_refresh_coordinator_for_tests(oauth_refresh),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/provider-oauth/keys/key-codex-oauth-refresh-tunnel/refresh"
))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let plans = execution_plans.lock().expect("mutex should lock");
let refresh_plan = plans
.iter()
.find(|plan| plan.request_id == "provider-oauth:local-refresh-token")
.expect("local refresh plan should exist");
assert_eq!(
refresh_plan
.proxy
.as_ref()
.and_then(|proxy| proxy.mode.as_deref()),
Some("tunnel")
);
assert_eq!(
refresh_plan
.headers
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER),
None
);
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_consecutive_manual_oauth_refresh_uses_rotated_refresh_token() {
let refresh_request_bodies = Arc::new(Mutex::new(Vec::<String>::new()));