Coverage for agentos/cache/embedder.py: 32%
108 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +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."""
19 vector: list[float]
20 tokens: int = 0
21 model: str = ""
23 def __len__(self) -> int:
24 return len(self.vector)
26 def __iter__(self):
27 return iter(self.vector)
29 def __getitem__(self, idx):
30 return self.vector[idx]
33class BaseEmbedder(ABC):
34 """Embedding提供者抽象基类。"""
36 @abstractmethod
37 async def embed(self, text: str) -> EmbeddingResult: ...
39 @abstractmethod
40 async def embed_batch(self, texts: list[str]) -> list[EmbeddingResult]: ...
42 @abstractmethod
43 def dimension(self) -> int: ...
46class OpenAIEmbedder(BaseEmbedder):
47 """OpenAI text-embedding-3-small / text-embedding-3-large."""
49 MODELS = {
50 "small": ("text-embedding-3-small", 1536),
51 "large": ("text-embedding-3-large", 3072),
52 "ada": ("text-embedding-ada-002", 1536),
53 }
55 def __init__(
56 self, model: str = "small", api_key: str = "", base_url: str = "https://api.openai.com/v1"
57 ):
58 info = self.MODELS.get(model)
59 if not info:
60 raise ValueError(f"Unknown model key: {model}. Use: {list(self.MODELS.keys())}")
61 self.model_id, self._dim = info
62 self.api_key = api_key or os.environ.get("OPENAI_API_KEY", "")
63 self.base_url = base_url
64 self._http = httpx.AsyncClient(
65 timeout=60,
66 headers={
67 "Authorization": f"Bearer {self.api_key}",
68 "Content-Type": "application/json",
69 },
70 )
72 def dimension(self) -> int:
73 return self._dim
75 async def embed(self, text: str) -> EmbeddingResult:
76 results = await self.embed_batch([text])
77 return results[0]
79 async def embed_batch(self, texts: list[str]) -> list[EmbeddingResult]:
80 body = {"model": self.model_id, "input": texts}
81 resp = await self._http.post(f"{self.base_url}/embeddings", json=body)
82 resp.raise_for_status()
83 data = resp.json()
84 results = []
85 for item in data["data"]:
86 results.append(
87 EmbeddingResult(
88 vector=item["embedding"],
89 model=self.model_id,
90 )
91 )
92 return results
94 async def close(self):
95 await self._http.aclose()
98class LocalEmbedder(BaseEmbedder):
99 """本地sentence-transformers模型。无API调用,零成本。"""
101 def __init__(self, model_name: str = "all-MiniLM-L6-v2"):
102 self.model_name = model_name
103 self._model = None
104 self._dim = 384
106 def _ensure_model(self):
107 if self._model is None:
108 from sentence_transformers import SentenceTransformer
110 self._model = SentenceTransformer(self.model_name)
111 self._dim = self._model.get_sentence_embedding_dimension()
113 def dimension(self) -> int:
114 if self._model is None:
115 if self.model_name == "all-MiniLM-L6-v2":
116 self._dim = 384
117 elif "large" in self.model_name:
118 self._dim = 1024
119 else:
120 self._dim = 768
121 return self._dim
123 async def embed(self, text: str) -> EmbeddingResult:
124 self._ensure_model()
125 vec = self._model.encode(text, normalize_embeddings=True)
126 return EmbeddingResult(vector=vec.tolist(), model=self.model_name)
128 async def embed_batch(self, texts: list[str]) -> list[EmbeddingResult]:
129 self._ensure_model()
130 vecs = self._model.encode(texts, normalize_embeddings=True)
131 return [EmbeddingResult(vector=v.tolist(), model=self.model_name) for v in vecs]
134class CohereEmbedder(BaseEmbedder):
135 """Cohere Embed API."""
137 def __init__(self, model: str = "embed-english-v3.0", api_key: str = ""):
138 self.model_id = model
139 self.api_key = api_key or os.environ.get("COHERE_API_KEY", "")
140 self._http = httpx.AsyncClient(
141 timeout=60,
142 headers={
143 "Authorization": f"Bearer {self.api_key}",
144 "Content-Type": "application/json",
145 },
146 )
147 self._dim = {"embed-english-v3.0": 1024, "embed-multilingual-v3.0": 1024}.get(model, 1024)
149 def dimension(self) -> int:
150 return self._dim
152 async def embed(self, text: str) -> EmbeddingResult:
153 body = {"model": self.model_id, "texts": [text], "input_type": "search_document"}
154 resp = await self._http.post("https://api.cohere.ai/v1/embed", json=body)
155 resp.raise_for_status()
156 data = resp.json()
157 return EmbeddingResult(vector=data["embeddings"][0], model=self.model_id)
159 async def embed_batch(self, texts: list[str]) -> list[EmbeddingResult]:
160 body = {"model": self.model_id, "texts": texts, "input_type": "search_document"}
161 resp = await self._http.post("https://api.cohere.ai/v1/embed", json=body)
162 resp.raise_for_status()
163 data = resp.json()
164 return [EmbeddingResult(vector=vec, model=self.model_id) for vec in data["embeddings"]]
166 async def close(self):
167 await self._http.aclose()
170async def get_embedder(provider: str = "openai", **kwargs) -> BaseEmbedder:
171 """工厂函数:获取embedder实例。"""
172 match provider:
173 case "openai":
174 return OpenAIEmbedder(**kwargs)
175 case "local":
176 return LocalEmbedder(**kwargs)
177 case "cohere":
178 return CohereEmbedder(**kwargs)
179 case _:
180 raise ValueError(f"Unknown embedder provider: {provider}. Use: openai/local/cohere")
183async def cosine_similarity(a: list[float], b: list[float]) -> float:
184 """余弦相似度。"""
185 dot = sum(x * y for x, y in zip(a, b))
186 norm_a = sum(x * x for x in a) ** 0.5
187 norm_b = sum(x * x for x in b) ** 0.5
188 if norm_a == 0 or norm_b == 0:
189 return 0.0
190 return dot / (norm_a * norm_b)