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