Coverage for agentos/cache/embedder.py: 32%

108 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 01:44 +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 

19 vector: list[float] 

20 tokens: int = 0 

21 model: str = "" 

22 

23 def __len__(self) -> int: 

24 return len(self.vector) 

25 

26 def __iter__(self): 

27 return iter(self.vector) 

28 

29 def __getitem__(self, idx): 

30 return self.vector[idx] 

31 

32 

33class BaseEmbedder(ABC): 

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

35 

36 @abstractmethod 

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

38 

39 @abstractmethod 

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

41 

42 @abstractmethod 

43 def dimension(self) -> int: ... 

44 

45 

46class OpenAIEmbedder(BaseEmbedder): 

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

48 

49 MODELS = { 

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

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

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

53 } 

54 

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 ) 

71 

72 def dimension(self) -> int: 

73 return self._dim 

74 

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

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

77 return results[0] 

78 

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 

93 

94 async def close(self): 

95 await self._http.aclose() 

96 

97 

98class LocalEmbedder(BaseEmbedder): 

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

100 

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 

105 

106 def _ensure_model(self): 

107 if self._model is None: 

108 from sentence_transformers import SentenceTransformer 

109 

110 self._model = SentenceTransformer(self.model_name) 

111 self._dim = self._model.get_sentence_embedding_dimension() 

112 

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 

122 

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) 

127 

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] 

132 

133 

134class CohereEmbedder(BaseEmbedder): 

135 """Cohere Embed API.""" 

136 

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) 

148 

149 def dimension(self) -> int: 

150 return self._dim 

151 

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) 

158 

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

165 

166 async def close(self): 

167 await self._http.aclose() 

168 

169 

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

181 

182 

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)