diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 26d36580cf0..1b060852f8e 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3504,6 +3504,7 @@ dependencies = [ "litellm-config", "litellm-secrets", "rstest", + "serde", "sha2 0.10.9", "subtle", "thiserror 2.0.19", @@ -3521,6 +3522,7 @@ dependencies = [ "futures-util", "litellm-auth", "litellm-core", + "litellm-gateway-auth", "litellm-host-http", "litellm-http", "litellm-llms", diff --git a/litellm-rust/crates/gateway-auth/AGENTS.md b/litellm-rust/crates/gateway-auth/AGENTS.md new file mode 100644 index 00000000000..cbdd7595a55 --- /dev/null +++ b/litellm-rust/crates/gateway-auth/AGENTS.md @@ -0,0 +1,23 @@ +# Gateway authentication + +Inbound authentication uses a verifier, identity resolver, and authorizer injected into `Auth`. `Auth::from_config` composes the current master-key implementation. `Auth::new` accepts alternate implementations and an injected clock + +`CredentialExtractor` selects a token or transport credential from HTTP request parts. The default `Bearer` extractor is strict; a host can supply another extractor with `Auth::with_extractor` for a separate route profile. A transport credential is only a selection marker, never proof of identity + +`Authenticator` receives that selected credential and HTTP request parts. Transport verifiers must obtain proof from trusted server extensions, not client-supplied identity headers. It returns `VerifiedIdentity` only after verification. Request parts provide headers and trusted transport extensions for integrations; the verifier must validate any client-controlled claims before treating them as identity. There is no fallback chain after verification failure + +`IdentityResolver` maps verified identity into an internal principal and permissions. The caller retains the original authentication evidence and credential restrictions independently of that mapping. Principals include an authority, subject, and kind. External subjects must remain authority-scoped unless the resolver explicitly maps them to a shared internal identity + +`AuthenticatedRequest::authorize` checks expiry and intersects resolved permissions with credential restrictions before calling the injected authorizer. The authorizer may impose additional policy but cannot expand either permission set. The returned operation grant binds the exact authenticated caller instance and access request. Consuming it for another caller or target fails + +`Permissions::All`, `None`, and `Only` provide the initial policy vocabulary. `Only` matches exact typed access requests. Model requests contain both the public name and resolved deployment model. MCP requests contain the resolved server ID and upstream operation, including bare tool/prompt names and resource URIs. More expressive policy adapters can be added without changing credential verification + +The Axum `authenticate` middleware currently accepts one Authorization header using the existing Bearer format. Missing, duplicate, empty, or malformed credentials fail. Authentication replaces any preexisting caller extension and checks the method and matched route before dispatch. Handlers extract `AuthenticatedRequest` and authorize their parsed operation before calling a provider. Missing authenticated context fails closed + +Authentication evidence and principals contain no raw token, password, request body, or mutable accounting state. Session ownership is scoped by principal authority, subject, verifier, and credential ID, so separate credentials do not silently share an MCP session. Credential rotation may retain ownership when the verifier preserves a stable credential ID. Scope and expiry checks still run for each operation + +Failures distinguish invalid or expired credentials, forbidden operations, unavailable authentication services, and missing server configuration/context. HTTP adapters map those outcomes to status codes; inference keeps its API-specific error envelopes + +Virtual-key storage, JWT verification, OAuth2 introspection, trusted-proxy validation, SSO, and custom Python hook adapters are not implemented here yet. They should supply these contracts rather than bypassing the shared authorization boundary. Request-body-dependent custom hooks will need a bounded endpoint adapter after parsing + +Budget reservations, rate limits, usage settlement, retries, and cancellation accounting belong after authorization in the execution lifecycle. This foundation adds no accounting backend or cached authorization decisions. Future reservations must bind the same caller and resolved operation and settle idempotently across streaming completion and cancellation diff --git a/litellm-rust/crates/gateway-auth/Cargo.toml b/litellm-rust/crates/gateway-auth/Cargo.toml index 340f7224618..87ba7df4569 100644 --- a/litellm-rust/crates/gateway-auth/Cargo.toml +++ b/litellm-rust/crates/gateway-auth/Cargo.toml @@ -6,13 +6,14 @@ license.workspace = true repository.workspace = true [dependencies] -axum.workspace = true +axum = { workspace = true, features = ["matched-path"] } litellm-auth-types.workspace = true litellm-config.workspace = true litellm-secrets.workspace = true sha2.workspace = true subtle.workspace = true thiserror.workspace = true +serde.workspace = true [dev-dependencies] futures-util.workspace = true diff --git a/litellm-rust/crates/gateway-auth/src/authentication.rs b/litellm-rust/crates/gateway-auth/src/authentication.rs new file mode 100644 index 00000000000..4ecfaa375b4 --- /dev/null +++ b/litellm-rust/crates/gateway-auth/src/authentication.rs @@ -0,0 +1,243 @@ +use std::{future::Future, pin::Pin, sync::Arc, time::SystemTime}; + +use axum::http::{Request, request::Parts}; +use litellm_auth_types::SecretValue; +use litellm_config::Config; +use litellm_secrets::source::SecretSource; +use sha2::{Digest, Sha256}; +use subtle::ConstantTimeEq; + +use crate::{ + AuthenticatedCaller, AuthenticatedRequest, Authentication, AuthenticationMethod, Authorizer, + Bearer, CredentialExtractor, Error, NoAdditionalPolicy, Permissions, Principal, PrincipalKind, + ResolvedIdentity, VerifiedIdentity, +}; + +pub type AuthFuture<'a, T> = Pin> + Send + 'a>>; + +pub trait Clock: Send + Sync { + fn now(&self) -> SystemTime; +} + +pub struct SystemClock; + +impl Clock for SystemClock { + fn now(&self) -> SystemTime { + SystemTime::now() + } +} + +pub enum Credential { + Token(SecretValue), + Transport, +} + +pub trait Authenticator: Send + Sync { + fn validate(&self) -> AuthFuture<'_, ()> { + Box::pin(async { Ok(()) }) + } + + fn verify<'a>( + &'a self, + credential: &'a Credential, + request: &'a Parts, + ) -> AuthFuture<'a, VerifiedIdentity>; +} + +pub trait IdentityResolver: Send + Sync { + fn resolve<'a>(&'a self, identity: &'a VerifiedIdentity) -> AuthFuture<'a, ResolvedIdentity>; +} + +pub struct LocalAdministrator; + +impl IdentityResolver for LocalAdministrator { + fn resolve<'a>(&'a self, identity: &'a VerifiedIdentity) -> AuthFuture<'a, ResolvedIdentity> { + Box::pin(async move { + let allowed = match identity.authentication.method { + AuthenticationMethod::MasterKey => identity.principal == master_principal(), + AuthenticationMethod::Session => { + identity.principal.authority() == "litellm:local-ui" + && identity.principal.kind() == PrincipalKind::Human + } + AuthenticationMethod::External(_) => false, + }; + if !allowed { + return Err(Error::Forbidden); + } + Ok(ResolvedIdentity { + principal: identity.principal.clone(), + permissions: Permissions::All, + }) + }) + } +} + +#[derive(Clone)] +pub struct Auth { + extractor: Arc, + authenticator: Arc, + identities: Arc, + authorizer: Arc, + clock: Arc, +} + +impl Auth { + pub fn new( + authenticator: Arc, + identities: Arc, + authorizer: Arc, + clock: Arc, + ) -> Self { + Self { + extractor: Arc::new(Bearer), + authenticator, + identities, + authorizer, + clock, + } + } + + pub fn from_config(config: &Config, secrets: Arc) -> Self { + let master = Arc::new(MasterKeyAuthenticator { + key: config.general_settings.master_key.clone(), + secrets, + }); + Self { + extractor: Arc::new(Bearer), + authenticator: master, + identities: Arc::new(LocalAdministrator), + authorizer: Arc::new(NoAdditionalPolicy), + clock: Arc::new(SystemClock), + } + } + + pub fn with_extractor(self, extractor: Arc) -> Self { + Self { extractor, ..self } + } + + pub async fn authenticate_parts(&self, request: &Parts) -> Result { + let credential = self.extractor.extract(request)?; + self.authenticate_request(&credential, request).await + } + + pub async fn validate(&self) -> Result<(), Error> { + self.authenticator.validate().await + } + + pub async fn authenticate( + &self, + credential: &SecretValue, + ) -> Result { + self.authenticate_request( + &Credential::Token(credential.clone()), + &Request::new(()).into_parts().0, + ) + .await + } + + pub async fn authenticate_request( + &self, + credential: &Credential, + request: &Parts, + ) -> Result { + let verified = self.authenticator.verify(credential, request).await?; + resolve( + verified, + self.identities.as_ref(), + self.authorizer.clone(), + self.clock.clone(), + ) + .await + } +} + +pub(crate) async fn resolve( + verified: VerifiedIdentity, + identities: &dyn IdentityResolver, + authorizer: Arc, + clock: Arc, +) -> Result { + if verified + .authentication + .expires_at + .is_some_and(|expiry| expiry <= clock.now()) + { + return Err(Error::Expired); + } + let resolved = identities.resolve(&verified).await?; + Ok(AuthenticatedRequest::new( + Arc::new(AuthenticatedCaller::new(verified, resolved)), + authorizer, + clock, + )) +} + +pub struct MasterKeyAuthenticator { + key: Option, + secrets: Arc, +} + +impl MasterKeyAuthenticator { + pub fn new(key: Option, secrets: Arc) -> Self { + Self { key, secrets } + } + + async fn key(&self) -> Result { + let configured = self.key.as_ref().ok_or(Error::Unconfigured)?; + let resolved = match configured.expose().strip_prefix("os.environ/") { + Some(name) if !name.is_empty() => self + .secrets + .get_secret_str(name) + .await? + .ok_or(Error::Unconfigured)?, + Some(_) => return Err(Error::Unconfigured), + None => configured.clone(), + }; + if resolved.expose().trim().is_empty() { + return Err(Error::Unconfigured); + } + Ok(resolved) + } +} + +impl Authenticator for MasterKeyAuthenticator { + fn validate(&self) -> AuthFuture<'_, ()> { + Box::pin(async { self.key().await.map(|_| ()) }) + } + + fn verify<'a>( + &'a self, + credential: &'a Credential, + _: &'a Parts, + ) -> AuthFuture<'a, VerifiedIdentity> { + Box::pin(async move { + let expected = self.key().await?; + let Credential::Token(credential) = credential else { + return Err(Error::InvalidToken); + }; + if !bool::from( + Sha256::digest(credential.expose()).ct_eq(&Sha256::digest(expected.expose())), + ) { + return Err(Error::InvalidToken); + } + Ok(VerifiedIdentity { + principal: master_principal(), + authentication: Authentication { + method: AuthenticationMethod::MasterKey, + verifier: "litellm:master-key".into(), + credential_id: "master".into(), + expires_at: None, + }, + restrictions: Permissions::All, + }) + }) + } +} + +fn master_principal() -> Principal { + Principal::new( + "litellm:master-key".into(), + "master".into(), + PrincipalKind::System, + ) +} diff --git a/litellm-rust/crates/gateway-auth/src/authorization.rs b/litellm-rust/crates/gateway-auth/src/authorization.rs new file mode 100644 index 00000000000..522fae64f3e --- /dev/null +++ b/litellm-rust/crates/gateway-auth/src/authorization.rs @@ -0,0 +1,139 @@ +use std::sync::Arc; + +use crate::{AuthFuture, AuthenticatedCaller, Clock, Error, SharedCaller}; + +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum McpAction { + Connect, + ListTools, + CallTool(String), + ListPrompts, + GetPrompt(String), + ListResources, + ListResourceTemplates, + ReadResource(String), +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum UiAction { + SessionInfo, + Logout, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum AccessRequest { + Route { method: String, path: String }, + Model { name: String, deployment: String }, + Mcp { server: String, action: McpAction }, + Ui(UiAction), +} + +#[derive(Clone, Debug, Default)] +pub enum Permissions { + All, + #[default] + None, + Only(Arc<[AccessRequest]>), +} + +impl Permissions { + pub fn allows(&self, request: &AccessRequest) -> bool { + match self { + Self::All => true, + Self::None => false, + Self::Only(allowed) => allowed.contains(request), + } + } +} + +pub trait Authorizer: Send + Sync { + fn authorize<'a>( + &'a self, + caller: &'a AuthenticatedCaller, + request: &'a AccessRequest, + ) -> AuthFuture<'a, ()>; +} + +pub struct NoAdditionalPolicy; + +impl Authorizer for NoAdditionalPolicy { + fn authorize<'a>( + &'a self, + _: &'a AuthenticatedCaller, + _: &'a AccessRequest, + ) -> AuthFuture<'a, ()> { + Box::pin(async { Ok(()) }) + } +} + +#[derive(Clone)] +pub struct AuthenticatedRequest { + caller: SharedCaller, + authorizer: Arc, + clock: Arc, +} + +impl AuthenticatedRequest { + pub(crate) fn new( + caller: SharedCaller, + authorizer: Arc, + clock: Arc, + ) -> Self { + Self { + caller, + authorizer, + clock, + } + } + + pub fn caller(&self) -> &AuthenticatedCaller { + &self.caller + } + + pub async fn authorize(&self, request: AccessRequest) -> Result { + if self + .caller + .authentication() + .expires_at + .is_some_and(|expiry| expiry <= self.clock.now()) + { + return Err(Error::Expired); + } + if !self.caller.restrictions().allows(&request) + || !self.caller.permissions().allows(&request) + { + return Err(Error::Forbidden); + } + self.authorizer.authorize(&self.caller, &request).await?; + Ok(AuthorizedOperation { + caller: self.caller.clone(), + request, + }) + } +} + +pub struct AuthorizedOperation { + caller: SharedCaller, + request: AccessRequest, +} + +impl AuthorizedOperation { + pub fn caller(&self) -> &AuthenticatedCaller { + &self.caller + } + pub fn request(&self) -> &AccessRequest { + &self.request + } + + pub fn consume( + self, + caller: &AuthenticatedCaller, + request: &AccessRequest, + ) -> Result<(), Error> { + if std::ptr::eq(self.caller.as_ref(), caller) && &self.request == request { + Ok(()) + } else { + Err(Error::Forbidden) + } + } +} diff --git a/litellm-rust/crates/gateway-auth/src/error.rs b/litellm-rust/crates/gateway-auth/src/error.rs index 735eb7741ad..33f02f63796 100644 --- a/litellm-rust/crates/gateway-auth/src/error.rs +++ b/litellm-rust/crates/gateway-auth/src/error.rs @@ -5,6 +5,14 @@ use axum::{ #[derive(Debug, thiserror::Error)] pub enum Error { + #[error("operation is not permitted")] + Forbidden, + #[error("credential has expired")] + Expired, + #[error("authenticated request context is missing")] + MissingIdentity, + #[error("authentication service unavailable")] + Unavailable, #[error("gateway auth not configured")] Unconfigured, #[error("missing or invalid bearer token")] @@ -13,12 +21,21 @@ pub enum Error { Secret(#[from] litellm_secrets::Error), } -impl IntoResponse for Error { - fn into_response(self) -> Response { - let status = match &self { - Self::InvalidToken => StatusCode::UNAUTHORIZED, - Self::Unconfigured | Self::Secret(_) => StatusCode::INTERNAL_SERVER_ERROR, - }; - (status, self.to_string()).into_response() +impl Error { + pub fn status(&self) -> StatusCode { + match self { + Self::InvalidToken | Self::Expired => StatusCode::UNAUTHORIZED, + Self::Forbidden => StatusCode::FORBIDDEN, + Self::Unavailable => StatusCode::SERVICE_UNAVAILABLE, + Self::MissingIdentity | Self::Unconfigured | Self::Secret(_) => { + StatusCode::INTERNAL_SERVER_ERROR + } + } + } +} + +impl IntoResponse for Error { + fn into_response(self) -> Response { + (self.status(), self.to_string()).into_response() } } diff --git a/litellm-rust/crates/gateway-auth/src/http.rs b/litellm-rust/crates/gateway-auth/src/http.rs new file mode 100644 index 00000000000..ddbf240509a --- /dev/null +++ b/litellm-rust/crates/gateway-auth/src/http.rs @@ -0,0 +1,90 @@ +use axum::{ + extract::{FromRequestParts, MatchedPath, Request, State}, + http::{header::AUTHORIZATION, request::Parts}, + middleware::Next, + response::Response, +}; +use litellm_auth_types::SecretValue; + +use crate::{AccessRequest, Auth, AuthenticatedRequest, Credential, Error}; + +pub trait CredentialExtractor: Send + Sync { + fn extract(&self, parts: &Parts) -> Result; +} + +pub struct Bearer; + +impl CredentialExtractor for Bearer { + fn extract(&self, parts: &Parts) -> Result { + bearer(parts).map(Credential::Token) + } +} + +fn bearer(parts: &Parts) -> Result { + let mut values = parts.headers.get_all(AUTHORIZATION).iter(); + let value = values.next().ok_or(Error::InvalidToken)?; + if values.next().is_some() { + return Err(Error::InvalidToken); + } + let token = value + .to_str() + .ok() + .and_then(|value| value.strip_prefix("Bearer ")) + .map(str::trim) + .filter(|value| !value.is_empty() && !value.bytes().any(|byte| byte.is_ascii_whitespace())) + .ok_or(Error::InvalidToken)?; + Ok(SecretValue::new(token)) +} + +pub async fn authenticate( + State(auth): State, + request: Request, + next: Next, +) -> Result { + let (mut parts, body) = request.into_parts(); + let identity = auth.authenticate_parts(&parts).await?; + let route = AccessRequest::Route { + method: parts.method.to_string(), + path: parts + .extensions + .get::() + .map_or(parts.uri.path(), MatchedPath::as_str) + .into(), + }; + identity + .authorize(route.clone()) + .await? + .consume(identity.caller(), &route)?; + parts.extensions.insert(identity); + Ok(next.run(Request::from_parts(parts, body)).await) +} + +impl FromRequestParts for AuthenticatedRequest { + type Rejection = Error; + + async fn from_request_parts(parts: &mut Parts, _: &S) -> Result { + parts + .extensions + .get::() + .cloned() + .ok_or(Error::MissingIdentity) + } +} + +pub struct RequireMasterKey; + +impl FromRequestParts for RequireMasterKey { + type Rejection = Error; + + async fn from_request_parts(parts: &mut Parts, state: &Auth) -> Result { + state.validate().await?; + let identity = state + .authenticate_request(&Credential::Token(bearer(parts)?), parts) + .await?; + if identity.caller().authentication().method != crate::AuthenticationMethod::MasterKey { + return Err(Error::InvalidToken); + } + parts.extensions.insert(identity); + Ok(Self) + } +} diff --git a/litellm-rust/crates/gateway-auth/src/identity.rs b/litellm-rust/crates/gateway-auth/src/identity.rs new file mode 100644 index 00000000000..286983ae647 --- /dev/null +++ b/litellm-rust/crates/gateway-auth/src/identity.rs @@ -0,0 +1,122 @@ +use std::{sync::Arc, time::SystemTime}; + +use sha2::{Digest, Sha256}; + +use crate::authorization::Permissions; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum PrincipalKind { + Human, + Service, + System, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct Principal { + authority: String, + subject: String, + kind: PrincipalKind, +} + +impl Principal { + pub fn new(authority: String, subject: String, kind: PrincipalKind) -> Self { + Self { + authority, + subject, + kind, + } + } + + pub fn authority(&self) -> &str { + &self.authority + } + pub fn subject(&self) -> &str { + &self.subject + } + pub fn kind(&self) -> PrincipalKind { + self.kind + } +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum AuthenticationMethod { + MasterKey, + Session, + External(String), +} + +#[derive(Clone, Debug)] +pub struct Authentication { + pub method: AuthenticationMethod, + pub verifier: String, + pub credential_id: String, + pub expires_at: Option, +} + +pub struct VerifiedIdentity { + pub principal: Principal, + pub authentication: Authentication, + pub restrictions: Permissions, +} + +pub struct ResolvedIdentity { + pub principal: Principal, + pub permissions: Permissions, +} + +#[derive(Clone, Debug)] +pub struct AuthenticatedCaller { + verified_principal: Principal, + principal: Principal, + authentication: Authentication, + permissions: Permissions, + restrictions: Permissions, +} + +impl AuthenticatedCaller { + pub(crate) fn new(verified: VerifiedIdentity, resolved: ResolvedIdentity) -> Self { + Self { + verified_principal: verified.principal, + principal: resolved.principal, + authentication: verified.authentication, + permissions: resolved.permissions, + restrictions: verified.restrictions, + } + } + + pub fn verified_principal(&self) -> &Principal { + &self.verified_principal + } + + pub fn principal(&self) -> &Principal { + &self.principal + } + pub fn authentication(&self) -> &Authentication { + &self.authentication + } + pub fn permissions(&self) -> &Permissions { + &self.permissions + } + pub fn restrictions(&self) -> &Permissions { + &self.restrictions + } + + pub fn session_owner(&self) -> String { + let fields = [ + self.verified_principal.authority(), + self.verified_principal.subject(), + self.principal.authority(), + self.principal.subject(), + self.authentication.verifier.as_str(), + self.authentication.credential_id.as_str(), + ]; + let digest = fields.into_iter().fold(Sha256::new(), |mut digest, field| { + digest.update((field.len() as u64).to_be_bytes()); + digest.update(field.as_bytes()); + digest + }); + format!("{:x}", digest.finalize()) + } +} + +pub type SharedCaller = Arc; diff --git a/litellm-rust/crates/gateway-auth/src/lib.rs b/litellm-rust/crates/gateway-auth/src/lib.rs index 7a1fb244654..971a43a3a5c 100644 --- a/litellm-rust/crates/gateway-auth/src/lib.rs +++ b/litellm-rust/crates/gateway-auth/src/lib.rs @@ -1,74 +1,26 @@ +mod authentication; +mod authorization; mod error; +mod http; +mod identity; -use std::sync::Arc; - -use axum::{ - extract::FromRequestParts, - http::{header::AUTHORIZATION, request::Parts}, -}; -use litellm_auth_types::SecretValue; -use litellm_config::Config; -use litellm_secrets::source::SecretSource; use sha2::{Digest, Sha256}; -use subtle::ConstantTimeEq; +pub use authentication::{ + Auth, AuthFuture, Authenticator, Clock, Credential, IdentityResolver, LocalAdministrator, + MasterKeyAuthenticator, SystemClock, +}; +pub use authorization::{ + AccessRequest, AuthenticatedRequest, AuthorizedOperation, Authorizer, McpAction, + NoAdditionalPolicy, Permissions, UiAction, +}; pub use error::Error; - -#[derive(Clone)] -pub struct Auth { - master_key: Option, - secrets: Arc, -} - -impl Auth { - pub fn from_config(config: &Config, secrets: Arc) -> Self { - Self { - master_key: config.general_settings.master_key.clone(), - secrets, - } - } - - async fn master_key(&self) -> Result { - let configured = self.master_key.as_ref().ok_or(Error::Unconfigured)?; - let resolved = match configured.expose().strip_prefix("os.environ/") { - Some(name) if !name.is_empty() => self - .secrets - .get_secret_str(name) - .await? - .ok_or(Error::Unconfigured)?, - Some(_) => return Err(Error::Unconfigured), - None => configured.clone(), - }; - if resolved.expose().trim().is_empty() { - return Err(Error::Unconfigured); - } - Ok(resolved) - } -} +pub use http::{Bearer, CredentialExtractor, RequireMasterKey, authenticate}; +pub use identity::{ + AuthenticatedCaller, Authentication, AuthenticationMethod, Principal, PrincipalKind, + ResolvedIdentity, SharedCaller, VerifiedIdentity, +}; pub fn hash_token(token: &str) -> String { format!("{:x}", Sha256::digest(token.as_bytes())) } - -pub struct RequireMasterKey; - -impl FromRequestParts for RequireMasterKey { - type Rejection = Error; - - async fn from_request_parts(parts: &mut Parts, state: &Auth) -> Result { - let expected = state.master_key().await?; - let provided = parts - .headers - .get(AUTHORIZATION) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.strip_prefix("Bearer ")) - .map(str::trim) - .ok_or(Error::InvalidToken)?; - let actual_hash = Sha256::digest(provided.as_bytes()); - let expected_hash = Sha256::digest(expected.expose().as_bytes()); - match bool::from(actual_hash.ct_eq(&expected_hash)) { - true => Ok(Self), - false => Err(Error::InvalidToken), - } - } -} diff --git a/litellm-rust/crates/gateway-auth/tests/auth.rs b/litellm-rust/crates/gateway-auth/tests/auth.rs index 58625296664..a12230e8d5c 100644 --- a/litellm-rust/crates/gateway-auth/tests/auth.rs +++ b/litellm-rust/crates/gateway-auth/tests/auth.rs @@ -98,3 +98,458 @@ fn hash_token_matches_python_sha256_hexdigest() { "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b" ); } + +use litellm_gateway_auth::{ + AccessRequest, AuthFuture, AuthenticatedCaller, AuthenticatedRequest, Authentication, + AuthenticationMethod, Authenticator, Authorizer, Clock, Error, IdentityResolver, Permissions, + Principal, PrincipalKind, ResolvedIdentity, VerifiedIdentity, authenticate, +}; +use std::{ + sync::atomic::{AtomicU64, AtomicUsize, Ordering}, + time::{Duration, SystemTime}, +}; + +struct TestClock(AtomicU64); + +impl Clock for TestClock { + fn now(&self) -> SystemTime { + SystemTime::UNIX_EPOCH + Duration::from_secs(self.0.load(Ordering::SeqCst)) + } +} + +struct Verifier { + authority: &'static str, + restrictions: Permissions, + failure: Option, +} + +impl Authenticator for Verifier { + fn verify<'a>( + &'a self, + credential: &'a litellm_gateway_auth::Credential, + _: &'a axum::http::request::Parts, + ) -> AuthFuture<'a, VerifiedIdentity> { + Box::pin(async move { + match self.failure { + Some(401) => return Err(Error::InvalidToken), + Some(_) => return Err(Error::Unavailable), + None => (), + } + let litellm_gateway_auth::Credential::Token(credential) = credential else { + return Err(Error::InvalidToken); + }; + if credential.expose() != "test-token" { + return Err(Error::InvalidToken); + } + Ok(VerifiedIdentity { + principal: Principal::new( + self.authority.into(), + "subject".into(), + PrincipalKind::Service, + ), + authentication: Authentication { + method: AuthenticationMethod::External("test".into()), + verifier: "test-verifier".into(), + credential_id: "credential".into(), + expires_at: Some(SystemTime::UNIX_EPOCH + Duration::from_secs(100)), + }, + restrictions: self.restrictions.clone(), + }) + }) + } +} + +struct Identities { + mapped_principal: Option, + calls: AtomicUsize, + permissions: Permissions, +} + +impl IdentityResolver for Identities { + fn resolve<'a>(&'a self, identity: &'a VerifiedIdentity) -> AuthFuture<'a, ResolvedIdentity> { + Box::pin(async move { + self.calls.fetch_add(1, Ordering::SeqCst); + Ok(ResolvedIdentity { + principal: self + .mapped_principal + .clone() + .unwrap_or_else(|| identity.principal.clone()), + permissions: self.permissions.clone(), + }) + }) + } +} + +struct Policy(AtomicUsize); + +impl Authorizer for Policy { + fn authorize<'a>( + &'a self, + _: &'a AuthenticatedCaller, + _: &'a AccessRequest, + ) -> AuthFuture<'a, ()> { + Box::pin(async move { + self.0.fetch_add(1, Ordering::SeqCst); + Ok(()) + }) + } +} + +struct Pipeline { + auth: Auth, + identities: Arc, + policy: Arc, + clock: Arc, +} + +fn pipeline( + authority: &'static str, + permissions: Permissions, + restrictions: Permissions, + failure: Option, +) -> Pipeline { + let identities = Arc::new(Identities { + mapped_principal: None, + calls: AtomicUsize::new(0), + permissions, + }); + let policy = Arc::new(Policy(AtomicUsize::new(0))); + let clock = Arc::new(TestClock(AtomicU64::new(10))); + let auth = Auth::new( + Arc::new(Verifier { + authority, + restrictions, + failure, + }), + identities.clone(), + policy.clone(), + clock.clone(), + ); + Pipeline { + auth, + identities, + policy, + clock, + } +} + +fn model(name: &str) -> AccessRequest { + AccessRequest::Model { + name: name.into(), + deployment: name.into(), + } +} + +#[rstest] +#[case::invalid_token(401)] +#[case::unavailable_verifier(503)] +#[tokio::test] +async fn failed_verification_never_resolves_identity_or_runs_the_handler(#[case] status: u16) { + let pipeline = pipeline("issuer-a", Permissions::All, Permissions::All, Some(status)); + let app = Router::new() + .route("/protected", get(|| async { StatusCode::NO_CONTENT })) + .layer(axum::middleware::from_fn_with_state( + pipeline.auth, + authenticate, + )); + let response = app + .oneshot( + Request::get("/protected") + .header("authorization", "Bearer test-token") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status().as_u16(), status); + assert_eq!(pipeline.identities.calls.load(Ordering::SeqCst), 0); + assert_eq!(pipeline.policy.0.load(Ordering::SeqCst), 0); +} + +#[rstest] +#[case::membership_restricts(Permissions::Only(Arc::from([model("allowed")])), Permissions::All)] +#[case::credential_restricts(Permissions::All, Permissions::Only(Arc::from([model("allowed")])))] +#[tokio::test] +async fn policy_cannot_expand_membership_or_credential_permissions( + #[case] permissions: Permissions, + #[case] restrictions: Permissions, +) { + let pipeline = pipeline("issuer-a", permissions, restrictions, None); + let identity = pipeline + .auth + .authenticate(&SecretValue::new("test-token")) + .await + .unwrap(); + assert!(matches!( + identity.authorize(model("denied")).await, + Err(Error::Forbidden) + )); + assert_eq!(pipeline.policy.0.load(Ordering::SeqCst), 0); + let allowed = model("allowed"); + identity + .authorize(allowed.clone()) + .await + .unwrap() + .consume(identity.caller(), &allowed) + .unwrap(); + assert_eq!(pipeline.policy.0.load(Ordering::SeqCst), 1); +} + +#[rstest] +#[tokio::test] +async fn authorization_grants_bind_the_exact_caller_and_target() { + let first = pipeline("issuer-a", Permissions::All, Permissions::All, None); + let second = pipeline("issuer-b", Permissions::All, Permissions::All, None); + let caller = first + .auth + .authenticate(&SecretValue::new("test-token")) + .await + .unwrap(); + let other = second + .auth + .authenticate(&SecretValue::new("test-token")) + .await + .unwrap(); + assert_ne!(caller.caller().principal(), other.caller().principal()); + assert_ne!( + caller.caller().session_owner(), + other.caller().session_owner() + ); + let access = model("allowed"); + assert!(matches!( + caller + .authorize(access.clone()) + .await + .unwrap() + .consume(other.caller(), &access), + Err(Error::Forbidden) + )); + assert!(matches!( + caller + .authorize(access) + .await + .unwrap() + .consume(caller.caller(), &model("different")), + Err(Error::Forbidden) + )); +} + +#[rstest] +#[tokio::test] +async fn expiry_is_checked_again_for_operations_on_existing_connections() { + let pipeline = pipeline("issuer-a", Permissions::All, Permissions::All, None); + let identity = pipeline + .auth + .authenticate(&SecretValue::new("test-token")) + .await + .unwrap(); + assert!(identity.authorize(model("allowed")).await.is_ok()); + pipeline.clock.0.store(100, Ordering::SeqCst); + assert!(matches!( + identity.authorize(model("allowed")).await, + Err(Error::Expired) + )); + assert!(matches!( + pipeline + .auth + .authenticate(&SecretValue::new("test-token")) + .await, + Err(Error::Expired) + )); + assert_eq!(pipeline.identities.calls.load(Ordering::SeqCst), 1); + assert_eq!(pipeline.policy.0.load(Ordering::SeqCst), 1); +} + +#[rstest] +#[tokio::test] +async fn bearer_extraction_rejects_duplicates_and_does_not_trust_preexisting_identity() { + let pipeline = pipeline("issuer-a", Permissions::All, Permissions::All, None); + let identity = pipeline + .auth + .authenticate(&SecretValue::new("test-token")) + .await + .unwrap(); + let app = Router::new() + .route( + "/protected", + get(|_: AuthenticatedRequest| async { StatusCode::NO_CONTENT }), + ) + .layer(axum::middleware::from_fn_with_state( + pipeline.auth, + authenticate, + )); + let duplicate = Request::get("/protected") + .header("authorization", "Bearer test-token") + .header("authorization", "Bearer test-token") + .body(Body::empty()) + .unwrap(); + assert_eq!(app.clone().oneshot(duplicate).await.unwrap().status(), 401); + let forged = Request::get("/protected") + .extension(identity) + .body(Body::empty()) + .unwrap(); + assert_eq!(app.clone().oneshot(forged).await.unwrap().status(), 401); + let valid = Request::get("/protected") + .header("authorization", "Bearer test-token") + .body(Body::empty()) + .unwrap(); + assert_eq!(app.oneshot(valid).await.unwrap().status(), 204); +} + +#[rstest] +#[tokio::test] +async fn route_permissions_are_checked_before_parsing_or_dispatch() { + let allowed = AccessRequest::Route { + method: "GET".into(), + path: "/allowed/{id}".into(), + }; + let pipeline = pipeline( + "issuer-a", + Permissions::Only(Arc::from([allowed])), + Permissions::All, + None, + ); + let app = Router::new() + .route("/allowed/{id}", get(|| async { StatusCode::NO_CONTENT })) + .route("/denied", get(|| async { StatusCode::NO_CONTENT })) + .layer(axum::middleware::from_fn_with_state( + pipeline.auth, + authenticate, + )); + let allowed = Request::get("/allowed/123") + .header("authorization", "Bearer test-token") + .body(Body::empty()) + .unwrap(); + let denied = Request::get("/denied") + .header("authorization", "Bearer test-token") + .body(Body::empty()) + .unwrap(); + assert_eq!(app.clone().oneshot(allowed).await.unwrap().status(), 204); + assert_eq!(app.oneshot(denied).await.unwrap().status(), 403); +} + +struct CustomHeader; + +impl litellm_gateway_auth::CredentialExtractor for CustomHeader { + fn extract( + &self, + parts: &axum::http::request::Parts, + ) -> Result { + let value = parts + .headers + .get("x-test-key") + .and_then(|value| value.to_str().ok()) + .ok_or(Error::InvalidToken)?; + Ok(litellm_gateway_auth::Credential::Token(SecretValue::new( + value, + ))) + } +} + +#[rstest] +#[case::selected_header("test-token", "wrong-token", 204)] +#[case::no_fallback("wrong-token", "test-token", 401)] +#[tokio::test] +async fn route_profile_uses_only_its_selected_credentials( + #[case] custom: &str, + #[case] bearer: &str, + #[case] status: u16, +) { + let pipeline = pipeline("issuer-a", Permissions::All, Permissions::All, None); + let app = Router::new() + .route("/protected", get(|| async { StatusCode::NO_CONTENT })) + .layer(axum::middleware::from_fn_with_state( + pipeline.auth.with_extractor(Arc::new(CustomHeader)), + authenticate, + )); + let response = app + .oneshot( + Request::get("/protected") + .header("x-test-key", custom) + .header("authorization", format!("Bearer {bearer}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status().as_u16(), status); +} + +#[rstest] +#[tokio::test] +async fn master_rotation_preserves_session_owner_but_rejects_the_old_key( + secrets: Arc, +) { + let first = + Config::from_yaml("model_list: []\ngeneral_settings:\n master_key: first-key\n").unwrap(); + let second = + Config::from_yaml("model_list: []\ngeneral_settings:\n master_key: second-key\n").unwrap(); + let original = Auth::from_config(&first, secrets.clone()) + .authenticate(&SecretValue::new("first-key")) + .await + .unwrap(); + let rotated = Auth::from_config(&second, secrets); + let replacement = rotated + .authenticate(&SecretValue::new("second-key")) + .await + .unwrap(); + assert_eq!( + original.caller().session_owner(), + replacement.caller().session_owner() + ); + assert!(matches!( + rotated.authenticate(&SecretValue::new("first-key")).await, + Err(Error::InvalidToken) + )); +} + +#[rstest] +#[tokio::test] +async fn mapping_two_issuers_to_one_account_does_not_merge_session_ownership() { + let resolver = Arc::new(Identities { + mapped_principal: Some(Principal::new( + "internal".into(), + "shared-account".into(), + PrincipalKind::Human, + )), + calls: AtomicUsize::new(0), + permissions: Permissions::All, + }); + let first = Auth::new( + Arc::new(Verifier { + authority: "issuer-a", + restrictions: Permissions::All, + failure: None, + }), + resolver.clone(), + Arc::new(Policy(AtomicUsize::new(0))), + Arc::new(TestClock(AtomicU64::new(10))), + ); + let second = Auth::new( + Arc::new(Verifier { + authority: "issuer-b", + restrictions: Permissions::All, + failure: None, + }), + resolver, + Arc::new(Policy(AtomicUsize::new(0))), + Arc::new(TestClock(AtomicU64::new(10))), + ); + let original = first + .authenticate(&SecretValue::new("test-token")) + .await + .unwrap(); + let mapped = second + .authenticate(&SecretValue::new("test-token")) + .await + .unwrap(); + assert_eq!(original.caller().principal(), mapped.caller().principal()); + assert_ne!( + original.caller().verified_principal(), + mapped.caller().verified_principal() + ); + assert_ne!( + original.caller().session_owner(), + mapped.caller().session_owner() + ); +} diff --git a/litellm-rust/crates/gateway-inference/Cargo.toml b/litellm-rust/crates/gateway-inference/Cargo.toml index 8217c273dbf..fd7c99204f7 100644 --- a/litellm-rust/crates/gateway-inference/Cargo.toml +++ b/litellm-rust/crates/gateway-inference/Cargo.toml @@ -10,6 +10,7 @@ axum = { workspace = true, features = ["json", "multipart", "original-uri"] } base64.workspace = true bytes.workspace = true litellm-auth.workspace = true +litellm-gateway-auth.workspace = true litellm-core.workspace = true litellm-host-http.workspace = true litellm-http.workspace = true diff --git a/litellm-rust/crates/gateway-inference/src/audio_transcription.rs b/litellm-rust/crates/gateway-inference/src/audio_transcription.rs index 94b4d6ce91f..dc762fb8643 100644 --- a/litellm-rust/crates/gateway-inference/src/audio_transcription.rs +++ b/litellm-rust/crates/gateway-inference/src/audio_transcription.rs @@ -1,3 +1,4 @@ +use litellm_gateway_auth::AuthenticatedRequest; use std::{path::Path, sync::Arc}; use axum::{Json, extract::State, response::IntoResponse}; @@ -12,19 +13,22 @@ use crate::{ pub(crate) async fn create( State(gateway): State>, + identity: AuthenticatedRequest, body: InferenceBody, ) -> Result { - handle(&gateway, body).await.map(Json) + handle(&gateway, &identity, body).await.map(Json) } async fn handle( gateway: &Gateway, + identity: &AuthenticatedRequest, InferenceBody { fields: body, upload, }: InferenceBody, ) -> Result { let deployment = request::resolve_deployment(gateway, &body)?; + request::authorize_model(identity, deployment, &body).await?; let audio = match upload { Some(upload) => { let format = upload diff --git a/litellm-rust/crates/gateway-inference/src/chat_completions.rs b/litellm-rust/crates/gateway-inference/src/chat_completions.rs index 16022e2f103..069bbece5fe 100644 --- a/litellm-rust/crates/gateway-inference/src/chat_completions.rs +++ b/litellm-rust/crates/gateway-inference/src/chat_completions.rs @@ -1,3 +1,4 @@ +use litellm_gateway_auth::AuthenticatedRequest; use std::sync::Arc; use axum::{ @@ -12,13 +13,15 @@ use crate::{Error, Gateway, JsonObject, request}; pub(crate) async fn create( State(gateway): State>, + identity: AuthenticatedRequest, JsonObject(body): JsonObject, ) -> Result { - handle(&gateway, body).await + handle(&gateway, &identity, body).await } pub(crate) async fn create_from_model_path( State(gateway): State>, + identity: AuthenticatedRequest, Path(model): Path, JsonObject(body): JsonObject, ) -> Result { @@ -29,11 +32,16 @@ pub(crate) async fn create_from_model_path( .collect(), Some(_) => body, }; - handle(&gateway, body).await + handle(&gateway, &identity, body).await } -async fn handle(gateway: &Gateway, body: Map) -> Result { +async fn handle( + gateway: &Gateway, + identity: &AuthenticatedRequest, + body: Map, +) -> Result { let deployment = request::resolve_deployment(gateway, &body)?; + request::authorize_model(identity, deployment, &body).await?; let messages = body.get("messages").cloned().unwrap_or_default(); let response = litellm_host_http::serve_unary( gateway diff --git a/litellm-rust/crates/gateway-inference/src/error.rs b/litellm-rust/crates/gateway-inference/src/error.rs index 96edc3a94c5..6b99069dc9d 100644 --- a/litellm-rust/crates/gateway-inference/src/error.rs +++ b/litellm-rust/crates/gateway-inference/src/error.rs @@ -10,6 +10,8 @@ use serde_json::{Map, Value, json}; #[derive(Debug, thiserror::Error)] pub enum Error { + #[error(transparent)] + Auth(#[from] litellm_gateway_auth::Error), #[error("invalid request body: {0}")] InvalidBody(String), #[error( @@ -46,6 +48,7 @@ impl From> for Error { impl Error { pub fn status(&self) -> StatusCode { match self { + Self::Auth(error) => error.status(), Self::Unsupported(_) | Self::Route(RouteError::Unsupported(_)) | Self::Ocr(OcrError::Unsupported(_)) => StatusCode::NOT_IMPLEMENTED, diff --git a/litellm-rust/crates/gateway-inference/src/messages.rs b/litellm-rust/crates/gateway-inference/src/messages.rs index 0338640e741..87572156502 100644 --- a/litellm-rust/crates/gateway-inference/src/messages.rs +++ b/litellm-rust/crates/gateway-inference/src/messages.rs @@ -1,5 +1,6 @@ //! `POST /v1/messages`, as the Python proxy's `anthropic_response` serves it. +use litellm_gateway_auth::AuthenticatedRequest; use std::sync::Arc; use axum::{ @@ -22,12 +23,13 @@ const ANTHROPIC_API_HEADER_PROVIDERS: &str = "anthropic,bedrock,bedrock_mantle,v pub async fn create( State(gateway): State>, + identity: AuthenticatedRequest, RequestId(request_id): RequestId, headers: HeaderMap, body: Result, ) -> impl IntoResponse { let result = match body { - Ok(JsonObject(body)) => handle(&gateway, &headers, body).await, + Ok(JsonObject(body)) => handle(&gateway, &identity, &headers, body).await, Err(error) => Err(error), }; result.map_err(|error| (error.status(), Json(error.body(request_id.as_deref())))) @@ -35,10 +37,12 @@ pub async fn create( async fn handle( gateway: &Gateway, + identity: &AuthenticatedRequest, headers: &HeaderMap, body: Map, ) -> Result { let deployment = request::resolve_deployment(gateway, &body)?; + request::authorize_model(identity, deployment, &body).await?; let call = project(deployment, body, headers)?; let machine = gateway.messages.clone().machine(call); let stream = diff --git a/litellm-rust/crates/gateway-inference/src/ocr.rs b/litellm-rust/crates/gateway-inference/src/ocr.rs index 787183ccb33..a7b9d26d656 100644 --- a/litellm-rust/crates/gateway-inference/src/ocr.rs +++ b/litellm-rust/crates/gateway-inference/src/ocr.rs @@ -1,3 +1,4 @@ +use litellm_gateway_auth::AuthenticatedRequest; use std::sync::Arc; use axum::{Json, extract::State, http::HeaderMap, response::IntoResponse}; @@ -13,14 +14,16 @@ use crate::{ pub(crate) async fn create( State(gateway): State>, + identity: AuthenticatedRequest, headers: HeaderMap, body: InferenceBody, ) -> Result { - handle(&gateway, &headers, body).await.map(Json) + handle(&gateway, &identity, &headers, body).await.map(Json) } async fn handle( gateway: &Gateway, + identity: &AuthenticatedRequest, headers: &HeaderMap, InferenceBody { fields: body, @@ -32,6 +35,7 @@ async fn handle( .and_then(|value| value.to_str().ok()) .map(str::to_owned); let deployment = request::resolve_deployment(gateway, &body)?; + request::authorize_model(identity, deployment, &body).await?; let document = match upload { Some(upload) => OcrDocumentInput::Bytes { bytes: upload.bytes, diff --git a/litellm-rust/crates/gateway-inference/src/request.rs b/litellm-rust/crates/gateway-inference/src/request.rs index dfedef93a76..b398da2013b 100644 --- a/litellm-rust/crates/gateway-inference/src/request.rs +++ b/litellm-rust/crates/gateway-inference/src/request.rs @@ -90,6 +90,26 @@ fn object(body: &[u8]) -> Result, Error> { } } +pub(crate) async fn authorize_model( + identity: &litellm_gateway_auth::AuthenticatedRequest, + deployment: &Deployment, + body: &Map, +) -> Result<(), Error> { + let name = body + .get("model") + .and_then(Value::as_str) + .ok_or_else(|| Error::InvalidBody("model is required".into()))?; + let access = litellm_gateway_auth::AccessRequest::Model { + name: name.into(), + deployment: deployment.model.clone(), + }; + identity + .authorize(access.clone()) + .await? + .consume(identity.caller(), &access)?; + Ok(()) +} + pub(crate) fn resolve_deployment<'a>( gateway: &'a Gateway, body: &Map, diff --git a/litellm-rust/crates/gateway-inference/src/responses.rs b/litellm-rust/crates/gateway-inference/src/responses.rs index 97040ce6656..a6ce4bad4ac 100644 --- a/litellm-rust/crates/gateway-inference/src/responses.rs +++ b/litellm-rust/crates/gateway-inference/src/responses.rs @@ -2,6 +2,7 @@ use std::sync::Arc; use axum::{Json, body::Bytes, extract::State, response::Response}; use litellm_core::responses::{route::Responses, types::ResponsesCall}; +use litellm_gateway_auth::AuthenticatedRequest; use litellm_host_http::Sse; use serde_json::json; @@ -9,9 +10,11 @@ use crate::{Error, Gateway, JsonObject, request}; pub(crate) async fn create( State(gateway): State>, + identity: AuthenticatedRequest, JsonObject(body): JsonObject, ) -> Result { let deployment = request::resolve_deployment(&gateway, &body)?; + request::authorize_model(&identity, deployment, &body).await?; let call = ResponsesCall { model: deployment.model.clone(), input: body.get("input").cloned().unwrap_or_default(), diff --git a/litellm-rust/crates/gateway-inference/tests/messages.rs b/litellm-rust/crates/gateway-inference/tests/messages.rs index 8d2578a954b..818fdab27f2 100644 --- a/litellm-rust/crates/gateway-inference/tests/messages.rs +++ b/litellm-rust/crates/gateway-inference/tests/messages.rs @@ -172,3 +172,22 @@ async fn hosted_provider_failure_preserves_status_body_and_request_id() { assert_eq!(body["error"], error["error"]); assert_eq!(body["request_id"], "host-http-request"); } + +#[rstest] +#[case::messages("/v1/messages")] +#[case::chat("/v1/chat/completions")] +#[case::model_path("/engines/path-model/chat/completions")] +#[case::ocr("/ocr")] +#[case::transcription("/audio/transcriptions")] +#[tokio::test] +async fn model_permissions_prevent_provider_calls(#[case] path: &str) { + let upstream = MockServer::start().await; + let app = support::app_with_permissions( + "anthropic/test-model", + &upstream.uri(), + litellm_gateway_auth::Permissions::None, + ); + let response = support::post(app, path, json!({"model": "public/model"})).await; + assert_eq!(response.status(), 403); + assert!(upstream.received_requests().await.unwrap().is_empty()); +} diff --git a/litellm-rust/crates/gateway-inference/tests/responses.rs b/litellm-rust/crates/gateway-inference/tests/responses.rs index 68b40a16472..ded7d73c5f3 100644 --- a/litellm-rust/crates/gateway-inference/tests/responses.rs +++ b/litellm-rust/crates/gateway-inference/tests/responses.rs @@ -97,3 +97,22 @@ async fn responses_preserve_upstream_errors_without_retry() { .contains("slow down") ); } + +#[rstest] +#[tokio::test] +async fn responses_authorize_the_public_model_before_provider_execution() { + let upstream = MockServer::start().await; + let app = support::app_with_permissions( + "openai/test-model", + &upstream.uri(), + litellm_gateway_auth::Permissions::None, + ); + let response = support::post( + app, + "/v1/responses", + json!({"model": "public/model", "input": "hello"}), + ) + .await; + assert_eq!(response.status(), 403); + assert!(upstream.received_requests().await.unwrap().is_empty()); +} diff --git a/litellm-rust/crates/gateway-inference/tests/support/mod.rs b/litellm-rust/crates/gateway-inference/tests/support/mod.rs index 7c35f73b82b..b59e335334e 100644 --- a/litellm-rust/crates/gateway-inference/tests/support/mod.rs +++ b/litellm-rust/crates/gateway-inference/tests/support/mod.rs @@ -26,6 +26,14 @@ impl SecretSource for NoSecrets { } pub fn app(model: &str, api_base: &str) -> Router { + app_with_permissions(model, api_base, litellm_gateway_auth::Permissions::All) +} + +pub fn app_with_permissions( + model: &str, + api_base: &str, + permissions: litellm_gateway_auth::Permissions, +) -> Router { let pool = Arc::new(HttpClientPool::new(Arc::new(PublicDnsResolver))); let http = Resolution::from(&HttpSettings::default()).config; let secrets = Arc::new(NoSecrets); @@ -50,6 +58,32 @@ pub fn app(model: &str, api_base: &str) -> Router { ) .unwrap(), )) + .layer(axum::middleware::from_fn_with_state( + permissions, + test_identity, + )) +} + +async fn test_identity( + axum::extract::State(permissions): axum::extract::State, + mut request: axum::extract::Request, + next: axum::middleware::Next, +) -> Response { + let auth = litellm_gateway_auth::Auth::new( + Arc::new(litellm_gateway_auth::MasterKeyAuthenticator::new( + Some(SecretValue::new("test-inbound-key")), + Arc::new(NoSecrets), + )), + Arc::new(TestPermissions(permissions)), + Arc::new(litellm_gateway_auth::NoAdditionalPolicy), + Arc::new(litellm_gateway_auth::SystemClock), + ); + let identity = auth + .authenticate(&SecretValue::new("test-inbound-key")) + .await + .unwrap(); + request.extensions_mut().insert(identity); + next.run(request).await } pub async fn post(app: Router, path: &str, body: Value) -> Response { @@ -66,3 +100,19 @@ pub async fn post(app: Router, path: &str, body: Value) -> Response { pub async fn json(response: Response) -> Value { serde_json::from_slice(&to_bytes(response.into_body(), 1024 * 1024).await.unwrap()).unwrap() } + +struct TestPermissions(litellm_gateway_auth::Permissions); + +impl litellm_gateway_auth::IdentityResolver for TestPermissions { + fn resolve<'a>( + &'a self, + identity: &'a litellm_gateway_auth::VerifiedIdentity, + ) -> litellm_gateway_auth::AuthFuture<'a, litellm_gateway_auth::ResolvedIdentity> { + Box::pin(async move { + Ok(litellm_gateway_auth::ResolvedIdentity { + principal: identity.principal.clone(), + permissions: self.0.clone(), + }) + }) + } +} diff --git a/litellm-rust/crates/gateway/src/lib.rs b/litellm-rust/crates/gateway/src/lib.rs index 33d96234015..9510e626f26 100644 --- a/litellm-rust/crates/gateway/src/lib.rs +++ b/litellm-rust/crates/gateway/src/lib.rs @@ -11,7 +11,7 @@ use http_body_util::BodyExt; use litellm_config::Config; use litellm_core::resources::CoreResources; -use litellm_gateway_auth::{Auth, RequireMasterKey}; +use litellm_gateway_auth::Auth; use litellm_gateway_inference::{Gateway, ModelList}; use litellm_http::{ ClientVariant, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver, @@ -37,10 +37,10 @@ pub fn build_inference(config: &Config) -> Result, litellm_http::Er pub fn router(inference: Arc, config: &Config) -> Router { let auth = Auth::from_config(config, inference.secrets.clone()); litellm_gateway_inference::router(inference) - .route_layer(axum::middleware::from_extractor_with_state::< - RequireMasterKey, - _, - >(auth)) + .route_layer(axum::middleware::from_fn_with_state( + auth, + litellm_gateway_auth::authenticate, + )) .layer(axum::middleware::from_fn(log_request)) }