Coverage for agentos/cache/embedder.py: 32%
108 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"""
2Embedding实现层 — 多种embedding provider的真实调用。
3v0.50: 新增模块。为语义缓存/向量数据库提供embedding实现。
4"""
6from __future__ import annotations
8import os
9from abc import ABC, abstractmethod
10from dataclasses import dataclass
12import httpx
15@dataclass
16class EmbeddingResult:
17 """Result of an embedding generation request."""
18 vector: list[float]
19 tokens: int = 0
20 model: str = ""
22 def __len__(self) -> int:
23 return len(self.vector)
25 def __iter__(self):
26 return iter(self.vector)
28 def __getitem__(self, idx):
29 return self.vector[idx]
32class BaseEmbedder(ABC):
33 """Embedding提供者抽象基类。"""
35 @abstractmethod
36 async def embed(self, text: str) -> EmbeddingResult:
37 ...
39 @abstractmethod
40 async def embed_batch(self, texts: list[str]) -> list[EmbeddingResult]:
41 ...
43 @abstractmethod
44 def dimension(self) -> int:
45 ...
48class OpenAIEmbedder(BaseEmbedder):
49 """OpenAI text-embedding-3-small / text-embedding-3-large."""
51 MODELS = {
52 "small": ("text-embedding-3-small", 1536),
53 "large": ("text-embedding-3-large", 3072),
54 "ada": ("text-embedding-ada-002", 1536),
55 }
57 def __init__(self, model: str = "small", api_key: str = "",
58 base_url: str = "https://api.openai.com/v1"):
59 info = self.MODELS.get(model)
60 if not info:
61 raise ValueError(f"Unknown model key: {model}. Use: {list(self.MODELS.keys())}")
62 self.model_id, self._dim = info
63 self.api_key = api_key or os.environ.get("OPENAI_API_KEY", "")
64 self.base_url = base_url
65 self._http = httpx.AsyncClient(timeout=60, headers={
66 "Authorization": f"Bearer {self.api_key}",
67 "Content-Type": "application/json",
68 })
70 def dimension(self) -> int:
71 return self._dim
73 async def embed(self, text: str) -> EmbeddingResult:
74 results = await self.embed_batch([text])
75 return results[0]
77 async def embed_batch(self, texts: list[str]) -> list[EmbeddingResult]:
78 body = {"model": self.model_id, "input": texts}
79 resp = await self._http.post(f"{self.base_url}/embeddings", json=body)
80 resp.raise_for_status()
81 data = resp.json()
82 results = []
83 for item in data["data"]:
84 results.append(EmbeddingResult(
85 vector=item["embedding"],
86 model=self.model_id,
87 ))
88 return results
90 async def close(self):
91 await self._http.aclose()
94class LocalEmbedder(BaseEmbedder):
95 """本地sentence-transformers模型。无API调用,零成本。"""
97 def __init__(self, model_name: str = "all-MiniLM-L6-v2"):
98 self.model_name = model_name
99 self._model = None
100 self._dim = 384
102 def _ensure_model(self):
103 if self._model is None:
104 from sentence_transformers import SentenceTransformer
105 self._model = SentenceTransformer(self.model_name)
106 self._dim = self._model.get_sentence_embedding_dimension()
108 def dimension(self) -> int:
109 if self._model is None:
110 if self.model_name == "all-MiniLM-L6-v2":
111 self._dim = 384
112 elif "large" in self.model_name:
113 self._dim = 1024
114 else:
115 self._dim = 768
116 return self._dim
118 async def embed(self, text: str) -> EmbeddingResult:
119 self._ensure_model()
120 vec = self._model.encode(text, normalize_embeddings=True)
121 return EmbeddingResult(vector=vec.tolist(), model=self.model_name)
123 async def embed_batch(self, texts: list[str]) -> list[EmbeddingResult]:
124 self._ensure_model()
125 vecs = self._model.encode(texts, normalize_embeddings=True)
126 return [
127 EmbeddingResult(vector=v.tolist(), model=self.model_name)
128 for v in vecs
129 ]
132class CohereEmbedder(BaseEmbedder):
133 """Cohere Embed API."""
135 def __init__(self, model: str = "embed-english-v3.0", api_key: str = ""):
136 self.model_id = model
137 self.api_key = api_key or os.environ.get("COHERE_API_KEY", "")
138 self._http = httpx.AsyncClient(timeout=60, headers={
139 "Authorization": f"Bearer {self.api_key}",
140 "Content-Type": "application/json",
141 })
142 self._dim = {"embed-english-v3.0": 1024, "embed-multilingual-v3.0": 1024}.get(model, 1024)
144 def dimension(self) -> int:
145 return self._dim
147 async def embed(self, text: str) -> EmbeddingResult:
148 body = {"model": self.model_id, "texts": [text], "input_type": "search_document"}
149 resp = await self._http.post("https://api.cohere.ai/v1/embed", json=body)
150 resp.raise_for_status()
151 data = resp.json()
152 return EmbeddingResult(vector=data["embeddings"][0], model=self.model_id)
154 async def embed_batch(self, texts: list[str]) -> list[EmbeddingResult]:
155 body = {"model": self.model_id, "texts": texts, "input_type": "search_document"}
156 resp = await self._http.post("https://api.cohere.ai/v1/embed", json=body)
157 resp.raise_for_status()
158 data = resp.json()
159 return [
160 EmbeddingResult(vector=vec, model=self.model_id)
161 for vec in data["embeddings"]
162 ]
164 async def close(self):
165 await self._http.aclose()
168async def get_embedder(provider: str = "openai", **kwargs) -> BaseEmbedder:
169 """工厂函数:获取embedder实例。"""
170 match provider:
171 case "openai":
172 return OpenAIEmbedder(**kwargs)
173 case "local":
174 return LocalEmbedder(**kwargs)
175 case "cohere":
176 return CohereEmbedder(**kwargs)
177 case _:
178 raise ValueError(f"Unknown embedder provider: {provider}. Use: openai/local/cohere")
181async def cosine_similarity(a: list[float], b: list[float]) -> float:
182 """余弦相似度。"""
183 dot = sum(x * y for x, y in zip(a, b))
184 norm_a = sum(x * x for x in a) ** 0.5
185 norm_b = sum(x * x for x in b) ** 0.5
186 if norm_a == 0 or norm_b == 0:
187 return 0.0
188 return dot / (norm_a * norm_b)