from __future__ import annotations import hashlib import math import re from typing import Protocol from app.config import Settings class EmbeddingProvider(Protocol): @property def dimension(self) -> int: ... def embed_query(self, text: str) -> list[float]: ... def embed_documents(self, texts: list[str]) -> list[list[float]]: ... class HashEmbeddingProvider: """无外部模型依赖的确定性向量,仅用于课程演示和自动化测试。""" def __init__(self, dimension: int = 256) -> None: self._dimension = dimension @property def dimension(self) -> int: return self._dimension @staticmethod def _features(text: str) -> list[str]: normalized = re.sub(r"\s+", "", text.lower()) chars = list(normalized) bigrams = [normalized[index : index + 2] for index in range(len(normalized) - 1)] words = re.findall(r"[a-z]+\d+|\d+|[a-z]+", normalized) return chars + bigrams + words def _embed(self, text: str) -> list[float]: vector = [0.0] * self._dimension for feature in self._features(text): # 稳定哈希保证相同文本在不同进程中生成相同演示向量。 digest = hashlib.blake2b(feature.encode("utf-8"), digest_size=8).digest() raw = int.from_bytes(digest, "big") index = raw % self._dimension sign = 1.0 if raw & 1 else -1.0 vector[index] += sign norm = math.sqrt(sum(value * value for value in vector)) if norm == 0: return vector # 单位化后可直接使用 COSINE 距离进行课程检索演示。 return [value / norm for value in vector] def embed_query(self, text: str) -> list[float]: return self._embed(text) def embed_documents(self, texts: list[str]) -> list[list[float]]: return [self._embed(text) for text in texts] class SentenceTransformerEmbeddingProvider: def __init__(self, model_name: str) -> None: try: from sentence_transformers import SentenceTransformer except ImportError as exc: raise RuntimeError( "缺少 sentence-transformers,请执行 uv sync --extra embeddings" ) from exc self._model = SentenceTransformer(model_name) self._dimension = self._model.get_sentence_embedding_dimension() @property def dimension(self) -> int: return int(self._dimension) def embed_query(self, text: str) -> list[float]: return self._model.encode(text, normalize_embeddings=True).tolist() def embed_documents(self, texts: list[str]) -> list[list[float]]: return self._model.encode(texts, normalize_embeddings=True).tolist() def create_embedding_provider(settings: Settings) -> EmbeddingProvider: # Provider 工厂隔离模型选择,Milvus Store 只依赖统一向量接口。 if settings.embedding_provider == "hash": return HashEmbeddingProvider(settings.hash_embedding_dim) if settings.embedding_provider == "sentence_transformers": return SentenceTransformerEmbeddingProvider(settings.embedding_model_name) raise ValueError(f"不支持的 EMBEDDING_PROVIDER: {settings.embedding_provider}")