use std::collections::{BTreeMap, HashMap, HashSet};
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_context::CommitContext;
use crate::classify::tiers::llm_prompt::{LlmOutcome, LlmUsage};
use crate::classify::tiers::ClassificationResult;
use crate::core::config::{LlmConfig, LlmContextItem, 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 {
if r.method == crate::core::models::ClassificationMethod::RepoMap {
return false;
}
match scope {
LlmFallbackScope::LowConfidence => r.confidence <= threshold,
LlmFallbackScope::Unanswered => is_unanswered(r),
}
}
impl super::pipeline::ClassificationPipeline {
pub(crate) fn llm_enabled(&self) -> bool {
self.config.llm.is_some()
|| self
.config
.classification
.as_ref()
.is_some_and(|c| c.use_llm)
}
}
pub(super) fn load_contexts(
db: &Database,
llm: Option<&LlmConfig>,
commits: &[CommitRow],
results: &[ClassificationResult],
scope: LlmFallbackScope,
threshold: f64,
) -> Result<BTreeMap<usize, CommitContext>> {
use crate::core::db::commit_context::{load_issue_types, load_paths, load_pr_titles};
let Some(cfg) = llm.filter(|c| !c.context.is_empty()) else {
return Ok(BTreeMap::new());
};
let wants = |item| cfg.context.contains(&item);
let sent: Vec<usize> = (0..commits.len())
.filter(|&i| !commits[i].is_merge && llm_eligible(scope, &results[i], threshold))
.collect();
let conn = db.connection();
let ids: Vec<i64> = sent.iter().map(|&i| commits[i].id).collect();
let shas: HashSet<&str> = sent.iter().map(|&i| commits[i].sha.as_str()).collect();
let db_err = crate::core::TgaError::from;
let mut paths = HashMap::new();
let mut prs = HashMap::new();
let mut issues = HashMap::new();
if wants(LlmContextItem::Paths) {
paths = load_paths(conn, &ids).map_err(db_err)?;
}
if wants(LlmContextItem::PrTitle) {
prs = load_pr_titles(conn, &shas).map_err(db_err)?;
}
if wants(LlmContextItem::IssueType) {
issues = load_issue_types(conn, &shas).map_err(db_err)?;
}
let contexts: BTreeMap<usize, CommitContext> = sent
.into_iter()
.filter_map(|i| {
let c = &commits[i];
let ctx = CommitContext::new(
paths.get(&c.id).map_or(&[][..], Vec::as_slice),
cfg.context_max_paths,
cfg.context_max_path_bytes,
prs.get(&c.sha).map(String::as_str),
issues.get(&c.sha).map(String::as_str),
);
(!ctx.is_empty()).then_some((i, ctx))
})
.collect();
info!(
with_context = contexts.len(),
items = ?cfg.context,
"LLM context loaded"
);
Ok(contexts)
}
#[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 skipped: usize,
pub calls_with_usage: usize,
pub input_tokens: u64,
pub output_tokens: u64,
pub skipped_shas: Vec<String>,
}
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,
o if o == LlmOutcome::Skipped.as_str() => self.skipped += 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>,
model: Option<String>,
text_mode: Option<&'static str>,
}
pub(super) async fn run_llm_fallback(
engine: &ClassificationEngine,
commits: &[CommitRow],
contexts: &BTreeMap<usize, CommitContext>,
results: &mut [ClassificationResult],
scope: LlmFallbackScope,
threshold: f64,
concurrency: usize,
) -> (LlmUsageTotals, Vec<UsageRow>) {
let mut skipped_merges = 0_usize;
let mut pending: Vec<(usize, &str, Option<&CommitContext>)> = 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(), contexts.get(&idx)));
}
}
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, ctx)| async move {
let call = engine.llm_classify_with_context(message, ctx).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,
model: call.model,
text_mode: call.text_mode,
});
}
rows.sort_by_key(|r| r.idx);
totals.skipped_shas = rows
.iter()
.filter(|r| r.outcome == LlmOutcome::Skipped.as_str())
.map(|r| commits[r.idx].sha.clone())
.collect();
(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, text_mode) \
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)",
)
.map_err(crate::core::TgaError::from)?;
for row in rows {
let commit = &commits[row.idx];
stmt.execute(params![
commit.id,
commit.sha,
identity.0,
row.model.as_deref().unwrap_or(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,
row.text_mode,
])
.map_err(crate::core::TgaError::from)?;
}
}
tx.commit().map_err(crate::core::TgaError::from)?;
Ok(())
}