use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum CitationSource {
Llm,
Extracted,
Fused,
None,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CitedField {
pub value: serde_json::Value,
pub page: Option<u32>,
pub bbox: Option<[f64; 4]>,
pub confidence: Option<f64>,
pub source: CitationSource,
}
pub struct CitationOutput {
pub structured_output: serde_json::Value,
pub structured_output_flat: serde_json::Value,
}
pub fn fuse(
merged: serde_json::Value,
ocr_elements: &[serde_json::Value],
element_metadata: &[serde_json::Value],
emit_citations: bool,
match_threshold: f64,
fused_confidence: f64,
) -> CitationOutput {
if !emit_citations {
return CitationOutput {
structured_output: merged.clone(),
structured_output_flat: merged,
};
}
let structured_output = envelope_with_citations(
&merged,
ocr_elements,
element_metadata,
match_threshold,
fused_confidence,
);
let structured_output_flat = flatten_cited(&structured_output);
CitationOutput {
structured_output,
structured_output_flat,
}
}
fn envelope_with_citations(
value: &serde_json::Value,
ocr_elements: &[serde_json::Value],
element_metadata: &[serde_json::Value],
match_threshold: f64,
fused_confidence: f64,
) -> serde_json::Value {
match value {
serde_json::Value::Object(obj) => {
let mut result = serde_json::Map::new();
for (k, v) in obj {
result.insert(
k.clone(),
envelope_with_citations(v, ocr_elements, element_metadata, match_threshold, fused_confidence),
);
}
serde_json::Value::Object(result)
}
serde_json::Value::Array(arr) => {
let result: Vec<_> = arr
.iter()
.map(|v| envelope_with_citations(v, ocr_elements, element_metadata, match_threshold, fused_confidence))
.collect();
serde_json::Value::Array(result)
}
leaf => {
if is_citation_envelope(leaf) {
leaf.clone()
} else {
let cited =
try_fuse_with_extracted(leaf, ocr_elements, element_metadata, match_threshold, fused_confidence);
serde_json::to_value(&cited).unwrap_or(leaf.clone())
}
}
}
}
fn is_citation_envelope(v: &serde_json::Value) -> bool {
if let serde_json::Value::Object(obj) = v {
obj.contains_key("value") && !obj.is_empty() && obj.len() <= 5
} else {
false
}
}
fn try_fuse_with_extracted(
value: &serde_json::Value,
ocr_elements: &[serde_json::Value],
_element_metadata: &[serde_json::Value],
match_threshold: f64,
fused_confidence: f64,
) -> CitedField {
let value_str = match value {
serde_json::Value::String(s) => s.clone(),
other => other.to_string(),
};
let value_lower = value_str.to_lowercase();
let value_trimmed = value_lower.trim();
for ocr in ocr_elements {
if let Some(text) = ocr.get("text").and_then(|t| t.as_str()) {
let ocr_lower = text.to_lowercase();
let ocr_text = ocr_lower.trim();
if text_similarity(value_trimmed, ocr_text) > match_threshold {
let page = ocr.get("page_number").and_then(|p| p.as_u64()).map(|p| p as u32);
let bbox = extract_bbox(ocr);
return CitedField {
value: value.clone(),
page,
bbox,
confidence: Some(fused_confidence),
source: CitationSource::Fused,
};
}
}
}
CitedField {
value: value.clone(),
page: None,
bbox: None,
confidence: None,
source: CitationSource::None,
}
}
fn text_similarity(a: &str, b: &str) -> f64 {
if a.is_empty() && b.is_empty() {
return 1.0;
}
if a.is_empty() || b.is_empty() {
return 0.0;
}
let max_len = a.chars().count().max(b.chars().count());
let matching = a.chars().zip(b.chars()).filter(|(ca, cb)| ca == cb).count();
matching as f64 / max_len as f64
}
fn extract_bbox(ocr: &serde_json::Value) -> Option<[f64; 4]> {
ocr.get("bbox").and_then(|b| {
if let serde_json::Value::Array(arr) = b {
if arr.len() >= 4 {
let coords: Option<Vec<f64>> = arr.iter().take(4).map(|v| v.as_f64()).collect();
coords.map(|c| [c[0], c[1], c[2], c[3]])
} else {
None
}
} else {
None
}
})
}
fn flatten_cited(value: &serde_json::Value) -> serde_json::Value {
match value {
serde_json::Value::Object(obj) => {
if is_citation_envelope(value)
&& let Some(inner_value) = obj.get("value")
{
return flatten_cited(inner_value);
}
let mut result = serde_json::Map::new();
for (k, v) in obj {
result.insert(k.clone(), flatten_cited(v));
}
serde_json::Value::Object(result)
}
serde_json::Value::Array(arr) => {
let result: Vec<_> = arr.iter().map(flatten_cited).collect();
serde_json::Value::Array(result)
}
leaf => leaf.clone(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn emit_citations_false_returns_passthrough() {
let merged = serde_json::json!({"name": "Alice", "age": 30});
let result = fuse(merged.clone(), &[], &[], false, 0.8, 0.95);
assert_eq!(result.structured_output, merged);
assert_eq!(result.structured_output_flat, merged);
}
#[test]
fn llm_envelope_passes_through() {
let merged = serde_json::json!({
"name": {
"value": "Alice",
"page": 1,
"bbox": [10.0, 20.0, 100.0, 30.0],
"source": "llm"
}
});
let result = fuse(merged.clone(), &[], &[], true, 0.8, 0.95);
assert!(result.structured_output.get("name").is_some());
}
#[test]
fn scalar_value_without_match_sets_source_none() {
let merged = serde_json::json!({"field": "unknown_value"});
let result = fuse(merged, &[], &[], true, 0.8, 0.95);
let field = result.structured_output.get("field").unwrap();
if let Ok(cited) = serde_json::from_value::<CitedField>(field.clone()) {
assert_eq!(cited.source, CitationSource::None);
assert_eq!(cited.value, serde_json::json!("unknown_value"));
}
}
#[test]
fn scalar_with_matching_ocr_fuses() {
let merged = serde_json::json!({"field": "Alice"});
let ocr = serde_json::json!({
"text": "Alice",
"page_number": 1,
"bbox": [10.0, 20.0, 100.0, 30.0]
});
let result = fuse(merged, &[ocr], &[], true, 0.8, 0.95);
let field = result.structured_output.get("field").unwrap();
if let Ok(cited) = serde_json::from_value::<CitedField>(field.clone()) {
assert_eq!(cited.source, CitationSource::Fused);
assert_eq!(cited.page, Some(1));
assert!(cited.bbox.is_some());
}
}
#[test]
fn flatten_cited_extracts_values_only() {
let cited = serde_json::json!({
"name": {
"value": "Alice",
"page": 1,
"source": "fused"
},
"age": {
"value": 30,
"source": "none"
}
});
let flattened = flatten_cited(&cited);
let name = flattened.get("name").unwrap();
let age = flattened.get("age").unwrap();
assert_eq!(name.as_str(), Some("Alice"));
assert_eq!(age.as_u64(), Some(30));
}
#[test]
fn text_similarity_identical_multibyte_strings_score_one() {
assert_eq!(text_similarity("世界", "世界"), 1.0);
assert_eq!(text_similarity("café", "café"), 1.0);
}
#[test]
fn scalar_with_matching_non_ascii_ocr_fuses() {
let merged = serde_json::json!({"field": "世界"});
let ocr = serde_json::json!({
"text": "世界",
"page_number": 2,
"bbox": [11.0, 22.0, 110.0, 33.0]
});
let result = fuse(merged, &[ocr], &[], true, 0.8, 0.95);
let field = result.structured_output.get("field").unwrap();
let cited: CitedField = serde_json::from_value(field.clone()).expect("field should deserialize as CitedField");
assert_eq!(cited.source, CitationSource::Fused);
assert_eq!(cited.page, Some(2));
assert_eq!(cited.bbox, Some([11.0, 22.0, 110.0, 33.0]));
assert_eq!(cited.confidence, Some(0.95));
}
#[test]
fn nested_objects_are_handled_recursively() {
let merged = serde_json::json!({
"person": {
"name": "Bob",
"contact": {
"email": "bob@example.com"
}
}
});
let result = fuse(merged, &[], &[], true, 0.8, 0.95);
assert!(result.structured_output.get("person").is_some());
assert!(
result
.structured_output
.get("person")
.and_then(|p| p.get("contact"))
.is_some()
);
}
}