use crate::record::kind::Kind;
use crate::record::payloads::UsageOutcomeCounts;
use crate::record::record::{Record, ThreadId};
use crate::record::refs::RecordId;
use crate::state::state::State;
use crate::verify::rules::VerifierRules;
use std::collections::{BTreeMap, BTreeSet};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ContextPolicy {
#[default]
ExcludeTainted,
IncludeTainted,
}
#[derive(Debug, Clone)]
pub struct Context {
pub records: Vec<Record>,
pub usage_feedback: BTreeMap<(RecordId, String), UsageOutcomeCounts>,
pub tainted_records: BTreeSet<RecordId>,
}
pub fn build_context(
records: &[Record],
state: &State,
rules: &VerifierRules,
thread: ThreadId,
) -> Context {
build_context_with(records, state, rules, thread, ContextPolicy::default())
}
pub fn build_context_with(
records: &[Record],
state: &State,
rules: &VerifierRules,
thread: ThreadId,
policy: ContextPolicy,
) -> Context {
let include_tainted = policy == ContextPolicy::IncludeTainted;
let mut selected: Vec<&Record> = records
.iter()
.filter(|r| {
r.thread == thread
&& r.kind != Kind::Verdict
&& state.accepted_records.contains(&r.id)
&& !state.replaced_records.contains(&r.id)
&& !state.retracted_records.contains(&r.id)
&& (include_tainted || !state.tainted_records.contains(&r.id))
})
.collect();
selected.sort_by(|a, b| b.time.cmp(&a.time).then_with(|| a.id.cmp(&b.id)));
selected.truncate(rules.max_context_records as usize);
let context_records: Vec<Record> = selected.into_iter().cloned().collect();
let context_ids: BTreeSet<RecordId> = context_records.iter().map(|r| r.id).collect();
let usage_feedback: BTreeMap<(RecordId, String), UsageOutcomeCounts> = state
.usage_counts
.iter()
.filter(|((used_record, _role), _counts)| context_ids.contains(used_record))
.map(|(k, v)| (k.clone(), v.clone()))
.collect();
let tainted_records: BTreeSet<RecordId> = context_records
.iter()
.filter(|r| state.tainted_records.contains(&r.id))
.map(|r| r.id)
.collect();
Context {
records: context_records,
usage_feedback,
tainted_records,
}
}