61 lines
2.0 KiB
Python
61 lines
2.0 KiB
Python
"""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]
|