systemprompt_security/auth/
validation.rs1use 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}