use std::collections::BTreeMap;
use std::path::PathBuf;
const CODEBOOK_KINDS: &[&str] = &["ENT", "REL", "GEO", "TIME"];
fn sentences(text: &str) -> Vec<String> {
text.split(['.', '\n', ';', '!', '?'])
.map(|s| s.trim())
.filter(|s| s.split_whitespace().count() >= 3)
.map(|s| s.to_string())
.collect()
}
fn main() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let args: Vec<String> = std::env::args().collect();
let dir = args.get(1).cloned().unwrap_or_else(|| "benchmark_corpus".into());
let out_path = args.get(2).cloned().unwrap_or_else(|| "spans.json".into());
let max_docs: usize = args.get(3).and_then(|s| s.parse().ok()).unwrap_or(24);
let ml = steeldb::paths::model_dir("step0_bundle_ml", "STEELDB_ML_BUNDLE", "spo.onnx")
.ok_or("SPO tagger bundle not found (models/step0_bundle_ml/spo.onnx)")?;
let m2v_dir = steeldb::paths::model_dir("model2vec", "STEELDB_MODEL2VEC", "potion.f32")
.ok_or("model2vec not found (models/model2vec/potion.f32)")?;
let mut tagger = steeldb::text::tagger::SpoTagger::load(&ml)?;
let embedder = steeldb::text::Model2Vec::load(&m2v_dir)?;
eprintln!("tagger: {} ยท embedder: {} dims", ml.display(), embedder.dim());
let path = PathBuf::from(&dir);
let raw = if path.is_file() {
std::fs::read_to_string(&path)?
} else {
let mut files: Vec<PathBuf> = std::fs::read_dir(&path)?
.filter_map(|e| e.ok().map(|e| e.path()))
.filter(|p| p.extension().map(|x| x == "md" || x == "txt").unwrap_or(false))
.collect();
files.sort();
files
.iter()
.take(max_docs)
.filter_map(|p| std::fs::read_to_string(p).ok())
.collect::<Vec<_>>()
.join("\n---\n")
};
let docs: Vec<String> = raw
.split("\n---")
.map(|d| d.trim().to_string())
.filter(|d| !d.is_empty())
.take(max_docs)
.collect();
eprintln!("documents: {}", docs.len());
let mut seen: BTreeMap<(String, String), usize> = BTreeMap::new();
let mut routed_away: BTreeMap<String, usize> = BTreeMap::new();
let mut records: Vec<serde_json::Value> = Vec::new();
let mut numeric_rules: Vec<serde_json::Value> = Vec::new();
for (di, doc) in docs.iter().enumerate() {
for sentence in sentences(doc) {
let spans = match tagger.tag(&sentence) {
Ok(s) => s,
Err(e) => {
eprintln!(" tag failed: {e}");
continue;
}
};
for sp in spans {
let (ss, se) = steeldb::spans::snap_to_words(&sentence, sp.start, sp.end);
let surface = sentence
.get(ss..se)
.map(|t| t.trim().to_string())
.unwrap_or_else(|| sp.text.trim().to_string());
if surface.chars().count() < 3 || !surface.chars().any(|c| c.is_alphanumeric()) {
continue;
}
if !CODEBOOK_KINDS.contains(&sp.kind.as_str()) {
*routed_away.entry(sp.kind.clone()).or_default() += 1;
if sp.kind == "QTY" {
for (qs, qe, field) in steeldb::emergent::quantity_spans(&surface) {
let text = &surface[qs..qe];
let digits: String = text
.chars()
.enumerate()
.take_while(|(i, c)| c.is_ascii_digit() || *c == '.' || (*i == 0 && *c == '-'))
.map(|(_, c)| c)
.collect();
if let Ok(v) = digits.parse::<f64>() {
numeric_rules.push(serde_json::json!({
"span": surface, "field": field, "value": v, "doc": di,
}));
}
}
}
continue;
}
let key = (sp.kind.clone(), surface.to_lowercase());
let count = seen.entry(key.clone()).or_insert(0);
*count += 1;
if *count > 1 {
continue; }
let Some(vec) = embedder.embed(&surface) else { continue };
let v: Vec<f32> = vec.iter().map(|x| (x * 10_000.0).round() / 10_000.0).collect();
records.push(serde_json::json!({
"text": surface,
"kind": sp.kind,
"doc": di,
"vec": v,
}));
}
}
}
for r in records.iter_mut() {
let k = (
r["kind"].as_str().unwrap_or("").to_string(),
r["text"].as_str().unwrap_or("").to_lowercase(),
);
r["count"] = serde_json::json!(seen.get(&k).copied().unwrap_or(1));
}
let mut by_kind: BTreeMap<&str, usize> = BTreeMap::new();
for r in &records {
*by_kind.entry(r["kind"].as_str().unwrap_or("?")).or_default() += 1;
}
let payload = serde_json::json!({
"documents": docs.len(),
"dim": embedder.dim(),
"codebook_kinds": CODEBOOK_KINDS,
"spans": records,
"routed_away": routed_away,
"numeric_rules": numeric_rules,
"source": "SPO tagger (spo.onnx) + model2vec embeddings, exported natively",
});
std::fs::write(&out_path, serde_json::to_string(&payload)?)?;
eprintln!("spans kept for the codebook:");
for (k, n) in &by_kind {
eprintln!(" {k:<6} {n}");
}
eprintln!("routed away (never enter the codebook):");
for (k, n) in &routed_away {
eprintln!(" {k:<6} {n}");
}
let bytes = std::fs::metadata(&out_path).map(|m| m.len()).unwrap_or(0);
eprintln!("wrote {} ({} KB)", out_path, bytes / 1024);
Ok(())
}