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//! Establishing a session reads and writes the database, so it is bounded and
14//! degrades rather than fails — see the `context` submodule. A request that
15//! cannot be given a session is served with an untracked, actor-less context
16//! instead of a 500: the alternative is that a database fault takes the public
17//! site down, and a page view is worth more than the analytics row describing
18//! it. Nothing is escalated by the degraded context — it carries no auth token
19//! and no user, so every gate above `public` still refuses it.
20//!
21//! Copyright (c) systemprompt.io — Business Source License 1.1.
22//! See <https://systemprompt.io> for licensing details.
23
24mod attestation;
25mod context;
26mod lifecycle;
27mod skip;
28
29pub use attestation::{SessionAttestationError, attest_session};
30pub use skip::should_skip_session_tracking;
31
32use axum::extract::{ConnectInfo, Request};
33use axum::http::header;
34use axum::middleware::Next;
35use axum::response::Response;
36use ipnet::IpNet;
37use std::net::SocketAddr;
38use std::sync::Arc;
39use systemprompt_analytics::{AnalyticsService, SessionAnalytics};
40use systemprompt_models::api::ApiError;
41use systemprompt_oauth::services::SessionCreationService;
42use systemprompt_runtime::AppContext;
43use systemprompt_security::{CookieExtractor, HeaderExtractor};
44use systemprompt_traits::{ExtractSignals, SessionProvider};
45use systemprompt_users::UserService;
46
47struct RequestMeta<'a> {
48    headers: &'a http::HeaderMap,
49    uri: &'a http::Uri,
50    analytics: &'a SessionAnalytics,
51}
52
53#[derive(Clone, Debug)]
54pub struct SessionMiddleware {
55    analytics_service: Arc<AnalyticsService>,
56    session_creation_service: Arc<SessionCreationService>,
57    trusted_proxies: Arc<Vec<IpNet>>,
58    jwt_issuer: Arc<str>,
59    ignored_forwarded_warn: Arc<systemprompt_logging::LogThrottle>,
60    degraded_warn: Arc<systemprompt_logging::LogThrottle>,
61}
62
63const IGNORED_FORWARDED_WARN_INTERVAL_SECS: u64 = 3600;
64const DEGRADED_WARN_INTERVAL_SECS: u64 = 60;
65
66impl SessionMiddleware {
67    pub fn new(ctx: &AppContext) -> Self {
68        let user_service = UserService::new(Arc::clone(ctx.user_repository()));
69        let concrete = Arc::clone(&ctx.analytics_repositories().session_store);
70        let analytics: Arc<dyn SessionProvider> = concrete;
71        let session_creation_service = Arc::new(SessionCreationService::new(
72            analytics,
73            Arc::new(user_service),
74        ));
75
76        Self {
77            analytics_service: Arc::clone(ctx.analytics_service()),
78            session_creation_service,
79            trusted_proxies: Arc::new(ctx.config().trusted_proxies.clone()),
80            jwt_issuer: Arc::from(ctx.config().jwt_issuer.as_str()),
81            ignored_forwarded_warn: Arc::new(systemprompt_logging::LogThrottle::new(
82                IGNORED_FORWARDED_WARN_INTERVAL_SECS,
83            )),
84            degraded_warn: Arc::new(systemprompt_logging::LogThrottle::new(
85                DEGRADED_WARN_INTERVAL_SECS,
86            )),
87        }
88    }
89
90    pub async fn handle(&self, mut request: Request, next: Next) -> Result<Response, ApiError> {
91        let caller_ip = super::client_addr::resolve_client_ip(
92            request.headers(),
93            request.extensions().get::<ConnectInfo<SocketAddr>>(),
94            &self.trusted_proxies,
95        );
96        if let Some(peer) = request.extensions().get::<ConnectInfo<SocketAddr>>()
97            && super::client_addr::forwarded_headers_ignored(
98                request.headers(),
99                peer.0.ip(),
100                &self.trusted_proxies,
101            )
102            && self.ignored_forwarded_warn.allow()
103        {
104            tracing::warn!(
105                peer_ip = %peer.0.ip(),
106                "ignoring forwarded client-IP headers from an untrusted peer; if this \
107                 server runs behind a proxy, add the peer's range to server.trusted_proxies"
108            );
109        }
110        let uri = request.uri().clone();
111        let headers = request.headers();
112        let analytics = self.analytics_service.extract_analytics(
113            headers,
114            ExtractSignals {
115                uri: Some(&uri),
116                caller_ip,
117            },
118        );
119        let meta = RequestMeta {
120            headers,
121            uri: &uri,
122            analytics: &analytics,
123        };
124
125        let should_skip = should_skip_session_tracking(uri.path());
126
127        tracing::debug!(
128            path = %uri.path(),
129            should_skip = should_skip,
130            "Session middleware evaluating request"
131        );
132
133        let trace_id = HeaderExtractor::extract_trace_id(headers);
134
135        let (req_ctx, jwt_cookie) = self
136            .establish_or_degrade(should_skip, trace_id, &meta, uri.path())
137            .await;
138
139        tracing::debug!(
140            path = %uri.path(),
141            session_id = %req_ctx.session_id(),
142            "Session middleware setting context"
143        );
144
145        request.extensions_mut().insert(req_ctx);
146
147        let mut response = next.run(request).await;
148
149        if let Some(token) = jwt_cookie {
150            let cookie = format!(
151                "{}={token}; HttpOnly; SameSite=Strict; Path=/; Max-Age=604800",
152                CookieExtractor::DEFAULT_COOKIE_NAME
153            );
154            if let Ok(cookie_value) = cookie.parse() {
155                response
156                    .headers_mut()
157                    .insert(header::SET_COOKIE, cookie_value);
158            }
159        }
160
161        Ok(response)
162    }
163}