Fix OAuth refresh through tunnel proxy

This commit is contained in:
fawney19
2026-04-28 19:45:34 +08:00
parent 9194f78e56
commit 3d20c05ef7
7 changed files with 185 additions and 18 deletions

View File

@@ -98,6 +98,8 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message(
}
if error_code == "invalid_grant"
|| error_code == "invalid_refresh_token"
|| error_code == "refresh_token_expired"
|| lowered.contains("could not validate your refresh token")
|| (lowered.contains("refresh token")
&& ["expired", "revoked", "invalid"]
.iter()
@@ -143,3 +145,18 @@ pub(crate) fn merge_provider_oauth_refresh_failure_reason(
}
Some(refresh_reason.to_string())
}
#[cfg(test)]
mod tests {
use super::normalize_provider_oauth_refresh_error_message;
#[test]
fn normalizes_openai_refresh_token_expired_response() {
let body = r#"{"error":{"message":"Could not validate your refresh token. Please try signing in again.","type":"invalid_request_error","param":null,"code":"refresh_token_expired"}}"#;
assert_eq!(
normalize_provider_oauth_refresh_error_message(Some(401), Some(body)),
"refresh_token 无效、已过期或已撤销,请重新登录授权"
);
}
}

View File

@@ -232,6 +232,8 @@ fn normalize_local_oauth_refresh_error_message(
}
if error_code == "invalid_grant"
|| error_code == "invalid_refresh_token"
|| error_code == "refresh_token_expired"
|| lowered.contains("could not validate your refresh token")
|| (lowered.contains("refresh token")
&& ["expired", "revoked", "invalid"]
.iter()
@@ -1283,6 +1285,12 @@ impl AppState {
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(),
"true".to_string(),
);
if proxy_is_tunnel {
headers.insert(
EXECUTION_REQUEST_HTTP1_ONLY_HEADER.to_string(),
"true".to_string(),
);
}
let plan = ExecutionPlan {
request_id: request.request_id.to_string(),
candidate_id: None,
@@ -1665,4 +1673,14 @@ mod tests {
.expect("snapshot should exist");
assert!(!snapshot.provider.enable_format_conversion);
}
#[test]
fn normalizes_local_openai_refresh_token_expired_response() {
let body = r#"{"error":{"message":"Could not validate your refresh token. Please try signing in again.","type":"invalid_request_error","param":null,"code":"refresh_token_expired"}}"#;
assert_eq!(
super::normalize_local_oauth_refresh_error_message(Some(401), Some(body)),
"refresh_token 无效、已过期或已撤销,请重新登录授权"
);
}
}

View File

@@ -4712,8 +4712,9 @@ async fn gateway_refreshes_admin_provider_oauth_key_tunnel_proxy_with_direct_ref
assert_eq!(
refresh_plan
.headers
.get(EXECUTION_REQUEST_HTTP1_ONLY_HEADER),
None
.get(EXECUTION_REQUEST_HTTP1_ONLY_HEADER)
.map(String::as_str),
Some("true")
);
gateway_handle.abort();

View File

@@ -446,10 +446,7 @@ fn empty_request_body() -> upstream_client::UpstreamRequestBody {
}
fn buffered_request_body(body: Bytes) -> upstream_client::UpstreamRequestBody {
if body.is_empty() {
return empty_request_body();
}
upstream_client::stream_request_body(stream::once(async move { Ok(BodyFrame::data(body)) }))
upstream_client::full_request_body(body)
}
// Drain tunnel body frames on a detached task so the shared dispatcher is no
@@ -486,6 +483,63 @@ fn prepare_request_body(
}
}
async fn collect_request_body_for_replay(
mut body_rx: mpsc::Receiver<TunnelFrame>,
body_size: Arc<AtomicUsize>,
deadline: Instant,
replay_budget_bytes: usize,
) -> Result<Bytes, String> {
let mut body = BytesMut::new();
loop {
let frame = recv_body_frame_with_deadline(&mut body_rx, deadline).await?;
let Some(frame) = frame else {
return Ok(body.freeze());
};
match frame.msg_type {
MsgType::RequestBody => {
let end_stream = frame.is_end_stream();
let payload = decompress_if_gzip(&frame)
.map_err(|error| format!("gzip decompress failed: {error}"))?;
if !payload.is_empty() {
if body.len().saturating_add(payload.len()) > replay_budget_bytes {
return Err(format!(
"request body exceeds redirect replay budget: {} > {}",
body.len().saturating_add(payload.len()),
replay_budget_bytes
));
}
body_size.fetch_add(payload.len(), Ordering::Relaxed);
body.extend_from_slice(&payload);
}
if end_stream {
return Ok(body.freeze());
}
}
MsgType::StreamError => {
return Err(String::from_utf8(frame.payload.to_vec())
.unwrap_or_else(|_| "client cancelled request body".to_string()));
}
MsgType::StreamEnd => return Ok(body.freeze()),
_ => continue,
}
}
}
fn replay_body_from_buffered(body: Bytes, replay_budget_bytes: usize) -> ReplayableRequestBody {
let state = Arc::new(RequestBodyReplayState::new(
replay_budget_bytes.max(body.len()).max(1),
));
if !body.is_empty() {
state.push_chunk(body);
}
state.finish();
ReplayableRequestBody::Pending(state)
}
async fn recv_body_frame_with_deadline(
body_rx: &mut mpsc::Receiver<TunnelFrame>,
deadline: Instant,
@@ -781,6 +835,7 @@ async fn relay_upstream_response(
request_timing: upstream_client::RequestTiming,
request_body_size: &AtomicUsize,
redirect_count: usize,
request_body_mode: &'static str,
) -> Option<Duration> {
let status = response.status().as_u16();
let ttfb_ms = total_elapsed.as_millis() as u64;
@@ -803,6 +858,7 @@ async fn relay_upstream_response(
"timing_source": "instrumented_connector",
"total_ms": total_elapsed.as_millis() as u64,
"body_size": request_body_size.load(Ordering::Relaxed),
"request_body_mode": request_body_mode,
"mode": "tunnel",
"redirect_count": redirect_count,
});
@@ -1094,18 +1150,49 @@ async fn handle_stream_inner(
let timeout = Duration::from_secs(meta.timeout.clamp(MIN_TIMEOUT_SECS, MAX_TIMEOUT_SECS));
let request_body_size = Arc::new(AtomicUsize::new(0));
let request_has_body = request_likely_has_body(&current_method, &meta.headers);
let mut prepared_body = if request_has_body {
let replay_budget_bytes = if follow_redirects {
state.config.redirect_replay_budget_bytes
} else {
0
};
prepare_request_body(
let replay_budget_bytes = state.config.redirect_replay_budget_bytes;
let can_buffer_redirect_body = request_has_body && follow_redirects && replay_budget_bytes > 0;
let overall_start = Instant::now();
let request_body_mode = if can_buffer_redirect_body {
"buffered_fixed"
} else if request_has_body {
"streaming"
} else {
"empty"
};
let mut prepared_body = if can_buffer_redirect_body {
let buffered_body = match collect_request_body_for_replay(
body_rx,
Arc::clone(&request_body_size),
deadline,
replay_budget_bytes,
)
.await
{
Ok(body) => body,
Err(message) => {
log_stream_failure(
stream_log_context(
server,
stream_id,
&current_method,
Some(&current_url),
0,
request_body_size.load(Ordering::Relaxed),
),
&message,
overall_start.elapsed(),
);
send_error(frame_tx, stream_id, &message).await;
return None;
}
};
PreparedRequestBody {
first_request_body: Some(buffered_request_body(buffered_body.clone())),
replay_body: replay_body_from_buffered(buffered_body, replay_budget_bytes),
}
} else if request_has_body {
prepare_request_body(body_rx, Arc::clone(&request_body_size), deadline, 0)
} else {
PreparedRequestBody {
first_request_body: Some(build_streaming_request_body(
@@ -1120,7 +1207,6 @@ async fn handle_stream_inner(
}
};
let overall_start = Instant::now();
let mut total_dns_ms = 0u64;
let mut redirects_followed = 0usize;
let mut next_request_body = None::<upstream_client::UpstreamRequestBody>;
@@ -1200,6 +1286,7 @@ async fn handle_stream_inner(
response_ctx.request_timing,
request_body_size.as_ref(),
redirects_followed,
request_body_mode,
)
.await;
}
@@ -1236,6 +1323,7 @@ async fn handle_stream_inner(
response_ctx.request_timing,
request_body_size.as_ref(),
redirects_followed,
request_body_mode,
)
.await;
}
@@ -1288,6 +1376,7 @@ async fn handle_stream_inner(
response_ctx.request_timing,
request_body_size.as_ref(),
redirects_followed,
request_body_mode,
)
.await;
}
@@ -1412,7 +1501,7 @@ mod tests {
use aether_runtime::{ConcurrencyGate, DistributedConcurrencyGate};
use arc_swap::ArcSwap;
use axum::body::Body;
use axum::http::{header, Response, StatusCode};
use axum::http::{header, HeaderMap, Response, StatusCode};
use axum::routing::{get, post};
use axum::Router;
use futures_util::Sink;
@@ -1800,7 +1889,14 @@ mod tests {
let app = Router::new()
.route(
"/start",
post(|body: Bytes| async move {
post(|headers: HeaderMap, body: Bytes| async move {
assert_eq!(
headers
.get(header::CONTENT_LENGTH)
.and_then(|value| value.to_str().ok()),
Some("5")
);
assert!(headers.get(header::TRANSFER_ENCODING).is_none());
assert_eq!(body, Bytes::from_static(b"hello"));
Response::builder()
.status(StatusCode::TEMPORARY_REDIRECT)
@@ -1811,7 +1907,14 @@ mod tests {
)
.route(
"/final",
post(|body: Bytes| async move {
post(|headers: HeaderMap, body: Bytes| async move {
assert_eq!(
headers
.get(header::CONTENT_LENGTH)
.and_then(|value| value.to_str().ok()),
Some("5")
);
assert!(headers.get(header::TRANSFER_ENCODING).is_none());
assert_eq!(body, Bytes::from_static(b"hello"));
Response::builder()
.status(StatusCode::OK)

View File

@@ -1,3 +1,4 @@
use std::convert::Infallible;
use std::future::Future;
use std::io;
use std::net::IpAddr;
@@ -9,7 +10,7 @@ use std::time::Duration;
use bytes::Bytes;
use futures_util::Stream;
use http_body_util::combinators::UnsyncBoxBody;
use http_body_util::{BodyExt, StreamBody};
use http_body_util::{BodyExt, Full, StreamBody};
use hyper::body::Frame;
use hyper::rt;
use hyper::Response;
@@ -43,6 +44,12 @@ where
StreamBody::new(stream).boxed_unsync()
}
pub fn full_request_body(body: Bytes) -> UpstreamRequestBody {
Full::new(body)
.map_err(|err: Infallible| match err {})
.boxed_unsync()
}
#[derive(Clone, Copy, Debug, Default)]
pub struct ConnectTiming {
pub connect_ms: u64,