systemprompt_api/services/middleware/jwt/
validation.rs1use 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}