use futures::stream::StreamExt;
use rusqlite::params;
use tracing::{info, warn};
use crate::classify::classifier::ClassificationEngine;
use crate::classify::errors::Result;
use crate::classify::tiers::llm_prompt::{LlmOutcome, LlmUsage};
use crate::classify::tiers::ClassificationResult;
use crate::core::config::LlmFallbackScope;
use crate::core::db::Database;
use super::pipeline_db::CommitRow;
pub(crate) fn is_unanswered(r: &ClassificationResult) -> bool {
r.category == "uncategorized" || r.subcategory.as_deref() == Some("uncategorized")
}
pub(crate) fn llm_eligible(
scope: LlmFallbackScope,
r: &ClassificationResult,
threshold: f64,
) -> bool {
match scope {
LlmFallbackScope::LowConfidence => r.confidence <= threshold,
LlmFallbackScope::Unanswered => is_unanswered(r),
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct LlmUsageTotals {
pub calls: usize,
pub adopted: usize,
pub not_adopted: usize,
pub abstained: usize,
pub out_of_set: usize,
pub failed: usize,
pub calls_with_usage: usize,
pub input_tokens: u64,
pub output_tokens: u64,
}
impl LlmUsageTotals {
fn add(&mut self, outcome: &str, usage: Option<LlmUsage>) {
self.calls += 1;
match outcome {
ADOPTED => self.adopted += 1,
NOT_ADOPTED => self.not_adopted += 1,
o if o == LlmOutcome::Abstained.as_str() => self.abstained += 1,
o if o == LlmOutcome::OutOfSet.as_str() => self.out_of_set += 1,
_ => self.failed += 1,
}
if let Some(u) = usage {
self.calls_with_usage += 1;
self.input_tokens += u.input_tokens;
self.output_tokens += u.output_tokens;
}
}
}
const ADOPTED: &str = "adopted";
const NOT_ADOPTED: &str = "not_adopted";
pub(super) struct UsageRow {
idx: usize,
outcome: &'static str,
usage: Option<LlmUsage>,
}
pub(super) async fn run_llm_fallback(
engine: &ClassificationEngine,
commits: &[CommitRow],
results: &mut [ClassificationResult],
scope: LlmFallbackScope,
threshold: f64,
concurrency: usize,
) -> (LlmUsageTotals, Vec<UsageRow>) {
let mut skipped_merges = 0_usize;
let mut pending: Vec<(usize, &str)> = Vec::new();
for (idx, c) in commits.iter().enumerate() {
if !llm_eligible(scope, &results[idx], threshold) {
continue;
}
if c.is_merge {
skipped_merges += 1;
} else {
pending.push((idx, c.message.as_str()));
}
}
info!(
pending = pending.len(),
skipped_merges,
scope = ?scope,
"LLM fallback selection"
);
let pb = super::pipeline_db::make_progress(pending.len() as u64, "LLM fallback");
let pb_ref = &pb;
let calls: Vec<_> =
futures::stream::iter(pending.into_iter().map(|(idx, message)| async move {
let call = engine.llm_classify_detailed(message).await;
pb_ref.inc(1);
(idx, call)
}))
.buffer_unordered(concurrency.max(1))
.collect()
.await;
pb.finish_and_clear();
let mut totals = LlmUsageTotals::default();
let mut rows = Vec::with_capacity(calls.len());
for (idx, call) in calls {
let Some(call) = call else { continue };
let outcome = match call.verdict {
Some(r) if r.confidence > results[idx].confidence => {
results[idx] = r;
ADOPTED
}
Some(r) => {
warn!(
commit_idx = idx,
original_conf = results[idx].confidence,
new_conf = r.confidence,
"LLM fallback did not improve confidence; keeping original verdict"
);
NOT_ADOPTED
}
None => call.outcome.as_str(),
};
totals.add(outcome, call.usage);
rows.push(UsageRow {
idx,
outcome,
usage: call.usage,
});
}
(totals, rows)
}
pub(super) fn record_usage(
db: &mut Database,
commits: &[CommitRow],
rows: &[UsageRow],
identity: (&str, &str),
run_started_at: &str,
) -> Result<()> {
if rows.is_empty() {
return Ok(());
}
let tx = db
.connection_mut()
.transaction()
.map_err(crate::core::TgaError::from)?;
{
let mut stmt = tx
.prepare(
"INSERT INTO llm_usage (commit_id, commit_sha, provider, model, outcome, \
input_tokens, output_tokens, run_started_at) \
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)",
)
.map_err(crate::core::TgaError::from)?;
for row in rows {
let commit = &commits[row.idx];
stmt.execute(params![
commit.id,
commit.sha,
identity.0,
identity.1,
row.outcome,
row.usage.map(|u| u.input_tokens as i64),
row.usage.map(|u| u.output_tokens as i64),
run_started_at,
])
.map_err(crate::core::TgaError::from)?;
}
}
tx.commit().map_err(crate::core::TgaError::from)?;
Ok(())
}