Skip to main content

mj_controller/database/
usage.rs

1use super::*;
2#[cfg(test)]
3use mj_core::usage::ProviderCost;
4
5pub fn load_session_usage(
6    session_id: &str,
7    after_seq: u64,
8    limit: usize,
9) -> Result<Option<UsagePage>> {
10    load_session_usage_from(&database_path(), session_id, after_seq, limit)
11}
12
13fn load_session_usage_from(
14    path: &Path,
15    session_id: &str,
16    after_seq: u64,
17    limit: usize,
18) -> Result<Option<UsagePage>> {
19    let mut connection = open_reader(path)?;
20    let tx = connection.transaction()?;
21    if read_materialized_session_fields(&tx, session_id)?.is_none() {
22        return Ok(None);
23    }
24    let latest_seq: u64 = tx.query_row(
25        "SELECT COALESCE(MAX(completed_ordinal), 0) FROM session_turn_usage WHERE session_id = ?1",
26        [session_id],
27        |r| r.get(0),
28    )?;
29    let mut statement = tx.prepare("SELECT body FROM session_turn_usage WHERE session_id = ?1 AND completed_ordinal > ?2 ORDER BY completed_ordinal LIMIT ?3")?;
30    let turns = statement
31        .query_map(
32            params![session_id, after_seq, limit.clamp(1, 1000) as i64],
33            |r| r.get::<_, String>(0),
34        )?
35        .map(|r| Ok(serde_json::from_str::<MaterializedTurnOutcome>(&r?)?))
36        .collect::<Result<Vec<_>>>()?;
37    let mut coverage = UsageCoverage::default();
38    let mut statement = tx.prepare("SELECT json_extract(body, '$.usage.scope'), COUNT(*) FROM session_turn_usage WHERE session_id = ?1 GROUP BY json_extract(body, '$.usage.scope')")?;
39    for row in statement.query_map([session_id], |r| {
40        Ok((r.get::<_, Option<String>>(0)?, r.get::<_, u64>(1)?))
41    })? {
42        let (scope, count) = row?;
43        coverage.recorded_turns += count;
44        match scope.as_deref() {
45            Some("turn") => coverage.full_turn_reports += count,
46            Some("last_request") => coverage.last_request_reports += count,
47            Some(_) => coverage.unspecified_reports += count,
48            None => coverage.missing_reports += count,
49        }
50    }
51    let mut totals = BTreeMap::new();
52    for counter in [
53        "total_tokens",
54        "input_tokens",
55        "output_tokens",
56        "thought_tokens",
57        "cached_read_tokens",
58        "cached_write_tokens",
59    ] {
60        let (tokens, reported_turns) = tx.query_row("SELECT SUM(json_extract(body, ?2)), COUNT(json_extract(body, ?2)) FROM session_turn_usage WHERE session_id = ?1 AND json_extract(body, '$.usage.scope') = 'turn'", params![session_id, format!("$.usage.{counter}")], |r| Ok((r.get::<_, Option<u64>>(0)?, r.get::<_, u64>(1)?)))?;
61        if let Some(tokens) = tokens {
62            totals.insert(
63                counter.into(),
64                UsageCounterTotal {
65                    tokens,
66                    reported_turns,
67                },
68            );
69        }
70    }
71    let cost: Option<String> = tx
72        .query_row(
73            "SELECT body FROM session_provider_cost WHERE session_id = ?1",
74            [session_id],
75            |r| r.get(0),
76        )
77        .optional()?;
78    Ok(Some(UsagePage {
79        session_id: session_id.into(),
80        next_after_seq: turns
81            .last()
82            .map_or(latest_seq.max(after_seq), |t| t.completed_ordinal),
83        latest_seq,
84        turns,
85        totals,
86        coverage,
87        provider_session_cost: cost.map(|s| serde_json::from_str(&s)).transpose()?,
88    }))
89}
90
91#[cfg(test)]
92mod tests {
93    use super::*;
94    use mj_core::state::TurnOutcomeKind;
95    use mj_core::usage::{TokenUsage, UsageScope};
96
97    #[test]
98    fn usage_survives_reopening_and_replay_with_honest_coverage() -> Result<()> {
99        let dir = tempfile::tempdir()?;
100        let path = dir.path().join("usage.sqlite");
101        let conn = open(&path)?;
102        // Use the same projection setup as the materialized database tests.
103        drop(conn);
104        save_session_to(&path, &super::super::tests::session("usage", "project"))?;
105
106        let mut turns = Vec::new();
107        for (i, scope) in [
108            Some(UsageScope::Turn),
109            Some(UsageScope::Turn),
110            Some(UsageScope::LastRequest),
111            None,
112        ]
113        .into_iter()
114        .enumerate()
115        {
116            let turn = MaterializedTurnOutcome {
117                diagnostic: Some(mj_core::diagnostic::TurnDiagnostic {
118                    message: "Usage limit exceeded".into(),
119                    code: Some("provider.auth_error".into()),
120                    http_status: Some(403),
121                    reset_at: None,
122                }),
123                command_id: format!("p{i}"),
124                accepted_ordinal: None,
125                turn_start_position: Some(i as u64 + 1),
126                completed_ordinal: i as u64 + 1,
127                completed_at_ms: 100,
128                outcome: TurnOutcomeKind::Completed {
129                    stop_reason: "EndTurn".into(),
130                },
131                usage: scope.map(|scope| TokenUsage {
132                    provider_details: Some(Box::new(mj_core::usage::ProviderTurnUsage {
133                        cost: Some(mj_core::usage::ProviderTurnCost::from_usd_ticks(
134                            88_767_200, false,
135                        )),
136                        model_calls: Some(2),
137                        api_duration_ms: Some(1700),
138                        elapsed_ms: Some(2377),
139                        model_usage: BTreeMap::new(),
140                    })),
141                    scope,
142                    total_tokens: 30,
143                    input_tokens: 20,
144                    output_tokens: 10,
145                    thought_tokens: None,
146                    cached_read_tokens: Some(0),
147                    cached_write_tokens: None,
148                }),
149            };
150            turns.push(turn);
151        }
152        for _ in 0..2 {
153            apply_projection_page_to(&path, "usage", |page| {
154                for turn in &turns {
155                    let prior = if turn.completed_ordinal == 1 {
156                        mj_core::relay::RELAY_EVENT_GENESIS_DIGEST.into()
157                    } else {
158                        format!("{:064x}", turn.completed_ordinal - 1)
159                    };
160                    page.apply(
161                        turn.completed_ordinal,
162                        &prior,
163                        &format!("{:064x}", turn.completed_ordinal),
164                        &MaterializedSessionMutation {
165                            last_turn_outcome: Some(turn.clone()),
166                            provider_cost: Some(ProviderCost {
167                                amount: 1.25,
168                                currency: "USD".into(),
169                                observed_at_ms: 100,
170                            }),
171                            ..Default::default()
172                        },
173                    )?;
174                }
175                Ok(())
176            })?;
177        }
178        let first = load_session_usage_from(&path, "usage", 0, 2)?.unwrap();
179        assert_eq!(first.turns.len(), 2);
180        let details = first.turns[0]
181            .usage
182            .as_ref()
183            .unwrap()
184            .provider_details
185            .as_ref()
186            .unwrap();
187        assert_eq!(details.cost.as_ref().unwrap().usd, "0.0088767200");
188        assert_eq!(details.elapsed_ms, Some(2377));
189        assert_eq!(
190            first.turns[0].diagnostic.as_ref().unwrap().http_status,
191            Some(403)
192        );
193        assert_eq!(first.provider_session_cost.as_ref().unwrap().amount, 1.25);
194        assert_eq!(first.coverage.recorded_turns, 4);
195        assert_eq!(first.coverage.full_turn_reports, 2);
196        assert_eq!(first.coverage.last_request_reports, 1);
197        assert_eq!(first.coverage.missing_reports, 1);
198        assert_eq!(first.totals["total_tokens"].tokens, 60);
199        assert_eq!(first.totals["cached_read_tokens"].tokens, 0);
200        assert!(!first.totals.contains_key("thought_tokens"));
201        let next = load_session_usage_from(&path, "usage", first.next_after_seq, 2)?.unwrap();
202        assert_eq!(next.turns.len(), 2);
203        assert_eq!(next.next_after_seq, next.latest_seq);
204        assert_eq!(next.totals, first.totals);
205        Ok(())
206    }
207}