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