use crate::corpus::Document;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::path::Path;
#[derive(Debug, Deserialize)]
pub struct TopicDef {
#[serde(default)]
pub keywords: Vec<String>,
}
#[derive(Debug, Deserialize)]
pub struct SourceDefault {
#[serde(default = "default_role")]
pub role: String,
#[serde(default)]
pub topics: Vec<String>,
}
fn default_role() -> String {
"style".to_string()
}
#[derive(Debug, Deserialize)]
pub struct Taxonomy {
#[serde(default)]
pub topics: HashMap<String, TopicDef>,
#[serde(default)]
pub source_defaults: HashMap<String, SourceDefault>,
}
#[derive(Debug, Deserialize, Default)]
pub struct Override {
pub role: Option<String>,
pub topics: Option<Vec<String>>,
pub lora_bucket: Option<String>,
pub rag: Option<bool>,
}
#[derive(Debug, Default)]
pub struct Overrides(pub HashMap<String, Override>);
#[derive(Serialize)]
pub struct CatalogEntry {
pub doc_id: String,
pub source: String,
pub title: String,
pub role: String,
pub topics: Vec<String>,
pub lora_bucket: String,
pub rag: bool,
pub is_reference: bool,
pub chars: usize,
#[serde(skip_serializing_if = "Option::is_none", default)] pub auto_topic: Option<String>,
#[serde(skip_serializing_if = "Option::is_none", default)] pub topic_confidence: Option<f32>,
}
fn lora_bucket_for(role: &str) -> String {
match role {
"code" => "code",
"knowledge" => "knowledge",
"capability" => "capability",
"craft" => "reference",
other => other, }
.to_string()
}
pub fn load_taxonomy(path: &Path) -> Taxonomy {
std::fs::read_to_string(path)
.ok()
.and_then(|s| serde_yaml::from_str(&s).ok())
.unwrap_or(Taxonomy { topics: HashMap::new(), source_defaults: HashMap::new() })
}
pub fn load_overrides(path: &Path) -> Overrides {
std::fs::read_to_string(path)
.ok()
.and_then(|s| serde_json::from_str::<HashMap<String, Override>>(&s).ok())
.map(Overrides)
.unwrap_or_default()
}
pub fn classify(doc: &Document, taxonomy: &Taxonomy, overrides: &Overrides) -> CatalogEntry {
classify_with_fallback(doc, taxonomy, overrides, None)
}
pub fn classify_with_fallback(doc: &Document, taxonomy: &Taxonomy, overrides: &Overrides, fallback_role: Option<&str>) -> CatalogEntry {
let default = taxonomy.source_defaults.get(&doc.source);
let mut role = default
.map(|d| d.role.clone())
.or_else(|| fallback_role.map(|r| r.to_string()))
.unwrap_or_else(|| if doc.is_reference { "reference".to_string() } else { "style".to_string() });
let mut topics: Vec<String> = default.map(|d| d.topics.clone()).unwrap_or_default();
let hay = {
let mut s = doc.title.to_lowercase();
s.push(' ');
let body: String = doc.text.chars().take(8000).collect();
s.push_str(&body.to_lowercase());
s
};
let mut matched: Vec<&String> = taxonomy
.topics
.iter()
.filter(|(topic, def)| {
def.keywords.iter().any(|k| hay.contains(&k.to_lowercase())) && !topics.contains(*topic)
})
.map(|(topic, _)| topic)
.collect();
matched.sort();
for topic in matched {
topics.push(topic.clone());
}
let mut lora_bucket = lora_bucket_for(&role);
let mut rag = matches!(role.as_str(), "reference" | "code" | "knowledge" | "craft" | "style");
if let Some(ov) = overrides.0.get(&doc.doc_id) {
if let Some(r) = &ov.role {
role = r.clone();
lora_bucket = lora_bucket_for(&role);
rag = matches!(role.as_str(), "reference" | "code" | "knowledge" | "craft" | "style");
}
if let Some(t) = &ov.topics {
topics = t.clone();
}
if let Some(lb) = &ov.lora_bucket {
lora_bucket = lb.clone();
}
if let Some(rg) = ov.rag {
rag = rg;
}
}
CatalogEntry {
doc_id: doc.doc_id.clone(),
source: doc.source.clone(),
title: doc.title.clone(),
role,
topics,
lora_bucket,
rag,
is_reference: doc.is_reference,
chars: doc.text.chars().count(),
auto_topic: None,
topic_confidence: None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::corpus::Document;
fn doc(source: &str, is_ref: bool, text: &str) -> Document {
Document {
doc_id: format!("{source}:x"),
text: text.to_string(),
source: source.to_string(),
title: "x".to_string(),
is_reference: is_ref,
}
}
#[test]
fn source_default_role_and_lora_bucket() {
let tax = Taxonomy {
topics: HashMap::new(),
source_defaults: HashMap::from([(
"twitter".to_string(),
SourceDefault { role: "style".to_string(), topics: vec!["humor".to_string()] },
)]),
};
let overrides = Overrides(HashMap::new());
let entry = classify(&doc("twitter", false, "lol meme"), &tax, &overrides);
assert_eq!(entry.role, "style");
assert_eq!(entry.lora_bucket, "style");
assert!(entry.rag);
}
#[test]
fn keyword_adds_topic_and_craft_maps_to_reference_bucket() {
let tax = Taxonomy {
topics: HashMap::from([(
"security".to_string(),
TopicDef { keywords: vec!["ransomware".to_string()] },
)]),
source_defaults: HashMap::from([(
"web".to_string(),
SourceDefault { role: "craft".to_string(), topics: vec![] },
)]),
};
let overrides = Overrides(HashMap::new());
let entry = classify(&doc("web", false, "a ransomware writeup"), &tax, &overrides);
assert_eq!(entry.lora_bucket, "reference"); assert!(entry.topics.contains(&"security".to_string()));
}
#[test]
fn override_wins() {
let tax = Taxonomy { topics: HashMap::new(), source_defaults: HashMap::new() };
let overrides = Overrides(HashMap::from([(
"twitter:x".to_string(),
Override { role: Some("code".to_string()), topics: None, lora_bucket: None, rag: None },
)]));
let entry = classify(&doc("twitter", false, "hi"), &tax, &overrides);
assert_eq!(entry.role, "code");
assert_eq!(entry.lora_bucket, "code");
}
#[test]
fn classify_leaves_auto_topic_none() {
let d = doc("me", false, "some text");
let e = classify(&d, &Taxonomy { topics: Default::default(), source_defaults: Default::default() }, &Overrides::default());
assert!(e.auto_topic.is_none() && e.topic_confidence.is_none());
}
}