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]]
|
[[package]]
|
||||||
name = "aether-hub"
|
name = "aether-hub"
|
||||||
version = "0.1.8"
|
version = "0.1.9"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"async-stream",
|
"async-stream",
|
||||||
"axum",
|
"axum",
|
||||||
@@ -19,6 +19,7 @@ dependencies = [
|
|||||||
"dashmap",
|
"dashmap",
|
||||||
"flate2",
|
"flate2",
|
||||||
"futures-util",
|
"futures-util",
|
||||||
|
"http-body-util",
|
||||||
"parking_lot",
|
"parking_lot",
|
||||||
"reqwest",
|
"reqwest",
|
||||||
"serde",
|
"serde",
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ flate2 = "1"
|
|||||||
futures-util = "0.3"
|
futures-util = "0.3"
|
||||||
bytes = "1"
|
bytes = "1"
|
||||||
async-stream = "0.3"
|
async-stream = "0.3"
|
||||||
|
http-body-util = "0.1"
|
||||||
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
|
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
|
||||||
|
|
||||||
[profile.release]
|
[profile.release]
|
||||||
|
|||||||
@@ -15,6 +15,8 @@ use tracing::{debug, info, warn};
|
|||||||
use crate::control_plane::ControlPlaneClient;
|
use crate::control_plane::ControlPlaneClient;
|
||||||
use crate::protocol;
|
use crate::protocol;
|
||||||
|
|
||||||
|
const MAX_REQUEST_BODY_FRAME_SIZE: usize = 32 * 1024;
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
pub enum SendStatus {
|
pub enum SendStatus {
|
||||||
Queued,
|
Queued,
|
||||||
@@ -421,7 +423,6 @@ impl HubRouter {
|
|||||||
&self,
|
&self,
|
||||||
node_id: &str,
|
node_id: &str,
|
||||||
meta: &protocol::RequestMeta,
|
meta: &protocol::RequestMeta,
|
||||||
body: Bytes,
|
|
||||||
) -> Result<Arc<LocalStream>, String> {
|
) -> Result<Arc<LocalStream>, String> {
|
||||||
let proxy_conn = self
|
let proxy_conn = self
|
||||||
.get_proxy_conn(node_id)
|
.get_proxy_conn(node_id)
|
||||||
@@ -454,20 +455,6 @@ impl HubRouter {
|
|||||||
&meta_payload,
|
&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.
|
// Frames encoded successfully -- now register the stream.
|
||||||
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
|
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
|
||||||
let local_stream = Arc::new(LocalStream::new(
|
let local_stream = Arc::new(LocalStream::new(
|
||||||
@@ -481,18 +468,75 @@ impl HubRouter {
|
|||||||
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
|
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
|
||||||
|
|
||||||
match proxy_conn.send(Message::Binary(header_frame.into())) {
|
match proxy_conn.send(Message::Binary(header_frame.into())) {
|
||||||
SendStatus::Queued => {}
|
SendStatus::Queued => Ok(local_stream),
|
||||||
SendStatus::Closed | SendStatus::Congested => {
|
SendStatus::Closed | SendStatus::Congested => {
|
||||||
self.cleanup_local_stream(local_stream_id);
|
self.cleanup_local_stream(local_stream_id);
|
||||||
proxy_conn.release_stream();
|
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())) {
|
match proxy_conn.send(Message::Binary(body_frame.into())) {
|
||||||
SendStatus::Queued => Ok(local_stream),
|
SendStatus::Queued => Ok(()),
|
||||||
SendStatus::Closed | SendStatus::Congested => {
|
SendStatus::Closed | SendStatus::Congested => {
|
||||||
self.cancel_local_stream(local_stream_id, "proxy connection congested");
|
|
||||||
Err("proxy connection congested".to_string())
|
Err("proxy connection congested".to_string())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -772,9 +816,11 @@ mod tests {
|
|||||||
hub.register_proxy(proxy);
|
hub.register_proxy(proxy);
|
||||||
|
|
||||||
let stream = hub
|
let stream = hub
|
||||||
.open_local_stream("node-1", &build_meta(), Bytes::new())
|
.open_local_stream("node-1", &build_meta())
|
||||||
.expect("open local stream");
|
.expect("open local stream");
|
||||||
let _ = proxy_rx.try_recv().expect("headers frame");
|
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");
|
let _ = proxy_rx.try_recv().expect("body frame");
|
||||||
|
|
||||||
hub.cancel_local_stream(stream.id, "client dropped");
|
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");
|
let header = protocol::FrameHeader::parse(&cancelled_data).expect("cancel frame header");
|
||||||
assert_eq!(header.msg_type, protocol::STREAM_ERROR);
|
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 async_stream::stream;
|
||||||
use axum::body::{Body, Bytes};
|
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::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
|
||||||
use axum::response::IntoResponse;
|
use axum::response::IntoResponse;
|
||||||
|
use bytes::BytesMut;
|
||||||
|
use futures_util::StreamExt;
|
||||||
use tracing::warn;
|
use tracing::warn;
|
||||||
|
|
||||||
use crate::hub::LocalBodyEvent;
|
use crate::hub::{LocalBodyEvent, LocalStream};
|
||||||
use crate::protocol;
|
use crate::protocol;
|
||||||
use crate::AppState;
|
use crate::AppState;
|
||||||
|
|
||||||
pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error";
|
pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error";
|
||||||
|
const MAX_RELAY_META_LEN: usize = 256 * 1024;
|
||||||
|
|
||||||
struct StreamGuard {
|
struct StreamGuard {
|
||||||
hub: std::sync::Arc<crate::hub::HubRouter>,
|
hub: std::sync::Arc<crate::hub::HubRouter>,
|
||||||
@@ -34,7 +37,7 @@ pub async fn relay_request(
|
|||||||
Path(node_id): Path<String>,
|
Path(node_id): Path<String>,
|
||||||
State(state): State<AppState>,
|
State(state): State<AppState>,
|
||||||
ConnectInfo(addr): ConnectInfo<SocketAddr>,
|
ConnectInfo(addr): ConnectInfo<SocketAddr>,
|
||||||
body: Bytes,
|
request: Request,
|
||||||
) -> impl IntoResponse {
|
) -> impl IntoResponse {
|
||||||
if !addr.ip().is_loopback() {
|
if !addr.ip().is_loopback() {
|
||||||
return tunnel_error_response(
|
return tunnel_error_response(
|
||||||
@@ -44,19 +47,104 @@ pub async fn relay_request(
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
let (meta, request_body) = match decode_envelope(body) {
|
let mut body_stream = request.into_body().into_data_stream();
|
||||||
Ok(value) => value,
|
let mut envelope_buf = BytesMut::new();
|
||||||
Err(error) => {
|
let mut meta: Option<protocol::RequestMeta> = None;
|
||||||
return tunnel_error_response(StatusCode::BAD_REQUEST, "bad_request", &error);
|
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) {
|
if let Err(error) = state
|
||||||
Ok(stream) => stream,
|
.hub
|
||||||
Err(error) => {
|
.push_local_request_body(stream.id, Bytes::new(), true)
|
||||||
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
|
{
|
||||||
}
|
state.hub.cancel_local_stream(stream.id, &error);
|
||||||
};
|
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
|
||||||
|
}
|
||||||
|
|
||||||
let request_guard = StreamGuard {
|
let request_guard = StreamGuard {
|
||||||
hub: state.hub.clone(),
|
hub: state.hub.clone(),
|
||||||
stream_id: stream.id,
|
stream_id: stream.id,
|
||||||
@@ -123,20 +211,25 @@ pub async fn relay_request(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn decode_envelope(body: Bytes) -> Result<(protocol::RequestMeta, Bytes), String> {
|
fn try_decode_envelope_meta(
|
||||||
if body.len() < 4 {
|
buffer: &BytesMut,
|
||||||
return Err("relay envelope too short".to_string());
|
) -> 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
|
let meta_end = 4usize
|
||||||
.checked_add(meta_len)
|
.checked_add(meta_len)
|
||||||
.ok_or_else(|| "relay envelope length overflow".to_string())?;
|
.ok_or_else(|| "relay envelope length overflow".to_string())?;
|
||||||
if body.len() < meta_end {
|
if buffer.len() < meta_end {
|
||||||
return Err("relay envelope metadata truncated".to_string());
|
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}"))?;
|
.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)]) {
|
fn append_headers(target: &mut HeaderMap, headers: &[(String, String)]) {
|
||||||
|
|||||||
@@ -182,7 +182,9 @@ where
|
|||||||
|
|
||||||
MsgType::StreamEnd | MsgType::StreamError => {
|
MsgType::StreamEnd | MsgType::StreamError => {
|
||||||
// Client-side cancellation or end
|
// 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 => {
|
MsgType::Ping => {
|
||||||
|
|||||||
@@ -3,22 +3,27 @@
|
|||||||
//! Receives request frames, executes the upstream HTTP request,
|
//! Receives request frames, executes the upstream HTTP request,
|
||||||
//! and sends response frames back through the writer channel.
|
//! 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::atomic::Ordering;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use std::time::{Duration, Instant};
|
use std::time::{Duration, Instant};
|
||||||
|
|
||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
|
use futures_util::stream;
|
||||||
use futures_util::StreamExt;
|
use futures_util::StreamExt;
|
||||||
use http_body_util::BodyExt;
|
use http_body_util::BodyExt;
|
||||||
|
use hyper::body::Frame as BodyFrame;
|
||||||
use tokio::sync::mpsc;
|
use tokio::sync::mpsc;
|
||||||
use tracing::{debug, warn};
|
use tracing::{debug, warn};
|
||||||
|
|
||||||
use crate::state::{AppState, ServerContext};
|
use crate::state::{AppState, ServerContext};
|
||||||
use crate::target_filter;
|
use crate::target_filter;
|
||||||
use crate::upstream_client::{self, UpstreamRequestBody};
|
use crate::upstream_client;
|
||||||
|
|
||||||
use super::protocol::{
|
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;
|
use super::writer::FrameSender;
|
||||||
|
|
||||||
@@ -64,13 +69,13 @@ pub async fn handle_stream(
|
|||||||
server: Arc<ServerContext>,
|
server: Arc<ServerContext>,
|
||||||
stream_id: u32,
|
stream_id: u32,
|
||||||
meta: RequestMeta,
|
meta: RequestMeta,
|
||||||
mut body_rx: mpsc::Receiver<Frame>,
|
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||||
frame_tx: FrameSender,
|
frame_tx: FrameSender,
|
||||||
) {
|
) {
|
||||||
server.active_connections.fetch_add(1, Ordering::Release);
|
server.active_connections.fetch_add(1, Ordering::Release);
|
||||||
|
|
||||||
let connect_elapsed =
|
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);
|
server.active_connections.fetch_sub(1, Ordering::Release);
|
||||||
if let Some(d) = connect_elapsed {
|
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.
|
/// 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 {
|
match tokio::time::timeout(FRAME_SEND_TIMEOUT, tx.send(frame)).await {
|
||||||
Ok(Ok(())) => true,
|
Ok(Ok(())) => true,
|
||||||
Ok(Err(_)) => {
|
Ok(Err(_)) => {
|
||||||
@@ -102,62 +107,9 @@ async fn handle_stream_inner(
|
|||||||
server: &ServerContext,
|
server: &ServerContext,
|
||||||
stream_id: u32,
|
stream_id: u32,
|
||||||
meta: RequestMeta,
|
meta: RequestMeta,
|
||||||
body_rx: &mut mpsc::Receiver<Frame>,
|
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||||
frame_tx: &FrameSender,
|
frame_tx: &FrameSender,
|
||||||
) -> Option<Duration> {
|
) -> 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
|
// Validate target
|
||||||
let target_url = match url::Url::parse(&meta.url) {
|
let target_url = match url::Url::parse(&meta.url) {
|
||||||
Ok(u) => u,
|
Ok(u) => u,
|
||||||
@@ -207,12 +159,14 @@ async fn handle_stream_inner(
|
|||||||
// Execute upstream request
|
// Execute upstream request
|
||||||
let client = &state.upstream_client;
|
let client = &state.upstream_client;
|
||||||
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 = build_streaming_request_body(body_rx, Arc::clone(&request_body_size));
|
||||||
|
|
||||||
let method: hyper::Method = meta.method.parse().unwrap_or(hyper::Method::GET);
|
let method: hyper::Method = meta.method.parse().unwrap_or(hyper::Method::GET);
|
||||||
let mut request = match hyper::Request::builder()
|
let mut request = match hyper::Request::builder()
|
||||||
.method(method)
|
.method(method)
|
||||||
.uri(meta.url.as_str())
|
.uri(meta.url.as_str())
|
||||||
.body(UpstreamRequestBody::new(body.clone()))
|
.body(request_body)
|
||||||
{
|
{
|
||||||
Ok(request) => request,
|
Ok(request) => request,
|
||||||
Err(e) => {
|
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 mut captured_connection = upstream_client::capture_connection(&mut request);
|
||||||
let connection_start = Instant::now();
|
let connection_start = Instant::now();
|
||||||
let connection_capture = tokio::spawn(async move {
|
let connection_capture = tokio::spawn(async move {
|
||||||
@@ -313,7 +266,7 @@ async fn handle_stream_inner(
|
|||||||
"upstream_processing_ms": request_timing.response_wait_ms,
|
"upstream_processing_ms": request_timing.response_wait_ms,
|
||||||
"timing_source": "instrumented_connector",
|
"timing_source": "instrumented_connector",
|
||||||
"total_ms": connect_elapsed.as_millis() as u64,
|
"total_ms": connect_elapsed.as_millis() as u64,
|
||||||
"body_size": body_size,
|
"body_size": request_body_size.load(Ordering::Relaxed),
|
||||||
"mode": "tunnel",
|
"mode": "tunnel",
|
||||||
});
|
});
|
||||||
resp_headers.push(("x-proxy-timing".to_string(), timing.to_string()));
|
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);
|
let (meta_payload, meta_flags) = compress_payload(meta_json);
|
||||||
if !send_frame(
|
if !send_frame(
|
||||||
frame_tx,
|
frame_tx,
|
||||||
Frame::new(
|
TunnelFrame::new(
|
||||||
stream_id,
|
stream_id,
|
||||||
MsgType::ResponseHeaders,
|
MsgType::ResponseHeaders,
|
||||||
meta_flags,
|
meta_flags,
|
||||||
@@ -350,7 +303,7 @@ async fn handle_stream_inner(
|
|||||||
let (payload, extra_flags) = compress_payload(chunk);
|
let (payload, extra_flags) = compress_payload(chunk);
|
||||||
if !send_frame(
|
if !send_frame(
|
||||||
frame_tx,
|
frame_tx,
|
||||||
Frame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
|
TunnelFrame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
@@ -365,7 +318,12 @@ async fn handle_stream_inner(
|
|||||||
let (payload, extra_flags) = compress_payload(slice);
|
let (payload, extra_flags) = compress_payload(slice);
|
||||||
if !send_frame(
|
if !send_frame(
|
||||||
frame_tx,
|
frame_tx,
|
||||||
Frame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
|
TunnelFrame::new(
|
||||||
|
stream_id,
|
||||||
|
MsgType::ResponseBody,
|
||||||
|
extra_flags,
|
||||||
|
payload,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
@@ -387,7 +345,7 @@ async fn handle_stream_inner(
|
|||||||
// Send STREAM_END
|
// Send STREAM_END
|
||||||
let _ = send_frame(
|
let _ = send_frame(
|
||||||
frame_tx,
|
frame_tx,
|
||||||
Frame::new(
|
TunnelFrame::new(
|
||||||
stream_id,
|
stream_id,
|
||||||
MsgType::StreamEnd,
|
MsgType::StreamEnd,
|
||||||
flags::END_STREAM,
|
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
|
// Error frames use best-effort delivery — don't block if writer is congested
|
||||||
let _ = send_frame(
|
let _ = send_frame(
|
||||||
tx,
|
tx,
|
||||||
Frame::new(
|
TunnelFrame::new(
|
||||||
stream_id,
|
stream_id,
|
||||||
MsgType::StreamError,
|
MsgType::StreamError,
|
||||||
0,
|
0,
|
||||||
@@ -413,3 +371,136 @@ async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
|
|||||||
)
|
)
|
||||||
.await;
|
.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 std::time::Duration;
|
||||||
|
|
||||||
use bytes::Bytes;
|
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::rt;
|
||||||
use hyper::Response;
|
use hyper::Response;
|
||||||
use hyper::Uri;
|
use hyper::Uri;
|
||||||
@@ -30,9 +33,16 @@ type BoxError = Box<dyn std::error::Error + Send + Sync>;
|
|||||||
type PlainStream = TokioIo<TcpStream>;
|
type PlainStream = TokioIo<TcpStream>;
|
||||||
type TlsStream = TokioIo<tokio_rustls::client::TlsStream<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 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)]
|
#[derive(Clone, Copy, Debug, Default)]
|
||||||
pub struct ConnectTiming {
|
pub struct ConnectTiming {
|
||||||
pub connect_ms: u64,
|
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.api_format.conversion.stream_state import StreamState
|
||||||
from src.core.logger import logger
|
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:
|
def _is_chat_completions_response(data: dict[str, Any]) -> bool:
|
||||||
"""检测数据是否为 OpenAI Chat Completions 格式(而非 Responses API 格式)。
|
"""检测数据是否为 OpenAI Chat Completions 格式(而非 Responses API 格式)。
|
||||||
@@ -101,6 +103,18 @@ def _get_openai_chat_normalizer() -> "FormatNormalizer | None":
|
|||||||
return format_conversion_registry.get_normalizer("openai:chat")
|
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):
|
class OpenAICliNormalizer(FormatNormalizer):
|
||||||
FORMAT_ID = "openai:cli"
|
FORMAT_ID = "openai:cli"
|
||||||
capabilities = FormatCapabilities(
|
capabilities = FormatCapabilities(
|
||||||
@@ -132,13 +146,13 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
request: dict[str, Any],
|
request: dict[str, Any],
|
||||||
variant: str,
|
variant: str,
|
||||||
) -> dict[str, Any] | None:
|
) -> dict[str, Any] | None:
|
||||||
"""Codex 同格式透传:直接在原始请求体上做最小补丁,跳过 internal 转换。"""
|
"""Codex 同格式透传:做最小补丁并保持稳定的请求前缀顺序。"""
|
||||||
if variant.lower() != "codex":
|
if variant.lower() != "codex":
|
||||||
return None
|
return None
|
||||||
out: dict[str, Any] = dict(request)
|
out: dict[str, Any] = dict(request)
|
||||||
# 内部路由标记:绝不能透传到上游。
|
# 内部路由标记:绝不能透传到上游。
|
||||||
out.pop("_aether_compact", None)
|
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:
|
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||||
model = str(request.get("model") or "")
|
model = str(request.get("model") or "")
|
||||||
@@ -2129,10 +2143,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
joined = "\n\n".join(parts)
|
joined = "\n\n".join(parts)
|
||||||
return joined or None
|
return joined or None
|
||||||
|
|
||||||
_REQUEST_PREFIX_KEYS = ("model", "instructions", "tools", "input")
|
|
||||||
|
|
||||||
def _reorder_request_prefix_keys(self, payload: dict[str, Any]) -> dict[str, Any]:
|
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:
|
def _error_type_from_value(self, value: str) -> ErrorType:
|
||||||
for t in ErrorType:
|
for t in ErrorType:
|
||||||
|
|||||||
@@ -9,6 +9,9 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import Any
|
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
|
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)
|
out: dict[str, Any] = dict(request_body)
|
||||||
# Internal routing marker; never send upstream.
|
# Internal routing marker; never send upstream.
|
||||||
out.pop("_aether_compact", None)
|
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(
|
def maybe_patch_request_for_codex(
|
||||||
|
|||||||
@@ -66,22 +66,21 @@ class HubTunnelTransport(httpx.AsyncBaseTransport):
|
|||||||
if key.lower() not in _HOP_BY_HOP_HEADERS_BYTES:
|
if key.lower() not in _HOP_BY_HOP_HEADERS_BYTES:
|
||||||
headers[key.decode("latin-1")] = value.decode("latin-1")
|
headers[key.decode("latin-1")] = value.decode("latin-1")
|
||||||
|
|
||||||
body = request.content or await request.aread() or b""
|
relay_content = _iter_relay_envelope(
|
||||||
envelope = _encode_relay_envelope(
|
|
||||||
{
|
{
|
||||||
"method": request.method,
|
"method": request.method,
|
||||||
"url": str(request.url),
|
"url": str(request.url),
|
||||||
"headers": headers,
|
"headers": headers,
|
||||||
"timeout": int(self._timeout),
|
"timeout": int(self._timeout),
|
||||||
},
|
},
|
||||||
body,
|
request,
|
||||||
)
|
)
|
||||||
|
|
||||||
relay_request = self._relay_client.build_request(
|
relay_request = self._relay_client.build_request(
|
||||||
"POST",
|
"POST",
|
||||||
config.local_relay_url(self._node_id),
|
config.local_relay_url(self._node_id),
|
||||||
headers={"content-type": _RELAY_CONTENT_TYPE},
|
headers={"content-type": _RELAY_CONTENT_TYPE},
|
||||||
content=envelope,
|
content=relay_content,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -124,9 +123,39 @@ class HubRelayResponseStream(httpx.AsyncByteStream):
|
|||||||
await self._response.aclose()
|
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")
|
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:
|
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"
|
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:
|
def test_patch_openai_cli_request_for_codex_does_not_inject_prompt_cache_key() -> None:
|
||||||
req = {"model": "gpt-test", "input": [], "_aether_compact": True}
|
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."
|
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:
|
def test_codex_envelope_extra_headers_does_not_inject_synthetic_headers() -> None:
|
||||||
from src.services.provider.adapters.codex.envelope import codex_oauth_envelope
|
from src.services.provider.adapters.codex.envelope import codex_oauth_envelope
|
||||||
|
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ class _FakeRelayClient:
|
|||||||
def __init__(self, response: httpx.Response) -> None:
|
def __init__(self, response: httpx.Response) -> None:
|
||||||
self.response = response
|
self.response = response
|
||||||
self.sent_request: httpx.Request | None = None
|
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:
|
def build_request(self, method: str, url: str, **kwargs: Any) -> httpx.Request:
|
||||||
return httpx.Request(method, url, **kwargs)
|
return httpx.Request(method, url, **kwargs)
|
||||||
@@ -22,6 +23,7 @@ class _FakeRelayClient:
|
|||||||
async def send(self, request: httpx.Request, *, stream: bool = False) -> httpx.Response:
|
async def send(self, request: httpx.Request, *, stream: bool = False) -> httpx.Response:
|
||||||
_ = stream
|
_ = stream
|
||||||
self.sent_request = request
|
self.sent_request = request
|
||||||
|
self.sent_body = await request.aread()
|
||||||
self.response.request = request
|
self.response.request = request
|
||||||
return self.response
|
return self.response
|
||||||
|
|
||||||
@@ -56,7 +58,7 @@ async def test_transport_encodes_local_relay_envelope(monkeypatch: pytest.Monkey
|
|||||||
await response.aclose()
|
await response.aclose()
|
||||||
|
|
||||||
assert fake_client.sent_request is not None
|
assert fake_client.sent_request is not None
|
||||||
payload = fake_client.sent_request.content
|
payload = fake_client.sent_body
|
||||||
assert payload is not None
|
assert payload is not None
|
||||||
meta_len = struct.unpack("!I", payload[:4])[0]
|
meta_len = struct.unpack("!I", payload[:4])[0]
|
||||||
meta = json.loads(payload[4 : 4 + meta_len].decode("utf-8"))
|
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"):
|
with pytest.raises(httpx.ReadTimeout, match="relay timed out"):
|
||||||
await transport.handle_async_request(request)
|
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