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::{AnalyticsProvider, AuthUser, 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 message: format!("Failed to check user existence: {e}"),
80 })?
81 .ok_or_else(|| {
82 tracing::info!(
83 session_id = %jwt_context.session_id.as_str(),
84 user_id = %jwt_context.user_id.as_str(),
85 route = %route_context,
86 "JWT validation failed: user no longer exists in database"
87 );
88 ContextExtractionError::UserNotFound(format!(
89 "User {} no longer exists",
90 jwt_context.user_id.as_str()
91 ))
92 })?;
93
94 cache.put(jwt_context.user_id.clone(), user.clone());
95 require_active(user, jwt_context, route_context)
96}
97
98fn require_active(
99 user: AuthUser,
100 jwt_context: &JwtUserContext,
101 route_context: &str,
102) -> Result<ValidatedUser, ContextExtractionError> {
103 if !user.is_active {
104 tracing::info!(
105 session_id = %jwt_context.session_id.as_str(),
106 user_id = %jwt_context.user_id.as_str(),
107 route = %route_context,
108 "JWT validation failed: user is not active"
109 );
110 return Err(ContextExtractionError::UserNotFound(format!(
111 "User {} is not active",
112 jwt_context.user_id.as_str()
113 )));
114 }
115 Ok(ValidatedUser { user })
116}
117
118pub fn user_is_admin(user: &AuthUser) -> bool {
119 user.roles
120 .iter()
121 .any(|r| r.as_str() == UserRole::Admin.as_str())
122}
123
124pub(super) async fn validate_session_exists(
125 analytics_provider: &Arc<dyn AnalyticsProvider>,
126 jwt_context: &JwtUserContext,
127 route_context: &str,
128) -> Result<(), ContextExtractionError> {
129 attest_session(
130 analytics_provider,
131 &jwt_context.session_id,
132 &jwt_context.user_id,
133 route_context,
134 )
135 .await
136 .map_err(|e| match e {
137 SessionAttestationError::Lookup(message) => ContextExtractionError::DatabaseError {
138 message: format!("Failed to check session: {message}"),
139 },
140 other => ContextExtractionError::InvalidToken(other.to_string()),
141 })
142}