"""Embedding provider 工厂:根据环境变量构造合适后端。""" from __future__ import annotations import os from .base import AsyncEmbeddingProvider, EmbeddingProvider from .models import EmbeddingError, EmbeddingProviderType from .remote import ( DashScopeAsyncEmbeddingProvider, DashScopeEmbeddingProvider, ) def _read_env(key: str, default: str | None = None) -> str | None: val = os.environ.get(key) if val is None or val.strip() == "": return default return val.strip() def resolve_provider_type(provider: str | None = None) -> EmbeddingProviderType: """根据 provider 参数 / env 解析出 EmbeddingProviderType。 映射: dashscope / qwen / remote -> DASHSCOPE local / local-bge / bge / bge-m3 -> LOCAL_BGE 默认 dashscope。 """ p = (provider or _read_env("EMBEDDING_PROVIDER", "dashscope") or "dashscope").lower() if p in ("dashscope", "qwen", "remote"): return EmbeddingProviderType.DASHSCOPE if p in ("local", "local-bge", "bge", "bge-m3"): return EmbeddingProviderType.LOCAL_BGE raise EmbeddingError(f"未知 embedding provider: {provider!r}") def make_sync_provider( provider: str | None = None, **kwargs: object, ) -> EmbeddingProvider: """构造同步 provider。""" pt = resolve_provider_type(provider) if pt == EmbeddingProviderType.DASHSCOPE: return DashScopeEmbeddingProvider(**kwargs) # type: ignore[arg-type] # 本地后端 from .local import LocalBGEEmbeddingProvider return LocalBGEEmbeddingProvider(**kwargs) # type: ignore[arg-type] def make_async_provider( provider: str | None = None, **kwargs: object, ) -> AsyncEmbeddingProvider: """构造异步 provider。""" pt = resolve_provider_type(provider) if pt == EmbeddingProviderType.DASHSCOPE: return DashScopeAsyncEmbeddingProvider(**kwargs) # type: ignore[arg-type] from .local import LocalBGEAsyncEmbeddingProvider return LocalBGEAsyncEmbeddingProvider(**kwargs) # type: ignore[arg-type]