From 016e43d3a410febdf47f25f53f83d6e135d43ecb Mon Sep 17 00:00:00 2001 From: Henrik Mertens Date: Tue, 28 Jul 2026 00:24:26 +0200 Subject: [PATCH 1/2] Implement OpenAI-compatible API gateway with key-based authentication - Add `gateway_projects` and `gateway_api_keys` tables for managing external API access - Implement OpenAI-compatible routes for `/v1/chat/completions`, `/v1/responses`, and `/v1/models` - Add bearer token authentication with scope validation (`inference:read`, `inference:write`) - Integrate `omniference` skins to bridge gateway requests to internal provider engines - Support team-based model access policies and usage tracking via gateway headers - Add automated tests for authentication, revoked/expired keys, and team policy enforcement - Bump `axum` compatibility by adding `axum07` bridge for `omniference` integration --- Cargo.lock | 1 + Cargo.toml | 3 +- migrations/20260727000000_gateway_auth.sql | 30 +++ src/routes/mod.rs | 7 + src/routes/public/mod.rs | 1 + src/routes/public/openai.rs | 155 ++++++++++++++++ src/routes/public/users.rs | 1 - src/tests/gateway.rs | 204 +++++++++++++++++++++ src/tests/mod.rs | 1 + src/types/gateway/mod.rs | 48 +++++ src/types/gateway/repository.rs | 115 ++++++++++++ src/types/gateway/responses.rs | 29 +++ src/types/mod.rs | 2 + src/utils/mod.rs | 1 + src/utils/openai_gateway.rs | 122 ++++++++++++ 15 files changed, 718 insertions(+), 2 deletions(-) create mode 100644 migrations/20260727000000_gateway_auth.sql create mode 100644 src/routes/public/openai.rs create mode 100644 src/tests/gateway.rs create mode 100644 src/types/gateway/mod.rs create mode 100644 src/types/gateway/repository.rs create mode 100644 src/types/gateway/responses.rs create mode 100644 src/utils/openai_gateway.rs diff --git a/Cargo.lock b/Cargo.lock index e79286a..2972f37 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -11,6 +11,7 @@ dependencies = [ "argon2", "async-stream", "async-trait", + "axum 0.7.9", "axum 0.8.9", "base64", "chrono", diff --git a/Cargo.toml b/Cargo.toml index 062a526..363a6f2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -10,7 +10,8 @@ keywords = ["chat", "axum", "websocket", "postgres"] categories = ["web-programming"] [dependencies] -axum = { version = "0.8" } +axum = { version = "0.8" } +axum07 = { package = "axum", version = "0.7" } tokio = { version = "1.48", features = ["full"] } serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" diff --git a/migrations/20260727000000_gateway_auth.sql b/migrations/20260727000000_gateway_auth.sql new file mode 100644 index 0000000..10473dd --- /dev/null +++ b/migrations/20260727000000_gateway_auth.sql @@ -0,0 +1,30 @@ +CREATE TABLE gateway_projects ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + owner_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE, + team_id UUID REFERENCES teams(id) ON DELETE CASCADE, + name VARCHAR(100) NOT NULL, + is_enabled BOOLEAN NOT NULL DEFAULT TRUE, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + +CREATE INDEX gateway_projects_owner_idx ON gateway_projects(owner_id); +CREATE INDEX gateway_projects_team_idx ON gateway_projects(team_id) WHERE team_id IS NOT NULL; + +CREATE TABLE gateway_api_keys ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + project_id UUID NOT NULL REFERENCES gateway_projects(id) ON DELETE CASCADE, + name VARCHAR(100) NOT NULL, + secret_hash TEXT NOT NULL, + key_prefix VARCHAR(48) NOT NULL, + last_four VARCHAR(4) NOT NULL, + scopes JSONB NOT NULL DEFAULT '["inference:read", "inference:write"]'::jsonb, + is_enabled BOOLEAN NOT NULL DEFAULT TRUE, + expires_at TIMESTAMPTZ, + revoked_at TIMESTAMPTZ, + last_used_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + +CREATE UNIQUE INDEX gateway_api_keys_prefix_idx ON gateway_api_keys(key_prefix); +CREATE INDEX gateway_api_keys_project_idx ON gateway_api_keys(project_id); diff --git a/src/routes/mod.rs b/src/routes/mod.rs index dc41621..56545a9 100644 --- a/src/routes/mod.rs +++ b/src/routes/mod.rs @@ -12,7 +12,14 @@ use std::sync::Arc; use tower_cookies::CookieManagerLayer; pub fn build_router() -> Router> { + let openai_router = Router::new() + .route("/models", get(public::openai::list_models)) + .route("/chat/completions", post(public::openai::chat_completions)) + .route("/responses", post(public::openai::responses)) + .method_not_allowed_fallback(public::openai::method_not_allowed) + .fallback(public::openai::not_found); Router::new() + .nest("/openai/v1", openai_router) .route("/api/v1/health", get(public::base::health)) .route("/api/v1/base", get(public::base::get_base)) // Admin i18n diff --git a/src/routes/public/mod.rs b/src/routes/public/mod.rs index e710d96..88223f5 100644 --- a/src/routes/public/mod.rs +++ b/src/routes/public/mod.rs @@ -6,6 +6,7 @@ pub mod mcp; pub mod messages; pub mod models; pub mod oauth; +pub mod openai; pub mod preferences; pub mod providers; pub mod streaming; diff --git a/src/routes/public/openai.rs b/src/routes/public/openai.rs new file mode 100644 index 0000000..687f371 --- /dev/null +++ b/src/routes/public/openai.rs @@ -0,0 +1,155 @@ +use crate::types::models::Model; +use crate::types::{JobState, OpenAiModel, OpenAiModelsResponse}; +use crate::utils::openai_gateway::{add_gateway_headers, authenticate, error_response, parse_json, run_chat, run_responses}; +use axum::body::Bytes; +use axum::extract::State; +use axum::http::{HeaderMap, StatusCode}; +use axum::response::{IntoResponse, Response}; +use omniference::types::providers::openai::{OpenAIChatRequest, OpenAIResponsesRequestPayload}; +use std::sync::Arc; +use uuid::Uuid; + +pub async fn not_found() -> Response { + add_gateway_headers( + error_response(StatusCode::NOT_FOUND, "The requested resource was not found", "not_found_error", "not_found"), + None, + Uuid::new_v4(), + ) +} + +pub async fn method_not_allowed() -> Response { + add_gateway_headers( + error_response( + StatusCode::METHOD_NOT_ALLOWED, + "Invalid HTTP method for this endpoint", + "invalid_request_error", + "method_not_allowed", + ), + None, + Uuid::new_v4(), + ) +} + +pub async fn list_models(State(state): State>, headers: HeaderMap) -> impl IntoResponse { + let request_id = Uuid::new_v4(); + let context = match authenticate(&state.db, &headers, "inference:read").await { + Ok(context) => context, + Err(response) => return add_gateway_headers(response, None, request_id), + }; + let models = match OpenAiModel::list_for_user(&state.db, &context.user_id).await { + Ok(models) => models, + Err(error) => { + tracing::error!(%error, "failed to list gateway models"); + return add_gateway_headers( + error_response(StatusCode::INTERNAL_SERVER_ERROR, "Failed to list models", "server_error", "internal_error"), + Some(&context), + request_id, + ); + } + }; + add_gateway_headers( + axum::Json(OpenAiModelsResponse { + object: "list".to_string(), + data: models, + }) + .into_response(), + Some(&context), + request_id, + ) +} + +pub async fn chat_completions(State(state): State>, headers: HeaderMap, body: Bytes) -> Response { + let request_id = Uuid::new_v4(); + let context = match authenticate(&state.db, &headers, "inference:write").await { + Ok(context) => context, + Err(response) => return add_gateway_headers(response, None, request_id), + }; + let request: OpenAIChatRequest = match parse_json(&body) { + Ok(request) => request, + Err(message) => { + return add_gateway_headers( + error_response(StatusCode::BAD_REQUEST, message, "invalid_request_error", "invalid_request_body"), + Some(&context), + request_id, + ); + } + }; + if let Some(response) = validate_model_access(&state, &context.user_id, &request.model).await { + return add_gateway_headers(response, Some(&context), request_id); + } + let response = run_chat(request).await; + add_gateway_headers(response, Some(&context), request_id) +} + +pub async fn responses(State(state): State>, headers: HeaderMap, body: Bytes) -> Response { + let request_id = Uuid::new_v4(); + let context = match authenticate(&state.db, &headers, "inference:write").await { + Ok(context) => context, + Err(response) => return add_gateway_headers(response, None, request_id), + }; + let request: OpenAIResponsesRequestPayload = match parse_json(&body) { + Ok(request) => request, + Err(message) => { + return add_gateway_headers( + error_response(StatusCode::BAD_REQUEST, message, "invalid_request_error", "invalid_request_body"), + Some(&context), + request_id, + ); + } + }; + let Some(model) = request.model.as_deref() else { + return add_gateway_headers( + error_response( + StatusCode::BAD_REQUEST, + "Missing required parameter: 'model'.", + "invalid_request_error", + "missing_required_parameter", + ), + Some(&context), + request_id, + ); + }; + if let Some(response) = validate_model_access(&state, &context.user_id, model).await { + return add_gateway_headers(response, Some(&context), request_id); + } + let response = run_responses(request).await; + add_gateway_headers(response, Some(&context), request_id) +} + +async fn validate_model_access(state: &JobState, user_id: &Uuid, model_id: &str) -> Option { + let model = match Model::find_by_model_id(&state.db, model_id).await { + Ok(Some(model)) if model.is_enabled => model, + Ok(_) => return Some(model_not_found(model_id)), + Err(error) => { + tracing::error!(%error, model_id, "failed to resolve gateway model"); + return Some(error_response( + StatusCode::INTERNAL_SERVER_ERROR, + "Failed to resolve model", + "server_error", + "internal_error", + )); + } + }; + match Model::can_user_use_model(&state.db, user_id, &model.id).await { + Ok(true) => None, + Ok(false) => Some(model_not_found(model_id)), + Err(error) => { + tracing::error!(%error, model_id, "failed to evaluate gateway model policy"); + Some(error_response( + StatusCode::INTERNAL_SERVER_ERROR, + "Failed to resolve model", + "server_error", + "internal_error", + )) + } + } +} + +fn model_not_found(model_id: &str) -> Response { + error_response( + StatusCode::NOT_FOUND, + format!("Model '{model_id}' not found"), + "invalid_request_error", + "model_not_found", + ) +} diff --git a/src/routes/public/users.rs b/src/routes/public/users.rs index 7ae152d..3011407 100644 --- a/src/routes/public/users.rs +++ b/src/routes/public/users.rs @@ -66,4 +66,3 @@ pub async fn get_my_analytics(State(state): State>, cookies: Cooki } } } - diff --git a/src/tests/gateway.rs b/src/tests/gateway.rs new file mode 100644 index 0000000..61c5446 --- /dev/null +++ b/src/tests/gateway.rs @@ -0,0 +1,204 @@ +#[cfg(test)] +mod tests { + use crate::types::{GatewayAuthError, GatewayCredential, OpenAiModel}; + use crate::utils::auth::hash_password; + use crate::utils::openai_gateway::bridge_response; + use axum::body::Bytes; + use omniference::skins::{OpenAIChatSkin, OpenAIResponsesSkin, Skin}; + use omniference::types::providers::openai::{OpenAIChatRequest, OpenAIResponsesRequestPayload}; + use omniference::types::{ModelRef, ProviderConfig, ProviderEndpoint, ProviderKind}; + use serde_json::json; + use sqlx::PgPool; + use std::collections::BTreeMap; + use std::convert::Infallible; + use std::sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }; + use uuid::Uuid; + + struct DropFlag(Arc); + + impl Drop for DropFlag { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + + async fn create_user(pool: &PgPool) -> Uuid { + sqlx::query_scalar("INSERT INTO users (email, username, password_hash) VALUES ('gateway@example.com', 'gateway', 'hash') RETURNING id") + .fetch_one(pool) + .await + .unwrap() + } + + async fn create_key(pool: &PgPool, user_id: Uuid, scopes: serde_json::Value) -> (Uuid, String) { + let project_id: Uuid = sqlx::query_scalar("INSERT INTO gateway_projects (owner_id, name) VALUES ($1, 'Gateway') RETURNING id") + .bind(user_id) + .fetch_one(pool) + .await + .unwrap(); + let key_id = Uuid::new_v4(); + let secret = "abcdefghijklmnopqrstuvwxyz0123456789"; + let token = format!("oxc_{}_{}", key_id.simple(), secret); + sqlx::query( + "INSERT INTO gateway_api_keys (id, project_id, name, secret_hash, key_prefix, last_four, scopes) + VALUES ($1, $2, 'Test', $3, $4, $5, $6)", + ) + .bind(key_id) + .bind(project_id) + .bind(hash_password(secret).unwrap()) + .bind(format!("oxc_{}", key_id.simple())) + .bind(&secret[secret.len() - 4..]) + .bind(scopes) + .execute(pool) + .await + .unwrap(); + (key_id, token) + } + + fn model_ref() -> ModelRef { + ModelRef { + alias: "provider/test-model".to_string(), + provider: ProviderConfig { + name: "provider".to_string(), + endpoint: ProviderEndpoint { + kind: ProviderKind::OpenAICompat, + base_url: "https://example.com".to_string(), + api_key: None, + extra_headers: BTreeMap::new(), + timeout: None, + }, + enabled: true, + catalog_provider_slug: None, + }, + model_id: "test-model".to_string(), + input_modalities: Vec::new(), + output_modalities: Vec::new(), + } + } + + #[sqlx::test(migrations = "./migrations")] + async fn bearer_key_authenticates_and_updates_last_used(pool: PgPool) { + let user_id = create_user(&pool).await; + let (key_id, token) = create_key(&pool, user_id, json!(["inference:read", "inference:write"])).await; + let context = GatewayCredential::authenticate(&pool, &token).await.unwrap(); + assert_eq!(context.key_id, key_id); + assert_eq!(context.user_id, user_id); + assert!(context.allows("inference:write")); + let last_used: Option> = sqlx::query_scalar("SELECT last_used_at FROM gateway_api_keys WHERE id = $1") + .bind(key_id) + .fetch_one(&pool) + .await + .unwrap(); + assert!(last_used.is_some()); + } + + #[sqlx::test(migrations = "./migrations")] + async fn invalid_revoked_and_expired_keys_are_rejected(pool: PgPool) { + let user_id = create_user(&pool).await; + let (key_id, token) = create_key(&pool, user_id, json!(["inference:write"])).await; + assert!(matches!( + GatewayCredential::authenticate(&pool, "oxc_00000000000000000000000000000000_abcdefghijklmnopqrstuvwxyz012345").await, + Err(GatewayAuthError::Invalid) + )); + sqlx::query("UPDATE gateway_api_keys SET revoked_at = NOW() WHERE id = $1") + .bind(key_id) + .execute(&pool) + .await + .unwrap(); + assert!(matches!(GatewayCredential::authenticate(&pool, &token).await, Err(GatewayAuthError::Invalid))); + sqlx::query("UPDATE gateway_api_keys SET revoked_at = NULL, expires_at = NOW() - INTERVAL '1 second' WHERE id = $1") + .bind(key_id) + .execute(&pool) + .await + .unwrap(); + assert!(matches!(GatewayCredential::authenticate(&pool, &token).await, Err(GatewayAuthError::Invalid))); + } + + #[sqlx::test(migrations = "./migrations")] + async fn model_listing_respects_team_policy(pool: PgPool) { + let user_id = create_user(&pool).await; + let team_id: Uuid = sqlx::query_scalar("INSERT INTO teams (name, allow_all_models) VALUES ('Gateway Team', false) RETURNING id") + .fetch_one(&pool) + .await + .unwrap(); + sqlx::query("INSERT INTO team_members (team_id, user_id) VALUES ($1, $2)") + .bind(team_id) + .bind(user_id) + .execute(&pool) + .await + .unwrap(); + let provider_id: Uuid = sqlx::query_scalar( + "INSERT INTO providers (kind, name, base_url, is_enabled) VALUES ('OPENAI', 'Gateway Provider', 'https://example.com', true) RETURNING id", + ) + .fetch_one(&pool) + .await + .unwrap(); + let visible_id: Uuid = + sqlx::query_scalar("INSERT INTO models (provider_id, model_id, display_name, is_enabled) VALUES ($1, 'visible', 'Visible', true) RETURNING id") + .bind(provider_id) + .fetch_one(&pool) + .await + .unwrap(); + sqlx::query("INSERT INTO models (provider_id, model_id, display_name, is_enabled) VALUES ($1, 'hidden', 'Hidden', true)") + .bind(provider_id) + .execute(&pool) + .await + .unwrap(); + sqlx::query("INSERT INTO team_model_access (team_id, model_id) VALUES ($1, $2)") + .bind(team_id) + .bind(visible_id) + .execute(&pool) + .await + .unwrap(); + let models = OpenAiModel::list_for_user(&pool, &user_id).await.unwrap(); + assert_eq!(models.len(), 1); + assert_eq!(models[0].id, "visible"); + } + + #[test] + fn chat_skin_converts_streaming_and_non_streaming_requests() { + for stream in [false, true] { + let request: OpenAIChatRequest = serde_json::from_value(json!({ + "model": "test-model", + "messages": [{"role": "user", "content": "hello"}], + "stream": stream + })) + .unwrap(); + let ir = OpenAIChatSkin::external_to_ir(request, model_ref()).unwrap(); + assert_eq!(ir.stream, stream); + assert_eq!(ir.messages.len(), 1); + assert!(ir.openai_chat_request.is_some()); + } + } + + #[test] + fn responses_skin_converts_streaming_and_non_streaming_requests() { + for stream in [false, true] { + let request: OpenAIResponsesRequestPayload = serde_json::from_value(json!({ + "model": "test-model", + "input": "hello", + "stream": stream + })) + .unwrap(); + let ir = OpenAIResponsesSkin::external_to_ir(request, model_ref()).unwrap(); + assert_eq!(ir.stream, stream); + assert_eq!(ir.messages.len(), 1); + } + } + + #[tokio::test] + async fn dropping_bridged_stream_releases_the_upstream_body() { + let dropped = Arc::new(AtomicBool::new(false)); + let guard = DropFlag(dropped.clone()); + let stream = futures_util::stream::once(async move { + let _guard = guard; + std::future::pending::>().await + }); + let response = axum07::response::Response::new(axum07::body::Body::from_stream(stream)); + let response = bridge_response(response); + drop(response); + assert!(dropped.load(Ordering::SeqCst)); + } +} diff --git a/src/tests/mod.rs b/src/tests/mod.rs index 003ee6b..09a7967 100644 --- a/src/tests/mod.rs +++ b/src/tests/mod.rs @@ -3,6 +3,7 @@ pub mod base; pub mod catalog; pub mod config; pub mod encryption; +pub mod gateway; pub mod i18n; pub mod images; pub mod logging; diff --git a/src/types/gateway/mod.rs b/src/types/gateway/mod.rs new file mode 100644 index 0000000..6097ea1 --- /dev/null +++ b/src/types/gateway/mod.rs @@ -0,0 +1,48 @@ +mod repository; +mod responses; + +pub use responses::*; + +use chrono::{DateTime, Utc}; +use sqlx::types::Json; +use uuid::Uuid; + +#[derive(Debug, sqlx::FromRow)] +pub struct GatewayCredential { + pub key_id: Uuid, + pub project_id: Uuid, + pub user_id: Uuid, + pub team_id: Option, + pub project_name: String, + pub secret_hash: String, + pub scopes: Json>, + pub key_enabled: bool, + pub project_enabled: bool, + pub expires_at: Option>, + pub revoked_at: Option>, +} + +#[derive(Clone, Debug)] +pub struct GatewayAuthContext { + pub key_id: Uuid, + pub project_id: Uuid, + pub user_id: Uuid, + pub team_id: Option, + pub project_name: String, + pub scopes: Vec, +} + +impl GatewayAuthContext { + #[must_use] + pub fn allows(&self, scope: &str) -> bool { + self.scopes.iter().any(|candidate| candidate == "*" || candidate == scope) + } +} + +#[derive(Debug, thiserror::Error)] +pub enum GatewayAuthError { + #[error("invalid_api_key")] + Invalid, + #[error("gateway_auth_unavailable")] + Unavailable, +} diff --git a/src/types/gateway/repository.rs b/src/types/gateway/repository.rs new file mode 100644 index 0000000..ced1b64 --- /dev/null +++ b/src/types/gateway/repository.rs @@ -0,0 +1,115 @@ +use super::{GatewayAuthContext, GatewayAuthError, GatewayCredential, OpenAiModel}; +use crate::utils::auth::{hash_password, verify_password}; +use chrono::Utc; +use sqlx::PgPool; +use std::sync::LazyLock; +use uuid::Uuid; + +static DUMMY_SECRET_HASH: LazyLock = LazyLock::new(|| hash_password("oxide-gateway-invalid-secret").unwrap_or_default()); + +impl GatewayCredential { + pub async fn authenticate(pool: &PgPool, token: &str) -> Result { + let key_id = parse_key_id(token); + let credential = match key_id { + Some(id) => Self::find(pool, &id).await.map_err(|_| GatewayAuthError::Unavailable)?, + None => None, + }; + let hash = credential.as_ref().map_or(DUMMY_SECRET_HASH.as_str(), |value| value.secret_hash.as_str()); + let secret = token.rsplit_once('_').map_or("", |(_, value)| value).to_owned(); + let hash = hash.to_owned(); + let verified = tokio::task::spawn_blocking(move || verify_password(&secret, &hash).unwrap_or(false)) + .await + .map_err(|_| GatewayAuthError::Unavailable)?; + let Some(credential) = credential else { + return Err(GatewayAuthError::Invalid); + }; + if !verified + || !credential.key_enabled + || !credential.project_enabled + || credential.revoked_at.is_some() + || credential.expires_at.is_some_and(|expires_at| expires_at <= Utc::now()) + { + return Err(GatewayAuthError::Invalid); + } + sqlx::query("UPDATE gateway_api_keys SET last_used_at = NOW() WHERE id = $1") + .bind(credential.key_id) + .execute(pool) + .await + .map_err(|_| GatewayAuthError::Unavailable)?; + Ok(GatewayAuthContext { + key_id: credential.key_id, + project_id: credential.project_id, + user_id: credential.user_id, + team_id: credential.team_id, + project_name: credential.project_name, + scopes: credential.scopes.0, + }) + } + + async fn find(pool: &PgPool, key_id: &Uuid) -> Result, sqlx::Error> { + sqlx::query_as::<_, Self>( + r#" + SELECT + k.id AS key_id, + k.project_id, + p.owner_id AS user_id, + p.team_id, + p.name AS project_name, + k.secret_hash, + k.scopes, + k.is_enabled AS key_enabled, + p.is_enabled AS project_enabled, + k.expires_at, + k.revoked_at + FROM gateway_api_keys k + JOIN gateway_projects p ON p.id = k.project_id + WHERE k.id = $1 + "#, + ) + .bind(key_id) + .fetch_optional(pool) + .await + } +} + +impl OpenAiModel { + pub async fn list_for_user(pool: &PgPool, user_id: &Uuid) -> Result, sqlx::Error> { + sqlx::query_as::<_, Self>( + r#" + SELECT + m.model_id AS id, + 'model'::text AS object, + EXTRACT(EPOCH FROM m.created_at)::bigint AS created, + p.name AS owned_by + FROM models m + JOIN providers p ON p.id = m.provider_id + WHERE m.is_enabled = TRUE + AND p.is_enabled = TRUE + AND EXISTS ( + SELECT 1 + FROM team_members tm + JOIN teams t ON t.id = tm.team_id + LEFT JOIN team_model_access model_access + ON model_access.team_id = t.id AND model_access.model_id = m.id + LEFT JOIN team_model_access provider_access + ON provider_access.team_id = t.id AND provider_access.provider_id = p.id + WHERE tm.user_id = $1 + AND (t.allow_all_models OR model_access.id IS NOT NULL OR provider_access.id IS NOT NULL) + ) + ORDER BY m.model_id + "#, + ) + .bind(user_id) + .fetch_all(pool) + .await + } +} + +fn parse_key_id(token: &str) -> Option { + let mut parts = token.splitn(3, '_'); + (parts.next() == Some("oxc")) + .then(|| parts.next()) + .flatten() + .and_then(|value| Uuid::parse_str(value).ok()) + .filter(|_| parts.next().is_some_and(|secret| secret.len() >= 32)) +} diff --git a/src/types/gateway/responses.rs b/src/types/gateway/responses.rs new file mode 100644 index 0000000..e051fc6 --- /dev/null +++ b/src/types/gateway/responses.rs @@ -0,0 +1,29 @@ +use serde::Serialize; + +#[derive(Debug, Serialize)] +pub struct OpenAiErrorResponse { + pub error: OpenAiError, +} + +#[derive(Debug, Serialize)] +pub struct OpenAiError { + pub message: String, + #[serde(rename = "type")] + pub kind: String, + pub param: Option, + pub code: String, +} + +#[derive(Debug, Serialize, sqlx::FromRow)] +pub struct OpenAiModel { + pub id: String, + pub object: String, + pub created: i64, + pub owned_by: String, +} + +#[derive(Debug, Serialize)] +pub struct OpenAiModelsResponse { + pub object: String, + pub data: Vec, +} diff --git a/src/types/mod.rs b/src/types/mod.rs index 6628680..93d8a16 100644 --- a/src/types/mod.rs +++ b/src/types/mod.rs @@ -30,6 +30,7 @@ pub mod base; pub mod budgets; pub mod catalog; pub mod chat; +pub mod gateway; pub mod i18n; pub mod images; pub mod models; @@ -49,6 +50,7 @@ pub use auth::*; pub use base::*; pub use budgets::*; pub use chat::*; +pub use gateway::*; pub use i18n::*; pub use images::*; pub use permissions::*; diff --git a/src/utils/mod.rs b/src/utils/mod.rs index b41a292..696f0a1 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -2,6 +2,7 @@ pub mod auth; pub mod encryption; pub mod images; pub mod oauth; +pub mod openai_gateway; pub mod provider_billing; pub mod providers; pub mod response; diff --git a/src/utils/openai_gateway.rs b/src/utils/openai_gateway.rs new file mode 100644 index 0000000..0fb0fba --- /dev/null +++ b/src/utils/openai_gateway.rs @@ -0,0 +1,122 @@ +use crate::ai; +use crate::types::{GatewayAuthContext, OpenAiError, OpenAiErrorResponse}; +use axum::body::Body; +use axum::http::{HeaderMap, HeaderValue, StatusCode, header}; +use axum::response::{IntoResponse, Response}; +use omniference::server::SkinAwareJson; +use omniference::skins::{OpenAIChatSkin, OpenAIResponsesSkin, SkinContext}; +use omniference::types::providers::openai::{OpenAIChatRequest, OpenAIResponsesRequestPayload}; +use sqlx::PgPool; +use uuid::Uuid; + +pub async fn authenticate(pool: &PgPool, headers: &HeaderMap, scope: &str) -> Result { + let Some(value) = headers.get(header::AUTHORIZATION).and_then(|value| value.to_str().ok()) else { + return Err(error_response( + StatusCode::UNAUTHORIZED, + "Missing bearer authentication", + "authentication_error", + "invalid_api_key", + )); + }; + let Some((scheme, token)) = value.split_once(' ') else { + return Err(error_response( + StatusCode::UNAUTHORIZED, + "Invalid bearer authentication", + "authentication_error", + "invalid_api_key", + )); + }; + if !scheme.eq_ignore_ascii_case("bearer") || token.is_empty() { + return Err(error_response( + StatusCode::UNAUTHORIZED, + "Invalid bearer authentication", + "authentication_error", + "invalid_api_key", + )); + } + let context = crate::types::GatewayCredential::authenticate(pool, token).await.map_err(|error| match error { + crate::types::GatewayAuthError::Invalid => error_response(StatusCode::UNAUTHORIZED, "Incorrect API key provided", "authentication_error", "invalid_api_key"), + crate::types::GatewayAuthError::Unavailable => error_response( + StatusCode::INTERNAL_SERVER_ERROR, + "Authentication service unavailable", + "server_error", + "gateway_auth_unavailable", + ), + })?; + if !context.allows(scope) { + return Err(error_response( + StatusCode::FORBIDDEN, + "API key is not permitted to use this endpoint", + "permission_error", + "insufficient_scope", + )); + } + Ok(context) +} + +pub async fn run_chat(request: OpenAIChatRequest) -> Response { + let engine = ai::get(); + let engine = engine.read().await; + let context = SkinContext::with_service(engine.service().clone()); + drop(engine); + let response = OpenAIChatSkin::handle_chat(axum07::extract::State(context), SkinAwareJson(request)).await; + bridge_response(response) +} + +pub async fn run_responses(request: OpenAIResponsesRequestPayload) -> Response { + let engine = ai::get(); + let engine = engine.read().await; + let context = SkinContext::with_service(engine.service().clone()); + drop(engine); + let response = OpenAIResponsesSkin::handle_responses(axum07::extract::State(context), SkinAwareJson(request)).await; + bridge_response(response) +} + +pub fn parse_json(body: &[u8]) -> Result { + serde_json::from_slice(body).map_err(|error| format!("Failed to parse request body: {error}")) +} + +#[must_use] +pub fn error_response(status: StatusCode, message: impl Into, kind: impl Into, code: impl Into) -> Response { + ( + status, + axum::Json(OpenAiErrorResponse { + error: OpenAiError { + message: message.into(), + kind: kind.into(), + param: None, + code: code.into(), + }, + }), + ) + .into_response() +} + +#[must_use] +pub fn add_gateway_headers(mut response: Response, context: Option<&GatewayAuthContext>, request_id: Uuid) -> Response { + if let Ok(value) = HeaderValue::from_str(&request_id.to_string()) { + response.headers_mut().insert("x-request-id", value); + } + if let Some(context) = context { + if let Ok(value) = HeaderValue::from_str(&context.project_id.to_string()) { + response.headers_mut().insert("x-oxide-project-id", value); + } + if let Ok(value) = HeaderValue::from_str(&context.key_id.to_string()) { + response.headers_mut().insert("x-oxide-api-key-id", value); + } + if let Some(team_id) = context.team_id + && let Ok(value) = HeaderValue::from_str(&team_id.to_string()) + { + response.headers_mut().insert("x-oxide-team-id", value); + } + if let Ok(value) = HeaderValue::from_str(&context.project_name) { + response.headers_mut().insert("x-oxide-project-name", value); + } + } + response +} + +pub(crate) fn bridge_response(response: axum07::response::Response) -> Response { + let (parts, body) = response.into_parts(); + Response::from_parts(parts, Body::new(body)) +} From e7b374aa0f45a7fc66361eca2a3f7a91bfd6713e Mon Sep 17 00:00:00 2001 From: Henrik Mertens Date: Sun, 2 Aug 2026 23:53:41 +0200 Subject: [PATCH 2/2] Refactor OpenAI gateway to use Axum 0.8 and implement budget enforcement - Upgrade `axum` to 0.8 and `tower-http` to 0.6, removing `axum07` compatibility bridge - Implement middleware-based authentication and request ID tracking for OpenAI routes - Integrate `QueuedCostSink` with `OmniferenceEngine` for asynchronous usage and cost tracking - Add budget validation to gateway requests to block access when user quotas are exceeded - Refactor model resolution to support `provider/model` slug format and team-aware access policies - Utilize `omniference` skins for standardized OpenAI error handling and metadata injection - Optimize API key usage updates by throttling database writes to once per minute per key --- Cargo.lock | 139 +++------------------- Cargo.toml | 3 +- src/ai.rs | 14 ++- src/main.rs | 12 +- src/routes/mod.rs | 14 ++- src/routes/public/openai.rs | 199 +++++++++++++------------------- src/tests/gateway.rs | 83 +++++++------ src/types/gateway/mod.rs | 25 +++- src/types/gateway/repository.rs | 100 +++++++++++----- src/types/gateway/responses.rs | 29 ----- src/utils/mod.rs | 1 + src/utils/omniference_cost.rs | 76 ++++++++++++ src/utils/openai_gateway.rs | 93 ++++++++------- 13 files changed, 393 insertions(+), 395 deletions(-) delete mode 100644 src/types/gateway/responses.rs create mode 100644 src/utils/omniference_cost.rs diff --git a/Cargo.lock b/Cargo.lock index 2972f37..e889853 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -11,8 +11,7 @@ dependencies = [ "argon2", "async-stream", "async-trait", - "axum 0.7.9", - "axum 0.8.9", + "axum", "base64", "chrono", "dotenv", @@ -32,7 +31,7 @@ dependencies = [ "tokio", "tokio-util", "tower-cookies", - "tower-http 0.5.2", + "tower-http", "tracing", "tracing-subscriber", "uuid", @@ -222,47 +221,13 @@ version = "1.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" -[[package]] -name = "axum" -version = "0.7.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "edca88bc138befd0323b20752846e6587272d3b03b0343c8ea28a6f819e6e71f" -dependencies = [ - "async-trait", - "axum-core 0.4.5", - "bytes", - "futures-util", - "http", - "http-body", - "http-body-util", - "hyper", - "hyper-util", - "itoa", - "matchit 0.7.3", - "memchr", - "mime", - "percent-encoding", - "pin-project-lite", - "rustversion", - "serde", - "serde_json", - "serde_path_to_error", - "serde_urlencoded", - "sync_wrapper", - "tokio", - "tower 0.5.3", - "tower-layer", - "tower-service", - "tracing", -] - [[package]] name = "axum" version = "0.8.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "31b698c5f9a010f6573133b09e0de5408834d0c82f8d7475a89fc1867a71cd90" dependencies = [ - "axum-core 0.5.6", + "axum-core", "bytes", "form_urlencoded", "futures-util", @@ -272,7 +237,7 @@ dependencies = [ "hyper", "hyper-util", "itoa", - "matchit 0.8.4", + "matchit", "memchr", "mime", "percent-encoding", @@ -283,28 +248,7 @@ dependencies = [ "serde_urlencoded", "sync_wrapper", "tokio", - "tower 0.5.3", - "tower-layer", - "tower-service", - "tracing", -] - -[[package]] -name = "axum-core" -version = "0.4.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "09f2bd6146b97ae3359fa0cc6d6b376d9539582c7b4220f041a33ec24c226199" -dependencies = [ - "async-trait", - "bytes", - "futures-util", - "http", - "http-body", - "http-body-util", - "mime", - "pin-project-lite", - "rustversion", - "sync_wrapper", + "tower", "tower-layer", "tower-service", "tracing", @@ -2072,12 +2016,6 @@ dependencies = [ "regex-automata", ] -[[package]] -name = "matchit" -version = "0.7.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0e7465ac9959cc2b1404e8e2367b43684a6d13790fe23056cc8c6c5a6b7bcb94" - [[package]] name = "matchit" version = "0.8.4" @@ -2232,7 +2170,7 @@ dependencies = [ "anyhow", "async-stream", "async-trait", - "axum 0.7.9", + "axum", "base64", "bytes", "dotenvy", @@ -2249,8 +2187,8 @@ dependencies = [ "tokio-stream", "tokio-util", "toml 0.8.23", - "tower 0.4.13", - "tower-http 0.5.2", + "tower", + "tower-http", "tracing", "tracing-subscriber", "uuid", @@ -2367,26 +2305,6 @@ dependencies = [ "indexmap", ] -[[package]] -name = "pin-project" -version = "1.1.13" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2466b2336ed02bcdca6b294417127b90ec92038d1d5c4fbeac971a922e0e0924" -dependencies = [ - "pin-project-internal", -] - -[[package]] -name = "pin-project-internal" -version = "1.1.13" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c96395f0a926bc13b1c17622aaddda1ecb55d49c8f1bf9777e4d877800a43f8b" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.119", -] - [[package]] name = "pin-project-lite" version = "0.2.17" @@ -2842,8 +2760,8 @@ dependencies = [ "tokio-native-tls", "tokio-rustls", "tokio-util", - "tower 0.5.3", - "tower-http 0.6.11", + "tower", + "tower-http", "tower-service", "url", "wasm-bindgen", @@ -3902,21 +3820,6 @@ version = "1.1.2+spec-1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7d56353a2a665ad0f41a421187180aab746c8c325620617ad883a99a1cbe66d2" -[[package]] -name = "tower" -version = "0.4.13" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b8fa9be0de6cf49e536ce1851f987bd21a43b771b09473c3549a6c853db37c1c" -dependencies = [ - "futures-core", - "futures-util", - "pin-project", - "pin-project-lite", - "tower-layer", - "tower-service", - "tracing", -] - [[package]] name = "tower" version = "0.5.3" @@ -3939,7 +3842,7 @@ version = "0.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "151b5a3e3c45df17466454bb74e9ecedecc955269bdedbf4d150dfa393b55a36" dependencies = [ - "axum-core 0.5.6", + "axum-core", "cookie", "futures-util", "http", @@ -3949,23 +3852,6 @@ dependencies = [ "tower-service", ] -[[package]] -name = "tower-http" -version = "0.5.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1e9cd434a998747dd2c4276bc96ee2e0c7a2eadf3cae88e52be55a05fa9053f5" -dependencies = [ - "bitflags", - "bytes", - "http", - "http-body", - "http-body-util", - "pin-project-lite", - "tower-layer", - "tower-service", - "tracing", -] - [[package]] name = "tower-http" version = "0.6.11" @@ -3978,9 +3864,10 @@ dependencies = [ "http", "http-body", "pin-project-lite", - "tower 0.5.3", + "tower", "tower-layer", "tower-service", + "tracing", "url", ] diff --git a/Cargo.toml b/Cargo.toml index 363a6f2..3e8d076 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -11,7 +11,6 @@ categories = ["web-programming"] [dependencies] axum = { version = "0.8" } -axum07 = { package = "axum", version = "0.7" } tokio = { version = "1.48", features = ["full"] } serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" @@ -37,7 +36,7 @@ reqwest = { version = "0.12", default-features = false, features = [ "json", "multipart", ] } -tower-http = { version = "0.5", features = ["cors"] } +tower-http = { version = "0.6", features = ["cors"] } omniference = { version = "0.3" } aes-gcm = "0.10" rand = "0.8" diff --git a/src/ai.rs b/src/ai.rs index c6300c5..9c338ee 100644 --- a/src/ai.rs +++ b/src/ai.rs @@ -8,12 +8,14 @@ use omniference::{ OmniferenceEngine, + middleware::cost::QueuedCostSink, types::{ProviderConfig, ProviderEndpoint}, }; use sqlx::PgPool; use std::{collections::BTreeMap, sync::Arc}; use tokio::sync::RwLock; +use crate::types::JobState; use crate::types::providers::Provider; use crate::utils::encryption::decrypt_api_key; @@ -21,8 +23,10 @@ use crate::utils::encryption::decrypt_api_key; pub static OF_ENGINE: std::sync::OnceLock>> = std::sync::OnceLock::new(); /// Initialize the AI engine with providers from the database -pub async fn init(pool: &PgPool) { - let engine = Arc::new(RwLock::new(OmniferenceEngine::new())); +pub async fn init(state: &Arc) { + let pool = &state.db; + let cost_sink = QueuedCostSink::spawn(Arc::new(crate::utils::omniference_cost::OxideCostSink::new(Arc::clone(state)))); + let engine = Arc::new(RwLock::new(OmniferenceEngine::with_cost_sink(cost_sink))); let providers = Provider::list_enabled_system(pool).await.unwrap_or_default(); @@ -69,13 +73,15 @@ pub async fn catalog() -> Option> { pub async fn sync_pricing_overrides(_pool: &PgPool) {} /// Reload providers from the database -pub async fn reload_providers(pool: &PgPool) { +pub async fn reload_providers(state: &Arc) { + let pool = &state.db; let providers = Provider::list_enabled_system(pool).await.unwrap_or_default(); let provider_count = providers.len(); // Create a fresh engine with new providers - let new_engine = OmniferenceEngine::new(); + let cost_sink = QueuedCostSink::spawn(Arc::new(crate::utils::omniference_cost::OxideCostSink::new(Arc::clone(state)))); + let new_engine = OmniferenceEngine::with_cost_sink(cost_sink); let engine_arc = get(); let mut engine_write = engine_arc.write().await; *engine_write = new_engine; diff --git a/src/main.rs b/src/main.rs index 5ecdfcc..802f9d5 100644 --- a/src/main.rs +++ b/src/main.rs @@ -69,19 +69,17 @@ async fn main() { i18n::I18n::init(&pool).await; println!("[I18N] Translations loaded"); - ai::init(&pool).await; - let app_state = Arc::new(JobState { - db: pool.clone(), + db: pool, mcp_pool: crate::utils::tools::McpConnectionPool::new(), client_tool_pending: crate::types::state::ClientToolPending::new(), }); + ai::init(&app_state).await; tokio::spawn(jobs::start_job_scheduler(app_state.clone())); let app = Router::new() - .merge(routes::build_router()) - .with_state(app_state) + .merge(routes::build_router(Arc::clone(&app_state))) .layer(DefaultBodyLimit::max(8 * 1024 * 1024)); let address = format!( @@ -94,7 +92,7 @@ async fn main() { println!("[SERVER] Listening on http://{}", address); - let pool_for_shutdown = pool.clone(); + let state_for_shutdown = Arc::clone(&app_state); let server = axum::serve(listener, app).with_graceful_shutdown(async move { shutdown_signal().await; @@ -102,7 +100,7 @@ async fn main() { println!("[SERVER] Shutdown signal received"); println!("[DATABASE] Closing pool..."); - pool_for_shutdown.close().await; + state_for_shutdown.db.close().await; println!("[DATABASE] Pool closed"); }); diff --git a/src/routes/mod.rs b/src/routes/mod.rs index 56545a9..0bc91a6 100644 --- a/src/routes/mod.rs +++ b/src/routes/mod.rs @@ -6,18 +6,25 @@ pub mod public; use crate::types::JobState; use axum::{ Router, + middleware::{from_fn, from_fn_with_state}, routing::{delete, get, patch, post, put}, }; use std::sync::Arc; use tower_cookies::CookieManagerLayer; -pub fn build_router() -> Router> { - let openai_router = Router::new() +pub fn build_router(state: Arc) -> Router { + let openai_read_router = Router::new() .route("/models", get(public::openai::list_models)) + .route_layer(from_fn_with_state(state.clone(), crate::utils::openai_gateway::authenticate_read)); + let openai_write_router = Router::new() .route("/chat/completions", post(public::openai::chat_completions)) .route("/responses", post(public::openai::responses)) + .route_layer(from_fn_with_state(state.clone(), crate::utils::openai_gateway::authenticate_write)); + let openai_router = openai_read_router + .merge(openai_write_router) .method_not_allowed_fallback(public::openai::method_not_allowed) - .fallback(public::openai::not_found); + .fallback(public::openai::not_found) + .layer(from_fn(crate::utils::openai_gateway::add_request_id)); Router::new() .nest("/openai/v1", openai_router) .route("/api/v1/health", get(public::base::health)) @@ -166,4 +173,5 @@ pub fn build_router() -> Router> { .route("/api/v1/images/{id}", get(public::images::serve_image)) .route("/api/v1/images", post(public::images::upload_image)) .layer(CookieManagerLayer::new()) + .with_state(state) } diff --git a/src/routes/public/openai.rs b/src/routes/public/openai.rs index 687f371..2a14c51 100644 --- a/src/routes/public/openai.rs +++ b/src/routes/public/openai.rs @@ -1,141 +1,95 @@ -use crate::types::models::Model; -use crate::types::{JobState, OpenAiModel, OpenAiModelsResponse}; -use crate::utils::openai_gateway::{add_gateway_headers, authenticate, error_response, parse_json, run_chat, run_responses}; -use axum::body::Bytes; -use axum::extract::State; -use axum::http::{HeaderMap, StatusCode}; +use crate::types::models::ModelPricing; +use crate::types::{Budget, GatewayAuthContext, GatewayModel, JobState}; +use crate::utils::openai_gateway::{error_response, run_chat, run_responses}; +use axum::extract::{Extension, State}; +use axum::http::StatusCode; use axum::response::{IntoResponse, Response}; +use omniference::server::SkinAwareJson; +use omniference::skins::{OpenAIErrorHandler, SkinErrorHandler, SkinRequestMetadata}; +use omniference::types::providers::OpenAIModelsResponse; use omniference::types::providers::openai::{OpenAIChatRequest, OpenAIResponsesRequestPayload}; use std::sync::Arc; use uuid::Uuid; pub async fn not_found() -> Response { - add_gateway_headers( - error_response(StatusCode::NOT_FOUND, "The requested resource was not found", "not_found_error", "not_found"), - None, - Uuid::new_v4(), - ) + OpenAIErrorHandler.handle_not_found() } pub async fn method_not_allowed() -> Response { - add_gateway_headers( - error_response( - StatusCode::METHOD_NOT_ALLOWED, - "Invalid HTTP method for this endpoint", - "invalid_request_error", - "method_not_allowed", - ), - None, - Uuid::new_v4(), - ) + OpenAIErrorHandler.handle_method_not_allowed() } -pub async fn list_models(State(state): State>, headers: HeaderMap) -> impl IntoResponse { - let request_id = Uuid::new_v4(); - let context = match authenticate(&state.db, &headers, "inference:read").await { - Ok(context) => context, - Err(response) => return add_gateway_headers(response, None, request_id), - }; - let models = match OpenAiModel::list_for_user(&state.db, &context.user_id).await { +pub async fn list_models(State(state): State>, Extension(context): Extension) -> Response { + let models = match GatewayModel::list_for_context(&state.db, &context).await { Ok(models) => models, Err(error) => { tracing::error!(%error, "failed to list gateway models"); - return add_gateway_headers( - error_response(StatusCode::INTERNAL_SERVER_ERROR, "Failed to list models", "server_error", "internal_error"), - Some(&context), - request_id, - ); + return error_response(StatusCode::INTERNAL_SERVER_ERROR, "Failed to list models", "server_error", "internal_error"); } }; - add_gateway_headers( - axum::Json(OpenAiModelsResponse { - object: "list".to_string(), - data: models, - }) - .into_response(), - Some(&context), - request_id, - ) + axum::Json(OpenAIModelsResponse { + object: Some("list".to_string()), + data: models, + }) + .into_response() } -pub async fn chat_completions(State(state): State>, headers: HeaderMap, body: Bytes) -> Response { - let request_id = Uuid::new_v4(); - let context = match authenticate(&state.db, &headers, "inference:write").await { - Ok(context) => context, - Err(response) => return add_gateway_headers(response, None, request_id), - }; - let request: OpenAIChatRequest = match parse_json(&body) { - Ok(request) => request, - Err(message) => { - return add_gateway_headers( - error_response(StatusCode::BAD_REQUEST, message, "invalid_request_error", "invalid_request_body"), - Some(&context), - request_id, - ); - } +pub async fn chat_completions( + State(state): State>, + Extension(context): Extension, + SkinAwareJson(request): SkinAwareJson, +) -> Response { + let model_id = match resolve_model_access(&state, &context, &request.model).await { + Ok(model_id) => model_id, + Err(response) => return response, }; - if let Some(response) = validate_model_access(&state, &context.user_id, &request.model).await { - return add_gateway_headers(response, Some(&context), request_id); - } - let response = run_chat(request).await; - add_gateway_headers(response, Some(&context), request_id) + run_chat(request, request_metadata(&context, model_id)).await } -pub async fn responses(State(state): State>, headers: HeaderMap, body: Bytes) -> Response { - let request_id = Uuid::new_v4(); - let context = match authenticate(&state.db, &headers, "inference:write").await { - Ok(context) => context, - Err(response) => return add_gateway_headers(response, None, request_id), - }; - let request: OpenAIResponsesRequestPayload = match parse_json(&body) { - Ok(request) => request, - Err(message) => { - return add_gateway_headers( - error_response(StatusCode::BAD_REQUEST, message, "invalid_request_error", "invalid_request_body"), - Some(&context), - request_id, - ); - } - }; +pub async fn responses( + State(state): State>, + Extension(context): Extension, + SkinAwareJson(request): SkinAwareJson, +) -> Response { let Some(model) = request.model.as_deref() else { - return add_gateway_headers( - error_response( - StatusCode::BAD_REQUEST, - "Missing required parameter: 'model'.", - "invalid_request_error", - "missing_required_parameter", - ), - Some(&context), - request_id, + return error_response( + StatusCode::BAD_REQUEST, + "Missing required parameter: 'model'.", + "invalid_request_error", + "missing_required_parameter", ); }; - if let Some(response) = validate_model_access(&state, &context.user_id, model).await { - return add_gateway_headers(response, Some(&context), request_id); - } - let response = run_responses(request).await; - add_gateway_headers(response, Some(&context), request_id) + let model_id = match resolve_model_access(&state, &context, model).await { + Ok(model_id) => model_id, + Err(response) => return response, + }; + run_responses(request, request_metadata(&context, model_id)).await } -async fn validate_model_access(state: &JobState, user_id: &Uuid, model_id: &str) -> Option { - let model = match Model::find_by_model_id(&state.db, model_id).await { - Ok(Some(model)) if model.is_enabled => model, - Ok(_) => return Some(model_not_found(model_id)), - Err(error) => { - tracing::error!(%error, model_id, "failed to resolve gateway model"); - return Some(error_response( - StatusCode::INTERNAL_SERVER_ERROR, - "Failed to resolve model", - "server_error", - "internal_error", - )); - } - }; - match Model::can_user_use_model(&state.db, user_id, &model.id).await { - Ok(true) => None, - Ok(false) => Some(model_not_found(model_id)), +async fn resolve_model_access(state: &JobState, context: &GatewayAuthContext, model_id: &str) -> Result { + match GatewayModel::resolve_accessible(&state.db, context, model_id).await { + Ok(Some(model_id)) => match budget_allows_model(state, &context.user_id, &model_id).await { + Ok(true) => Ok(model_id), + Ok(false) => Err(error_response( + StatusCode::TOO_MANY_REQUESTS, + "Budget exceeded for this model", + "insufficient_quota", + "budget_exceeded", + )), + Err(error) => { + tracing::error!(%error, %model_id, "failed to evaluate gateway budget policy"); + Err(error_response( + StatusCode::INTERNAL_SERVER_ERROR, + "Failed to check budget status", + "server_error", + "internal_error", + )) + } + }, + Ok(None) => Err(OpenAIErrorHandler.handle_model_not_found(model_id)), Err(error) => { tracing::error!(%error, model_id, "failed to evaluate gateway model policy"); - Some(error_response( + Err(error_response( StatusCode::INTERNAL_SERVER_ERROR, "Failed to resolve model", "server_error", @@ -145,11 +99,22 @@ async fn validate_model_access(state: &JobState, user_id: &Uuid, model_id: &str) } } -fn model_not_found(model_id: &str) -> Response { - error_response( - StatusCode::NOT_FOUND, - format!("Model '{model_id}' not found"), - "invalid_request_error", - "model_not_found", - ) +async fn budget_allows_model(state: &JobState, user_id: &Uuid, model_id: &Uuid) -> Result { + if ModelPricing::is_free(&state.db, model_id).await? { + return Ok(true); + } + let status = Budget::status_for_user(&state.db, user_id).await?; + Ok(!status.blocked_model_ids.contains(model_id)) +} + +fn request_metadata(context: &GatewayAuthContext, model_id: Uuid) -> SkinRequestMetadata { + let mut metadata = std::collections::BTreeMap::new(); + metadata.insert("oxide_user_id".to_string(), context.user_id.to_string()); + metadata.insert("oxide_project_id".to_string(), context.project_id.to_string()); + metadata.insert("oxide_api_key_id".to_string(), context.key_id.to_string()); + metadata.insert("oxide_model_id".to_string(), model_id.to_string()); + if let Some(team_id) = context.team_id { + metadata.insert("oxide_team_id".to_string(), team_id.to_string()); + } + SkinRequestMetadata(metadata) } diff --git a/src/tests/gateway.rs b/src/tests/gateway.rs index 61c5446..67b0f3b 100644 --- a/src/tests/gateway.rs +++ b/src/tests/gateway.rs @@ -1,30 +1,15 @@ #[cfg(test)] mod tests { - use crate::types::{GatewayAuthError, GatewayCredential, OpenAiModel}; + use crate::types::{GatewayAuthError, GatewayCredential, GatewayModel}; use crate::utils::auth::hash_password; - use crate::utils::openai_gateway::bridge_response; - use axum::body::Bytes; use omniference::skins::{OpenAIChatSkin, OpenAIResponsesSkin, Skin}; use omniference::types::providers::openai::{OpenAIChatRequest, OpenAIResponsesRequestPayload}; use omniference::types::{ModelRef, ProviderConfig, ProviderEndpoint, ProviderKind}; use serde_json::json; use sqlx::PgPool; use std::collections::BTreeMap; - use std::convert::Infallible; - use std::sync::{ - Arc, - atomic::{AtomicBool, Ordering}, - }; use uuid::Uuid; - struct DropFlag(Arc); - - impl Drop for DropFlag { - fn drop(&mut self) { - self.0.store(true, Ordering::SeqCst); - } - } - async fn create_user(pool: &PgPool) -> Uuid { sqlx::query_scalar("INSERT INTO users (email, username, password_hash) VALUES ('gateway@example.com', 'gateway', 'hash') RETURNING id") .fetch_one(pool) @@ -33,13 +18,18 @@ mod tests { } async fn create_key(pool: &PgPool, user_id: Uuid, scopes: serde_json::Value) -> (Uuid, String) { - let project_id: Uuid = sqlx::query_scalar("INSERT INTO gateway_projects (owner_id, name) VALUES ($1, 'Gateway') RETURNING id") + create_team_key(pool, user_id, None, scopes).await + } + + async fn create_team_key(pool: &PgPool, user_id: Uuid, team_id: Option, scopes: serde_json::Value) -> (Uuid, String) { + let project_id: Uuid = sqlx::query_scalar("INSERT INTO gateway_projects (owner_id, team_id, name) VALUES ($1, $2, 'Gateway') RETURNING id") .bind(user_id) + .bind(team_id) .fetch_one(pool) .await .unwrap(); let key_id = Uuid::new_v4(); - let secret = "abcdefghijklmnopqrstuvwxyz0123456789"; + let secret = "abcdefghijklmnopqrstuvwxyz_0123456789"; let token = format!("oxc_{}_{}", key_id.simple(), secret); sqlx::query( "INSERT INTO gateway_api_keys (id, project_id, name, secret_hash, key_prefix, last_four, scopes) @@ -152,9 +142,48 @@ mod tests { .execute(&pool) .await .unwrap(); - let models = OpenAiModel::list_for_user(&pool, &user_id).await.unwrap(); + let (_, token) = create_key(&pool, user_id, json!(["inference:read"])).await; + let context = GatewayCredential::authenticate(&pool, &token).await.unwrap(); + let models = GatewayModel::list_for_context(&pool, &context).await.unwrap(); assert_eq!(models.len(), 1); - assert_eq!(models[0].id, "visible"); + assert_eq!(models[0].id, "gateway provider/visible"); + assert!(GatewayModel::resolve_accessible(&pool, &context, "gateway provider/visible").await.unwrap().is_some()); + assert!(GatewayModel::resolve_accessible(&pool, &context, "gateway provider/hidden").await.unwrap().is_none()); + } + + #[sqlx::test(migrations = "./migrations")] + async fn team_project_only_uses_its_team_policy(pool: PgPool) { + let user_id = create_user(&pool).await; + let project_team_id: Uuid = sqlx::query_scalar("INSERT INTO teams (name, allow_all_models) VALUES ('Project Team', false) RETURNING id") + .fetch_one(&pool) + .await + .unwrap(); + let other_team_id: Uuid = sqlx::query_scalar("INSERT INTO teams (name, allow_all_models) VALUES ('Other Team', true) RETURNING id") + .fetch_one(&pool) + .await + .unwrap(); + for team_id in [project_team_id, other_team_id] { + sqlx::query("INSERT INTO team_members (team_id, user_id) VALUES ($1, $2)") + .bind(team_id) + .bind(user_id) + .execute(&pool) + .await + .unwrap(); + } + let provider_id: Uuid = + sqlx::query_scalar("INSERT INTO providers (kind, name, base_url, is_enabled) VALUES ('OPENAI', 'Scoped Provider', 'https://example.com', true) RETURNING id") + .fetch_one(&pool) + .await + .unwrap(); + sqlx::query("INSERT INTO models (provider_id, model_id, display_name, is_enabled) VALUES ($1, 'scoped', 'Scoped', true)") + .bind(provider_id) + .execute(&pool) + .await + .unwrap(); + let (_, token) = create_team_key(&pool, user_id, Some(project_team_id), json!(["inference:read"])).await; + let context = GatewayCredential::authenticate(&pool, &token).await.unwrap(); + assert!(GatewayModel::list_for_context(&pool, &context).await.unwrap().is_empty()); + assert!(GatewayModel::resolve_accessible(&pool, &context, "scoped provider/scoped").await.unwrap().is_none()); } #[test] @@ -187,18 +216,4 @@ mod tests { assert_eq!(ir.messages.len(), 1); } } - - #[tokio::test] - async fn dropping_bridged_stream_releases_the_upstream_body() { - let dropped = Arc::new(AtomicBool::new(false)); - let guard = DropFlag(dropped.clone()); - let stream = futures_util::stream::once(async move { - let _guard = guard; - std::future::pending::>().await - }); - let response = axum07::response::Response::new(axum07::body::Body::from_stream(stream)); - let response = bridge_response(response); - drop(response); - assert!(dropped.load(Ordering::SeqCst)); - } } diff --git a/src/types/gateway/mod.rs b/src/types/gateway/mod.rs index 6097ea1..cfdf517 100644 --- a/src/types/gateway/mod.rs +++ b/src/types/gateway/mod.rs @@ -1,12 +1,13 @@ mod repository; -mod responses; - -pub use responses::*; use chrono::{DateTime, Utc}; +use omniference::types::providers::OpenAIModel; use sqlx::types::Json; use uuid::Uuid; +pub const INFERENCE_READ_SCOPE: &str = "inference:read"; +pub const INFERENCE_WRITE_SCOPE: &str = "inference:write"; + #[derive(Debug, sqlx::FromRow)] pub struct GatewayCredential { pub key_id: Uuid, @@ -39,6 +40,24 @@ impl GatewayAuthContext { } } +#[derive(Debug, sqlx::FromRow)] +pub struct GatewayModel { + pub provider_name: String, + pub model_id: String, + pub created_at: DateTime, +} + +impl From for OpenAIModel { + fn from(model: GatewayModel) -> Self { + Self { + id: format!("{}/{}", model.provider_name.to_ascii_lowercase(), model.model_id), + object: Some("model".to_string()), + created: Some(model.created_at.timestamp().max(0) as u64), + owned_by: Some(model.provider_name), + } + } +} + #[derive(Debug, thiserror::Error)] pub enum GatewayAuthError { #[error("invalid_api_key")] diff --git a/src/types/gateway/repository.rs b/src/types/gateway/repository.rs index ced1b64..861e2b4 100644 --- a/src/types/gateway/repository.rs +++ b/src/types/gateway/repository.rs @@ -1,6 +1,7 @@ -use super::{GatewayAuthContext, GatewayAuthError, GatewayCredential, OpenAiModel}; +use super::{GatewayAuthContext, GatewayAuthError, GatewayCredential, GatewayModel}; use crate::utils::auth::{hash_password, verify_password}; use chrono::Utc; +use omniference::types::providers::OpenAIModel; use sqlx::PgPool; use std::sync::LazyLock; use uuid::Uuid; @@ -9,13 +10,13 @@ static DUMMY_SECRET_HASH: LazyLock = LazyLock::new(|| hash_password("oxi impl GatewayCredential { pub async fn authenticate(pool: &PgPool, token: &str) -> Result { - let key_id = parse_key_id(token); - let credential = match key_id { - Some(id) => Self::find(pool, &id).await.map_err(|_| GatewayAuthError::Unavailable)?, + let parsed = parse_token(token); + let credential = match parsed.as_ref() { + Some(parsed) => Self::find(pool, &parsed.key_id).await.map_err(|_| GatewayAuthError::Unavailable)?, None => None, }; let hash = credential.as_ref().map_or(DUMMY_SECRET_HASH.as_str(), |value| value.secret_hash.as_str()); - let secret = token.rsplit_once('_').map_or("", |(_, value)| value).to_owned(); + let secret = parsed.map_or_else(String::new, |parsed| parsed.secret.to_owned()); let hash = hash.to_owned(); let verified = tokio::task::spawn_blocking(move || verify_password(&secret, &hash).unwrap_or(false)) .await @@ -31,11 +32,14 @@ impl GatewayCredential { { return Err(GatewayAuthError::Invalid); } - sqlx::query("UPDATE gateway_api_keys SET last_used_at = NOW() WHERE id = $1") - .bind(credential.key_id) - .execute(pool) - .await - .map_err(|_| GatewayAuthError::Unavailable)?; + if let Err(error) = + sqlx::query("UPDATE gateway_api_keys SET last_used_at = NOW() WHERE id = $1 AND (last_used_at IS NULL OR last_used_at < NOW() - INTERVAL '1 minute')") + .bind(credential.key_id) + .execute(pool) + .await + { + tracing::warn!(%error, key_id = %credential.key_id, "failed to update gateway API key usage timestamp"); + } Ok(GatewayAuthContext { key_id: credential.key_id, project_id: credential.project_id, @@ -72,15 +76,14 @@ impl GatewayCredential { } } -impl OpenAiModel { - pub async fn list_for_user(pool: &PgPool, user_id: &Uuid) -> Result, sqlx::Error> { - sqlx::query_as::<_, Self>( +impl GatewayModel { + pub async fn list_for_context(pool: &PgPool, context: &GatewayAuthContext) -> Result, sqlx::Error> { + let models = sqlx::query_as::<_, Self>( r#" SELECT - m.model_id AS id, - 'model'::text AS object, - EXTRACT(EPOCH FROM m.created_at)::bigint AS created, - p.name AS owned_by + p.name AS provider_name, + m.model_id, + m.created_at FROM models m JOIN providers p ON p.id = m.provider_id WHERE m.is_enabled = TRUE @@ -94,22 +97,65 @@ impl OpenAiModel { LEFT JOIN team_model_access provider_access ON provider_access.team_id = t.id AND provider_access.provider_id = p.id WHERE tm.user_id = $1 + AND ($2::uuid IS NULL OR t.id = $2) AND (t.allow_all_models OR model_access.id IS NOT NULL OR provider_access.id IS NOT NULL) ) - ORDER BY m.model_id + ORDER BY p.name, m.model_id "#, ) - .bind(user_id) + .bind(context.user_id) + .bind(context.team_id) .fetch_all(pool) - .await + .await?; + Ok(models.into_iter().map(OpenAIModel::from).collect()) + } + + pub async fn resolve_accessible(pool: &PgPool, context: &GatewayAuthContext, requested_id: &str) -> Result, sqlx::Error> { + let (provider_name, model_id) = requested_id.split_once('/').map_or((None, requested_id), |(provider, model)| (Some(provider), model)); + let matches = sqlx::query_scalar::<_, Uuid>( + r#" + SELECT m.id + FROM models m + JOIN providers p ON p.id = m.provider_id + WHERE m.model_id = $3 + AND ($4::text IS NULL OR LOWER(p.name) = LOWER($4)) + AND m.is_enabled = TRUE + AND p.is_enabled = TRUE + AND EXISTS ( + SELECT 1 + FROM team_members tm + JOIN teams t ON t.id = tm.team_id + LEFT JOIN team_model_access model_access + ON model_access.team_id = t.id AND model_access.model_id = m.id + LEFT JOIN team_model_access provider_access + ON provider_access.team_id = t.id AND provider_access.provider_id = p.id + WHERE tm.user_id = $1 + AND ($2::uuid IS NULL OR t.id = $2) + AND (t.allow_all_models OR model_access.id IS NOT NULL OR provider_access.id IS NOT NULL) + ) + LIMIT 2 + "#, + ) + .bind(context.user_id) + .bind(context.team_id) + .bind(model_id) + .bind(provider_name) + .fetch_all(pool) + .await?; + Ok((matches.len() == 1).then(|| matches[0])) } } -fn parse_key_id(token: &str) -> Option { - let mut parts = token.splitn(3, '_'); - (parts.next() == Some("oxc")) - .then(|| parts.next()) - .flatten() - .and_then(|value| Uuid::parse_str(value).ok()) - .filter(|_| parts.next().is_some_and(|secret| secret.len() >= 32)) +struct ParsedToken<'a> { + key_id: Uuid, + secret: &'a str, +} + +fn parse_token(token: &str) -> Option> { + let remainder = token.strip_prefix("oxc_")?; + let (key_id, secret) = remainder.split_once('_')?; + (secret.len() >= 32).then_some(ParsedToken { + key_id: Uuid::parse_str(key_id).ok()?, + secret, + }) } diff --git a/src/types/gateway/responses.rs b/src/types/gateway/responses.rs deleted file mode 100644 index e051fc6..0000000 --- a/src/types/gateway/responses.rs +++ /dev/null @@ -1,29 +0,0 @@ -use serde::Serialize; - -#[derive(Debug, Serialize)] -pub struct OpenAiErrorResponse { - pub error: OpenAiError, -} - -#[derive(Debug, Serialize)] -pub struct OpenAiError { - pub message: String, - #[serde(rename = "type")] - pub kind: String, - pub param: Option, - pub code: String, -} - -#[derive(Debug, Serialize, sqlx::FromRow)] -pub struct OpenAiModel { - pub id: String, - pub object: String, - pub created: i64, - pub owned_by: String, -} - -#[derive(Debug, Serialize)] -pub struct OpenAiModelsResponse { - pub object: String, - pub data: Vec, -} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 696f0a1..791a7ac 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -2,6 +2,7 @@ pub mod auth; pub mod encryption; pub mod images; pub mod oauth; +pub mod omniference_cost; pub mod openai_gateway; pub mod provider_billing; pub mod providers; diff --git a/src/utils/omniference_cost.rs b/src/utils/omniference_cost.rs new file mode 100644 index 0000000..d7a0c95 --- /dev/null +++ b/src/utils/omniference_cost.rs @@ -0,0 +1,76 @@ +use crate::types::models::{Model, ModelPricing}; +use crate::types::{Budget, JobState, UsageEvent, UsageEventRecord}; +use async_trait::async_trait; +use omniference::middleware::cost::{AsyncCostSink, CostRecord}; +use rust_decimal::Decimal; +use rust_decimal::prelude::FromPrimitive; +use std::sync::Arc; +use uuid::Uuid; + +pub struct OxideCostSink { + state: Arc, +} + +impl OxideCostSink { + #[must_use] + pub fn new(state: Arc) -> Self { + Self { state } + } +} + +#[async_trait] +impl AsyncCostSink for OxideCostSink { + async fn record(&self, record: CostRecord) { + if let Err(error) = self.persist(&record).await { + tracing::error!(%error, provider = record.provider, model = record.model, "failed to persist Omniference usage"); + } + } +} + +impl OxideCostSink { + async fn persist(&self, record: &CostRecord) -> Result<(), sqlx::Error> { + let Some(user_id) = metadata_uuid(record, "oxide_user_id") else { + return Ok(()); + }; + let Some(model_id) = metadata_uuid(record, "oxide_model_id") else { + return Ok(()); + }; + let Some(model) = Model::find_by_id(&self.state.db, &model_id).await? else { + return Ok(()); + }; + let input_tokens = token_count(record.usage.input_tokens); + let output_tokens = token_count(record.usage.output_tokens); + let reasoning_tokens = token_count(record.usage.reasoning_tokens); + let cost_total = ModelPricing::usage_cost(&self.state.db, &model.id, input_tokens, output_tokens, reasoning_tokens) + .await? + .unwrap_or_else(|| Decimal::from_f64(record.cost.total).unwrap_or(Decimal::ZERO)); + let team_id = match metadata_uuid(record, "oxide_team_id") { + Some(team_id) => Some(team_id), + None => Budget::primary_team_id(&self.state.db, &user_id).await?, + }; + UsageEvent::record( + &self.state.db, + UsageEventRecord { + user_id: &user_id, + team_id, + model_id: &model.id, + provider_id: &model.provider_id, + request_type: "gateway", + input_tokens, + output_tokens, + reasoning_tokens, + cost_total, + }, + ) + .await?; + Ok(()) + } +} + +fn metadata_uuid(record: &CostRecord, key: &str) -> Option { + record.metadata.get(key).and_then(|value| Uuid::parse_str(value).ok()) +} + +fn token_count(tokens: u32) -> i32 { + tokens.min(i32::MAX as u32) as i32 +} diff --git a/src/utils/openai_gateway.rs b/src/utils/openai_gateway.rs index 0fb0fba..cc6523e 100644 --- a/src/utils/openai_gateway.rs +++ b/src/utils/openai_gateway.rs @@ -1,12 +1,14 @@ use crate::ai; -use crate::types::{GatewayAuthContext, OpenAiError, OpenAiErrorResponse}; -use axum::body::Body; +use crate::types::{GatewayAuthContext, INFERENCE_READ_SCOPE, INFERENCE_WRITE_SCOPE, JobState}; +use axum::extract::{Request, State}; use axum::http::{HeaderMap, HeaderValue, StatusCode, header}; +use axum::middleware::Next; use axum::response::{IntoResponse, Response}; use omniference::server::SkinAwareJson; -use omniference::skins::{OpenAIChatSkin, OpenAIResponsesSkin, SkinContext}; +use omniference::skins::{OpenAIChatSkin, OpenAIResponsesSkin, SkinContext, SkinRequestMetadata, openai_error_response}; use omniference::types::providers::openai::{OpenAIChatRequest, OpenAIResponsesRequestPayload}; use sqlx::PgPool; +use std::sync::Arc; use uuid::Uuid; pub async fn authenticate(pool: &PgPool, headers: &HeaderMap, scope: &str) -> Result { @@ -54,69 +56,74 @@ pub async fn authenticate(pool: &PgPool, headers: &HeaderMap, scope: &str) -> Re Ok(context) } -pub async fn run_chat(request: OpenAIChatRequest) -> Response { +pub async fn run_chat(request: OpenAIChatRequest, metadata: SkinRequestMetadata) -> Response { let engine = ai::get(); let engine = engine.read().await; let context = SkinContext::with_service(engine.service().clone()); drop(engine); - let response = OpenAIChatSkin::handle_chat(axum07::extract::State(context), SkinAwareJson(request)).await; - bridge_response(response) + OpenAIChatSkin::handle_chat(State(context), Some(axum::Extension(metadata)), SkinAwareJson(request)).await } -pub async fn run_responses(request: OpenAIResponsesRequestPayload) -> Response { +pub async fn run_responses(request: OpenAIResponsesRequestPayload, metadata: SkinRequestMetadata) -> Response { let engine = ai::get(); let engine = engine.read().await; let context = SkinContext::with_service(engine.service().clone()); drop(engine); - let response = OpenAIResponsesSkin::handle_responses(axum07::extract::State(context), SkinAwareJson(request)).await; - bridge_response(response) + OpenAIResponsesSkin::handle_responses(State(context), Some(axum::Extension(metadata)), SkinAwareJson(request)).await } -pub fn parse_json(body: &[u8]) -> Result { - serde_json::from_slice(body).map_err(|error| format!("Failed to parse request body: {error}")) +pub async fn authenticate_read(State(state): State>, request: Request, next: Next) -> Response { + authenticate_request(&state, request, next, INFERENCE_READ_SCOPE).await +} + +pub async fn authenticate_write(State(state): State>, request: Request, next: Next) -> Response { + authenticate_request(&state, request, next, INFERENCE_WRITE_SCOPE).await +} + +async fn authenticate_request(state: &JobState, mut request: Request, next: Next, scope: &str) -> Response { + let context = match authenticate(&state.db, request.headers(), scope).await { + Ok(context) => context, + Err(response) => return response, + }; + request.extensions_mut().insert(context.clone()); + let response = next.run(request).await; + add_gateway_context_headers(response, &context) +} + +pub async fn add_request_id(request: Request, next: Next) -> Response { + let request_id = Uuid::new_v4(); + let response = next.run(request).await; + add_request_id_header(response, request_id) } #[must_use] pub fn error_response(status: StatusCode, message: impl Into, kind: impl Into, code: impl Into) -> Response { - ( - status, - axum::Json(OpenAiErrorResponse { - error: OpenAiError { - message: message.into(), - kind: kind.into(), - param: None, - code: code.into(), - }, - }), - ) - .into_response() + (status, axum::Json(openai_error_response(message, kind, code))).into_response() } #[must_use] -pub fn add_gateway_headers(mut response: Response, context: Option<&GatewayAuthContext>, request_id: Uuid) -> Response { +fn add_request_id_header(mut response: Response, request_id: Uuid) -> Response { if let Ok(value) = HeaderValue::from_str(&request_id.to_string()) { response.headers_mut().insert("x-request-id", value); } - if let Some(context) = context { - if let Ok(value) = HeaderValue::from_str(&context.project_id.to_string()) { - response.headers_mut().insert("x-oxide-project-id", value); - } - if let Ok(value) = HeaderValue::from_str(&context.key_id.to_string()) { - response.headers_mut().insert("x-oxide-api-key-id", value); - } - if let Some(team_id) = context.team_id - && let Ok(value) = HeaderValue::from_str(&team_id.to_string()) - { - response.headers_mut().insert("x-oxide-team-id", value); - } - if let Ok(value) = HeaderValue::from_str(&context.project_name) { - response.headers_mut().insert("x-oxide-project-name", value); - } - } response } -pub(crate) fn bridge_response(response: axum07::response::Response) -> Response { - let (parts, body) = response.into_parts(); - Response::from_parts(parts, Body::new(body)) +#[must_use] +fn add_gateway_context_headers(mut response: Response, context: &GatewayAuthContext) -> Response { + if let Ok(value) = HeaderValue::from_str(&context.project_id.to_string()) { + response.headers_mut().insert("x-oxide-project-id", value); + } + if let Ok(value) = HeaderValue::from_str(&context.key_id.to_string()) { + response.headers_mut().insert("x-oxide-api-key-id", value); + } + if let Some(team_id) = context.team_id + && let Ok(value) = HeaderValue::from_str(&team_id.to_string()) + { + response.headers_mut().insert("x-oxide-team-id", value); + } + if let Ok(value) = HeaderValue::from_str(&context.project_name) { + response.headers_mut().insert("x-oxide-project-name", value); + } + response }