systemprompt-runtime 0.65.0

Application runtime for systemprompt.io AI governance infrastructure. AppContext, lifecycle builder, extension registry, and module wiring for the MCP governance pipeline.
Documentation
//! Audit-view queries joining messages, tool calls, and linked MCP rows.
//!
//! Copyright (c) systemprompt.io — Business Source License 1.1.
//! See <https://systemprompt.io> for licensing details.

use systemprompt_identifiers::{AiRequestId, McpServerId, McpToolName, TaskId, TraceId};

use super::{Result, TraceRepository};
use crate::trace::models::{
    AuditLookupResult, AuditPage, AuditToolCallRow, ConversationMessage, LinkedMcpCall,
};

struct AuditRow {
    id: AiRequestId,
    provider: Option<String>,
    model: Option<String>,
    requested_model: Option<String>,
    input_tokens: Option<i32>,
    output_tokens: Option<i32>,
    cache_read_tokens: Option<i32>,
    cache_creation_tokens: Option<i32>,
    reasoning_tokens: Option<i32>,
    cost_microdollars: i64,
    latency_ms: Option<i32>,
    status: String,
    finish_reason: Option<String>,
    error_message: Option<String>,
    task_id: Option<TaskId>,
    trace_id: Option<TraceId>,
}

struct MsgRow {
    role: String,
    content: String,
    sequence_number: i32,
}

fn audit_row_to_result(r: AuditRow) -> AuditLookupResult {
    AuditLookupResult {
        id: r.id,
        provider: r.provider,
        model: r.model,
        requested_model: r.requested_model,
        input_tokens: r.input_tokens,
        output_tokens: r.output_tokens,
        cache_read_tokens: r.cache_read_tokens,
        cache_creation_tokens: r.cache_creation_tokens,
        reasoning_tokens: r.reasoning_tokens,
        cost_microdollars: r.cost_microdollars,
        latency_ms: r.latency_ms,
        status: r.status,
        finish_reason: r.finish_reason,
        error_message: r.error_message,
        task_id: r.task_id,
        trace_id: r.trace_id,
    }
}

impl TraceRepository {
    pub async fn find_ai_request_for_audit(&self, id: &str) -> Result<Option<AuditLookupResult>> {
        let partial = format!("{id}%");

        if let Some(row) = self.find_audit_by_request_id(id, &partial).await? {
            return Ok(Some(row));
        }
        if let Some(row) = self.find_audit_by_task_id(id, &partial).await? {
            return Ok(Some(row));
        }
        self.find_audit_by_trace_id(id, &partial).await
    }

    async fn find_audit_by_request_id(
        &self,
        id: &str,
        partial: &str,
    ) -> Result<Option<AuditLookupResult>> {
        let row = sqlx::query_as!(
            AuditRow,
            r#"
        SELECT id as "id!: AiRequestId", provider, model,
            requested_model,
            input_tokens, output_tokens,
            cache_read_tokens, cache_creation_tokens, reasoning_tokens,
            cost_microdollars as "cost_microdollars!",
            latency_ms,
            status as "status!",
            finish_reason,
            error_message,
            task_id as "task_id: TaskId",
            trace_id as "trace_id: TraceId"
        FROM ai_requests WHERE id = $1 OR id LIKE $2 LIMIT 1
        "#,
            id,
            partial
        )
        .fetch_optional(&*self.pool)
        .await?;

        Ok(row.map(audit_row_to_result))
    }

    async fn find_audit_by_task_id(
        &self,
        id: &str,
        partial: &str,
    ) -> Result<Option<AuditLookupResult>> {
        let row = sqlx::query_as!(
            AuditRow,
            r#"
        SELECT id as "id!: AiRequestId", provider, model,
            requested_model,
            input_tokens, output_tokens,
            cache_read_tokens, cache_creation_tokens, reasoning_tokens,
            cost_microdollars as "cost_microdollars!",
            latency_ms,
            status as "status!",
            finish_reason,
            error_message,
            task_id as "task_id: TaskId",
            trace_id as "trace_id: TraceId"
        FROM ai_requests WHERE task_id = $1 OR task_id LIKE $2
        ORDER BY created_at DESC LIMIT 1
        "#,
            id,
            partial
        )
        .fetch_optional(&*self.pool)
        .await?;

        Ok(row.map(audit_row_to_result))
    }

