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::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,
}
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(config: &OcrConfig) -> PaddleOcrVlOptions {
let mut task = PaddleOcrVlTask::default();
let mut model_path: Option<String> = None;
let mut model_id = DEFAULT_MODEL_ID.to_string();
let mut hf_revision = None;
let mut cache_dir = 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" => PaddleOcrVlTask::Table,
"formula" => PaddleOcrVlTask::Formula,
"chart" => PaddleOcrVlTask::Chart,
_ => PaddleOcrVlTask::Ocr,
};
}
if let Some(p) = opts.get("model_path").and_then(|v| v.as_str()) {
model_path = Some(p.to_string());
}
if let Some(id) = opts.get("model_id").and_then(|v| v.as_str()) {
model_id = id.to_string();
}
hf_revision = opts
.get("hf_revision")
.or_else(|| opts.get("revision"))
.and_then(|value| value.as_str())
.map(str::to_string);
cache_dir = opts
.get("cache_dir")
.and_then(|value| value.as_str())
.map(PathBuf::from);
}
PaddleOcrVlOptions {
task,
model_path,
model_id,
hf_revision,
cache_dir,
device: super::resolve_device_preference(config),
}
}
}
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 = 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)
.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(ExtractedDocument {
content,
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(test)]
mod tests {
use super::*;
#[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 = PaddleOcrVlBackend::parse_options(&config);
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 = PaddleOcrVlBackend::parse_options(&config);
assert_eq!(options.task, PaddleOcrVlTask::Table);
}
#[test]
fn test_parse_options_custom_device() {
let config = OcrConfig {
backend_options: Some(serde_json::json!({
"device": "cpu"
})),
..Default::default()
};
let options = PaddleOcrVlBackend::parse_options(&config);
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 = PaddleOcrVlBackend::parse_options(&config);
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 = PaddleOcrVlBackend::parse_options(&config);
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 = PaddleOcrVlBackend::parse_options(&config);
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_initialize_and_shutdown() {
let backend = PaddleOcrVlBackend::default_task();
assert!(backend.initialize().is_ok());
assert!(backend.shutdown().is_ok());
}
}