use std::io::{Read as _, Write as _};
use std::time::Duration;
use axum::body::Body;
use axum::http::{Request, StatusCode};
use futures_util::StreamExt as _;
use serde_json::json;
use tower::ServiceExt;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use crate::config::Config;
use crate::http::app;
fn catalog_toml(base_url: &str) -> String {
format!(
r#"
[payment]
enabled = false
[[upstreams]]
name = "stub"
base_url = "{base_url}"
api_key = "sk-upstream"
timeout_secs = 5
connect_timeout_secs = 1
[[models]]
id = "gpt-4o-mini"
upstream = "stub"
owned_by = "openai"
"#
)
}
fn catalog_with_upstream_model(base_url: &str) -> String {
format!(
r#"
[payment]
enabled = false
[[upstreams]]
name = "stub"
base_url = "{base_url}"
api_key = "sk-upstream"
timeout_secs = 5
connect_timeout_secs = 1
[[models]]
id = "gpt-4o-mini"
upstream = "stub"
upstream_model = "vendor-model"
"#
)
}
fn router(toml: &str) -> axum::Router {
app(Config::from_toml_str(toml).expect("config")).expect("router")
}
async fn json_body(response: axum::http::Response<Body>) -> serde_json::Value {
let bytes = axum::body::to_bytes(response.into_body(), 1 << 20)
.await
.expect("body");
serde_json::from_slice(&bytes).expect("json")
}
fn chat_body() -> Body {
Body::from(r#"{"model":"gpt-4o-mini","messages":[{"role":"user","content":"hi"}]}"#)
}
#[tokio::test]
async fn authorization_is_not_forwarded() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/chat/completions"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"id":"ok"})))
.mount(&server)
.await;
let response = router(&catalog_toml(&server.uri()))
.oneshot(
Request::post("/v1/chat/completions")
.header("content-type", "application/json")
.header("authorization", "Bearer sk-client")
.header("payment-signature", "sig")
.header("payment-required", "req")
.header("payment-response", "resp")
.header("sign-in-with-x", "siwx")
.body(chat_body())
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::OK, "status");
let received = server.received_requests().await.expect("recorded");
let request = received.first().expect("upstream request");
let auths: Vec<_> = request
.headers
.get_all("authorization")
.iter()
.map(|value| value.to_str().unwrap_or(""))
.collect();
assert_eq!(auths, ["Bearer sk-upstream"], "injected key only");
assert!(
request.headers.get("payment-signature").is_none(),
"signature stripped"
);
assert!(
request.headers.get("payment-required").is_none(),
"required stripped"
);
assert!(
request.headers.get("payment-response").is_none(),
"response stripped"
);
assert!(
request.headers.get("sign-in-with-x").is_none(),
"siwx stripped"
);
}
#[tokio::test]
async fn no_double_v1_when_base_url_has_path() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/chat/completions"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"id":"ok"})))
.mount(&server)
.await;
let base = format!("{}/v1/", server.uri());
let response = router(&catalog_toml(&base))
.oneshot(
Request::post("/v1/chat/completions")
.header("content-type", "application/json")
.body(chat_body())
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::OK, "status");
let received = server.received_requests().await.expect("recorded");
let request = received.first().expect("upstream request");
assert_eq!(request.url.path(), "/v1/chat/completions", "single /v1");
}
#[tokio::test]
async fn unknown_model_is_proxied() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/chat/completions"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"id":"ok"})))
.mount(&server)
.await;
let response = router(&catalog_toml(&server.uri()))
.oneshot(
Request::post("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(r#"{"model":"foo"}"#))
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::OK, "status");
}
#[tokio::test]
async fn unknown_v1_path_is_proxied() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/images/generations"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"created": 1})))
.mount(&server)
.await;
let response = router(&catalog_toml(&server.uri()))
.oneshot(
Request::post("/v1/images/generations")
.header("content-type", "application/json")
.body(Body::from(r#"{"model":"gpt-image-1","prompt":"hi"}"#))
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::OK, "status");
}
#[tokio::test]
async fn models_get_is_proxied_unpaid() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/models"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"data": []})))
.mount(&server)
.await;
let response = router(&catalog_toml(&server.uri()))
.oneshot(
Request::get("/v1/models")
.body(Body::empty())
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::OK, "status");
}
#[tokio::test]
async fn realtime_is_not_proxied() {
let server = MockServer::start().await;
let response = router(&catalog_toml(&server.uri()))
.oneshot(
Request::get("/v1/realtime")
.body(Body::empty())
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::NOT_FOUND, "status");
assert!(
server
.received_requests()
.await
.expect("recorded")
.is_empty(),
"no upstream"
);
}
#[tokio::test]
async fn embeddings_completions_responses_passthrough() {
let server = MockServer::start().await;
for path_str in ["/v1/embeddings", "/v1/completions", "/v1/responses"] {
Mock::given(method("POST"))
.and(path(path_str))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"ok": path_str})))
.mount(&server)
.await;
}
let app = router(&catalog_toml(&server.uri()));
for path_str in ["/v1/embeddings", "/v1/completions", "/v1/responses"] {
let response = app
.clone()
.oneshot(
Request::post(path_str)
.header("content-type", "application/json")
.body(Body::from(r#"{"model":"gpt-4o-mini"}"#))
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::OK, "{path_str}");
let body = json_body(response).await;
assert_eq!(
body.get("ok").and_then(serde_json::Value::as_str),
Some(path_str),
"{path_str} body"
);
}
}
#[tokio::test]
async fn rewrites_upstream_model() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/chat/completions"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"id":"ok"})))
.mount(&server)
.await;
let response = router(&catalog_with_upstream_model(&server.uri()))
.oneshot(
Request::post("/v1/chat/completions")
.header("content-type", "application/json")
.body(chat_body())
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::OK, "status");
let received = server.received_requests().await.expect("recorded");
let request = received.first().expect("upstream request");
let body: serde_json::Value = serde_json::from_slice(&request.body).expect("json");
assert_eq!(
body.get("model").and_then(serde_json::Value::as_str),
Some("vendor-model"),
"rewritten"
);
}
#[tokio::test]
async fn upstream_error_status_is_passthrough() {
for (status, envelope) in [
(
StatusCode::BAD_REQUEST,
json!({"error":{"message":"bad request","type":"invalid_request_error"}}),
),
(
StatusCode::INTERNAL_SERVER_ERROR,
json!({"error":{"message":"overloaded","type":"server_error"}}),
),
] {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/chat/completions"))
.respond_with(ResponseTemplate::new(status.as_u16()).set_body_json(envelope.clone()))
.mount(&server)
.await;
let response = router(&catalog_toml(&server.uri()))
.oneshot(
Request::post("/v1/chat/completions")
.header("content-type", "application/json")
.body(chat_body())
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), status, "status {status}");
assert_eq!(json_body(response).await, envelope, "body {status}");
assert_eq!(
server.received_requests().await.expect("recorded").len(),
1,
"upstream called {status}"
);
}
}
#[tokio::test]
async fn invalid_json_is_400() {
let server = MockServer::start().await;
let response = router(&catalog_toml(&server.uri()))
.oneshot(
Request::post("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from("not-json"))
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::BAD_REQUEST, "status");
assert!(
server
.received_requests()
.await
.expect("recorded")
.is_empty(),
"no upstream"
);
}
#[tokio::test]
async fn sse_first_byte_arrives_before_stream_end() {
let (base, release, thread) = spawn_held_sse();
let response = router(&catalog_toml(&base))
.oneshot(
Request::post("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(
r#"{"model":"gpt-4o-mini","messages":[{"role":"user","content":"hi"}],"stream":true}"#,
))
.expect("request"),
)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::OK, "status");
assert!(
response
.headers()
.get("content-type")
.and_then(|value| value.to_str().ok())
.is_some_and(|ct| ct.contains("text/event-stream")),
"sse content-type"
);
let mut stream = response.into_body().into_data_stream();
let first = tokio::time::timeout(
Duration::from_secs(2),
next_containing(&mut stream, b"first"),
)
.await
.expect("first SSE byte waited on stream end")
.expect("first chunk");
assert!(
first
.windows(b"first".len())
.any(|window| window == b"first"),
"first event: {}",
String::from_utf8_lossy(&first)
);
release.send(()).expect("release");
let last = tokio::time::timeout(
Duration::from_secs(2),
next_containing(&mut stream, b"last"),
)
.await
.expect("last event")
.expect("last chunk");
assert!(
last.windows(b"last".len()).any(|window| window == b"last"),
"last event"
);
thread.join().expect("sse thread");
}
async fn next_containing(
stream: &mut axum::body::BodyDataStream,
needle: &[u8],
) -> Option<Vec<u8>> {
let mut buf = Vec::new();
loop {
let chunk = stream.next().await?.ok()?;
buf.extend_from_slice(&chunk);
if buf.windows(needle.len()).any(|window| window == needle) {
return Some(buf);
}
}
}
fn spawn_held_sse() -> (
String,
std::sync::mpsc::Sender<()>,
std::thread::JoinHandle<()>,
) {
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let (tx, rx) = std::sync::mpsc::channel();
let thread = std::thread::spawn(move || {
let (mut sock, _) = listener.accept().expect("accept");
sock.set_nodelay(true).expect("nodelay");
let mut buf = Vec::new();
let mut tmp = [0_u8; 1024];
loop {
let n = sock.read(&mut tmp).expect("read request");
if n == 0 {
break;
}
let Some(chunk) = tmp.get(..n) else {
break;
};
buf.extend_from_slice(chunk);
if buf.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
sock.write_all(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nCache-Control: no-cache\r\n\r\ndata: first\n\n",
)
.expect("headers");
sock.flush().expect("flush first");
rx.recv().expect("hold until first byte is observed");
sock.write_all(b"data: last\n\n").expect("last");
sock.flush().expect("flush last");
});
(format!("http://{addr}"), tx, thread)
}