use std::sync::{Arc, Mutex, RwLock};
use std::time::{SystemTime, UNIX_EPOCH};
use crate::trace::{TraceEnvelope, TraceEvent, TraceSink};
use crate::usage::{Capability, IntentGraph};
struct Pending {
query: String,
}
pub struct UsageLearner {
inner: Arc<dyn TraceSink>,
graph: Arc<RwLock<IntentGraph>>,
pending: Mutex<Option<Pending>>,
}
impl UsageLearner {
pub fn new(graph: Arc<RwLock<IntentGraph>>, inner: Arc<dyn TraceSink>) -> Self {
Self {
inner,
graph,
pending: Mutex::new(None),
}
}
pub fn graph(&self) -> Arc<RwLock<IntentGraph>> {
self.graph.clone()
}
fn remember_query(&self, query: &str) {
if let Ok(mut pending) = self.pending.lock() {
*pending = Some(Pending {
query: query.to_string(),
});
}
if let Ok(graph) = self.graph.read() {
graph.arm_credit(query);
}
}
fn confirm(&self, kind: Capability, capability_id: &str, ts_ms: u64) {
let Ok(pending) = self.pending.lock() else {
return;
};
let Some(query) = pending.as_ref().map(|p| p.query.clone()) else {
return; };
drop(pending);
if let Ok(mut graph) = self.graph.write() {
let first_confirmation = graph.claim_credit(&query);
graph.observe(&query, kind, capability_id, ts_ms, first_confirmation);
}
}
pub fn replay(&self, envelope: &TraceEnvelope) {
self.learn_from(&envelope.event, envelope.ts);
}
fn learn_from(&self, event: &TraceEvent, ts_ms: u64) {
match event {
TraceEvent::Search { query, .. } | TraceEvent::SkillSearch { query, .. } => {
self.remember_query(query)
}
TraceEvent::InvokeStart { tool_id, .. } => {
self.confirm(Capability::Tool, tool_id, ts_ms)
}
TraceEvent::SkillInvoke { skill_id, .. } => {
self.confirm(Capability::Skill, skill_id, ts_ms)
}
_ => {}
}
}
}
impl TraceSink for UsageLearner {
fn record(&self, event: TraceEvent) {
self.learn_from(&event, now_ms());
self.inner.record(event);
}
fn sample_rate(&self) -> f64 {
self.inner.sample_rate()
}
}
fn now_ms() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or(0)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::trace::{MemorySink, NoopSink, Origin};
fn learner() -> (Arc<UsageLearner>, Arc<RwLock<IntentGraph>>) {
let graph = Arc::new(RwLock::new(IntentGraph::empty()));
let l = Arc::new(UsageLearner::new(graph.clone(), Arc::new(NoopSink)));
(l, graph)
}
fn search(query: &str) -> TraceEvent {
TraceEvent::Search {
query: query.into(),
origin: Origin::Agent,
top_k: 5,
hits: Vec::new(),
stages: Vec::new(),
took_ms: 0,
}
}
fn invoke(tool_id: &str) -> TraceEvent {
TraceEvent::InvokeStart {
tool_id: tool_id.into(),
args_size_bytes: 0,
}
}
#[test]
fn a_search_then_invoke_becomes_one_observation() {
let (l, graph) = learner();
l.record(search("why is the build broken"));
l.record(invoke("gh_run_list"));
let g = graph.read().unwrap();
assert_eq!(g.len(), 1);
assert_eq!(g.intents[0].support, 1);
assert_eq!(g.intents[0].tools.get("gh_run_list"), Some(&1.0));
}
#[test]
fn a_search_nobody_acts_on_teaches_nothing() {
let (l, graph) = learner();
l.record(search("why is the build broken"));
assert!(graph.read().unwrap().is_empty());
}
#[test]
fn an_invoke_with_no_preceding_search_teaches_nothing() {
let (l, graph) = learner();
l.record(invoke("gh_run_list"));
assert!(graph.read().unwrap().is_empty());
}
#[test]
fn what_retrieval_returned_never_becomes_an_edge() {
let (l, graph) = learner();
l.record(TraceEvent::Search {
query: "why is the build broken".into(),
origin: Origin::Agent,
top_k: 5,
hits: vec![crate::trace::SearchHitTrace {
tool_id: "docker_build".into(),
score: 9.9,
}],
stages: Vec::new(),
took_ms: 0,
});
l.record(invoke("gh_run_list"));
let g = graph.read().unwrap();
assert_eq!(
g.intents[0].tools.keys().collect::<Vec<_>>(),
vec!["gh_run_list"]
);
}
#[test]
fn several_invokes_after_one_search_all_count_as_capabilities() {
let (l, graph) = learner();
l.record(search("why is the build broken"));
l.record(invoke("gh_run_list"));
l.record(invoke("gh_run_view"));
l.record(invoke("read_file"));
let g = graph.read().unwrap();
assert_eq!(g.len(), 1);
assert_eq!(g.intents[0].tools.len(), 3, "three capabilities were used");
assert_eq!(g.intents[0].support, 1, "but only one question was asked");
for (id, w) in &g.intents[0].tools {
assert_eq!(*w, 1.0, "{id} was used once");
}
}
#[test]
fn the_same_question_asked_twice_counts_twice() {
let (l, graph) = learner();
l.record(search("why is the build broken"));
l.record(invoke("gh_run_list"));
l.record(search("why is the build broken"));
l.record(invoke("gh_run_list"));
let g = graph.read().unwrap();
assert_eq!(g.intents[0].support, 2);
assert_eq!(g.intents[0].tools["gh_run_list"], 2.0);
}
#[test]
fn separate_searches_each_count() {
let (l, graph) = learner();
l.record(search("why is the build broken"));
l.record(invoke("gh_run_list"));
l.record(search("is the build broken again"));
l.record(invoke("gh_run_list"));
let g = graph.read().unwrap();
assert_eq!(g.intents[0].support, 2, "two questions, two observations");
}
#[test]
fn a_capability_search_across_both_registries_counts_once() {
let (l, graph) = learner();
l.record(search("why is the build broken"));
l.record(TraceEvent::SkillSearch {
query: "why is the build broken".into(),
origin: Origin::Agent,
top_k: 5,
hits: Vec::new(),
stages: Vec::new(),
took_ms: 0,
});
l.record(invoke("gh_run_list"));
l.record(TraceEvent::SkillInvoke {
skill_id: "ci-triage".into(),
took_ms: 1,
});
let g = graph.read().unwrap();
assert_eq!(g.len(), 1);
assert_eq!(
g.intents[0].support, 1,
"one question, however many catalogs it hit"
);
assert_eq!(g.intents[0].tools.len(), 1);
assert_eq!(g.intents[0].skills.len(), 1);
}
#[test]
fn two_learners_sharing_a_graph_count_a_capability_search_once() {
let graph = Arc::new(RwLock::new(IntentGraph::empty()));
let tools = Arc::new(UsageLearner::new(graph.clone(), Arc::new(NoopSink)));
let skills = Arc::new(UsageLearner::new(graph.clone(), Arc::new(NoopSink)));
tools.record(search("why is the build broken"));
skills.record(TraceEvent::SkillSearch {
query: "why is the build broken".into(),
origin: Origin::Agent,
top_k: 5,
hits: Vec::new(),
stages: Vec::new(),
took_ms: 0,
});
tools.record(invoke("gh_run_list"));
skills.record(TraceEvent::SkillInvoke {
skill_id: "ci-triage".into(),
took_ms: 1,
});
let g = graph.read().unwrap();
assert_eq!(g.len(), 1);
assert_eq!(
g.intents[0].support, 1,
"one question, even across two per-catalog learners"
);
assert_eq!(g.intents[0].tools.get("gh_run_list"), Some(&1.0));
assert_eq!(g.intents[0].skills.get("ci-triage"), Some(&1.0));
}
#[test]
fn a_new_search_replaces_the_pending_query() {
let (l, graph) = learner();
l.record(search("why is the build broken"));
l.record(search("rotate the signing key"));
l.record(invoke("vault_rotate"));
let g = graph.read().unwrap();
assert_eq!(g.len(), 1, "only the later query should have been credited");
assert!(
g.intents[0]
.members
.contains(&"rotate the signing key".to_string())
);
}
#[test]
fn skill_searches_and_skill_invokes_pair_on_the_skill_edges() {
let (l, graph) = learner();
l.record(TraceEvent::SkillSearch {
query: "why is the build broken".into(),
origin: Origin::Agent,
top_k: 5,
hits: Vec::new(),
stages: Vec::new(),
took_ms: 0,
});
l.record(TraceEvent::SkillInvoke {
skill_id: "ci-triage".into(),
took_ms: 1,
});
let g = graph.read().unwrap();
assert_eq!(g.intents[0].skills.get("ci-triage"), Some(&1.0));
assert!(g.intents[0].tools.is_empty());
}
#[test]
fn every_event_is_forwarded_to_the_inner_sink() {
let inner = Arc::new(MemorySink::new("s"));
let graph = Arc::new(RwLock::new(IntentGraph::empty()));
let l = UsageLearner::new(graph, inner.clone());
l.record(search("why is the build broken"));
l.record(invoke("gh_run_list"));
l.record(TraceEvent::AuthNeeds {
upstream: "gh".into(),
});
assert_eq!(inner.snapshot().len(), 3);
}
#[test]
fn unrelated_events_are_forwarded_without_learning() {
let (l, graph) = learner();
l.record(TraceEvent::AuthNeeds {
upstream: "gh".into(),
});
assert!(graph.read().unwrap().is_empty());
}
}