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