embeddings.py 2.6 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576
  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. """Invalidate stored vectors when either the text or model config changes."""
  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. """Embedding client for DashScope and other OpenAI-compatible endpoints."""
  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. if not self.configured:
  26. raise EmbeddingUnavailable(
  27. "未配置 EMBEDDING_API_KEY / EMBEDDING_BASE_URL / EMBEDDING_MODEL"
  28. )
  29. if not inputs:
  30. return []
  31. payload: dict[str, Any] = {
  32. "model": settings.embedding_model,
  33. "input": inputs,
  34. "encoding_format": "float",
  35. }
  36. if settings.embedding_dimensions:
  37. payload["dimensions"] = settings.embedding_dimensions
  38. async with httpx.AsyncClient(timeout=60) as client:
  39. response = await client.post(
  40. f"{settings.embedding_base_url}/embeddings",
  41. headers={"Authorization": f"Bearer {settings.embedding_api_key}"},
  42. json=payload,
  43. )
  44. response.raise_for_status()
  45. body = response.json()
  46. data = body.get("data")
  47. if not isinstance(data, list):
  48. raise EmbeddingUnavailable("Embedding API 返回中缺少 data 数组")
  49. ordered = sorted(data, key=lambda item: item.get("index", 0))
  50. vectors = [item.get("embedding") for item in ordered]
  51. if len(vectors) != len(inputs) or not all(isinstance(item, list) for item in vectors):
  52. raise EmbeddingUnavailable("Embedding API 返回的向量数量不匹配")
  53. if settings.embedding_dimensions and any(
  54. len(vector) != settings.embedding_dimensions for vector in vectors
  55. ):
  56. raise EmbeddingUnavailable(
  57. f"Embedding 向量维度与配置不一致,期望 {settings.embedding_dimensions}"
  58. )
  59. return [[float(value) for value in vector] for vector in vectors]
  60. embedding_client = OpenAICompatibleEmbeddingClient()