relay-knowledge 1.1.17

Graph-database-based knowledge graph project.
Documentation
//! Projects bounded MCP, QoS, graph, and index diagnostics as Prometheus metrics.

#[cfg(test)]
mod mod_tests;

use axum::{
    extract::State,
    http::{HeaderMap, StatusCode, header},
    response::{IntoResponse, Response},
};
use std::time::Duration;

use super::{
    http_contract::validate_origin,
    runtime::{McpMethodError, McpServer, admit_mcp_request},
    state::SessionLookup,
    tool_contract::request_context,
};

pub(super) fn metrics_endpoint(endpoint: &str) -> String {
    endpoint_child(endpoint, "metrics")
}

fn endpoint_child(endpoint: &str, child: &str) -> String {
    if endpoint == "/" {
        format!("/{child}")
    } else {
        format!("{}/{child}", endpoint.trim_end_matches('/'))
    }
}

pub(super) async fn handle_metrics_get(
    State(server): State<McpServer>,
    headers: HeaderMap,
) -> Response {
    if let Err(status) = validate_origin(&server, &headers) {
        return status.into_response();
    }
    let permit = match admit_mcp_request(&server) {
        Ok(permit) => permit,
        Err(_) => return StatusCode::TOO_MANY_REQUESTS.into_response(),
    };
    let timeout = Duration::from_millis(server.agent.access_policy.max_runtime_ms);
    let result = tokio::time::timeout(timeout, prometheus_metrics(&server, "metrics-get")).await;
    drop(permit);

    match result {
        Ok(Ok(metrics)) => (
            StatusCode::OK,
            [(header::CONTENT_TYPE, "text/plain; version=0.0.4")],
            metrics,
        )
            .into_response(),
        Ok(Err(error)) => (
            StatusCode::INTERNAL_SERVER_ERROR,
            [(header::CONTENT_TYPE, "text/plain")],
            error.message,
        )
            .into_response(),
        Err(_) => {
            server.qos.record_timed_out();
            (
                StatusCode::REQUEST_TIMEOUT,
                [(header::CONTENT_TYPE, "text/plain")],
                "metrics endpoint exceeded max_runtime_ms".to_owned(),
            )
                .into_response()
        }
    }
}

pub(super) async fn prometheus_metrics(
    server: &McpServer,
    request_id: &str,
) -> Result<String, McpMethodError> {
    let health = server
        .service
        .health(request_context(request_id.to_owned()))
        .await
        .map_err(McpMethodError::api)?;
    let qos = server.qos.diagnostics_snapshot();
    let agent_metrics = server.metrics.snapshot();
    let mut output = String::new();
    push_metric(
        &mut output,
        "relay_knowledge_graph_version",
        "Current committed graph version.",
        health.graph.graph_version.get(),
    );
    push_metric(
        &mut output,
        "relay_knowledge_index_refresh_queue_depth",
        "Pending index refresh task count.",
        health.index_refresh.queue_depth,
    );
    push_metric(
        &mut output,
        "relay_knowledge_index_refresh_dead_letter_count",
        "Dead-lettered index refresh task count.",
        health.index_refresh.dead_letter_count,
    );
    push_metric(
        &mut output,
        "relay_knowledge_qos_in_flight_requests",
        "Current admitted MCP request count.",
        qos.usage.in_flight_requests,
    );
    push_metric(
        &mut output,
        "relay_knowledge_qos_queued_requests",
        "Current queued MCP request count.",
        qos.usage.queued_requests,
    );
    push_metric(
        &mut output,
        "relay_knowledge_qos_admitted_total",
        "Cumulative admitted network work count.",
        qos.admitted_total,
    );
    push_metric(
        &mut output,
        "relay_knowledge_qos_queued_total",
        "Cumulative work admitted through the queued QoS path.",
        qos.queued_total,
    );
    push_metric(
        &mut output,
        "relay_knowledge_qos_rejected_total",
        "Cumulative QoS rejection count.",
        qos.rejected_total,
    );
    push_metric(
        &mut output,
        "relay_knowledge_qos_timed_out_total",
        "Cumulative network timeout count.",
        qos.timed_out_total,
    );
    push_metric(
        &mut output,
        "relay_knowledge_qos_cancelled_total",
        "Cumulative cancelled network work count.",
        qos.cancelled_total,
    );
    push_metric(
        &mut output,
        "relay_knowledge_qos_dropped_total",
        "Cumulative dropped network work count.",
        qos.dropped_total,
    );
    push_metric(
        &mut output,
        "relay_knowledge_mcp_cold_start_total",
        "Recorded MCP initialize-to-tools-list cold start samples.",
        agent_metrics.cold_start_total,
    );
    push_metric(
        &mut output,
        "relay_knowledge_mcp_cold_start_duration_ms_total",
        "Total MCP initialize-to-tools-list cold start latency in milliseconds.",
        agent_metrics.cold_start_duration_ms_total,
    );
    for index in &health.indexes {
        output.push_str(&format!(
            "relay_knowledge_index_stale{{kind=\"{}\"}} {}\n",
            index.kind.as_str(),
            usize::from(index.is_stale_for(health.graph.graph_version))
        ));
    }

    Ok(output)
}

pub(super) fn record_tools_list_cold_start(server: &McpServer, session: &SessionLookup) {
    if let Some(duration_ms) = server
        .sessions
        .record_tools_list_cold_start(session.session_id())
    {
        server.metrics.record_cold_start("mcp", duration_ms);
        tracing::info!(
            protocol = "mcp",
            duration_ms,
            "recorded initialize_to_tools_list cold start"
        );
    }
}

pub(super) fn tools_list_result(server: &McpServer, session: &SessionLookup) -> serde_json::Value {
    record_tools_list_cold_start(server, session);
    super::tool_registry::tools_list_result()
}

fn push_metric(
    output: &mut String,
    name: &'static str,
    description: &'static str,
    value: impl ToString,
) {
    output.push_str(&format!("# HELP {name} {description}\n"));
    output.push_str(&format!("# TYPE {name} gauge\n"));
    output.push_str(&format!("{name} {}\n", value.to_string()));
}