mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
feat(rust): separate gateway authentication and authorization (#43467)
Co-authored-by: Yujong Lee <yujong@berri.ai>
This commit is contained in:
parent
876539e1b3
commit
5f637a2b11
22 changed files with 1263 additions and 84 deletions
2
litellm-rust/Cargo.lock
generated
2
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
|
|
|
|||
23
litellm-rust/crates/gateway-auth/AGENTS.md
Normal file
23
litellm-rust/crates/gateway-auth/AGENTS.md
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
243
litellm-rust/crates/gateway-auth/src/authentication.rs
Normal file
243
litellm-rust/crates/gateway-auth/src/authentication.rs
Normal file
|
|
@ -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<Box<dyn Future<Output = Result<T, Error>> + 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<dyn CredentialExtractor>,
|
||||
authenticator: Arc<dyn Authenticator>,
|
||||
identities: Arc<dyn IdentityResolver>,
|
||||
authorizer: Arc<dyn Authorizer>,
|
||||
clock: Arc<dyn Clock>,
|
||||
}
|
||||
|
||||
impl Auth {
|
||||
pub fn new(
|
||||
authenticator: Arc<dyn Authenticator>,
|
||||
identities: Arc<dyn IdentityResolver>,
|
||||
authorizer: Arc<dyn Authorizer>,
|
||||
clock: Arc<dyn Clock>,
|
||||
) -> Self {
|
||||
Self {
|
||||
extractor: Arc::new(Bearer),
|
||||
authenticator,
|
||||
identities,
|
||||
authorizer,
|
||||
clock,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn from_config(config: &Config, secrets: Arc<dyn SecretSource>) -> 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<dyn CredentialExtractor>) -> Self {
|
||||
Self { extractor, ..self }
|
||||
}
|
||||
|
||||
pub async fn authenticate_parts(&self, request: &Parts) -> Result<AuthenticatedRequest, Error> {
|
||||
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<AuthenticatedRequest, Error> {
|
||||
self.authenticate_request(
|
||||
&Credential::Token(credential.clone()),
|
||||
&Request::new(()).into_parts().0,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn authenticate_request(
|
||||
&self,
|
||||
credential: &Credential,
|
||||
request: &Parts,
|
||||
) -> Result<AuthenticatedRequest, Error> {
|
||||
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<dyn Authorizer>,
|
||||
clock: Arc<dyn Clock>,
|
||||
) -> Result<AuthenticatedRequest, Error> {
|
||||
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<SecretValue>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
impl MasterKeyAuthenticator {
|
||||
pub fn new(key: Option<SecretValue>, secrets: Arc<dyn SecretSource>) -> Self {
|
||||
Self { key, secrets }
|
||||
}
|
||||
|
||||
async fn key(&self) -> Result<SecretValue, Error> {
|
||||
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,
|
||||
)
|
||||
}
|
||||
139
litellm-rust/crates/gateway-auth/src/authorization.rs
Normal file
139
litellm-rust/crates/gateway-auth/src/authorization.rs
Normal file
|
|
@ -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<dyn Authorizer>,
|
||||
clock: Arc<dyn Clock>,
|
||||
}
|
||||
|
||||
impl AuthenticatedRequest {
|
||||
pub(crate) fn new(
|
||||
caller: SharedCaller,
|
||||
authorizer: Arc<dyn Authorizer>,
|
||||
clock: Arc<dyn Clock>,
|
||||
) -> Self {
|
||||
Self {
|
||||
caller,
|
||||
authorizer,
|
||||
clock,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn caller(&self) -> &AuthenticatedCaller {
|
||||
&self.caller
|
||||
}
|
||||
|
||||
pub async fn authorize(&self, request: AccessRequest) -> Result<AuthorizedOperation, Error> {
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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()
|
||||
}
|
||||
}
|
||||
|
|
|
|||
90
litellm-rust/crates/gateway-auth/src/http.rs
Normal file
90
litellm-rust/crates/gateway-auth/src/http.rs
Normal file
|
|
@ -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<Credential, Error>;
|
||||
}
|
||||
|
||||
pub struct Bearer;
|
||||
|
||||
impl CredentialExtractor for Bearer {
|
||||
fn extract(&self, parts: &Parts) -> Result<Credential, Error> {
|
||||
bearer(parts).map(Credential::Token)
|
||||
}
|
||||
}
|
||||
|
||||
fn bearer(parts: &Parts) -> Result<SecretValue, Error> {
|
||||
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<Auth>,
|
||||
request: Request,
|
||||
next: Next,
|
||||
) -> Result<Response, Error> {
|
||||
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::<MatchedPath>()
|
||||
.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<S: Send + Sync> FromRequestParts<S> for AuthenticatedRequest {
|
||||
type Rejection = Error;
|
||||
|
||||
async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Error> {
|
||||
parts
|
||||
.extensions
|
||||
.get::<Self>()
|
||||
.cloned()
|
||||
.ok_or(Error::MissingIdentity)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct RequireMasterKey;
|
||||
|
||||
impl FromRequestParts<Auth> for RequireMasterKey {
|
||||
type Rejection = Error;
|
||||
|
||||
async fn from_request_parts(parts: &mut Parts, state: &Auth) -> Result<Self, Error> {
|
||||
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)
|
||||
}
|
||||
}
|
||||
122
litellm-rust/crates/gateway-auth/src/identity.rs
Normal file
122
litellm-rust/crates/gateway-auth/src/identity.rs
Normal file
|
|
@ -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<SystemTime>,
|
||||
}
|
||||
|
||||
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<AuthenticatedCaller>;
|
||||
|
|
@ -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<SecretValue>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
impl Auth {
|
||||
pub fn from_config(config: &Config, secrets: Arc<dyn SecretSource>) -> Self {
|
||||
Self {
|
||||
master_key: config.general_settings.master_key.clone(),
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
async fn master_key(&self) -> Result<SecretValue, Error> {
|
||||
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<Auth> for RequireMasterKey {
|
||||
type Rejection = Error;
|
||||
|
||||
async fn from_request_parts(parts: &mut Parts, state: &Auth) -> Result<Self, Error> {
|
||||
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),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<u16>,
|
||||
}
|
||||
|
||||
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<Principal>,
|
||||
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<Identities>,
|
||||
policy: Arc<Policy>,
|
||||
clock: Arc<TestClock>,
|
||||
}
|
||||
|
||||
fn pipeline(
|
||||
authority: &'static str,
|
||||
permissions: Permissions,
|
||||
restrictions: Permissions,
|
||||
failure: Option<u16>,
|
||||
) -> 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<litellm_gateway_auth::Credential, Error> {
|
||||
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<dyn SecretSource>,
|
||||
) {
|
||||
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()
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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<Arc<Gateway>>,
|
||||
identity: AuthenticatedRequest,
|
||||
body: InferenceBody,
|
||||
) -> Result<impl IntoResponse, Error> {
|
||||
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<Value, Error> {
|
||||
let deployment = request::resolve_deployment(gateway, &body)?;
|
||||
request::authorize_model(identity, deployment, &body).await?;
|
||||
let audio = match upload {
|
||||
Some(upload) => {
|
||||
let format = upload
|
||||
|
|
|
|||
|
|
@ -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<Arc<Gateway>>,
|
||||
identity: AuthenticatedRequest,
|
||||
JsonObject(body): JsonObject,
|
||||
) -> Result<impl IntoResponse, Error> {
|
||||
handle(&gateway, body).await
|
||||
handle(&gateway, &identity, body).await
|
||||
}
|
||||
|
||||
pub(crate) async fn create_from_model_path(
|
||||
State(gateway): State<Arc<Gateway>>,
|
||||
identity: AuthenticatedRequest,
|
||||
Path(model): Path<String>,
|
||||
JsonObject(body): JsonObject,
|
||||
) -> Result<impl IntoResponse, Error> {
|
||||
|
|
@ -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<String, Value>) -> Result<Response, Error> {
|
||||
async fn handle(
|
||||
gateway: &Gateway,
|
||||
identity: &AuthenticatedRequest,
|
||||
body: Map<String, Value>,
|
||||
) -> Result<Response, Error> {
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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<litellm_host_http::Error<RouteError>> 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,
|
||||
|
|
|
|||
|
|
@ -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<Arc<Gateway>>,
|
||||
identity: AuthenticatedRequest,
|
||||
RequestId(request_id): RequestId,
|
||||
headers: HeaderMap,
|
||||
body: Result<JsonObject, Error>,
|
||||
) -> 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<String, Value>,
|
||||
) -> Result<Response, Error> {
|
||||
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 =
|
||||
|
|
|
|||
|
|
@ -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<Arc<Gateway>>,
|
||||
identity: AuthenticatedRequest,
|
||||
headers: HeaderMap,
|
||||
body: InferenceBody,
|
||||
) -> Result<impl IntoResponse, Error> {
|
||||
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,
|
||||
|
|
|
|||
|
|
@ -90,6 +90,26 @@ fn object(body: &[u8]) -> Result<Map<String, Value>, Error> {
|
|||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn authorize_model(
|
||||
identity: &litellm_gateway_auth::AuthenticatedRequest,
|
||||
deployment: &Deployment,
|
||||
body: &Map<String, Value>,
|
||||
) -> 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<String, Value>,
|
||||
|
|
|
|||
|
|
@ -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<Arc<Gateway>>,
|
||||
identity: AuthenticatedRequest,
|
||||
JsonObject(body): JsonObject,
|
||||
) -> Result<Response, Error> {
|
||||
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(),
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<litellm_gateway_auth::Permissions>,
|
||||
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(),
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<Arc<Gateway>, litellm_http::Er
|
|||
pub fn router(inference: Arc<Gateway>, 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))
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue