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::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>,
hf_revision: Option<String>,
}
#[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 {
let trimmed = content.trim();
let stripped = trimmed.strip_prefix("$$").unwrap_or(trimmed).trim_start();
stripped.strip_suffix("$$").unwrap_or(stripped).trim_end().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_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) -> GlmOcrOptions {
let mut task = self.default_task;
let mut layout_mode = self.layout_mode;
let mut enable_chart_understanding = false;
let mut cache_dir = None;
let mut hf_revision = None;
if let Some(opts) = &config.backend_options {
if let Some(t) = opts.get("task").and_then(|v| v.as_str()) {
task = match t {
"table" => GlmOcrTask::Table,
"formula" => GlmOcrTask::Formula,
"chart" => GlmOcrTask::Chart,
"caption" => GlmOcrTask::Caption,
_ => GlmOcrTask::Ocr,
};
}
if let Some(m) = opts.get("layout_mode").and_then(|v| v.as_str()) {
layout_mode = match m {
#[cfg(feature = "layout-detection")]
"paired" => LayoutMode::Paired,
_ => LayoutMode::WholePage,
};
}
if let Some(e) = opts.get("enable_chart_understanding").and_then(|v| v.as_bool()) {
enable_chart_understanding = e;
}
cache_dir = opts
.get("cache_dir")
.and_then(|value| value.as_str())
.map(PathBuf::from);
hf_revision = opts
.get("hf_revision")
.or_else(|| opts.get("revision"))
.and_then(|value| value.as_str())
.map(str::to_string);
}
let device = super::resolve_device_preference(config);
GlmOcrOptions {
task,
device,
layout_mode,
enable_chart_understanding,
cache_dir,
hf_revision,
}
}
}
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 = image_bytes.to_vec();
let dtype = self.dtype;
let cache_dir = opts.cache_dir.unwrap_or_else(hf_hub::resolve_cache_dir);
let revision = opts.hf_revision.unwrap_or_else(|| GlmOcrEngine::revision().to_string());
let (content, formulas) = match opts.layout_mode {
LayoutMode::WholePage => {
let task = opts.task;
let device = opts.device;
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, 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())
}
#[cfg(feature = "layout-detection")]
LayoutMode::Paired => {
let enable_chart_understanding = opts.enable_chart_understanding;
process_paired(
image_bytes,
opts.device,
dtype,
enable_chart_understanding,
cache_dir,
revision,
)
.await?
}
};
Ok(ExtractedDocument {
content,
formulas,
mime_type: Cow::Borrowed("text/markdown"),
..Default::default()
})
}
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
}
}
#[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>)> {
use crate::layout::LayoutModelManager;
use crate::layout::models::LayoutModel;
tokio::task::spawn_blocking(move || {
let img = image::load_from_memory(&image_bytes)
.map_err(|e| crate::XbergError::Ocr {
message: format!("GLM-OCR paired: image decode failed: {e}"),
source: Some(Box::new(e)),
})?
.to_rgb8();
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>), crate::XbergError>((output.content, 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();
for detection in &sorted {
let bbox = &detection.bbox;
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 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: crate::types::extraction::BoundingBox {
x0: bbox.x1 as f64,
y0: bbox.y1 as f64,
x1: bbox.x2 as f64,
y1: bbox.y2 as f64,
},
page: 1,
});
}
parts.push(wrapped);
}
Ok::<(String, Vec<crate::types::Formula>), crate::XbergError>((parts.join("\n\n"), formulas))
})
.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()));
}
#[test]
fn test_parse_options_defaults() {
let config = OcrConfig::default();
let opts = GlmOcrBackend::new(GlmOcrTask::default(), LayoutMode::default()).parse_options(&config);
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);
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);
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);
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);
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);
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);
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);
assert_eq!(opts.task, GlmOcrTask::Chart);
assert_eq!(opts.device, DevicePreference::Cuda);
assert!(opts.enable_chart_understanding);
}
#[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);
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);
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());
}
}