pub mod chunk_classifier;
pub mod page_classifier;
pub use chunk_classifier::{classify_chunks, classify_chunks_owned};
pub use page_classifier::{classify_pages, classify_text};
pub async fn classify_document(
pages: &[&str],
config: &crate::core::config::PageClassificationConfig,
) -> crate::Result<Vec<crate::ClassificationLabel>> {
if config.labels.is_empty() {
return Err(crate::XbergError::validation(
"PageClassificationConfig.labels must contain at least one entry",
));
}
if pages.is_empty() {
return Ok(Vec::new());
}
let ctx = page_classifier::ClassifyContext::new(config);
let mut per_page_labels: Vec<Vec<crate::ClassificationLabel>> = Vec::new();
for page_text in pages {
if page_text.is_empty() {
continue;
}
let (labels, _usage) = page_classifier::classify_one(page_text, &ctx, config).await?;
per_page_labels.push(labels);
}
Ok(finalize_document_labels(&per_page_labels, config.multi_label))
}
pub(crate) async fn classify_document_onto(
result: &mut crate::types::ExtractedDocument,
config: &crate::core::config::PageClassificationConfig,
) -> crate::Result<Vec<crate::ClassificationLabel>> {
if config.labels.is_empty() {
return Err(crate::XbergError::validation(
"PageClassificationConfig.labels must contain at least one entry",
));
}
let pages: Vec<String> = match result.pages.as_deref() {
Some(pages) if !pages.is_empty() => pages.iter().map(|page| page.content.clone()).collect(),
_ => vec![result.content.clone()],
};
let ctx = page_classifier::ClassifyContext::new(config);
let mut per_page_labels: Vec<Vec<crate::ClassificationLabel>> = Vec::new();
let mut classifications: Vec<crate::types::classification::PageClassification> = Vec::new();
let mut usages: Vec<crate::types::LlmUsage> = Vec::new();
for (index, page_text) in pages.iter().enumerate() {
if page_text.is_empty() {
continue;
}
let (labels, usage) = page_classifier::classify_one(page_text, &ctx, config).await?;
if let Some(usage) = usage {
usages.push(usage);
}
classifications.push(crate::types::classification::PageClassification {
page_number: index as u32 + 1,
labels: labels.clone(),
});
per_page_labels.push(labels);
}
if !classifications.is_empty() {
result.page_classifications = Some(classifications);
}
if !usages.is_empty() {
result.llm_usage.get_or_insert_with(Vec::new).extend(usages);
}
Ok(finalize_document_labels(&per_page_labels, config.multi_label))
}
fn finalize_document_labels(
per_page_labels: &[Vec<crate::ClassificationLabel>],
multi_label: bool,
) -> Vec<crate::ClassificationLabel> {
let aggregated = aggregate_page_labels(per_page_labels);
if multi_label {
let mut labels = aggregated;
labels.sort_by(|a, b| a.label.cmp(&b.label));
labels
} else {
let best = aggregated.into_iter().max_by(|a, b| {
let a_score = a.confidence.unwrap_or(0.0);
let b_score = b.confidence.unwrap_or(0.0);
a_score.partial_cmp(&b_score).unwrap_or(std::cmp::Ordering::Equal)
});
best.into_iter().collect()
}
}
fn aggregate_page_labels(per_page_labels: &[Vec<crate::ClassificationLabel>]) -> Vec<crate::ClassificationLabel> {
let mut order: Vec<String> = Vec::new();
let mut counts: std::collections::HashMap<String, (f32, u32)> = std::collections::HashMap::new();
for labels in per_page_labels {
for label in labels {
let entry = counts.entry(label.label.clone()).or_insert_with(|| {
order.push(label.label.clone());
(0.0, 0)
});
if let Some(conf) = label.confidence {
entry.0 += conf;
entry.1 += 1;
}
}
}
order
.into_iter()
.map(|label| {
let (sum, count) = counts[&label];
let confidence = if count > 0 { Some(sum / count as f32) } else { None };
crate::ClassificationLabel { label, confidence }
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ClassificationLabel;
fn label(name: &str, confidence: Option<f32>) -> ClassificationLabel {
ClassificationLabel {
label: name.to_string(),
confidence,
}
}
#[test]
fn should_average_confidence_across_pages_reporting_the_same_label() {
let per_page = vec![
vec![label("invoice", Some(0.6))],
vec![label("invoice", Some(0.8)), label("memo", Some(0.2))],
];
let aggregated = aggregate_page_labels(&per_page);
assert_eq!(aggregated.len(), 2);
assert_eq!(aggregated[0].label, "invoice");
assert_eq!(aggregated[0].confidence, Some(0.700_000_05));
assert_eq!(aggregated[1].label, "memo");
assert_eq!(aggregated[1].confidence, Some(0.2));
}
#[test]
fn should_return_none_confidence_when_no_page_reported_one() {
let per_page = vec![vec![label("invoice", None)], vec![label("invoice", None)]];
let aggregated = aggregate_page_labels(&per_page);
assert_eq!(aggregated.len(), 1);
assert_eq!(aggregated[0].label, "invoice");
assert_eq!(aggregated[0].confidence, None);
}
#[test]
fn should_average_only_over_pages_that_reported_a_confidence() {
let per_page = vec![vec![label("invoice", Some(0.9))], vec![label("invoice", None)]];
let aggregated = aggregate_page_labels(&per_page);
assert_eq!(aggregated.len(), 1);
assert_eq!(
aggregated[0].confidence,
Some(0.9),
"the None entry must not dilute the average"
);
}
#[test]
fn should_preserve_first_seen_order() {
let per_page = vec![vec![label("zeta", Some(0.5)), label("alpha", Some(0.5))]];
let aggregated = aggregate_page_labels(&per_page);
let names: Vec<&str> = aggregated.iter().map(|l| l.label.as_str()).collect();
assert_eq!(names, vec!["zeta", "alpha"]);
}
}