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) -> Self {
62 let user_service = UserService::new(Arc::clone(ctx.user_repository()));
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 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(
160 &self,
161 session_prefix: &str,
162 trace_id: systemprompt_identifiers::TraceId,
163 meta: &RequestMeta<'_>,
164 ) -> Result<RequestContext, ApiError> {
165 let (user_id, fingerprint) = self
166 .session_creation_service
167 .ensure_anonymous_user(meta.analytics)
168 .await
169 .map_err(|e| {
170 tracing::error!(error = %e, session_prefix, "Failed to ensure anonymous user");
171 ApiError::internal_error("Service temporarily unavailable")
172 })?;
173
174 let session_id = SessionId::new(format!("{session_prefix}_{}", Uuid::new_v4()));
175 let context_id = ContextId::derived_from_session(&session_id);
176 Ok(
177 RequestContext::new(session_id, trace_id, context_id, AgentName::system())
178 .with_actor(systemprompt_identifiers::Actor::anonymous(user_id))
179 .with_user_type(UserType::Anon)
180 .with_tracked(false)
181 .with_fingerprint_hash(fingerprint),
182 )
183 }
184
185 async fn tracked_context(
186 &self,
187 trace_id: systemprompt_identifiers::TraceId,
188 meta: &RequestMeta<'_>,
189 ) -> Result<(RequestContext, Option<String>), ApiError> {
190 tracing::debug!(
191 path = %meta.uri.path(),
192 skip_tracking = meta.analytics.skip_tracking,
193 user_agent = ?meta.analytics.user_agent,
194 "Session middleware bot check"
195 );
196
197 if meta.analytics.skip_tracking {
198 return Ok((self.anonymous_context("bot", trace_id, meta).await?, None));
199 }
200
201 let token_result = TokenExtractor::browser_only().extract(meta.headers).ok();
202
203 let (session_id, user_id, jwt_token, jwt_cookie, fingerprint_hash) =
204 self.resolve_session(token_result, meta).await?;
205
206 let context_id = HeaderExtractor::extract_context_id(meta.headers)
207 .unwrap_or_else(|| ContextId::derived_from_session(&session_id));
208
209 let mut ctx = RequestContext::new(session_id, trace_id, context_id, AgentName::system())
210 .with_actor(systemprompt_identifiers::Actor::user(user_id))
211 .with_auth_token(jwt_token)
212 .with_user_type(UserType::Anon)
213 .with_tracked(true);
214 if let Some(fp) = fingerprint_hash {
215 ctx = ctx.with_fingerprint_hash(fp);
216 }
217 Ok((ctx, jwt_cookie))
218 }
219
220 async fn resolve_session(
221 &self,
222 token_result: Option<String>,
223 meta: &RequestMeta<'_>,
224 ) -> Result<(SessionId, UserId, String, Option<String>, Option<String>), ApiError> {
225 let Some(token) = token_result else {
226 let (sid, uid, token, is_new, fp) =
227 lifecycle::create_new_session(&self.session_creation_service, meta).await?;
228 let jwt_cookie = if is_new { Some(token.clone()) } else { None };
229 return Ok((sid, uid, token, jwt_cookie, Some(fp)));
230 };
231
232 let Ok(jwt_context) = extract_user_context(&token) else {
233 let (sid, uid, token, is_new, fp) =
234 lifecycle::create_new_session(&self.session_creation_service, meta).await?;
235 let jwt_cookie = if is_new { Some(token.clone()) } else { None };
236 return Ok((sid, uid, token, jwt_cookie, Some(fp)));
237 };
238
239 let analytics_provider: Arc<dyn AnalyticsProvider> =
240 Arc::<AnalyticsService>::clone(&self.analytics_service);
241
242 match attest_session(
243 &analytics_provider,
244 &jwt_context.session_id,
245 &jwt_context.user_id,
246 "session_middleware",
247 )
248 .await
249 {
250 Ok(()) => {
251 return Ok((
252 jwt_context.session_id,
253 jwt_context.user_id,
254 token,
255 None,
256 None,
257 ));
258 },
259 Err(e) => tracing::info!(
260 old_session_id = %jwt_context.session_id,
261 user_id = %jwt_context.user_id,
262 reason = %e,
263 "JWT session failed attestation, refreshing with new session"
264 ),
265 }
266
267 match lifecycle::refresh_session_for_user(
268 &self.session_creation_service,
269 &jwt_context.user_id,
270 meta,
271 )
272 .await
273 {
274 Ok((sid, uid, new_token, _, fp)) => {
275 Ok((sid, uid, new_token.clone(), Some(new_token), Some(fp)))
276 },
277 Err(e) if e.error_key.as_deref() == Some("user_not_found") => {
278 tracing::warn!(
279 user_id = %jwt_context.user_id,
280 "JWT references non-existent user, creating new anonymous session"
281 );
282 let (sid, uid, token, _, fp) =
283 lifecycle::create_new_session(&self.session_creation_service, meta).await?;
284 Ok((sid, uid, token.clone(), Some(token), Some(fp)))
285 },
286 Err(e) => Err(e),
287 }
288 }
289}