Skip to main content

systemprompt_security/auth/
validation.rs

1//! JWT validation service producing a `RequestContext` from session claims.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use axum::http::HeaderMap;
7use systemprompt_identifiers::{AccessTokenId, Actor, ContextId, JwtToken, SessionId, UserId};
8use systemprompt_models::auth::JwtAudience;
9use systemprompt_models::execution::context::RequestContext;
10
11use crate::error::{AuthError, AuthResult};
12use crate::extraction::{HeaderExtractor, TokenExtractor};
13use crate::jwt::{ValidationPolicy, decode_session_claims};
14use crate::session::ValidatedSessionClaims;
15
16#[derive(Debug)]
17pub struct AuthValidationService {
18    issuer: String,
19    audiences: Vec<JwtAudience>,
20}
21
22impl AuthValidationService {
23    #[must_use]
24    pub const fn new(issuer: String, audiences: Vec<JwtAudience>) -> Self {
25        Self { issuer, audiences }
26    }
27
28    pub fn validate_request(&self, headers: &HeaderMap) -> AuthResult<RequestContext> {
29        let token = TokenExtractor::extract_from_authorization(headers)
30            .map_err(|_e| AuthError::MissingAuthorization)?;
31        let claims = self.validate_token(&token)?;
32        Ok(Self::create_context_from_claims(&claims, &token, headers))
33    }
34
35    fn validate_token(&self, token: &str) -> AuthResult<ValidatedSessionClaims> {
36        let policy = ValidationPolicy::issuer_scoped(&self.issuer, &self.audiences);
37        let claims = decode_session_claims(token, &policy)?;
38
39        Ok(ValidatedSessionClaims {
40            user_id: UserId::try_new(claims.sub).map_err(AuthError::InvalidSubject)?,
41            session_id: claims
42                .session_id
43                .map(SessionId::new)
44                .ok_or(AuthError::MissingSessionId)?,
45            user_type: claims.user_type,
46            jti: (!claims.jti.is_empty()).then(|| AccessTokenId::new(claims.jti)),
47            exp: claims.exp,
48        })
49    }
50
51    fn create_context_from_claims(
52        claims: &ValidatedSessionClaims,
53        token: &str,
54        headers: &HeaderMap,
55    ) -> RequestContext {
56        let session_id = claims.session_id.clone();
57        let user_id = claims.user_id.clone();
58
59        let context_id = HeaderExtractor::extract_context_id(headers)
60            .unwrap_or_else(|| ContextId::derived_from_session(&session_id));
61
62        let ctx = RequestContext::new(
63            session_id,
64            HeaderExtractor::extract_trace_id(headers),
65            context_id,
66            HeaderExtractor::extract_agent_name(headers),
67            Actor::user(user_id),
68        )
69        .with_auth_token(JwtToken::new(token))
70        .with_user_type(claims.user_type)
71        .with_token_exp(claims.exp);
72        match &claims.jti {
73            Some(jti) => ctx.with_jti(jti.clone()),
74            None => ctx,
75        }
76    }
77}