Skip to main content

kimetsu_brain/
roi.rs

1//! ROI ledger — v1.5 story "kimetsu pays for itself".
2//!
3//! Estimates the token savings that Kimetsu's memory brain delivered
4//! by surfacing relevant knowledge before a coding session, so the model
5//! didn't have to (re-)discover it through expensive exploration.
6//!
7//! # Assumption-based estimates
8//!
9//! Savings constants are nominal assumptions, not measured counterfactuals,
10//! calibrated guarantees, or lower bounds. Delivered cost observations retain
11//! their producer units separately; byte bounds are not model token counts.
12
13use kimetsu_core::{KimetsuResult, memory::MemoryKind};
14use rusqlite::{OptionalExtension, params};
15use serde::Serialize;
16
17// ---------------------------------------------------------------------------
18// S2.4(b): Output-token accounting
19// ---------------------------------------------------------------------------
20
21/// Assumed ratio of output tokens to input tokens for a typical coding
22/// assistant response. This ratio has not been calibrated against sessions.
23/// We use 0.25 as a nominal assumption
24/// and expose it in the public report.
25///
26/// **Audited limitation**: this is a ratio-based *estimate* because Claude Code
27/// does not expose per-session output token counts to the Stop hook.  The
28/// estimate will be off for short responses (low ratio) and long code-gen runs
29/// (higher ratio).  We document this in the `--json` `output_token_estimate`
30/// field with an `"estimate_method": "ratio_0.25"` annotation.
31pub const OUTPUT_TOKEN_INPUT_RATIO: f64 = 0.25;
32
33/// Estimate output tokens from the brain-injected input token count.
34///
35/// This is a conservative proxy for sessions where the model is guided by
36/// brain context — more relevant context → fewer wasted generation tokens.
37/// See [`OUTPUT_TOKEN_INPUT_RATIO`] for calibration notes.
38pub fn estimate_output_tokens(input_tokens: u64) -> u64 {
39    (input_tokens as f64 * OUTPUT_TOKEN_INPUT_RATIO).round() as u64
40}
41
42// ---------------------------------------------------------------------------
43// S2.4(c): New event kind savings constants
44// ---------------------------------------------------------------------------
45
46/// Conservative token savings per `digest_served` event.
47///
48/// Assumption: a digest saves the model from re-reading the CLAUDE.md +
49/// searching for the top conventions at session start.  Estimated equivalent:
50/// ~2 search calls × 600 tokens/call = ~1 200 tokens.  We claim 800 as a
51/// nominal assumption.
52pub const SAVED_TOKENS_PER_DIGEST_SERVED: u64 = 800;
53
54/// Conservative token savings per `resume_served` event.
55///
56/// Assumption: an episodic resume avoids the model asking "what were you
57/// working on?" + 1–2 file reads to reconstruct context.  Estimated
58/// equivalent: ~2 tool calls × 400 tokens/call = ~800 tokens.  We claim 500.
59pub const SAVED_TOKENS_PER_RESUME_SERVED: u64 = 500;
60
61/// Conservative token savings per `skill.served` event (future-proof).
62///
63/// Assumption: a synthesized skill file avoids the model re-deriving the
64/// composite procedure from individual memories.  We claim 300 as a
65/// nominal assumption.
66pub const SAVED_TOKENS_PER_SKILL_SERVED: u64 = 300;
67
68// ---------------------------------------------------------------------------
69// Per-kind nominal assumptions
70// ---------------------------------------------------------------------------
71
72/// Assumed estimate of tokens saved per citation, by memory
73/// kind.  These are deliberate *under*-estimates of the exploration cost the
74/// model would have incurred without the brain context.
75///
76/// Assumed methodology (see <https://kimetsu.dev/docs/roi-methodology/> for details):
77/// - `failure_pattern`: avoids the "try → fail → diagnose → fix" loop.
78///   Typical loop: ~3 tool calls × ~500 tokens/call = ~1 500 tokens.
79/// - `command`: avoids a web/docs lookup or `--help` trial.  ~1–2 tool
80///   calls = ~400 tokens.
81/// - `convention`: avoids a code-search to find the project pattern.
82///   ~1–2 searches = ~300 tokens.
83/// - `fact`: avoids asking the user or searching docs. ~1 exchange = ~500 t.
84/// - `preference`: avoids one clarifying question. ~1 exchange = ~200 t.
85///
86/// These constants are the source of truth for the methodology doc.
87pub const SAVED_TOKENS_PER_CITATION: &[(MemoryKind, u32)] = &[
88    (MemoryKind::FailurePattern, 1500),
89    (MemoryKind::Command, 400),
90    (MemoryKind::Convention, 300),
91    (MemoryKind::Fact, 500),
92    (MemoryKind::Preference, 200),
93];
94
95// ---------------------------------------------------------------------------
96// Built-in price table (input tokens, conservative single number per family)
97// ---------------------------------------------------------------------------
98
99/// Conservative input-token price in USD per million tokens ($/MTok) for
100/// known model families.  These are approximate and marked as such in the
101/// `--json` output (`usd` field carries them as estimates).
102///
103/// Matched against the project's `[model] model` config value by prefix
104/// (longest match wins).  Unknown model → `usd: None` unless
105/// `[model] price_per_mtok` is set in `project.toml`.
106///
107/// Last updated: 2026-06, approximate retail/API-key pricing.
108const BUILTIN_PRICE_TABLE: &[(&str, f64)] = &[
109    // Anthropic Claude 4 family (Opus > Sonnet > Haiku)
110    ("claude-opus-4", 15.00),
111    ("claude-sonnet-4", 3.00),
112    ("claude-haiku-4", 0.80),
113    // Anthropic Claude 3 family
114    ("claude-3-opus", 15.00),
115    ("claude-3-5-sonnet", 3.00),
116    ("claude-3-5-haiku", 0.80),
117    ("claude-3-sonnet", 3.00),
118    ("claude-3-haiku", 0.25),
119    // Anthropic Bedrock cross-region routing prefixes
120    ("us.anthropic.claude-opus-4", 15.00),
121    ("us.anthropic.claude-sonnet-4", 3.00),
122    ("us.anthropic.claude-haiku-4", 0.80),
123    // OpenAI gpt-5 family
124    ("gpt-5", 2.00),
125    ("gpt-4o", 2.50),
126    ("gpt-4-turbo", 10.00),
127    ("gpt-4", 30.00),
128];
129
130/// Resolve a $/MTok price for the given model id.
131///
132/// Precedence: `price_override` (from `[model] price_per_mtok`) >
133/// longest-prefix match in [`BUILTIN_PRICE_TABLE`].
134/// Returns `None` when neither applies.
135pub fn resolve_price_per_mtok(model: &str, price_override: Option<f64>) -> Option<f64> {
136    if let Some(p) = price_override {
137        return Some(p);
138    }
139    let model_lower = model.to_lowercase();
140    // Longest prefix match wins.
141    let mut best: Option<(&str, f64)> = None;
142    for (prefix, price) in BUILTIN_PRICE_TABLE {
143        if model_lower.starts_with(prefix) && best.is_none_or(|(b, _)| prefix.len() > b.len()) {
144            best = Some((prefix, *price));
145        }
146    }
147    best.map(|(_, p)| p)
148}
149
150// ---------------------------------------------------------------------------
151// Pure savings estimator
152// ---------------------------------------------------------------------------
153
154/// Estimate total tokens saved from a citation summary.
155///
156/// `citations` is a slice of `(kind, count)` pairs — how many times each
157/// memory kind was cited in the window.  The function is intentionally pure
158/// (no I/O) so it can be unit-tested without a DB.
159///
160/// The result is a assumption-based estimate: if a kind has no entry in
161/// [`SAVED_TOKENS_PER_CITATION`] it contributes 0 (fail-safe).
162pub fn estimate_savings(citations: &[(MemoryKind, u32)]) -> u64 {
163    citations
164        .iter()
165        .map(|(kind, count)| {
166            let per = SAVED_TOKENS_PER_CITATION
167                .iter()
168                .find(|(k, _)| k == kind)
169                .map(|(_, v)| *v as u64)
170                .unwrap_or(0);
171            per * (*count as u64)
172        })
173        .sum()
174}
175
176// ---------------------------------------------------------------------------
177// Report types
178// ---------------------------------------------------------------------------
179
180/// USD sub-report — only present when a price is resolvable.
181#[derive(Debug, Clone, Serialize)]
182pub struct RoiUsd {
183    /// Estimated USD value of the tokens saved by the brain.
184    pub saved: f64,
185    /// Estimated USD cost of the brain overhead (injected tokens consumed
186    /// at inference time).  Uses the same $/MTok price as `saved`.
187    pub spent: f64,
188    /// `saved − spent`.  Can be negative when overhead exceeds the
189    /// estimated savings.
190    pub net: f64,
191}
192
193/// Full ROI report for a time window.
194#[derive(Debug, Clone, Serialize)]
195pub struct RoiReport {
196    pub estimate_label: &'static str,
197    pub model: String,
198    pub assumptions: serde_json::Value,
199    /// Observed producer costs grouped by units. Not summed as measured tokens.
200    pub delivered_cost_by_unit: std::collections::BTreeMap<String, u64>,
201    /// Window length in days, or `None` for "all time".
202    pub window_days: Option<u32>,
203    /// Total tokens injected by the brain (sum of `used_tokens` from
204    /// `context.injected` events in the window).
205    pub injected_tokens: u64,
206    /// S2.4(b): Estimated output tokens generated in the window.
207    ///
208    /// Computed as `injected_tokens × OUTPUT_TOKEN_INPUT_RATIO`.
209    /// **Audited limitation**: ratio-based estimate; Claude Code does not
210    /// expose per-session output token counts.
211    pub estimated_output_tokens: u64,
212    /// Number of `context.served` events in the window.
213    pub served_events: u64,
214    /// S2.4(c): Number of `digest_served` events in the window.
215    pub digest_served_events: u64,
216    /// S2.4(c): Number of `resume_served` events in the window.
217    pub resume_served_events: u64,
218    /// S2.4(c): Tokens saved from warm-start digests and resumes.
219    pub warmstart_saved_tokens: u64,
220    /// Total citation count (rows in `memory_citations` for runs in the
221    /// window).
222    pub citations: u64,
223    /// Estimated tokens saved (assumption-based estimate).
224    pub estimated_saved_tokens: u64,
225    /// `estimated_saved_tokens − injected_tokens`.  Can be negative.
226    pub net_tokens: i64,
227    /// USD sub-report; `None` when price is unknown.
228    pub usd: Option<RoiUsd>,
229}
230
231// ---------------------------------------------------------------------------
232// S2.4(a): Per-memory ROI
233// ---------------------------------------------------------------------------
234
235/// Per-memory ROI entry for `kimetsu brain roi --top`.
236#[derive(Debug, Clone, Serialize)]
237pub struct MemoryRoiEntry {
238    pub memory_id: String,
239    pub kind: String,
240    /// First ~80 chars of the memory text (for human readability).
241    pub text_head: String,
242    /// Total number of times this memory has been cited in the window.
243    pub citation_count: u64,
244    /// Estimated tokens saved by this memory's citations.
245    pub estimated_saved_tokens: u64,
246}
247
248/// Compute per-memory ROI for the top `limit` memories by estimated savings.
249///
250/// Only memories with ≥1 citation in the window are returned.
251pub fn per_memory_roi(
252    conn: &rusqlite::Connection,
253    window: RoiWindow,
254    limit: usize,
255) -> KimetsuResult<Vec<MemoryRoiEntry>> {
256    let window_since: Option<String> = match window {
257        RoiWindow::All => None,
258        RoiWindow::Days(days) => {
259            let secs = days as i64 * 86_400;
260            let now = time::OffsetDateTime::now_utc();
261            let cutoff = now - time::Duration::seconds(secs);
262            let fmt = time::format_description::well_known::Rfc3339;
263            Some(cutoff.format(&fmt).unwrap_or_default())
264        }
265    };
266
267    // Collect (memory_id, citation_count) pairs.
268    struct Row {
269        memory_id: String,
270        count: u64,
271    }
272    let rows: Vec<Row> = match &window_since {
273        Some(ts) => {
274            let mut stmt = conn.prepare(
275                "SELECT mc.memory_id, COUNT(*) \
276                 FROM memory_citations mc \
277                 LEFT JOIN runs r ON mc.run_id = r.run_id \
278                 WHERE r.started_at >= ?1 \
279                    OR (r.run_id IS NULL AND mc.cited_at >= ?1) \
280                 GROUP BY mc.memory_id \
281                 ORDER BY COUNT(*) DESC",
282            )?;
283            let rows = stmt.query_map(params![ts], |r| {
284                Ok(Row {
285                    memory_id: r.get(0)?,
286                    count: r.get(1)?,
287                })
288            })?;
289            rows.collect::<Result<Vec<_>, _>>()?
290        }
291        None => {
292            let mut stmt = conn.prepare(
293                "SELECT memory_id, COUNT(*) FROM memory_citations \
294                 GROUP BY memory_id ORDER BY COUNT(*) DESC",
295            )?;
296            let rows = stmt.query_map([], |r| {
297                Ok(Row {
298                    memory_id: r.get(0)?,
299                    count: r.get(1)?,
300                })
301            })?;
302            rows.collect::<Result<Vec<_>, _>>()?
303        }
304    };
305
306    let mut entries: Vec<MemoryRoiEntry> = Vec::new();
307    for row in rows.into_iter().take(limit) {
308        // Resolve kind and text from the memories table.
309        let memory_row: Option<(String, String)> = conn
310            .query_row(
311                "SELECT kind, text FROM memories WHERE memory_id = ?1",
312                params![row.memory_id],
313                |r| Ok((r.get(0)?, r.get(1)?)),
314            )
315            .optional()?;
316        let (kind_str, text) = memory_row.unwrap_or_else(|| ("fact".to_string(), String::new()));
317        let mk = kind_str.parse::<MemoryKind>().unwrap_or(MemoryKind::Fact);
318        let per_cite = SAVED_TOKENS_PER_CITATION
319            .iter()
320            .find(|(k, _)| k == &mk)
321            .map(|(_, v)| *v as u64)
322            .unwrap_or(0);
323        let estimated_saved = per_cite * row.count;
324        let text_head: String = text.chars().take(80).collect();
325
326        entries.push(MemoryRoiEntry {
327            memory_id: row.memory_id,
328            kind: kind_str,
329            text_head,
330            citation_count: row.count,
331            estimated_saved_tokens: estimated_saved,
332        });
333    }
334
335    // Sort descending by estimated_saved_tokens.
336    entries.sort_by_key(|e| std::cmp::Reverse(e.estimated_saved_tokens));
337    Ok(entries)
338}
339
340// ---------------------------------------------------------------------------
341// Window parsing
342// ---------------------------------------------------------------------------
343
344/// Recognised window strings → days.  Mirrors the CLI arg values.
345#[derive(Debug, Clone, Copy, PartialEq, Eq)]
346pub enum RoiWindow {
347    Days(u32),
348    All,
349}
350
351impl RoiWindow {
352    /// Parse "7d", "30d", or "all".
353    pub fn parse(s: &str) -> Result<Self, String> {
354        match s.trim().to_lowercase().as_str() {
355            "all" => Ok(Self::All),
356            other => {
357                let digits = other.trim_end_matches('d');
358                digits
359                    .parse::<u32>()
360                    .map(Self::Days)
361                    .map_err(|_| format!("invalid window '{s}'; expected '7d', '30d', or 'all'"))
362            }
363        }
364    }
365
366    pub fn days(self) -> Option<u32> {
367        match self {
368            Self::Days(d) => Some(d),
369            Self::All => None,
370        }
371    }
372}
373
374impl Default for RoiWindow {
375    fn default() -> Self {
376        Self::Days(30)
377    }
378}
379
380// ---------------------------------------------------------------------------
381// DB-backed report
382// ---------------------------------------------------------------------------
383
384/// Compute a full ROI report from the project brain.
385///
386/// `window` controls how far back to look.  `price_per_mtok_override`
387/// comes from `[model] price_per_mtok` in `project.toml`; `model_name` is
388/// `[model] model`.
389pub fn roi_report(
390    conn: &rusqlite::Connection,
391    window: RoiWindow,
392    model_name: &str,
393    price_per_mtok_override: Option<f64>,
394) -> KimetsuResult<RoiReport> {
395    // Compute window boundary timestamp (ISO-8601 string).
396    let window_since: Option<String> = match window {
397        RoiWindow::All => None,
398        RoiWindow::Days(days) => {
399            // Compute `now − days` as an ISO string using the `time` crate
400            // that is already a workspace dep.
401            let secs = days as i64 * 86_400;
402            let now = time::OffsetDateTime::now_utc();
403            let cutoff = now - time::Duration::seconds(secs);
404            let fmt = time::format_description::well_known::Rfc3339;
405            Some(cutoff.format(&fmt).unwrap_or_default())
406        }
407    };
408
409    // --- served_events (context.served) ---
410    let served_events: u64 = match &window_since {
411        Some(ts) => conn.query_row(
412            "SELECT COUNT(*) FROM events WHERE kind = 'context.served' AND ts >= ?1",
413            params![ts],
414            |r| r.get(0),
415        )?,
416        None => conn.query_row(
417            "SELECT COUNT(*) FROM events WHERE kind = 'context.served'",
418            [],
419            |r| r.get(0),
420        )?,
421    };
422
423    // S2.4(c): digest_served and resume_served event counts.
424    let digest_served_events: u64 = match &window_since {
425        Some(ts) => conn.query_row(
426            "SELECT COUNT(*) FROM events WHERE kind = 'digest_served' AND ts >= ?1",
427            params![ts],
428            |r| r.get(0),
429        )?,
430        None => conn.query_row(
431            "SELECT COUNT(*) FROM events WHERE kind = 'digest_served'",
432            [],
433            |r| r.get(0),
434        )?,
435    };
436    let resume_served_events: u64 = match &window_since {
437        Some(ts) => conn.query_row(
438            "SELECT COUNT(*) FROM events WHERE kind = 'resume_served' AND ts >= ?1",
439            params![ts],
440            |r| r.get(0),
441        )?,
442        None => conn.query_row(
443            "SELECT COUNT(*) FROM events WHERE kind = 'resume_served'",
444            [],
445            |r| r.get(0),
446        )?,
447    };
448    let warmstart_saved_tokens = digest_served_events * SAVED_TOKENS_PER_DIGEST_SERVED
449        + resume_served_events * SAVED_TOKENS_PER_RESUME_SERVED;
450
451    // --- injected_tokens (sum of used_tokens across context.injected events) ---
452    let mut delivered_cost_by_unit = std::collections::BTreeMap::<String, u64>::new();
453    let injected_tokens: u64 = {
454        let payloads: Vec<String> = match &window_since {
455            Some(ts) => {
456                let mut stmt = conn.prepare(
457                    "SELECT payload_json FROM events WHERE kind = 'context.injected' AND ts >= ?1",
458                )?;
459                let rows = stmt.query_map(params![ts], |r| r.get::<_, String>(0))?;
460                rows.collect::<Result<Vec<_>, _>>()?
461            }
462            None => {
463                let mut stmt = conn
464                    .prepare("SELECT payload_json FROM events WHERE kind = 'context.injected'")?;
465                let rows = stmt.query_map([], |r| r.get::<_, String>(0))?;
466                rows.collect::<Result<Vec<_>, _>>()?
467            }
468        };
469        let mut sum: u64 = 0;
470        for p in &payloads {
471            let v: serde_json::Value = serde_json::from_str(p)?;
472            if let Some(t) = v.get("used_tokens").and_then(|x| x.as_u64()) {
473                *delivered_cost_by_unit
474                    .entry(
475                        v.get("cost_unit")
476                            .and_then(|x| x.as_str())
477                            .unwrap_or("legacy_token_estimate")
478                            .to_string(),
479                    )
480                    .or_default() += t;
481                sum += t;
482            }
483        }
484        sum
485    };
486
487    // --- citations per kind ---
488    // We join memory_citations → memories to get the kind for each citation.
489    // Window applied via run_id → runs.started_at for pipeline runs, or
490    // via the event ts for hook-originated citations.
491    //
492    // Strategy: collect citation rows that belong to runs in the window
493    // (started_at >= window_since) OR, for the sentinel hook run_id
494    // (all-zeroes), filter by the cited_at timestamp.
495    let citations_by_kind: Vec<(MemoryKind, u32)> = {
496        // Collect all (memory_id, count) pairs from citations in window.
497        struct Row {
498            memory_id: String,
499            count: u32,
500        }
501
502        let rows: Vec<Row> = match &window_since {
503            Some(ts) => {
504                let mut stmt = conn.prepare(
505                    "SELECT mc.memory_id, COUNT(*) \
506                     FROM memory_citations mc \
507                     LEFT JOIN runs r ON mc.run_id = r.run_id \
508                     WHERE r.started_at >= ?1 \
509                        OR (r.run_id IS NULL AND mc.cited_at >= ?1) \
510                     GROUP BY mc.memory_id",
511                )?;
512                let rows = stmt.query_map(params![ts], |r| {
513                    Ok(Row {
514                        memory_id: r.get(0)?,
515                        count: r.get(1)?,
516                    })
517                })?;
518                rows.collect::<Result<Vec<_>, _>>()?
519            }
520            None => {
521                let mut stmt = conn.prepare(
522                    "SELECT memory_id, COUNT(*) FROM memory_citations GROUP BY memory_id",
523                )?;
524                let rows = stmt.query_map([], |r| {
525                    Ok(Row {
526                        memory_id: r.get(0)?,
527                        count: r.get(1)?,
528                    })
529                })?;
530                rows.collect::<Result<Vec<_>, _>>()?
531            }
532        };
533
534        // For each memory_id, resolve its kind from the memories table.
535        // Unknown / invalidated memories default to Fact (conservative).
536        let mut by_kind: std::collections::HashMap<MemoryKind, u32> =
537            std::collections::HashMap::new();
538        for row in &rows {
539            let kind: Option<String> = conn
540                .query_row(
541                    "SELECT kind FROM memories WHERE memory_id = ?1",
542                    params![row.memory_id],
543                    |r| r.get(0),
544                )
545                .optional()?;
546            let mk = kind
547                .as_deref()
548                .and_then(|s| s.parse::<MemoryKind>().ok())
549                .unwrap_or(MemoryKind::Fact);
550            *by_kind.entry(mk).or_insert(0) += row.count;
551        }
552        by_kind.into_iter().collect()
553    };
554
555    let total_citations: u64 = citations_by_kind.iter().map(|(_, c)| *c as u64).sum();
556    let citation_saved_tokens = estimate_savings(&citations_by_kind);
557    // S2.4(c): include warm-start savings in the total estimate.
558    let estimated_saved_tokens = citation_saved_tokens + warmstart_saved_tokens;
559    let net_tokens = estimated_saved_tokens as i64 - injected_tokens as i64;
560    // S2.4(b): output token estimate.
561    let estimated_output_tokens = estimate_output_tokens(injected_tokens);
562
563    // --- USD ---
564    let price = resolve_price_per_mtok(model_name, price_per_mtok_override);
565    let usd = price.map(|p_per_mtok| {
566        let saved_usd = estimated_saved_tokens as f64 / 1_000_000.0 * p_per_mtok;
567        let spent_usd = injected_tokens as f64 / 1_000_000.0 * p_per_mtok;
568        RoiUsd {
569            saved: saved_usd,
570            spent: spent_usd,
571            net: saved_usd - spent_usd,
572        }
573    });
574
575    Ok(RoiReport {
576        estimate_label: "Assumption-based estimate; savings are not measured or guaranteed",
577        model: model_name.to_string(),
578        assumptions: serde_json::json!({"tokens_per_citation":SAVED_TOKENS_PER_CITATION.iter().map(|(kind,n)|(kind.to_string(),*n)).collect::<std::collections::BTreeMap<_,_>>(),"digest":SAVED_TOKENS_PER_DIGEST_SERVED,"resume":SAVED_TOKENS_PER_RESUME_SERVED,"output_input_ratio":OUTPUT_TOKEN_INPUT_RATIO,"price_per_mtok":price,"overhead":"legacy estimate combines producer costs; see delivered_cost_by_unit for observed units"}),
579        delivered_cost_by_unit,
580        window_days: window.days(),
581        injected_tokens,
582        estimated_output_tokens,
583        served_events,
584        digest_served_events,
585        resume_served_events,
586        warmstart_saved_tokens,
587        citations: total_citations,
588        estimated_saved_tokens,
589        net_tokens,
590        usd,
591    })
592}
593
594// ---------------------------------------------------------------------------
595// Per-session mini-report (for the Stop hook)
596// ---------------------------------------------------------------------------
597
598/// Compute a per-session ROI mini-report for the Stop hook.
599///
600/// Attribution strategy (conservative):
601/// 1. `context.served` events with matching `session_id` in payload → used
602///    for `served_events` and `injected_tokens`.
603/// 2. Citations: we cannot directly attribute `memory_citations` rows to a
604///    session_id (citations are keyed by run_id, not session_id).  Instead
605///    we fall back to a time-window bounded by the earliest and latest
606///    `context.served` event timestamps for this session.  If session_id
607///    is absent (old hook payload), the time window covers the last 24 hours
608///    as a rough proxy.
609/// 3. ZERO citations → returns `None` (silence; no savings line emitted).
610///
611/// All errors are swallowed and `None` returned — the hook must never fail.
612pub fn session_roi(
613    conn: &rusqlite::Connection,
614    session_id: Option<&str>,
615    model_name: &str,
616    price_per_mtok_override: Option<f64>,
617) -> Option<SessionRoi> {
618    session_roi_inner(conn, session_id, model_name, price_per_mtok_override).unwrap_or(None)
619}
620
621fn session_roi_inner(
622    conn: &rusqlite::Connection,
623    session_id: Option<&str>,
624    model_name: &str,
625    price_per_mtok_override: Option<f64>,
626) -> KimetsuResult<Option<SessionRoi>> {
627    // 1. Find context.served events for this session.
628    let (served_events, injected_tokens, earliest_ts, latest_ts) =
629        session_served_stats(conn, session_id)?;
630
631    // 2. Determine citation time window.
632    let (ts_lo, ts_hi) = match (earliest_ts.as_deref(), latest_ts.as_deref()) {
633        (Some(lo), Some(hi)) => (lo.to_string(), hi.to_string()),
634        _ => {
635            // Fall back to last 24h.
636            let now = time::OffsetDateTime::now_utc();
637            let fmt = time::format_description::well_known::Rfc3339;
638            let lo = (now - time::Duration::seconds(86_400))
639                .format(&fmt)
640                .unwrap_or_default();
641            let hi = now.format(&fmt).unwrap_or_default();
642            (lo, hi)
643        }
644    };
645
646    // 3. Collect citations in the time window.
647    let citations_by_kind = citations_in_window(conn, &ts_lo, &ts_hi)?;
648    let total_citations: u64 = citations_by_kind.iter().map(|(_, c)| *c as u64).sum();
649
650    // Silence when nothing was cited this session.
651    if total_citations == 0 {
652        return Ok(None);
653    }
654
655    let estimated_saved_tokens = estimate_savings(&citations_by_kind);
656    let net_tokens = estimated_saved_tokens as i64 - injected_tokens as i64;
657
658    let price = resolve_price_per_mtok(model_name, price_per_mtok_override);
659    let usd = price.map(|p_per_mtok| {
660        let saved_usd = estimated_saved_tokens as f64 / 1_000_000.0 * p_per_mtok;
661        let spent_usd = injected_tokens as f64 / 1_000_000.0 * p_per_mtok;
662        RoiUsd {
663            saved: saved_usd,
664            spent: spent_usd,
665            net: saved_usd - spent_usd,
666        }
667    });
668
669    Ok(Some(SessionRoi {
670        served_events,
671        injected_tokens,
672        citations: total_citations,
673        estimated_saved_tokens,
674        net_tokens,
675        usd,
676    }))
677}
678
679/// Lightweight per-session ROI summary used by the Stop hook.
680#[derive(Debug, Clone)]
681pub struct SessionRoi {
682    pub served_events: u64,
683    pub injected_tokens: u64,
684    pub citations: u64,
685    pub estimated_saved_tokens: u64,
686    pub net_tokens: i64,
687    pub usd: Option<RoiUsd>,
688}
689
690impl SessionRoi {
691    /// Build a one-line savings sentence for the Stop hook `systemMessage`.
692    /// Returns a human-readable string like:
693    ///   "[Kimetsu] Estimated savings (nominal assumptions): ~1 200 tokens (~$0.004) this session."
694    pub fn savings_sentence(&self) -> String {
695        match &self.usd {
696            Some(u) if u.net >= 0.0 => format!(
697                "[Kimetsu] Estimated savings (nominal assumptions): ~{} tokens (~${:.4}) this session.",
698                format_tokens(self.estimated_saved_tokens),
699                u.saved,
700            ),
701            Some(u) => format!(
702                "[Kimetsu] Estimated overhead (nominal assumptions): ~{} tokens (net −${:.4}) this session.",
703                format_tokens(self.injected_tokens),
704                u.spent - u.saved,
705            ),
706            None => format!(
707                "[Kimetsu] Estimated savings (nominal assumptions): ~{} tokens this session.",
708                format_tokens(self.estimated_saved_tokens),
709            ),
710        }
711    }
712}
713
714fn format_tokens(n: u64) -> String {
715    // Human-friendly: thousands separator via simple manual formatting.
716    if n < 1_000 {
717        return n.to_string();
718    }
719    let s = n.to_string();
720    let mut out = String::new();
721    let rem = s.len() % 3;
722    for (i, ch) in s.chars().enumerate() {
723        if i > 0 && (i % 3 == rem) {
724            out.push(' ');
725        }
726        out.push(ch);
727    }
728    out
729}
730
731// ---------------------------------------------------------------------------
732// Internal helpers
733// ---------------------------------------------------------------------------
734
735/// Returns (served_events, injected_tokens, earliest_ts, latest_ts) for a
736/// given session_id.  When session_id is None the query returns all events.
737fn session_served_stats(
738    conn: &rusqlite::Connection,
739    session_id: Option<&str>,
740) -> KimetsuResult<(u64, u64, Option<String>, Option<String>)> {
741    // context.served events for this session.
742    let served_payloads: Vec<String> = match session_id {
743        Some(sid) => {
744            let mut stmt = conn.prepare(
745                "SELECT payload_json FROM events \
746                 WHERE kind = 'context.served' \
747                   AND json_extract(payload_json, '$.session_id') = ?1",
748            )?;
749            let rows = stmt.query_map(params![sid], |r| r.get::<_, String>(0))?;
750            rows.collect::<Result<Vec<_>, _>>()?
751        }
752        None => {
753            // No session_id available — return empty; caller falls back to 24h window.
754            return Ok((0, 0, None, None));
755        }
756    };
757
758    // Also collect context.injected in the same session window by time.
759    // Derive earliest/latest ts from the served events first.
760    let mut earliest: Option<String> = None;
761    let mut latest: Option<String> = None;
762    let served_count = served_payloads.len() as u64;
763
764    // Parse timestamps from the events table for the served events.
765    if let Some(sid) = session_id {
766        if served_count > 0 {
767            let ts_row: (Option<String>, Option<String>) = conn.query_row(
768                "SELECT MIN(ts), MAX(ts) FROM events \
769                 WHERE kind = 'context.served' \
770                   AND json_extract(payload_json, '$.session_id') = ?1",
771                params![sid],
772                |r| Ok((r.get(0)?, r.get(1)?)),
773            )?;
774            earliest = ts_row.0;
775            latest = ts_row.1;
776        }
777    }
778
779    // Injected tokens: sum from context.injected events in [earliest, latest].
780    let injected_tokens: u64 = match (earliest.as_deref(), latest.as_deref()) {
781        (Some(lo), Some(hi)) => {
782            let payloads: Vec<String> = {
783                let mut stmt = conn.prepare(
784                    "SELECT payload_json FROM events \
785                     WHERE kind = 'context.injected' AND ts >= ?1 AND ts <= ?2",
786                )?;
787                let rows = stmt.query_map(params![lo, hi], |r| r.get::<_, String>(0))?;
788                rows.collect::<Result<Vec<_>, _>>()?
789            };
790            let mut sum: u64 = 0;
791            for p in &payloads {
792                let v: serde_json::Value = serde_json::from_str(p)?;
793                if let Some(t) = v.get("used_tokens").and_then(|x| x.as_u64()) {
794                    sum += t;
795                }
796            }
797            sum
798        }
799        _ => 0,
800    };
801
802    Ok((served_count, injected_tokens, earliest, latest))
803}
804
805/// Collect (MemoryKind, count) citation pairs for citations whose `cited_at`
806/// falls in `[ts_lo, ts_hi]`.
807fn citations_in_window(
808    conn: &rusqlite::Connection,
809    ts_lo: &str,
810    ts_hi: &str,
811) -> KimetsuResult<Vec<(MemoryKind, u32)>> {
812    struct Row {
813        memory_id: String,
814        count: u32,
815    }
816    let mut stmt = conn.prepare(
817        "SELECT memory_id, COUNT(*) FROM memory_citations \
818         WHERE cited_at >= ?1 AND cited_at <= ?2 \
819         GROUP BY memory_id",
820    )?;
821    let rows = stmt.query_map(params![ts_lo, ts_hi], |r| {
822        Ok(Row {
823            memory_id: r.get(0)?,
824            count: r.get(1)?,
825        })
826    })?;
827    let rows: Vec<Row> = rows.collect::<Result<Vec<_>, _>>()?;
828
829    let mut by_kind: std::collections::HashMap<MemoryKind, u32> = std::collections::HashMap::new();
830    for row in &rows {
831        let kind: Option<String> = conn
832            .query_row(
833                "SELECT kind FROM memories WHERE memory_id = ?1",
834                params![row.memory_id],
835                |r| r.get(0),
836            )
837            .optional()?;
838        let mk = kind
839            .as_deref()
840            .and_then(|s| s.parse::<MemoryKind>().ok())
841            .unwrap_or(MemoryKind::Fact);
842        *by_kind.entry(mk).or_insert(0) += row.count;
843    }
844    Ok(by_kind.into_iter().collect())
845}
846
847// ---------------------------------------------------------------------------
848// Tests
849// ---------------------------------------------------------------------------
850
851#[cfg(test)]
852mod tests {
853    use super::*;
854    use kimetsu_core::memory::MemoryKind;
855
856    // --- Pure function tests ---
857
858    #[test]
859    fn estimate_savings_zero_when_empty() {
860        assert_eq!(estimate_savings(&[]), 0);
861    }
862
863    #[test]
864    fn estimate_savings_single_kind() {
865        // 2 failure_pattern citations × 1500 = 3000
866        assert_eq!(estimate_savings(&[(MemoryKind::FailurePattern, 2)]), 3_000);
867    }
868
869    #[test]
870    fn estimate_savings_multi_kind() {
871        let citations = vec![
872            (MemoryKind::FailurePattern, 1), // 1500
873            (MemoryKind::Command, 2),        // 800
874            (MemoryKind::Convention, 1),     // 300
875            (MemoryKind::Fact, 1),           // 500
876            (MemoryKind::Preference, 3),     // 600
877        ];
878        assert_eq!(estimate_savings(&citations), 1500 + 800 + 300 + 500 + 600);
879    }
880
881    #[test]
882    fn estimate_savings_all_kinds_covered() {
883        // Every kind must appear in SAVED_TOKENS_PER_CITATION.
884        for kind in [
885            MemoryKind::FailurePattern,
886            MemoryKind::Command,
887            MemoryKind::Convention,
888            MemoryKind::Fact,
889            MemoryKind::Preference,
890        ] {
891            let v = SAVED_TOKENS_PER_CITATION
892                .iter()
893                .find(|(k, _)| k == &kind)
894                .map(|(_, v)| *v);
895            assert!(
896                v.is_some(),
897                "kind {:?} missing from SAVED_TOKENS_PER_CITATION",
898                kind
899            );
900            assert!(v.unwrap() > 0, "kind {:?} has zero constant", kind);
901        }
902    }
903
904    #[test]
905    fn resolve_price_override_wins() {
906        assert_eq!(
907            resolve_price_per_mtok("claude-sonnet-4-7", Some(5.0)),
908            Some(5.0)
909        );
910    }
911
912    #[test]
913    fn resolve_price_known_model() {
914        let p = resolve_price_per_mtok("claude-sonnet-4-7", None);
915        assert!(p.is_some(), "claude-sonnet-4 should match");
916        assert!((p.unwrap() - 3.0).abs() < 1e-9);
917    }
918
919    #[test]
920    fn resolve_price_unknown_model_none() {
921        assert!(resolve_price_per_mtok("my-custom-llm-v9", None).is_none());
922    }
923
924    #[test]
925    fn resolve_price_longest_prefix_wins() {
926        // "claude-opus-4" and "claude-opus-4" — make sure haiku doesn't
927        // match opus prefix.
928        let opus_p = resolve_price_per_mtok("claude-opus-4-5", None).unwrap_or(0.0);
929        let haiku_p = resolve_price_per_mtok("claude-haiku-4-5", None).unwrap_or(0.0);
930        assert!(opus_p > haiku_p, "opus should be more expensive than haiku");
931    }
932
933    #[test]
934    fn roi_window_parse() {
935        assert_eq!(RoiWindow::parse("7d").unwrap(), RoiWindow::Days(7));
936        assert_eq!(RoiWindow::parse("30d").unwrap(), RoiWindow::Days(30));
937        assert_eq!(RoiWindow::parse("all").unwrap(), RoiWindow::All);
938        assert_eq!(RoiWindow::parse("ALL").unwrap(), RoiWindow::All);
939        assert!(RoiWindow::parse("bad").is_err());
940    }
941
942    #[test]
943    fn format_tokens_below_1000() {
944        assert_eq!(format_tokens(42), "42");
945        assert_eq!(format_tokens(999), "999");
946    }
947
948    #[test]
949    fn format_tokens_thousands() {
950        assert_eq!(format_tokens(1_000), "1 000");
951        assert_eq!(format_tokens(12_345), "12 345");
952        assert_eq!(format_tokens(1_234_567), "1 234 567");
953    }
954
955    // --- DB-backed tests ---
956
957    use crate::{
958        project::{init_project, load_project},
959        projector,
960        user_brain::with_user_brain_disabled,
961    };
962    use kimetsu_core::{event::Event, ids::RunId, memory::MemoryScope};
963    use ulid::Ulid;
964
965    fn test_root() -> std::path::PathBuf {
966        let root = std::env::temp_dir().join(format!("kimetsu-roi-test-{}", Ulid::new()));
967        kimetsu_core::paths::git_init_boundary(&root);
968        root
969    }
970
971    fn seed_memory(root: &std::path::Path, kind: MemoryKind, text: &str) -> String {
972        crate::project::add_memory(root, MemoryScope::Project, kind, text).expect("add_memory")
973    }
974
975    fn seed_injected_event(conn: &rusqlite::Connection, run_id: RunId, used_tokens: u64) {
976        let ev = Event::new(
977            run_id,
978            "context.injected",
979            serde_json::json!({
980                "stage": "localization",
981                "memory_ids": [],
982                "used_tokens": used_tokens,
983                "capsule_count": 1,
984            }),
985        );
986        projector::apply_events(conn, &[ev]).expect("seed injected");
987    }
988
989    fn seed_citation(conn: &rusqlite::Connection, run_id: RunId, memory_id: &str, turn: i64) {
990        let ev = Event::new(
991            run_id,
992            "memory.cited",
993            serde_json::json!({
994                "memory_id": memory_id,
995                "turn": turn,
996            }),
997        );
998        projector::apply_events(conn, &[ev]).expect("seed citation");
999    }
1000
1001    #[test]
1002    fn roi_report_empty_db_returns_zeros() {
1003        with_user_brain_disabled(|| {
1004            let root = test_root();
1005            init_project(&root, false).expect("init");
1006            let (_paths, config, conn) = load_project(&root).expect("load");
1007            let report =
1008                roi_report(&conn, RoiWindow::All, &config.model.model, None).expect("roi_report");
1009            assert_eq!(report.injected_tokens, 0);
1010            assert_eq!(report.citations, 0);
1011            assert_eq!(report.estimated_saved_tokens, 0);
1012            assert_eq!(report.net_tokens, 0);
1013        });
1014    }
1015
1016    #[test]
1017    fn roi_report_with_citations_computes_savings() {
1018        with_user_brain_disabled(|| {
1019            let root = test_root();
1020            init_project(&root, false).expect("init");
1021            let m1 = seed_memory(&root, MemoryKind::FailurePattern, "fp1");
1022            let (_paths, config, conn) = load_project(&root).expect("load");
1023            let run_id = RunId::new();
1024            seed_injected_event(&conn, run_id, 300);
1025            seed_citation(&conn, run_id, &m1, 1);
1026
1027            let report =
1028                roi_report(&conn, RoiWindow::All, &config.model.model, None).expect("roi_report");
1029            // 1 failure_pattern citation = 1500 saved tokens.
1030            assert_eq!(report.estimated_saved_tokens, 1500);
1031            assert_eq!(report.injected_tokens, 300);
1032            assert_eq!(report.net_tokens, 1500 - 300);
1033            assert_eq!(report.citations, 1);
1034        });
1035    }
1036
1037    #[test]
1038    fn roi_report_negative_net_when_overhead_exceeds_savings() {
1039        with_user_brain_disabled(|| {
1040            let root = test_root();
1041            init_project(&root, false).expect("init");
1042            let m1 = seed_memory(&root, MemoryKind::Preference, "pref1");
1043            let (_paths, config, conn) = load_project(&root).expect("load");
1044            let run_id = RunId::new();
1045            // Inject 500 tokens but cite 1 preference (200 saved) → net = -300.
1046            seed_injected_event(&conn, run_id, 500);
1047            seed_citation(&conn, run_id, &m1, 1);
1048
1049            let report =
1050                roi_report(&conn, RoiWindow::All, &config.model.model, None).expect("roi_report");
1051            assert_eq!(report.estimated_saved_tokens, 200);
1052            assert_eq!(report.net_tokens, 200 - 500); // −300
1053        });
1054    }
1055
1056    #[test]
1057    fn roi_report_usd_with_known_model() {
1058        with_user_brain_disabled(|| {
1059            let root = test_root();
1060            init_project(&root, false).expect("init");
1061            let m1 = seed_memory(&root, MemoryKind::Command, "cmd1");
1062            let (_paths, _config, conn) = load_project(&root).expect("load");
1063            let run_id = RunId::new();
1064            seed_injected_event(&conn, run_id, 200);
1065            seed_citation(&conn, run_id, &m1, 1);
1066
1067            // Use a known model directly.
1068            let report =
1069                roi_report(&conn, RoiWindow::All, "claude-sonnet-4-7", None).expect("roi_report");
1070            let usd = report.usd.expect("usd must be Some for known model");
1071            // 400 saved tokens @ $3/MTok = $0.0012
1072            assert!((usd.saved - 400.0 / 1_000_000.0 * 3.0).abs() < 1e-9);
1073            // 200 injected tokens @ $3/MTok
1074            assert!((usd.spent - 200.0 / 1_000_000.0 * 3.0).abs() < 1e-9);
1075            assert!((usd.net - (usd.saved - usd.spent)).abs() < 1e-12);
1076        });
1077    }
1078
1079    #[test]
1080    fn roi_report_usd_with_override() {
1081        with_user_brain_disabled(|| {
1082            let root = test_root();
1083            init_project(&root, false).expect("init");
1084            let m1 = seed_memory(&root, MemoryKind::Fact, "fact1");
1085            let (_paths, _config, conn) = load_project(&root).expect("load");
1086            let run_id = RunId::new();
1087            seed_injected_event(&conn, run_id, 0);
1088            seed_citation(&conn, run_id, &m1, 1);
1089
1090            let report =
1091                roi_report(&conn, RoiWindow::All, "my-custom-llm", Some(10.0)).expect("roi_report");
1092            // 500 saved tokens @ $10/MTok = $0.005
1093            let usd = report.usd.expect("usd with override");
1094            assert!((usd.saved - 500.0 / 1_000_000.0 * 10.0).abs() < 1e-9);
1095        });
1096    }
1097
1098    #[test]
1099    fn roi_report_unknown_model_no_usd() {
1100        with_user_brain_disabled(|| {
1101            let root = test_root();
1102            init_project(&root, false).expect("init");
1103            let (_paths, _config, conn) = load_project(&root).expect("load");
1104
1105            let report = roi_report(&conn, RoiWindow::All, "totally-unknown-llm-xyz", None)
1106                .expect("roi_report");
1107            assert!(report.usd.is_none(), "usd must be None for unknown model");
1108        });
1109    }
1110
1111    // ── S2.4 tests ────────────────────────────────────────────────────────────
1112
1113    fn seed_event(conn: &rusqlite::Connection, kind: &str, payload: serde_json::Value) {
1114        let ev = Event::new(RunId::new(), kind, payload);
1115        projector::apply_events(conn, &[ev]).expect("seed event");
1116    }
1117
1118    #[test]
1119    fn roi_report_output_token_estimate_is_quarter_of_input() {
1120        with_user_brain_disabled(|| {
1121            let root = test_root();
1122            init_project(&root, false).expect("init");
1123            let (_paths, _config, conn) = load_project(&root).expect("load");
1124            let run_id = RunId::new();
1125            seed_injected_event(&conn, run_id, 4_000);
1126
1127            let report =
1128                roi_report(&conn, RoiWindow::All, "claude-sonnet-4", None).expect("roi_report");
1129            // 4000 * 0.25 = 1000
1130            assert_eq!(
1131                report.estimated_output_tokens, 1_000,
1132                "output token estimate must be 0.25 × input"
1133            );
1134        });
1135    }
1136
1137    #[test]
1138    fn roi_report_digest_served_adds_savings() {
1139        with_user_brain_disabled(|| {
1140            let root = test_root();
1141            init_project(&root, false).expect("init");
1142            let (_paths, _config, conn) = load_project(&root).expect("load");
1143
1144            seed_event(
1145                &conn,
1146                "digest_served",
1147                serde_json::json!({"digest_chars": 800, "approx_tokens": 200}),
1148            );
1149            seed_event(
1150                &conn,
1151                "resume_served",
1152                serde_json::json!({"resume_chars": 400, "approx_tokens": 100}),
1153            );
1154
1155            let report =
1156                roi_report(&conn, RoiWindow::All, "unknown-model", None).expect("roi_report");
1157            assert_eq!(report.digest_served_events, 1);
1158            assert_eq!(report.resume_served_events, 1);
1159            let expected_warmstart =
1160                SAVED_TOKENS_PER_DIGEST_SERVED + SAVED_TOKENS_PER_RESUME_SERVED;
1161            assert_eq!(
1162                report.warmstart_saved_tokens, expected_warmstart,
1163                "warmstart_saved_tokens must sum digest+resume"
1164            );
1165            assert_eq!(
1166                report.estimated_saved_tokens, expected_warmstart,
1167                "total savings must include warmstart (no citations here)"
1168            );
1169        });
1170    }
1171
1172    #[test]
1173    fn per_memory_roi_top_entries_sorted_by_savings() {
1174        with_user_brain_disabled(|| {
1175            let root = test_root();
1176            init_project(&root, false).expect("init");
1177            // Add two memories of different kinds.
1178            let fp_id = seed_memory(&root, MemoryKind::FailurePattern, "fp roi test");
1179            let cmd_id = seed_memory(&root, MemoryKind::Command, "cmd roi test");
1180            let (_paths, _config, conn) = load_project(&root).expect("load");
1181
1182            let run_id = RunId::new();
1183            // 1 failure_pattern cite (1500 saved) + 3 command cites (3×400=1200).
1184            seed_citation(&conn, run_id, &fp_id, 1);
1185            seed_citation(&conn, run_id, &cmd_id, 2);
1186            seed_citation(&conn, run_id, &cmd_id, 3);
1187            seed_citation(&conn, run_id, &cmd_id, 4);
1188
1189            let entries = per_memory_roi(&conn, RoiWindow::All, 10).expect("per_memory_roi");
1190            assert!(!entries.is_empty(), "must have entries");
1191            // FailurePattern (1500) > Command×3 (1200) → fp must come first.
1192            assert_eq!(
1193                entries[0].memory_id, fp_id,
1194                "failure_pattern cite must rank first by savings"
1195            );
1196            assert_eq!(entries[0].estimated_saved_tokens, 1500);
1197            assert_eq!(entries[0].citation_count, 1);
1198
1199            let cmd_entry = entries
1200                .iter()
1201                .find(|e| e.memory_id == cmd_id)
1202                .expect("cmd entry");
1203            assert_eq!(cmd_entry.citation_count, 3);
1204            assert_eq!(cmd_entry.estimated_saved_tokens, 1200);
1205
1206            std::fs::remove_dir_all(&root).ok();
1207        });
1208    }
1209
1210    #[test]
1211    fn per_memory_roi_respects_top_limit() {
1212        with_user_brain_disabled(|| {
1213            let root = test_root();
1214            init_project(&root, false).expect("init");
1215            let m1 = seed_memory(&root, MemoryKind::Fact, "fact1");
1216            let m2 = seed_memory(&root, MemoryKind::Fact, "fact2");
1217            let m3 = seed_memory(&root, MemoryKind::Fact, "fact3");
1218            let (_paths, _config, conn) = load_project(&root).expect("load");
1219            let run_id = RunId::new();
1220            seed_citation(&conn, run_id, &m1, 1);
1221            seed_citation(&conn, run_id, &m2, 2);
1222            seed_citation(&conn, run_id, &m3, 3);
1223
1224            let entries = per_memory_roi(&conn, RoiWindow::All, 2).expect("per_memory_roi limit");
1225            assert_eq!(entries.len(), 2, "must respect top limit");
1226            std::fs::remove_dir_all(&root).ok();
1227        });
1228    }
1229
1230    #[test]
1231    fn estimate_output_tokens_quarter_ratio() {
1232        assert_eq!(estimate_output_tokens(4_000), 1_000);
1233        assert_eq!(estimate_output_tokens(0), 0);
1234        assert_eq!(estimate_output_tokens(1_000), 250);
1235    }
1236
1237    #[test]
1238    fn session_roi_returns_none_when_no_citations() {
1239        with_user_brain_disabled(|| {
1240            let root = test_root();
1241            init_project(&root, false).expect("init");
1242            let (_paths, _config, conn) = load_project(&root).expect("load");
1243
1244            let result = session_roi(&conn, Some("sess-abc"), "claude-sonnet-4", None);
1245            assert!(result.is_none(), "no citations → no session roi");
1246        });
1247    }
1248
1249    #[test]
1250    fn savings_sentence_positive_no_usd() {
1251        let sr = SessionRoi {
1252            served_events: 3,
1253            injected_tokens: 100,
1254            citations: 2,
1255            estimated_saved_tokens: 1200,
1256            net_tokens: 1100,
1257            usd: None,
1258        };
1259        let s = sr.savings_sentence();
1260        assert!(s.contains("1 200"), "expected formatted token count");
1261        assert!(s.contains("[Kimetsu]"), "must have brand prefix");
1262    }
1263
1264    #[test]
1265    fn savings_sentence_positive_with_usd() {
1266        let sr = SessionRoi {
1267            served_events: 3,
1268            injected_tokens: 100,
1269            citations: 2,
1270            estimated_saved_tokens: 1500,
1271            net_tokens: 1400,
1272            usd: Some(RoiUsd {
1273                saved: 0.0045,
1274                spent: 0.0003,
1275                net: 0.0042,
1276            }),
1277        };
1278        let s = sr.savings_sentence();
1279        assert!(s.contains("$"), "must include dollar sign when usd present");
1280    }
1281}