scryer-mcp 0.2.1

Model Context Protocol (MCP) server for Scryer code intelligence
use std::collections::BTreeMap;
use std::sync::Arc;
use std::time::Instant;
use tiktoken_rs::CoreBPE;
use tracing::{debug, error};

use scryer_db::{ScryerDb, ToolInvocationMetric};

pub const MAX_PAYLOAD_TOKENS: usize = 1000;
pub const COST_PER_MILLION_TOKENS: f64 = 2.50;

/// Tool categories determining the heuristic naive baseline cost formula.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ToolCategory {
    OutlineOrScope,
    DefinitionOrType,
    ReferencesOrHierarchy,
    InspectSymbol,
    CrateOutline,
    SearchSymbols,
    Default,
}

impl ToolCategory {
    pub fn for_tool(tool_name: &str) -> Self {
        match tool_name {
            "get_file_outline" | "get_enclosing_scope" => Self::OutlineOrScope,
            "resolve_definition" | "get_type_contract" => Self::DefinitionOrType,
            "find_references" | "trace_call_hierarchy" | "calculate_blast_radius" => {
                Self::ReferencesOrHierarchy
            }
            "inspect_symbol" => Self::InspectSymbol,
            "get_crate_outline" => Self::CrateOutline,
            "search_symbols" => Self::SearchSymbols,
            _ => Self::Default,
        }
    }
}

/// Token savings telemetry middleware backed by `tiktoken-rs` and Turso persistence.
#[derive(Clone)]
pub struct TokenSavingsMiddleware {
    db: ScryerDb,
    bpe: Arc<CoreBPE>,
}

impl TokenSavingsMiddleware {
    /// Initialize the middleware using `cl100k_base` tokenizer.
    pub fn new(db: ScryerDb) -> anyhow::Result<Self> {
        let bpe = tiktoken_rs::cl100k_base()?;
        Ok(Self {
            db,
            bpe: Arc::new(bpe),
        })
    }

    /// Access the underlying tokenizer.
    pub fn bpe(&self) -> &CoreBPE {
        &self.bpe
    }

    /// Count exact tokens for a text string using BPE.
    pub fn count_tokens(&self, text: &str) -> usize {
        self.bpe.encode_with_special_tokens(text).len()
    }

    /// Enforce the payload token limit.
    /// - If `no_truncate` is true, bypasses truncation regardless of token count.
    /// - If `max_tokens` is provided and > 0, enforces that limit; if 0, truncation is disabled.
    /// - If `max_tokens` is None, defaults to `MAX_PAYLOAD_TOKENS` (1,000).
    pub fn enforce_limit(
        &self,
        payload: &str,
        max_tokens: Option<usize>,
        no_truncate: bool,
    ) -> (String, usize, bool) {
        let count = self.count_tokens(payload);
        if no_truncate {
            return (payload.to_string(), count, false);
        }

        let budget = max_tokens.unwrap_or(MAX_PAYLOAD_TOKENS);
        if budget == 0 || count <= budget {
            return (payload.to_string(), count, false);
        }

        // Truncate line by line to keep under limit
        let mut truncated = String::new();
        let budget_str = if budget == 1000 {
            "1,000".to_string()
        } else {
            budget.to_string()
        };
        let notice = format!(
            "\n\n... [Output truncated to enforce hard {budget_str} token limit. Pass 'no_truncate: true', increase 'max_tokens', or use pagination/filtering (offset, limit, query).]"
        );
        let notice_tokens = self.count_tokens(&notice);
        let token_budget = budget.saturating_sub(notice_tokens);

        for line in payload.lines() {
            let next = if truncated.is_empty() {
                line.to_string()
            } else {
                format!("{truncated}\n{line}")
            };

            if self.count_tokens(&next) > token_budget {
                break;
            }
            truncated = next;
        }

        truncated.push_str(&notice);
        let actual = self.count_tokens(&truncated);
        (truncated, actual, true)
    }

    /// Backward-compatible limit enforcement using default 1,000 token limit.
    pub fn enforce_default_limit(&self, payload: &str) -> (String, usize, bool) {
        self.enforce_limit(payload, None, false)
    }