    async fn find_audit_by_trace_id(
        &self,
        id: &str,
        partial: &str,
    ) -> Result<Option<AuditLookupResult>> {
        let row = sqlx::query_as!(
            AuditRow,
            r#"
        SELECT id as "id!: AiRequestId", provider, model,
            requested_model,
            input_tokens, output_tokens,
            cache_read_tokens, cache_creation_tokens, reasoning_tokens,
            cost_microdollars as "cost_microdollars!",
            latency_ms,
            status as "status!",
            finish_reason,
            error_message,
            task_id as "task_id: TaskId",
            trace_id as "trace_id: TraceId"
        FROM ai_requests WHERE trace_id = $1 OR trace_id LIKE $2
        ORDER BY created_at DESC LIMIT 1
        "#,
            id,
            partial
        )
        .fetch_optional(&*self.pool)
        .await?;

        Ok(row.map(audit_row_to_result))
    }

    pub async fn count_audit_messages(&self, request_id: &AiRequestId) -> Result<i64> {
        let count = sqlx::query_scalar!(
            r#"SELECT COUNT(*) as "count!" FROM ai_request_messages WHERE request_id = $1"#,
            request_id.as_str()
        )
        .fetch_one(&*self.pool)
        .await?;
        Ok(count)
    }

    pub async fn count_audit_tool_calls(&self, request_id: &AiRequestId) -> Result<i64> {
        let count = sqlx::query_scalar!(
            r#"SELECT COUNT(*) as "count!" FROM ai_request_tool_calls WHERE request_id = $1"#,
            request_id.as_str()
        )
        .fetch_one(&*self.pool)
        .await?;
        Ok(count)
    }

    pub async fn list_audit_messages(
        &self,
        request_id: &AiRequestId,
        page: AuditPage,
    ) -> Result<Vec<ConversationMessage>> {
        let rows = sqlx::query_as!(
            MsgRow,
            r#"
        SELECT role as "role!", content as "content!", sequence_number as "sequence_number!"
        FROM ai_request_messages WHERE request_id = $1 ORDER BY sequence_number
        OFFSET $2 LIMIT $3
        "#,
            request_id.as_str(),
            page.offset,
            page.sql_limit()
        )
        .fetch_all(&*self.pool)
        .await?;

        Ok(rows
            .into_iter()
            .map(|m| ConversationMessage {
                role: m.role,
                content: m.content,
                sequence_number: m.sequence_number,
            })
            .collect())
    }

    pub async fn list_audit_tool_calls(
        &self,
        request_id: &AiRequestId,
        page: AuditPage,
    ) -> Result<Vec<AuditToolCallRow>> {
        let rows = sqlx::query!(
            r#"
        SELECT tool_name as "tool_name!", tool_input as "tool_input!",
            sequence_number as "sequence_number!"
        FROM ai_request_tool_calls WHERE request_id = $1 ORDER BY sequence_number
        OFFSET $2 LIMIT $3
        "#,
            request_id.as_str(),
            page.offset,
            page.sql_limit()
        )
        .fetch_all(&*self.pool)
        .await?;

        Ok(rows
            .into_iter()
            .map(|t| AuditToolCallRow {
                tool_name: McpToolName::new(t.tool_name),
                tool_input: t.tool_input,
                sequence_number: t.sequence_number,
            })
            .collect())
    }

    pub async fn list_linked_mcp_calls(
        &self,
        request_id: &AiRequestId,
    ) -> Result<Vec<LinkedMcpCall>> {
        let rows = sqlx::query!(
            r#"
        SELECT
            mte.tool_name as "tool_name!",
            mte.server_name as "server_name!",
            mte.status as "status!",
            mte.execution_time_ms
        FROM mcp_tool_executions mte
        JOIN ai_request_tool_calls artc ON artc.mcp_execution_id = mte.mcp_execution_id
        WHERE artc.request_id = $1
        "#,
            request_id.as_str()
        )
        .fetch_all(&*self.pool)
        .await?;

        Ok(rows
            .into_iter()
            .map(|r| LinkedMcpCall {
                tool_name: McpToolName::new(r.tool_name),
                server_name: McpServerId::new(r.server_name),
                status: r.status,
                execution_time_ms: r.execution_time_ms,
            })
            .collect())
    }
}