Skip to main content

systemprompt_api/services/middleware/session/
mod.rs

1//! Session-establishment middleware.
2//!
3//! [`SessionMiddleware`] resolves or mints the per-request session: it skips
4//! untracked paths, short-circuits detected bots into anonymous contexts,
5//! validates an existing JWT session, and refreshes or recreates the session
6//! when the token is stale, issuing a `Set-Cookie` for newly minted tokens.
7//!
8//! Copyright (c) systemprompt.io — Business Source License 1.1.
9//! See <https://systemprompt.io> for licensing details.
10
11mod lifecycle;
12mod skip;
13
14pub use skip::should_skip_session_tracking;
15
16use axum::extract::{ConnectInfo, Request};
17use axum::http::header;
18use axum::middleware::Next;
19use axum::response::Response;
20use ipnet::IpNet;
21use std::net::SocketAddr;
22use std::sync::Arc;
23use systemprompt_analytics::{AnalyticsService, SessionAnalytics};
24use systemprompt_identifiers::{AgentName, ContextId, SessionId, UserId};
25use systemprompt_models::api::ApiError;
26use systemprompt_models::auth::UserType;
27use systemprompt_models::execution::context::RequestContext;
28use systemprompt_oauth::services::SessionCreationService;
29use systemprompt_runtime::AppContext;
30use systemprompt_security::{
31    CookieExtractor, HeaderExtractor, TokenExtractor, extract_user_context,
32};
33use systemprompt_traits::{AnalyticsProvider, ExtractSignals};
34use systemprompt_users::UserService;
35use uuid::Uuid;
36
37/// The request, extracted exactly once. `analytics` is derived at the top of
38/// [`SessionMiddleware::handle`] and passed by reference from there on — no
39/// downstream step re-derives it from `headers`.
40struct RequestMeta<'a> {
41    headers: &'a http::HeaderMap,
42    uri: &'a http::Uri,
43    analytics: &'a SessionAnalytics,
44}
45
46#[derive(Clone, Debug)]
47pub struct SessionMiddleware {
48    analytics_service: Arc<AnalyticsService>,
49    session_creation_service: Arc<SessionCreationService>,
50    trusted_proxies: Arc<Vec<IpNet>>,
51}
52
53impl SessionMiddleware {
54    pub fn new(ctx: &AppContext) -> anyhow::Result<Self> {
55        let user_service = UserService::new(ctx.db_pool())?;
56        let concrete = Arc::clone(ctx.analytics_service());
57        let analytics: Arc<dyn AnalyticsProvider> = concrete;
58        let session_creation_service = Arc::new(SessionCreationService::new(
59            analytics,
60            Arc::new(user_service),
61        ));
62
63        Ok(Self {
64            analytics_service: Arc::clone(ctx.analytics_service()),
65            session_creation_service,
66            trusted_proxies: Arc::new(ctx.config().trusted_proxies.clone()),
67        })
68    }
69
70    pub async fn handle(&self, mut request: Request, next: Next) -> Result<Response, ApiError> {
71        let caller_ip = super::client_addr::resolve_client_ip(
72            request.headers(),
73            request.extensions().get::<ConnectInfo<SocketAddr>>(),
74            &self.trusted_proxies,
75        );
76        let uri = request.uri().clone();
77        let headers = request.headers();
78        let analytics = self.analytics_service.extract_analytics(
79            headers,
80            ExtractSignals {
81                uri: Some(&uri),
82                caller_ip,
83            },
84        );
85        let meta = RequestMeta {
86            headers,
87            uri: &uri,
88            analytics: &analytics,
89        };
90
91        let should_skip = should_skip_session_tracking(uri.path());
92
93        tracing::debug!(
94            path = %uri.path(),
95            should_skip = should_skip,
96            "Session middleware evaluating request"
97        );
98
99        let trace_id = HeaderExtractor::extract_trace_id(headers);
100
101        let (req_ctx, jwt_cookie) = if should_skip {
102            (
103                self.anonymous_context("untracked", trace_id, &meta).await?,
104                None,
105            )
106        } else {
107            self.tracked_context(trace_id, &meta).await?
108        };
109
110        tracing::debug!(
111            path = %uri.path(),
112            session_id = %req_ctx.session_id(),
113            "Session middleware setting context"
114        );
115
116        request.extensions_mut().insert(req_ctx);
117
118        let mut response = next.run(request).await;
119
120        if let Some(token) = jwt_cookie {
121            let cookie = format!(
122                "{}={token}; HttpOnly; SameSite=Strict; Path=/; Max-Age=604800",
123                CookieExtractor::DEFAULT_COOKIE_NAME
124            );
125            if let Ok(cookie_value) = cookie.parse() {
126                response
127                    .headers_mut()
128                    .insert(header::SET_COOKIE, cookie_value);
129            }
130        }
131
132        Ok(response)
133    }
134
135    /// Builds an untracked anonymous context. `session_prefix` distinguishes
136    /// the synthetic session id (`untracked_*` for skip-tracking paths,
137    /// `bot_*` for detected crawlers) so the two cases stay legible in logs
138    /// and analytics without two near-identical constructors.
139    async fn anonymous_context(
140        &self,
141        session_prefix: &str,
142        trace_id: systemprompt_identifiers::TraceId,
143        meta: &RequestMeta<'_>,
144    ) -> Result<RequestContext, ApiError> {
145        let (user_id, fingerprint) = self
146            .session_creation_service
147            .ensure_anonymous_user(meta.analytics)
148            .await
149            .map_err(|e| {
150                tracing::error!(error = %e, session_prefix, "Failed to ensure anonymous user");
151                ApiError::internal_error("Service temporarily unavailable")
152            })?;
153
154        Ok(RequestContext::new(
155            SessionId::new(format!("{session_prefix}_{}", Uuid::new_v4())),
156            trace_id,
157            ContextId::generate(),
158            AgentName::system(),
159        )
160        .with_actor(systemprompt_identifiers::Actor::anonymous(user_id))
161        .with_user_type(UserType::Anon)
162        .with_tracked(false)
163        .with_fingerprint_hash(fingerprint))
164    }
165
166    async fn tracked_context(
167        &self,
168        trace_id: systemprompt_identifiers::TraceId,
169        meta: &RequestMeta<'_>,
170    ) -> Result<(RequestContext, Option<String>), ApiError> {
171        tracing::debug!(
172            path = %meta.uri.path(),
173            skip_tracking = meta.analytics.skip_tracking,
174            user_agent = ?meta.analytics.user_agent,
175            "Session middleware bot check"
176        );
177
178        if meta.analytics.skip_tracking {
179            return Ok((self.anonymous_context("bot", trace_id, meta).await?, None));
180        }
181
182        let token_result = TokenExtractor::browser_only().extract(meta.headers).ok();
183
184        let (session_id, user_id, jwt_token, jwt_cookie, fingerprint_hash) =
185            self.resolve_session(token_result, meta).await?;
186
187        let context_id =
188            HeaderExtractor::extract_context_id(meta.headers).unwrap_or_else(ContextId::generate);
189
190        let mut ctx = RequestContext::new(session_id, trace_id, context_id, AgentName::system())
191            .with_actor(systemprompt_identifiers::Actor::user(user_id))
192            .with_auth_token(jwt_token)
193            .with_user_type(UserType::Anon)
194            .with_tracked(true);
195        if let Some(fp) = fingerprint_hash {
196            ctx = ctx.with_fingerprint_hash(fp);
197        }
198        Ok((ctx, jwt_cookie))
199    }
200
201    async fn resolve_session(
202        &self,
203        token_result: Option<String>,
204        meta: &RequestMeta<'_>,
205    ) -> Result<(SessionId, UserId, String, Option<String>, Option<String>), ApiError> {
206        let Some(token) = token_result else {
207            let (sid, uid, token, is_new, fp) =
208                lifecycle::create_new_session(&self.session_creation_service, meta).await?;
209            let jwt_cookie = if is_new { Some(token.clone()) } else { None };
210            return Ok((sid, uid, token, jwt_cookie, Some(fp)));
211        };
212
213        let Ok(jwt_context) = extract_user_context(&token) else {
214            let (sid, uid, token, is_new, fp) =
215                lifecycle::create_new_session(&self.session_creation_service, meta).await?;
216            let jwt_cookie = if is_new { Some(token.clone()) } else { None };
217            return Ok((sid, uid, token, jwt_cookie, Some(fp)));
218        };
219
220        let session_exists = self
221            .analytics_service
222            .find_active_session_by_id(&jwt_context.session_id)
223            .await
224            .map_err(|e| {
225                tracing::warn!(error = %e, "find_active_session_by_id failed");
226                e
227            })
228            .ok()
229            .flatten()
230            .is_some();
231
232        if session_exists {
233            return Ok((
234                jwt_context.session_id,
235                jwt_context.user_id,
236                token,
237                None,
238                None,
239            ));
240        }
241
242        tracing::info!(
243            old_session_id = %jwt_context.session_id,
244            user_id = %jwt_context.user_id,
245            "JWT valid but session missing, refreshing with new session"
246        );
247        match lifecycle::refresh_session_for_user(
248            &self.session_creation_service,
249            &jwt_context.user_id,
250            meta,
251        )
252        .await
253        {
254            Ok((sid, uid, new_token, _, fp)) => {
255                Ok((sid, uid, new_token.clone(), Some(new_token), Some(fp)))
256            },
257            Err(e) if e.error_key.as_deref() == Some("user_not_found") => {
258                tracing::warn!(
259                    user_id = %jwt_context.user_id,
260                    "JWT references non-existent user, creating new anonymous session"
261                );
262                let (sid, uid, token, _, fp) =
263                    lifecycle::create_new_session(&self.session_creation_service, meta).await?;
264                Ok((sid, uid, token.clone(), Some(token), Some(fp)))
265            },
266            Err(e) => Err(e),
267        }
268    }
269}