Skip to main content

xz_embed/
traits.rs

1use async_trait::async_trait;
2use serde::{Deserialize, Serialize};
3use std::fmt::Debug;
4
5use crate::error::{EmbedError, StoreError};
6use crate::types::{MetadataFilter, SearchResult, StoreStats, VectorEntry};
7
8/// 文本截断策略
9#[derive(Debug, Clone, Serialize, Deserialize)]
10pub enum TruncationStrategy {
11    /// 直接报错(严格要求)
12    Error,
13    /// 从开头截断(保留前 N 个 token)
14    TruncateStart,
15    /// 从末尾截断(保留后 N 个 token)
16    TruncateEnd,
17    /// 保留首尾各 N/2 个 token(适合 LLM 输出)
18    TruncateBoth,
19}
20
21/// 嵌入模型定价信息
22#[derive(Debug, Clone, Serialize, Deserialize)]
23pub struct EmbedPricing {
24    /// 每百万 token 价格(输入)
25    pub input_per_million: f64,
26}
27
28/// 嵌入模型元信息
29#[derive(Debug, Clone, Serialize, Deserialize)]
30pub struct EmbedModelInfo {
31    pub name: String,
32    pub display_name: String,
33    /// 支持的维度列表(None 表示不支持维度选择)
34    pub supported_dimensions: Option<Vec<usize>>,
35    /// 当前维度
36    pub current_dimension: usize,
37    /// 最大输入 token 限制
38    pub max_input_tokens: usize,
39    /// 最大批次大小
40    pub max_batch_size: usize,
41    /// 定价信息
42    pub pricing: EmbedPricing,
43}
44
45/// 统一的文本向量嵌入接口
46#[async_trait]
47pub trait EmbeddingModel: Send + Sync + Debug {
48    /// 核心嵌入方法。对一批文本生成向量。
49    async fn embed(&self, input: &[&str]) -> Result<Vec<Vec<f32>>, EmbedError>;
50
51    /// 便捷方法:对单个文本生成向量
52    async fn embed_single(&self, text: &str) -> Result<Vec<f32>, EmbedError> {
53        let mut results = self.embed(&[text]).await?;
54        if results.is_empty() {
55            return Err(EmbedError::Model("embed_single 返回空结果".into()));
56        }
57        Ok(results.remove(0))
58    }
59
60    /// 模型信息
61    fn model_info(&self) -> &EmbedModelInfo;
62
63    /// 最大批次大小
64    fn max_batch_size(&self) -> usize;
65
66    /// 向量维度
67    fn dimensions(&self) -> usize;
68}
69
70/// 向量存储抽象 — 可插拔后端
71#[async_trait]
72pub trait VectorStore: Send + Sync + Debug {
73    /// 插入单条向量
74    async fn insert(&self, entry: VectorEntry) -> Result<(), StoreError>;
75
76    /// 批量插入向量(推荐的高吞吐写入路径)
77    async fn insert_batch(&self, entries: Vec<VectorEntry>) -> Result<(), StoreError>;
78
79    /// 相似度搜索(cosine 相似度,返回 Top-K)
80    async fn search(&self, query: &[f32], limit: usize) -> Result<Vec<SearchResult>, StoreError>;
81
82    /// 带元数据过滤的相似度搜索
83    async fn search_with_filter(
84        &self,
85        query: &[f32],
86        filter: &MetadataFilter,
87        limit: usize,
88    ) -> Result<Vec<SearchResult>, StoreError>;
89
90    /// 按 ID 批量删除
91    async fn delete(&self, ids: &[String]) -> Result<usize, StoreError>;
92
93    /// 按元数据过滤条件删除
94    async fn delete_by_filter(&self, filter: &MetadataFilter) -> Result<usize, StoreError>;
95
96    /// 清空存储
97    async fn clear(&self) -> Result<(), StoreError>;
98
99    /// 存储中的总条目数
100    async fn count(&self) -> Result<usize, StoreError>;
101
102    /// 创建/重建索引(后台线程,不阻塞写入)
103    async fn rebuild_index(&self) -> Result<(), StoreError>;
104
105    /// 存储统计信息
106    async fn stats(&self) -> Result<StoreStats, StoreError>;
107}
108
109/// VectorStore 生命周期 trait
110#[async_trait]
111pub trait StoreLifecycle: VectorStore {
112    /// 初始化存储(创建表结构等)
113    async fn initialize(&self) -> Result<(), StoreError>;
114
115    /// 优雅关闭
116    async fn close(&self) -> Result<(), StoreError>;
117
118    /// 强制持久化检查点
119    async fn checkpoint(&self) -> Result<(), StoreError>;
120
121    /// 是否健康
122    async fn health_check(&self) -> Result<bool, StoreError>;
123}