#[cfg(any(
test,
feature = "candle-trocr",
feature = "candle-paddleocr-vl",
feature = "candle-glm-ocr",
feature = "candle-deepseek-ocr"
))]
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use xberg_candle_ocr::DevicePreference;
#[cfg(any(
test,
feature = "candle-trocr",
feature = "candle-paddleocr-vl",
feature = "candle-glm-ocr",
feature = "candle-deepseek-ocr"
))]
use crate::XbergError;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum CandleDevicePreference {
#[default]
Auto,
Cpu,
Cuda,
Metal,
}
impl From<CandleDevicePreference> for DevicePreference {
fn from(value: CandleDevicePreference) -> Self {
match value {
CandleDevicePreference::Auto => Self::Auto,
CandleDevicePreference::Cpu => Self::Cpu,
CandleDevicePreference::Cuda => Self::Cuda,
CandleDevicePreference::Metal => Self::Metal,
}
}
}
impl From<DevicePreference> for CandleDevicePreference {
fn from(value: DevicePreference) -> Self {
match value {
DevicePreference::Auto => Self::Auto,
DevicePreference::Cpu => Self::Cpu,
DevicePreference::Cuda => Self::Cuda,
DevicePreference::Metal => Self::Metal,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum CandleTrocrVariant {
#[default]
BasePrinted,
LargePrinted,
BaseHandwritten,
LargeHandwritten,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum PaddleOcrVlTaskKind {
#[default]
Ocr,
Table,
Formula,
Chart,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum GlmOcrTaskKind {
#[default]
Ocr,
Table,
Formula,
Chart,
Caption,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum GlmOcrLayoutMode {
#[default]
WholePage,
Paired,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct TrocrBackendOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub variant: Option<CandleTrocrVariant>,
#[serde(skip_serializing_if = "Option::is_none")]
pub device: Option<CandleDevicePreference>,
#[serde(alias = "cache-dir", skip_serializing_if = "Option::is_none")]
pub cache_dir: Option<String>,
#[serde(alias = "hf-revision", alias = "revision", skip_serializing_if = "Option::is_none")]
pub hf_revision: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct PaddleOcrVlBackendOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub task: Option<PaddleOcrVlTaskKind>,
#[serde(alias = "model-path", skip_serializing_if = "Option::is_none")]
pub model_path: Option<String>,
#[serde(alias = "model-id", skip_serializing_if = "Option::is_none")]
pub model_id: Option<String>,
#[serde(alias = "hf-revision", alias = "revision", skip_serializing_if = "Option::is_none")]
pub hf_revision: Option<String>,
#[serde(alias = "cache-dir", skip_serializing_if = "Option::is_none")]
pub cache_dir: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub device: Option<CandleDevicePreference>,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct GlmOcrBackendOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub task: Option<GlmOcrTaskKind>,
#[serde(skip_serializing_if = "Option::is_none")]
pub device: Option<CandleDevicePreference>,
#[serde(alias = "layout-mode", skip_serializing_if = "Option::is_none")]
pub layout_mode: Option<GlmOcrLayoutMode>,
#[serde(alias = "enable-chart-understanding", skip_serializing_if = "Option::is_none")]
pub enable_chart_understanding: Option<bool>,
#[serde(alias = "cache-dir", skip_serializing_if = "Option::is_none")]
pub cache_dir: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct DeepseekOcrBackendOptions {
#[serde(alias = "model-path", skip_serializing_if = "Option::is_none")]
pub model_path: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub device: Option<CandleDevicePreference>,
#[serde(skip_serializing_if = "Option::is_none")]
pub version: Option<u32>,
}
#[cfg(any(
test,
feature = "candle-trocr",
feature = "candle-paddleocr-vl",
feature = "candle-glm-ocr",
feature = "candle-deepseek-ocr"
))]
pub(crate) fn parse_backend_options<T>(value: Option<&serde_json::Value>, backend_name: &str) -> crate::Result<T>
where
T: DeserializeOwned + Default,
{
let Some(value) = value else {
return Ok(T::default());
};
if !value.is_object() {
return Err(XbergError::validation(format!(
"invalid {backend_name} backend_options: expected a JSON object"
)));
}
serde_json::from_value(value.clone())
.map_err(|error| XbergError::validation(format!("invalid {backend_name} backend_options: {error}")))
}
#[cfg(any(
test,
feature = "candle-trocr",
feature = "candle-paddleocr-vl",
feature = "candle-glm-ocr",
feature = "candle-deepseek-ocr"
))]
pub(crate) fn validate_optional_non_empty(
value: Option<&str>,
backend_name: &str,
field_name: &str,
) -> crate::Result<()> {
if value.is_some_and(|value| value.trim().is_empty()) {
return Err(XbergError::validation(format!(
"invalid {backend_name} backend_options.{field_name}: value must not be empty"
)));
}
Ok(())
}
#[cfg_attr(alef, alef(skip))]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum CandleModelId {
#[default]
Trocr,
PaddleocrVl,
}
#[cfg_attr(alef, alef(skip))]
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default)]
pub struct CandleOcrConfig {
pub model: CandleModelId,
pub device: DevicePreference,
#[serde(alias = "cache-dir")]
pub cache_dir: Option<PathBuf>,
#[serde(alias = "hf-revision", alias = "revision")]
pub hf_revision: Option<String>,
#[serde(alias = "max-new-tokens")]
pub max_new_tokens: u32,
pub temperature: f32,
}
impl CandleOcrConfig {
#[must_use]
pub fn backend_name(&self) -> &'static str {
match self.model {
CandleModelId::Trocr => "candle-trocr",
CandleModelId::PaddleocrVl => "candle-paddleocr-vl",
}
}
#[must_use]
pub fn backend_options(&self) -> serde_json::Value {
let device = CandleDevicePreference::from(self.device);
let mut options = serde_json::Map::new();
options.insert("device".to_string(), serde_json::json!(device));
if let Some(cache_dir) = &self.cache_dir {
options.insert(
"cache_dir".to_string(),
serde_json::Value::String(cache_dir.to_string_lossy().into_owned()),
);
}
if let Some(hf_revision) = &self.hf_revision {
options.insert(
"hf_revision".to_string(),
serde_json::Value::String(hf_revision.clone()),
);
}
serde_json::Value::Object(options)
}
}
impl Default for CandleOcrConfig {
fn default() -> Self {
Self {
model: CandleModelId::default(),
device: DevicePreference::Auto,
cache_dir: None,
hf_revision: None,
max_new_tokens: 4096,
temperature: 0.0,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn should_serialize_canonical_snake_case_candle_options() {
let trocr = TrocrBackendOptions {
variant: Some(CandleTrocrVariant::LargeHandwritten),
device: Some(CandleDevicePreference::Cuda),
cache_dir: Some("/models/trocr".to_string()),
hf_revision: Some("trocr-revision".to_string()),
};
assert_eq!(
serde_json::to_value(trocr).unwrap(),
serde_json::json!({
"variant": "large-handwritten",
"device": "cuda",
"cache_dir": "/models/trocr",
"hf_revision": "trocr-revision"
})
);
let paddle = PaddleOcrVlBackendOptions {
task: Some(PaddleOcrVlTaskKind::Chart),
model_path: Some("/models/paddle".to_string()),
model_id: Some("example/paddle".to_string()),
hf_revision: Some("paddle-revision".to_string()),
cache_dir: Some("/cache/paddle".to_string()),
device: Some(CandleDevicePreference::Metal),
};
assert_eq!(
serde_json::to_value(paddle).unwrap(),
serde_json::json!({
"task": "chart",
"model_path": "/models/paddle",
"model_id": "example/paddle",
"hf_revision": "paddle-revision",
"cache_dir": "/cache/paddle",
"device": "metal"
})
);
let glm = GlmOcrBackendOptions {
task: Some(GlmOcrTaskKind::Caption),
device: Some(CandleDevicePreference::Cpu),
layout_mode: Some(GlmOcrLayoutMode::Paired),
enable_chart_understanding: Some(true),
cache_dir: Some("/cache/glm".to_string()),
};
assert_eq!(
serde_json::to_value(glm).unwrap(),
serde_json::json!({
"task": "caption",
"device": "cpu",
"layout_mode": "paired",
"enable_chart_understanding": true,
"cache_dir": "/cache/glm"
})
);
let deepseek = DeepseekOcrBackendOptions {
model_path: Some("/models/deepseek".to_string()),
device: Some(CandleDevicePreference::Auto),
version: Some(3),
};
assert_eq!(
serde_json::to_value(deepseek).unwrap(),
serde_json::json!({
"model_path": "/models/deepseek",
"device": "auto",
"version": 3
})
);
}
#[test]
fn should_accept_legacy_candle_option_key_aliases() {
let value = serde_json::json!({
"cache-dir": "/legacy-cache",
"hf-revision": "legacy-revision"
});
let options: TrocrBackendOptions = parse_backend_options(Some(&value), "candle-trocr").unwrap();
assert_eq!(options.cache_dir.as_deref(), Some("/legacy-cache"));
assert_eq!(options.hf_revision.as_deref(), Some("legacy-revision"));
let revision_alias = serde_json::json!({"revision": "short-alias"});
let options: PaddleOcrVlBackendOptions =
parse_backend_options(Some(&revision_alias), "candle-paddleocr-vl").unwrap();
assert_eq!(options.hf_revision.as_deref(), Some("short-alias"));
}
#[test]
fn should_reject_invalid_candle_backend_options_with_context() {
let invalid_device = serde_json::json!({"device": "quantum"});
let error = parse_backend_options::<TrocrBackendOptions>(Some(&invalid_device), "candle-trocr")
.unwrap_err()
.to_string();
assert!(error.contains("candle-trocr backend_options"));
assert!(error.contains("unknown variant `quantum`"));
let invalid_shape = serde_json::json!(["cpu"]);
let error = parse_backend_options::<GlmOcrBackendOptions>(Some(&invalid_shape), "candle-glm-ocr")
.unwrap_err()
.to_string();
assert!(error.contains("candle-glm-ocr backend_options"));
let error = validate_optional_non_empty(Some(" "), "candle-paddleocr-vl", "model_id")
.unwrap_err()
.to_string();
assert!(error.contains("candle-paddleocr-vl backend_options.model_id"));
let unknown_field = serde_json::json!({"model-path": "/models/deepseek", "versoin": 2});
let error = parse_backend_options::<DeepseekOcrBackendOptions>(Some(&unknown_field), "candle-deepseek-ocr")
.unwrap_err()
.to_string();
assert!(error.contains("unknown field `versoin`"));
assert!(error.contains("candle-deepseek-ocr backend_options"));
}
#[test]
fn should_convert_legacy_config_to_runtime_backend_options() {
let config = CandleOcrConfig {
model: CandleModelId::PaddleocrVl,
device: DevicePreference::Cuda,
cache_dir: Some(PathBuf::from("/legacy-cache")),
hf_revision: Some("legacy-revision".to_string()),
max_new_tokens: 128,
temperature: 0.75,
};
assert_eq!(
config.backend_options(),
serde_json::json!({
"device": "cuda",
"cache_dir": "/legacy-cache",
"hf_revision": "legacy-revision"
})
);
assert_eq!(config.backend_name(), "candle-paddleocr-vl");
}
}