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