ReMe/reme/components/as_embedding/__init__.py
jinliyl f44f52d919
Some checks failed
CI / Documentation / Test and build documentation (push) Waiting to run
CI / Python quality / Pre-commit (push) Waiting to run
CI / Python tests / Unit Tests - py3.11 (push) Waiting to run
CI / Python tests / Unit Tests - py3.12 (push) Waiting to run
CI / Python tests / Unit Tests - py3.13 (push) Waiting to run
CI / Windows / CLI smoke - py3.11 (push) Waiting to run
Deploy / Documentation / Build documentation (push) Waiting to run
Deploy / Documentation / deploy (push) Blocked by required conditions
Security / CodeQL / Analyze javascript-typescript (push) Waiting to run
Security / CodeQL / Analyze python (push) Waiting to run
CI / Website / Website checks (push) Has been cancelled
CI / Python packages / Build and verify distributions (push) Has been cancelled
fix(embedding): exclude provider init from health timeout (#484)
2026-08-21 11:23:14 +08:00

217 lines
7.4 KiB
Python

"""AgentScope embedding model wrappers."""
import hashlib
import os
from typing import Any
from agentscope.credential import (
CredentialBase,
DashScopeCredential,
GeminiCredential,
OllamaCredential,
OpenAICredential,
)
from agentscope.embedding import (
EmbeddingModelBase,
)
from ..base_component import BaseComponent
from ..component_registry import R
from ...enumeration import ComponentEnum
class BaseAsEmbedding(BaseComponent):
"""Base wrapper for AgentScope embedding models."""
component_type = ComponentEnum.AS_EMBEDDING
credential_cls: type[CredentialBase]
endpoint_fields = ("base_url", "host")
endpoint_env: str | None = None
default_endpoint = ""
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.model: EmbeddingModelBase[Any] | None = None
@property
def dimensions(self) -> int:
"""Return configured dimensions without forcing provider construction."""
if self.model is not None:
return self.model.dimensions
dimensions = self.kwargs.get("dimensions")
if dimensions is None:
raise RuntimeError("Embedding dimensions are required before provider initialization.")
return int(dimensions)
@property
def vector_space(self) -> tuple[str, ...]:
"""Return the fields that make vectors from two embedding setups incompatible.
The wrapper identifies the provider consistently before and after lazy model
construction. Model details come from the live provider when one has been
injected at runtime, otherwise they come from the configured kwargs.
"""
if self.model is not None:
return (
self.backend or self.credential_cls.__name__,
str(getattr(self.model, "model", self.kwargs.get("model") or "")),
str(self.dimensions),
self._model_endpoint(),
)
return (
self.backend or self.credential_cls.__name__,
str(self.kwargs.get("model") or ""),
str(self.dimensions),
self._configured_endpoint(),
)
@property
def vector_space_id(self) -> str:
"""Return a short digest of :attr:`vector_space` for naming persisted vectors.
Cache consumers use this digest to avoid reusing vectors produced by a
different embedding setup.
"""
return hashlib.sha256("\x1f".join(self.vector_space).encode()).hexdigest()[:12]
@staticmethod
def _endpoint(credential: Any) -> str:
"""Read the provider endpoint from a credential object or a raw kwargs dict."""
for field in ("base_url", "host"):
value = credential.get(field) if isinstance(credential, dict) else getattr(credential, field, None)
if value:
return str(value).rstrip("/")
return ""
def _configured_endpoint(self) -> str:
"""Resolve the provider endpoint without constructing it eagerly."""
credential = self.kwargs.get("credential")
if isinstance(credential, dict):
fields = getattr(self.credential_cls, "model_fields", {})
for name in self.endpoint_fields:
if name in credential:
value = credential[name]
if value is not None:
return str(value).rstrip("/")
continue
field = fields.get(name)
if field is not None and not field.is_required():
value = field.get_default(call_default_factory=True)
if value:
return str(value).rstrip("/")
else:
endpoint = self._endpoint(credential)
if endpoint:
return endpoint
if self.endpoint_env:
endpoint = os.environ.get(self.endpoint_env)
if endpoint:
return endpoint.rstrip("/")
return self.default_endpoint
def _model_endpoint(self) -> str:
"""Return the endpoint used by the constructed provider."""
assert self.model is not None
endpoint = self._endpoint(getattr(self.model, "credential", self.kwargs.get("credential")))
return endpoint or self._configured_endpoint()
async def __call__(self, inputs: list[Any], **kwargs) -> list[list[float]]:
self._ensure_model()
assert self.model is not None
response = await self.model(inputs, **kwargs) # pylint: disable=not-callable
return response.embeddings
def initialize_model(self) -> None:
"""Construct the provider without making a remote request.
Callers that apply their own request timeout can initialize first so
one-time SDK imports and client construction do not consume that
timeout budget. Normal embedding calls remain lazily initialized.
"""
self._ensure_model()
async def _start(self) -> None:
"""Defer provider construction until the first remote embedding call."""
return None
def _ensure_model(self) -> None:
"""Construct the provider on demand while keeping dimensions locally available."""
if self.model is not None:
return
kwargs = dict(self.kwargs)
credential = self.credential_cls(**kwargs.pop("credential", {}))
model_cls = self.credential_cls.get_embedding_model_class()
if model_cls is None:
raise ValueError(f"{self.credential_cls.__name__} does not support embeddings.")
dimensions = self.dimensions
kwargs.pop("dimensions", None)
params_dict = kwargs.pop("parameters", None)
parameters = model_cls.Parameters(**params_dict) if params_dict else None
self.model = model_cls(
credential=credential,
dimensions=dimensions,
parameters=parameters,
**kwargs,
)
@R.register("openai")
class OpenAIAsEmbedding(BaseAsEmbedding):
"""OpenAI embedding model wrapper."""
credential_cls = OpenAICredential
endpoint_fields = ("base_url",)
endpoint_env = "OPENAI_BASE_URL"
default_endpoint = "https://api.openai.com/v1"
def _model_endpoint(self) -> str:
"""Read the endpoint the OpenAI client actually resolved."""
assert self.model is not None
client = getattr(self.model, "client", None)
base_url = getattr(client, "base_url", None)
return str(base_url).rstrip("/") if base_url is not None else super()._model_endpoint()
@R.register("dashscope")
class DashScopeAsEmbedding(BaseAsEmbedding):
"""DashScope embedding model wrapper."""
credential_cls = DashScopeCredential
@R.register("dashscope_multimodal")
class DashScopeMultiModalAsEmbedding(BaseAsEmbedding):
"""DashScope multimodal embedding model wrapper."""
credential_cls = DashScopeCredential
@R.register("gemini")
class GeminiAsEmbedding(BaseAsEmbedding):
"""Gemini embedding model wrapper."""
credential_cls = GeminiCredential
@R.register("ollama")
class OllamaAsEmbedding(BaseAsEmbedding):
"""Ollama embedding model wrapper."""
credential_cls = OllamaCredential
endpoint_fields = ("host",)
endpoint_env = "OLLAMA_HOST"
default_endpoint = "http://127.0.0.1:11434"
__all__ = [
"BaseAsEmbedding",
"OpenAIAsEmbedding",
"DashScopeAsEmbedding",
"DashScopeMultiModalAsEmbedding",
"GeminiAsEmbedding",
"OllamaAsEmbedding",
]