mod sse;
use axum::body::{Body, Bytes};
use axum::extract::{Request, State};
use axum::http::header::{self, HeaderName, HeaderValue};
use axum::http::{HeaderMap, StatusCode};
use axum::response::Response;
use http_body_util::BodyExt as _;
use url::Url;
use crate::config::{ModelConfig, UpstreamConfig};
use crate::http::{RequestContext, openai_error};
use crate::state::AppState;
pub(crate) async fn handler(State(state): State<AppState>, request: Request) -> Response {
let bill =
crate::http::bill::classify(request.method(), request.uri().path(), request.headers());
match bill {
crate::http::bill::Bill::Reject => {
return openai_error::not_found(format!("Invalid URL ({})", request.uri().path()));
}
crate::http::bill::Bill::Exact | crate::http::bill::Bill::Upto
if state.payment_enabled() =>
{
return crate::payment::serve::gate(State(state), request).await;
}
crate::http::bill::Bill::Unpaid
| crate::http::bill::Bill::Exact
| crate::http::bill::Bill::Upto => {}
}
let model_id = request
.extensions()
.get::<RequestContext>()
.map(|ctx| ctx.model.as_str());
let Some((upstream, client, catalog)) = state.resolve_upstream(model_id) else {
return openai_error::bad_gateway("upstream is not configured");
};
tracing::info!(
method = %request.method(),
path = %request.uri().path(),
model = model_id.unwrap_or("*"),
"proxy request"
);
forward(client, upstream, catalog, request).await
}
#[must_use]
pub(crate) fn join_origin(base: &Url, path: &str, query: Option<&str>) -> Url {
let mut url = base.clone();
url.set_path(path);
url.set_query(query);
url
}
fn strip_request(name: &str) -> bool {
matches!(
name,
"host"
| "connection"
| "transfer-encoding"
| "keep-alive"
| "proxy-authenticate"
| "proxy-authorization"
| "proxy-connection"
| "te"
| "trailer"
| "upgrade"
| "content-length"
| "accept-encoding"
| "authorization"
| "payment-signature"
| "payment-required"
| "payment-response"
| "sign-in-with-x"
)
}
fn strip_response(name: &str) -> bool {
matches!(
name,
"connection"
| "transfer-encoding"
| "keep-alive"
| "proxy-authenticate"
| "proxy-authorization"
| "proxy-connection"
| "te"
| "trailer"
| "upgrade"
| "content-length"
| "content-encoding"
)
}
#[derive(Debug)]
pub(crate) struct BufferedUpstream {
pub(crate) status: StatusCode,
pub(crate) headers: HeaderMap,
pub(crate) body: Bytes,
}
#[allow(
clippy::result_large_err,
reason = "HTTP error responses are returned as Err"
)]
pub(crate) async fn send_upstream(
client: &reqwest::Client,
upstream: &UpstreamConfig,
model: Option<&ModelConfig>,
request: Request,
) -> Result<reqwest::Response, Response> {
let (parts, body) = request.into_parts();
let Some(body) = collect_body(body).await else {
return Err(openai_error::invalid_request("failed to read request body"));
};
let body = match model {
Some(model) => maybe_rewrite_model(body, model),
None => body,
};
let url = join_origin(&upstream.base_url, parts.uri.path(), parts.uri.query());
let Some(builder) =
build_upstream_request(client, &parts.method, upstream, &parts.headers, body, url)
else {
return Err(openai_error::bad_gateway("invalid upstream API key"));
};
builder.send().await.map_err(|error| {
tracing::warn!(error = %error, "upstream request failed");
openai_error::bad_gateway("upstream request failed")
})
}
#[allow(
clippy::result_large_err,
reason = "HTTP error responses are returned as Err"
)]
pub(crate) async fn buffer_upstream(
upstream: reqwest::Response,
) -> Result<BufferedUpstream, Response> {
let status = status_of(&upstream);
let headers = filter_response_headers(upstream.headers());
match upstream.bytes().await {
Ok(body) => Ok(BufferedUpstream {
status,
headers,
body,
}),
Err(error) => {
tracing::warn!(error = %error, "upstream body failed");
Err(openai_error::bad_gateway("upstream request failed"))
}
}
}
#[must_use]
pub(crate) fn upstream_status(upstream: &reqwest::Response) -> StatusCode {
status_of(upstream)
}
#[must_use]
pub(crate) fn response_headers(upstream: &reqwest::Response) -> HeaderMap {
filter_response_headers(upstream.headers())
}
#[must_use]
pub(crate) fn stream_upstream(upstream: reqwest::Response) -> Response {
let status = status_of(&upstream);
let headers = filter_response_headers(upstream.headers());
sse::stream_response(upstream, status, headers)
}
async fn forward(
client: &reqwest::Client,
upstream: &UpstreamConfig,
model: Option<&ModelConfig>,
request: Request,
) -> Response {
match send_upstream(client, upstream, model, request).await {
Ok(resp) => passthrough(resp).await,
Err(response) => response,
}
}
async fn collect_body(body: Body) -> Option<Bytes> {
body.collect()
.await
.ok()
.map(http_body_util::Collected::to_bytes)
}
fn maybe_rewrite_model(body: Bytes, model: &ModelConfig) -> Bytes {
let Some(upstream_model) = model.upstream_model.as_deref() else {
return body;
};
if upstream_model == model.id {
return body;
}
rewrite_model_field(&body, upstream_model).unwrap_or(body)
}
fn rewrite_model_field(body: &[u8], upstream_model: &str) -> Option<Bytes> {
let mut value: serde_json::Value = serde_json::from_slice(body).ok()?;
value.as_object_mut()?.insert(
"model".to_owned(),
serde_json::Value::String(upstream_model.to_owned()),
);
serde_json::to_vec(&value).ok().map(Bytes::from)
}
fn build_upstream_request(
client: &reqwest::Client,
method: &axum::http::Method,
upstream: &UpstreamConfig,
headers: &HeaderMap,
body: Bytes,
url: Url,
) -> Option<reqwest::RequestBuilder> {
let auth = HeaderValue::from_str(&format!("Bearer {}", upstream.api_key.expose())).ok()?;
let method =
reqwest::Method::from_bytes(method.as_str().as_bytes()).unwrap_or(reqwest::Method::POST);
let mut builder = client
.request(method, url)
.header(header::AUTHORIZATION, auth);
for (name, value) in headers {
if strip_request(name.as_str()) {
continue;
}
builder = builder.header(name.clone(), value.clone());
}
Some(builder.body(body))
}
async fn passthrough(upstream: reqwest::Response) -> Response {
let status = status_of(&upstream);
let headers = filter_response_headers(upstream.headers());
if sse::is_event_stream(&headers) {
return sse::stream_response(upstream, status, headers);
}
match upstream.bytes().await {
Ok(body) => buffered_response(status, headers, body),
Err(error) => {
tracing::warn!(error = %error, "upstream body failed");
openai_error::bad_gateway("upstream request failed")
}
}
}
pub(crate) fn buffered_response(status: StatusCode, headers: HeaderMap, body: Bytes) -> Response {
let mut response = Response::new(Body::from(body));
*response.status_mut() = status;
*response.headers_mut() = headers;
response
}
fn status_of(resp: &reqwest::Response) -> StatusCode {
StatusCode::from_u16(resp.status().as_u16()).unwrap_or(StatusCode::BAD_GATEWAY)
}
fn filter_response_headers(headers: &HeaderMap) -> HeaderMap {
let mut out = HeaderMap::new();
for (name, value) in headers {
if strip_response(name.as_str()) {
continue;
}
let Ok(name) = HeaderName::from_bytes(name.as_str().as_bytes()) else {
continue;
};
let Ok(value) = HeaderValue::from_bytes(value.as_bytes()) else {
continue;
};
out.append(name, value);
}
out
}
#[cfg(test)]
mod tests {
use url::Url;
use super::{join_origin, strip_request, strip_response};
fn url(raw: &str) -> Url {
raw.parse().expect("url")
}
#[test]
fn join_origin_replaces_base_path() {
let joined = join_origin(
&url("https://api.openai.com/v1/"),
"/v1/chat/completions",
None,
);
assert_eq!(
joined.as_str(),
"https://api.openai.com/v1/chat/completions",
"no double /v1"
);
}
#[test]
fn join_origin_keeps_query() {
let joined = join_origin(
&url("https://api.example.com"),
"/v1/models",
Some("limit=10"),
);
assert_eq!(
joined.as_str(),
"https://api.example.com/v1/models?limit=10",
"query"
);
}
#[test]
fn strips_client_authorization_and_payment_headers() {
assert!(strip_request("authorization"), "authorization");
assert!(strip_request("payment-signature"), "signature");
assert!(strip_request("payment-required"), "required");
assert!(strip_request("payment-response"), "response");
assert!(strip_request("sign-in-with-x"), "siwx");
assert!(strip_request("host"), "host");
assert!(!strip_request("content-type"), "content-type");
}
#[test]
fn strips_response_length_and_encoding() {
assert!(strip_response("content-length"), "content-length");
assert!(strip_response("content-encoding"), "content-encoding");
assert!(!strip_response("content-type"), "content-type");
}
}
#[cfg(test)]
mod proxy_tests;