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