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:
fawney19
2026-03-17 22:07:09 +08:00
parent 59840fa419
commit 0342f609d0
12 changed files with 541 additions and 125 deletions

3
aether-hub/Cargo.lock generated
View File

@@ -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",

View File

@@ -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]

View File

@@ -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);
}
}

View File

@@ -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)]) {

View File

@@ -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 => {

View File

@@ -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);
}
}

View File

@@ -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,

View File

@@ -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:

View File

@@ -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(

View File

@@ -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:

View File

@@ -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

View File

@@ -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"}'