Skip to main content

lc_embeddings/
local.rs

1// lc-embeddings/src/local.rs
2//! Local embedding implementations
3//!
4//! Contains two implementations:
5//! - `BagOfWordsEmbeddings`: Lightweight word-frequency hash embedding (pure Rust, no external deps), always available
6//! - `LocalEmbeddings`: ONNX Runtime-based neural network embedding (requires `local-embeddings` feature)
7//!
8//! `BagOfWordsEmbeddings` is suitable for offline, privacy, zero-cost coarse-grained retrieval;
9//! `LocalEmbeddings` is suitable for high-quality semantic embedding scenarios (e.g., BGE/E5 models).
10
11use async_trait::async_trait;
12
13#[cfg(feature = "local-embeddings")]
14use std::path::Path;
15
16use crate::{EmbeddingError, Embeddings};
17
18// ---------------------------------------------------------------------------
19// BagOfWordsEmbeddings — word-frequency hash + L2 normalization (always available)
20// ---------------------------------------------------------------------------
21
22/// Lightweight local embedding (word-frequency hash + L2 normalization)
23///
24/// Based on word frequency + hashing, no API calls, suitable for offline, privacy, zero-cost coarse-grained retrieval.
25///
26/// Note: This is a lightweight implementation (bag-of-words hash) with limited semantic quality;
27/// for high-quality neural network embeddings (BGE/E5 via `ort`), enable the `local-embeddings` feature
28/// and use [`LocalEmbeddings`].
29pub struct BagOfWordsEmbeddings {
30    dim: usize,
31}
32
33impl BagOfWordsEmbeddings {
34    /// Create local embedding with specified dimension
35    pub fn new(dim: usize) -> Self {
36        Self { dim: dim.max(1) }
37    }
38
39    /// Default dimension 256
40    pub fn default_dim() -> Self {
41        Self::new(256)
42    }
43
44    /// Tokenize: English by non-alphanumeric split (lowercased), Chinese/non-ASCII by single character
45    fn tokenize(text: &str) -> Vec<String> {
46        let mut tokens = Vec::new();
47        let mut current = String::new();
48        for c in text.chars() {
49            if c.is_alphanumeric() {
50                if c.is_ascii() {
51                    current.push(c.to_ascii_lowercase());
52                } else {
53                    // Non-ASCII (Chinese etc.) single character as token
54                    if !current.is_empty() {
55                        tokens.push(std::mem::take(&mut current));
56                    }
57                    tokens.push(c.to_string());
58                }
59            } else if !current.is_empty() {
60                tokens.push(std::mem::take(&mut current));
61            }
62        }
63        if !current.is_empty() {
64            tokens.push(current);
65        }
66        tokens
67    }
68
69    /// FNV-1a hash
70    fn hash(s: &str) -> u64 {
71        let mut h: u64 = 0xcbf29ce484222325;
72        for b in s.bytes() {
73            h ^= b as u64;
74            h = h.wrapping_mul(0x100000001b3);
75        }
76        h
77    }
78
79    /// Compute embedding vector (word-frequency hash + L2 normalization)
80    fn embed(&self, text: &str) -> Vec<f32> {
81        let mut v = vec![0.0f32; self.dim];
82        for token in Self::tokenize(text) {
83            let idx = (Self::hash(&token) as usize) % self.dim;
84            v[idx] += 1.0;
85        }
86        // L2 normalization
87        let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
88        if norm > 0.0 {
89            for x in &mut v {
90                *x /= norm;
91            }
92        }
93        v
94    }
95}
96
97impl Default for BagOfWordsEmbeddings {
98    fn default() -> Self {
99        Self::default_dim()
100    }
101}
102
103#[async_trait]
104impl Embeddings for BagOfWordsEmbeddings {
105    async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
106        if text.trim().is_empty() {
107            return Err(EmbeddingError::EmptyInput);
108        }
109        Ok(self.embed(text))
110    }
111
112    fn dimension(&self) -> usize {
113        self.dim
114    }
115
116    fn model_name(&self) -> &str {
117        "local-bow"
118    }
119}
120
121// ---------------------------------------------------------------------------
122// LocalEmbeddings — ONNX Runtime neural network embedding (requires local-embeddings feature)
123// ---------------------------------------------------------------------------
124
125#[cfg(feature = "local-embeddings")]
126mod nn {
127    use super::*;
128    use ort::value::Tensor;
129    use std::path::PathBuf;
130    use std::sync::{Arc, Condvar, Mutex};
131
132    // ⚠️ 未验证声明(P2-2/3/4):
133    //
134    // 本机(Windows,无 MSVC Build Tools,GNU 链上 ort-sys 无 x86_64-pc-windows-gnu
135    // 预编译产物)无法编译 `local-embeddings` feature。以下代码全部按 vendored 的
136    // ort 2.0.0-rc.13 与 tokenizers 0.22.2 源码逐项核对 API 写成,但**未经编译器验证**。
137    // 请在支持 ort 的环境执行:
138    //
139    //   cargo test -p lc-embeddings --features local-embeddings --lib
140    //   cargo clippy -p lc-embeddings --features local-embeddings --lib
141    //
142    // 已核对的关键 API:
143    // - ort::session::Session::builder().commit_from_memory(&[u8])(impl_commit.rs:93)
144    // - Session::run(impl Into<SessionInputs>) 需 &mut self;SessionInputs: From<Vec<(K,V)>>
145    //   其中 K: Into<Cow<str>>, V: Into<SessionInputValue>(input.rs:62)
146    // - Tensor::from_array((Vec<i64>, Vec<i64>));try_extract_tensor::<f32>()
147    // - tokenizers::Tokenizer::from_file/encode_batch(..., add_special_tokens)
148    // - Encoding::get_ids/get_attention_mask/get_type_ids;tokenizer.get_padding()/token_to_id()
149
150    /// 动态 batch 模型的默认单批文本数上限。
151    const DEFAULT_MAX_BATCH: usize = 32;
152    /// 序列维为动态时,单条文本 token 数的默认截断上限(超出截断)。
153    const DEFAULT_MAX_SEQ_LEN: usize = 512;
154
155    /// 默认会话池大小:与 CPU 逻辑核数对齐但封顶 8,避免多个 ONNX session
156    /// 各自持有整份模型权重导致内存失控。
157    fn default_pool_size() -> usize {
158        std::thread::available_parallelism()
159            .map(|n| n.get().clamp(1, 8))
160            .unwrap_or(2)
161    }
162
163    /// 锁中毒时映射为 `ApiError`(内部状态异常,显式报错而非 panic)。
164    fn lock_error<T>(_: std::sync::PoisonError<T>) -> EmbeddingError {
165        EmbeddingError::ApiError("session pool lock poisoned".to_string())
166    }
167
168    /// 有界 ONNX Session 并发池(P2-3)。
169    ///
170    /// 每个 `ort::session::Session` 持有独立的 ORT 运行环境;`Session::run` 要求
171    /// `&mut self`,多线程并发推理不能共享单个 session,池让并发请求复用至多
172    /// `capacity` 个会话。模型以字节保存在内存中,会话按需 `commit_from_memory`
173    /// 惰性创建——既避免模型文件事后被删除导致 `commit_from_file` 失效,也让池
174    /// 完全自持、可跨线程共享。`acquire` 超容量时经 Condvar 阻塞等待。
175    struct SessionPool {
176        model_bytes: Arc<Vec<u8>>,
177        idle: Mutex<Vec<ort::session::Session>>,
178        live: Mutex<usize>,
179        capacity: usize,
180        notify: Condvar,
181    }
182
183    impl SessionPool {
184        /// 预建第一个会话(live=1),后续会话按需惰性创建。
185        fn new(model_bytes: Vec<u8>, capacity: usize) -> Result<Arc<Self>, EmbeddingError> {
186            let capacity = capacity.max(1);
187            let pool = Arc::new(Self {
188                model_bytes: Arc::new(model_bytes),
189                idle: Mutex::new(Vec::new()),
190                live: Mutex::new(0),
191                capacity,
192                notify: Condvar::new(),
193            });
194            let session = pool.build_session()?;
195            *pool.live.lock().map_err(lock_error)? = 1;
196            pool.idle.lock().map_err(lock_error)?.push(session);
197            Ok(pool)
198        }
199
200        fn build_session(&self) -> Result<ort::session::Session, EmbeddingError> {
201            ort::session::Session::builder()
202                .map_err(|e| {
203                    EmbeddingError::ApiError(format!("Failed to create ONNX SessionBuilder: {e}"))
204                })?
205                .commit_from_memory(&self.model_bytes)
206                .map_err(|e| {
207                    EmbeddingError::ApiError(format!("Failed to load ONNX model from memory: {e}"))
208                })
209        }
210
211        /// 借出一个会话。空闲队列有现成的直接拿;池未满则惰性新建一个;满则
212        /// Condvar 阻塞等待归还。调用方应在 `spawn_blocking` 线程内使用本方法。
213        fn acquire(self: &Arc<Self>) -> Result<SessionGuard, EmbeddingError> {
214            // 快速路径:空闲队列非空。
215            {
216                let mut idle = self.idle.lock().map_err(lock_error)?;
217                if let Some(session) = idle.pop() {
218                    return Ok(SessionGuard {
219                        pool: self.clone(),
220                        session: Some(session),
221                    });
222                }
223            }
224            // 慢速路径:需要新建会话,或等待其他调用方归还。
225            let mut live = self.live.lock().map_err(lock_error)?;
226            loop {
227                // 等待期间可能有会话被归还,重新检查空闲队列。
228                {
229                    let mut idle = self.idle.lock().map_err(lock_error)?;
230                    if let Some(session) = idle.pop() {
231                        return Ok(SessionGuard {
232                            pool: self.clone(),
233                            session: Some(session),
234                        });
235                    }
236                }
237                if *live < self.capacity {
238                    *live += 1;
239                    match self.build_session() {
240                        Ok(session) => {
241                            return Ok(SessionGuard {
242                                pool: self.clone(),
243                                session: Some(session),
244                            });
245                        }
246                        Err(e) => {
247                            // 回滚 live 计数并唤醒等待者,让其他人有机会重试。
248                            *live -= 1;
249                            self.notify.notify_one();
250                            return Err(e);
251                        }
252                    }
253                }
254                live = self.notify.wait(live).map_err(lock_error)?;
255            }
256        }
257
258        /// 归还会话:入空闲队列并唤醒一个等待者。锁中毒时直接丢弃会话。
259        fn release(&self, session: ort::session::Session) {
260            let Ok(mut idle) = self.idle.lock() else {
261                return;
262            };
263            idle.push(session);
264            drop(idle);
265            self.notify.notify_one();
266        }
267    }
268
269    /// RAII 会话借出句柄:离开作用域(含 `?` 提前返回)自动归还池,不泄漏 session。
270    struct SessionGuard {
271        pool: Arc<SessionPool>,
272        session: Option<ort::session::Session>,
273    }
274
275    impl SessionGuard {
276        fn session(&mut self) -> Result<&mut ort::session::Session, EmbeddingError> {
277            self.session.as_mut().ok_or_else(|| {
278                EmbeddingError::ApiError("session guard already released".to_string())
279            })
280        }
281    }
282
283    impl Drop for SessionGuard {
284        fn drop(&mut self) {
285            if let Some(session) = self.session.take() {
286                self.pool.release(session);
287            }
288        }
289    }
290
291    /// 内部共享状态:池 + tokenizer + 静态元信息。
292    ///
293    /// 用 `Arc` 包一层,让 `spawn_blocking` 闭包捕获 owned 引用(满足 `'static`),
294    /// 修掉旧实现 `move || self.embed_single(&text)` 捕获 `&self` 导致的潜伏编译错误。
295    struct LocalInner {
296        pool: Arc<SessionPool>,
297        tokenizer: tokenizers::Tokenizer,
298        dim: usize,
299        model_name: String,
300        seq_limit: usize,
301        max_batch: usize,
302    }
303
304    impl LocalInner {
305        /// 真实 tokenizers 编码(P2-2):加载 HuggingFace WordPiece/BPE tokenizer,
306        /// 产出 `input_ids + attention_mask + token_type_ids` 三件套。
307        fn tokenize(&self, texts: &[String]) -> Result<Vec<tokenizers::Encoding>, EmbeddingError> {
308            self.tokenizer
309                .encode_batch(texts.to_vec(), true)
310                .map_err(|e| {
311                    EmbeddingError::Config(format!("Tokenizer failed to encode input: {e}"))
312                })
313        }
314
315        /// 把一批 encodings 整理成按 batch 内最长序列 pad 对齐的张量数据(P2-4)。
316        ///
317        /// 返回 `(input_ids, attention_mask, token_type_ids, 对齐后序列长 max_len)`。
318        /// 纯逻辑、不依赖 ONNX 会话,可直接单测。pad 位置 mask=0、type_id=0。
319        fn build_batch_tensors(
320            encodings: &[tokenizers::Encoding],
321            seq_limit: usize,
322            pad_id: u32,
323        ) -> (Vec<i64>, Vec<i64>, Vec<i64>, usize) {
324            let batch = encodings.len();
325            let lens: Vec<usize> = encodings
326                .iter()
327                .map(|e| e.get_ids().len().min(seq_limit))
328                .collect();
329            let max_len = lens.iter().copied().max().unwrap_or(0);
330            let pad = pad_id as i64;
331            let mut input_ids = vec![pad; batch * max_len];
332            let mut attention_mask = vec![0i64; batch * max_len];
333            let mut token_type_ids = vec![0i64; batch * max_len];
334            for (b, enc) in encodings.iter().enumerate() {
335                let ids = enc.get_ids();
336                let mask = enc.get_attention_mask();
337                let types = enc.get_type_ids();
338                let len = lens[b];
339                let row = b * max_len;
340                for i in 0..len {
341                    input_ids[row + i] = ids[i] as i64;
342                    attention_mask[row + i] = mask.get(i).copied().unwrap_or(1) as i64;
343                    token_type_ids[row + i] = types.get(i).copied().unwrap_or(0) as i64;
344                }
345            }
346            (input_ids, attention_mask, token_type_ids, max_len)
347        }
348
349        /// 批量推理:pad 对齐喂入,masked mean pooling 出每行向量,L2 归一化(P2-4)。
350        fn infer_rows(
351            &self,
352            encodings: &[tokenizers::Encoding],
353        ) -> Result<Vec<Vec<f32>>, EmbeddingError> {
354            if encodings.is_empty() {
355                return Err(EmbeddingError::EmptyInput);
356            }
357            let pad_id = Self::resolve_pad_id(&self.tokenizer);
358            let (input_ids, attention_mask, token_type_ids, fed_seq_len) =
359                Self::build_batch_tensors(encodings, self.seq_limit, pad_id);
360            if fed_seq_len == 0 {
361                return Err(EmbeddingError::EmptyInput);
362            }
363            let batch = encodings.len();
364            let mask_rows: Vec<Vec<i64>> = (0..batch)
365                .map(|b| attention_mask[b * fed_seq_len..(b + 1) * fed_seq_len].to_vec())
366                .collect();
367
368            let input_shape = vec![batch as i64, fed_seq_len as i64];
369            let input_tensor = Tensor::from_array((input_shape, input_ids)).map_err(|e| {
370                EmbeddingError::ApiError(format!("Failed to construct input_ids tensor: {e}"))
371            })?;
372            let attention_tensor = Tensor::from_array((input_shape.clone(), attention_mask))
373                .map_err(|e| {
374                    EmbeddingError::ApiError(format!(
375                        "Failed to construct attention_mask tensor: {e}"
376                    ))
377                })?;
378            let type_tensor =
379                Tensor::from_array((input_shape.clone(), token_type_ids)).map_err(|e| {
380                    EmbeddingError::ApiError(format!(
381                        "Failed to construct token_type_ids tensor: {e}"
382                    ))
383                })?;
384
385            let mut guard = self.pool.acquire()?;
386            let input_names: Vec<String> = {
387                let session = guard.session()?;
388                session
389                    .inputs()
390                    .iter()
391                    .map(|o| o.name().to_string())
392                    .collect()
393            };
394            // 只喂模型声明存在的输入,按名对齐;未知输入名显式报错而非静默跳过(P2-2)。
395            let mut input_ids_slot = Some(input_tensor);
396            let mut attention_slot = Some(attention_tensor);
397            let mut type_slot = Some(type_tensor);
398            let mut named: Vec<(String, Tensor<i64>)> = Vec::with_capacity(input_names.len());
399            for name in input_names {
400                let tensor = match name.as_str() {
401                    "input_ids" => input_ids_slot.take().ok_or_else(|| {
402                        EmbeddingError::ParseError(
403                            "ONNX model declares duplicate 'input_ids' input".to_string(),
404                        )
405                    })?,
406                    "attention_mask" => attention_slot.take().ok_or_else(|| {
407                        EmbeddingError::ParseError(
408                            "ONNX model declares duplicate 'attention_mask' input".to_string(),
409                        )
410                    })?,
411                    "token_type_ids" => type_slot.take().ok_or_else(|| {
412                        EmbeddingError::ParseError(
413                            "ONNX model declares duplicate 'token_type_ids' input".to_string(),
414                        )
415                    })?,
416                    other => {
417                        return Err(EmbeddingError::ParseError(format!(
418                            "Unsupported ONNX model input '{other}': the `local-embeddings` \
419                             feature only supports input_ids / attention_mask / token_type_ids"
420                        )));
421                    }
422                };
423                named.push((name, tensor));
424            }
425            if named.is_empty() {
426                return Err(EmbeddingError::ParseError(
427                    "ONNX model declares no supported inputs".to_string(),
428                ));
429            }
430
431            let outputs = guard
432                .session()?
433                .run(named)
434                .map_err(|e| EmbeddingError::ApiError(format!("ONNX inference failed: {e}")))?;
435
436            let output_value = outputs.get(0).ok_or_else(|| {
437                EmbeddingError::ParseError("ONNX model has no output".to_string())
438            })?;
439            let (shape, data) = output_value.try_extract_tensor::<f32>().map_err(|e| {
440                EmbeddingError::ParseError(format!("Failed to extract output tensor: {e}"))
441            })?;
442            let shape_vec: Vec<usize> = shape.iter().map(|&d| d as usize).collect();
443
444            let mut rows = Self::pool_rows(&shape_vec, data, &mask_rows, batch, fed_seq_len)?;
445            for row in &mut rows {
446                crate::l2_normalize(row);
447            }
448            Ok(rows)
449        }
450
451        /// 从输出张量提取每行向量。
452        ///
453        /// - 3D `[batch, seq, dim]`:按喂入的 attention_mask 做 masked mean pooling
454        ///   (mask=0 的 pad 位置不参与均值);输出序列长与喂入长不同时取较小者。
455        /// - 2D `[batch, dim]`:直接按行切。
456        /// - batch 行数与输入不一致 → 显式 `BatchMismatch`(P0-1 对齐契约)。
457        fn pool_rows(
458            shape: &[usize],
459            data: &[f32],
460            masks: &[Vec<i64>],
461            batch: usize,
462            fed_seq_len: usize,
463        ) -> Result<Vec<Vec<f32>>, EmbeddingError> {
464            match shape.len() {
465                3 => {
466                    let out_batch = shape[0];
467                    if out_batch != batch {
468                        return Err(EmbeddingError::BatchMismatch {
469                            expected: batch,
470                            actual: out_batch,
471                        });
472                    }
473                    let out_seq = shape[1];
474                    let dim = shape[2];
475                    let seq = out_seq.min(fed_seq_len);
476                    let mut result = vec![vec![0.0f32; dim]; batch];
477                    for b in 0..batch {
478                        let mask = &masks[b];
479                        let mut count = 0usize;
480                        for s in 0..seq {
481                            if s >= mask.len() || mask[s] == 0 {
482                                continue; // pad 位置不参与均值
483                            }
484                            count += 1;
485                            let base = (b * out_seq + s) * dim;
486                            for d in 0..dim {
487                                result[b][d] += data[base + d];
488                            }
489                        }
490                        if count > 0 {
491                            for d in 0..dim {
492                                result[b][d] /= count as f32;
493                            }
494                        }
495                    }
496                    Ok(result)
497                }
498                2 => {
499                    let out_batch = shape[0];
500                    if out_batch != batch {
501                        return Err(EmbeddingError::BatchMismatch {
502                            expected: batch,
503                            actual: out_batch,
504                        });
505                    }
506                    let dim = shape[1];
507                    let mut result = Vec::with_capacity(batch);
508                    for b in 0..batch {
509                        let base = b * dim;
510                        result.push(data[base..base + dim].to_vec());
511                    }
512                    Ok(result)
513                }
514                _ => Err(EmbeddingError::ParseError(format!(
515                    "Unsupported output dimension count: {}",
516                    shape.len()
517                ))),
518            }
519        }
520
521        /// 批量执行完整嵌入管线(chunk 内逐批 pad 对齐推理)。
522        fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
523            if texts.is_empty() {
524                return Ok(Vec::new());
525            }
526            let mut results = Vec::with_capacity(texts.len());
527            for chunk in texts.chunks(self.max_batch.max(1)) {
528                let encodings = self.tokenize(chunk)?;
529                results.extend(self.infer_rows(&encodings)?);
530            }
531            Ok(results)
532        }
533
534        /// 单条文本嵌入。
535        fn embed_single(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
536            let mut rows = self.embed_batch(&[text.to_string()])?;
537            rows.pop().ok_or_else(|| {
538                EmbeddingError::ParseError(
539                    "model returned no embedding for single text".to_string(),
540                )
541            })
542        }
543
544        /// 推断输出维度(取输出 shape 最后一个正维度,同旧实现)。
545        fn infer_dimension(session: &ort::session::Session) -> Result<usize, EmbeddingError> {
546            let outputs = session.outputs();
547            if outputs.is_empty() {
548                return Err(EmbeddingError::ParseError(
549                    "ONNX model has no output nodes".to_string(),
550                ));
551            }
552            let dtype = outputs[0].dtype();
553            let shape = dtype.tensor_shape().ok_or_else(|| {
554                EmbeddingError::ParseError("Output is not a Tensor type".to_string())
555            })?;
556            let dim = shape
557                .iter()
558                .rev()
559                .find_map(|&d| if d > 0 { Some(d as usize) } else { None })
560                .ok_or_else(|| {
561                    EmbeddingError::ParseError(format!(
562                        "Cannot infer embedding dimension from model output shape: {:?}",
563                        *shape
564                    ))
565                })?;
566            Ok(dim)
567        }
568
569        /// 从模型第一个输入的静态 shape 推断 batch 上限与序列截断上限(P2-4)。
570        ///
571        /// - 静态 batch=1 → 只能顺序执行(batch_cap=1);
572        /// - 静态 batch=N → 单批上限 `min(N, max_batch)`;
573        /// - 动态 batch(-1)→ 直接用 `max_batch`。
574        /// 序列维同理:静态给值则取该值,动态取 `max_seq_len`。
575        fn infer_input_capability(
576            session: &ort::session::Session,
577            max_batch: usize,
578            max_seq_len: usize,
579        ) -> Result<(usize, usize), EmbeddingError> {
580            let inputs = session.inputs();
581            let input = inputs.first().ok_or_else(|| {
582                EmbeddingError::ParseError("ONNX model has no input nodes".to_string())
583            })?;
584            let shape = input.dtype().tensor_shape().ok_or_else(|| {
585                EmbeddingError::ParseError("Model input is not a Tensor type".to_string())
586            })?;
587            let batch_dim = shape.first().copied().unwrap_or(-1);
588            let seq_dim = shape.get(1).copied().unwrap_or(-1);
589            let batch_cap = if batch_dim > 0 {
590                (batch_dim as usize).min(max_batch)
591            } else {
592                max_batch
593            };
594            let seq_limit = if seq_dim > 0 {
595                seq_dim as usize
596            } else {
597                max_seq_len
598            };
599            Ok((batch_cap, seq_limit))
600        }
601
602        /// 解析 pad token id:优先 tokenizer 显式 padding 配置,其次词表 `[PAD]`,
603        /// 最后回退 0([UNK]/[PAD] 之外的保守选择)。
604        fn resolve_pad_id(tokenizer: &tokenizers::Tokenizer) -> u32 {
605            tokenizer
606                .get_padding()
607                .map(|p| p.pad_id)
608                .or_else(|| tokenizer.token_to_id("[PAD]"))
609                .unwrap_or(0)
610        }
611    }
612
613    /// 在模型同目录发现 tokenizer.json:`<模型名>.json` 优先,其次 `tokenizer.json`。
614    fn discover_tokenizer(model_path: &Path) -> Result<PathBuf, EmbeddingError> {
615        let dir = model_path.parent().unwrap_or_else(|| Path::new("."));
616        let stem = model_path
617            .file_stem()
618            .map(|s| s.to_string_lossy().to_string())
619            .unwrap_or_default();
620        for candidate in [dir.join(format!("{stem}.json")), dir.join("tokenizer.json")] {
621            if candidate.is_file() {
622                return Ok(candidate);
623            }
624        }
625        Err(EmbeddingError::Config(format!(
626            "No tokenizer.json found next to ONNX model '{}'. The `local-embeddings` feature \
627             uses the HuggingFace `tokenizers` crate and requires a real tokenizer.json \
628             (WordPiece/BPE vocab). Place one as '{stem}.json' or 'tokenizer.json' next to the \
629             model, or pass it explicitly via `LocalEmbeddings::from_file_with_tokenizer`. The old \
630             byte-hash fake tokenizer was removed in P2-2: feeding fake token IDs to a neural \
631             embedding model yields garbage vectors.",
632            model_path.display()
633        )))
634    }
635
636    /// ONNX Runtime-based local neural network embedding
637    ///
638    /// 从 ONNX 模型 + HuggingFace `tokenizers`(WordPiece/BPE)加载,本地推理,
639    /// 无外部 API 调用,适合隐私敏感或离线场景。
640    ///
641    /// P2-2:以真实 tokenizer 替换旧的字节 hash 伪 tokenizer,组齐
642    /// `input_ids + attention_mask + token_type_ids`。
643    /// P2-3:`RwLock<Session>` 改为有界 Session 并发池,支持多线程并发推理。
644    /// P2-4:改为按最长序列 pad 对齐的批量推理 + masked mean pooling。
645    ///
646    /// # Example
647    ///
648    /// ```ignore
649    /// use lc_embeddings::LocalEmbeddings;
650    ///
651    /// let embedder = LocalEmbeddings::from_file("model.onnx")?;
652    /// let vec = embedder.embed_query("hello world").await?;
653    /// ```
654    pub struct LocalEmbeddings {
655        inner: Arc<LocalInner>,
656    }
657
658    impl LocalEmbeddings {
659        /// 从 ONNX 模型文件加载;自动在模型同目录发现 tokenizer.json
660        /// (`<模型名>.json` 或 `tokenizer.json`,见 [`LocalEmbeddingsBuilder`])。
661        pub fn from_file(model_path: impl AsRef<Path>) -> Result<Self, EmbeddingError> {
662            Self::builder().model_path(model_path).build()
663        }
664
665        /// 从 ONNX 模型 + 显式 tokenizer.json 加载。
666        pub fn from_file_with_tokenizer(
667            model_path: impl AsRef<Path>,
668            tokenizer_path: impl AsRef<Path>,
669        ) -> Result<Self, EmbeddingError> {
670            Self::builder()
671                .model_path(model_path)
672                .tokenizer_path(tokenizer_path)
673                .build()
674        }
675
676        /// 创建构建器,可定制会话池大小 / 批量上限 / 序列截断上限。
677        pub fn builder() -> LocalEmbeddingsBuilder {
678            LocalEmbeddingsBuilder {
679                model_path: None,
680                tokenizer_path: None,
681                pool_size: default_pool_size(),
682                max_batch: DEFAULT_MAX_BATCH,
683                max_seq_len: DEFAULT_MAX_SEQ_LEN,
684            }
685        }
686    }
687
688    /// [`LocalEmbeddings`] 构建器,用于定制 ONNX 会话池与批量策略。
689    pub struct LocalEmbeddingsBuilder {
690        model_path: Option<PathBuf>,
691        tokenizer_path: Option<PathBuf>,
692        pool_size: usize,
693        max_batch: usize,
694        max_seq_len: usize,
695    }
696
697    impl LocalEmbeddingsBuilder {
698        /// 设置 ONNX 模型路径(必填)。
699        pub fn model_path(mut self, path: impl AsRef<Path>) -> Self {
700            self.model_path = Some(path.as_ref().to_path_buf());
701            self
702        }
703
704        /// 设置 tokenizer.json 路径;缺省时自动在模型同目录发现。
705        pub fn tokenizer_path(mut self, path: impl AsRef<Path>) -> Self {
706            self.tokenizer_path = Some(path.as_ref().to_path_buf());
707            self
708        }
709
710        /// 设置 Session 并发池大小(默认 = CPU 逻辑核数,封顶 8)。
711        pub fn pool_size(mut self, size: usize) -> Self {
712            self.pool_size = size;
713            self
714        }
715
716        /// 设置单次推理的最大文本数(默认 32)。
717        pub fn max_batch(mut self, n: usize) -> Self {
718            self.max_batch = n;
719            self
720        }
721
722        /// 设置序列截断上限(默认 512),超出的 token 会被截断。
723        pub fn max_seq_len(mut self, n: usize) -> Self {
724            self.max_seq_len = n;
725            self
726        }
727
728        /// 构建 `LocalEmbeddings`:加载模型字节 + tokenizer,预建会话池,推断维度与批量能力。
729        pub fn build(self) -> Result<LocalEmbeddings, EmbeddingError> {
730            let model_path = self.model_path.ok_or_else(|| {
731                EmbeddingError::Config(
732                    "model path is required: call LocalEmbeddings::builder().model_path(path)"
733                        .to_string(),
734                )
735            })?;
736            let model_bytes = std::fs::read(&model_path).map_err(|e| {
737                EmbeddingError::ApiError(format!(
738                    "Failed to read ONNX model '{}': {e}",
739                    model_path.display()
740                ))
741            })?;
742            let tokenizer_path = match self.tokenizer_path {
743                Some(p) => p,
744                None => discover_tokenizer(&model_path)?,
745            };
746            let tokenizer = tokenizers::Tokenizer::from_file(&tokenizer_path).map_err(|e| {
747                EmbeddingError::Config(format!(
748                    "Failed to load tokenizer '{}': {e}. Expected a HuggingFace tokenizer.json \
749                     (WordPiece/BPE).",
750                    tokenizer_path.display()
751                ))
752            })?;
753
754            let model_name = model_path
755                .file_stem()
756                .and_then(|s| s.to_str())
757                .unwrap_or("unknown")
758                .to_string();
759
760            let pool = SessionPool::new(model_bytes, self.pool_size)?;
761
762            // 借用首个会话推断输出维度与输入 batch/序列能力。
763            let (dim, max_batch, seq_limit) = {
764                let mut guard = pool.acquire()?;
765                let session = guard.session()?;
766                let dim = LocalInner::infer_dimension(session)?;
767                let (max_batch, seq_limit) =
768                    LocalInner::infer_input_capability(session, self.max_batch, self.max_seq_len)?;
769                (dim, max_batch, seq_limit)
770            };
771
772            Ok(LocalEmbeddings {
773                inner: Arc::new(LocalInner {
774                    pool,
775                    tokenizer,
776                    dim,
777                    model_name,
778                    seq_limit,
779                    max_batch,
780                }),
781            })
782        }
783    }
784
785    #[async_trait]
786    impl Embeddings for LocalEmbeddings {
787        async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
788            // ONNX 推理是 CPU 密集,放入阻塞线程池。捕获 owned `Arc<LocalInner>`
789            // 以满足 spawn_blocking 的 `'static` 约束(修掉旧实现 `&self` 捕获的潜伏错误)。
790            let inner = self.inner.clone();
791            let text = text.to_string();
792            tokio::task::spawn_blocking(move || inner.embed_single(&text))
793                .await
794                .map_err(|e| EmbeddingError::ApiError(format!("Task execution failed: {e}")))?
795        }
796
797        async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
798            if texts.is_empty() {
799                return Ok(Vec::new());
800            }
801            // P1-1: 任一空/全空白文本都报错,与 trait 默认契约一致。
802            if texts.iter().any(|t| t.trim().is_empty()) {
803                return Err(EmbeddingError::EmptyInput);
804            }
805
806            let inner = self.inner.clone();
807            let texts: Vec<String> = texts.iter().map(|s| s.to_string()).collect();
808            tokio::task::spawn_blocking(move || inner.embed_batch(&texts))
809                .await
810                .map_err(|e| EmbeddingError::ApiError(format!("Task execution failed: {e}")))?
811        }
812
813        fn dimension(&self) -> usize {
814            self.inner.dim
815        }
816
817        fn model_name(&self) -> &str {
818            &self.inner.model_name
819        }
820    }
821
822    #[cfg(test)]
823    mod nn_tests {
824        use super::*;
825
826        /// 构建一个最小 WordPiece tokenizer(JSON 走 `Tokenizer::from_bytes`,
827        /// 与真实 tokenizer.json 的加载路径一致)。词表:`[UNK]=0, hello=1, world=2`,
828        /// `with_pad_in_vocab=true` 时追加 `[PAD]=3`。
829        fn tiny_tokenizer(with_pad_in_vocab: bool) -> tokenizers::Tokenizer {
830            let mut vocab = serde_json::json!({
831                "[UNK]": 0,
832                "hello": 1,
833                "world": 2,
834            });
835            if with_pad_in_vocab {
836                vocab["[PAD]"] = serde_json::json!(3);
837            }
838            let json = serde_json::json!({
839                "version": "1.0",
840                "truncation": null,
841                "padding": null,
842                "added_tokens": [],
843                "normalizer": null,
844                "pre_tokenizer": { "type": "Whitespace" },
845                "post_processor": null,
846                "decoder": null,
847                "model": {
848                    "type": "WordPiece",
849                    "vocab": vocab,
850                    "unk_token": "[UNK]",
851                    "continuing_subword_prefix": "##",
852                    "max_input_chars_per_word": 100
853                }
854            });
855            tokenizers::Tokenizer::from_bytes(json.to_string().as_bytes())
856                .expect("tiny tokenizer should deserialize")
857        }
858
859        #[test]
860        fn test_l2_normalize() {
861            let mut v = vec![3.0, 4.0];
862            crate::l2_normalize(&mut v);
863            let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
864            assert!((norm - 1.0).abs() < 1e-5);
865            assert!((v[0] - 0.6).abs() < 1e-5);
866            assert!((v[1] - 0.8).abs() < 1e-5);
867        }
868
869        #[test]
870        fn test_l2_normalize_zero() {
871            let mut v = vec![0.0, 0.0, 0.0];
872            crate::l2_normalize(&mut v);
873            assert!(v.iter().all(|x| *x == 0.0));
874        }
875
876        /// P2-2: 真实 WordPiece tokenizer 输出确定 token ID(无 post_processor,
877        /// `add_special_tokens=true` 也不会追加 [CLS]/[SEP])。
878        #[test]
879        fn test_tokenize_real_wordpiece() {
880            let tok = tiny_tokenizer(false);
881            let enc = tok.encode("hello world", true).unwrap();
882            assert_eq!(enc.get_ids(), &[1u32, 2u32]);
883            assert_eq!(enc.get_attention_mask(), &[1u32, 1u32]);
884        }
885
886        /// P2-2: 词表外单词回退到 [UNK]=0。
887        #[test]
888        fn test_tokenize_unknown_word_uses_unk() {
889            let tok = tiny_tokenizer(false);
890            let enc = tok.encode("zzzznotinvocab", true).unwrap();
891            assert_eq!(enc.get_ids(), &[0u32]);
892        }
893
894        /// P2-4: 批量 pad 对齐——短行补 pad_id、mask 补 0,长行不动。
895        #[test]
896        fn test_build_batch_tensors_pads_to_longest() {
897            let tok = tiny_tokenizer(false);
898            let encodings = tok
899                .encode_batch(vec!["hello".to_string(), "hello world".to_string()], true)
900                .unwrap();
901            let (input_ids, attention_mask, token_type_ids, max_len) =
902                LocalInner::build_batch_tensors(&encodings, 8, 0);
903            assert_eq!(max_len, 2);
904            // "hello" → [1, PAD(0)],"hello world" → [1, 2]
905            assert_eq!(input_ids, vec![1, 0, 1, 2]);
906            assert_eq!(attention_mask, vec![1, 0, 1, 1]);
907            assert_eq!(token_type_ids, vec![0, 0, 0, 0]);
908        }
909
910        #[test]
911        fn test_resolve_pad_id_with_pad_token() {
912            let tok = tiny_tokenizer(true);
913            assert_eq!(LocalInner::resolve_pad_id(&tok), 3);
914        }
915
916        #[test]
917        fn test_resolve_pad_id_defaults_zero() {
918            let tok = tiny_tokenizer(false);
919            assert_eq!(LocalInner::resolve_pad_id(&tok), 0);
920        }
921
922        /// P2-4: 3D masked mean pooling——mask=0 的 pad 位置不参与均值。
923        #[test]
924        fn test_pool_rows_3d_masked() {
925            // shape [2, 3, 2]:两行各 3 个位置、dim=2。
926            let shape = vec![2usize, 3, 2];
927            let data = vec![
928                // row0: tokens [1,2,3]
929                1.0, 10.0, 2.0, 20.0, 3.0, 30.0, // row1
930                4.0, 40.0, 5.0, 50.0, 6.0, 60.0,
931            ];
932            let masks = vec![vec![1, 1, 1], vec![1, 0, 0]];
933            let rows = LocalInner::pool_rows(&shape, &data, &masks, 2, 3).unwrap();
934            assert_eq!(rows.len(), 2);
935            // row0 均值 = ((1+2+3)/3, (10+20+30)/3) = (2, 20)
936            assert!((rows[0][0] - 2.0).abs() < 1e-5);
937            assert!((rows[0][1] - 20.0).abs() < 1e-5);
938            // row1 只取第一个位置 = (4, 40)
939            assert!((rows[1][0] - 4.0).abs() < 1e-5);
940            assert!((rows[1][1] - 40.0).abs() < 1e-5);
941        }
942
943        /// P2-4: 2D `[batch, dim]` 输出直接按行切。
944        #[test]
945        fn test_pool_rows_2d() {
946            let shape = vec![2usize, 2];
947            let data = vec![1.0, 2.0, 3.0, 4.0];
948            let masks = vec![vec![1], vec![1]];
949            let rows = LocalInner::pool_rows(&shape, &data, &masks, 2, 1).unwrap();
950            assert_eq!(rows, vec![vec![1.0, 2.0], vec![3.0, 4.0]]);
951        }
952
953        /// P0-1 对齐契约:模型输出行数与输入 batch 不一致 → 显式 BatchMismatch。
954        #[test]
955        fn test_pool_rows_batch_mismatch() {
956            let shape = vec![3usize, 2, 2];
957            let data = vec![0.0; 12];
958            let masks = vec![vec![1], vec![1]];
959            let err = LocalInner::pool_rows(&shape, &data, &masks, 2, 1).unwrap_err();
960            assert!(matches!(
961                err,
962                EmbeddingError::BatchMismatch {
963                    expected: 2,
964                    actual: 3
965                }
966            ));
967        }
968    }
969}
970
971// When local-embeddings feature is enabled, re-export LocalEmbeddings and its builder
972#[cfg(feature = "local-embeddings")]
973pub use nn::{LocalEmbeddings, LocalEmbeddingsBuilder};
974
975// ---------------------------------------------------------------------------
976// Backward compatibility: LocalEmbeddings without feature points to BagOfWordsEmbeddings
977// ---------------------------------------------------------------------------
978
979/// Without the `local-embeddings` feature, `LocalEmbeddings` is a type alias for `BagOfWordsEmbeddings`,
980/// maintaining backward compatibility.
981///
982/// With the `local-embeddings` feature enabled, `LocalEmbeddings` becomes the ONNX Runtime-based neural network implementation.
983///
984/// P2-1: 消除静默降级。无 feature 时 `LocalEmbeddings` 静默退化为词袋哈希嵌入,
985/// 用户以为在用语义向量、实际是词频——"好像能用,但不对"。这里加
986/// `#[deprecated]` 让降级在编译期可见:使用者需显式改用 `BagOfWordsEmbeddings`,
987/// 或开启 `local-embeddings` feature 使用真正的 ONNX 神经嵌入。
988#[cfg(not(feature = "local-embeddings"))]
989#[deprecated(
990    note = "LocalEmbeddings without the `local-embeddings` feature degrades to \
991            BagOfWordsEmbeddings (bag-of-words hash), not semantic neural embedding. \
992            Enable the `local-embeddings` feature, or use BagOfWordsEmbeddings explicitly."
993)]
994pub type LocalEmbeddings = BagOfWordsEmbeddings;
995
996// ---------------------------------------------------------------------------
997// Tests
998// ---------------------------------------------------------------------------
999
1000#[cfg(test)]
1001mod tests {
1002    use super::*;
1003    use crate::cosine_similarity;
1004
1005    // ---- BagOfWordsEmbeddings tests ----
1006
1007    #[tokio::test]
1008    async fn test_bow_dimension() {
1009        let e = BagOfWordsEmbeddings::new(128);
1010        let v = e.embed_query("hello world").await.unwrap();
1011        assert_eq!(v.len(), 128);
1012        assert_eq!(e.dimension(), 128);
1013    }
1014
1015    #[tokio::test]
1016    async fn test_bow_same_text_same_vector() {
1017        let e = BagOfWordsEmbeddings::new(64);
1018        let a = e.embed_query("rust programming").await.unwrap();
1019        let b = e.embed_query("rust programming").await.unwrap();
1020        assert_eq!(a, b);
1021    }
1022
1023    #[tokio::test]
1024    async fn test_bow_different_text_different_vector() {
1025        let e = BagOfWordsEmbeddings::new(64);
1026        let a = e.embed_query("rust programming").await.unwrap();
1027        let b = e.embed_query("cooking recipe pasta").await.unwrap();
1028        assert_ne!(a, b);
1029    }
1030
1031    #[tokio::test]
1032    async fn test_bow_shared_words_more_similar() {
1033        let e = BagOfWordsEmbeddings::new(256);
1034        let base = e.embed_query("rust programming language").await.unwrap();
1035        let similar = e.embed_query("rust programming tutorial").await.unwrap();
1036        let different = e.embed_query("cooking pasta recipe").await.unwrap();
1037
1038        let sim_similar = cosine_similarity(&base, &similar).unwrap_or(0.0);
1039        let sim_different = cosine_similarity(&base, &different).unwrap_or(0.0);
1040        assert!(
1041            sim_similar > sim_different,
1042            "Shared words should be more similar: {} vs {}",
1043            sim_similar,
1044            sim_different
1045        );
1046    }
1047
1048    #[tokio::test]
1049    async fn test_bow_normalized() {
1050        let e = BagOfWordsEmbeddings::new(64);
1051        let v = e.embed_query("some text here").await.unwrap();
1052        let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
1053        assert!((norm - 1.0).abs() < 1e-5, "norm = {}", norm);
1054    }
1055
1056    #[tokio::test]
1057    async fn test_bow_empty_text_returns_error() {
1058        let e = BagOfWordsEmbeddings::new(64);
1059        let result = e.embed_query("").await;
1060        assert!(result.is_err());
1061        assert!(matches!(result.unwrap_err(), EmbeddingError::EmptyInput));
1062    }
1063
1064    #[tokio::test]
1065    async fn test_bow_chinese_tokenize() {
1066        let e = BagOfWordsEmbeddings::new(128);
1067        let a = e.embed_query("机器学习").await.unwrap();
1068        let b = e.embed_query("机器学习").await.unwrap();
1069        assert_eq!(a, b);
1070        let c = e.embed_query("深度学习").await.unwrap();
1071        let sim = cosine_similarity(&a, &c).unwrap_or(0.0);
1072        assert!(
1073            sim > 0.0,
1074            "Shared '学习' should have positive similarity: {}",
1075            sim
1076        );
1077    }
1078
1079    #[test]
1080    fn test_bow_tokenize_english() {
1081        let t = BagOfWordsEmbeddings::tokenize("Hello, World! 123");
1082        assert!(t.contains(&"hello".to_string()));
1083        assert!(t.contains(&"world".to_string()));
1084        assert!(t.contains(&"123".to_string()));
1085    }
1086
1087    #[test]
1088    fn test_bow_tokenize_chinese() {
1089        let t = BagOfWordsEmbeddings::tokenize("机器学习");
1090        assert!(t.contains(&"机".to_string()));
1091        assert!(t.contains(&"学".to_string()));
1092        assert_eq!(t.len(), 4);
1093    }
1094
1095    #[test]
1096    fn test_bow_model_name() {
1097        let e = BagOfWordsEmbeddings::default_dim();
1098        assert_eq!(e.model_name(), "local-bow");
1099    }
1100
1101    // ---- LocalEmbeddings backward compatibility test (without feature, is BagOfWordsEmbeddings alias) ----
1102
1103    /// P2-1: 该测试正是验证"无 feature 时 LocalEmbeddings = BagOfWordsEmbeddings",
1104    /// 是有意使用已弃用别名,`#[allow(deprecated)]` 豁免降级警告。
1105    #[allow(deprecated)]
1106    #[tokio::test]
1107    async fn test_local_embeddings_backward_compat() {
1108        // Without feature, LocalEmbeddings = BagOfWordsEmbeddings
1109        let e = LocalEmbeddings::new(64);
1110        let v = e.embed_query("test backward compat").await.unwrap();
1111        assert_eq!(v.len(), 64);
1112        assert_eq!(e.model_name(), "local-bow");
1113    }
1114}