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