o402 0.1.5

OpenAI-compatible gateway, paid with x402.
//! Reverse proxy to configured OpenAI-compatible upstreams.

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;

/// Proxy a classified request. Paid routes go through the x402 gate.
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
}

/// Origin of `base_url` plus request path and query. Ignores any path on `base_url`.
#[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"
    )
}

/// Buffered upstream JSON (or other non-SSE) body.
#[derive(Debug)]
pub(crate) struct BufferedUpstream {
    /// HTTP status from upstream.
    pub(crate) status: StatusCode,
    /// Hop-by-hop-stripped response headers.
    pub(crate) headers: HeaderMap,
    /// Collected body bytes.
    pub(crate) body: Bytes,
}

/// Send the request to the configured upstream.
///
/// # Errors
///
/// Returns an HTTP [`Response`] when the body cannot be read, the API key is
/// not a valid header, or the upstream request fails to start.
#[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")
    })
}

/// Collect the upstream body. Used by Sequential settle after status is known.
///
/// # Errors
///
/// Returns an HTTP [`Response`] when the body cannot be read.
#[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"))
        }
    }
}

/// HTTP status from headers already received by [`send_upstream`].
#[must_use]
pub(crate) fn upstream_status(upstream: &reqwest::Response) -> StatusCode {
    status_of(upstream)
}

/// Hop-by-hop-stripped response headers from [`send_upstream`].
#[must_use]
pub(crate) fn response_headers(upstream: &reqwest::Response) -> HeaderMap {
    filter_response_headers(upstream.headers())
}

/// Forward upstream bytes as they arrive. Does not collect the body.
#[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;