systemprompt-api 0.63.0

Axum-based HTTP server and API gateway for systemprompt.io AI governance infrastructure. Exposes governed agents, MCP, A2A, and admin endpoints with rate limiting and RBAC.
Documentation
//! Request analytics middleware.
//!
//! [`AnalyticsMiddleware`] records tracked requests after the response is
//! produced, running session activity, velocity-based scanner detection,
//! behavioural bot scoring, and analytics-event capture on the process's
//! [`BackgroundTasks`] so the request path is never blocked on persistence and
//! shutdown drains every pending write.
//!
//! Copyright (c) systemprompt.io — Business Source License 1.1.
//! See <https://systemprompt.io> for licensing details.

pub mod detection;
pub mod events;

use axum::extract::Request;
use axum::http::StatusCode;
use axum::middleware::Next;
use axum::response::Response;
use std::sync::Arc;

use systemprompt_analytics::SessionSignalsRepository;
use systemprompt_identifiers::SessionId;
use systemprompt_logging::AnalyticsRepository;
use systemprompt_models::{RequestContext, RouteClassifier};
use systemprompt_runtime::AppContext;
use systemprompt_security::ScannerDetector;
use systemprompt_traits::{BackgroundTasks, DynSessionStore};

pub use events::AnalyticsEventParams;

struct TrackingParams<'a> {
    req_ctx: &'a RequestContext,
    uri: &'a http::Uri,
    method: &'a http::Method,
    status_code: u16,
    response_time_ms: u64,
    user_agent: Option<String>,
    referer: Option<String>,
    is_scanner: bool,
    html_response: bool,
}

#[derive(Clone)]
pub struct AnalyticsMiddleware {
    sessions: DynSessionStore,
    signals: Arc<SessionSignalsRepository>,
    analytics_repo: Arc<AnalyticsRepository>,
    route_classifier: Arc<RouteClassifier>,
    background: BackgroundTasks,
}

impl std::fmt::Debug for AnalyticsMiddleware {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        f.debug_struct("AnalyticsMiddleware")
            .field("signals", &self.signals)
            .field("analytics_repo", &self.analytics_repo)
            .field("route_classifier", &self.route_classifier)
            .field("background", &self.background)
            .finish_non_exhaustive()
    }
}

impl AnalyticsMiddleware {
    pub fn new(app_context: &AppContext) -> Self {
        let repositories = app_context.analytics_repositories();
        let sessions = Arc::clone(&repositories.session_store);
        let signals = Arc::new(repositories.session_signals.clone());
        let analytics_repo = Arc::new(AnalyticsRepository::new(app_context.db_pool()));
        let route_classifier = Arc::clone(app_context.route_classifier());

        Self {
            sessions,
            signals,
            analytics_repo,
            route_classifier,
            background: app_context.background_tasks().clone(),
        }
    }

    pub async fn track_request(
        &self,
        request: Request,
        next: Next,
    ) -> Result<Response, StatusCode> {
        let method = request.method().clone();
        let uri = request.uri().clone();

        let Some(req_ctx) = request.extensions().get::<RequestContext>().cloned() else {
            return Ok(next.run(request).await);
        };

        if !req_ctx.request.is_tracked {
            return Ok(next.run(request).await);
        }

        let user_agent = request
            .headers()
            .get("user-agent")
            .and_then(|v| v.to_str().ok())
            .map(str::to_owned);

        let referer = request
            .headers()
            .get("referer")
            .and_then(|v| v.to_str().ok())
            .map(str::to_owned);

        let start_time = std::time::Instant::now();
        let response = next.run(request).await;
        let response_time_ms = start_time.elapsed().as_millis() as u64;
        let status_code = response.status();
        let html_response = response
            .headers()
            .get(http::header::CONTENT_TYPE)
            .and_then(|v| v.to_str().ok())
            .is_some_and(|ct| ct.trim_start().starts_with("text/html"));

        let should_track = self
            .route_classifier
            .should_track_analytics(uri.path(), method.as_str());

        let is_scanner =
            ScannerDetector::is_scanner(Some(uri.path()), user_agent.as_deref(), None, None);

        if should_track {
            self.spawn_tracking_tasks(TrackingParams {
                req_ctx: &req_ctx,
                uri: &uri,
                method: &method,
                status_code: status_code.as_u16(),
                response_time_ms,
                user_agent,
                referer,
                is_scanner,
                html_response,
            });
        }

        Ok(response)
    }

