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
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
1"""
2AgentOS v0.30 向量数据库集成 — Chroma + FAISS。
3语义记忆检索、知识库索引。
4"""
6from dataclasses import dataclass, field
7import os
8import pickle
9import uuid
12@dataclass
13class VectorEntry:
14 """向量条目。"""
15 id: str
16 text: str
17 metadata: dict = field(default_factory=dict)
18 score: float = 0.0
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: ...
29class FAISSVectorStore(BaseVectorStore):
30 """基于 FAISS 的轻量向量存储。"""
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()
41 def _init_index(self):
42 try:
43 import faiss
44 self._index = faiss.IndexFlatIP(self.dim)
45 except ImportError:
46 self._index = None
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)
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
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]]
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
98 def delete(self, ids: list[str]):
99 for rid in ids:
100 self._store.pop(rid, None)
102 def count(self) -> int:
103 return len(self._store)
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)
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
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)
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"]
139 def __del__(self):
140 if self.index_path:
141 self._save()
144class ChromaVectorStore(BaseVectorStore):
145 """Chroma 向量存储。"""
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()
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
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
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
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
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)
201 def count(self) -> int:
202 if self._collection:
203 return self._collection.count()
204 return len(self._fallback_store)
206 @property
207 def _fallback_store(self) -> dict:
208 if not hasattr(self, "_fb"):
209 self._fb = {}
210 return self._fb