Coverage for agentos/vectorstore/db.py: 22%

153 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 10:59 +0800

1""" 

2AgentOS v0.30 向量数据库集成 — Chroma + FAISS。 

3语义记忆检索、知识库索引。 

4""" 

5 

6from dataclasses import dataclass, field 

7import os 

8import pickle 

9import uuid 

10 

11 

12@dataclass 

13class VectorEntry: 

14 """向量条目。""" 

15 id: str 

16 text: str 

17 metadata: dict = field(default_factory=dict) 

18 score: float = 0.0 

19 

20 

21class BaseVectorStore: 

22 """向量存储基类。""" 

23 def add(self, texts: list[str], metadatas: list[dict] | None = None, ids: list[str] | None = None) -> list[str]: ... 

24 def search(self, query: str, top_k: int = 5) -> list[VectorEntry]: ... 

25 def delete(self, ids: list[str]): ... 

26 def count(self) -> int: ... 

27 

28 

29class FAISSVectorStore(BaseVectorStore): 

30 """基于 FAISS 的轻量向量存储。""" 

31 

32 def __init__(self, dim: int = 768, index_path: str = ""): 

33 self.dim = dim 

34 self.index_path = index_path 

35 self._index = None 

36 self._store: dict[str, tuple[list[float], str, dict]] = {} 

37 self._next_id = 0 

38 if index_path and os.path.exists(index_path): 

39 self._load() 

40 

41 def _init_index(self): 

42 try: 

43 import faiss 

44 self._index = faiss.IndexFlatIP(self.dim) 

45 except ImportError: 

46 self._index = None 

47 

48 def add(self, texts: list[str], metadatas: list[dict] | None = None, ids: list[str] | None = None) -> list[str]: 

49 embeddings = self._embed(texts) 

50 if not self._index: 

51 self._init_index() 

52 if self._index: 

53 import numpy as np 

54 vecs = np.array(embeddings, dtype=np.float32) 

55 self._index.add(vecs) 

56 

57 res_ids = [] 

58 for i, text in enumerate(texts): 

59 rid = ids[i] if ids else f"v{self._next_id}" 

60 self._next_id += 1 

61 self._store[rid] = (embeddings[i], text, metadatas[i] if metadatas else {}) 

62 res_ids.append(rid) 

63 return res_ids 

64 

65 def _fallback_search(self, q_vec, top_k): 

66 """Fallback余弦相似度搜索(无faiss时使用)。""" 

67 import math 

68 scores = [] 

69 for rid, (vec, text, meta) in self._store.items(): 

70 dot = sum(a*b for a,b in zip(q_vec, vec)) 

71 na = math.sqrt(sum(a*a for a in q_vec)) 

72 nb = math.sqrt(sum(b*b for b in vec)) 

73 sim = dot/(na*nb) if na*nb > 0 else 0.0 

74 scores.append((sim, rid, text, meta)) 

75 scores.sort(key=lambda x: x[0], reverse=True) 

76 return [VectorEntry(id=rid, text=text, metadata=meta, score=float(s)) 

77 for s, rid, text, meta in scores[:top_k]] 

78 

79 def search(self, query: str, top_k: int = 5) -> list[VectorEntry]: 

80 if not self._store: 

81 return [] 

82 q_vec = self._embed([query])[0] 

83 if not self._index: 

84 return self._fallback_search(q_vec, top_k) 

85 q_vec = self._embed([query])[0] 

86 import numpy as np 

87 D, I = self._index.search(np.array([q_vec], dtype=np.float32), min(top_k, self.count())) 

88 results = [] 

89 for score, idx in zip(D[0], I[0]): 

90 if idx < 0: 

91 continue 

92 rid = f"v{idx}" 

93 if rid in self._store: 

94 _, text, meta = self._store[rid] 

95 results.append(VectorEntry(id=rid, text=text, metadata=meta, score=float(score))) 

96 return results 

97 

98 def delete(self, ids: list[str]): 

99 for rid in ids: 

100 self._store.pop(rid, None) 

101 

102 def count(self) -> int: 

103 return len(self._store) 

104 

105 def _embed(self, texts: list[str]) -> list[list[float]]: 

