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::{
PaddleOcrVlBackendOptions, PaddleOcrVlTaskKind, 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::DType;
use xberg_candle_ocr::DevicePreference;
use xberg_candle_ocr::models::PaddleOcrVlEngine;
use xberg_candle_ocr::models::PaddleOcrVlTask;
type PoolKey = (String, PaddleOcrVlTask, DevicePreference);
type PooledEngine = Arc<Mutex<PaddleOcrVlEngine>>;
static ENGINE_POOL: LazyLock<RwLock<AHashMap<PoolKey, PooledEngine>>> = LazyLock::new(|| RwLock::new(AHashMap::new()));
fn get_or_init_engine(
model_path: &str,
task: PaddleOcrVlTask,
preference: DevicePreference,
) -> crate::Result<PooledEngine> {
let key: PoolKey = (model_path.to_string(), task, preference);
{
let pool = ENGINE_POOL.read();
if let Some(engine) = pool.get(&key) {
return Ok(Arc::clone(engine));
}
}
let candle_device = preference.select().map_err(|e| crate::XbergError::Ocr {
message: format!("Failed to select compute device: {e}"),
source: Some(Box::new(e)),
})?;
tracing::info!(
task = ?task,
preference = ?preference,
"Initialising PaddleOCR-VL engine (cold start)"
);
let new_engine =
PaddleOcrVlEngine::new(model_path, task, candle_device, DType::F32).map_err(|e| crate::XbergError::Ocr {
message: format!("PaddleOCR-VL engine initialisation failed: {e}"),
source: Some(Box::new(e)),
})?;
let new_engine = Arc::new(Mutex::new(new_engine));
let mut pool = ENGINE_POOL.write();
if let Some(existing) = pool.get(&key) {
return Ok(Arc::clone(existing));
}
pool.insert(key, Arc::clone(&new_engine));
Ok(new_engine)
}
const DEFAULT_MODEL_ID: &str = "xberg-io/paddleocr-vl-1.6";
#[cfg_attr(alef, alef(skip))]
pub struct PaddleOcrVlBackend {
task: PaddleOcrVlTask,
}
#[derive(Debug)]
struct PaddleOcrVlOptions {
task: PaddleOcrVlTask,
model_path: Option<String>,
model_id: String,
hf_revision: Option<String>,
cache_dir: Option<PathBuf>,
device: DevicePreference,
}
impl PaddleOcrVlBackend {
pub fn new(task: PaddleOcrVlTask) -> Self {
Self { task }
}
pub fn default_task() -> Self {
Self::new(PaddleOcrVlTask::default())
}
fn parse_options(&self, config: &OcrConfig) -> Result<PaddleOcrVlOptions> {
let options: PaddleOcrVlBackendOptions =
parse_backend_options(config.backend_options.as_ref(), "candle-paddleocr-vl")?;
for (field, value) in [
("model_path", options.model_path.as_deref()),
("model_id", options.model_id.as_deref()),
("hf_revision", options.hf_revision.as_deref()),
("cache_dir", options.cache_dir.as_deref()),
] {
validate_optional_non_empty(value, "candle-paddleocr-vl", field)?;
}
let task = match options.task {
Some(PaddleOcrVlTaskKind::Ocr) => PaddleOcrVlTask::Ocr,
Some(PaddleOcrVlTaskKind::Table) => PaddleOcrVlTask::Table,
Some(PaddleOcrVlTaskKind::Formula) => PaddleOcrVlTask::Formula,
Some(PaddleOcrVlTaskKind::Chart) => PaddleOcrVlTask::Chart,
None => self.task,
};
Ok(PaddleOcrVlOptions {
task,
model_path: options.model_path,
model_id: options.model_id.unwrap_or_else(|| DEFAULT_MODEL_ID.to_string()),
hf_revision: options.hf_revision,
cache_dir: options.cache_dir.map(PathBuf::from),
device: super::resolve_device_preference(config, options.device),
})
}
}
impl Plugin for PaddleOcrVlBackend {
fn name(&self) -> &str {
"candle-paddleocr-vl"
}
fn version(&self) -> String {
"0.1.0".to_string()
}
fn initialize(&self) -> Result<()> {
tracing::debug!("Initializing PaddleOCR-VL backend: {} task", self.task);
Ok(())
}
fn shutdown(&self) -> Result<()> {
Ok(())
}
}
#[async_trait]
impl OcrBackend for PaddleOcrVlBackend {
async fn process_image(&self, image_bytes: &[u8], config: &OcrConfig) -> Result<ExtractedDocument> {
let options = self.parse_options(config)?;
if image_bytes.is_empty() {
return Err(crate::XbergError::Validation {
message: "Empty image data provided to PaddleOCR-VL".to_string(),
source: None,
});
}
let image_bytes_owned = image_bytes.to_vec();
let content = tokio::task::spawn_blocking(move || {
let model_path = match options.model_path {
Some(p) => p,
None => super::model_stager::ensure_paddleocr_vl_16(
&options.model_id,
options.hf_revision.as_deref(),
options.cache_dir.as_deref(),
)
.map(|dir| dir.to_string_lossy().into_owned())
.map_err(|e| crate::XbergError::Ocr {
message: format!("PaddleOCR-VL weight download failed: {e}"),
source: None,
})?,
};
let engine = get_or_init_engine(&model_path, options.task, options.device)?;
let mut engine_guard = engine.lock();
let output = engine_guard
.process_image(&image_bytes_owned)
.map_err(|e| crate::XbergError::Ocr {
message: format!("PaddleOCR-VL inference failed: {e}"),
source: Some(Box::new(e)),
})?;
Ok::<String, crate::XbergError>(output.content)
})
.await
.map_err(|e| crate::XbergError::Ocr {
message: format!("PaddleOCR-VL task execution failed: {e}"),
source: None,
})??;
Ok(super::ocr_result::build_ocr_document(
content,
Vec::new(),
Cow::Borrowed("text/markdown"),
image_bytes,
config,
"candle-paddleocr-vl",
))
}
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(test)]
mod tests {
use super::*;
fn parse_default_options(config: &OcrConfig) -> Result<PaddleOcrVlOptions> {
PaddleOcrVlBackend::default_task().parse_options(config)
}
#[test]
fn test_paddleocr_vl_backend_creation() {
let backend = PaddleOcrVlBackend::default_task();
assert_eq!(backend.name(), "candle-paddleocr-vl");
assert_eq!(backend.backend_type(), OcrBackendType::Candle);
}
#[test]
fn test_paddleocr_vl_emits_structured_markdown() {
let backend = PaddleOcrVlBackend::default_task();
assert!(backend.emits_structured_markdown());
}
#[test]
fn test_paddleocr_vl_language_support() {
let backend = PaddleOcrVlBackend::default_task();
assert!(backend.supports_language("eng"));
assert!(backend.supports_language("zho"));
assert!(backend.supports_language("jpn"));
assert!(backend.supports_language("fra"));
assert!(backend.supports_language("unknown"));
}
#[test]
fn test_paddleocr_vl_supported_languages() {
let backend = PaddleOcrVlBackend::default_task();
let langs = backend.supported_languages();
assert!(langs.contains(&"eng".to_string()));
assert!(langs.contains(&"zho".to_string()));
assert!(langs.contains(&"jpn".to_string()));
}
#[test]
fn test_parse_options_defaults() {
let config = OcrConfig::default();
let options = parse_default_options(&config).unwrap();
assert_eq!(options.task, PaddleOcrVlTask::Ocr);
assert!(options.model_path.is_none());
assert_eq!(options.model_id, DEFAULT_MODEL_ID);
assert!(options.hf_revision.is_none());
assert!(options.cache_dir.is_none());
assert_eq!(options.device, DevicePreference::Auto);
}
#[test]
fn test_parse_options_custom_task() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({
"task": "table"
})),
..Default::default()
};
let options = parse_default_options(&config).unwrap();
assert_eq!(options.task, PaddleOcrVlTask::Table);
}
#[test]
fn should_use_constructor_task_when_backend_options_omit_task() {
let backend = PaddleOcrVlBackend::new(PaddleOcrVlTask::Table);
let options = backend.parse_options(&OcrConfig::default()).unwrap();
assert_eq!(options.task, backend.task);
}
#[test]
fn should_prefer_explicit_task_over_constructor_task() {
let backend = PaddleOcrVlBackend::new(PaddleOcrVlTask::Chart);
let config = OcrConfig {
backend_options: Some(serde_json::json!({
"task": "formula"
})),
..Default::default()
};
let options = backend.parse_options(&config).unwrap();
assert_eq!(backend.task, PaddleOcrVlTask::Chart);
assert_eq!(options.task, PaddleOcrVlTask::Formula);
}
#[test]
fn test_parse_options_custom_device() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({
"device": "cpu"
})),
..Default::default()
};
let options = parse_default_options(&config).unwrap();
assert_eq!(options.device, DevicePreference::Cpu);
}
#[test]
fn test_parse_options_model_path() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({
"model_path": "/models/paddleocr-vl"
})),
..Default::default()
};
let options = parse_default_options(&config).unwrap();
assert_eq!(options.model_path.as_deref(), Some("/models/paddleocr-vl"));
}
#[test]
fn test_parse_options_custom_model_id() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({
"model_id": "some-org/custom-paddleocr-vl"
})),
..Default::default()
};
let options = parse_default_options(&config).unwrap();
assert!(options.model_path.is_none());
assert_eq!(options.model_id, "some-org/custom-paddleocr-vl");
}
#[test]
fn test_parse_options_hf_cache_and_revision() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({
"hf_revision": "0123456789abcdef",
"cache_dir": "/tmp/hf-hub"
})),
..Default::default()
};
let options = parse_default_options(&config).unwrap();
assert_eq!(options.hf_revision.as_deref(), Some("0123456789abcdef"));
assert_eq!(options.cache_dir.as_deref(), Some(Path::new("/tmp/hf-hub")));
}
#[test]
fn test_parse_options_non_object_json_returns_contextual_error() {
let config = OcrConfig {
backend_options: Some(serde_json::json!(false)),
..Default::default()
};
let error = parse_default_options(&config).unwrap_err().to_string();
assert!(error.contains("candle-paddleocr-vl backend_options"));
}
#[test]
fn test_parse_options_empty_object_returns_defaults() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({})),
..Default::default()
};
let options = parse_default_options(&config).unwrap();
assert_eq!(options.task, PaddleOcrVlTask::Ocr);
assert!(options.model_path.is_none());
assert_eq!(options.model_id, DEFAULT_MODEL_ID);
assert!(options.hf_revision.is_none());
assert!(options.cache_dir.is_none());
assert_eq!(options.device, DevicePreference::Auto);
}
#[test]
fn test_initialize_and_shutdown() {
let backend = PaddleOcrVlBackend::default_task();
assert!(backend.initialize().is_ok());
assert!(backend.shutdown().is_ok());
}
#[test]
fn should_stay_on_requires_upright_until_rotation_handling_is_measured() {
let backend = PaddleOcrVlBackend::default_task();
let dynamic: &dyn OcrBackend = &backend;
assert_eq!(
dynamic.page_orientation_handling(),
crate::plugins::PageOrientationHandling::RequiresUpright
);
}
}