diff --git a/Cargo.lock b/Cargo.lock index e79286a..e889853 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -11,7 +11,7 @@ dependencies = [ "argon2", "async-stream", "async-trait", - "axum 0.8.9", + "axum", "base64", "chrono", "dotenv", @@ -31,7 +31,7 @@ dependencies = [ "tokio", "tokio-util", "tower-cookies", - "tower-http 0.5.2", + "tower-http", "tracing", "tracing-subscriber", "uuid", @@ -221,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", @@ -271,7 +237,7 @@ dependencies = [ "hyper", "hyper-util", "itoa", - "matchit 0.8.4", + "matchit", "memchr", "mime", "percent-encoding", @@ -282,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", @@ -2071,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" @@ -2231,7 +2170,7 @@ dependencies = [ "anyhow", "async-stream", "async-trait", - "axum 0.7.9", + "axum", "base64", "bytes", "dotenvy", @@ -2248,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", @@ -2366,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" @@ -2841,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", @@ -3901,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" @@ -3938,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", @@ -3948,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" @@ -3977,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 062a526..3e8d076 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -10,7 +10,7 @@ keywords = ["chat", "axum", "websocket", "postgres"] categories = ["web-programming"] [dependencies] -axum = { version = "0.8" } +axum = { version = "0.8" } tokio = { version = "1.48", features = ["full"] } serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" @@ -36,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/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/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 dc41621..0bc91a6 100644 --- a/src/routes/mod.rs +++ b/src/routes/mod.rs @@ -6,13 +6,27 @@ 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> { +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) + .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)) .route("/api/v1/base", get(public::base::get_base)) // Admin i18n @@ -159,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/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..2a14c51 --- /dev/null +++ b/src/routes/public/openai.rs @@ -0,0 +1,120 @@ +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 { + OpenAIErrorHandler.handle_not_found() +} + +pub async fn method_not_allowed() -> Response { + OpenAIErrorHandler.handle_method_not_allowed() +} + +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 error_response(StatusCode::INTERNAL_SERVER_ERROR, "Failed to list models", "server_error", "internal_error"); + } + }; + axum::Json(OpenAIModelsResponse { + object: Some("list".to_string()), + data: models, + }) + .into_response() +} + +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, + }; + run_chat(request, request_metadata(&context, model_id)).await +} + +pub async fn responses( + State(state): State>, + Extension(context): Extension, + SkinAwareJson(request): SkinAwareJson, +) -> Response { + let Some(model) = request.model.as_deref() else { + return error_response( + StatusCode::BAD_REQUEST, + "Missing required parameter: 'model'.", + "invalid_request_error", + "missing_required_parameter", + ); + }; + 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 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"); + Err(error_response( + StatusCode::INTERNAL_SERVER_ERROR, + "Failed to resolve model", + "server_error", + "internal_error", + )) + } + } +} + +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/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..67b0f3b --- /dev/null +++ b/src/tests/gateway.rs @@ -0,0 +1,219 @@ +#[cfg(test)] +mod tests { + use crate::types::{GatewayAuthError, GatewayCredential, GatewayModel}; + use crate::utils::auth::hash_password; + 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 uuid::Uuid; + + 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) { + 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 = "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) + 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 (_, 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, "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] + 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); + } + } +} 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..cfdf517 --- /dev/null +++ b/src/types/gateway/mod.rs @@ -0,0 +1,67 @@ +mod repository; + +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, + 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, 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")] + 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..861e2b4 --- /dev/null +++ b/src/types/gateway/repository.rs @@ -0,0 +1,161 @@ +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; + +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 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 = 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 + .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); + } + 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, + 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 GatewayModel { + pub async fn list_for_context(pool: &PgPool, context: &GatewayAuthContext) -> Result, sqlx::Error> { + let models = sqlx::query_as::<_, Self>( + r#" + SELECT + 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 + 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) + ) + ORDER BY p.name, m.model_id + "#, + ) + .bind(context.user_id) + .bind(context.team_id) + .fetch_all(pool) + .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])) + } +} + +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/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..791a7ac 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -2,6 +2,8 @@ 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; pub mod response; 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 new file mode 100644 index 0000000..cc6523e --- /dev/null +++ b/src/utils/openai_gateway.rs @@ -0,0 +1,129 @@ +use crate::ai; +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, 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 { + 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, metadata: SkinRequestMetadata) -> Response { + let engine = ai::get(); + let engine = engine.read().await; + let context = SkinContext::with_service(engine.service().clone()); + drop(engine); + OpenAIChatSkin::handle_chat(State(context), Some(axum::Extension(metadata)), SkinAwareJson(request)).await +} + +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); + OpenAIResponsesSkin::handle_responses(State(context), Some(axum::Extension(metadata)), SkinAwareJson(request)).await +} + +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(openai_error_response(message, kind, code))).into_response() +} + +#[must_use] +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); + } + response +} + +#[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 +}