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

1""" 

2Embedding实现层 — 多种embedding provider的真实调用。 

3v0.50: 新增模块。为语义缓存/向量数据库提供embedding实现。 

4""" 

5 

6from __future__ import annotations 

7 

8import os 

9from abc import ABC, abstractmethod 

10from dataclasses import dataclass 

11 

12import httpx 

13 

14 

15@dataclass 

16class EmbeddingResult: 

17 """Result of an embedding generation request.""" 

18 vector: list[float] 

19 tokens: int = 0 

20 model: str = "" 

21 

22 def __len__(self) -> int: 

23 return len(self.vector) 

24 

25 def __iter__(self): 

26 return iter(self.vector) 

27 

28 def __getitem__(self, idx): 

29 return self.vector[idx] 

30 

31 

32class BaseEmbedder(ABC): 

33 """Embedding提供者抽象基类。""" 

34 

35 @abstractmethod 

36 async def embed(self, text: str) -> EmbeddingResult: 

37 ... 

38 

39 @abstractmethod 

40 async def embed_batch(self, texts: list[str]) -> list[EmbeddingResult]: 

41 ... 

42 

43 @abstractmethod 

44 def dimension(self) -> int: 

45 ... 

46 

47 

48class OpenAIEmbedder(BaseEmbedder): 

49 """OpenAI text-embedding-3-small / text-embedding-3-large.""" 

50 

51 MODELS = { 

52 "small": ("text-embedding-3-small", 1536), 

53 "large": ("text-embedding-3-large", 3072), 

54 "ada": ("text-embedding-ada-002", 1536), 

55 } 

56 

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 }) 

69 

70 def dimension(self) -> int: 

71 return self._dim 

72 

73 async def embed(self, text: str) -> EmbeddingResult: 

74 results = await self.embed_batch([text]) 

75 return results[0] 

76 

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 

89 

90 async def close(self): 

91 await self._http.aclose() 

92 

93 

94class LocalEmbedder(BaseEmbedder): 

95 """本地sentence-transformers模型。无API调用,零成本。""" 

96 

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 

101 

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() 

107 

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 

117 

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) 

122 

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 ] 

130 

131 

132class CohereEmbedder(BaseEmbedder): 

133 """Cohere Embed API.""" 

134 

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) 

143 

144 def dimension(self) -> int: 

145 return self._dim 

146 

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) 

153 

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 ] 

163 

164 async def close(self): 

165 await self._http.aclose() 

166 

167 

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") 

179 

180 

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)