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, spawning detached tasks for session activity, velocity-based
5//! scanner detection, behavioural bot scoring, and analytics-event capture so
6//! the request path is never blocked on persistence.
7//!
8//! Copyright (c) systemprompt.io — Business Source License 1.1.
9//! See <https://systemprompt.io> for licensing details.
10
11pub 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}