Skip to main content

systemprompt_api/services/middleware/jwt/
validation.rs

1//! JWT validation middleware with a short-TTL user cache.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use std::collections::HashMap;
7use std::sync::{Arc, Mutex};
8use std::time::{Duration, Instant};
9
10use systemprompt_identifiers::UserId;
11use systemprompt_models::auth::UserRole;
12use systemprompt_models::execution::context::ContextExtractionError;
13use systemprompt_security::JwtUserContext;
14use systemprompt_traits::{AuthUser, SessionProvider, UserProvider};
15
16use crate::services::middleware::session::{SessionAttestationError, attest_session};
17
18const USER_CACHE_TTL: Duration = Duration::from_secs(30);
19
20#[derive(Debug)]
21pub struct ValidatedUser {
22    pub user: AuthUser,
23}
24
25#[derive(Debug)]
26pub struct UserCache {
27    entries: Mutex<HashMap<UserId, (AuthUser, Instant)>>,
28    ttl: Duration,
29}
30
31impl Default for UserCache {
32    fn default() -> Self {
33        Self::with_ttl(USER_CACHE_TTL)
34    }
35}
36
37impl UserCache {
38    pub fn new() -> Arc<Self> {
39        Arc::new(Self::default())
40    }
41
42    pub fn with_ttl(ttl: Duration) -> Self {
43        Self {
44            entries: Mutex::new(HashMap::new()),
45            ttl,
46        }
47    }
48
49    pub fn get_fresh(&self, user_id: &UserId) -> Option<AuthUser> {
50        let guard = self.entries.lock().ok()?;
51        let fresh = guard
52            .get(user_id)
53            .and_then(|(user, fetched_at)| (fetched_at.elapsed() < self.ttl).then(|| user.clone()));
54        drop(guard);
55        fresh
56    }
57
58    pub fn put(&self, user_id: UserId, user: AuthUser) {
59        if let Ok(mut guard) = self.entries.lock() {
60            guard.insert(user_id, (user, Instant::now()));
61        }
62    }
63}
64
65pub async fn validate_user_exists(
66    user_provider: &Arc<dyn UserProvider>,
67    cache: &Arc<UserCache>,
68    jwt_context: &JwtUserContext,
69    route_context: &str,
70) -> Result<ValidatedUser, ContextExtractionError> {
71    if let Some(user) = cache.get_fresh(&jwt_context.user_id) {
72        return require_active(user, jwt_context, route_context);
73    }
74
75    let user = user_provider
76        .find_by_id(&jwt_context.user_id)
77        .await
78        .map_err(|e| ContextExtractionError::DatabaseError {
79            context: "Failed to check user existence".to_owned(),
80            source: e.into(),
81        })?
82        .ok_or_else(|| {
83            tracing::info!(
84                session_id = %jwt_context.session_id.as_str(),
85                user_id = %jwt_context.user_id.as_str(),
86                route = %route_context,
87                "JWT validation failed: user no longer exists in database"
88            );
89            ContextExtractionError::UserNotFound(format!(
90                "User {} no longer exists",
91                jwt_context.user_id.as_str()
92            ))
93        })?;
94
95    cache.put(jwt_context.user_id.clone(), user.clone());
96    require_active(user, jwt_context, route_context)
97}
98
99fn require_active(
100    user: AuthUser,
101    jwt_context: &JwtUserContext,
102    route_context: &str,
103) -> Result<ValidatedUser, ContextExtractionError> {
104    if !user.is_active {
105        tracing::info!(
106            session_id = %jwt_context.session_id.as_str(),
107            user_id = %jwt_context.user_id.as_str(),
108            route = %route_context,
109            "JWT validation failed: user is not active"
110        );
111        return Err(ContextExtractionError::UserNotFound(format!(
112            "User {} is not active",
113            jwt_context.user_id.as_str()
114        )));
115    }
116    Ok(ValidatedUser { user })
117}
118
119pub fn user_is_admin(user: &AuthUser) -> bool {
120    user.roles
121        .iter()
122        .any(|r| r.as_str() == UserRole::Admin.as_str())
123}
124
125pub(super) async fn validate_session_exists(
126    session_provider: &Arc<dyn SessionProvider>,
127    jwt_context: &JwtUserContext,
128    route_context: &str,
129) -> Result<(), ContextExtractionError> {
130    attest_session(
131        session_provider,
132        &jwt_context.session_id,
133        &jwt_context.user_id,
134        route_context,
135    )
136    .await
137    .map_err(|e| match e {
138        SessionAttestationError::Lookup(source) => ContextExtractionError::DatabaseError {
139            context: "Failed to check session".to_owned(),
140            source: Box::new(source),
141        },
142        other => ContextExtractionError::InvalidToken(other.into()),
143    })
144}