106 """轻量嵌入:使用 all-MiniLM-L6-v2 或回退到 TF-IDF。""" 

107 try: 

108 from sentence_transformers import SentenceTransformer 

109 model = SentenceTransformer("all-MiniLM-L6-v2") 

110 embeddings = model.encode(texts, normalize_embeddings=True) 

111 return embeddings.tolist() 

112 except ImportError: 

113 return self._tfidf_embed(texts) 

114 

115 def _tfidf_embed(self, texts: list[str]) -> list[list[float]]: 

116 """TF-IDF 回退,仅作占位。""" 

117 import hashlib 

118 dim = self.dim 

119 result = [] 

120 for t in texts: 

121 h = hashlib.sha256(t.encode()).digest() 

122 vec = [(h[i] / 255.0) for i in range(min(len(h), dim))] 

123 vec += [0.0] * (dim - len(vec)) 

124 result.append(vec) 

125 return result 

126 

127 def _save(self): 

128 if self.index_path: 

129 os.makedirs(os.path.dirname(self.index_path) or ".", exist_ok=True) 

130 with open(self.index_path, "wb") as f: 

131 pickle.dump({"store": self._store, "next_id": self._next_id}, f) 

132 

133 def _load(self): 

134 with open(self.index_path, "rb") as f: 

135 data = pickle.load(f) 

136 self._store = data["store"] 

137 self._next_id = data["next_id"] 

138 

139 def __del__(self): 

140 if self.index_path: 

141 self._save() 

142 

143 

144class ChromaVectorStore(BaseVectorStore): 

145 """Chroma 向量存储。""" 

146 

147 def __init__(self, collection_name: str = "agentos", persist_dir: str = "./chroma_data"): 

148 self.collection_name = collection_name 

149 self.persist_dir = persist_dir 

150 self._client = None 

151 self._collection = None 

152 self._init() 

153 

154 def _init(self): 

155 try: 

156 import chromadb 

157 self._client = chromadb.PersistentClient(path=self.persist_dir) 

158 self._collection = self._client.get_or_create_collection(self.collection_name) 

159 except ImportError: 

160 self._collection = None 

161 

162 def add(self, texts: list[str], metadatas: list[dict] | None = None, ids: list[str] | None = None) -> list[str]: 

163 if not self._collection: 

164 ids = ids or [f"v{len(self._fallback_store)}-{i}" for i in range(len(texts))] 

165 for i, t in enumerate(texts): 

166 self._fallback_store[ids[i]] = {"text": t, "metadata": metadatas[i] if metadatas else {}} 

167 return ids 

168 

169 ids = ids or [str(uuid.uuid4())[:8] for _ in texts] 

170 self._collection.add(documents=texts, metadatas=metadatas or [{}] * len(texts), ids=ids) 

171 return ids 

172 

173 def search(self, query: str, top_k: int = 5) -> list[VectorEntry]: 

174 if not self._collection: 

175 if self._fallback_store: 

176 return [ 

177 VectorEntry(id=k, text=v["text"], metadata=v["metadata"], score=0.5) 

178 for k, v in list(self._fallback_store.items())[:top_k] 

179 ] 

180 return [] 

181 results = self._collection.query(query_texts=[query], n_results=top_k) 

182 entries = [] 

183 for i, rid in enumerate(results.get("ids", [[]])[0]): 

184 entries.append( 

185 VectorEntry( 

186 id=rid, 

187 text=results["documents"][0][i] if results.get("documents") else "", 

188 metadata=results["metadatas"][0][i] if results.get("metadatas") else {}, 

189 score=1.0 - results["distances"][0][i] if results.get("distances") else 0.0, 

190 ) 

191 ) 

192 return entries 

193 

194 def delete(self, ids: list[str]): 

195 if self._collection: 

196 self._collection.delete(ids=ids) 

197 else: 

198 for rid in ids: 

199 self._fallback_store.pop(rid, None) 

200 

201 def count(self) -> int: 

202 if self._collection: 

203 return self._collection.count() 

204 return len(self._fallback_store) 

205 

206 @property 

207 def _fallback_store(self) -> dict: 

208 if not hasattr(self, "_fb"): 

209 self._fb = {} 

210 return self._fb