use ratatui::crossterm::event::{self, Event, KeyCode, KeyModifiers};
use ratatui::crossterm::execute;
use ratatui::crossterm::terminal::{disable_raw_mode, enable_raw_mode, EnterAlternateScreen, LeaveAlternateScreen};
use ratatui::prelude::*;
use ratatui::widgets::{Block, Borders, Paragraph};
use std::io::stdout;
use std::path::{Path, PathBuf};
use std::time::Duration;
use steeldb::discover_ontology::{candidate_terms, Cluster, OtDiscover};
use steeldb::text::Model2Vec;
const PALETTE: [Color; 8] = [Color::Cyan, Color::Green, Color::Yellow, Color::Magenta, Color::Blue, Color::Red, Color::LightGreen, Color::LightMagenta];
const SCHEDULE: [usize; 14] = [1, 2, 3, 5, 8, 12, 18, 26, 40, 60, 90, 130, 180, 200];
fn read_corpus(path: &Path) -> String {
let mut files = Vec::new();
if path.is_file() {
files.push(path.to_path_buf());
} else {
let mut stack = vec![path.to_path_buf()];
while let Some(d) = stack.pop() {
let Ok(rd) = std::fs::read_dir(&d) else { continue };
for e in rd.flatten() {
let name = e.file_name().to_string_lossy().to_string();
if name.starts_with('.') || name == "node_modules" || name == "target" {
continue;
}
let p = e.path();
if p.is_dir() {
stack.push(p);
} else {
files.push(p);
}
}
}
}
let mut text = String::new();
for f in files {
if let Ok(s) = std::fs::read_to_string(&f) {
text.push_str(&s);
text.push('\n');
}
if text.len() > 4_000_000 {
break;
}
}
text
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let args: Vec<String> = std::env::args().collect();
let path = match args.get(1) {
Some(p) => p.clone(),
None => {
eprintln!("usage: ontology <dir|file> [k]");
std::process::exit(2);
}
};
let k: usize = args.get(2).and_then(|s| s.parse().ok()).unwrap_or(8);
let m2v_dir = steeldb::paths::model_dir("model2vec", "STEELDB_MODEL2VEC", "potion.f32")
.ok_or_else(|| "model2vec models not found (set STEELDB_MODEL2VEC or place under models/model2vec)".to_string())?;
eprint!("loading model2vec … ");
let m2v = Model2Vec::load(&m2v_dir).map_err(|e| format!("model2vec ({}): {e}", m2v_dir.display()))?;
eprintln!("ok ({}d)", m2v.dim());
eprint!("reading corpus + extracting terms … ");
let text = read_corpus(Path::new(&path));
let terms = candidate_terms(&text, 300);
let mut kept_terms = Vec::new();
let mut embs = Vec::new();
for t in terms {
if let Some(e) = m2v.embed(&t) {
kept_terms.push(t);
embs.push(e);
}
}
eprintln!("{} embeddable terms", kept_terms.len());
if kept_terms.len() < 4 {
eprintln!("not enough terms to discover an ontology");
std::process::exit(1);
}
let disc = OtDiscover::new(kept_terms, embs, k);
if args.iter().any(|a| a == "--print") || std::env::var("STEELDB_PRINT").is_ok() {
let (assign, cost) = disc.assign(200);
println!("{} terms → {} facets · cost {:.4}", disc.terms.len(), disc.k, cost);
for c in disc.clusters(&assign, 8) {
println!(" [{:>3}] {:<16} {}", c.size, c.label, c.terms.iter().skip(1).cloned().collect::<Vec<_>>().join(" "));
}
return Ok(());
}
let mut app = App { disc, step: 0, clusters: Vec::new(), cost: 0.0, cost_hist: Vec::new(), nterms: 0 };
app.nterms = app.disc.terms.len();
app.advance();
enable_raw_mode()?;
let mut out = stdout();
execute!(out, EnterAlternateScreen)?;
let mut terminal = Terminal::new(CrosstermBackend::new(out))?;
let res = run(&mut terminal, &mut app);
disable_raw_mode()?;
execute!(terminal.backend_mut(), LeaveAlternateScreen)?;
terminal.show_cursor()?;
res.map_err(Into::into)
}
struct App {
disc: OtDiscover,
step: usize,
clusters: Vec<Cluster>,
cost: f32,
cost_hist: Vec<f32>,
nterms: usize,
}
impl App {
fn iters(&self) -> usize {
SCHEDULE[self.step.min(SCHEDULE.len() - 1)]
}
fn converged(&self) -> bool {
self.step >= SCHEDULE.len() - 1
}
fn advance(&mut self) {
let (assign, cost) = self.disc.assign(self.iters());
self.clusters = self.disc.clusters(&assign, 8);
self.cost = cost;
self.cost_hist.push(cost);
if !self.converged() {
self.step += 1;
}
}
}
fn run<B: Backend>(terminal: &mut Terminal<B>, app: &mut App) -> std::io::Result<()> {
let mut tick = 0usize;
loop {
terminal.draw(|f| draw(f, app))?;
if event::poll(Duration::from_millis(150))? {
if let Event::Key(k) = event::read()? {
if matches!(k.code, KeyCode::Esc | KeyCode::Char('q')) || (k.modifiers.contains(KeyModifiers::CONTROL) && matches!(k.code, KeyCode::Char('c'))) {
return Ok(());
}
if matches!(k.code, KeyCode::Char('r')) {
app.step = 0;
app.cost_hist.clear();
app.advance();
}
}
}
tick += 1;
if !app.converged() && tick % 2 == 0 {
app.advance();
}
if app.converged() && tick % 2 == 0 {
let (assign, cost) = app.disc.assign(SCHEDULE[SCHEDULE.len() - 1]);
app.clusters = app.disc.clusters(&assign, 8);
app.cost = cost;
}
}
}
fn spark(hist: &[f32]) -> String {
if hist.is_empty() {
return String::new();
}
let bars = ['▁', '▂', '▃', '▄', '▅', '▆', '▇', '█'];
let (lo, hi) = hist.iter().fold((f32::INFINITY, f32::NEG_INFINITY), |(l, h), &v| (l.min(v), h.max(v)));
let rng = (hi - lo).max(1e-9);
hist.iter().map(|&v| bars[(((v - lo) / rng) * 7.0).round() as usize]).collect()
}
fn draw(f: &mut Frame, app: &App) {
let chunks = Layout::default()
.direction(Direction::Vertical)
.constraints([Constraint::Length(1), Constraint::Min(1), Constraint::Length(3)])
.split(f.area());
let status = if app.converged() { "converged" } else { "converging…" };
let header = format!(" Sinkhorn-OT ontology discovery · {} terms → {} facets · iter {} · cost {:.4} · {}", app.nterms, app.disc.k, app.iters(), app.cost, status);
f.render_widget(
Paragraph::new(Line::from(Span::styled(header, Style::default().fg(Color::Black).bg(Color::Cyan).add_modifier(Modifier::BOLD)))).style(Style::default().bg(Color::Cyan)),
chunks[0],
);
let maxsize = app.clusters.iter().map(|c| c.size).max().unwrap_or(1).max(1);
let barw = 16;
let mut lines: Vec<Line> = Vec::new();
for (i, c) in app.clusters.iter().enumerate() {
let color = PALETTE[i % PALETTE.len()];
let filled = (c.size * barw / maxsize).max(1);
let bar: String = "█".repeat(filled);
let pad: String = "·".repeat(barw - filled);
let vocab = c.terms.iter().skip(1).cloned().collect::<Vec<_>>().join(" ");
lines.push(Line::from(vec![
Span::styled(format!(" {bar}"), Style::default().fg(color)),
Span::styled(pad, Style::default().fg(Color::DarkGray)),
Span::styled(format!(" {:>3} ", c.size), Style::default().fg(Color::DarkGray)),
Span::styled(format!("{:<16}", c.label), Style::default().fg(color).add_modifier(Modifier::BOLD)),
Span::styled(vocab, Style::default().fg(Color::Gray)),
]));
}
f.render_widget(
Paragraph::new(lines).block(Block::default().borders(Borders::ALL).title(" discovered facets (cluster → top terms) ")),
chunks[1],
);
let foot = Line::from(vec![
Span::styled(" transport cost ", Style::default().fg(Color::DarkGray)),
Span::styled(spark(&app.cost_hist), Style::default().fg(Color::Green)),
Span::styled(" q/Esc quit · r replay", Style::default().fg(Color::DarkGray)),
]);
f.render_widget(Paragraph::new(foot).block(Block::default().borders(Borders::ALL).title(" convergence ")), chunks[2]);
}