    fn spawn_tracking_tasks(&self, params: TrackingParams<'_>) {
        let req_ctx = params.req_ctx;
        let uri = params.uri;
        let method = params.method;
        let status_code = params.status_code;
        let response_time_ms = params.response_time_ms;
        let user_agent = params.user_agent;
        let referer = params.referer;
        let is_scanner = params.is_scanner;
        let html_response = params.html_response;
        let endpoint = format!("{} {}", method, uri.path());
        let path = uri.path().to_owned();

        if is_scanner {
            self.spawn_mark_scanner_task(req_ctx.request.session_id.clone());
        }

        self.spawn_velocity_scanner_check(req_ctx.request.session_id.clone());

        self.spawn_session_tracking_task(req_ctx.request.session_id.clone());

        detection::spawn_behavioral_detection_task(
            &self.background,
            Arc::clone(&self.sessions),
            Arc::clone(&self.signals),
            detection::DetectionSubject {
                session_id: req_ctx.request.session_id.clone(),
                fingerprint_hash: req_ctx.request.fingerprint_hash.clone(),
                user_agent: user_agent.clone(),
                request_count: 1,
            },
        );

        events::spawn_analytics_event_task(
            &self.background,
            Arc::clone(&self.analytics_repo),
            Arc::clone(&self.route_classifier),
            AnalyticsEventParams {
                req_ctx: req_ctx.clone(),
                endpoint,
                path,
                method: method.to_string(),
                uri: uri.clone(),
                status_code,
                response_time_ms,
                user_agent,
                referer,
                html_response,
            },
        );
    }

    fn spawn_session_tracking_task(&self, session_id: SessionId) {
        let sessions = Arc::clone(&self.sessions);

        // Why: one UPDATE per request — the increment also stamps
        // last_activity_at and duration_seconds.
        self.background
            .spawn("analytics_session_activity", async move {
                if let Err(e) = sessions.increment_request_count(&session_id).await {
                    tracing::error!(error = %e, "Failed to record session request");
                }
            });
    }

    fn spawn_velocity_scanner_check(&self, session_id: SessionId) {
        let sessions = Arc::clone(&self.sessions);

        self.background
            .spawn("analytics_velocity_check", async move {
                let (request_count, duration_seconds) =
                    match sessions.get_session_velocity(&session_id).await {
                        Ok(velocity) => velocity,
                        Err(e) => {
                            tracing::warn!(
                                error = %e,
                                session_id = %session_id,
                                "Failed to read session velocity; scanner check skipped"
                            );
                            return;
                        },
                    };

                if let (Some(count), Some(duration)) = (request_count, duration_seconds)
                    && ScannerDetector::is_high_velocity(count, duration)
                    && let Err(e) = sessions.mark_as_scanner(&session_id).await
                {
                    tracing::warn!(
                        error = %e,
                        session_id = %session_id,
                        "Failed to mark high-velocity session as scanner"
                    );
                }
            });
    }

    fn spawn_mark_scanner_task(&self, session_id: SessionId) {
        let sessions = Arc::clone(&self.sessions);

        self.background.spawn("analytics_mark_scanner", async move {
            if let Err(e) = sessions.mark_as_scanner(&session_id).await {
                tracing::warn!(error = %e, session_id = %session_id, "Failed to mark session as scanner");
            }
        });
    }
}