systemprompt_api/services/middleware/session/
mod.rs1mod lifecycle;
12mod skip;
13
14pub use skip::should_skip_session_tracking;
15
16use axum::extract::{ConnectInfo, Request};
17use axum::http::header;
18use axum::middleware::Next;
19use axum::response::Response;
20use ipnet::IpNet;
21use std::net::SocketAddr;
22use std::sync::Arc;
23use systemprompt_analytics::{AnalyticsService, SessionAnalytics};
24use systemprompt_identifiers::{AgentName, ContextId, SessionId, UserId};
25use systemprompt_models::api::ApiError;
26use systemprompt_models::auth::UserType;
27use systemprompt_models::execution::context::RequestContext;
28use systemprompt_oauth::services::SessionCreationService;
29use systemprompt_runtime::AppContext;
30use systemprompt_security::{
31 CookieExtractor, HeaderExtractor, TokenExtractor, extract_user_context,
32};
33use systemprompt_traits::{AnalyticsProvider, ExtractSignals};
34use systemprompt_users::UserService;
35use uuid::Uuid;
36
37struct RequestMeta<'a> {
41 headers: &'a http::HeaderMap,
42 uri: &'a http::Uri,
43 analytics: &'a SessionAnalytics,
44}
45
46#[derive(Clone, Debug)]
47pub struct SessionMiddleware {
48 analytics_service: Arc<AnalyticsService>,
49 session_creation_service: Arc<SessionCreationService>,
50 trusted_proxies: Arc<Vec<IpNet>>,
51}
52
53impl SessionMiddleware {
54 pub fn new(ctx: &AppContext) -> anyhow::Result<Self> {
55 let user_service = UserService::new(ctx.db_pool())?;
56 let concrete = Arc::clone(ctx.analytics_service());
57 let analytics: Arc<dyn AnalyticsProvider> = concrete;
58 let session_creation_service = Arc::new(SessionCreationService::new(
59 analytics,
60 Arc::new(user_service),
61 ));
62
63 Ok(Self {
64 analytics_service: Arc::clone(ctx.analytics_service()),
65 session_creation_service,
66 trusted_proxies: Arc::new(ctx.config().trusted_proxies.clone()),
67 })
68 }
69
70 pub async fn handle(&self, mut request: Request, next: Next) -> Result<Response, ApiError> {
71 let caller_ip = super::client_addr::resolve_client_ip(
72 request.headers(),
73 request.extensions().get::<ConnectInfo<SocketAddr>>(),
74 &self.trusted_proxies,
75 );
76 let uri = request.uri().clone();
77 let headers = request.headers();
78 let analytics = self.analytics_service.extract_analytics(
79 headers,
80 ExtractSignals {
81 uri: Some(&uri),
82 caller_ip,
83 },
84 );
85 let meta = RequestMeta {
86 headers,
87 uri: &uri,
88 analytics: &analytics,
89 };
90
91 let should_skip = should_skip_session_tracking(uri.path());
92
93 tracing::debug!(
94 path = %uri.path(),
95 should_skip = should_skip,
96 "Session middleware evaluating request"
97 );
98
99 let trace_id = HeaderExtractor::extract_trace_id(headers);
100
101 let (req_ctx, jwt_cookie) = if should_skip {
102 (
103 self.anonymous_context("untracked", trace_id, &meta).await?,
104 None,
105 )
106 } else {
107 self.tracked_context(trace_id, &meta).await?
108 };
109
110 tracing::debug!(
111 path = %uri.path(),
112 session_id = %req_ctx.session_id(),
113 "Session middleware setting context"
114 );
115
116 request.extensions_mut().insert(req_ctx);
117
118 let mut response = next.run(request).await;
119
120 if let Some(token) = jwt_cookie {
121 let cookie = format!(
122 "{}={token}; HttpOnly; SameSite=Strict; Path=/; Max-Age=604800",
123 CookieExtractor::DEFAULT_COOKIE_NAME
124 );
125 if let Ok(cookie_value) = cookie.parse() {
126 response
127 .headers_mut()
128 .insert(header::SET_COOKIE, cookie_value);
129 }
130 }
131
132 Ok(response)
133 }
134
135 async fn anonymous_context(
140 &self,
141 session_prefix: &str,
142 trace_id: systemprompt_identifiers::TraceId,
143 meta: &RequestMeta<'_>,
144 ) -> Result<RequestContext, ApiError> {
145 let (user_id, fingerprint) = self
146 .session_creation_service
147 .ensure_anonymous_user(meta.analytics)
148 .await
149 .map_err(|e| {
150 tracing::error!(error = %e, session_prefix, "Failed to ensure anonymous user");
151 ApiError::internal_error("Service temporarily unavailable")
152 })?;
153
154 Ok(RequestContext::new(
155 SessionId::new(format!("{session_prefix}_{}", Uuid::new_v4())),
156 trace_id,
157 ContextId::generate(),
158 AgentName::system(),
159 )
160 .with_actor(systemprompt_identifiers::Actor::anonymous(user_id))
161 .with_user_type(UserType::Anon)
162 .with_tracked(false)
163 .with_fingerprint_hash(fingerprint))
164 }
165
166 async fn tracked_context(
167 &self,
168 trace_id: systemprompt_identifiers::TraceId,
169 meta: &RequestMeta<'_>,
170 ) -> Result<(RequestContext, Option<String>), ApiError> {
171 tracing::debug!(
172 path = %meta.uri.path(),
173 skip_tracking = meta.analytics.skip_tracking,
174 user_agent = ?meta.analytics.user_agent,
175 "Session middleware bot check"
176 );
177
178 if meta.analytics.skip_tracking {
179 return Ok((self.anonymous_context("bot", trace_id, meta).await?, None));
180 }
181
182 let token_result = TokenExtractor::browser_only().extract(meta.headers).ok();
183
184 let (session_id, user_id, jwt_token, jwt_cookie, fingerprint_hash) =
185 self.resolve_session(token_result, meta).await?;
186
187 let context_id =
188 HeaderExtractor::extract_context_id(meta.headers).unwrap_or_else(ContextId::generate);
189
190 let mut ctx = RequestContext::new(session_id, trace_id, context_id, AgentName::system())
191 .with_actor(systemprompt_identifiers::Actor::user(user_id))
192 .with_auth_token(jwt_token)
193 .with_user_type(UserType::Anon)
194 .with_tracked(true);
195 if let Some(fp) = fingerprint_hash {
196 ctx = ctx.with_fingerprint_hash(fp);
197 }
198 Ok((ctx, jwt_cookie))
199 }
200
201 async fn resolve_session(
202 &self,
203 token_result: Option<String>,
204 meta: &RequestMeta<'_>,
205 ) -> Result<(SessionId, UserId, String, Option<String>, Option<String>), ApiError> {
206 let Some(token) = token_result else {
207 let (sid, uid, token, is_new, fp) =
208 lifecycle::create_new_session(&self.session_creation_service, meta).await?;
209 let jwt_cookie = if is_new { Some(token.clone()) } else { None };
210 return Ok((sid, uid, token, jwt_cookie, Some(fp)));
211 };
212
213 let Ok(jwt_context) = extract_user_context(&token) else {
214 let (sid, uid, token, is_new, fp) =
215 lifecycle::create_new_session(&self.session_creation_service, meta).await?;
216 let jwt_cookie = if is_new { Some(token.clone()) } else { None };
217 return Ok((sid, uid, token, jwt_cookie, Some(fp)));
218 };
219
220 let session_exists = self
221 .analytics_service
222 .find_active_session_by_id(&jwt_context.session_id)
223 .await
224 .map_err(|e| {
225 tracing::warn!(error = %e, "find_active_session_by_id failed");
226 e
227 })
228 .ok()
229 .flatten()
230 .is_some();
231
232 if session_exists {
233 return Ok((
234 jwt_context.session_id,
235 jwt_context.user_id,
236 token,
237 None,
238 None,
239 ));
240 }
241
242 tracing::info!(
243 old_session_id = %jwt_context.session_id,
244 user_id = %jwt_context.user_id,
245 "JWT valid but session missing, refreshing with new session"
246 );
247 match lifecycle::refresh_session_for_user(
248 &self.session_creation_service,
249 &jwt_context.user_id,
250 meta,
251 )
252 .await
253 {
254 Ok((sid, uid, new_token, _, fp)) => {
255 Ok((sid, uid, new_token.clone(), Some(new_token), Some(fp)))
256 },
257 Err(e) if e.error_key.as_deref() == Some("user_not_found") => {
258 tracing::warn!(
259 user_id = %jwt_context.user_id,
260 "JWT references non-existent user, creating new anonymous session"
261 );
262 let (sid, uid, token, _, fp) =
263 lifecycle::create_new_session(&self.session_creation_service, meta).await?;
264 Ok((sid, uid, token.clone(), Some(token), Some(fp)))
265 },
266 Err(e) => Err(e),
267 }
268 }
269}