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())
}
}