use std::path::PathBuf;
use async_trait::async_trait;
use crate::Result;
use crate::XbergError;
use crate::layout::{LayoutError, LayoutModelManager};
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct ModelId {
pub kind: String,
}
impl ModelId {
pub fn new(kind: impl Into<String>) -> Self {
Self { kind: kind.into() }
}
}
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
pub trait ModelProvider: Send + Sync + 'static {
async fn ensure_model(&self, model: &ModelId) -> Result<PathBuf>;
}
#[derive(Debug, Clone)]
pub struct DefaultModelProvider {
manager: LayoutModelManager,
}
impl Default for DefaultModelProvider {
fn default() -> Self {
Self {
manager: LayoutModelManager::new(None),
}
}
}
fn ensure_by_kind(manager: &LayoutModelManager, kind: &str) -> Result<PathBuf> {
let resolved = match kind {
"rtdetr" => manager.ensure_rtdetr_model(),
"tatr" => manager.ensure_tatr_model(),
"table_classifier" => manager.ensure_table_classifier(),
"pp_doclayout_v3" => manager.ensure_pp_doclayout_v3_model(),
variant => manager.ensure_slanet_model(variant),
};
resolved.map_err(|error| map_layout_error(kind, error))
}
fn map_layout_error(kind: &str, error: LayoutError) -> XbergError {
let message = format!("layout model '{kind}' unavailable: {error}");
match &error {
LayoutError::ModelDownload(detail) if detail.starts_with("Unknown model type:") => {
XbergError::validation(message)
}
_ => XbergError::MissingDependency(message),
}
}
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
impl ModelProvider for DefaultModelProvider {
async fn ensure_model(&self, model: &ModelId) -> Result<PathBuf> {
let manager = self.manager.clone();
let kind = model.kind.clone();
tokio::task::spawn_blocking(move || ensure_by_kind(&manager, &kind))
.await
.map_err(|error| XbergError::Other(format!("layout model task failed to join: {error}")))?
}
}