systemprompt_api/services/middleware/session/
mod.rs1mod 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 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 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}