    /// Calculate estimated naive cost based on Phase 4 specification.
    ///
    /// - Outline / Scope: `File Bytes / 3.8`
    /// - Definition / Type: `max(File Bytes / 3.8, 1500)`
    /// - References / Hierarchy: `min(Matches, 5) * max(File Bytes / 3.8, 1200)`
    pub fn calculate_naive_cost(
        category: ToolCategory,
        file_bytes: Option<usize>,
        match_count: Option<usize>,
    ) -> u64 {
        let bytes = file_bytes.unwrap_or(0) as f64;
        let naive_bytes_tokens = bytes / 3.8;

        match category {
            ToolCategory::OutlineOrScope => {
                if file_bytes.is_some() {
                    naive_bytes_tokens.round() as u64
                } else {
                    1000
                }
            }
            ToolCategory::DefinitionOrType => naive_bytes_tokens.max(1500.0).round() as u64,
            ToolCategory::ReferencesOrHierarchy => {
                let matches = (match_count.unwrap_or(1).min(5)) as f64;
                let per_file = naive_bytes_tokens.max(1200.0);
                (matches * per_file).round() as u64
            }
            // Docs lookup + snippet read + a 1-hop call trace.
            ToolCategory::InspectSymbol => 6000,
            ToolCategory::CrateOutline => 8000,
            ToolCategory::SearchSymbols => {
                let matches = (match_count.unwrap_or(1).min(5)) as f64;
                (matches * 1200.0).max(3500.0).round() as u64
            }
            ToolCategory::Default => 1000,
        }
    }

    /// Record a tool invocation and asynchronously persist to Turso `ToolInvocationMetric`.
    pub fn record_invocation(
        &self,
        session_id: String,
        project_id: Option<u64>,
        tool_name: String,
        payload_tokens: u64,
        estimated_naive_tokens: u64,
        duration: Instant,
    ) {
        let duration_ms = duration.elapsed().as_millis() as u32;
        let tokens_saved = estimated_naive_tokens.saturating_sub(payload_tokens);
        let db = self.db.clone();
        let now = scryer_db::time::now_rfc3339();

        tokio::spawn(async move {
            let mut guard = db.lock().await;
            let result = ToolInvocationMetric::create()
                .session_id(session_id)
                .project_id(project_id)
                .tool_name(tool_name)
                .payload_tokens(payload_tokens)
                .estimated_naive_tokens(estimated_naive_tokens)
                .tokens_saved(tokens_saved)
                .execution_duration_ms(duration_ms)
                .created_at(now)
                .exec(&mut *guard)
                .await;

            if let Err(e) = result {
                error!("Failed to persist ToolInvocationMetric: {e}");
            } else {
                debug!(tokens_saved, "Persisted ToolInvocationMetric successfully");
            }
        });
    }

    /// Aggregate token savings metrics across session or all sessions.
    pub async fn get_metrics(
        &self,
        session_id: Option<&str>,
    ) -> anyhow::Result<TokenSavingsSummary> {
        let mut guard = self.db.lock().await;
        let metrics = match session_id {
            Some(sid) => {
                ToolInvocationMetric::filter(ToolInvocationMetric::fields().session_id().eq(sid))
                    .exec(&mut *guard)
                    .await?
            }
            None => ToolInvocationMetric::all().exec(&mut *guard).await?,
        };

        let total_invocations = metrics.len() as u64;
        let mut total_payload_tokens = 0u64;
        let mut total_naive_tokens = 0u64;
        let mut total_tokens_saved = 0u64;
        let mut total_duration_ms = 0u64;
        let mut by_tool: BTreeMap<String, u64> = BTreeMap::new();

        for m in &metrics {
            *by_tool.entry(m.tool_name.clone()).or_default() += 1;
            total_payload_tokens += m.payload_tokens;
            total_naive_tokens += m.estimated_naive_tokens;
            total_tokens_saved += m.tokens_saved;
            total_duration_ms += m.execution_duration_ms as u64;
        }

        let average_duration_ms = if total_invocations > 0 {
            total_duration_ms as f64 / total_invocations as f64
        } else {
            0.0
        };

        let billable_cost_reduction_dollars =
            (total_tokens_saved as f64 / 1_000_000.0) * COST_PER_MILLION_TOKENS;

        Ok(TokenSavingsSummary {
            total_invocations,
            total_payload_tokens,
            total_estimated_naive_tokens: total_naive_tokens,
            total_tokens_saved,
            average_duration_ms,
            billable_cost_reduction_dollars,
            by_tool,
        })
    }
}

#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, schemars::JsonSchema)]
pub struct TokenSavingsSummary {
    pub total_invocations: u64,
    pub total_payload_tokens: u64,
    pub total_estimated_naive_tokens: u64,
    pub total_tokens_saved: u64,
    pub average_duration_ms: f64,
    pub billable_cost_reduction_dollars: f64,
    /// Invocation count per tool name, for measuring which tools agents actually use.
    #[serde(default)]
    pub by_tool: BTreeMap<String, u64>,
}