use async_trait::async_trait;
use std::borrow::Cow;
use std::path::{Path, PathBuf};
use std::sync::{Arc, LazyLock};
use ahash::AHashMap;
use parking_lot::{Mutex, RwLock};
use crate::Result;
use crate::candle_ocr::config::{
GlmOcrBackendOptions, GlmOcrLayoutMode, GlmOcrTaskKind, parse_backend_options, validate_optional_non_empty,
};
use crate::core::config::OcrConfig;
use crate::plugins::{OcrBackend, OcrBackendType, Plugin};
use crate::types::ExtractedDocument;
use xberg_candle_ocr::CandleOcrError;
use xberg_candle_ocr::DType;
use xberg_candle_ocr::DevicePreference;
use xberg_candle_ocr::models::GlmOcrEngine;
use xberg_candle_ocr::models::GlmOcrTask;
#[cfg_attr(alef, alef(skip))]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LayoutMode {
WholePage,
#[cfg(feature = "layout-detection")]
Paired,
}
#[allow(clippy::derivable_impls)]
impl Default for LayoutMode {
fn default() -> Self {
#[cfg(feature = "layout-detection")]
{
LayoutMode::Paired
}
#[cfg(not(feature = "layout-detection"))]
{
LayoutMode::WholePage
}
}
}
type EnginePool = RwLock<AHashMap<(DevicePreference, DType, PathBuf, String), Arc<GlmOcrEngine>>>;
static ENGINE_POOL: LazyLock<EnginePool> = LazyLock::new(|| RwLock::new(AHashMap::new()));
#[cfg(feature = "layout-detection")]
type LayoutPool = RwLock<
AHashMap<(String, DevicePreference), Arc<Mutex<crate::layout::models::pp_doclayout_v3::PpDocLayoutV3Model>>>,
>;
#[cfg(feature = "layout-detection")]
static LAYOUT_POOL: LazyLock<LayoutPool> = LazyLock::new(|| RwLock::new(AHashMap::new()));
#[inline]
fn pool_get_or_init<K, V, E>(
pool: &RwLock<AHashMap<K, Arc<V>>>,
key: K,
init: impl FnOnce() -> std::result::Result<V, E>,
) -> std::result::Result<Arc<V>, E>
where
K: std::hash::Hash + Eq + Clone,
V: Send + 'static,
{
{
let pool_guard = pool.read();
if let Some(value) = pool_guard.get(&key) {
return Ok(Arc::clone(value));
}
}
let new_value = Arc::new(init()?);
let mut pool_guard = pool.write();
if let Some(existing) = pool_guard.get(&key) {
return Ok(Arc::clone(existing));
}
pool_guard.insert(key, Arc::clone(&new_value));
Ok(new_value)
}
fn get_or_init_engine(
preference: DevicePreference,
dtype: DType,
cache_dir: PathBuf,
revision: String,
) -> crate::Result<Arc<GlmOcrEngine>> {
let key = (preference, dtype, cache_dir.clone(), revision.clone());
pool_get_or_init::<(DevicePreference, DType, PathBuf, String), GlmOcrEngine, crate::XbergError>(
&ENGINE_POOL,
key,
|| {
let device = preference.select().map_err(|e| crate::XbergError::Ocr {
message: format!("Failed to select compute device: {e}"),
source: Some(Box::new(e)),
})?;
tracing::info!(
preference = ?preference,
?dtype,
"Initialising GLM-OCR engine (cold start)"
);
GlmOcrEngine::new_with_hf(GlmOcrTask::default(), device, dtype, Some(&cache_dir), Some(&revision)).map_err(
|e| crate::XbergError::Ocr {
message: format!("GLM-OCR engine initialisation failed: {e}"),
source: Some(Box::new(e)),
},
)
},
)
}
#[cfg(feature = "layout-detection")]
fn get_or_init_layout_model(
model_path: &Path,
device: DevicePreference,
) -> crate::Result<Arc<Mutex<crate::layout::models::pp_doclayout_v3::PpDocLayoutV3Model>>> {
use crate::layout::models::pp_doclayout_v3::PpDocLayoutV3Model;
let model_path_str = model_path
.to_str()
.ok_or_else(|| crate::XbergError::Ocr {
message: format!("Model path contains invalid UTF-8: {}", model_path.display()),
source: None,
})?
.to_string();
let key = (model_path_str.clone(), device);
pool_get_or_init::<(String, DevicePreference), Mutex<PpDocLayoutV3Model>, crate::XbergError>(
&LAYOUT_POOL,
key,
|| {
tracing::info!(
path = model_path_str.as_str(),
?device,
"Initialising PP-DocLayout-V3 model (cold start)"
);
PpDocLayoutV3Model::from_file(&model_path_str, None)
.map_err(|e| crate::XbergError::Ocr {
message: format!("PP-DocLayout-V3 model initialisation failed: {e}"),
source: Some(Box::new(e)),
})
.map(Mutex::new)
},
)
}
#[cfg_attr(alef, alef(skip))]
#[derive(Debug, Clone)]
struct GlmOcrOptions {
task: GlmOcrTask,
device: DevicePreference,
layout_mode: LayoutMode,
enable_chart_understanding: bool,
cache_dir: Option<PathBuf>,
}
#[cfg(feature = "layout-detection")]
fn task_for_label(label: crate::layout::LayoutClass, enable_chart_understanding: bool) -> GlmOcrTask {
use crate::layout::LayoutClass;
match label {
LayoutClass::Table => GlmOcrTask::Table,
LayoutClass::Formula => GlmOcrTask::Formula,
LayoutClass::Chart => {
if enable_chart_understanding {
GlmOcrTask::Chart
} else {
GlmOcrTask::Caption
}
}
LayoutClass::Picture => GlmOcrTask::Caption,
LayoutClass::Text
| LayoutClass::Title
| LayoutClass::SectionHeader
| LayoutClass::Caption
| LayoutClass::ListItem
| LayoutClass::Footnote
| LayoutClass::PageHeader
| LayoutClass::PageFooter
| LayoutClass::DocumentIndex
| LayoutClass::Code
| LayoutClass::CheckboxSelected
| LayoutClass::CheckboxUnselected
| LayoutClass::Form
| LayoutClass::KeyValueRegion => GlmOcrTask::Ocr,
}
}
#[cfg(feature = "layout-detection")]
fn strip_formula_delimiters(content: &str) -> String {
crate::extraction::derive::strip_math_delimiters(content).to_string()
}
fn wrap_output(task: GlmOcrTask, content: &str) -> String {
match task {
GlmOcrTask::Table => content.to_string(),
GlmOcrTask::Formula => format!("$$\n{}\n$$", content.trim()),
GlmOcrTask::Chart => format!("```json\n{}\n```", content.trim()),
GlmOcrTask::Ocr | GlmOcrTask::Caption => content.to_string(),
}
}
#[cfg(feature = "layout-detection")]
fn is_table_region(class: crate::layout::LayoutClass) -> bool {
class == crate::layout::LayoutClass::Table
}
#[cfg(feature = "layout-detection")]
fn merge_table_bounding_boxes(
tables: &mut [crate::types::Table],
detection_bboxes: &[crate::types::extraction::BoundingBox],
) {
for (table, bbox) in tables.iter_mut().zip(detection_bboxes.iter()) {
table.bounding_box = Some(*bbox);
}
}
#[cfg(feature = "layout-detection")]
fn checkbox_marker_for_class(class: crate::layout::LayoutClass) -> Option<&'static str> {
use crate::layout::LayoutClass;
match class {
LayoutClass::CheckboxSelected => Some("[x]"),
LayoutClass::CheckboxUnselected => Some("[ ]"),
_ => None,
}
}
#[cfg_attr(alef, alef(skip))]
pub struct GlmOcrBackend {
default_task: GlmOcrTask,
layout_mode: LayoutMode,
dtype: DType,
}
impl GlmOcrBackend {
pub fn new(default_task: GlmOcrTask, layout_mode: LayoutMode) -> Self {
Self {
default_task,
layout_mode,
dtype: DType::F32,
}
}
pub fn with_dtype(mut self, dtype: DType) -> Self {
self.dtype = dtype;
self
}
fn parse_options(&self, config: &OcrConfig) -> Result<GlmOcrOptions> {
let options: GlmOcrBackendOptions = parse_backend_options(config.backend_options.as_ref(), "candle-glm-ocr")?;
validate_optional_non_empty(options.cache_dir.as_deref(), "candle-glm-ocr", "cache_dir")?;
let task = match options.task {
None => self.default_task,
Some(GlmOcrTaskKind::Ocr) => GlmOcrTask::Ocr,
Some(GlmOcrTaskKind::Table) => GlmOcrTask::Table,
Some(GlmOcrTaskKind::Formula) => GlmOcrTask::Formula,
Some(GlmOcrTaskKind::Chart) => GlmOcrTask::Chart,
Some(GlmOcrTaskKind::Caption) => GlmOcrTask::Caption,
};
let layout_mode = match options.layout_mode {
None => self.layout_mode,
Some(GlmOcrLayoutMode::WholePage) => LayoutMode::WholePage,
#[cfg(feature = "layout-detection")]
Some(GlmOcrLayoutMode::Paired) => LayoutMode::Paired,
#[cfg(not(feature = "layout-detection"))]
Some(GlmOcrLayoutMode::Paired) => {
return Err(crate::XbergError::validation(
"invalid candle-glm-ocr backend_options.layout_mode: paired requires the layout-detection feature"
.to_string(),
));
}
};
Ok(GlmOcrOptions {
task,
device: super::resolve_device_preference(config, options.device),
layout_mode,
enable_chart_understanding: options.enable_chart_understanding.unwrap_or(false),
cache_dir: options.cache_dir.map(PathBuf::from),
})
}
}
impl Plugin for GlmOcrBackend {
fn name(&self) -> &str {
"candle-glm-ocr"
}
fn version(&self) -> String {
"0.1.0".to_string()
}
fn initialize(&self) -> Result<()> {
tracing::debug!(
task = %self.default_task,
"Initializing GLM-OCR backend"
);
Ok(())
}
fn shutdown(&self) -> Result<()> {
Ok(())
}
}
#[async_trait]
impl OcrBackend for GlmOcrBackend {
async fn process_image(&self, image_bytes: &[u8], config: &OcrConfig) -> Result<ExtractedDocument> {
let opts = self.parse_options(config)?;
if image_bytes.is_empty() {
return Err(crate::XbergError::Validation {
message: "Empty image data provided to GLM-OCR".to_string(),
source: None,
});
}
let image_bytes_owned = image_bytes.to_vec();
let dtype = self.dtype;
let cache_dir = opts.cache_dir.unwrap_or_else(hf_hub::resolve_cache_dir);
let revision = GlmOcrEngine::revision().to_string();
let (content, formulas, table_bboxes) = match opts.layout_mode {
LayoutMode::WholePage => {
let task = opts.task;
let device = opts.device;
crate::extraction::image_decode::validate_standard_image_with_default_security_limits(
&image_bytes_owned,
)?;
let content = tokio::task::spawn_blocking(move || {
let engine = get_or_init_engine(device, dtype, cache_dir, revision)?;
let output = engine.process_image_with_task(&image_bytes_owned, task).map_err(|e| {
crate::XbergError::Ocr {
message: format!("GLM-OCR inference failed: {e}"),
source: Some(Box::new(e)),
}
})?;
Ok::<String, crate::XbergError>(output.content)
})
.await
.map_err(|e| crate::XbergError::Ocr {
message: format!("GLM-OCR task execution failed: {e}"),
source: None,
})??;
(content, Vec::new(), Vec::new())
}
#[cfg(feature = "layout-detection")]
LayoutMode::Paired => {
let enable_chart_understanding = opts.enable_chart_understanding;
process_paired(
image_bytes_owned,
opts.device,
dtype,
enable_chart_understanding,
cache_dir,
revision,
)
.await?
}
};
let mut document = super::ocr_result::build_ocr_document(
content,
formulas,
Cow::Borrowed("text/markdown"),
image_bytes,
config,
"candle-glm-ocr",
);
#[cfg(feature = "layout-detection")]
merge_table_bounding_boxes(&mut document.tables, &table_bboxes);
#[cfg(not(feature = "layout-detection"))]
let _ = table_bboxes;
Ok(document)
}
async fn process_image_file(&self, path: &Path, config: &OcrConfig) -> Result<ExtractedDocument> {
let bytes = crate::core::io::read_file_async(path).await?;
self.process_image(&bytes, config).await
}
fn supports_language(&self, _lang: &str) -> bool {
true
}
fn supported_languages(&self) -> Vec<String> {
vec![
"eng", "en", "zho", "zh", "jpn", "ja", "kor", "ko", "fra", "fr", "deu", "de", "spa", "es", "ita", "it",
"por", "pt", "rus", "ru", "ara", "ar", "hin", "hi", "tha", "th", "vie", "vi",
]
.iter()
.map(|s| s.to_string())
.collect()
}
fn backend_type(&self) -> OcrBackendType {
OcrBackendType::Candle
}
fn emits_structured_markdown(&self) -> bool {
true
}
fn confidence_semantics(&self) -> crate::plugins::ConfidenceSemantics {
crate::plugins::ConfidenceSemantics::None
}
}
#[cfg(feature = "layout-detection")]
async fn process_paired(
image_bytes: Vec<u8>,
device: DevicePreference,
dtype: DType,
enable_chart_understanding: bool,
cache_dir: PathBuf,
revision: String,
) -> crate::Result<(
String,
Vec<crate::types::Formula>,
Vec<crate::types::extraction::BoundingBox>,
)> {
use crate::layout::LayoutModelManager;
use crate::layout::models::LayoutModel;
const CROP_PNG_ENCODE_BYTES_PER_PIXEL: u64 = 4;
const CROP_PNG_ENCODE_FIXED_BYTES: u64 = 256 * 1024;
tokio::task::spawn_blocking(move || {
let img = crate::extraction::image_decode::decode_standard_rgb8_with_default_security_limits(&image_bytes)
.map_err(|error| crate::XbergError::Ocr {
message: format!("GLM-OCR paired: image decode failed: {error}"),
source: Some(Box::new(error)),
})?;
let security_limits = crate::extractors::security::SecurityLimits::default();
crate::layout::engine::validate_layout_batch_peak(&[&img], &security_limits)?;
let manager = LayoutModelManager::new(None);
let model_path = manager
.ensure_pp_doclayout_v3_model()
.map_err(|e| crate::XbergError::Ocr {
message: format!("GLM-OCR paired: layout model unavailable: {e}"),
source: Some(Box::new(e)),
})?;
let layout_model = get_or_init_layout_model(&model_path, device).map_err(|e| crate::XbergError::Ocr {
message: format!("GLM-OCR paired: layout detection init failed: {e}"),
source: Some(Box::new(e)),
})?;
let detections = layout_model.lock().detect(&img).map_err(|e| crate::XbergError::Ocr {
message: format!("GLM-OCR paired: layout detection failed: {e}"),
source: Some(Box::new(e)),
})?;
let mut sorted = detections;
sorted.sort_by(|a, b| a.bbox.y1.total_cmp(&b.bbox.y1).then(a.bbox.x1.total_cmp(&b.bbox.x1)));
let engine = get_or_init_engine(device, dtype, cache_dir, revision)?;
if sorted.is_empty() {
tracing::debug!("GLM-OCR paired: no layout regions detected, falling back to whole-page inference");
let output = engine
.process_image_with_task(&image_bytes, GlmOcrTask::Ocr)
.map_err(|e| crate::XbergError::Ocr {
message: format!("GLM-OCR paired (fallback): whole-page inference failed: {e}"),
source: Some(Box::new(e)),
})?;
return Ok::<
(
String,
Vec<crate::types::Formula>,
Vec<crate::types::extraction::BoundingBox>,
),
crate::XbergError,
>((output.content, Vec::new(), Vec::new()));
}
let img_width = img.width();
let img_height = img.height();
let mut parts: Vec<String> = Vec::with_capacity(sorted.len());
let mut formulas: Vec<crate::types::Formula> = Vec::new();
let mut table_bboxes: Vec<crate::types::extraction::BoundingBox> = Vec::new();
for detection in &sorted {
let bbox = &detection.bbox;
let region_bbox = crate::types::extraction::BoundingBox {
x0: bbox.x1 as f64,
y0: bbox.y1 as f64,
x1: bbox.x2 as f64,
y1: bbox.y2 as f64,
};
if let Some(marker) = checkbox_marker_for_class(detection.class_name) {
parts.push(marker.to_string());
continue;
}
let x = (bbox.x1.max(0.0) as u32).min(img_width.saturating_sub(1));
let y = (bbox.y1.max(0.0) as u32).min(img_height.saturating_sub(1));
let w = ((bbox.x2 - bbox.x1).max(1.0) as u32).min(img_width - x);
let h = ((bbox.y2 - bbox.y1).max(1.0) as u32).min(img_height - y);
let current_bytes = u64::try_from(img.as_raw().len()).map_err(|_| {
crate::extraction::image_decode::image_dimension_error(img_width, img_height, u64::MAX, u64::MAX)
})?;
let crop_and_encode_bytes =
crate::extraction::image_decode::decoded_byte_count(w, h, 3 + CROP_PNG_ENCODE_BYTES_PER_PIXEL)?
.checked_add(CROP_PNG_ENCODE_FIXED_BYTES)
.ok_or_else(|| {
crate::extraction::image_decode::image_dimension_error(
img_width,
img_height,
u64::MAX,
u64::MAX,
)
})?;
crate::extraction::image_decode::validate_image_live_bytes(
img_width,
img_height,
current_bytes,
crop_and_encode_bytes,
&security_limits,
)?;
let crop = image::imageops::crop_imm(&img, x, y, w, h).to_image();
let mut crop_bytes: Vec<u8> = Vec::new();
crop.write_to(&mut std::io::Cursor::new(&mut crop_bytes), image::ImageFormat::Png)
.map_err(|e| crate::XbergError::Ocr {
message: format!("GLM-OCR paired: crop encode failed: {e}"),
source: Some(Box::new(e)),
})?;
let region_task = task_for_label(detection.class_name, enable_chart_understanding);
let output = match engine.process_image_with_task(&crop_bytes, region_task) {
Ok(out) => out,
Err(CandleOcrError::UnsupportedConfig(ref msg)) => {
tracing::warn!(
class = ?detection.class_name,
bbox = ?bbox,
reason = %msg,
"GLM-OCR paired: skipping region (unsupported config)"
);
continue;
}
Err(e) => {
return Err(crate::XbergError::Ocr {
message: format!("GLM-OCR paired: region inference failed: {e}"),
source: Some(Box::new(e)),
});
}
};
let latex_clean = if detection.class_name == crate::layout::LayoutClass::Formula {
strip_formula_delimiters(&output.content)
} else {
output.content.clone()
};
let wrapped = wrap_output(region_task, &latex_clean);
if detection.class_name == crate::layout::LayoutClass::Formula && !latex_clean.is_empty() {
formulas.push(crate::types::Formula {
latex: latex_clean,
bbox: Some(region_bbox),
page: Some(1),
});
} else if is_table_region(detection.class_name) && !output.content.trim().is_empty() {
table_bboxes.push(region_bbox);
}
parts.push(wrapped);
}
Ok::<
(
String,
Vec<crate::types::Formula>,
Vec<crate::types::extraction::BoundingBox>,
),
crate::XbergError,
>((parts.join("\n\n"), formulas, table_bboxes))
})
.await
.map_err(|e| crate::XbergError::Ocr {
message: format!("GLM-OCR paired task execution failed: {e}"),
source: Some(Box::new(e)),
})?
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_glm_ocr_backend_creation() {
let backend = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default());
assert_eq!(backend.name(), "candle-glm-ocr");
assert_eq!(backend.backend_type(), OcrBackendType::Candle);
}
#[test]
fn test_glm_ocr_emits_structured_markdown() {
let backend = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default());
assert!(backend.emits_structured_markdown());
}
#[test]
fn test_glm_ocr_language_support() {
let backend = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default());
assert!(backend.supports_language("eng"));
assert!(backend.supports_language("zho"));
assert!(backend.supports_language("jpn"));
assert!(backend.supports_language("unknown"));
}
#[test]
fn test_glm_ocr_supported_languages() {
let backend = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default());
let langs = backend.supported_languages();
assert!(langs.contains(&"eng".to_string()));
assert!(langs.contains(&"zho".to_string()));
assert!(langs.contains(&"fra".to_string()));
}
#[tokio::test]
async fn should_reject_oversized_declared_dimensions_before_loading_glm_model() {
let backend = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::WholePage);
let bytes = crate::extraction::image_decode::bmp_with_declared_dimensions(6000, 6000);
let error = backend
.process_image(&bytes, &OcrConfig::default())
.await
.expect_err("GLM-OCR must validate the decoded-byte budget before model initialization");
assert!(matches!(error, crate::XbergError::Validation { .. }));
assert!(error.to_string().contains("6000x6000"));
assert!(error.to_string().contains("security_limits.max_content_size"));
}
#[test]
fn test_parse_options_defaults() {
let config = OcrConfig::default();
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
assert_eq!(opts.task, GlmOcrTask::Ocr);
assert_eq!(opts.device, DevicePreference::Auto);
assert!(!opts.enable_chart_understanding);
}
#[test]
fn test_parse_options_custom_task() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({"task": "table"})),
..Default::default()
};
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
assert_eq!(opts.task, GlmOcrTask::Table);
}
#[test]
fn test_parse_options_formula_task() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({"task": "formula"})),
..Default::default()
};
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
assert_eq!(opts.task, GlmOcrTask::Formula);
}
#[test]
fn test_parse_options_custom_device() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({"device": "cpu"})),
..Default::default()
};
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
assert_eq!(opts.device, DevicePreference::Cpu);
}
#[test]
fn test_parse_options_enable_chart_understanding_true() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({"enable_chart_understanding": true})),
..Default::default()
};
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
assert!(opts.enable_chart_understanding);
}
#[test]
fn test_parse_options_enable_chart_understanding_false() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({"enable_chart_understanding": false})),
..Default::default()
};
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
assert!(!opts.enable_chart_understanding);
}
#[test]
fn test_parse_options_chart_understanding_default() {
let config = OcrConfig::default();
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
assert!(!opts.enable_chart_understanding);
}
#[test]
fn test_parse_options_combined() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({
"task": "chart",
"device": "cuda",
"enable_chart_understanding": true
})),
..Default::default()
};
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
assert_eq!(opts.task, GlmOcrTask::Chart);
assert_eq!(opts.device, DevicePreference::Cuda);
assert!(opts.enable_chart_understanding);
}
#[test]
fn test_parse_options_non_object_json_returns_contextual_errors() {
let config = OcrConfig {
backend_options: Some(serde_json::json!([1, 2, 3])),
..Default::default()
};
let error = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap_err()
.to_string();
assert!(error.contains("candle-glm-ocr backend_options"));
let config = OcrConfig {
backend_options: Some(serde_json::json!("ocr")),
..Default::default()
};
let error = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap_err()
.to_string();
assert!(error.contains("candle-glm-ocr backend_options"));
}
#[test]
fn test_parse_options_empty_object_returns_defaults() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({})),
..Default::default()
};
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
assert_eq!(opts.task, GlmOcrTask::Ocr);
assert_eq!(opts.device, DevicePreference::Auto);
assert!(!opts.enable_chart_understanding);
assert!(opts.cache_dir.is_none());
}
#[test]
fn test_parse_options_layout_mode_and_cache_dir() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({
"layout_mode": "paired",
"cache_dir": "/tmp/glm-cache"
})),
..Default::default()
};
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
assert_eq!(opts.layout_mode, LayoutMode::Paired);
assert_eq!(opts.cache_dir.as_deref(), Some(Path::new("/tmp/glm-cache")));
}
#[test]
fn test_initialize_and_shutdown() {
let backend = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default());
assert!(backend.initialize().is_ok());
assert!(backend.shutdown().is_ok());
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_task_for_label_table() {
use crate::layout::LayoutClass;
assert_eq!(task_for_label(LayoutClass::Table, false), GlmOcrTask::Table);
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_task_for_label_formula() {
use crate::layout::LayoutClass;
assert_eq!(task_for_label(LayoutClass::Formula, false), GlmOcrTask::Formula);
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_task_for_label_text() {
use crate::layout::LayoutClass;
assert_eq!(task_for_label(LayoutClass::Text, false), GlmOcrTask::Ocr);
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_task_for_label_chart_disabled() {
use crate::layout::LayoutClass;
assert_eq!(task_for_label(LayoutClass::Chart, false), GlmOcrTask::Caption);
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_task_for_label_chart_enabled() {
use crate::layout::LayoutClass;
assert_eq!(task_for_label(LayoutClass::Chart, true), GlmOcrTask::Chart);
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_parse_and_route_chart_with_understanding_enabled() {
use crate::layout::LayoutClass;
let config = OcrConfig {
backend_options: Some(serde_json::json!({"enable_chart_understanding": true})),
..Default::default()
};
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
let routed_task = task_for_label(LayoutClass::Chart, opts.enable_chart_understanding);
assert_eq!(routed_task, GlmOcrTask::Chart);
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_parse_and_route_chart_with_understanding_disabled() {
use crate::layout::LayoutClass;
let config = OcrConfig {
backend_options: Some(serde_json::json!({"enable_chart_understanding": false})),
..Default::default()
};
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default())
.parse_options(&config)
.unwrap();
let routed_task = task_for_label(LayoutClass::Chart, opts.enable_chart_understanding);
assert_eq!(routed_task, GlmOcrTask::Caption);
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_wrap_output_formula() {
let wrapped = wrap_output(GlmOcrTask::Formula, "x^2 + y^2 = r^2");
assert!(wrapped.starts_with("$$\n"));
assert!(wrapped.ends_with("\n$$"));
assert!(wrapped.contains("x^2 + y^2 = r^2"));
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_strip_formula_delimiters_removes_wrapping_dollars() {
let wrapped = "$$\nE = mc^2\n$$";
let result = strip_formula_delimiters(wrapped);
assert_eq!(result, "E = mc^2");
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_strip_formula_delimiters_handles_pre_wrapped_content() {
let pre_wrapped = "$$x^2 + y^2 = z^2$$";
let result = strip_formula_delimiters(pre_wrapped);
assert_eq!(result, "x^2 + y^2 = z^2");
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_strip_formula_delimiters_preserves_undecorated_content() {
let plain = "a + b = c";
let result = strip_formula_delimiters(plain);
assert_eq!(result, "a + b = c");
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_formula_extraction_from_wrapped_output() {
let task = GlmOcrTask::Formula;
let raw_latex = "E = mc^2";
let wrapped = wrap_output(task, raw_latex);
let stripped = strip_formula_delimiters(&wrapped);
assert_eq!(stripped, raw_latex);
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_wrap_output_chart() {
let wrapped = wrap_output(GlmOcrTask::Chart, r#"{"type":"bar"}"#);
assert!(wrapped.starts_with("```json\n"));
assert!(wrapped.ends_with("\n```"));
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_wrap_output_table_passthrough() {
let table = "| A | B |\n|---|---|\n| 1 | 2 |";
let wrapped = wrap_output(GlmOcrTask::Table, table);
assert_eq!(wrapped, table);
}
#[test]
fn test_pool_get_or_init_caches_on_first_miss() {
use std::sync::Arc as StdArc;
use std::sync::atomic::{AtomicUsize, Ordering};
let pool = RwLock::new(AHashMap::new());
let init_count = StdArc::new(AtomicUsize::new(0));
let init_count_clone = StdArc::clone(&init_count);
let result1 = pool_get_or_init(&pool, "test_key", || {
init_count_clone.fetch_add(1, Ordering::SeqCst);
Ok::<u32, String>(42)
});
assert!(result1.is_ok());
assert_eq!(init_count.load(Ordering::SeqCst), 1, "Initializer should run once");
let init_count_clone = StdArc::clone(&init_count);
let result2 = pool_get_or_init(&pool, "test_key", || {
init_count_clone.fetch_add(1, Ordering::SeqCst);
Ok::<u32, String>(99)
});
assert!(result2.is_ok());
assert_eq!(
init_count.load(Ordering::SeqCst),
1,
"Initializer should still have run exactly once"
);
let v1 = result1.unwrap();
let v2 = result2.unwrap();
assert!(Arc::ptr_eq(&v1, &v2), "Cached values should be the same Arc instance");
assert_eq!(*v1, 42, "First initializer's value should be stored");
}
#[test]
fn test_pool_get_or_init_concurrent_access() {
use std::sync::Arc as StdArc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::thread;
let pool = StdArc::new(RwLock::new(AHashMap::new()));
let init_count = StdArc::new(AtomicUsize::new(0));
let mut handles = vec![];
for _ in 0..5 {
let pool_clone = StdArc::clone(&pool);
let init_count_clone = StdArc::clone(&init_count);
let handle = thread::spawn(move || {
let result = pool_get_or_init(&pool_clone, "concurrent_key", || {
init_count_clone.fetch_add(1, Ordering::SeqCst);
std::thread::sleep(std::time::Duration::from_millis(1));
Ok::<u32, String>(42)
});
result.unwrap()
});
handles.push(handle);
}
let results: Vec<_> = handles.into_iter().map(|h| h.join().unwrap()).collect();
for i in 1..results.len() {
assert!(
Arc::ptr_eq(&results[0], &results[i]),
"All concurrent callers should receive the same Arc instance"
);
}
let final_count = init_count.load(Ordering::SeqCst);
assert!(final_count >= 1, "Initializer must run at least once");
}
#[test]
fn test_glm_ocr_zero_regions_fallback_guard() {
assert_eq!(GlmOcrTask::Ocr, GlmOcrTask::default());
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_is_table_region_true_for_table_class() {
use crate::layout::LayoutClass;
assert!(is_table_region(LayoutClass::Table));
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_is_table_region_false_for_text_class() {
use crate::layout::LayoutClass;
assert!(!is_table_region(LayoutClass::Text));
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_merge_table_bounding_boxes_attaches_bbox_matching_detection_order() {
use crate::types::Table;
use crate::types::extraction::BoundingBox;
let mut tables = vec![
Table {
cells: vec![vec!["A".to_string()]],
markdown: "| A |\n|---|".to_string(),
page_number: 1,
..Default::default()
},
Table {
cells: vec![vec!["B".to_string()]],
markdown: "| B |\n|---|".to_string(),
page_number: 1,
..Default::default()
},
];
let bboxes = vec![
BoundingBox {
x0: 10.0,
y0: 20.0,
x1: 100.0,
y1: 200.0,
},
BoundingBox {
x0: 5.0,
y0: 6.0,
x1: 7.0,
y1: 8.0,
},
];
merge_table_bounding_boxes(&mut tables, &bboxes);
assert_eq!(tables.len(), 2, "table count must be unchanged by the merge");
assert_eq!(
tables[0].bounding_box,
Some(BoundingBox {
x0: 10.0,
y0: 20.0,
x1: 100.0,
y1: 200.0
})
);
assert_eq!(
tables[1].bounding_box,
Some(BoundingBox {
x0: 5.0,
y0: 6.0,
x1: 7.0,
y1: 8.0
})
);
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_process_paired_table_detection_yields_exactly_one_table_with_source_bbox() {
use crate::core::config::OcrConfig;
use crate::layout::LayoutClass;
use crate::types::extraction::BoundingBox;
let table_bbox = BoundingBox {
x0: 12.0,
y0: 34.0,
x1: 512.0,
y1: 734.0,
};
let table_output = "| Name | Age |\n|------|-----|\n| Alice | 30 |";
let text_output = "Just some plain OCR'd prose.";
let regions: Vec<(LayoutClass, BoundingBox, &str)> = vec![
(LayoutClass::Table, table_bbox, table_output),
(LayoutClass::Text, BoundingBox::default(), text_output),
];
let mut parts: Vec<String> = Vec::with_capacity(regions.len());
let mut table_bboxes: Vec<BoundingBox> = Vec::new();
for (class, bbox, output) in ®ions {
if is_table_region(*class) && !output.trim().is_empty() {
table_bboxes.push(*bbox);
}
parts.push(output.to_string());
}
let content = parts.join("\n\n");
let config = OcrConfig::default();
let mut doc = super::super::ocr_result::build_ocr_document(
content,
Vec::new(),
std::borrow::Cow::Borrowed("text/markdown"),
&[],
&config,
"candle-glm-ocr",
);
merge_table_bounding_boxes(&mut doc.tables, &table_bboxes);
assert_eq!(
doc.tables.len(),
1,
"the Text-class region must not produce a table entry"
);
assert_eq!(doc.tables[0].bounding_box, Some(table_bbox));
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_checkbox_marker_for_class_selected_is_x_marker() {
use crate::layout::LayoutClass;
assert_eq!(checkbox_marker_for_class(LayoutClass::CheckboxSelected), Some("[x]"));
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_checkbox_marker_for_class_unselected_is_empty_marker() {
use crate::layout::LayoutClass;
assert_eq!(checkbox_marker_for_class(LayoutClass::CheckboxUnselected), Some("[ ]"));
}
#[cfg(feature = "layout-detection")]
#[test]
fn test_checkbox_marker_for_class_none_for_non_checkbox_class() {
use crate::layout::LayoutClass;
assert_eq!(checkbox_marker_for_class(LayoutClass::Text), None);
assert_eq!(checkbox_marker_for_class(LayoutClass::Table), None);
}
}