systemprompt_api/services/middleware/analytics/
mod.rs1pub mod detection;
13pub mod events;
14
15use axum::extract::Request;
16use axum::http::StatusCode;
17use axum::middleware::Next;
18use axum::response::Response;
19use std::sync::Arc;
20
21use systemprompt_analytics::SessionSignalsRepository;
22use systemprompt_identifiers::SessionId;
23use systemprompt_logging::AnalyticsRepository;
24use systemprompt_models::{RequestContext, RouteClassifier};
25use systemprompt_runtime::AppContext;
26use systemprompt_security::ScannerDetector;
27use systemprompt_traits::{BackgroundTasks, DynSessionStore};
28
29pub use events::AnalyticsEventParams;
30
31struct TrackingParams<'a> {
32 req_ctx: &'a RequestContext,
33 uri: &'a http::Uri,
34 method: &'a http::Method,
35 status_code: u16,
36 response_time_ms: u64,
37 user_agent: Option<String>,
38 referer: Option<String>,
39 is_scanner: bool,
40 html_response: bool,
41}
42
43#[derive(Clone)]
44pub struct AnalyticsMiddleware {
45 sessions: DynSessionStore,
46 signals: Arc<SessionSignalsRepository>,
47 analytics_repo: Arc<AnalyticsRepository>,
48 route_classifier: Arc<RouteClassifier>,
49 background: BackgroundTasks,
50}
51
52impl std::fmt::Debug for AnalyticsMiddleware {
53 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
54 f.debug_struct("AnalyticsMiddleware")
55 .field("signals", &self.signals)
56 .field("analytics_repo", &self.analytics_repo)
57 .field("route_classifier", &self.route_classifier)
58 .field("background", &self.background)
59 .finish_non_exhaustive()
60 }
61}
62
63impl AnalyticsMiddleware {
64 pub fn new(app_context: &AppContext) -> Self {
65 let repositories = app_context.analytics_repositories();
66 let sessions = Arc::clone(&repositories.session_store);
67 let signals = Arc::new(repositories.session_signals.clone());
68 let analytics_repo = Arc::new(AnalyticsRepository::new(app_context.db_pool()));
69 let route_classifier = Arc::clone(app_context.route_classifier());
70
71 Self {
72 sessions,
73 signals,
74 analytics_repo,
75 route_classifier,
76 background: app_context.background_tasks().clone(),
77 }
78 }
79
80 pub async fn track_request(
81 &self,
82 request: Request,
83 next: Next,
84 ) -> Result<Response, StatusCode> {
85 let method = request.method().clone();
86 let uri = request.uri().clone();
87
88 let Some(req_ctx) = request.extensions().get::<RequestContext>().cloned() else {
89 return Ok(next.run(request).await);
90 };
91
92 if !req_ctx.request.is_tracked {
93 return Ok(next.run(request).await);
94 }
95
96 let user_agent = request
97 .headers()
98 .get("user-agent")
99 .and_then(|v| v.to_str().ok())
100 .map(str::to_owned);
101
102 let referer = request
103 .headers()
104 .get("referer")
105 .and_then(|v| v.to_str().ok())
106 .map(str::to_owned);
107
108 let start_time = std::time::Instant::now();
109 let response = next.run(request).await;
110 let response_time_ms = start_time.elapsed().as_millis() as u64;
111 let status_code = response.status();
112 let html_response = response
113 .headers()
114 .get(http::header::CONTENT_TYPE)
115 .and_then(|v| v.to_str().ok())
116 .is_some_and(|ct| ct.trim_start().starts_with("text/html"));
117
118 let should_track = self
119 .route_classifier
120 .should_track_analytics(uri.path(), method.as_str());
121
122 let is_scanner =
123 ScannerDetector::is_scanner(Some(uri.path()), user_agent.as_deref(), None, None);
124
125 if should_track {
126 self.spawn_tracking_tasks(TrackingParams {
127 req_ctx: &req_ctx,
128 uri: &uri,
129 method: &method,
130 status_code: status_code.as_u16(),
131 response_time_ms,
132 user_agent,
133 referer,
134 is_scanner,
135 html_response,
136 });
137 }
138
139 Ok(response)
140 }
141
142 fn spawn_tracking_tasks(&self, params: TrackingParams<'_>) {
143 let req_ctx = params.req_ctx;
144 let uri = params.uri;
145 let method = params.method;
146 let status_code = params.status_code;
147 let response_time_ms = params.response_time_ms;
148 let user_agent = params.user_agent;
149 let referer = params.referer;
150 let is_scanner = params.is_scanner;
151 let html_response = params.html_response;
152 let endpoint = format!("{} {}", method, uri.path());
153 let path = uri.path().to_owned();
154
155 if is_scanner {
156 self.spawn_mark_scanner_task(req_ctx.request.session_id.clone());
157 }
158
159 self.spawn_velocity_scanner_check(req_ctx.request.session_id.clone());
160
161 self.spawn_session_tracking_task(req_ctx.request.session_id.clone());
162
163 detection::spawn_behavioral_detection_task(
164 &self.background,
165 Arc::clone(&self.sessions),
166 Arc::clone(&self.signals),
167 detection::DetectionSubject {
168 session_id: req_ctx.request.session_id.clone(),
169 fingerprint_hash: req_ctx.request.fingerprint_hash.clone(),
170 user_agent: user_agent.clone(),
171 request_count: 1,
172 },
173 );
174
175 events::spawn_analytics_event_task(
176 &self.background,
177 Arc::clone(&self.analytics_repo),
178 Arc::clone(&self.route_classifier),
179 AnalyticsEventParams {
180 req_ctx: req_ctx.clone(),
181 endpoint,
182 path,
183 method: method.to_string(),
184 uri: uri.clone(),
185 status_code,
186 response_time_ms,
187 user_agent,
188 referer,
189 html_response,
190 },
191 );
192 }
193
194 fn spawn_session_tracking_task(&self, session_id: SessionId) {
195 let sessions = Arc::clone(&self.sessions);
196
197 self.background
200 .spawn("analytics_session_activity", async move {
201 if let Err(e) = sessions.increment_request_count(&session_id).await {
202 tracing::error!(error = %e, "Failed to record session request");
203 }
204 });
205 }
206
207 fn spawn_velocity_scanner_check(&self, session_id: SessionId) {
208 let sessions = Arc::clone(&self.sessions);
209
210 self.background
211 .spawn("analytics_velocity_check", async move {
212 let (request_count, duration_seconds) =
213 match sessions.get_session_velocity(&session_id).await {
214 Ok(velocity) => velocity,
215 Err(e) => {
216 tracing::warn!(
217 error = %e,
218 session_id = %session_id,
219 "Failed to read session velocity; scanner check skipped"
220 );
221 return;
222 },
223 };
224
225 if let (Some(count), Some(duration)) = (request_count, duration_seconds)
226 && ScannerDetector::is_high_velocity(count, duration)
227 && let Err(e) = sessions.mark_as_scanner(&session_id).await
228 {
229 tracing::warn!(
230 error = %e,
231 session_id = %session_id,
232 "Failed to mark high-velocity session as scanner"
233 );
234 }
235 });
236 }
237
238 fn spawn_mark_scanner_task(&self, session_id: SessionId) {
239 let sessions = Arc::clone(&self.sessions);
240
241 self.background.spawn("analytics_mark_scanner", async move {
242 if let Err(e) = sessions.mark_as_scanner(&session_id).await {
243 tracing::warn!(error = %e, session_id = %session_id, "Failed to mark session as scanner");
244 }
245 });
246 }
247}