1use std::collections::BTreeMap;
4use std::sync::Arc;
5
6use crate::engine::{
7 AsrEngine, DepthEngine, DetectEngine, EngineInfo, OcrEngine, Task, TtsEngine, VlmEngine,
8};
9use crate::error::{Error, Result};
10
11#[derive(Default)]
18pub struct EngineRegistry {
19 asr: BTreeMap<String, Arc<dyn AsrEngine>>,
20 tts: BTreeMap<String, Arc<dyn TtsEngine>>,
21 ocr: BTreeMap<String, Arc<dyn OcrEngine>>,
22 vlm: BTreeMap<String, Arc<dyn VlmEngine>>,
23 detect: BTreeMap<String, Arc<dyn DetectEngine>>,
24 depth: BTreeMap<String, Arc<dyn DepthEngine>>,
25 asr_default: Option<String>,
28 tts_default: Option<String>,
29 ocr_default: Option<String>,
30 vlm_default: Option<String>,
31 detect_default: Option<String>,
32 depth_default: Option<String>,
33}
34
35impl EngineRegistry {
36 #[must_use]
37 pub fn new() -> Self {
38 Self::default()
39 }
40
41 pub fn register_asr(&mut self, engine: Arc<dyn AsrEngine>) {
42 let name = engine.info().name;
43 self.asr_default.get_or_insert_with(|| name.clone());
44 self.asr.insert(name, engine);
45 }
46
47 pub fn register_tts(&mut self, engine: Arc<dyn TtsEngine>) {
48 let name = engine.info().name;
49 self.tts_default.get_or_insert_with(|| name.clone());
50 self.tts.insert(name, engine);
51 }
52
53 pub fn register_ocr(&mut self, engine: Arc<dyn OcrEngine>) {
54 let name = engine.info().name;
55 self.ocr_default.get_or_insert_with(|| name.clone());
56 self.ocr.insert(name, engine);
57 }
58
59 pub fn register_vlm(&mut self, engine: Arc<dyn VlmEngine>) {
60 let name = engine.info().name;
61 self.vlm_default.get_or_insert_with(|| name.clone());
62 self.vlm.insert(name, engine);
63 }
64
65 pub fn register_depth(&mut self, engine: Arc<dyn DepthEngine>) {
66 let name = engine.info().name;
67 self.depth_default.get_or_insert_with(|| name.clone());
68 self.depth.insert(name, engine);
69 }
70
71 pub fn register_detect(&mut self, engine: Arc<dyn DetectEngine>) {
72 let name = engine.info().name;
73 self.detect_default.get_or_insert_with(|| name.clone());
74 self.detect.insert(name, engine);
75 }
76
77 pub fn asr(&self, name: Option<&str>) -> Result<Arc<dyn AsrEngine>> {
79 resolve(&self.asr, name, self.asr_default.as_deref(), Task::Asr)
80 }
81
82 pub fn tts(&self, name: Option<&str>) -> Result<Arc<dyn TtsEngine>> {
83 resolve(&self.tts, name, self.tts_default.as_deref(), Task::Tts)
84 }
85
86 pub fn ocr(&self, name: Option<&str>) -> Result<Arc<dyn OcrEngine>> {
87 resolve(&self.ocr, name, self.ocr_default.as_deref(), Task::Ocr)
88 }
89
90 pub fn vlm(&self, name: Option<&str>) -> Result<Arc<dyn VlmEngine>> {
91 resolve(&self.vlm, name, self.vlm_default.as_deref(), Task::Vlm)
92 }
93
94 pub fn depth(&self, name: Option<&str>) -> Result<Arc<dyn DepthEngine>> {
95 resolve(
96 &self.depth,
97 name,
98 self.depth_default.as_deref(),
99 Task::Depth,
100 )
101 }
102
103 pub fn detect(&self, name: Option<&str>) -> Result<Arc<dyn DetectEngine>> {
104 resolve(
105 &self.detect,
106 name,
107 self.detect_default.as_deref(),
108 Task::Detect,
109 )
110 }
111
112 #[must_use]
114 pub fn list(&self) -> Vec<EngineInfo> {
115 let mut out: Vec<EngineInfo> = Vec::new();
116 out.extend(self.asr.values().map(|e| e.info()));
117 out.extend(self.tts.values().map(|e| e.info()));
118 out.extend(self.ocr.values().map(|e| e.info()));
119 out.extend(self.vlm.values().map(|e| e.info()));
120 out.extend(self.detect.values().map(|e| e.info()));
121 out.sort_by(|a, b| a.task.cmp(&b.task).then_with(|| a.name.cmp(&b.name)));
122 out
123 }
124}
125
126fn resolve<E: ?Sized>(
127 map: &BTreeMap<String, Arc<E>>,
128 name: Option<&str>,
129 default: Option<&str>,
130 task: Task,
131) -> Result<Arc<E>> {
132 match name.or(default) {
133 Some(n) => map.get(n).cloned().ok_or_else(|| Error::UnknownEngine {
134 task,
135 name: n.to_string(),
136 }),
137 None => Err(Error::UnknownEngine {
138 task,
139 name: "<no engines registered>".to_string(),
140 }),
141 }
142}