mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
Fix OAuth refresh through tunnel proxy
This commit is contained in:
@@ -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 无效、已过期或已撤销,请重新登录授权"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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 无效、已过期或已撤销,请重新登录授权"
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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();
|
||||||
|
|||||||
@@ -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(¤t_method, &meta.headers);
|
let request_has_body = request_likely_has_body(¤t_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,
|
||||||
|
¤t_method,
|
||||||
|
Some(¤t_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)
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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 无效、已过期或已撤销,请重新登录授权',
|
||||||
|
)
|
||||||
|
})
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -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:')
|
||||||
|
|||||||
Reference in New Issue
Block a user