embeddings.py 2.6 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677
  1. from __future__ import annotations
  2. import hashlib
  3. from typing import Any
  4. import httpx
  5. from .config import settings
  6. class EmbeddingUnavailable(RuntimeError):
  7. pass
  8. def embedding_fingerprint(content: str) -> str:
  9. """文本或模型配置变化时生成新的向量指纹。"""
  10. payload = (
  11. f"{settings.embedding_base_url}\0{settings.embedding_model}\0"
  12. f"{settings.embedding_dimensions}\0{content}"
  13. )
  14. return hashlib.sha256(payload.encode("utf-8")).hexdigest()
  15. class OpenAICompatibleEmbeddingClient:
  16. """调用 OpenAI 兼容的 Embedding 接口。"""
  17. @property
  18. def configured(self) -> bool:
  19. return bool(
  20. settings.embedding_api_key
  21. and settings.embedding_base_url
  22. and settings.embedding_model
  23. )
  24. async def embed(self, inputs: list[str]) -> list[list[float]]:
  25. """批量生成向量,并校验返回数量和维度。"""
  26. if not self.configured:
  27. raise EmbeddingUnavailable(
  28. "未配置 EMBEDDING_API_KEY / EMBEDDING_BASE_URL / EMBEDDING_MODEL"
  29. )
  30. if not inputs:
  31. return []
  32. payload: dict[str, Any] = {
  33. "model": settings.embedding_model,
  34. "input": inputs,
  35. "encoding_format": "float",
  36. }
  37. if settings.embedding_dimensions:
  38. payload["dimensions"] = settings.embedding_dimensions
  39. async with httpx.AsyncClient(timeout=60) as client:
  40. response = await client.post(
  41. f"{settings.embedding_base_url}/embeddings",
  42. headers={"Authorization": f"Bearer {settings.embedding_api_key}"},
  43. json=payload,
  44. )
  45. response.raise_for_status()
  46. body = response.json()
  47. data = body.get("data")
  48. if not isinstance(data, list):
  49. raise EmbeddingUnavailable("Embedding API 返回中缺少 data 数组")
  50. ordered = sorted(data, key=lambda item: item.get("index", 0))
  51. vectors = [item.get("embedding") for item in ordered]
  52. if len(vectors) != len(inputs) or not all(isinstance(item, list) for item in vectors):
  53. raise EmbeddingUnavailable("Embedding API 返回的向量数量不匹配")
  54. if settings.embedding_dimensions and any(
  55. len(vector) != settings.embedding_dimensions for vector in vectors
  56. ):
  57. raise EmbeddingUnavailable(
  58. f"Embedding 向量维度与配置不一致,期望 {settings.embedding_dimensions}"
  59. )
  60. return [[float(value) for value in vector] for vector in vectors]
  61. embedding_client = OpenAICompatibleEmbeddingClient()