use crate::api::SteelDb;
#[derive(Debug, Clone, PartialEq)]
pub struct Candidate {
pub name: String,
pub words: Vec<String>,
pub rationale: String,
}
#[derive(Debug, Clone, Default)]
pub struct Proposal {
pub candidates: Vec<Candidate>,
pub source: String,
}
impl std::fmt::Display for Proposal {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
writeln!(f, "proposal from {} — {} candidate(s)", self.source, self.candidates.len())?;
for c in &self.candidates {
writeln!(f, " {} — {}", c.name, c.rationale)?;
writeln!(f, " words: {}", c.words.join(", "))?;
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct Verdict {
pub name: String,
pub kept: bool,
pub reason: String,
pub coverage: f64,
pub overlap: f64,
}
impl SteelDb {
pub fn adopt(&mut self, proposal: &Proposal) -> Vec<Verdict> {
let mut verdicts = Vec::new();
for (round, cand) in proposal.candidates.iter().enumerate() {
let c = crate::grow::Candidate {
name: cand.name.clone(),
parent: None,
description: cand.rationale.clone(),
examples: cand.words.clone(),
worth_adding: true,
};
let docs = self.documents().to_vec();
let spec = self.spec_snapshot();
let scored = crate::grow::score_candidate_full(&spec, &docs, &c);
let (score, dup) = match scored {
Some((s, d)) => (Some(s), d),
None => (None, None),
};
let ev = crate::grow::gate_full(&spec, &c, score.as_ref(), dup, self.min_gain(), round);
if ev.kept {
self.push_category(cand.name.clone(), cand.words.clone());
}
verdicts.push(Verdict {
name: cand.name.clone(),
kept: ev.kept,
reason: ev.reason,
coverage: ev.coverage,
overlap: ev.maxcos,
});
}
if verdicts.iter().any(|v| v.kept) {
self.reproject();
}
verdicts
}
}
pub struct Teacher {
kind: Kind,
}
enum Kind {
Fixed(Proposal),
#[cfg(feature = "paddock")]
Local { base_url: String, model: String },
#[cfg(feature = "bedrock")]
Bedrock { model_id: String },
}
#[derive(Debug)]
pub enum LearnError {
NotConfigured(String),
BadResponse(String),
Transport(String),
}
impl std::fmt::Display for LearnError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
LearnError::NotConfigured(m) => write!(f, "not configured for learning: {m}"),
LearnError::BadResponse(m) => write!(f, "unusable response: {m}"),
LearnError::Transport(m) => write!(f, "call failed: {m}"),
}
}
}
impl std::error::Error for LearnError {}
impl Teacher {
pub fn fixed(proposal: Proposal) -> Teacher {
Teacher { kind: Kind::Fixed(proposal) }
}
#[cfg(feature = "paddock")]
pub fn local(base_url: impl Into<String>, model: impl Into<String>) -> Result<Teacher, LearnError> {
let base_url = base_url.into();
if !base_url.starts_with("http") {
return Err(LearnError::NotConfigured(format!(
"base_url should be an http(s) endpoint, got {base_url:?}"
)));
}
Ok(Teacher { kind: Kind::Local { base_url, model: model.into() } })
}
#[cfg(feature = "paddock")]
pub fn ollama(model: impl Into<String>) -> Result<Teacher, LearnError> {
Teacher::local("http://localhost:11434/v1", model)
}
#[cfg(feature = "bedrock")]
pub fn bedrock(model_id: impl Into<String>) -> Result<Teacher, LearnError> {
if std::env::var("AWS_REGION").is_err() && std::env::var("AWS_DEFAULT_REGION").is_err() {
return Err(LearnError::NotConfigured(
"set AWS_REGION (or AWS_DEFAULT_REGION) to the region hosting the model".into(),
));
}
Ok(Teacher { kind: Kind::Bedrock { model_id: model_id.into() } })
}
pub async fn propose_categories(&self, db: &SteelDb) -> Result<Proposal, LearnError> {
let _ = db;
match &self.kind {
Kind::Fixed(p) => Ok(p.clone()),
#[cfg(feature = "paddock")]
Kind::Local { base_url, model } => {
let cfg = crate::agent::config::ProviderConfig::Paddock {
base_url: base_url.clone(),
model: model.clone(),
api_key: None,
};
Self::propose_via(cfg, db).await
}
#[cfg(feature = "bedrock")]
Kind::Bedrock { model_id } => {
let cfg = crate::agent::config::ProviderConfig::Bedrock {
model_id: model_id.clone(),
region: std::env::var("AWS_REGION").ok(),
};
Self::propose_via(cfg, db).await
}
}
}
#[cfg(any(feature = "paddock", feature = "bedrock"))]
async fn propose_via(
cfg: crate::agent::config::ProviderConfig,
db: &SteelDb,
) -> Result<Proposal, LearnError> {
let label = match &cfg {
#[cfg(feature = "paddock")]
crate::agent::config::ProviderConfig::Paddock { model, .. } => format!("local:{model}"),
#[cfg(feature = "bedrock")]
crate::agent::config::ProviderConfig::Bedrock { model_id, .. } => format!("bedrock:{model_id}"),
_ => "model".to_string(),
};
let provider = cfg.build().await.map_err(LearnError::Transport)?;
let sample: Vec<String> = db.documents().iter().take(48).cloned().collect();
let spec = crate::vocabulary::propose(provider.as_ref(), "documents", &sample)
.await
.map_err(LearnError::BadResponse)?;
Ok(Proposal {
source: label,
candidates: spec
.entity_facets
.into_iter()
.map(|f| Candidate { name: f.name, words: f.examples, rationale: f.description })
.collect(),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn docs() -> Vec<String> {
[
"Morty Shade defeated Wallace Gale at Ecruteak City during the Indigo Invitational in 2025.",
"Bea Strike defeated Iris Draco at Ecruteak City during the Indigo Invitational in 2025.",
"A habitat survey recorded Aggron near Sootopolis City at an elevation of 1082 m.",
"A habitat survey recorded Salamence near Sootopolis City at an elevation of 2369 m.",
"Milotic is not permitted in Series 1 play for the 2025 season.",
]
.iter()
.map(|s| s.to_string())
.collect()
}
#[test]
fn a_proposal_changes_nothing_until_adopted() {
let db = SteelDb::ingest(docs()).unwrap();
let before = db.categories().len();
let _p = Proposal {
source: "test".into(),
candidates: vec![Candidate {
name: "trainer".into(),
words: vec!["defeated".into(), "Shade".into()],
rationale: "people who compete".into(),
}],
};
assert_eq!(db.categories().len(), before);
}
#[test]
fn a_fixed_teacher_needs_no_credentials() {
let db = SteelDb::ingest(docs()).unwrap();
let p = Proposal {
source: "fixed".into(),
candidates: vec![Candidate {
name: "ruling".into(),
words: vec!["permitted".into(), "Series".into(), "season".into()],
rationale: "competition rules".into(),
}],
};
let teacher = Teacher::fixed(p.clone());
let got = block_on(teacher.propose_categories(&db)).unwrap();
assert_eq!(got.candidates, p.candidates);
}
fn block_on<F: std::future::Future>(mut fut: F) -> F::Output {
use std::task::{Context, Poll, RawWaker, RawWakerVTable, Waker};
fn noop(_: *const ()) {}
fn clone(p: *const ()) -> RawWaker {
RawWaker::new(p, &VTABLE)
}
static VTABLE: RawWakerVTable = RawWakerVTable::new(clone, noop, noop, noop);
let waker = unsafe { Waker::from_raw(RawWaker::new(std::ptr::null(), &VTABLE)) };
let mut cx = Context::from_waker(&waker);
let mut fut = unsafe { std::pin::Pin::new_unchecked(&mut fut) };
loop {
match fut.as_mut().poll(&mut cx) {
Poll::Ready(v) => return v,
Poll::Pending => panic!("the fixed teacher must not yield"),
}
}
}
#[test]
fn adoption_reports_a_verdict_per_candidate_and_can_reject() {
let mut db = SteelDb::ingest(docs()).unwrap();
let proposal = Proposal {
source: "test".into(),
candidates: vec![
Candidate {
name: "ruling".into(),
words: vec!["permitted".into(), "Series".into()],
rationale: "rules".into(),
},
Candidate {
name: "ruling2".into(),
words: vec!["permitted".into(), "Series".into()],
rationale: "the same thing again".into(),
},
],
};
let verdicts = db.adopt(&proposal);
assert_eq!(verdicts.len(), 2, "one verdict per candidate");
for v in &verdicts {
assert!(!v.reason.is_empty(), "a rejection must be explicable: {v:?}");
}
assert!(
!verdicts[1].kept || verdicts[1].overlap < 0.99,
"an exact duplicate should not be adopted unexamined: {:?}",
verdicts[1]
);
}
#[test]
fn an_adopted_category_becomes_queryable() {
let mut db = SteelDb::ingest(docs()).unwrap();
let proposal = Proposal {
source: "test".into(),
candidates: vec![Candidate {
name: "ruling".into(),
words: vec!["permitted".into(), "season".into(), "Series".into()],
rationale: "rules".into(),
}],
};
let verdicts = db.adopt(&proposal);
if verdicts[0].kept {
let answer = db.query("ruling/*").expect("an adopted category must be queryable");
assert!(!answer.is_empty(), "and must actually match documents");
}
}
#[test]
fn proposals_display_for_review_before_adoption() {
let p = Proposal {
source: "bedrock:test".into(),
candidates: vec![Candidate {
name: "trainer".into(),
words: vec!["defeated".into()],
rationale: "competitors".into(),
}],
};
let shown = p.to_string();
assert!(shown.contains("bedrock:test"), "{shown}");
assert!(shown.contains("trainer"), "{shown}");
assert!(shown.contains("competitors"), "the rationale must be reviewable: {shown}");
}
}