use axum::http::HeaderMap;
use systemprompt_identifiers::{AccessTokenId, Actor, ContextId, JwtToken, SessionId, UserId};
use systemprompt_models::auth::JwtAudience;
use systemprompt_models::execution::context::RequestContext;
use crate::error::{AuthError, AuthResult};
use crate::extraction::{HeaderExtractor, TokenExtractor};
use crate::jwt::{ValidationPolicy, decode_session_claims};
use crate::session::ValidatedSessionClaims;
#[derive(Debug)]
pub struct AuthValidationService {
issuer: String,
audiences: Vec<JwtAudience>,
}
impl AuthValidationService {
#[must_use]
pub const fn new(issuer: String, audiences: Vec<JwtAudience>) -> Self {
Self { issuer, audiences }
}
pub fn validate_request(&self, headers: &HeaderMap) -> AuthResult<RequestContext> {
let token = TokenExtractor::extract_from_authorization(headers)
.map_err(|_e| AuthError::MissingAuthorization)?;
let claims = self.validate_token(&token)?;
Ok(Self::create_context_from_claims(&claims, &token, headers))
}
fn validate_token(&self, token: &str) -> AuthResult<ValidatedSessionClaims> {
let policy = ValidationPolicy::issuer_scoped(&self.issuer, &self.audiences);
let claims = decode_session_claims(token, &policy)?;
Ok(ValidatedSessionClaims {
user_id: UserId::try_new(claims.sub).map_err(AuthError::InvalidSubject)?,
session_id: claims
.session_id
.map(SessionId::new)
.ok_or(AuthError::MissingSessionId)?,
user_type: claims.user_type,
jti: (!claims.jti.is_empty()).then(|| AccessTokenId::new(claims.jti)),
exp: claims.exp,
})
}
fn create_context_from_claims(
claims: &ValidatedSessionClaims,
token: &str,
headers: &HeaderMap,
) -> RequestContext {
let session_id = claims.session_id.clone();
let user_id = claims.user_id.clone();
let context_id = HeaderExtractor::extract_context_id(headers)
.unwrap_or_else(|| ContextId::derived_from_session(&session_id));
let ctx = RequestContext::new(
session_id,
HeaderExtractor::extract_trace_id(headers),
context_id,
HeaderExtractor::extract_agent_name(headers),
Actor::user(user_id),
)
.with_auth_token(JwtToken::new(token))
.with_user_type(claims.user_type)
.with_token_exp(claims.exp);
match &claims.jti {
Some(jti) => ctx.with_jti(jti.clone()),
None => ctx,
}
}
}