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            Some(UsageScope::Unspecified),
113        ]
114        .into_iter()
115        .enumerate()
116        {
117            let turn = MaterializedTurnOutcome {
118                diagnostic: Some(mj_core::diagnostic::TurnDiagnostic {
119                    message: "Usage limit exceeded".into(),
120                    code: Some("provider.auth_error".into()),
121                    http_status: Some(403),
122                    reset_at: None,
123                }),
124                command_id: format!("p{i}"),
125                accepted_ordinal: None,
126                turn_start_position: Some(i as u64 + 1),
127                completed_ordinal: i as u64 + 1,
128                completed_at_ms: 100,
129                outcome: TurnOutcomeKind::Completed {
130                    stop_reason: "EndTurn".into(),
131                },
132                usage: scope.map(|scope| TokenUsage {
133                    provider_details: Some(Box::new(mj_core::usage::ProviderTurnUsage {
134                        cost: Some(mj_core::usage::ProviderTurnCost::from_usd_ticks(
135                            88_767_200, false,
136                        )),
137                        model_calls: Some(2),
138                        api_duration_ms: Some(1700),
139                        elapsed_ms: Some(2377),
140                        model_usage: BTreeMap::new(),
141                    })),
142                    scope,
143                    total_tokens: 30,
144                    input_tokens: 20,
145                    output_tokens: 10,
146                    thought_tokens: None,
147                    cached_read_tokens: Some(0),
148                    cached_write_tokens: None,
149                }),
150            };
151            turns.push(turn);
152        }
153        for _ in 0..2 {
154            apply_projection_page_to(&path, "usage", |page| {
155                for turn in &turns {
156                    let prior = if turn.completed_ordinal == 1 {
157                        mj_core::relay::RELAY_EVENT_GENESIS_DIGEST.into()
158                    } else {
159                        format!("{:064x}", turn.completed_ordinal - 1)
160                    };
161                    page.apply(
162                        turn.completed_ordinal,
163                        &prior,
164                        &format!("{:064x}", turn.completed_ordinal),
165                        &MaterializedSessionMutation {
166                            last_turn_outcome: Some(turn.clone()),
167                            provider_cost: Some(ProviderCost {
168                                amount: 1.25,
169                                currency: "USD".into(),
170                                observed_at_ms: 100,
171                            }),
172                            ..Default::default()
173                        },
174                    )?;
175                }
176                Ok(())
177            })?;
178        }
179        let first = load_session_usage_from(&path, "usage", 0, 2)?.unwrap();
180        assert_eq!(first.turns.len(), 2);
181        let details = first.turns[0]
182            .usage
183            .as_ref()
184            .unwrap()
185            .provider_details
186            .as_ref()
187            .unwrap();
188        assert_eq!(details.cost.as_ref().unwrap().usd, "0.0088767200");
189        assert_eq!(details.elapsed_ms, Some(2377));
190        assert_eq!(
191            first.turns[0].diagnostic.as_ref().unwrap().http_status,
192            Some(403)
193        );
194        assert_eq!(first.provider_session_cost.as_ref().unwrap().amount, 1.25);
195        assert_eq!(first.coverage.recorded_turns, 5);
196        assert_eq!(first.coverage.full_turn_reports, 2);
197        assert_eq!(first.coverage.last_request_reports, 1);
198        assert_eq!(first.coverage.unspecified_reports, 1);
199        assert_eq!(first.coverage.missing_reports, 1);
200        // Only the two whole-turn reports are summed: the last-request and
201        // unspecified turns stay out of the totals.
202        assert_eq!(first.totals["total_tokens"].tokens, 60);
203        assert_eq!(first.totals["total_tokens"].reported_turns, 2);
204        assert_eq!(first.totals["cached_read_tokens"].tokens, 0);
205        assert!(!first.totals.contains_key("thought_tokens"));
206        let next = load_session_usage_from(&path, "usage", first.next_after_seq, 2)?.unwrap();
207        assert_eq!(next.turns.len(), 2);
208        assert_eq!(next.latest_seq, 5);
209        assert_eq!(next.totals, first.totals);
210        let last = load_session_usage_from(&path, "usage", next.next_after_seq, 2)?.unwrap();
211        assert_eq!(last.turns.len(), 1);
212        assert_eq!(
213            last.turns[0].usage.as_ref().unwrap().scope,
214            UsageScope::Unspecified
215        );
216        assert_eq!(last.next_after_seq, last.latest_seq);
217        assert_eq!(last.totals, first.totals);
218        Ok(())
219    }
220}