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