Skip to main content

systemprompt_api/services/middleware/analytics/
mod.rs

1//! Request analytics middleware.
2//!
3//! [`AnalyticsMiddleware`] records tracked requests after the response is
4//! produced, running session activity, velocity-based scanner detection,
5//! behavioural bot scoring, and analytics-event capture on the process's
6//! [`BackgroundTasks`] so the request path is never blocked on persistence and
7//! shutdown drains every pending write.
8//!
9//! Copyright (c) systemprompt.io — Business Source License 1.1.
10//! See <https://systemprompt.io> for licensing details.
11
12pub 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        // Why: one UPDATE per request — the increment also stamps
198        // last_activity_at and duration_seconds.
199        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}