mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat(tunnel): 请求体流式传输 & OpenAI CLI 请求 key 排序优化
Hub 端: - open_local_stream 不再接收 body 参数,改为通过 push_local_request_body 分块推送 - 请求体按 32KB 分帧发送,避免大请求一次性压缩和传输 - local_relay 改为流式解析 envelope 和转发请求体 Proxy 端: - stream_handler 改为流式传输请求体到上游,不再预先收集完整 body - upstream_client 请求体类型从 Full<Bytes> 改为 UnsyncBoxBody 以支持流式传输 - dispatcher 将 StreamEnd/StreamError 事件转发给 stream handler Python 端: - hub_transport relay envelope 改为异步生成器流式发送 - 提取 reorder_openai_cli_request_prefix_keys 为公共函数 - Codex passthrough 路径也应用稳定的前缀 key 排序
This commit is contained in:
3
aether-hub/Cargo.lock
generated
3
aether-hub/Cargo.lock
generated
@@ -10,7 +10,7 @@ checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa"
|
||||
|
||||
[[package]]
|
||||
name = "aether-hub"
|
||||
version = "0.1.8"
|
||||
version = "0.1.9"
|
||||
dependencies = [
|
||||
"async-stream",
|
||||
"axum",
|
||||
@@ -19,6 +19,7 @@ dependencies = [
|
||||
"dashmap",
|
||||
"flate2",
|
||||
"futures-util",
|
||||
"http-body-util",
|
||||
"parking_lot",
|
||||
"reqwest",
|
||||
"serde",
|
||||
|
||||
@@ -18,6 +18,7 @@ flate2 = "1"
|
||||
futures-util = "0.3"
|
||||
bytes = "1"
|
||||
async-stream = "0.3"
|
||||
http-body-util = "0.1"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
|
||||
|
||||
[profile.release]
|
||||
|
||||
@@ -15,6 +15,8 @@ use tracing::{debug, info, warn};
|
||||
use crate::control_plane::ControlPlaneClient;
|
||||
use crate::protocol;
|
||||
|
||||
const MAX_REQUEST_BODY_FRAME_SIZE: usize = 32 * 1024;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum SendStatus {
|
||||
Queued,
|
||||
@@ -421,7 +423,6 @@ impl HubRouter {
|
||||
&self,
|
||||
node_id: &str,
|
||||
meta: &protocol::RequestMeta,
|
||||
body: Bytes,
|
||||
) -> Result<Arc<LocalStream>, String> {
|
||||
let proxy_conn = self
|
||||
.get_proxy_conn(node_id)
|
||||
@@ -454,20 +455,6 @@ impl HubRouter {
|
||||
&meta_payload,
|
||||
);
|
||||
|
||||
let (body_payload, body_flags) = match protocol::compress_payload(body.as_ref()) {
|
||||
Ok(result) => result,
|
||||
Err(e) => {
|
||||
proxy_conn.release_stream();
|
||||
return Err(format!("failed to compress request body: {e}"));
|
||||
}
|
||||
};
|
||||
let body_frame = protocol::encode_frame(
|
||||
proxy_stream_id,
|
||||
protocol::REQUEST_BODY,
|
||||
body_flags | protocol::FLAG_END_STREAM,
|
||||
&body_payload,
|
||||
);
|
||||
|
||||
// Frames encoded successfully -- now register the stream.
|
||||
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
|
||||
let local_stream = Arc::new(LocalStream::new(
|
||||
@@ -481,18 +468,75 @@ impl HubRouter {
|
||||
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
|
||||
|
||||
match proxy_conn.send(Message::Binary(header_frame.into())) {
|
||||
SendStatus::Queued => {}
|
||||
SendStatus::Queued => Ok(local_stream),
|
||||
SendStatus::Closed | SendStatus::Congested => {
|
||||
self.cleanup_local_stream(local_stream_id);
|
||||
proxy_conn.release_stream();
|
||||
return Err("proxy connection congested".to_string());
|
||||
Err("proxy connection congested".to_string())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn push_local_request_body(
|
||||
&self,
|
||||
local_stream_id: u64,
|
||||
payload: Bytes,
|
||||
end_stream: bool,
|
||||
) -> Result<(), String> {
|
||||
let stream = self
|
||||
.local_streams
|
||||
.get(&local_stream_id)
|
||||
.map(|entry| entry.value().clone())
|
||||
.ok_or_else(|| "local stream not found".to_string())?;
|
||||
let proxy_conn = self
|
||||
.proxy_conns_by_id
|
||||
.get(&stream.proxy_conn_id)
|
||||
.map(|entry| entry.value().clone())
|
||||
.ok_or_else(|| "proxy connection unavailable".to_string())?;
|
||||
|
||||
let total_chunks = payload.len().div_ceil(MAX_REQUEST_BODY_FRAME_SIZE);
|
||||
if total_chunks == 0 {
|
||||
if end_stream {
|
||||
self.send_request_body_frame(&proxy_conn, stream.proxy_stream_id, &[], true)?;
|
||||
}
|
||||
} else {
|
||||
for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() {
|
||||
let is_last_chunk = index + 1 == total_chunks;
|
||||
self.send_request_body_frame(
|
||||
&proxy_conn,
|
||||
stream.proxy_stream_id,
|
||||
chunk,
|
||||
end_stream && is_last_chunk,
|
||||
)?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn send_request_body_frame(
|
||||
&self,
|
||||
proxy_conn: &Arc<ProxyConn>,
|
||||
proxy_stream_id: u32,
|
||||
payload: &[u8],
|
||||
end_stream: bool,
|
||||
) -> Result<(), String> {
|
||||
let (body_payload, body_flags) = protocol::compress_payload(payload)
|
||||
.map_err(|e| format!("failed to compress request body: {e}"))?;
|
||||
let body_frame = protocol::encode_frame(
|
||||
proxy_stream_id,
|
||||
protocol::REQUEST_BODY,
|
||||
body_flags
|
||||
| if end_stream {
|
||||
protocol::FLAG_END_STREAM
|
||||
} else {
|
||||
0
|
||||
},
|
||||
&body_payload,
|
||||
);
|
||||
match proxy_conn.send(Message::Binary(body_frame.into())) {
|
||||
SendStatus::Queued => Ok(local_stream),
|
||||
SendStatus::Queued => Ok(()),
|
||||
SendStatus::Closed | SendStatus::Congested => {
|
||||
self.cancel_local_stream(local_stream_id, "proxy connection congested");
|
||||
Err("proxy connection congested".to_string())
|
||||
}
|
||||
}
|
||||
@@ -772,9 +816,11 @@ mod tests {
|
||||
hub.register_proxy(proxy);
|
||||
|
||||
let stream = hub
|
||||
.open_local_stream("node-1", &build_meta(), Bytes::new())
|
||||
.open_local_stream("node-1", &build_meta())
|
||||
.expect("open local stream");
|
||||
let _ = proxy_rx.try_recv().expect("headers frame");
|
||||
hub.push_local_request_body(stream.id, Bytes::new(), true)
|
||||
.expect("finish empty body");
|
||||
let _ = proxy_rx.try_recv().expect("body frame");
|
||||
|
||||
hub.cancel_local_stream(stream.id, "client dropped");
|
||||
@@ -787,4 +833,46 @@ mod tests {
|
||||
let header = protocol::FrameHeader::parse(&cancelled_data).expect("cancel frame header");
|
||||
assert_eq!(header.msg_type, protocol::STREAM_ERROR);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn push_local_request_body_splits_large_payload_and_marks_end() {
|
||||
let hub = HubRouter::new(ControlPlaneClient::disabled());
|
||||
|
||||
let (proxy_tx, mut proxy_rx) = mpsc::channel(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
let proxy = Arc::new(ProxyConn::new(
|
||||
200,
|
||||
"node-2".to_string(),
|
||||
"Node 2".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
));
|
||||
hub.register_proxy(proxy);
|
||||
|
||||
let stream = hub
|
||||
.open_local_stream("node-2", &build_meta())
|
||||
.expect("open local stream");
|
||||
let _ = proxy_rx.try_recv().expect("headers frame");
|
||||
|
||||
let payload = Bytes::from(vec![b'x'; MAX_REQUEST_BODY_FRAME_SIZE + 17]);
|
||||
hub.push_local_request_body(stream.id, payload, true)
|
||||
.expect("push request body");
|
||||
|
||||
let first = match proxy_rx.try_recv().expect("first body frame") {
|
||||
Message::Binary(data) => data.to_vec(),
|
||||
other => panic!("unexpected message: {other:?}"),
|
||||
};
|
||||
let first_header = protocol::FrameHeader::parse(&first).expect("first body header");
|
||||
assert_eq!(first_header.msg_type, protocol::REQUEST_BODY);
|
||||
assert_eq!(first_header.flags & protocol::FLAG_END_STREAM, 0);
|
||||
|
||||
let second = match proxy_rx.try_recv().expect("second body frame") {
|
||||
Message::Binary(data) => data.to_vec(),
|
||||
other => panic!("unexpected message: {other:?}"),
|
||||
};
|
||||
let second_header = protocol::FrameHeader::parse(&second).expect("second body header");
|
||||
assert_eq!(second_header.msg_type, protocol::REQUEST_BODY);
|
||||
assert_ne!(second_header.flags & protocol::FLAG_END_STREAM, 0);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,16 +4,19 @@ use std::time::Duration;
|
||||
|
||||
use async_stream::stream;
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::extract::{ConnectInfo, Path, State};
|
||||
use axum::extract::{ConnectInfo, Path, Request, State};
|
||||
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
|
||||
use axum::response::IntoResponse;
|
||||
use bytes::BytesMut;
|
||||
use futures_util::StreamExt;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::hub::LocalBodyEvent;
|
||||
use crate::hub::{LocalBodyEvent, LocalStream};
|
||||
use crate::protocol;
|
||||
use crate::AppState;
|
||||
|
||||
pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error";
|
||||
const MAX_RELAY_META_LEN: usize = 256 * 1024;
|
||||
|
||||
struct StreamGuard {
|
||||
hub: std::sync::Arc<crate::hub::HubRouter>,
|
||||
@@ -34,7 +37,7 @@ pub async fn relay_request(
|
||||
Path(node_id): Path<String>,
|
||||
State(state): State<AppState>,
|
||||
ConnectInfo(addr): ConnectInfo<SocketAddr>,
|
||||
body: Bytes,
|
||||
request: Request,
|
||||
) -> impl IntoResponse {
|
||||
if !addr.ip().is_loopback() {
|
||||
return tunnel_error_response(
|
||||
@@ -44,19 +47,104 @@ pub async fn relay_request(
|
||||
);
|
||||
}
|
||||
|
||||
let (meta, request_body) = match decode_envelope(body) {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
return tunnel_error_response(StatusCode::BAD_REQUEST, "bad_request", &error);
|
||||
let mut body_stream = request.into_body().into_data_stream();
|
||||
let mut envelope_buf = BytesMut::new();
|
||||
let mut meta: Option<protocol::RequestMeta> = None;
|
||||
let mut stream: Option<std::sync::Arc<LocalStream>> = None;
|
||||
|
||||
while let Some(chunk_result) = body_stream.next().await {
|
||||
let chunk = match chunk_result {
|
||||
Ok(chunk) => chunk,
|
||||
Err(error) => {
|
||||
if let Some(active_stream) = &stream {
|
||||
state
|
||||
.hub
|
||||
.cancel_local_stream(active_stream.id, "failed to read relay request body");
|
||||
}
|
||||
warn!(error = %error, "failed to read local relay request body");
|
||||
return tunnel_error_response(
|
||||
StatusCode::BAD_GATEWAY,
|
||||
"relay",
|
||||
"failed to read relay request body",
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
if stream.is_none() {
|
||||
envelope_buf.extend_from_slice(&chunk);
|
||||
let Some((parsed_meta, body_offset)) = (match try_decode_envelope_meta(&envelope_buf) {
|
||||
Ok(result) => result,
|
||||
Err(error) => {
|
||||
return tunnel_error_response(StatusCode::BAD_REQUEST, "bad_request", &error);
|
||||
}
|
||||
}) else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let opened_stream = match state.hub.open_local_stream(&node_id, &parsed_meta) {
|
||||
Ok(stream) => stream,
|
||||
Err(error) => {
|
||||
return tunnel_error_response(
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"connect",
|
||||
&error,
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
if envelope_buf.len() > body_offset {
|
||||
let first_body_chunk = Bytes::copy_from_slice(&envelope_buf[body_offset..]);
|
||||
if let Err(error) =
|
||||
state
|
||||
.hub
|
||||
.push_local_request_body(opened_stream.id, first_body_chunk, false)
|
||||
{
|
||||
state.hub.cancel_local_stream(opened_stream.id, &error);
|
||||
return tunnel_error_response(
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"connect",
|
||||
&error,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
envelope_buf.clear();
|
||||
meta = Some(parsed_meta);
|
||||
stream = Some(opened_stream);
|
||||
continue;
|
||||
}
|
||||
|
||||
let Some(active_stream) = &stream else {
|
||||
continue;
|
||||
};
|
||||
if let Err(error) = state
|
||||
.hub
|
||||
.push_local_request_body(active_stream.id, chunk, false)
|
||||
{
|
||||
state.hub.cancel_local_stream(active_stream.id, &error);
|
||||
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
|
||||
}
|
||||
}
|
||||
|
||||
let (meta, stream) = match (meta, stream) {
|
||||
(Some(meta), Some(stream)) => (meta, stream),
|
||||
_ => {
|
||||
return tunnel_error_response(
|
||||
StatusCode::BAD_REQUEST,
|
||||
"bad_request",
|
||||
"relay envelope metadata truncated",
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
let stream = match state.hub.open_local_stream(&node_id, &meta, request_body) {
|
||||
Ok(stream) => stream,
|
||||
Err(error) => {
|
||||
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
|
||||
}
|
||||
};
|
||||
if let Err(error) = state
|
||||
.hub
|
||||
.push_local_request_body(stream.id, Bytes::new(), true)
|
||||
{
|
||||
state.hub.cancel_local_stream(stream.id, &error);
|
||||
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
|
||||
}
|
||||
|
||||
let request_guard = StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
@@ -123,20 +211,25 @@ pub async fn relay_request(
|
||||
}
|
||||
}
|
||||
|
||||
fn decode_envelope(body: Bytes) -> Result<(protocol::RequestMeta, Bytes), String> {
|
||||
if body.len() < 4 {
|
||||
return Err("relay envelope too short".to_string());
|
||||
fn try_decode_envelope_meta(
|
||||
buffer: &BytesMut,
|
||||
) -> Result<Option<(protocol::RequestMeta, usize)>, String> {
|
||||
if buffer.len() < 4 {
|
||||
return Ok(None);
|
||||
}
|
||||
let meta_len = u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize;
|
||||
if meta_len > MAX_RELAY_META_LEN {
|
||||
return Err("relay metadata too large".to_string());
|
||||
}
|
||||
let meta_len = u32::from_be_bytes([body[0], body[1], body[2], body[3]]) as usize;
|
||||
let meta_end = 4usize
|
||||
.checked_add(meta_len)
|
||||
.ok_or_else(|| "relay envelope length overflow".to_string())?;
|
||||
if body.len() < meta_end {
|
||||
return Err("relay envelope metadata truncated".to_string());
|
||||
if buffer.len() < meta_end {
|
||||
return Ok(None);
|
||||
}
|
||||
let meta = serde_json::from_slice::<protocol::RequestMeta>(&body[4..meta_end])
|
||||
let meta = serde_json::from_slice::<protocol::RequestMeta>(&buffer[4..meta_end])
|
||||
.map_err(|e| format!("invalid relay metadata: {e}"))?;
|
||||
Ok((meta, body.slice(meta_end..)))
|
||||
Ok(Some((meta, meta_end)))
|
||||
}
|
||||
|
||||
fn append_headers(target: &mut HeaderMap, headers: &[(String, String)]) {
|
||||
|
||||
@@ -182,7 +182,9 @@ where
|
||||
|
||||
MsgType::StreamEnd | MsgType::StreamError => {
|
||||
// Client-side cancellation or end
|
||||
streams.remove(&frame.stream_id);
|
||||
if let Some(tx) = streams.remove(&frame.stream_id) {
|
||||
let _ = tx.send(frame).await;
|
||||
}
|
||||
}
|
||||
|
||||
MsgType::Ping => {
|
||||
|
||||
@@ -3,22 +3,27 @@
|
||||
//! Receives request frames, executes the upstream HTTP request,
|
||||
//! and sends response frames back through the writer channel.
|
||||
|
||||
use std::io;
|
||||
use std::sync::atomic::AtomicUsize;
|
||||
use std::sync::atomic::Ordering;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::stream;
|
||||
use futures_util::StreamExt;
|
||||
use http_body_util::BodyExt;
|
||||
use hyper::body::Frame as BodyFrame;
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::{debug, warn};
|
||||
|
||||
use crate::state::{AppState, ServerContext};
|
||||
use crate::target_filter;
|
||||
use crate::upstream_client::{self, UpstreamRequestBody};
|
||||
use crate::upstream_client;
|
||||
|
||||
use super::protocol::{
|
||||
compress_payload, decompress_if_gzip, flags, Frame, MsgType, RequestMeta, ResponseMeta,
|
||||
compress_payload, decompress_if_gzip, flags, Frame as TunnelFrame, MsgType, RequestMeta,
|
||||
ResponseMeta,
|
||||
};
|
||||
use super::writer::FrameSender;
|
||||
|
||||
@@ -64,13 +69,13 @@ pub async fn handle_stream(
|
||||
server: Arc<ServerContext>,
|
||||
stream_id: u32,
|
||||
meta: RequestMeta,
|
||||
mut body_rx: mpsc::Receiver<Frame>,
|
||||
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||
frame_tx: FrameSender,
|
||||
) {
|
||||
server.active_connections.fetch_add(1, Ordering::Release);
|
||||
|
||||
let connect_elapsed =
|
||||
handle_stream_inner(&state, &server, stream_id, meta, &mut body_rx, &frame_tx).await;
|
||||
handle_stream_inner(&state, &server, stream_id, meta, body_rx, &frame_tx).await;
|
||||
|
||||
server.active_connections.fetch_sub(1, Ordering::Release);
|
||||
if let Some(d) = connect_elapsed {
|
||||
@@ -79,7 +84,7 @@ pub async fn handle_stream(
|
||||
}
|
||||
|
||||
/// Send a frame to the writer with a timeout. Returns false if send failed.
|
||||
async fn send_frame(tx: &FrameSender, frame: Frame) -> bool {
|
||||
async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
||||
match tokio::time::timeout(FRAME_SEND_TIMEOUT, tx.send(frame)).await {
|
||||
Ok(Ok(())) => true,
|
||||
Ok(Err(_)) => {
|
||||
@@ -102,62 +107,9 @@ async fn handle_stream_inner(
|
||||
server: &ServerContext,
|
||||
stream_id: u32,
|
||||
meta: RequestMeta,
|
||||
body_rx: &mut mpsc::Receiver<Frame>,
|
||||
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||
frame_tx: &FrameSender,
|
||||
) -> Option<Duration> {
|
||||
// Collect request body
|
||||
let mut body_parts: Vec<Bytes> = Vec::new();
|
||||
let mut body_done = false;
|
||||
|
||||
// Drain body frames
|
||||
while !body_done {
|
||||
match body_rx.recv().await {
|
||||
Some(frame) => {
|
||||
if frame.msg_type == MsgType::RequestBody {
|
||||
let payload = match decompress_if_gzip(&frame) {
|
||||
Ok(d) => d,
|
||||
Err(e) => {
|
||||
send_error(
|
||||
frame_tx,
|
||||
stream_id,
|
||||
&format!("gzip decompress failed: {e}"),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
}
|
||||
};
|
||||
if !payload.is_empty() {
|
||||
body_parts.push(payload);
|
||||
}
|
||||
if frame.is_end_stream() {
|
||||
body_done = true;
|
||||
}
|
||||
} else if frame.msg_type == MsgType::StreamEnd
|
||||
|| frame.msg_type == MsgType::StreamError
|
||||
{
|
||||
body_done = true;
|
||||
if frame.msg_type == MsgType::StreamError {
|
||||
return None; // Client cancelled
|
||||
}
|
||||
}
|
||||
}
|
||||
None => return None, // Channel closed
|
||||
}
|
||||
}
|
||||
|
||||
let body: Bytes = if body_parts.is_empty() {
|
||||
Bytes::new()
|
||||
} else if body_parts.len() == 1 {
|
||||
body_parts.into_iter().next().unwrap()
|
||||
} else {
|
||||
let total: usize = body_parts.iter().map(|b| b.len()).sum();
|
||||
let mut combined = Vec::with_capacity(total);
|
||||
for part in &body_parts {
|
||||
combined.extend_from_slice(part);
|
||||
}
|
||||
Bytes::from(combined)
|
||||
};
|
||||
|
||||
// Validate target
|
||||
let target_url = match url::Url::parse(&meta.url) {
|
||||
Ok(u) => u,
|
||||
@@ -207,12 +159,14 @@ async fn handle_stream_inner(
|
||||
// Execute upstream request
|
||||
let client = &state.upstream_client;
|
||||
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 = build_streaming_request_body(body_rx, Arc::clone(&request_body_size));
|
||||
|
||||
let method: hyper::Method = meta.method.parse().unwrap_or(hyper::Method::GET);
|
||||
let mut request = match hyper::Request::builder()
|
||||
.method(method)
|
||||
.uri(meta.url.as_str())
|
||||
.body(UpstreamRequestBody::new(body.clone()))
|
||||
.body(request_body)
|
||||
{
|
||||
Ok(request) => request,
|
||||
Err(e) => {
|
||||
@@ -240,7 +194,6 @@ async fn handle_stream_inner(
|
||||
}
|
||||
}
|
||||
|
||||
let body_size = body.len();
|
||||
let mut captured_connection = upstream_client::capture_connection(&mut request);
|
||||
let connection_start = Instant::now();
|
||||
let connection_capture = tokio::spawn(async move {
|
||||
@@ -313,7 +266,7 @@ async fn handle_stream_inner(
|
||||
"upstream_processing_ms": request_timing.response_wait_ms,
|
||||
"timing_source": "instrumented_connector",
|
||||
"total_ms": connect_elapsed.as_millis() as u64,
|
||||
"body_size": body_size,
|
||||
"body_size": request_body_size.load(Ordering::Relaxed),
|
||||
"mode": "tunnel",
|
||||
});
|
||||
resp_headers.push(("x-proxy-timing".to_string(), timing.to_string()));
|
||||
@@ -325,7 +278,7 @@ async fn handle_stream_inner(
|
||||
let (meta_payload, meta_flags) = compress_payload(meta_json);
|
||||
if !send_frame(
|
||||
frame_tx,
|
||||
Frame::new(
|
||||
TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::ResponseHeaders,
|
||||
meta_flags,
|
||||
@@ -350,7 +303,7 @@ async fn handle_stream_inner(
|
||||
let (payload, extra_flags) = compress_payload(chunk);
|
||||
if !send_frame(
|
||||
frame_tx,
|
||||
Frame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
|
||||
TunnelFrame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -365,7 +318,12 @@ async fn handle_stream_inner(
|
||||
let (payload, extra_flags) = compress_payload(slice);
|
||||
if !send_frame(
|
||||
frame_tx,
|
||||
Frame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
|
||||
TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::ResponseBody,
|
||||
extra_flags,
|
||||
payload,
|
||||
),
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -387,7 +345,7 @@ async fn handle_stream_inner(
|
||||
// Send STREAM_END
|
||||
let _ = send_frame(
|
||||
frame_tx,
|
||||
Frame::new(
|
||||
TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::StreamEnd,
|
||||
flags::END_STREAM,
|
||||
@@ -404,7 +362,7 @@ async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
|
||||
// Error frames use best-effort delivery — don't block if writer is congested
|
||||
let _ = send_frame(
|
||||
tx,
|
||||
Frame::new(
|
||||
TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
@@ -413,3 +371,136 @@ async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
fn build_streaming_request_body(
|
||||
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||
body_size: Arc<AtomicUsize>,
|
||||
) -> upstream_client::UpstreamRequestBody {
|
||||
let body_stream = stream::unfold(
|
||||
(body_rx, body_size, false),
|
||||
|(mut body_rx, body_size, finished)| async move {
|
||||
if finished {
|
||||
return None;
|
||||
}
|
||||
|
||||
loop {
|
||||
let frame = match body_rx.recv().await {
|
||||
Some(frame) => frame,
|
||||
None => return None,
|
||||
};
|
||||
|
||||
match frame.msg_type {
|
||||
MsgType::RequestBody => {
|
||||
let end_stream = frame.is_end_stream();
|
||||
let payload = match decompress_if_gzip(&frame) {
|
||||
Ok(payload) => payload,
|
||||
Err(error) => {
|
||||
let err =
|
||||
io::Error::other(format!("gzip decompress failed: {error}"));
|
||||
return Some((Err(err), (body_rx, body_size, true)));
|
||||
}
|
||||
};
|
||||
|
||||
if payload.is_empty() {
|
||||
if end_stream {
|
||||
return None;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
body_size.fetch_add(payload.len(), Ordering::Relaxed);
|
||||
return Some((
|
||||
Ok(BodyFrame::data(payload)),
|
||||
(body_rx, body_size, end_stream),
|
||||
));
|
||||
}
|
||||
MsgType::StreamError => {
|
||||
let message = String::from_utf8(frame.payload.to_vec())
|
||||
.unwrap_or_else(|_| "client cancelled request body".to_string());
|
||||
return Some((Err(io::Error::other(message)), (body_rx, body_size, true)));
|
||||
}
|
||||
MsgType::StreamEnd => return None,
|
||||
_ => continue,
|
||||
}
|
||||
}
|
||||
},
|
||||
);
|
||||
|
||||
upstream_client::stream_request_body(body_stream)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn streaming_request_body_yields_chunks_and_tracks_size() {
|
||||
let (tx, rx) = mpsc::channel(4);
|
||||
let body_size = Arc::new(AtomicUsize::new(0));
|
||||
let mut body = build_streaming_request_body(rx, Arc::clone(&body_size));
|
||||
|
||||
tx.send(TunnelFrame::new(
|
||||
1,
|
||||
MsgType::RequestBody,
|
||||
0,
|
||||
Bytes::from_static(b"abc"),
|
||||
))
|
||||
.await
|
||||
.expect("send first chunk");
|
||||
tx.send(TunnelFrame::new(
|
||||
1,
|
||||
MsgType::RequestBody,
|
||||
flags::END_STREAM,
|
||||
Bytes::from_static(b"def"),
|
||||
))
|
||||
.await
|
||||
.expect("send final chunk");
|
||||
drop(tx);
|
||||
|
||||
let first = body
|
||||
.frame()
|
||||
.await
|
||||
.expect("first frame")
|
||||
.expect("first frame ok")
|
||||
.into_data()
|
||||
.expect("first data frame");
|
||||
let second = body
|
||||
.frame()
|
||||
.await
|
||||
.expect("second frame")
|
||||
.expect("second frame ok")
|
||||
.into_data()
|
||||
.expect("second data frame");
|
||||
|
||||
assert_eq!(first, Bytes::from_static(b"abc"));
|
||||
assert_eq!(second, Bytes::from_static(b"def"));
|
||||
assert!(body.frame().await.is_none());
|
||||
assert_eq!(body_size.load(Ordering::Relaxed), 6);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn streaming_request_body_surfaces_client_cancel_as_error() {
|
||||
let (tx, rx) = mpsc::channel(4);
|
||||
let body_size = Arc::new(AtomicUsize::new(0));
|
||||
let mut body = build_streaming_request_body(rx, Arc::clone(&body_size));
|
||||
|
||||
tx.send(TunnelFrame::new(
|
||||
1,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from_static(b"client cancelled"),
|
||||
))
|
||||
.await
|
||||
.expect("send cancel frame");
|
||||
drop(tx);
|
||||
|
||||
let err = body
|
||||
.frame()
|
||||
.await
|
||||
.expect("error frame present")
|
||||
.expect_err("body should surface cancellation error");
|
||||
assert!(err.to_string().contains("client cancelled"));
|
||||
assert!(body.frame().await.is_none());
|
||||
assert_eq!(body_size.load(Ordering::Relaxed), 0);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,7 +7,10 @@ use std::task::{Context, Poll};
|
||||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use http_body_util::Full;
|
||||
use futures_util::Stream;
|
||||
use http_body_util::combinators::UnsyncBoxBody;
|
||||
use http_body_util::{BodyExt, StreamBody};
|
||||
use hyper::body::Frame;
|
||||
use hyper::rt;
|
||||
use hyper::Response;
|
||||
use hyper::Uri;
|
||||
@@ -30,9 +33,16 @@ type BoxError = Box<dyn std::error::Error + Send + Sync>;
|
||||
type PlainStream = TokioIo<TcpStream>;
|
||||
type TlsStream = TokioIo<tokio_rustls::client::TlsStream<TcpStream>>;
|
||||
|
||||
pub type UpstreamRequestBody = Full<Bytes>;
|
||||
pub type UpstreamRequestBody = UnsyncBoxBody<Bytes, io::Error>;
|
||||
pub type UpstreamClient = Client<InstrumentedConnector, UpstreamRequestBody>;
|
||||
|
||||
pub fn stream_request_body<S>(stream: S) -> UpstreamRequestBody
|
||||
where
|
||||
S: Stream<Item = Result<Frame<Bytes>, io::Error>> + Send + 'static,
|
||||
{
|
||||
StreamBody::new(stream).boxed_unsync()
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
pub struct ConnectTiming {
|
||||
pub connect_ms: u64,
|
||||
|
||||
@@ -76,6 +76,8 @@ from src.core.api_format.conversion.stream_events import (
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
from src.core.logger import logger
|
||||
|
||||
_OPENAI_CLI_REQUEST_PREFIX_KEYS = ("model", "instructions", "tools", "input")
|
||||
|
||||
|
||||
def _is_chat_completions_response(data: dict[str, Any]) -> bool:
|
||||
"""检测数据是否为 OpenAI Chat Completions 格式(而非 Responses API 格式)。
|
||||
@@ -101,6 +103,18 @@ def _get_openai_chat_normalizer() -> "FormatNormalizer | None":
|
||||
return format_conversion_registry.get_normalizer("openai:chat")
|
||||
|
||||
|
||||
def reorder_openai_cli_request_prefix_keys(payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Keep the stable OpenAI Responses prefix ahead of the typically dynamic tail."""
|
||||
ordered: dict[str, Any] = {}
|
||||
for key in _OPENAI_CLI_REQUEST_PREFIX_KEYS:
|
||||
if key in payload:
|
||||
ordered[key] = payload[key]
|
||||
for key, value in payload.items():
|
||||
if key not in ordered:
|
||||
ordered[key] = value
|
||||
return ordered
|
||||
|
||||
|
||||
class OpenAICliNormalizer(FormatNormalizer):
|
||||
FORMAT_ID = "openai:cli"
|
||||
capabilities = FormatCapabilities(
|
||||
@@ -132,13 +146,13 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
request: dict[str, Any],
|
||||
variant: str,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Codex 同格式透传:直接在原始请求体上做最小补丁,跳过 internal 转换。"""
|
||||
"""Codex 同格式透传:做最小补丁并保持稳定的请求前缀顺序。"""
|
||||
if variant.lower() != "codex":
|
||||
return None
|
||||
out: dict[str, Any] = dict(request)
|
||||
# 内部路由标记:绝不能透传到上游。
|
||||
out.pop("_aether_compact", None)
|
||||
return out
|
||||
return reorder_openai_cli_request_prefix_keys(out)
|
||||
|
||||
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||
model = str(request.get("model") or "")
|
||||
@@ -2129,10 +2143,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
joined = "\n\n".join(parts)
|
||||
return joined or None
|
||||
|
||||
_REQUEST_PREFIX_KEYS = ("model", "instructions", "tools", "input")
|
||||
|
||||
def _reorder_request_prefix_keys(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
return self._reorder_request_keys(payload, self._REQUEST_PREFIX_KEYS)
|
||||
return reorder_openai_cli_request_prefix_keys(payload)
|
||||
|
||||
def _error_type_from_value(self, value: str) -> ErrorType:
|
||||
for t in ErrorType:
|
||||
|
||||
@@ -9,6 +9,9 @@ from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from src.core.api_format.conversion.normalizers.openai_cli import (
|
||||
reorder_openai_cli_request_prefix_keys,
|
||||
)
|
||||
from src.core.provider_types import ProviderType
|
||||
|
||||
|
||||
@@ -23,7 +26,8 @@ def patch_openai_cli_request_for_codex(
|
||||
out: dict[str, Any] = dict(request_body)
|
||||
# Internal routing marker; never send upstream.
|
||||
out.pop("_aether_compact", None)
|
||||
return out
|
||||
# Match the normalizer's stable prefix ordering even on same-format passthrough.
|
||||
return reorder_openai_cli_request_prefix_keys(out)
|
||||
|
||||
|
||||
def maybe_patch_request_for_codex(
|
||||
|
||||
@@ -66,22 +66,21 @@ class HubTunnelTransport(httpx.AsyncBaseTransport):
|
||||
if key.lower() not in _HOP_BY_HOP_HEADERS_BYTES:
|
||||
headers[key.decode("latin-1")] = value.decode("latin-1")
|
||||
|
||||
body = request.content or await request.aread() or b""
|
||||
envelope = _encode_relay_envelope(
|
||||
relay_content = _iter_relay_envelope(
|
||||
{
|
||||
"method": request.method,
|
||||
"url": str(request.url),
|
||||
"headers": headers,
|
||||
"timeout": int(self._timeout),
|
||||
},
|
||||
body,
|
||||
request,
|
||||
)
|
||||
|
||||
relay_request = self._relay_client.build_request(
|
||||
"POST",
|
||||
config.local_relay_url(self._node_id),
|
||||
headers={"content-type": _RELAY_CONTENT_TYPE},
|
||||
content=envelope,
|
||||
content=relay_content,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -124,9 +123,39 @@ class HubRelayResponseStream(httpx.AsyncByteStream):
|
||||
await self._response.aclose()
|
||||
|
||||
|
||||
def _encode_relay_envelope(meta: dict[str, object], body: bytes) -> bytes:
|
||||
async def _iter_relay_envelope(
|
||||
meta: dict[str, object],
|
||||
request: httpx.Request,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
meta_json = json.dumps(meta, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
|
||||
return struct.pack("!I", len(meta_json)) + meta_json + body
|
||||
yield struct.pack("!I", len(meta_json)) + meta_json
|
||||
async for chunk in _iter_request_body(request):
|
||||
if chunk:
|
||||
yield chunk
|
||||
|
||||
|
||||
async def _iter_request_body(request: httpx.Request) -> AsyncGenerator[bytes, None]:
|
||||
stream = request.stream
|
||||
try:
|
||||
if hasattr(stream, "__aiter__"):
|
||||
async for chunk in stream:
|
||||
if chunk:
|
||||
yield bytes(chunk)
|
||||
return
|
||||
|
||||
if hasattr(stream, "__iter__"):
|
||||
for chunk in stream:
|
||||
if chunk:
|
||||
yield bytes(chunk)
|
||||
return
|
||||
|
||||
body = request.content
|
||||
if body:
|
||||
yield body
|
||||
finally:
|
||||
aclose = getattr(stream, "aclose", None)
|
||||
if callable(aclose):
|
||||
await aclose()
|
||||
|
||||
|
||||
async def _read_error_message(response: httpx.Response) -> str:
|
||||
|
||||
@@ -58,6 +58,31 @@ def test_patch_openai_cli_request_for_codex_preserves_existing_prompt_cache_key(
|
||||
assert out["prompt_cache_key"] == "client-cache-key"
|
||||
|
||||
|
||||
def test_patch_openai_cli_request_for_codex_reorders_stable_prefix_keys() -> None:
|
||||
req = {
|
||||
"temperature": 0.7,
|
||||
"input": [],
|
||||
"metadata": {"request_id": "abc"},
|
||||
"model": "gpt-test",
|
||||
"tools": [{"type": "function", "name": "demo"}],
|
||||
"instructions": "keep",
|
||||
"store": True,
|
||||
"_aether_compact": True,
|
||||
}
|
||||
|
||||
out = patch_openai_cli_request_for_codex(req)
|
||||
|
||||
assert list(out.keys()) == [
|
||||
"model",
|
||||
"instructions",
|
||||
"tools",
|
||||
"input",
|
||||
"temperature",
|
||||
"metadata",
|
||||
"store",
|
||||
]
|
||||
|
||||
|
||||
def test_patch_openai_cli_request_for_codex_does_not_inject_prompt_cache_key() -> None:
|
||||
req = {"model": "gpt-test", "input": [], "_aether_compact": True}
|
||||
|
||||
@@ -151,6 +176,34 @@ def test_openai_cli_normalizer_codex_variant_keeps_instructions_missing_for_defa
|
||||
assert patched["instructions"] == "You are GPT-5."
|
||||
|
||||
|
||||
def test_openai_cli_normalizer_patch_for_codex_reorders_stable_prefix_keys() -> None:
|
||||
from src.core.api_format.conversion.normalizers.openai_cli import OpenAICliNormalizer
|
||||
|
||||
normalizer = OpenAICliNormalizer()
|
||||
out = normalizer.patch_for_variant(
|
||||
{
|
||||
"temperature": 0.7,
|
||||
"input": [],
|
||||
"metadata": {"request_id": "abc"},
|
||||
"model": "gpt-test",
|
||||
"tools": [{"type": "function", "name": "demo"}],
|
||||
"instructions": "keep",
|
||||
"_aether_compact": True,
|
||||
},
|
||||
"codex",
|
||||
)
|
||||
|
||||
assert out is not None
|
||||
assert list(out.keys()) == [
|
||||
"model",
|
||||
"instructions",
|
||||
"tools",
|
||||
"input",
|
||||
"temperature",
|
||||
"metadata",
|
||||
]
|
||||
|
||||
|
||||
def test_codex_envelope_extra_headers_does_not_inject_synthetic_headers() -> None:
|
||||
from src.services.provider.adapters.codex.envelope import codex_oauth_envelope
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ class _FakeRelayClient:
|
||||
def __init__(self, response: httpx.Response) -> None:
|
||||
self.response = response
|
||||
self.sent_request: httpx.Request | None = None
|
||||
self.sent_body: bytes | None = None
|
||||
|
||||
def build_request(self, method: str, url: str, **kwargs: Any) -> httpx.Request:
|
||||
return httpx.Request(method, url, **kwargs)
|
||||
@@ -22,6 +23,7 @@ class _FakeRelayClient:
|
||||
async def send(self, request: httpx.Request, *, stream: bool = False) -> httpx.Response:
|
||||
_ = stream
|
||||
self.sent_request = request
|
||||
self.sent_body = await request.aread()
|
||||
self.response.request = request
|
||||
return self.response
|
||||
|
||||
@@ -56,7 +58,7 @@ async def test_transport_encodes_local_relay_envelope(monkeypatch: pytest.Monkey
|
||||
await response.aclose()
|
||||
|
||||
assert fake_client.sent_request is not None
|
||||
payload = fake_client.sent_request.content
|
||||
payload = fake_client.sent_body
|
||||
assert payload is not None
|
||||
meta_len = struct.unpack("!I", payload[:4])[0]
|
||||
meta = json.loads(payload[4 : 4 + meta_len].decode("utf-8"))
|
||||
@@ -88,3 +90,33 @@ async def test_transport_maps_relay_timeout_to_read_timeout(
|
||||
|
||||
with pytest.raises(httpx.ReadTimeout, match="relay timed out"):
|
||||
await transport.handle_async_request(request)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transport_streams_request_body_from_async_generator(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
transport = HubTunnelTransport("node-1", timeout=12.0)
|
||||
fake_client = _FakeRelayClient(httpx.Response(200, content=b"ok"))
|
||||
monkeypatch.setattr("src.services.proxy_node.hub_transport.get_hub_config", _relay_config)
|
||||
monkeypatch.setattr(transport, "_relay_client", fake_client)
|
||||
|
||||
async def body() -> Any:
|
||||
yield b'{"hello":'
|
||||
yield b'"world"}'
|
||||
|
||||
request = httpx.Request(
|
||||
"POST",
|
||||
"https://example.com/v1/chat/completions",
|
||||
headers={"content-type": "application/json"},
|
||||
content=body(),
|
||||
)
|
||||
|
||||
response = await transport.handle_async_request(request)
|
||||
assert response.status_code == 200
|
||||
await response.aclose()
|
||||
|
||||
payload = fake_client.sent_body
|
||||
assert payload is not None
|
||||
meta_len = struct.unpack("!I", payload[:4])[0]
|
||||
assert payload[4 + meta_len :] == b'{"hello":"world"}'
|
||||
|
||||
Reference in New Issue
Block a user