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//! Validation runs through [`attest_session`], the same predicate the JWT and
9//! gateway credential paths use: a cookie must name a session the server issued
10//! *to that user*. An existence-only check would let a signed token borrow
11//! another user's live session for analytics attribution.
12//!
13//! Copyright (c) systemprompt.io — Business Source License 1.1.
14//! See <https://systemprompt.io> for licensing details.
15
16mod attestation;
17mod lifecycle;
18mod skip;
19
20pub use attestation::{SessionAttestationError, attest_session};
21pub use skip::should_skip_session_tracking;
22
23use axum::extract::{ConnectInfo, Request};
24use axum::http::header;
25use axum::middleware::Next;
26use axum::response::Response;
27use ipnet::IpNet;
28use std::net::SocketAddr;
29use std::sync::Arc;
30use systemprompt_analytics::{AnalyticsService, SessionAnalytics};
31use systemprompt_identifiers::{AgentName, ContextId, SessionId, UserId};
32use systemprompt_models::api::ApiError;
33use systemprompt_models::auth::UserType;
34use systemprompt_models::execution::context::RequestContext;
35use systemprompt_oauth::services::SessionCreationService;
36use systemprompt_runtime::AppContext;
37use systemprompt_security::{
38    CookieExtractor, HeaderExtractor, TokenExtractor, extract_user_context,
39};
40use systemprompt_traits::{AnalyticsProvider, ExtractSignals};
41use systemprompt_users::UserService;
42use uuid::Uuid;
43
44struct RequestMeta<'a> {
45    headers: &'a http::HeaderMap,
46    uri: &'a http::Uri,
47    analytics: &'a SessionAnalytics,
48}
49
50#[derive(Clone, Debug)]
51pub struct SessionMiddleware {
52    analytics_service: Arc<AnalyticsService>,
53    session_creation_service: Arc<SessionCreationService>,
54    trusted_proxies: Arc<Vec<IpNet>>,
55    ignored_forwarded_warn: Arc<systemprompt_logging::LogThrottle>,
56}
57
58const IGNORED_FORWARDED_WARN_INTERVAL_SECS: u64 = 3600;
59
60impl SessionMiddleware {
61    pub fn new(ctx: &AppContext) -> anyhow::Result<Self> {
62        let user_service = UserService::new(ctx.db_pool())?;
63        let concrete = Arc::clone(ctx.analytics_service());
64        let analytics: Arc<dyn AnalyticsProvider> = concrete;
65        let session_creation_service = Arc::new(SessionCreationService::new(
66            analytics,
67            Arc::new(user_service),
68        ));
69
70        Ok(Self {
71            analytics_service: Arc::clone(ctx.analytics_service()),
72            session_creation_service,
73            trusted_proxies: Arc::new(ctx.config().trusted_proxies.clone()),
74            ignored_forwarded_warn: Arc::new(systemprompt_logging::LogThrottle::new(
75                IGNORED_FORWARDED_WARN_INTERVAL_SECS,
76            )),
77        })
78    }
79
80    pub async fn handle(&self, mut request: Request, next: Next) -> Result<Response, ApiError> {
81        let caller_ip = super::client_addr::resolve_client_ip(
82            request.headers(),
83            request.extensions().get::<ConnectInfo<SocketAddr>>(),
84            &self.trusted_proxies,
85        );
86        if let Some(peer) = request.extensions().get::<ConnectInfo<SocketAddr>>()
87            && super::client_addr::forwarded_headers_ignored(
88                request.headers(),
89                peer.0.ip(),
90                &self.trusted_proxies,
91            )
92            && self.ignored_forwarded_warn.allow()
93        {
94            tracing::warn!(
95                peer_ip = %peer.0.ip(),
96                "ignoring forwarded client-IP headers from untrusted private peer; if this \
97                 server runs behind a proxy, add the peer's range to server.trusted_proxies"
98            );
99        }
100        let uri = request.uri().clone();
101        let headers = request.headers();
102        let analytics = self.analytics_service.extract_analytics(
103            headers,
104            ExtractSignals {
105                uri: Some(&uri),
106                caller_ip,
107            },
108        );
109        let meta = RequestMeta {
110            headers,
111            uri: &uri,
112            analytics: &analytics,
113        };
114
115        let should_skip = should_skip_session_tracking(uri.path());
116
117        tracing::debug!(
118            path = %uri.path(),
119            should_skip = should_skip,
120            "Session middleware evaluating request"
121        );
122
123        let trace_id = HeaderExtractor::extract_trace_id(headers);
124
125        let (req_ctx, jwt_cookie) = if should_skip {
126            (
127                self.anonymous_context("untracked", trace_id, &meta).await?,
128                None,
129            )
130        } else {
131            self.tracked_context(trace_id, &meta).await?
132        };
133
134        tracing::debug!(
135            path = %uri.path(),
136            session_id = %req_ctx.session_id(),
137            "Session middleware setting context"
138        );
139
140        request.extensions_mut().insert(req_ctx);
141
142        let mut response = next.run(request).await;
143
144        if let Some(token) = jwt_cookie {
145            let cookie = format!(
146                "{}={token}; HttpOnly; SameSite=Strict; Path=/; Max-Age=604800",
147                CookieExtractor::DEFAULT_COOKIE_NAME
148            );
149            if let Ok(cookie_value) = cookie.parse() {
150                response
151                    .headers_mut()
152                    .insert(header::SET_COOKIE, cookie_value);
153            }
154        }
155
156        Ok(response)
157    }
158
159    // Why: Builds an untracked anonymous context. `session_prefix` distinguishes
160    // the synthetic session id (`untracked_*` for skip-tracking paths,
161    // `bot_*` for detected crawlers) so the two cases stay legible in logs
162    // and analytics without two near-identical constructors.
163    async fn anonymous_context(
164        &self,
165        session_prefix: &str,
166        trace_id: systemprompt_identifiers::TraceId,
167        meta: &RequestMeta<'_>,
168    ) -> Result<RequestContext, ApiError> {
169        let (user_id, fingerprint) = self
170            .session_creation_service
171            .ensure_anonymous_user(meta.analytics)
172            .await
173            .map_err(|e| {
174                tracing::error!(error = %e, session_prefix, "Failed to ensure anonymous user");
175                ApiError::internal_error("Service temporarily unavailable")
176            })?;
177
178        Ok(RequestContext::new(
179            SessionId::new(format!("{session_prefix}_{}", Uuid::new_v4())),
180            trace_id,
181            ContextId::generate(),
182            AgentName::system(),
183        )
184        .with_actor(systemprompt_identifiers::Actor::anonymous(user_id))
185        .with_user_type(UserType::Anon)
186        .with_tracked(false)
187        .with_fingerprint_hash(fingerprint))
188    }
189
190    async fn tracked_context(
191        &self,
192        trace_id: systemprompt_identifiers::TraceId,
193        meta: &RequestMeta<'_>,
194    ) -> Result<(RequestContext, Option<String>), ApiError> {
195        tracing::debug!(
196            path = %meta.uri.path(),
197            skip_tracking = meta.analytics.skip_tracking,
198            user_agent = ?meta.analytics.user_agent,
199            "Session middleware bot check"
200        );
201
202        if meta.analytics.skip_tracking {
203            return Ok((self.anonymous_context("bot", trace_id, meta).await?, None));
204        }
205
206        let token_result = TokenExtractor::browser_only().extract(meta.headers).ok();
207
208        let (session_id, user_id, jwt_token, jwt_cookie, fingerprint_hash) =
209            self.resolve_session(token_result, meta).await?;
210
211        let context_id =
212            HeaderExtractor::extract_context_id(meta.headers).unwrap_or_else(ContextId::generate);
213
214        let mut ctx = RequestContext::new(session_id, trace_id, context_id, AgentName::system())
215            .with_actor(systemprompt_identifiers::Actor::user(user_id))
216            .with_auth_token(jwt_token)
217            .with_user_type(UserType::Anon)
218            .with_tracked(true);
219        if let Some(fp) = fingerprint_hash {
220            ctx = ctx.with_fingerprint_hash(fp);
221        }
222        Ok((ctx, jwt_cookie))
223    }
224
225    async fn resolve_session(
226        &self,
227        token_result: Option<String>,
228        meta: &RequestMeta<'_>,
229    ) -> Result<(SessionId, UserId, String, Option<String>, Option<String>), ApiError> {
230        let Some(token) = token_result else {
231            let (sid, uid, token, is_new, fp) =
232                lifecycle::create_new_session(&self.session_creation_service, meta).await?;
233            let jwt_cookie = if is_new { Some(token.clone()) } else { None };
234            return Ok((sid, uid, token, jwt_cookie, Some(fp)));
235        };
236
237        let Ok(jwt_context) = extract_user_context(&token) else {
238            let (sid, uid, token, is_new, fp) =
239                lifecycle::create_new_session(&self.session_creation_service, meta).await?;
240            let jwt_cookie = if is_new { Some(token.clone()) } else { None };
241            return Ok((sid, uid, token, jwt_cookie, Some(fp)));
242        };
243
244        let analytics_provider: Arc<dyn AnalyticsProvider> =
245            Arc::<AnalyticsService>::clone(&self.analytics_service);
246
247        // Why: a lookup failure is infrastructure, not evidence of a bad
248        // session, so it takes the same mint-a-replacement recovery as a
249        // missing one rather than failing the request.
250        match attest_session(
251            &analytics_provider,
252            &jwt_context.session_id,
253            &jwt_context.user_id,
254            "session_middleware",
255        )
256        .await
257        {
258            Ok(()) => {
259                return Ok((
260                    jwt_context.session_id,
261                    jwt_context.user_id,
262                    token,
263                    None,
264                    None,
265                ));
266            },
267            Err(e) => tracing::info!(
268                old_session_id = %jwt_context.session_id,
269                user_id = %jwt_context.user_id,
270                reason = %e,
271                "JWT session failed attestation, refreshing with new session"
272            ),
273        }
274
275        match lifecycle::refresh_session_for_user(
276            &self.session_creation_service,
277            &jwt_context.user_id,
278            meta,
279        )
280        .await
281        {
282            Ok((sid, uid, new_token, _, fp)) => {
283                Ok((sid, uid, new_token.clone(), Some(new_token), Some(fp)))
284            },
285            Err(e) if e.error_key.as_deref() == Some("user_not_found") => {
286                tracing::warn!(
287                    user_id = %jwt_context.user_id,
288                    "JWT references non-existent user, creating new anonymous session"
289                );
290                let (sid, uid, token, _, fp) =
291                    lifecycle::create_new_session(&self.session_creation_service, meta).await?;
292                Ok((sid, uid, token.clone(), Some(token), Some(fp)))
293            },
294            Err(e) => Err(e),
295        }
296    }
297}