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::{Actor, ContextId, SessionId, UserId};
8use systemprompt_models::auth::{JwtAudience, MAX_ACT_CHAIN_DEPTH, UserType};
9use systemprompt_models::execution::context::RequestContext;
10
11use crate::error::{AuthError, AuthResult};
12use crate::extraction::{HeaderExtractor, TokenExtractor};
13use crate::jwt::{ValidationPolicy, decode_rs256_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_rs256_claims(token, &policy)?;
38
39        if let Some(ref act) = claims.act {
40            let depth = act.depth();
41            if depth > MAX_ACT_CHAIN_DEPTH {
42                return Err(AuthError::ActChainTooDeep {
43                    depth,
44                    max: MAX_ACT_CHAIN_DEPTH,
45                });
46            }
47        }
48
49        let derived_type = UserType::from_permissions(&claims.scope);
50        if derived_type != claims.user_type {
51            return Err(AuthError::UserTypeMismatch {
52                claimed: claims.user_type,
53                derived: derived_type,
54            });
55        }
56
57        Ok(ValidatedSessionClaims {
58            user_id: UserId::new(claims.sub),
59            session_id: claims
60                .session_id
61                .map(SessionId::new)
62                .ok_or(AuthError::MissingSessionId)?,
63            user_type: derived_type,
64            jti: claims.jti,
65            exp: claims.exp,
66        })
67    }
68
69    fn create_context_from_claims(
70        claims: &ValidatedSessionClaims,
71        token: &str,
72        headers: &HeaderMap,
73    ) -> RequestContext {
74        let session_id = claims.session_id.clone();
75        let user_id = claims.user_id.clone();
76
77        let context_id = HeaderExtractor::extract_context_id(headers)
78            .unwrap_or_else(|| ContextId::derived_from_session(&session_id));
79
80        RequestContext::new(
81            session_id,
82            HeaderExtractor::extract_trace_id(headers),
83            context_id,
84            HeaderExtractor::extract_agent_name(headers),
85        )
86        .with_actor(Actor::user(user_id))
87        .with_auth_token(token)
88        .with_user_type(claims.user_type)
89        .with_jti(claims.jti.clone())
90        .with_token_exp(claims.exp)
91    }
92}