1use std::borrow::Cow;
7use std::collections::HashMap;
8use std::path::{Path, PathBuf};
9use std::sync::Arc;
10use std::time::Duration;
11
12use async_trait::async_trait;
13use ort::session::{Session, SessionInputValue};
14use ort::value::Tensor;
15use parking_lot::Mutex;
16
17use crate::ai::provider::{
18 InferenceOutputs, InferenceProvider, InferenceRequest, InferenceResponse, ProviderError, Usage,
19};
20use crate::ai::registry::Task;
21
22const INFERENCE_TIMEOUT: Duration = Duration::from_secs(60);
25
26struct LoadedModel {
29 session: Mutex<Session>,
30 tokenizer: tokenizers::Tokenizer,
31 input_names: Vec<String>,
32 labels: Vec<String>,
34}
35
36pub struct LocalProvider {
38 cache_dir: PathBuf,
39 loaded: Mutex<HashMap<String, Arc<LoadedModel>>>,
40 load_lock: tokio::sync::Mutex<()>,
42}
43
44impl LocalProvider {
45 #[must_use]
47 pub fn new(cache_dir: impl Into<PathBuf>) -> Self {
48 Self {
49 cache_dir: cache_dir.into(),
50 loaded: Mutex::new(HashMap::new()),
51 load_lock: tokio::sync::Mutex::new(()),
52 }
53 }
54
55 async fn ensure_model(&self, source: &str) -> Result<Arc<LoadedModel>, ProviderError> {
57 if let Some(model) = self.loaded.lock().get(source) {
58 return Ok(Arc::clone(model));
59 }
60 let _load = self.load_lock.lock().await;
62 if let Some(model) = self.loaded.lock().get(source) {
63 return Ok(Arc::clone(model));
64 }
65
66 let loaded = if let Some(repo_id) = source.strip_prefix("hf:") {
67 let api = hf_hub::api::tokio::ApiBuilder::from_env()
68 .with_cache_dir(self.cache_dir.clone())
69 .build()
70 .map_err(|e| {
71 ProviderError::Transport(format!("failed to initialize hf-hub client: {e}"))
72 })?;
73 let repo = api.model(repo_id.to_string());
74
75 let tokenizer_path = repo.get("tokenizer.json").await.map_err(|e| {
76 ProviderError::Transport(format!("failed to download tokenizer.json: {e}"))
77 })?;
78
79 let onnx_path = match repo.get("onnx/model.onnx").await {
80 Ok(path) => path,
81 Err(_) => repo.get("model.onnx").await.map_err(|e| {
82 ProviderError::Transport(format!("failed to download model.onnx: {e}"))
83 })?,
84 };
85
86 let config_path = repo.get("config.json").await.ok();
87
88 tokio::task::spawn_blocking(move || {
90 load_model_from_paths(&onnx_path, &tokenizer_path, config_path.as_deref())
91 })
92 .await
93 .map_err(|e| ProviderError::Transport(format!("model load task: {e}")))??
94 } else {
95 let dir = model_dir(&self.cache_dir, source);
96 tokio::task::spawn_blocking(move || load_model(&dir))
98 .await
99 .map_err(|e| ProviderError::Transport(format!("model load task: {e}")))??
100 };
101
102 let loaded = Arc::new(loaded);
103 self.loaded
104 .lock()
105 .insert(source.to_string(), Arc::clone(&loaded));
106 Ok(loaded)
107 }
108}
109
110#[must_use]
113pub fn model_dir(cache_dir: &Path, source: &str) -> PathBuf {
114 if let Some(repo) = source.strip_prefix("hf:") {
115 cache_dir.join(repo)
116 } else if let Some(path) = source.strip_prefix("file://") {
117 PathBuf::from(path)
118 } else {
119 PathBuf::from(source)
120 }
121}
122
123#[must_use]
126pub fn load_labels(cache_dir: &Path, source: &str) -> Vec<String> {
127 if let Ok(text) = std::fs::read_to_string(model_dir(cache_dir, source).join("config.json")) {
128 return parse_id2label(&text);
129 }
130 if let Some(repo_id) = source.strip_prefix("hf:") {
132 let cache = hf_hub::Cache::new(cache_dir.to_path_buf());
133 let repo = cache.repo(hf_hub::Repo::model(repo_id.to_string()));
134 if let Some(path) = repo.get("config.json") {
135 if let Ok(text) = std::fs::read_to_string(path) {
136 return parse_id2label(&text);
137 }
138 }
139 }
140 Vec::new()
141}
142
143fn parse_id2label(config_json: &str) -> Vec<String> {
144 let Ok(json) = serde_json::from_str::<serde_json::Value>(config_json) else {
145 return Vec::new();
146 };
147 let Some(map) = json.get("id2label").and_then(serde_json::Value::as_object) else {
148 return Vec::new();
149 };
150 let mut indexed: Vec<(usize, String)> = map
151 .iter()
152 .filter_map(|(k, v)| Some((k.parse().ok()?, v.as_str()?.to_string())))
153 .collect();
154 indexed.sort_by_key(|(index, _)| *index);
155 indexed.into_iter().map(|(_, label)| label).collect()
156}
157
158#[async_trait]
159impl InferenceProvider for LocalProvider {
160 async fn infer_batch(
161 &self,
162 request: InferenceRequest,
163 ) -> Result<InferenceResponse, ProviderError> {
164 if matches!(
165 request.task,
166 Task::Complete | Task::Summarize | Task::Translate | Task::Gen | Task::Extract
167 ) {
168 return Err(ProviderError::UnsupportedTask(request.task));
169 }
170 let loaded = self.ensure_model(&request.model).await?;
171 let task = request.task;
172 let inputs = request.inputs;
173 let run = tokio::task::spawn_blocking(move || run(&loaded, task, &inputs));
175 let outputs = match tokio::time::timeout(INFERENCE_TIMEOUT, run).await {
176 Ok(joined) => {
177 joined.map_err(|e| ProviderError::Transport(format!("inference task: {e}")))??
178 }
179 Err(_) => {
180 return Err(ProviderError::Timeout(
181 u64::try_from(INFERENCE_TIMEOUT.as_millis()).unwrap_or(u64::MAX),
182 ))
183 }
184 };
185 Ok(InferenceResponse {
186 outputs,
187 usage: Usage::ZERO,
188 })
189 }
190
191 fn name(&self) -> &'static str {
192 "local"
193 }
194
195 fn intrinsic_labels(&self, model: &str) -> Option<Vec<String>> {
196 let labels = self
198 .loaded
199 .lock()
200 .get(model)
201 .map_or_else(|| load_labels(&self.cache_dir, model), |m| m.labels.clone());
202 (!labels.is_empty()).then_some(labels)
203 }
204}
205
206fn onnx_path(dir: &Path) -> PathBuf {
209 let nested = dir.join("onnx").join("model.onnx");
210 if nested.exists() {
211 nested
212 } else {
213 dir.join("model.onnx")
214 }
215}
216
217fn load_model_from_paths(
218 onnx_path: &Path,
219 tokenizer_path: &Path,
220 config_path: Option<&Path>,
221) -> Result<LoadedModel, ProviderError> {
222 let session = Session::builder()
223 .map_err(|e| ProviderError::Transport(format!("ort init: {e}")))?
224 .commit_from_file(onnx_path)
225 .map_err(|e| ProviderError::Transport(format!("load onnx: {e}")))?;
226 let input_names = session
227 .inputs()
228 .iter()
229 .map(|i| i.name().to_string())
230 .collect();
231 let mut tokenizer = tokenizers::Tokenizer::from_file(tokenizer_path)
232 .map_err(|e| ProviderError::Transport(format!("load tokenizer: {e}")))?;
233 if tokenizer.get_padding().is_none() {
234 tokenizer.with_padding(Some(tokenizers::PaddingParams {
235 strategy: tokenizers::PaddingStrategy::BatchLongest,
236 direction: tokenizers::PaddingDirection::Right,
237 pad_to_multiple_of: None,
238 pad_id: 0,
239 pad_type_id: 0,
240 pad_token: "[PAD]".to_string(),
241 }));
242 }
243 let labels = if let Some(path) = config_path {
244 std::fs::read_to_string(path)
245 .ok()
246 .map(|text| parse_id2label(&text))
247 .unwrap_or_default()
248 } else {
249 Vec::new()
250 };
251 Ok(LoadedModel {
252 session: Mutex::new(session),
253 tokenizer,
254 input_names,
255 labels,
256 })
257}
258
259fn load_model(dir: &Path) -> Result<LoadedModel, ProviderError> {
260 let onnx = onnx_path(dir);
261 let tokenizer_path = dir.join("tokenizer.json");
262 if !onnx.exists() || !tokenizer_path.exists() {
263 return Err(ProviderError::Transport(format!(
264 "local model files not found in {} (expected onnx/model.onnx + tokenizer.json)",
265 dir.display()
266 )));
267 }
268 let config_path = dir.join("config.json");
269 let config_path_opt = config_path.exists().then_some(config_path);
270 load_model_from_paths(&onnx, &tokenizer_path, config_path_opt.as_deref())
271}
272
273fn run(
274 loaded: &LoadedModel,
275 task: Task,
276 inputs: &[String],
277) -> Result<InferenceOutputs, ProviderError> {
278 if inputs.is_empty() {
279 return Ok(InferenceOutputs::Vectors(vec![]));
280 }
281
282 let encodings = loaded
283 .tokenizer
284 .encode_batch(inputs.to_vec(), true)
285 .map_err(|e| ProviderError::BadResponse(format!("tokenize batch: {e}")))?;
286
287 let batch_size = encodings.len();
288 if batch_size == 0 {
289 return Ok(InferenceOutputs::Vectors(vec![]));
290 }
291 let seq = encodings[0].len();
292
293 let mut stacked_ids = Vec::with_capacity(batch_size * seq);
294 let mut stacked_mask = Vec::with_capacity(batch_size * seq);
295 for encoding in &encodings {
296 stacked_ids.extend(encoding.get_ids().iter().map(|&u| i64::from(u)));
297 stacked_mask.extend(encoding.get_attention_mask().iter().map(|&u| i64::from(u)));
298 }
299
300 let batch_size_i64 = i64::try_from(batch_size).unwrap_or(i64::MAX);
301 let seq_i64 = i64::try_from(seq).unwrap_or(i64::MAX);
302 let shape = vec![batch_size_i64, seq_i64];
303
304 let mut feeds: Vec<(Cow<str>, SessionInputValue)> =
305 Vec::with_capacity(loaded.input_names.len());
306 for name in &loaded.input_names {
307 let row = match name.as_str() {
308 "input_ids" => stacked_ids.clone(),
309 "attention_mask" => stacked_mask.clone(),
310 "token_type_ids" => vec![0i64; batch_size * seq],
311 other => {
312 return Err(ProviderError::BadResponse(format!(
313 "model expects unsupported input '{other}'"
314 )))
315 }
316 };
317 let tensor = Tensor::from_array((shape.clone(), row))
318 .map_err(|e| ProviderError::Transport(format!("build tensor: {e}")))?;
319 feeds.push((Cow::Owned(name.clone()), SessionInputValue::from(tensor)));
320 }
321
322 let mut session = loaded.session.lock();
323 let outputs = session
324 .run(feeds)
325 .map_err(|e| ProviderError::Transport(format!("inference: {e}")))?;
326 let (shape, data) = outputs[0]
327 .try_extract_tensor::<f32>()
328 .map_err(|e| ProviderError::BadResponse(format!("read output: {e}")))?;
329
330 if shape.first().copied().and_then(|d| usize::try_from(d).ok()) != Some(batch_size) {
331 return Err(ProviderError::BadResponse(format!(
332 "expected batch size {batch_size}, got output shape {shape:?}"
333 )));
334 }
335
336 let mut vectors = Vec::with_capacity(batch_size);
337 if task == Task::Embed && shape.len() == 3 {
338 let seq_out = usize::try_from(shape[1]).unwrap_or(0);
340 let hidden = usize::try_from(shape[2]).unwrap_or(0);
341 for (i, encoding) in encodings.iter().enumerate() {
342 let start = i * seq_out * hidden;
343 let end = (i + 1) * seq_out * hidden;
344 let slice = &data[start..end];
345 let mask: Vec<i64> = encoding
346 .get_attention_mask()
347 .iter()
348 .map(|&u| i64::from(u))
349 .collect();
350 vectors.push(mean_pool(slice, seq_out, hidden, &mask));
351 }
352 } else {
353 let dim: usize = shape[1..]
355 .iter()
356 .map(|&d| usize::try_from(d).unwrap_or(0))
357 .product();
358 for i in 0..batch_size {
359 let start = i * dim;
360 let end = (i + 1) * dim;
361 let slice = &data[start..end];
362 vectors.push(slice.to_vec());
363 }
364 }
365
366 Ok(InferenceOutputs::Vectors(vectors))
367}
368
369fn mean_pool(data: &[f32], seq: usize, hidden: usize, mask: &[i64]) -> Vec<f32> {
371 let mut pooled = vec![0.0_f32; hidden];
372 let mut count = 0.0_f32;
373 for t in 0..seq {
374 if mask.get(t).copied().unwrap_or(0) == 0 {
375 continue;
376 }
377 count += 1.0;
378 for h in 0..hidden {
379 pooled[h] += data[t * hidden + h];
380 }
381 }
382 if count > 0.0 {
383 for value in &mut pooled {
384 *value /= count;
385 }
386 }
387 pooled
388}
389
390#[cfg(test)]
391mod tests {
392 use super::*;
393
394 #[test]
395 fn model_dir_resolution() {
396 let cache = Path::new("/models");
397 assert_eq!(
398 model_dir(cache, "hf:onnx-community/finbert"),
399 Path::new("/models/onnx-community/finbert")
400 );
401 assert_eq!(model_dir(cache, "file:///abs/dir"), Path::new("/abs/dir"));
402 assert_eq!(model_dir(cache, "/some/path"), Path::new("/some/path"));
403 }
404
405 #[test]
406 fn parse_id2label_orders_by_index() {
407 let config = r#"{"id2label": {"2": "positive", "0": "negative", "1": "neutral"}}"#;
408 assert_eq!(
409 parse_id2label(config),
410 vec!["negative", "neutral", "positive"]
411 );
412 assert!(parse_id2label("{}").is_empty());
413 }
414
415 #[test]
416 fn intrinsic_labels_read_from_a_cached_models_config() {
417 let cache = tempfile::tempdir().expect("tempdir");
418 let provider = LocalProvider::new(cache.path());
419 let source = "hf:org/repo";
420
421 assert!(provider.intrinsic_labels(source).is_none());
423
424 let dir = model_dir(cache.path(), source);
425 std::fs::create_dir_all(&dir).unwrap();
426 std::fs::write(
427 dir.join("config.json"),
428 r#"{"id2label": {"0": "NEGATIVE", "1": "POSITIVE"}}"#,
429 )
430 .unwrap();
431 assert_eq!(
432 provider.intrinsic_labels(source),
433 Some(vec!["NEGATIVE".into(), "POSITIVE".into()])
434 );
435 }
436
437 #[test]
438 fn mean_pool_ignores_masked_tokens() {
439 let data = [1.0, 2.0, 3.0, 4.0, 100.0, 100.0];
441 let pooled = mean_pool(&data, 3, 2, &[1, 1, 0]);
442 assert_eq!(pooled, vec![2.0, 3.0]); }
444
445 #[tokio::test]
453 #[ignore = "downloads a model + needs ORT_DYLIB_PATH; run with --ignored"]
454 async fn classifies_with_a_real_onnx_community_model() {
455 use crate::ai::adapter::parse_response;
456 use crate::ai::provider::InferenceParams;
457 use crate::ai::registry::BackendKind;
458
459 let cache = tempfile::tempdir().expect("tempdir");
460 let provider = LocalProvider::new(cache.path());
461 let source = "hf:onnx-community/distilbert-base-uncased-finetuned-sst-2-english-ONNX";
462 let request = InferenceRequest {
463 task: Task::Classify,
464 model: source.to_string(),
465 inputs: vec![
466 "this film was absolutely wonderful, I loved every minute".into(),
467 "a complete waste of time, dull and disappointing".into(),
468 ],
469 params: InferenceParams::default(),
470 };
471
472 let response = provider.infer_batch(request).await.expect("inference");
473 let InferenceOutputs::Vectors(rows) = &response.outputs else {
474 panic!("local classify returns logits");
475 };
476 let labels = load_labels(cache.path(), source);
477 assert!(!labels.is_empty(), "id2label should load from config.json");
478 assert_eq!(rows.len(), 2);
479 assert!(
480 rows.iter().all(|r| r.len() == labels.len()),
481 "logit dimension must equal the label count",
482 );
483
484 let InferenceOutputs::Text(out) = parse_response(
485 Task::Classify,
486 BackendKind::Local,
487 response.outputs,
488 Some(&labels),
489 )
490 .expect("adapt") else {
491 panic!("argmax yields labels");
492 };
493 assert_eq!(out, vec!["POSITIVE".to_string(), "NEGATIVE".to_string()]);
494 }
495}