use saya_agent::OverrideFindingDto;
use saya_connectors::sql_references;
use saya_types::SqlDialect;
use std::sync::{Arc, Mutex};
use crate::contracts::{OverrideFinding, RecallReceipt, detect_overrides};
use super::DatabaseTools;
const MAX_FINDINGS: usize = 64;
pub(crate) struct OverrideLog {
findings: Mutex<Vec<OverrideFindingDto>>,
}
impl OverrideLog {
pub(crate) fn new() -> Self {
Self {
findings: Mutex::new(Vec::new()),
}
}
pub(crate) fn record(&self, finding: OverrideFindingDto) {
let mut guard = self.findings.lock().expect("override log not poisoned");
if guard.len() < MAX_FINDINGS {
guard.push(finding);
}
}
pub(crate) fn drain(&self) -> Vec<OverrideFindingDto> {
let all = std::mem::take(&mut *self.findings.lock().expect("override log not poisoned"));
dedupe_by_claim(all)
}
}
impl Default for OverrideLog {
fn default() -> Self {
Self::new()
}
}
pub(super) fn override_dto(finding: OverrideFinding) -> OverrideFindingDto {
OverrideFindingDto {
claim_id: finding.claim_id,
kind: finding.kind.to_string(),
claimed_value: finding.claimed_value,
observed_columns: finding.observed_columns,
}
}
impl DatabaseTools {
pub(crate) fn with_recall_receipt(
mut self,
receipt: Option<Arc<RecallReceipt>>,
override_log: Option<Arc<OverrideLog>>,
) -> Self {
self.recall_receipt = receipt;
self.override_log = override_log;
self
}
pub(in crate::agent::tools) fn detect_and_record_overrides(
&self,
sql: &str,
dialect: SqlDialect,
) {
let (Some(receipt), Some(log)) = (&self.recall_receipt, &self.override_log) else {
return;
};
let refs = sql_references(sql, dialect);
for finding in detect_overrides(receipt, refs.as_ref()) {
log.record(override_dto(finding));
}
}
}
fn dedupe_by_claim(findings: Vec<OverrideFindingDto>) -> Vec<OverrideFindingDto> {
let mut seen: Vec<String> = Vec::with_capacity(findings.len());
let mut kept = Vec::with_capacity(findings.len());
for finding in findings {
let key = finding.claim_id.as_str().to_string();
if seen.iter().any(|s| s == &key) {
continue;
}
seen.push(key);
kept.push(finding);
}
kept
}
#[cfg(test)]
mod tests {
use super::*;
fn finding(id: &str, observed: &str) -> OverrideFindingDto {
OverrideFindingDto {
claim_id: ClaimId::parse(id).unwrap(),
kind: "default_time_column".into(),
claimed_value: "return_date".into(),
observed_columns: vec![observed.into()],
}
}
use saya_types::ClaimId;
#[test]
fn drain_returns_findings_in_order_and_clears() {
let log = OverrideLog::new();
log.record(finding("c-a", "rental_date"));
log.record(finding("c-b", "checkout_date"));
let drained = log.drain();
assert_eq!(drained.len(), 2);
assert_eq!(drained[0].claim_id.as_str(), "c-a");
assert_eq!(drained[1].claim_id.as_str(), "c-b");
assert!(log.drain().is_empty());
}
#[test]
fn drain_dedupes_by_claim_id_keeping_the_first() {
let log = OverrideLog::new();
log.record(finding("c-a", "rental_date"));
log.record(finding("c-b", "checkout_date"));
log.record(finding("c-a", "checkout_date"));
let drained = log.drain();
assert_eq!(drained.len(), 2, "one finding per contradicted claim");
assert_eq!(drained[0].claim_id.as_str(), "c-a");
assert_eq!(drained[0].observed_columns, vec!["rental_date".to_string()]);
assert_eq!(drained[1].claim_id.as_str(), "c-b");
}
#[test]
fn the_log_caps_at_max_findings() {
let log = OverrideLog::new();
for i in 0..(MAX_FINDINGS + 5) {
log.record(finding(&format!("c-{i}"), "rental_date"));
}
assert_eq!(log.drain().len(), MAX_FINDINGS, "the log honors its cap");
}
}