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