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" if error_code == "invalid_grant"
|| error_code == "invalid_refresh_token" || error_code == "invalid_refresh_token"
|| error_code == "refresh_token_expired"
|| lowered.contains("could not validate your refresh token")
|| (lowered.contains("refresh token") || (lowered.contains("refresh token")
&& ["expired", "revoked", "invalid"] && ["expired", "revoked", "invalid"]
.iter() .iter()
@@ -143,3 +145,18 @@ pub(crate) fn merge_provider_oauth_refresh_failure_reason(
} }
Some(refresh_reason.to_string()) 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" if error_code == "invalid_grant"
|| error_code == "invalid_refresh_token" || error_code == "invalid_refresh_token"
|| error_code == "refresh_token_expired"
|| lowered.contains("could not validate your refresh token")
|| (lowered.contains("refresh token") || (lowered.contains("refresh token")
&& ["expired", "revoked", "invalid"] && ["expired", "revoked", "invalid"]
.iter() .iter()
@@ -1283,6 +1285,12 @@ impl AppState {
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(), EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(),
"true".to_string(), "true".to_string(),
); );
if proxy_is_tunnel {
headers.insert(
EXECUTION_REQUEST_HTTP1_ONLY_HEADER.to_string(),
"true".to_string(),
);
}
let plan = ExecutionPlan { let plan = ExecutionPlan {
request_id: request.request_id.to_string(), request_id: request.request_id.to_string(),
candidate_id: None, candidate_id: None,
@@ -1665,4 +1673,14 @@ mod tests {
.expect("snapshot should exist"); .expect("snapshot should exist");
assert!(!snapshot.provider.enable_format_conversion); 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!( assert_eq!(
refresh_plan refresh_plan
.headers .headers
.get(EXECUTION_REQUEST_HTTP1_ONLY_HEADER), .get(EXECUTION_REQUEST_HTTP1_ONLY_HEADER)
None .map(String::as_str),
Some("true")
); );
gateway_handle.abort(); 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 { fn buffered_request_body(body: Bytes) -> upstream_client::UpstreamRequestBody {
if body.is_empty() { upstream_client::full_request_body(body)
return empty_request_body();
}
upstream_client::stream_request_body(stream::once(async move { Ok(BodyFrame::data(body)) }))
} }
// Drain tunnel body frames on a detached task so the shared dispatcher is no // 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( async fn recv_body_frame_with_deadline(
body_rx: &mut mpsc::Receiver<TunnelFrame>, body_rx: &mut mpsc::Receiver<TunnelFrame>,
deadline: Instant, deadline: Instant,
@@ -781,6 +835,7 @@ async fn relay_upstream_response(
request_timing: upstream_client::RequestTiming, request_timing: upstream_client::RequestTiming,
request_body_size: &AtomicUsize, request_body_size: &AtomicUsize,
redirect_count: usize, redirect_count: usize,
request_body_mode: &'static str,
) -> Option<Duration> { ) -> Option<Duration> {
let status = response.status().as_u16(); let status = response.status().as_u16();
let ttfb_ms = total_elapsed.as_millis() as u64; let ttfb_ms = total_elapsed.as_millis() as u64;
@@ -803,6 +858,7 @@ async fn relay_upstream_response(
"timing_source": "instrumented_connector", "timing_source": "instrumented_connector",
"total_ms": total_elapsed.as_millis() as u64, "total_ms": total_elapsed.as_millis() as u64,
"body_size": request_body_size.load(Ordering::Relaxed), "body_size": request_body_size.load(Ordering::Relaxed),
"request_body_mode": request_body_mode,
"mode": "tunnel", "mode": "tunnel",
"redirect_count": redirect_count, "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 timeout = Duration::from_secs(meta.timeout.clamp(MIN_TIMEOUT_SECS, MAX_TIMEOUT_SECS));
let request_body_size = Arc::new(AtomicUsize::new(0)); let request_body_size = Arc::new(AtomicUsize::new(0));
let request_has_body = request_likely_has_body(&current_method, &meta.headers); let request_has_body = request_likely_has_body(&current_method, &meta.headers);
let mut prepared_body = if request_has_body { let replay_budget_bytes = state.config.redirect_replay_budget_bytes;
let replay_budget_bytes = if follow_redirects { let can_buffer_redirect_body = request_has_body && follow_redirects && replay_budget_bytes > 0;
state.config.redirect_replay_budget_bytes let overall_start = Instant::now();
} else { let request_body_mode = if can_buffer_redirect_body {
0 "buffered_fixed"
}; } else if request_has_body {
prepare_request_body( "streaming"
} else {
"empty"
};
let mut prepared_body = if can_buffer_redirect_body {
let buffered_body = match collect_request_body_for_replay(
body_rx, body_rx,
Arc::clone(&request_body_size), Arc::clone(&request_body_size),
deadline, deadline,
replay_budget_bytes, 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 { } else {
PreparedRequestBody { PreparedRequestBody {
first_request_body: Some(build_streaming_request_body( 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 total_dns_ms = 0u64;
let mut redirects_followed = 0usize; let mut redirects_followed = 0usize;
let mut next_request_body = None::<upstream_client::UpstreamRequestBody>; let mut next_request_body = None::<upstream_client::UpstreamRequestBody>;
@@ -1200,6 +1286,7 @@ async fn handle_stream_inner(
response_ctx.request_timing, response_ctx.request_timing,
request_body_size.as_ref(), request_body_size.as_ref(),
redirects_followed, redirects_followed,
request_body_mode,
) )
.await; .await;
} }
@@ -1236,6 +1323,7 @@ async fn handle_stream_inner(
response_ctx.request_timing, response_ctx.request_timing,
request_body_size.as_ref(), request_body_size.as_ref(),
redirects_followed, redirects_followed,
request_body_mode,
) )
.await; .await;
} }
@@ -1288,6 +1376,7 @@ async fn handle_stream_inner(
response_ctx.request_timing, response_ctx.request_timing,
request_body_size.as_ref(), request_body_size.as_ref(),
redirects_followed, redirects_followed,
request_body_mode,
) )
.await; .await;
} }
@@ -1412,7 +1501,7 @@ mod tests {
use aether_runtime::{ConcurrencyGate, DistributedConcurrencyGate}; use aether_runtime::{ConcurrencyGate, DistributedConcurrencyGate};
use arc_swap::ArcSwap; use arc_swap::ArcSwap;
use axum::body::Body; use axum::body::Body;
use axum::http::{header, Response, StatusCode}; use axum::http::{header, HeaderMap, Response, StatusCode};
use axum::routing::{get, post}; use axum::routing::{get, post};
use axum::Router; use axum::Router;
use futures_util::Sink; use futures_util::Sink;
@@ -1800,7 +1889,14 @@ mod tests {
let app = Router::new() let app = Router::new()
.route( .route(
"/start", "/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")); assert_eq!(body, Bytes::from_static(b"hello"));
Response::builder() Response::builder()
.status(StatusCode::TEMPORARY_REDIRECT) .status(StatusCode::TEMPORARY_REDIRECT)
@@ -1811,7 +1907,14 @@ mod tests {
) )
.route( .route(
"/final", "/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")); assert_eq!(body, Bytes::from_static(b"hello"));
Response::builder() Response::builder()
.status(StatusCode::OK) .status(StatusCode::OK)

View File

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

View File

@@ -30,4 +30,18 @@ describe('errorParser', () => {
'Token 刷新失败refresh_token 已被使用并轮换,请重新登录授权', 'Token 刷新失败refresh_token 已被使用并轮换,请重新登录授权',
) )
}) })
it('normalizes expired refresh token errors', () => {
const error = {
response: {
data: {
detail: '{"error":{"message":"Could not validate your refresh token. Please try signing in again.","type":"invalid_request_error","code":"refresh_token_expired"}}',
},
},
}
expect(parseApiError(error, 'Token 刷新失败')).toBe(
'Token 刷新失败refresh_token 无效、已过期或已撤销,请重新登录授权',
)
})
}) })

View File

@@ -145,6 +145,13 @@ function normalizeKnownApiErrorMessage(message: string): string {
return 'Token 刷新失败refresh_token 已被使用并轮换,请重新登录授权' return 'Token 刷新失败refresh_token 已被使用并轮换,请重新登录授权'
} }
if (
lowered.includes('refresh_token_expired')
|| lowered.includes('could not validate your refresh token')
) {
return 'Token 刷新失败refresh_token 无效、已过期或已撤销,请重新登录授权'
}
if ( if (
lowered.includes('token refresh 失败:') lowered.includes('token refresh 失败:')
|| lowered.includes('token refresh failed:') || lowered.includes('token refresh failed:')