From 57e56881a7893d6f9f57503e301d7d12068ebb84 Mon Sep 17 00:00:00 2001 From: jeffwu Date: Fri, 14 Aug 2026 10:27:13 +0800 Subject: [PATCH 1/2] feat(gateway): add modality-agnostic core + bridge skeleton (B0) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Base branch for the incremental gateway split. Adds the gateway SDK core (self-contained, stdlib-only) and an empty modality/__init__ so importing nexent.core.gateway works but the adapter registry is empty. - sdk/nexent/core/gateway/: __init__ (imports modality for side-effect registration), registry, transport, multimodal_adapter, multimodal_gateway, model_context (all ModelContext subclasses) - modality/__init__.py: empty aggregator (no per-vendor imports yet; each feature branch appends its modality's lines) - backend/services/model_gateway_service.py: bridge skeleton — generic machinery only (_FACTORY_NORMALIZE, _normalize_factory, _coalesce, _config_to_context with common-kwargs+raise, get_adapter_from_config, build_adapter_fresh, _fetch_slot_config, _fetch_voice_config). No per-modality wrappers/branches; nothing on develop imports this file yet. - sdk/pyproject.toml: namespaces=true (modality subpackages are PEP 420) Verified: py_compile OK; registry empty (has(openai,vlm)=False) as expected. Co-Authored-By: Claude --- backend/services/model_gateway_service.py | 222 ++++++++++++++++++ sdk/nexent/core/gateway/__init__.py | 55 +++++ sdk/nexent/core/gateway/modality/__init__.py | 18 ++ sdk/nexent/core/gateway/model_context.py | 89 +++++++ sdk/nexent/core/gateway/multimodal_adapter.py | 107 +++++++++ sdk/nexent/core/gateway/multimodal_gateway.py | 106 +++++++++ sdk/nexent/core/gateway/registry.py | 108 +++++++++ sdk/nexent/core/gateway/transport.py | 115 +++++++++ sdk/pyproject.toml | 3 + 9 files changed, 823 insertions(+) create mode 100644 backend/services/model_gateway_service.py create mode 100644 sdk/nexent/core/gateway/__init__.py create mode 100644 sdk/nexent/core/gateway/modality/__init__.py create mode 100644 sdk/nexent/core/gateway/model_context.py create mode 100644 sdk/nexent/core/gateway/multimodal_adapter.py create mode 100644 sdk/nexent/core/gateway/multimodal_gateway.py create mode 100644 sdk/nexent/core/gateway/registry.py create mode 100644 sdk/nexent/core/gateway/transport.py diff --git a/backend/services/model_gateway_service.py b/backend/services/model_gateway_service.py new file mode 100644 index 0000000000..66cfe3934b --- /dev/null +++ b/backend/services/model_gateway_service.py @@ -0,0 +1,222 @@ +"""Backend bridge: DB model config → :class:`ModelContext` → :class:`MultimodalGateway`. + +This is the *thin* Phase 2 bridge. Existing service factory functions keep +their public signatures; they fetch the model config dict (unchanged) and +delegate construction to the gateway via :func:`get_adapter_from_config`:: + + cfg = tenant_config_manager.get_model_config(...) + model = get_adapter_from_config(cfg, "llm", "llm", tenant_id, + temperature=0.3, top_p=0.95) + +The vendor ``if model_factory == ...`` dispatch is replaced by registry +resolution keyed on the normalized factory, so adding a vendor becomes one +``@register_adapter`` decorator + one ``_FACTORY_NORMALIZE`` entry — the +service layer is untouched. + +This is the **base branch** skeleton: only the modality-agnostic machinery +lives here. ``_config_to_context`` keeps the common-kwargs construction but +no per-modality subclass branch yet, so any call raises — no consumer calls +it in this branch. Each feature branch (``feat/gw-vlm``, ``-llm``, …) appends +its ``if modality == ...: return SubClass(**common, ...)`` branch and the +modality-specific ``get_*_adapter`` wrappers that consume it. +""" + +from __future__ import annotations + +import logging +from typing import Any, Dict, Optional + +from nexent import MessageObserver +from nexent.core.gateway import ( + EmbeddingContext, + LLMContext, + LongContextLLMContext, + ModelContext, + STTContext, + TTSContext, + VLMContext, + get_gateway, +) +from nexent.core.gateway.registry import get_registry +from consts.const import MODEL_CONFIG_MAPPING, TEST_PCM_PATH +from database.model_management_db import get_model_by_model_id, get_model_records +from utils.config_utils import get_model_name_from_config, tenant_config_manager + +logger = logging.getLogger("model_gateway_service") + +# Normalize vendor aliases to canonical registry factory names. +_FACTORY_NORMALIZE: Dict[str, str] = { + "volc": "volc", + "volcano": "volc", + "volcengine": "volc", + "火山引擎": "volc", + "dashscope": "dashscope", + "ali": "ali", + "alibaba": "ali", + "阿里云": "ali", + "silicon": "siliconflow", + "siliconflow": "siliconflow", + "openai": "openai", + "tokenpony": "tokenpony", + "jina": "jina", + "cohere": "cohere", + "modelengine": "modelengine", +} + +# Modality-specific default factory when the raw factory is empty/unknown. +_MODALITY_DEFAULT_FACTORY: Dict[str, str] = { + "llm": "openai", + "llm_long_context": "openai", + "vlm": "openai", + "embedding": "openai", + "rerank": "openai", + "stt": "ali", + "tts": "ali", + "multi_embedding": "jina", +} + + +def _normalize_factory(raw: Optional[str], modality: str) -> str: + """Return the canonical registry factory for ``raw`` under ``modality``.""" + cleaned = (raw or "").strip().lower() + factory = _FACTORY_NORMALIZE.get(cleaned, cleaned) + # STT/TTS historically route DashScope through the Ali client. + if modality in ("stt", "tts") and factory in ("dashscope", "ali", "alibaba"): + factory = "ali" + if get_registry().has(factory, modality): + return factory + default = _MODALITY_DEFAULT_FACTORY.get(modality, "openai") + if factory: + logger.debug( + "factory %r has no %s adapter; falling back to %r", factory, modality, default + ) + return default + + +def _coalesce(*vals: Any) -> Any: + """Return the first non-``None`` value, or ``None`` if all are ``None``. + + Unlike ``a or b``, this preserves falsy-but-valid values such as + ``temperature=0`` or ``top_p=0`` — an explicit ``0`` must reach the + adapter rather than being silently replaced by the cfg/default fallback. + """ + for v in vals: + if v is not None: + return v + return None + + +def _config_to_context( + cfg: Optional[dict], + modality: str, + slot: str, + tenant_id: Optional[str], + **construct_extras: Any, +) -> ModelContext: + """Build a modality-specific :class:`ModelContext` from a DB config + per-call extras. + + ``construct_extras`` carries per-call-site tuning (temperature, top_p, + max_output_tokens, stream, observer, display_name, timeout_seconds, + language, speed_ratio, ...) so construction is behavior-preserving. Known + keys are mapped to subclass fields directly. + + Base branch: the common kwargs are constructed but no per-modality + subclass branch is present yet, so this raises. Feature branches append + ``if modality == "vlm": return VLMContext(**common, ...)`` etc. + """ + cfg = cfg or {} + factory = _normalize_factory(cfg.get("model_factory"), modality) + needs_observer = modality in ("vlm", "llm", "llm_long_context") + observer = construct_extras.pop("observer", None) + if needs_observer and observer is None: + observer = MessageObserver() + + # ---- common kwargs (base class fields) ---- + common: Dict[str, Any] = dict( + model_name=construct_extras.pop("model_name", None) or get_model_name_from_config(cfg) or "", + base_url=cfg.get("base_url", ""), + api_key=cfg.get("api_key", ""), + modality=modality, + factory=factory, + tenant_id=tenant_id, + slot=slot, + ssl_verify=cfg.get("ssl_verify", True), + observer=observer, + display_name=_coalesce(construct_extras.pop("display_name", None), cfg.get("display_name")), + timeout_seconds=_coalesce(construct_extras.pop("timeout_seconds", None), cfg.get("timeout_seconds")), + ) + + # ---- modality-specific subclass construction (added per feature branch) ---- + # e.g. `if modality == "vlm": return VLMContext(**common, ...)` in feat/gw-vlm. + raise ValueError(f"Unknown modality: {modality}") + + +def get_adapter_from_config( + cfg: Optional[dict], + modality: str, + slot: str, + tenant_id: Optional[str] = None, + **construct_extras: Any, +): + """Resolve and return the adapter for ``cfg`` (cached by the gateway).""" + context = _config_to_context(cfg, modality, slot, tenant_id, **construct_extras) + return get_gateway().get_adapter(context) + + +def build_adapter_fresh( + cfg: Optional[dict], + modality: str, + slot: str, + tenant_id: Optional[str] = None, + **construct_extras: Any, +): + """Build a fresh adapter for ``cfg`` WITHOUT the gateway instance cache. + + Used by per-call construction sites (e.g. voice streaming sessions) where + vendor config carries per-request params (api_key, ws_url, voice, …) that + must not collide across tenants under a shared cache key. + """ + context = _config_to_context(cfg, modality, slot, tenant_id, **construct_extras) + cls = get_registry().resolve(context.factory, modality) + return cls(context) + + +# ---- Generic config-fetch helpers (consumed by per-modality wrappers in feature branches) --- + +def _fetch_slot_config(tenant_id, model_id, expected_type, slot_key): + """Fetch a model config by model_id (with type check) or by slot key.""" + if model_id: + cfg = get_model_by_model_id(int(model_id), tenant_id) + if not cfg: + raise ValueError(f"Model not found: {model_id}") + if cfg.get("model_type") != expected_type: + raise ValueError( + f"Selected model {model_id} is not a {expected_type} model" + ) + return cfg + return tenant_config_manager.get_model_config( + key=MODEL_CONFIG_MAPPING.get(slot_key, slot_key), tenant_id=tenant_id + ) + + +def _fetch_voice_config(tenant_id, model_type): + """Fetch an STT/TTS config from tenant_config or model_records (voice fallback).""" + try: + cfg = tenant_config_manager.get_model_config(tenant_id, model_type) + if cfg and isinstance(cfg, dict): + return cfg + except Exception: + pass + try: + records = get_model_records({"model_type": model_type}, tenant_id) + if records: + return records[0] + except Exception: + pass + return None + + +# Per-modality convenience wrappers (get_llm_adapter / get_vlm_adapter / +# get_stt_adapter_* / get_tts_adapter_* / get_embedding_adapter_from_config / +# get_rerank_adapter_from_config) are added by their respective feature +# branches alongside the matching ``_config_to_context`` subclass branch. diff --git a/sdk/nexent/core/gateway/__init__.py b/sdk/nexent/core/gateway/__init__.py new file mode 100644 index 0000000000..4128b7fb6b --- /dev/null +++ b/sdk/nexent/core/gateway/__init__.py @@ -0,0 +1,55 @@ +"""Multimodal model unified adaptation gateway. + +Provides a protocol-agnostic adapter layer that *composes* the existing, +stable model classes (``OpenAIModel`` / embedding / rerank …) behind a single +:class:`MultimodalGateway` entry point. For STT/TTS/VLM the protocol lives +directly in the adapter (no wrapped model class); LLM/LongContext stay as thin +``has-a`` delegation to ``OpenAIModel``. Adding a vendor becomes one +``@register_adapter(factory, modality)`` decorator — backend services no +longer hardcode ``if model_factory == ...`` dispatch. + +Importing :mod:`nexent.core.gateway` (or its :mod:`.modality` subpackage) +registers all built-in adapters with the process-wide registry. + +See ``doc/multimodal-gateway-design.md`` for the full design. +""" + +from .multimodal_adapter import ModelInfo, MultimodalAdapter +from .model_context import ( + EmbeddingContext, + LLMContext, + LongContextLLMContext, + ModelContext, + STTContext, + TTSContext, + VLMContext, +) +from .multimodal_gateway import MultimodalGateway, get_gateway +from .registry import AdapterRegistry, get_registry, register_adapter +from .transport import ( + HttpTransportMixin, + WebSocketTransportMixin, +) + +# Importing .modality registers all built-in adapters via @register_adapter. +from . import modality # noqa: F401 (side-effect: registration) + +__all__ = [ + "ModelInfo", + "MultimodalAdapter", + "ModelContext", + "LLMContext", + "LongContextLLMContext", + "VLMContext", + "EmbeddingContext", + "STTContext", + "TTSContext", + "MultimodalGateway", + "get_gateway", + "AdapterRegistry", + "get_registry", + "register_adapter", + "HttpTransportMixin", + "WebSocketTransportMixin", + "modality", +] diff --git a/sdk/nexent/core/gateway/modality/__init__.py b/sdk/nexent/core/gateway/modality/__init__.py new file mode 100644 index 0000000000..4606379cf7 --- /dev/null +++ b/sdk/nexent/core/gateway/modality/__init__.py @@ -0,0 +1,18 @@ +"""Modality adapter aggregation layer. + +This is the single place that re-exports the public adapter API and triggers +built-in adapter registration: importing :mod:`nexent.core.gateway.modality` +imports every built-in adapter module, whose ``@register_adapter`` decorators +populate the process-wide :class:`AdapterRegistry`. + +Modality subpackages (``llm`` / ``vlm`` / ``stt`` / ``tts`` / ``embedding`` / +``rerank``) are namespace packages on purpose: they ship no ``__init__.py`` so +this module stays the single aggregation point. Import concrete classes via +this layer or via their leaf module (``modality.vlm.openai``). + +This is the **base branch**: no adapters are registered yet, so the registry is +empty. Each feature branch (``feat/gw-vlm``, ``-llm``, …) appends its modality's +import lines here, which registers that modality's adapters. +""" + +__all__: list[str] = [] diff --git a/sdk/nexent/core/gateway/model_context.py b/sdk/nexent/core/gateway/model_context.py new file mode 100644 index 0000000000..b6f7d77f3d --- /dev/null +++ b/sdk/nexent/core/gateway/model_context.py @@ -0,0 +1,89 @@ +"""Modality-specific construction contexts for the gateway.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Dict, Optional + + +@dataclass +class ModelContext: + """Base construction context — common + cross-cutting fields. + + Subclasses add modality-specific fields. Passing a field the subclass + doesn't declare raises TypeError at construction time. + """ + + model_name: str + base_url: str + api_key: str + modality: str + factory: str + tenant_id: Optional[str] = None + slot: Optional[str] = None # VLM slot key (vlm/vlm3/...) + ssl_verify: bool = True + display_name: Optional[str] = None + observer: Any = None # cross-cutting: LLM/VLM/ModelEngine STT/TTS + timeout_seconds: Optional[float] = None # cross-cutting: all HTTP-backed adapters + + def cache_key(self) -> tuple: + return (self.tenant_id or "", self.modality, self.slot or "", + self.model_name, self.factory) + + +@dataclass +class LLMContext(ModelContext): + temperature: Optional[float] = None + top_p: Optional[float] = None + stream: Optional[bool] = None + max_output_tokens: Optional[int] = None + frequency_penalty: Optional[float] = None + extra_body: Optional[dict] = None # OpenAI API passthrough + + +@dataclass +class LongContextLLMContext(LLMContext): + max_tokens: Optional[int] = None # context window size + truncation_strategy: Optional[str] = None # "start" | "end" | ... + + +@dataclass +class VLMContext(LLMContext): + capabilities: Dict[str, bool] = field(default_factory=dict) # {"audio": False} + max_tokens: Optional[int] = None # max output tokens for image analysis + + +@dataclass +class EmbeddingContext(ModelContext): + embedding_dim: Optional[int] = None + model_type: Optional[str] = None # "embedding" | "multi_embedding" + + +@dataclass +class STTContext(ModelContext): + language: str = "zh" + audio_file_path: Optional[str] = None + model_appid: Optional[str] = None # Volc + access_token: Optional[str] = None # Volc + ws_url: Optional[str] = None # WS variants + auth_headers: Optional[dict] = None + format: str = "pcm" # Ali/Volc + rate: int = 16000 # Ali/Volc + resourceid: Optional[str] = None # Volc + enable_vad: bool = True # Ali default + sample_rate: Optional[int] = None # Ali/Volc + timeout: Optional[int] = None # Ali: per-operation WS timeout + + +@dataclass +class TTSContext(ModelContext): + speed_ratio: float = 1.0 + voice: Optional[str] = None + audio_file_path: Optional[str] = None + model_appid: Optional[str] = None # Volc + access_token: Optional[str] = None # Volc + ws_url: Optional[str] = None # WS variants + auth_headers: Optional[dict] = None + voice_type: Optional[str] = None # Volc + format: str = "mp3" # Ali + sample_rate: int = 16000 # Ali diff --git a/sdk/nexent/core/gateway/multimodal_adapter.py b/sdk/nexent/core/gateway/multimodal_adapter.py new file mode 100644 index 0000000000..390bc99fc1 --- /dev/null +++ b/sdk/nexent/core/gateway/multimodal_adapter.py @@ -0,0 +1,107 @@ +"""Multimodal adapter root ABC and model capability declaration.""" + +from abc import ABC, abstractmethod +from dataclasses import dataclass +from typing import Any, AsyncIterator, Dict + +from .model_context import ModelContext + + +@dataclass +class ModelInfo: + """Model capability declaration, replacing hardcoded URL sniffing. + + Attributes: + model_id: The model identifier passed to the provider API. + display_name: Human-readable name shown in the UI. + provider: The normalized factory name (e.g. ``"openai"``). + capabilities: Per-capability flags, e.g. ``{"image": True, + "audio": False, "video": True}``. + """ + + model_id: str + display_name: str + provider: str + capabilities: Dict[str, bool] + + +class MultimodalAdapter(ABC): + """Root interface for every modality adapter. + + Subclasses set the class-level ``modality`` (``"llm"`` | ``"vlm"`` | + ``"stt"`` | ``"tts"`` | ``"embedding"`` | ``"rerank"`` | ...) and + ``factory`` (``"openai"`` | ``"ali"`` | ``"volc"`` | ...), then implement + :meth:`invoke`, :meth:`health_check`, :meth:`get_model_info`. + + The root carries no wrapped-model state. Callers reach the model only + through the uniform interface above — never by tunnelling into a wrapped + instance's attributes. + """ + + modality: str + factory: str + + def __init__(self, context: ModelContext) -> None: + """Stores the construction context for the adapter. + + Args: + context: The unified construction context for this adapter. + """ + self._context = context + + @abstractmethod + async def invoke(self, request: Any) -> Any: + """Batch/synchronous entry point. + + LLM/VLM return a ChatMessage, Embedding returns a list of vectors, + Rerank returns a list of dicts, TTS returns audio bytes, and STT + returns a transcription dict. + + Args: + request: The modality-specific request payload. + + Returns: + The modality-specific response. + """ + raise NotImplementedError + + async def stream(self, request: Any) -> AsyncIterator[Any]: + """Streaming entry point. + + STT/TTS override this as an async generator. + + Args: + request: The modality-specific request payload. + + Yields: + Modality-specific stream chunks. + + Raises: + NotImplementedError: If the adapter does not support streaming. + """ + raise NotImplementedError( + f"{self.modality} adapter does not support streaming" + ) + + @abstractmethod + async def health_check(self) -> bool: + """Unified health check. + + Replaces the three legacy method names (``check_connectivity`` / + ``dimension_check`` / ``connectivity_check``). + + Returns: + True if the model is reachable, False otherwise. + """ + raise NotImplementedError + + @abstractmethod + def get_model_info(self) -> ModelInfo: + """Returns the capability declaration. + + Replaces analyze_audio_tool's ``getattr`` + URL sniffing. + + Returns: + The model's capability declaration. + """ + raise NotImplementedError diff --git a/sdk/nexent/core/gateway/multimodal_gateway.py b/sdk/nexent/core/gateway/multimodal_gateway.py new file mode 100644 index 0000000000..1a9462c356 --- /dev/null +++ b/sdk/nexent/core/gateway/multimodal_gateway.py @@ -0,0 +1,106 @@ +"""MultimodalGateway: the unified entry point replacing hardcoded dispatch. + +Backend services (``image_service``/``voice_service``/``vectordatabase_service``) +resolve an adapter through :meth:`get_adapter` instead of ``if model_factory`` +branches. The gateway caches adapter instances per ``(tenant, modality, slot, +model_name, factory)`` so a given model is constructed once. +""" + +from __future__ import annotations + +from typing import Any, Dict, Tuple + +from .multimodal_adapter import MultimodalAdapter +from .model_context import ModelContext +from .registry import AdapterRegistry, get_registry + + +class MultimodalGateway: + """Resolve and cache :class:`MultimodalAdapter` instances by context.""" + + def __init__(self, registry: AdapterRegistry = None) -> None: + """Initializes the gateway with a registry and empty cache. + + Args: + registry: The adapter registry to resolve from. Defaults to the + process-wide singleton. + """ + self._registry = registry or get_registry() + self._adapter_cache: Dict[Tuple, MultimodalAdapter] = {} + + def get_adapter(self, context: ModelContext) -> MultimodalAdapter: + """Returns the adapter for ``context``, building and caching it once. + + Args: + context: The construction context identifying the desired model. + + Returns: + The cached or newly built adapter instance. + """ + cls = self._registry.resolve(context.factory, context.modality) + key = context.cache_key() + if key not in self._adapter_cache: + self._adapter_cache[key] = cls(context) + return self._adapter_cache[key] + + async def invoke(self, context: ModelContext, request: Any) -> Any: + """Resolves the adapter for ``context`` and invokes it. + + Args: + context: The construction context identifying the desired model. + request: The modality-specific request payload. + + Returns: + The modality-specific response. + """ + return await self.get_adapter(context).invoke(request) + + def stream(self, context: ModelContext, request: Any): + """Returns the adapter's async iterator (not awaited — it's a generator). + + Args: + context: The construction context identifying the desired model. + request: The modality-specific request payload. + + Returns: + The adapter's async stream object. + """ + return self.get_adapter(context).stream(request) + + async def health_check(self, context: ModelContext) -> bool: + """Resolves the adapter for ``context`` and checks its health. + + Args: + context: The construction context identifying the desired model. + + Returns: + True if the model is reachable, False otherwise. + """ + return await self.get_adapter(context).health_check() + + def invalidate(self, context: ModelContext = None) -> None: + """Drops cached adapter instances. + + Args: + context: If provided, drops only that context's cached adapter. + If None, drops the entire cache. + """ + if context is None: + self._adapter_cache.clear() + else: + self._adapter_cache.pop(context.cache_key(), None) + + +_gateway: MultimodalGateway = None + + +def get_gateway() -> MultimodalGateway: + """Returns the process-wide gateway singleton (lazy). + + Returns: + The shared :class:`MultimodalGateway` instance. + """ + global _gateway + if _gateway is None: + _gateway = MultimodalGateway() + return _gateway diff --git a/sdk/nexent/core/gateway/registry.py b/sdk/nexent/core/gateway/registry.py new file mode 100644 index 0000000000..bfd7c51a8e --- /dev/null +++ b/sdk/nexent/core/gateway/registry.py @@ -0,0 +1,108 @@ +"""Adapter registry: maps ``(factory, modality)`` → adapter class. + +Paradigm aligned with :mod:`sdk.nexent.memory.providers.registry`. Vendors opt +in via the ``@register_adapter(factory, modality)`` decorator on their adapter +class; the backend never hardcodes vendor dispatch (``if model_factory == ...``). +""" + +from __future__ import annotations + +import logging +from typing import Dict, Tuple, Type + +from .multimodal_adapter import MultimodalAdapter + +logger = logging.getLogger("adapter_registry") + + +class AdapterRegistry: + """Registry of adapter classes keyed by ``(factory, modality)``.""" + + def __init__(self) -> None: + """Initializes an empty registry.""" + self._adapter_map: Dict[Tuple[str, str], Type[MultimodalAdapter]] = {} + + def register(self, factory: str, modality: str): + """Class decorator: register an adapter under ``(factory, modality)``. + + Args: + factory: The normalized provider name. + modality: The capability family identifier. + + Returns: + The class decorator that registers and returns the class. + """ + + def deco(cls: Type[MultimodalAdapter]) -> Type[MultimodalAdapter]: + key = (factory.lower().strip(), modality) + self._adapter_map[key] = cls + logger.debug("Registered adapter %s for %s", key, cls.__name__) + return cls + + return deco + + def resolve(self, factory: str, modality: str) -> Type[MultimodalAdapter]: + """Returns the adapter class for ``(factory, modality)``. + + Args: + factory: The normalized provider name. + modality: The capability family identifier. + + Returns: + The registered adapter class. + + Raises: + KeyError: If no adapter is registered for the pair. + """ + key = (factory.lower().strip(), modality) + if key not in self._adapter_map: + raise KeyError( + f"No adapter registered for factory={factory!r} modality={modality!r}; " + f"registered: {self.list_adapters()}" + ) + return self._adapter_map[key] + + def has(self, factory: str, modality: str) -> bool: + """Returns whether a ``(factory, modality)`` pair is registered. + + Args: + factory: The normalized provider name. + modality: The capability family identifier. + + Returns: + True if a pair is registered, False otherwise. + """ + return (factory.lower().strip(), modality) in self._adapter_map + + def list_adapters(self) -> list: + """Returns all registered ``(factory, modality)`` pairs. + + Returns: + A list of registered key tuples. + """ + return list(self._adapter_map.keys()) + + +_registry = AdapterRegistry() + + +def get_registry() -> AdapterRegistry: + """Returns the process-wide adapter registry singleton. + + Returns: + The shared :class:`AdapterRegistry` instance. + """ + return _registry + + +def register_adapter(factory: str, modality: str): + """Module-level convenience alias for ``AdapterRegistry.register``. + + Args: + factory: The normalized provider name. + modality: The capability family identifier. + + Returns: + The class decorator that registers the adapter. + """ + return _registry.register(factory, modality) diff --git a/sdk/nexent/core/gateway/transport.py b/sdk/nexent/core/gateway/transport.py new file mode 100644 index 0000000000..426a6b87ba --- /dev/null +++ b/sdk/nexent/core/gateway/transport.py @@ -0,0 +1,115 @@ +"""Transport-layer abstraction, orthogonal to modality logic. + +Inspired by Pipecat's ``WebsocketService`` mixin design: transport concerns +(HTTP base_url/api_key vs WebSocket ws_url/auth_headers) live in mixins that +are multiply-inherited alongside a modality ABC, decoupling the STT/TTS +adapters from a hardcoded WebSocket assumption. An HTTP-only vendor (e.g. +ModelEngine STT/TTS) can therefore take :class:`HttpTransportMixin` instead of +``WebSocketTransportMixin``. +""" + +from typing import Optional + + +class HttpTransportMixin: + """HTTP REST transport capability. + + Adapter classes multiply-inherit this alongside their modality ABC to gain + HTTP transport attributes without polluting the modality interface. + + Attributes: + transport_type: Always ``"http"``. + """ + + transport_type = "http" + + def __init__( + self, + *, + base_url: str, + api_key: str, + ssl_verify: bool = True, + timeout: float = 30.0, + ) -> None: + """Stores HTTP transport state for later per-call use. + + Args: + base_url: The HTTP endpoint URL. + api_key: The bearer token used for authorization. + ssl_verify: Whether to verify TLS certificates. + timeout: Default request timeout in seconds. + """ + self._base_url = base_url + self._api_key = api_key + self._ssl_verify = ssl_verify + self._timeout = timeout + + async def connect(self) -> None: + """No-op: HTTP clients are created lazily per-call.""" + return None + + async def close(self) -> None: + """No-op: there is no persistent HTTP session to close.""" + return None + + async def health_check(self) -> bool: + """Returns True; connectivity is the adapter's responsibility. + + Returns: + Always True (the mixin only carries state). + """ + # Delegated to the adapter's wrapped model; the mixin only carries state. + return True + + +class WebSocketTransportMixin: + """WebSocket transport capability. + + WS-specific parameters (ws_url, auth_headers) are managed here rather than + on the modality ABC, so STT/TTS adapters no longer need to hardcode + WebSocket assumptions. The session is created lazily. + + Attributes: + transport_type: Always ``"websocket"``. + """ + + transport_type = "websocket" + + def __init__( + self, + *, + ws_url: Optional[str] = None, + auth_headers: Optional[dict] = None, + ) -> None: + """Stores WebSocket transport state for later lazy connection. + + Args: + ws_url: The WebSocket endpoint URL. + auth_headers: Optional authentication headers sent on connect. + """ + self._ws_url = ws_url + self._auth_headers = auth_headers or {} + self._ws_connection = None # websockets.ClientConnection, created lazily + + async def connect(self) -> None: + """No-op: the wrapped model owns its WS lifecycle.""" + return None + + async def close(self) -> None: + """Closes the lazily-created WebSocket connection, if any. + + Idempotent: clears ``_ws_connection`` even on close failure. + """ + if self._ws_connection is not None: + try: + await self._ws_connection.close() + finally: + self._ws_connection = None + + async def health_check(self) -> bool: + """Returns whether a WebSocket URL is configured. + + Returns: + True if ``ws_url`` is set, False otherwise. + """ + return self._ws_url is not None diff --git a/sdk/pyproject.toml b/sdk/pyproject.toml index 755ef54efd..3a6689c749 100644 --- a/sdk/pyproject.toml +++ b/sdk/pyproject.toml @@ -96,6 +96,9 @@ dev = [ [tool.setuptools.packages.find] include = ["nexent*"] exclude = ["tests*", "examples*"] +# Modality subpackages (llm/vlm/stt/tts/embedding/rerank) intentionally have no +# __init__.py (namespace packages): modality/__init__.py is the single aggregator. +namespaces = true [tool.setuptools.package-data] "nexent.core.prompts" = ["*.yaml"] From 021e935ff32e49e3101583eb2bbdb06a116db6b4 Mon Sep 17 00:00:00 2001 From: jeffwu Date: Fri, 14 Aug 2026 10:38:54 +0800 Subject: [PATCH 2/2] feat(gateway): wire VLM image/video/audio understanding via adapter (B1) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit First feature branch on top of the B0 core. Switches VLM model access from image_service.get_vlm_model / get_video_understanding_model to the gateway adapter (get_vlm_adapter / OpenAIVLMAdapter). LLM access is untouched here (still get_llm_model — lands in B2). - modality/vlm/{vlm_adapter,openai,modelengine}.py + 3 vlm import lines in modality/__init__ (registers openai/modelengine vlm adapters) - bridge: _config_to_context vlm branch + get_vlm_adapter / get_vlm_adapter_from_config wrappers - create_agent_info: AnalyzeImageTool (slot="vlm") / AnalyzeAudio|VideoTool (slot="vlm3") injection switched to get_vlm_adapter - tool_configuration_service: analyze_image/analyze_audio|video validation switched to get_vlm_adapter - model_health_service: vlm connectivity check via build_adapter_fresh.health_check - SDK tools analyze_image/video/audio (VLMRequest + invoke_sync) + test_vlm_adapter + test_analyze_image_tool - tests: switched only the VLM patch sites in test_create_agent_info / test_tool_configuration to get_vlm_adapter (slot="vlm"/"vlm3") + added the model_gateway_service stub; LLM patches left as get_llm_model (B2) Verified: py_compile OK; working-tree vlm slice == ac91318af final (diff vs final shows only llm differences). E2E image-understanding path matches the previously-verified source branch (analyze_image -> get_vlm_adapter -> OpenAIVLMAdapter -> SiliconFlow). Co-Authored-By: Claude --- backend/agents/create_agent_info.py | 11 +- backend/services/model_gateway_service.py | 113 ++-- backend/services/model_health_service.py | 16 +- .../services/tool_configuration_service.py | 7 +- sdk/nexent/core/gateway/__init__.py | 14 +- sdk/nexent/core/gateway/modality/__init__.py | 24 +- .../core/gateway/modality/vlm/modelengine.py | 17 + .../core/gateway/modality/vlm/openai.py | 320 +++++++++++ .../core/gateway/modality/vlm/vlm_adapter.py | 49 ++ sdk/nexent/core/gateway/model_context.py | 2 +- sdk/nexent/core/gateway/multimodal_adapter.py | 2 +- sdk/nexent/core/gateway/multimodal_gateway.py | 10 +- sdk/nexent/core/gateway/registry.py | 7 +- sdk/nexent/core/gateway/transport.py | 10 +- sdk/nexent/core/tools/analyze_audio_tool.py | 31 +- sdk/nexent/core/tools/analyze_image_tool.py | 15 +- sdk/nexent/core/tools/analyze_video_tool.py | 17 +- test/backend/agents/test_create_agent_info.py | 21 +- .../services/test_model_gateway_service.py | 291 ++++++++++ .../services/test_model_health_service.py | 38 +- .../test_tool_configuration_service.py | 27 +- test/sdk/core/gateway/__init__.py | 0 test/sdk/core/gateway/test_model_context.py | 44 ++ .../core/gateway/test_multimodal_adapter.py | 82 +++ .../core/gateway/test_multimodal_gateway.py | 143 +++++ test/sdk/core/gateway/test_registry.py | 38 ++ test/sdk/core/gateway/test_transport.py | 97 ++++ test/sdk/core/models/test_vlm_adapter.py | 501 ++++++++++++++++++ .../tools/test_analyze_audio_video_tool.py | 33 +- .../sdk/core/tools/test_analyze_image_tool.py | 24 +- 30 files changed, 1801 insertions(+), 203 deletions(-) create mode 100644 sdk/nexent/core/gateway/modality/vlm/modelengine.py create mode 100644 sdk/nexent/core/gateway/modality/vlm/openai.py create mode 100644 sdk/nexent/core/gateway/modality/vlm/vlm_adapter.py create mode 100644 test/backend/services/test_model_gateway_service.py create mode 100644 test/sdk/core/gateway/__init__.py create mode 100644 test/sdk/core/gateway/test_model_context.py create mode 100644 test/sdk/core/gateway/test_multimodal_adapter.py create mode 100644 test/sdk/core/gateway/test_multimodal_gateway.py create mode 100644 test/sdk/core/gateway/test_registry.py create mode 100644 test/sdk/core/gateway/test_transport.py create mode 100644 test/sdk/core/models/test_vlm_adapter.py diff --git a/backend/agents/create_agent_info.py b/backend/agents/create_agent_info.py index e97c3c2afe..2575b5c7ba 100644 --- a/backend/agents/create_agent_info.py +++ b/backend/agents/create_agent_info.py @@ -43,7 +43,7 @@ from database.a2a_agent_db import PROTOCOL_JSONRPC from services.memory_config_service import build_memory_context -from services.image_service import get_video_understanding_model, get_vlm_model +from services.model_gateway_service import get_vlm_adapter from database.agent_db import ( search_agent_info_by_agent_id, query_sub_agent_relations, @@ -254,7 +254,7 @@ def _resolve_safe_input_budget( except UncertaintyReserveBasisUnknown as exc: # W2 uncertainty reserve needs context_window_tokens as the 10% basis. # Falls through here when a model row has max_input_tokens set but - # context_window_tokens is NULL — possible for rows imported before + # context_window_tokens is NULL - possible for rows imported before # W11 V1 save-time defaults landed, or for rows written directly via # SQL/legacy import. Degrade to the same "no W2 snapshot" branch the # caller already handles (falls back to W1 input_budget). @@ -289,7 +289,7 @@ def _resolve_input_budget( Calls ModelCapacityResolver with the catalog + operator overrides. Returns snapshot.provider_input_limit_tokens and monitoring fields on success. Falls back to _TOKEN_THRESHOLD_LEGACY_FALLBACK with no snapshot when - capacity is unknown — this is the migration-window behavior before all + capacity is unknown - this is the migration-window behavior before all model rows are backfilled. """ if not isinstance(model_info, dict): @@ -1600,15 +1600,14 @@ async def create_tool_config_list( elif tool_config.class_name == "AnalyzeImageTool": selected_model_id = param_dict.get("selected_model_id") tool_config.metadata = { - # get_vlm_model reads the first multimodal slot, now shown as image understanding. - "vlm_model": get_vlm_model(tenant_id=tenant_id, model_id=selected_model_id), + "vlm_model": get_vlm_adapter(tenant_id, selected_model_id, slot="vlm"), "storage_client": minio_client, "validate_url_access": lambda urls: validate_urls_access(urls, user_id) } elif tool_config.class_name in ["AnalyzeAudioTool", "AnalyzeVideoTool"]: selected_model_id = param_dict.get("selected_model_id") tool_config.metadata = { - "vlm_model": get_video_understanding_model(tenant_id=tenant_id, model_id=selected_model_id), + "vlm_model": get_vlm_adapter(tenant_id, selected_model_id, slot="vlm3"), "storage_client": minio_client, "validate_url_access": lambda urls: validate_urls_access(urls, user_id) } diff --git a/backend/services/model_gateway_service.py b/backend/services/model_gateway_service.py index 66cfe3934b..0c60443827 100644 --- a/backend/services/model_gateway_service.py +++ b/backend/services/model_gateway_service.py @@ -1,24 +1,7 @@ -"""Backend bridge: DB model config → :class:`ModelContext` → :class:`MultimodalGateway`. - -This is the *thin* Phase 2 bridge. Existing service factory functions keep -their public signatures; they fetch the model config dict (unchanged) and -delegate construction to the gateway via :func:`get_adapter_from_config`:: - - cfg = tenant_config_manager.get_model_config(...) - model = get_adapter_from_config(cfg, "llm", "llm", tenant_id, - temperature=0.3, top_p=0.95) - -The vendor ``if model_factory == ...`` dispatch is replaced by registry -resolution keyed on the normalized factory, so adding a vendor becomes one -``@register_adapter`` decorator + one ``_FACTORY_NORMALIZE`` entry — the -service layer is untouched. - -This is the **base branch** skeleton: only the modality-agnostic machinery -lives here. ``_config_to_context`` keeps the common-kwargs construction but -no per-modality subclass branch yet, so any call raises — no consumer calls -it in this branch. Each feature branch (``feat/gw-vlm``, ``-llm``, …) appends -its ``if modality == ...: return SubClass(**common, ...)`` branch and the -modality-specific ``get_*_adapter`` wrappers that consume it. +"""Backend bridge: turn DB model configs into gateway adapters. + +Service factory functions keep their signatures but delegate adapter +construction to the gateway via :func:`get_adapter_from_config`. """ from __future__ import annotations @@ -97,7 +80,7 @@ def _coalesce(*vals: Any) -> Any: """Return the first non-``None`` value, or ``None`` if all are ``None``. Unlike ``a or b``, this preserves falsy-but-valid values such as - ``temperature=0`` or ``top_p=0`` — an explicit ``0`` must reach the + ``temperature=0`` or ``top_p=0`` - an explicit ``0`` must reach the adapter rather than being silently replaced by the cfg/default fallback. """ for v in vals: @@ -115,15 +98,17 @@ def _config_to_context( ) -> ModelContext: """Build a modality-specific :class:`ModelContext` from a DB config + per-call extras. - ``construct_extras`` carries per-call-site tuning (temperature, top_p, - max_output_tokens, stream, observer, display_name, timeout_seconds, - language, speed_ratio, ...) so construction is behavior-preserving. Known - keys are mapped to subclass fields directly. +Args: + cfg: Raw model config dict from the DB (may be empty). + modality: The capability family identifier. + slot: The config slot key (e.g. "vlm" / "vlm3"). + tenant_id: The tenant identifier. + **construct_extras: Per-call-site tuning (temperature, top_p, stream, ...); + known keys map to subclass fields directly. - Base branch: the common kwargs are constructed but no per-modality - subclass branch is present yet, so this raises. Feature branches append - ``if modality == "vlm": return VLMContext(**common, ...)`` etc. - """ +Returns: + The modality-specific :class:`ModelContext` instance. +""" cfg = cfg or {} factory = _normalize_factory(cfg.get("model_factory"), modality) needs_observer = modality in ("vlm", "llm", "llm_long_context") @@ -131,7 +116,6 @@ def _config_to_context( if needs_observer and observer is None: observer = MessageObserver() - # ---- common kwargs (base class fields) ---- common: Dict[str, Any] = dict( model_name=construct_extras.pop("model_name", None) or get_model_name_from_config(cfg) or "", base_url=cfg.get("base_url", ""), @@ -146,8 +130,19 @@ def _config_to_context( timeout_seconds=_coalesce(construct_extras.pop("timeout_seconds", None), cfg.get("timeout_seconds")), ) - # ---- modality-specific subclass construction (added per feature branch) ---- - # e.g. `if modality == "vlm": return VLMContext(**common, ...)` in feat/gw-vlm. + if modality == "vlm": + caps = construct_extras.pop("capabilities", None) or {} + return VLMContext( + **common, + temperature=_coalesce(construct_extras.pop("temperature", None), cfg.get("temperature")), + top_p=_coalesce(construct_extras.pop("top_p", None), cfg.get("top_p")), + stream=construct_extras.pop("stream", None), + max_output_tokens=_coalesce(construct_extras.pop("max_output_tokens", None), cfg.get("max_output_tokens")), + frequency_penalty=cfg.get("frequency_penalty"), + extra_body=cfg.get("extra_body"), + max_tokens=cfg.get("max_tokens"), + capabilities=caps, + ) raise ValueError(f"Unknown modality: {modality}") @@ -172,17 +167,25 @@ def build_adapter_fresh( ): """Build a fresh adapter for ``cfg`` WITHOUT the gateway instance cache. - Used by per-call construction sites (e.g. voice streaming sessions) where - vendor config carries per-request params (api_key, ws_url, voice, …) that - must not collide across tenants under a shared cache key. - """ +Used by per-call construction sites (e.g. voice streaming sessions) where +vendor config carries per-request params (api_key, ws_url, voice, ...) that +must not collide across tenants under a shared cache key. + +Args: + cfg: Raw model config dict. + modality: The capability family identifier. + slot: The config slot key. + tenant_id: The tenant identifier. + **construct_extras: Extra fields forwarded to the context. + +Returns: + A newly constructed adapter instance. +""" context = _config_to_context(cfg, modality, slot, tenant_id, **construct_extras) cls = get_registry().resolve(context.factory, modality) return cls(context) -# ---- Generic config-fetch helpers (consumed by per-modality wrappers in feature branches) --- - def _fetch_slot_config(tenant_id, model_id, expected_type, slot_key): """Fetch a model config by model_id (with type check) or by slot key.""" if model_id: @@ -216,7 +219,31 @@ def _fetch_voice_config(tenant_id, model_type): return None -# Per-modality convenience wrappers (get_llm_adapter / get_vlm_adapter / -# get_stt_adapter_* / get_tts_adapter_* / get_embedding_adapter_from_config / -# get_rerank_adapter_from_config) are added by their respective feature -# branches alongside the matching ``_config_to_context`` subclass branch. +def get_vlm_adapter_from_config( + cfg: Optional[dict], + tenant_id: Optional[str] = None, + slot: str = "vlm", + **construct_extras: Any, +): + return get_adapter_from_config(cfg, "vlm", slot, tenant_id, **construct_extras) + + +def get_vlm_adapter(tenant_id: str, model_id: Optional[int] = None, slot: str = "vlm"): + """Resolve the VLM adapter directly (bridge owns config-fetch). + +Replaces ``image_service.get_vlm_model`` / ``get_video_understanding_model``. + +Args: + tenant_id: The tenant identifier. + model_id: Optional model id; defaults to the slot config when omitted. + slot: "vlm" (image) or "vlm3" (video/audio). + +Returns: + A VLM adapter instance, or ``None`` when no config is available. +""" + cfg = _fetch_slot_config(tenant_id, model_id, expected_type=slot, slot_key=slot) + if not cfg: + return None + return get_gateway().get_adapter(_config_to_context(cfg, "vlm", slot, tenant_id)) + + diff --git a/backend/services/model_health_service.py b/backend/services/model_health_service.py index 2d0c43d09f..a5b2c34aa3 100644 --- a/backend/services/model_health_service.py +++ b/backend/services/model_health_service.py @@ -2,10 +2,11 @@ from typing import Optional from nexent.core import MessageObserver -from nexent.core.models import OpenAIModel, OpenAIVLModel +from nexent.core.models import OpenAIModel from nexent.core.models.embedding_model import JinaEmbedding, OpenAICompatibleEmbedding, DashScopeMultimodalEmbedding, SiliconflowMultimodalEmbedding from nexent.monitor import set_monitoring_context, set_monitoring_operation from nexent.core.models.rerank_model import OpenAICompatibleRerank +from services.model_gateway_service import build_adapter_fresh from services.voice_service import get_voice_service from consts.const import LOCALHOST_IP, LOCALHOST_NAME, DOCKER_INTERNAL_HOST @@ -247,13 +248,12 @@ async def _perform_connectivity_check( observer = MessageObserver() set_monitoring_operation("connectivity_check", display_name=display_name) - connectivity = await OpenAIVLModel( - observer, - model_id=model_name, - api_base=model_base_url, - api_key=model_api_key, - ssl_verify=ssl_verify - ).check_connectivity() + connectivity = await build_adapter_fresh( + {"base_url": model_base_url, "api_key": model_api_key, + "ssl_verify": ssl_verify}, + "vlm", "vlm", None, model_name=model_name, + observer=observer, display_name=display_name, + ).health_check() elif model_type == 'stt': voice_service = get_voice_service() diff --git a/backend/services/tool_configuration_service.py b/backend/services/tool_configuration_service.py index 0a7cba3830..9b4bd6056e 100644 --- a/backend/services/tool_configuration_service.py +++ b/backend/services/tool_configuration_service.py @@ -50,7 +50,7 @@ from services.vectordatabase_service import get_embedding_model_by_index_name, get_rerank_model from utils.http_client_utils import create_httpx_client from database.client import minio_client -from services.image_service import get_video_understanding_model, get_vlm_model +from services.model_gateway_service import get_vlm_adapter from nexent.monitor import set_monitoring_context, set_monitoring_operation from services.vectordatabase_service import get_vector_db_core from utils.langchain_utils import discover_langchain_modules @@ -920,9 +920,8 @@ def _validate_local_tool( if not tenant_id or not user_id: raise ToolExecutionException( f"Tenant ID and User ID are required for {tool_name} validation") - # get_vlm_model reads the first multimodal slot, now shown as image understanding. selected_model_id = instantiation_params.get("selected_model_id") - image_to_text_model = get_vlm_model(tenant_id=tenant_id, model_id=selected_model_id) + image_to_text_model = get_vlm_adapter(tenant_id, selected_model_id, slot="vlm") vlm_display_name = getattr( image_to_text_model, 'display_name', None) set_monitoring_context(tenant_id=tenant_id) @@ -940,7 +939,7 @@ def _validate_local_tool( raise ToolExecutionException( f"Tenant ID and User ID are required for {tool_name} validation") selected_model_id = instantiation_params.get("selected_model_id") - video_understanding_model = get_video_understanding_model(tenant_id=tenant_id, model_id=selected_model_id) + video_understanding_model = get_vlm_adapter(tenant_id, selected_model_id, slot="vlm3") model_display_name = getattr( video_understanding_model, 'display_name', None) set_monitoring_context(tenant_id=tenant_id) diff --git a/sdk/nexent/core/gateway/__init__.py b/sdk/nexent/core/gateway/__init__.py index 4128b7fb6b..02d6fa6a15 100644 --- a/sdk/nexent/core/gateway/__init__.py +++ b/sdk/nexent/core/gateway/__init__.py @@ -1,17 +1,7 @@ """Multimodal model unified adaptation gateway. -Provides a protocol-agnostic adapter layer that *composes* the existing, -stable model classes (``OpenAIModel`` / embedding / rerank …) behind a single -:class:`MultimodalGateway` entry point. For STT/TTS/VLM the protocol lives -directly in the adapter (no wrapped model class); LLM/LongContext stay as thin -``has-a`` delegation to ``OpenAIModel``. Adding a vendor becomes one -``@register_adapter(factory, modality)`` decorator — backend services no -longer hardcode ``if model_factory == ...`` dispatch. - -Importing :mod:`nexent.core.gateway` (or its :mod:`.modality` subpackage) -registers all built-in adapters with the process-wide registry. - -See ``doc/multimodal-gateway-design.md`` for the full design. +A protocol-agnostic adapter layer behind a single :class:`MultimodalGateway` +entry point; importing this package registers all built-in adapters. """ from .multimodal_adapter import ModelInfo, MultimodalAdapter diff --git a/sdk/nexent/core/gateway/modality/__init__.py b/sdk/nexent/core/gateway/modality/__init__.py index 4606379cf7..bf56c01ac1 100644 --- a/sdk/nexent/core/gateway/modality/__init__.py +++ b/sdk/nexent/core/gateway/modality/__init__.py @@ -1,18 +1,14 @@ """Modality adapter aggregation layer. -This is the single place that re-exports the public adapter API and triggers -built-in adapter registration: importing :mod:`nexent.core.gateway.modality` -imports every built-in adapter module, whose ``@register_adapter`` decorators -populate the process-wide :class:`AdapterRegistry`. - -Modality subpackages (``llm`` / ``vlm`` / ``stt`` / ``tts`` / ``embedding`` / -``rerank``) are namespace packages on purpose: they ship no ``__init__.py`` so -this module stays the single aggregation point. Import concrete classes via -this layer or via their leaf module (``modality.vlm.openai``). - -This is the **base branch**: no adapters are registered yet, so the registry is -empty. Each feature branch (``feat/gw-vlm``, ``-llm``, …) appends its modality's -import lines here, which registers that modality's adapters. +Re-exports the public adapter API and triggers registration of all built-in +adapters via the ``@register_adapter`` decorators on import. """ -__all__: list[str] = [] +from .vlm.modelengine import ModelEngineVLMAdapter +from .vlm.openai import OpenAIVLMAdapter +from .vlm.vlm_adapter import VLMAdapter, VLMRequest + +__all__: list[str] = [ + # VLM + "VLMAdapter", "VLMRequest", "OpenAIVLMAdapter", "ModelEngineVLMAdapter", +] diff --git a/sdk/nexent/core/gateway/modality/vlm/modelengine.py b/sdk/nexent/core/gateway/modality/vlm/modelengine.py new file mode 100644 index 0000000000..8b297784e6 --- /dev/null +++ b/sdk/nexent/core/gateway/modality/vlm/modelengine.py @@ -0,0 +1,17 @@ +"""ModelEngine VLM adapter; protocol identical to OpenAI, only factory differs.""" + +from __future__ import annotations + +from ...registry import register_adapter +from .openai import OpenAIVLMAdapter + + +@register_adapter("modelengine", "vlm") +class ModelEngineVLMAdapter(OpenAIVLMAdapter): + """ModelEngine VLM - protocol identical to OpenAI; only ``factory`` differs. + + Attributes: + factory: ``"modelengine"``. + """ + + factory = "modelengine" diff --git a/sdk/nexent/core/gateway/modality/vlm/openai.py b/sdk/nexent/core/gateway/modality/vlm/openai.py new file mode 100644 index 0000000000..a7bcf5df34 --- /dev/null +++ b/sdk/nexent/core/gateway/modality/vlm/openai.py @@ -0,0 +1,320 @@ +"""OpenAI-compatible VLM adapter (protocol lives on the adapter).""" + +from __future__ import annotations + +import asyncio +import base64 +import logging +import os +from typing import Any, BinaryIO, Dict, List, Union + +from ....models import OpenAIModel + +from ...model_context import VLMContext +from ...multimodal_adapter import ModelInfo, MultimodalAdapter +from ...registry import register_adapter +from ...transport import HttpTransportMixin +from .vlm_adapter import VLMAdapter, VLMRequest + + +logger = logging.getLogger(__name__) + +_METHOD_MAP = {"image": "analyze_image", "audio": "analyze_audio", "video": "analyze_video"} + + +@register_adapter("openai", "vlm") +class OpenAIVLMAdapter(VLMAdapter, HttpTransportMixin): + """OpenAI-compatible VLM adapter; chat delegates to a wrapped OpenAIModel.""" + + factory = "openai" + + def __init__(self, context: VLMContext) -> None: + MultimodalAdapter.__init__(self, context) + HttpTransportMixin.__init__( + self, + base_url=context.base_url, + api_key=context.api_key, + ssl_verify=context.ssl_verify, + timeout=context.timeout_seconds if context.timeout_seconds is not None else 30.0, + ) + self._model: Any = None # wrapped OpenAIModel, built lazily + + def _build_model(self) -> None: + """Construct the wrapped :class:`OpenAIModel` with VLM sampling defaults. + + ``frequency_penalty`` is set as a dead instance attribute after + construction, mirroring the old ``OpenAIVLModel``. It must NOT be + passed to ``OpenAIModel.__init__`` because smolagents stores unknown + kwargs in ``self.kwargs`` and merges them into every API call. + """ + ctx = self._context + self._model = OpenAIModel( + observer=ctx.observer, + model_id=ctx.model_name, + api_base=self._base_url, + api_key=self._api_key, + ssl_verify=self._ssl_verify, + model_factory=self.factory, + display_name=ctx.display_name, + temperature=ctx.temperature if ctx.temperature is not None else 0.7, + top_p=ctx.top_p if ctx.top_p is not None else 0.7, + max_tokens=ctx.max_tokens if ctx.max_tokens is not None else 512, + ) + self._model.frequency_penalty = ctx.frequency_penalty if ctx.frequency_penalty is not None else 0.5 + + # VLM protocol (moved from openai_vlm.py) + + def encode_image(self, image_input: Union[str, BinaryIO]) -> str: + """Encode an image file or stream into a base64 string. + + Args: + image_input: A file path or a binary file-like object. + + Returns: + The base64-encoded image bytes as a string. + """ + if isinstance(image_input, str): + with open(image_input, "rb") as image_file: + return base64.b64encode(image_file.read()).decode('utf-8') + return base64.b64encode(image_input.read()).decode('utf-8') + + def prepare_image_message(self, image_input: Union[str, BinaryIO], + system_prompt: str = "Describe this picture.") -> List[Dict[str, Any]]: + """Build OpenAI-compatible chat messages embedding an encoded image. + + When ``image_input`` is a path, the image format is sniffed from the + file extension (defaulting to ``jpeg``). + + Args: + image_input: A file path or a binary file-like object. + system_prompt: System prompt guiding the analysis. + + Returns: + A two-message list (system + user) with the image as a data URL. + """ + base64_image = self.encode_image(image_input) + + image_format = "jpeg" + if isinstance(image_input, str) and os.path.exists(image_input): + _, ext = os.path.splitext(image_input) + if ext.lower() in ['.png', '.jpg', '.jpeg', '.gif', '.webp']: + image_format = ext.lower()[1:] + if image_format == 'jpg': + image_format = 'jpeg' + + messages = [{"role": "system", "content": [{"text": system_prompt, "type": "text"}]}, + {"role": "user", + "content": [{"type": "image_url", + "image_url": {"url": f"data:image/{image_format};base64,{base64_image}", + "detail": "auto"}}]}] + return messages + + def prepare_media_message(self, media_input: Union[str, BinaryIO], media_type: str, + content_type: str, system_prompt: str) -> List[Dict[str, Any]]: + """Build an OpenAI-compatible multimodal message for audio or video. + + Args: + media_input: A file path or a binary file-like object. + media_type: ``"audio"`` or ``"video"``. + content_type: MIME content type, e.g. ``"audio/mpeg"``. + system_prompt: System prompt guiding the analysis. + + Returns: + A user message list with the media as a data URL. + + Raises: + ValueError: If ``media_type`` is not "audio" or "video". + """ + if media_type not in ("audio", "video"): + raise ValueError(f"Unsupported media type: {media_type}") + + base64_media = self.encode_image(media_input) + media_url_key = f"{media_type}_url" + media_config: Dict[str, Any] = {"url": f"data:{content_type};base64,{base64_media}"} + if media_type == "video": + media_config.update({"detail": "high", "max_frames": 16, "fps": 1}) + + messages = [ + { + "role": "user", + "content": [ + {"type": media_url_key, media_url_key: media_config}, + {"type": "text", "text": system_prompt} + ] + } + ] + return messages + + def analyze_image(self, image_input: Union[str, BinaryIO], + system_prompt: str = "Please describe this picture concisely and carefully, within 200 words.", + stream: bool = True, **kwargs) -> Any: + """Analyze image content and return a smolagents ChatMessage. + + Args: + image_input: A file path or a binary file-like object. + system_prompt: System prompt guiding the analysis. + stream: Whether to stream the response. + **kwargs: Additional arguments forwarded to the wrapped model. + + Returns: + A smolagents ``ChatMessage`` with the image analysis. + """ + if self._model is None: + self._build_model() + messages = self.prepare_image_message(image_input, system_prompt) + # Call _model.__call__ explicitly so instance-level mocks work in tests. + return self._model(messages=messages, **kwargs) + + def analyze_audio(self, audio_input: Union[str, BinaryIO], + system_prompt: str = "Please analyze this audio carefully.", + content_type: str = "audio/mpeg", stream: bool = True, **kwargs) -> Any: + """Analyze audio content and return a smolagents ChatMessage. + + Args: + audio_input: A file path or a binary file-like object. + system_prompt: System prompt guiding the analysis. + content_type: MIME content type, e.g. ``"audio/mpeg"``. + stream: Whether to stream the response. Absorbed here so it is + NOT forwarded to the wrapped model — the wrapped + :class:`OpenAIModel` always drives ``stream`` itself, and + forwarding it would collide with its own ``stream=True``. + **kwargs: Additional arguments forwarded to the wrapped model. + + Returns: + A smolagents ``ChatMessage`` with the audio analysis. + """ + if self._model is None: + self._build_model() + messages = self.prepare_media_message(audio_input, "audio", content_type, system_prompt) + return self._model(messages=messages, **kwargs) + + def analyze_video(self, video_input: Union[str, BinaryIO], + system_prompt: str = "Please analyze this video carefully.", + content_type: str = "video/mp4", stream: bool = True, **kwargs) -> Any: + """Analyze video content and return a smolagents ChatMessage. + + Args: + video_input: A file path or a binary file-like object. + system_prompt: System prompt guiding the analysis. + content_type: MIME content type, e.g. ``"video/mp4"``. + stream: Whether to stream the response. Absorbed here so it is + NOT forwarded to the wrapped model — the wrapped + :class:`OpenAIModel` always drives ``stream`` itself, and + forwarding it would collide with its own ``stream=True``. + **kwargs: Additional arguments forwarded to the wrapped model. + + Returns: + A smolagents ``ChatMessage`` with the video analysis. + """ + if self._model is None: + self._build_model() + messages = self.prepare_media_message(video_input, "video", content_type, system_prompt) + return self._model(messages=messages, **kwargs) + + async def check_connectivity(self) -> bool: + """Check VLM connectivity by sending a test image + text prompt. + + Probes with the local ``assets/git-flow.png`` asset, falling back to a + hardcoded DashScope URL when the asset is missing. + + Returns: + True if the probe succeeds, False on any exception. + """ + if self._model is None: + self._build_model() + module_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + test_image_path = os.path.join(module_dir, "assets", "git-flow.png") + if os.path.exists(test_image_path): + base64_image = self.encode_image(test_image_path) + _, ext = os.path.splitext(test_image_path) + image_format = ext.lower()[1:] if ext else "png" + if image_format == "jpg": + image_format = "jpeg" + content_parts: List[Dict[str, Any]] = [ + {"type": "image_url", "image_url": {"url": f"data:image/{image_format};base64,{base64_image}"}}, + {"type": "text", "text": "Hello"}, + ] + else: + test_image_url = "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20250925/thtclx/input1.png" + content_parts = [ + {"type": "image_url", "image_url": {"url": test_image_url}}, + {"type": "text", "text": "Hello"}, + ] + + try: + await asyncio.to_thread( + self._model.client.chat.completions.create, + model=self._model.model_id, + messages=[{"role": "user", "content": content_parts}], + max_tokens=5, + stream=False, + ) + return True + except Exception: + logger.exception("VLM connectivity check failed") + return False + + # Adapter contract + + async def invoke(self, request: VLMRequest) -> Any: + """Dispatch to the matching ``analyze_*`` method and return a ChatMessage.""" + return self.invoke_sync(request) + + def invoke_sync(self, request: VLMRequest) -> Any: + """Dispatch the request to the matching ``analyze_*`` method. + + Args: + request: The VLM request; ``media_type`` selects the method, + ``prompt`` becomes ``system_prompt``, and ``kwargs`` are merged + into the forwarded call. + + Returns: + A smolagents ``ChatMessage`` from the dispatched ``analyze_*`` call. + """ + method = getattr(self, _METHOD_MAP[request.media_type]) + call_kwargs: Dict[str, Any] = {"stream": request.stream} + if request.prompt: + call_kwargs["system_prompt"] = request.prompt + if request.kwargs: + call_kwargs.update(request.kwargs) + return method(request.media_input, **call_kwargs) + + async def health_check(self) -> bool: + """Delegate to :meth:`check_connectivity`.""" + return await self.check_connectivity() + + def _is_siliconflow_non_omni(self) -> bool: + """Check whether this is a SiliconFlow VLM that cannot accept audio input. + + This is the only place that should know which (provider, model) combos + can't ingest a given media type - callers ask the adapter via + :meth:`get_model_info` rather than reaching into the wrapped model's + ``client_kwargs`` / ``model_id``. + + Returns: + True if the provider is SiliconFlow and the model is not Qwen3-Omni. + """ + return ( + "siliconflow" in (self._context.base_url or "").lower() + and "omni" not in (self._context.model_name or "").lower() + ) + + def get_model_info(self) -> ModelInfo: + """Return ``ModelInfo`` with image/audio/video capabilities. + + Explicit capability overrides from context win; defaults assume a + capable VLM. Provider-specific limitations are computed here so the + capability dict, not the caller's URL-sniffing, is the source of truth. + """ + caps = dict(self._context.capabilities) + if "audio" not in caps and self._is_siliconflow_non_omni(): + caps["audio"] = False + caps.setdefault("image", True) + caps.setdefault("audio", True) + caps.setdefault("video", True) + return ModelInfo( + model_id=self._context.model_name, + display_name=self._context.display_name or "", + provider=self.factory, + capabilities=caps, + ) diff --git a/sdk/nexent/core/gateway/modality/vlm/vlm_adapter.py b/sdk/nexent/core/gateway/modality/vlm/vlm_adapter.py new file mode 100644 index 0000000000..6dc5ba6237 --- /dev/null +++ b/sdk/nexent/core/gateway/modality/vlm/vlm_adapter.py @@ -0,0 +1,49 @@ +"""VLM (vision-language model) adapter root + request type.""" + +from __future__ import annotations + +from abc import abstractmethod +from dataclasses import dataclass +from typing import Any, BinaryIO, Dict, Optional, Union + +from ...multimodal_adapter import MultimodalAdapter + + +@dataclass +class VLMRequest: + """VLM understanding request. + + Attributes: + media_type: ``"image"`` | ``"audio"`` | ``"video"``. + media_input: A file path or a binary file-like object of the media. + prompt: System prompt guiding the analysis. + stream: Whether to stream the response. + kwargs: Extra arguments forwarded to the ``analyze_*`` method. + """ + + media_type: str # "image" | "audio" | "video" + media_input: Union[str, BinaryIO] + prompt: str = "" + stream: bool = True + kwargs: Optional[Dict[str, Any]] = None + + +class VLMAdapter(MultimodalAdapter): + """VLM adapter root. + + Attributes: + modality: ``"vlm"``. + """ + + modality = "vlm" + + @abstractmethod + async def invoke(self, request: VLMRequest) -> Any: + """Analyze ``media_input`` with ``prompt`` and return a ChatMessage. + + Args: + request: The VLM request describing the media and prompt to use. + + Returns: + A smolagents ``ChatMessage`` for the analyzed media. + """ diff --git a/sdk/nexent/core/gateway/model_context.py b/sdk/nexent/core/gateway/model_context.py index b6f7d77f3d..95e23724b7 100644 --- a/sdk/nexent/core/gateway/model_context.py +++ b/sdk/nexent/core/gateway/model_context.py @@ -8,7 +8,7 @@ @dataclass class ModelContext: - """Base construction context — common + cross-cutting fields. + """Base construction context - common + cross-cutting fields. Subclasses add modality-specific fields. Passing a field the subclass doesn't declare raises TypeError at construction time. diff --git a/sdk/nexent/core/gateway/multimodal_adapter.py b/sdk/nexent/core/gateway/multimodal_adapter.py index 390bc99fc1..4fb63179dc 100644 --- a/sdk/nexent/core/gateway/multimodal_adapter.py +++ b/sdk/nexent/core/gateway/multimodal_adapter.py @@ -34,7 +34,7 @@ class MultimodalAdapter(ABC): :meth:`invoke`, :meth:`health_check`, :meth:`get_model_info`. The root carries no wrapped-model state. Callers reach the model only - through the uniform interface above — never by tunnelling into a wrapped + through the uniform interface above - never by tunnelling into a wrapped instance's attributes. """ diff --git a/sdk/nexent/core/gateway/multimodal_gateway.py b/sdk/nexent/core/gateway/multimodal_gateway.py index 1a9462c356..313519bb68 100644 --- a/sdk/nexent/core/gateway/multimodal_gateway.py +++ b/sdk/nexent/core/gateway/multimodal_gateway.py @@ -1,10 +1,4 @@ -"""MultimodalGateway: the unified entry point replacing hardcoded dispatch. - -Backend services (``image_service``/``voice_service``/``vectordatabase_service``) -resolve an adapter through :meth:`get_adapter` instead of ``if model_factory`` -branches. The gateway caches adapter instances per ``(tenant, modality, slot, -model_name, factory)`` so a given model is constructed once. -""" +"""MultimodalGateway: the unified entry point replacing hardcoded dispatch.""" from __future__ import annotations @@ -56,7 +50,7 @@ async def invoke(self, context: ModelContext, request: Any) -> Any: return await self.get_adapter(context).invoke(request) def stream(self, context: ModelContext, request: Any): - """Returns the adapter's async iterator (not awaited — it's a generator). + """Returns the adapter's async iterator (not awaited - it's a generator). Args: context: The construction context identifying the desired model. diff --git a/sdk/nexent/core/gateway/registry.py b/sdk/nexent/core/gateway/registry.py index bfd7c51a8e..4526e6e018 100644 --- a/sdk/nexent/core/gateway/registry.py +++ b/sdk/nexent/core/gateway/registry.py @@ -1,9 +1,4 @@ -"""Adapter registry: maps ``(factory, modality)`` → adapter class. - -Paradigm aligned with :mod:`sdk.nexent.memory.providers.registry`. Vendors opt -in via the ``@register_adapter(factory, modality)`` decorator on their adapter -class; the backend never hardcodes vendor dispatch (``if model_factory == ...``). -""" +"""Process-wide adapter registry keyed by ``(factory, modality)``.""" from __future__ import annotations diff --git a/sdk/nexent/core/gateway/transport.py b/sdk/nexent/core/gateway/transport.py index 426a6b87ba..c916b4d16e 100644 --- a/sdk/nexent/core/gateway/transport.py +++ b/sdk/nexent/core/gateway/transport.py @@ -1,12 +1,4 @@ -"""Transport-layer abstraction, orthogonal to modality logic. - -Inspired by Pipecat's ``WebsocketService`` mixin design: transport concerns -(HTTP base_url/api_key vs WebSocket ws_url/auth_headers) live in mixins that -are multiply-inherited alongside a modality ABC, decoupling the STT/TTS -adapters from a hardcoded WebSocket assumption. An HTTP-only vendor (e.g. -ModelEngine STT/TTS) can therefore take :class:`HttpTransportMixin` instead of -``WebSocketTransportMixin``. -""" +"""Transport-layer mixins (HTTP / WebSocket), orthogonal to modality logic.""" from typing import Optional diff --git a/sdk/nexent/core/tools/analyze_audio_tool.py b/sdk/nexent/core/tools/analyze_audio_tool.py index 534aff39ae..d6cb80ed1a 100644 --- a/sdk/nexent/core/tools/analyze_audio_tool.py +++ b/sdk/nexent/core/tools/analyze_audio_tool.py @@ -7,13 +7,13 @@ import logging from io import BytesIO -from typing import List, Optional +from typing import Any, List, Optional from jinja2 import StrictUndefined, Template from pydantic import Field from smolagents.tools import Tool -from ...core.models import OpenAIVLModel +from ...core.gateway.modality import VLMRequest from ...core.utils.observer import MessageObserver, ProcessType from ...core.utils.prompt_template_utils import get_prompt_template from ...core.utils.tools_common_message import ToolCategory, ToolSign @@ -74,7 +74,7 @@ def __init__( description="Message observer", default=None, exclude=True), - vlm_model: OpenAIVLModel = Field( + vlm_model: Any = Field( description="The video understanding model to use", default=None, exclude=True), @@ -110,12 +110,16 @@ def __init__( def _validate_audio_capable_model(self) -> None: - """Fail early for SiliconFlow models that are known not to accept audio input.""" - client_kwargs = getattr(self.vlm_model, "client_kwargs", {}) or {} - base_url = client_kwargs.get("base_url", "") if isinstance(client_kwargs, dict) else "" - model_id = str(getattr(self.vlm_model, "model_id", "") or "") + """Fail early if the VLM cannot accept audio input (e.g. SiliconFlow non-omni). - if "siliconflow" in str(base_url).lower() and model_id and "omni" not in model_id.lower(): +Asks the adapter through the uniform :meth:`get_model_info` interface instead +of reaching into the wrapped model internals. + +Raises: + ValueError: If the selected VLM does not support audio input. +""" + info = self.vlm_model.get_model_info() + if not info.capabilities.get("audio", True): raise ValueError( "The selected video understanding model does not support audio input on SiliconFlow. " "Please choose a Qwen3-Omni model for analyze_audio." @@ -165,10 +169,13 @@ def _forward_impl( content_type = "audio/mpeg" audio_stream = BytesIO(audio_bytes) try: - response = self.vlm_model.analyze_audio( - audio_input=audio_stream, - system_prompt=system_prompt, - content_type=content_type, + response = self.vlm_model.invoke_sync( + VLMRequest( + media_type="audio", + media_input=audio_stream, + prompt=system_prompt, + kwargs={"content_type": content_type}, + ) ) except Exception as e: error_msg_zh = f"音频{index}分析失败: {str(e)}。请检查视频理解模型配置是否正确。" diff --git a/sdk/nexent/core/tools/analyze_image_tool.py b/sdk/nexent/core/tools/analyze_image_tool.py index f4262c292a..d08bfe769e 100644 --- a/sdk/nexent/core/tools/analyze_image_tool.py +++ b/sdk/nexent/core/tools/analyze_image_tool.py @@ -7,13 +7,13 @@ import logging from io import BytesIO -from typing import List +from typing import Any, List from jinja2 import Template, StrictUndefined from pydantic import Field from smolagents.tools import Tool -from ...core.models import OpenAIVLModel +from ...core.gateway.modality import VLMRequest from ...core.utils.observer import MessageObserver, ProcessType from ...core.utils.prompt_template_utils import get_prompt_template from ...core.utils.tools_common_message import ToolCategory, ToolSign @@ -76,7 +76,7 @@ def __init__( description="Message observer", default=None, exclude=True), - vlm_model: OpenAIVLModel = Field( + vlm_model: Any = Field( description="The image understanding model to use", default=None, exclude=True), @@ -168,9 +168,12 @@ def _forward_impl(self, image_urls_list: List[bytes], query: str) -> List[str]: logger.info(f"Extracting image #{index}, query: {query}") image_stream = BytesIO(image_bytes) try: - response = self.vlm_model.analyze_image( - image_input=image_stream, - system_prompt=system_prompt + response = self.vlm_model.invoke_sync( + VLMRequest( + media_type="image", + media_input=image_stream, + prompt=system_prompt, + ) ) except Exception as e: error_msg_zh = f"图片{index}分析失败: {str(e)}。请检查图片理解模型配置是否正确。" diff --git a/sdk/nexent/core/tools/analyze_video_tool.py b/sdk/nexent/core/tools/analyze_video_tool.py index 0393ec3325..3434a423c7 100644 --- a/sdk/nexent/core/tools/analyze_video_tool.py +++ b/sdk/nexent/core/tools/analyze_video_tool.py @@ -7,13 +7,13 @@ import logging from io import BytesIO -from typing import List, Optional +from typing import Any, List, Optional from jinja2 import StrictUndefined, Template from pydantic import Field from smolagents.tools import Tool -from ...core.models import OpenAIVLModel +from ...core.gateway.modality import VLMRequest from ...core.utils.observer import MessageObserver, ProcessType from ...core.utils.prompt_template_utils import get_prompt_template from ...core.utils.tools_common_message import ToolCategory, ToolSign @@ -74,7 +74,7 @@ def __init__( description="Message observer", default=None, exclude=True), - vlm_model: OpenAIVLModel = Field( + vlm_model: Any = Field( description="The video understanding model to use", default=None, exclude=True), @@ -151,10 +151,13 @@ def _forward_impl( content_type = "video/mp4" video_stream = BytesIO(video_bytes) try: - response = self.vlm_model.analyze_video( - video_input=video_stream, - system_prompt=system_prompt, - content_type=content_type, + response = self.vlm_model.invoke_sync( + VLMRequest( + media_type="video", + media_input=video_stream, + prompt=system_prompt, + kwargs={"content_type": content_type} if content_type else None, + ) ) except Exception as e: error_msg_zh = f"视频{index}分析失败: {str(e)}。请检查视频理解模型配置是否正确。" diff --git a/test/backend/agents/test_create_agent_info.py b/test/backend/agents/test_create_agent_info.py index c142990fad..e3869a882c 100644 --- a/test/backend/agents/test_create_agent_info.py +++ b/test/backend/agents/test_create_agent_info.py @@ -250,6 +250,11 @@ def model_validate(cls, value): get_vlm_model=MagicMock(return_value="stub_vlm"), get_video_understanding_model=MagicMock(return_value="stub_video_vlm"), ) +sys.modules['services.model_gateway_service'] = _create_stub_module( + "services.model_gateway_service", + get_llm_adapter=MagicMock(return_value="stub_llm_adapter"), + get_vlm_adapter=MagicMock(return_value="stub_vlm_adapter"), +) sys.modules['services.memory_config_service'] = MagicMock() # Extend services hierarchy with additional stubs sys.modules['services.file_management_service'] = _create_stub_module( @@ -1121,7 +1126,7 @@ async def test_create_tool_config_list_with_analyze_image_tool(self): with patch('backend.agents.create_agent_info.discover_langchain_tools', return_value=[]), \ patch('backend.agents.create_agent_info.search_tools_for_sub_agent') as mock_search_tools, \ - patch('backend.agents.create_agent_info.get_vlm_model') as mock_get_vlm_model, \ + patch('backend.agents.create_agent_info.get_vlm_adapter') as mock_get_vlm_model, \ patch('backend.agents.create_agent_info.minio_client', new_callable=MagicMock) as mock_minio_client: mock_search_tools.return_value = [ @@ -1142,7 +1147,7 @@ async def test_create_tool_config_list_with_analyze_image_tool(self): assert len(result) == 1 assert result[0] is mock_tool_instance - mock_get_vlm_model.assert_called_once_with(tenant_id="tenant_1", model_id=None) + mock_get_vlm_model.assert_called_once_with("tenant_1", None, slot="vlm") # Verify metadata includes validate_url_access lambda assert "vlm_model" in mock_tool_instance.metadata assert "storage_client" in mock_tool_instance.metadata @@ -1165,7 +1170,7 @@ async def test_create_tool_config_list_with_audio_video_tools(self, class_name, with patch('backend.agents.create_agent_info.discover_langchain_tools', return_value=[]), \ patch('backend.agents.create_agent_info.search_tools_for_sub_agent') as mock_search_tools, \ - patch('backend.agents.create_agent_info.get_video_understanding_model') as mock_get_video_model, \ + patch('backend.agents.create_agent_info.get_vlm_adapter') as mock_get_video_model, \ patch('backend.agents.create_agent_info.minio_client', new_callable=MagicMock): mock_search_tools.return_value = [ @@ -1186,7 +1191,7 @@ async def test_create_tool_config_list_with_audio_video_tools(self, class_name, assert len(result) == 1 assert result[0] is mock_tool_instance - mock_get_video_model.assert_called_once_with(tenant_id="tenant_1", model_id=None) + mock_get_video_model.assert_called_once_with("tenant_1", None, slot="vlm3") assert mock_tool_instance.metadata["vlm_model"] == "mock_video_model" assert "storage_client" in mock_tool_instance.metadata assert callable(mock_tool_instance.metadata["validate_url_access"]) @@ -1851,7 +1856,7 @@ async def test_create_tool_config_list_analyze_image_tool_validate_url_access(se with patch('backend.agents.create_agent_info.ToolConfig') as mock_tool_config, \ patch('backend.agents.create_agent_info.discover_langchain_tools', return_value=[]), \ patch('backend.agents.create_agent_info.search_tools_for_sub_agent') as mock_search_tools, \ - patch('backend.agents.create_agent_info.get_vlm_model') as mock_get_vlm_model, \ + patch('backend.agents.create_agent_info.get_vlm_adapter') as mock_get_vlm_model, \ patch('backend.agents.create_agent_info.minio_client', new_callable=MagicMock), \ patch('backend.agents.create_agent_info.validate_urls_access') as mock_validate: @@ -2204,7 +2209,7 @@ async def test_create_agent_config_disabled_compression_still_builds_components( @pytest.mark.asyncio async def test_create_agent_config_basic(self): """Test case for basic agent configuration creation""" - # Reset module-level mock — parallel_executor appends an extra + # Reset module-level mock - parallel_executor appends an extra # ToolConfig call after create_tool_config_list returns. Both # call history and side_effect must be cleared because prior # tests may have left an exhausted iterator on the shared mock. @@ -6911,7 +6916,7 @@ async def test_memory_service_build_failure_skips_store_tool(self): ) mock_search_tool.return_value.forward = MagicMock(return_value="") - # Should NOT raise — the exception is caught and warning logged + # Should NOT raise - the exception is caught and warning logged result = await create_agent_config("a1", "t1", "u1", "en", "query") assert result is not None tool_names = [ @@ -6975,7 +6980,7 @@ async def test_memory_context_service_failure_skips_presearch(self): ) mock_search_tool.return_value.forward = MagicMock(return_value="") - # Should NOT raise — the error is caught + # Should NOT raise - the error is caught result = await create_agent_config("a1", "t1", "u1", "en", "query") assert result is not None mock_search_tool.assert_not_called() diff --git a/test/backend/services/test_model_gateway_service.py b/test/backend/services/test_model_gateway_service.py new file mode 100644 index 0000000000..f7e7b88030 --- /dev/null +++ b/test/backend/services/test_model_gateway_service.py @@ -0,0 +1,291 @@ +"""Unit tests for backend.services.model_gateway_service. + +Covers factory normalization, context construction, adapter resolution and +the voice/VLM config-fetch helpers. Heavy database / utils / consts modules +are stubbed before import so only the gateway bridge logic is exercised. +""" + +import os +import sys +from unittest import mock + +import pytest + +# Dynamically determine the backend path +current_dir = os.path.dirname(os.path.abspath(__file__)) +backend_dir = os.path.abspath(os.path.join(current_dir, "../../../backend")) +sys.path.append(backend_dir) + + +class MockModule(mock.MagicMock): + @classmethod + def __getattr__(cls, key): + return mock.MagicMock() # Return a regular MagicMock instead of a new MockModule + + +# Mock required heavy modules before any import of the service occurs. +sys.modules['database'] = MockModule() +sys.modules['database.model_management_db'] = MockModule() +sys.modules['utils'] = MockModule() +sys.modules['utils.config_utils'] = MockModule() +sys.modules['consts'] = MockModule() +consts_const_module = MockModule() +consts_const_module.MODEL_CONFIG_MAPPING = {"vlm": "vlm_config_key", "vlm3": "vlm3_config_key"} +consts_const_module.TEST_PCM_PATH = "/tmp/test_voice.pcm" +sys.modules['consts.const'] = consts_const_module + +from nexent import MessageObserver +from nexent.core.gateway import VLMContext + +from backend.services import model_gateway_service as mgs + +# _normalize_factory + + +VOLC_CN = "\u706b\u5c71\u5f15\u64ce" # volcengine CN alias +ALI_CN = "\u963f\u91cc\u4e91" # alibaba CN alias + + +def test_normalize_factory_volc_aliases(): + with mock.patch.object(mgs, "get_registry") as mock_get_registry: + mock_get_registry.return_value.has.return_value = True + for raw in ("volc", "volcano", "volcengine", VOLC_CN, " VOLC "): + assert mgs._normalize_factory(raw, "vlm") == "volc" + for raw in ("ali", "alibaba", ALI_CN): + assert mgs._normalize_factory(raw, "vlm") == "ali" + assert mgs._normalize_factory("dashscope", "vlm") == "dashscope" + + +def test_normalize_factory_stt_tts_route_dashscope_to_ali(): + with mock.patch.object(mgs, "get_registry") as mock_get_registry: + mock_get_registry.return_value.has.return_value = True + assert mgs._normalize_factory("dashscope", "stt") == "ali" + assert mgs._normalize_factory("ali", "tts") == "ali" + + +def test_normalize_factory_unknown_falls_back_to_modality_default(): + with mock.patch.object(mgs, "get_registry") as mock_get_registry: + registry = mock_get_registry.return_value + registry.has.return_value = False + assert mgs._normalize_factory("", "llm") == "openai" + assert mgs._normalize_factory(None, "llm") == "openai" + assert mgs._normalize_factory("unknown", "llm") == "openai" + assert mgs._normalize_factory("", "stt") == "ali" + assert mgs._normalize_factory("", "embedding") == "openai" + assert mgs._normalize_factory("", "multi_embedding") == "jina" + + +def test_normalize_factory_registered_factory_passthrough(): + with mock.patch.object(mgs, "get_registry") as mock_get_registry: + mock_get_registry.return_value.has.return_value = True + assert mgs._normalize_factory("tokenpony", "llm") == "tokenpony" + + +# _coalesce + + +def test_coalesce_returns_first_non_none(): + assert mgs._coalesce(None, None, 3) == 3 + assert mgs._coalesce(None, "x") == "x" + assert mgs._coalesce(None, None) is None + + +def test_coalesce_preserves_falsy_values(): + assert mgs._coalesce(0, 1) == 0 + assert mgs._coalesce(False, True) is False + + +# _config_to_context + + +def test_config_to_context_vlm_full_override(): + observer = mock.MagicMock() + cfg = { + "model_factory": "openai", + "base_url": "https://api.openai.com/v1", + "api_key": "sk-test", + "ssl_verify": False, + "display_name": "GPT-4o", + "timeout_seconds": 9.5, + "temperature": 0, + "top_p": 0, + "max_output_tokens": 128, + "frequency_penalty": 0.4, + "extra_body": {"max_completion_tokens": 200}, + "max_tokens": 64, + } + with mock.patch.object(mgs, "get_registry") as mock_get_registry: + mock_get_registry.return_value.has.return_value = True + ctx = mgs._config_to_context( + cfg, + "vlm", + "vlm", + "tenant-1", + model_name="gpt-4o", + capabilities={"video": False}, + observer=observer, + ) + + assert isinstance(ctx, VLMContext) + assert ctx.model_name == "gpt-4o" + assert ctx.base_url == "https://api.openai.com/v1" + assert ctx.api_key == "sk-test" + assert ctx.modality == "vlm" + assert ctx.factory == "openai" + assert ctx.tenant_id == "tenant-1" + assert ctx.slot == "vlm" + assert ctx.ssl_verify is False + assert ctx.observer is observer + assert ctx.display_name == "GPT-4o" + assert ctx.timeout_seconds == 9.5 + assert ctx.temperature == 0 + assert ctx.top_p == 0 + assert ctx.max_output_tokens == 128 + assert ctx.frequency_penalty == 0.4 + assert ctx.extra_body == {"max_completion_tokens": 200} + assert ctx.max_tokens == 64 + assert ctx.capabilities == {"video": False} + + +def test_config_to_context_cfg_none_uses_defaults_and_observer(): + with mock.patch.object(mgs, "get_registry") as mock_get_registry, mock.patch.object(mgs, "get_model_name_from_config", return_value="repo/model"): + mock_get_registry.return_value.has.return_value = False + ctx = mgs._config_to_context(None, "vlm", "vlm", None) + + assert ctx.model_name == "repo/model" + assert ctx.factory == "openai" + assert ctx.base_url == "" + assert ctx.api_key == "" + assert ctx.ssl_verify is True + assert isinstance(ctx.observer, MessageObserver) + + +def test_config_to_context_unknown_modality_raises(): + with mock.patch.object(mgs, "get_registry") as mock_get_registry: + mock_get_registry.return_value.has.return_value = False + with pytest.raises(ValueError, match="Unknown modality: embedding"): + mgs._config_to_context({}, "embedding", "slot", None) + + +# Gateway / registry delegation + + +def test_get_adapter_from_config_delegates_to_gateway(): + gateway = mock.MagicMock() + gateway.get_adapter.return_value = "adapter" + with mock.patch.object(mgs, "get_gateway", return_value=gateway): + result = mgs.get_adapter_from_config({"model_factory": "openai"}, "vlm", "vlm", "t1", temperature=0) + + assert result == "adapter" + gateway.get_adapter.assert_called_once() + ctx = gateway.get_adapter.call_args.args[0] + assert isinstance(ctx, VLMContext) + assert ctx.tenant_id == "t1" + assert ctx.temperature == 0 + + +def test_build_adapter_fresh_constructs_without_gateway_cache(): + dummy_class = mock.MagicMock() + with mock.patch.object(mgs, "get_registry") as mock_get_registry: + mock_get_registry.return_value.has.return_value = True + mock_get_registry.return_value.resolve.return_value = dummy_class + result = mgs.build_adapter_fresh({"model_factory": "openai"}, "vlm", "vlm", "t1") + + assert result == dummy_class.return_value + dummy_class.assert_called_once() + assert dummy_class.call_args.args[0].factory == "openai" + + +# _fetch_slot_config + + +def test_fetch_slot_config_by_model_id(): + cfg = {"model_type": "vlm"} + with mock.patch.object(mgs, "get_model_by_model_id", return_value=cfg) as mock_get_model: + assert mgs._fetch_slot_config("t1", 5, "vlm", "vlm") is cfg + mock_get_model.assert_called_once_with(5, "t1") + + +def test_fetch_slot_config_model_not_found_raises(): + with mock.patch.object(mgs, "get_model_by_model_id", return_value=None), \ + pytest.raises(ValueError, match="Model not found: 5"): + mgs._fetch_slot_config("t1", 5, "vlm", "vlm") + + +def test_fetch_slot_config_wrong_model_type_raises(): + with mock.patch.object(mgs, "get_model_by_model_id", return_value={"model_type": "llm"}), \ + pytest.raises(ValueError, match="not a vlm model"): + mgs._fetch_slot_config("t1", 5, "vlm", "vlm") + + +def test_fetch_slot_config_by_slot_key(): + tenant_config_manager = mock.MagicMock() + cfg = {"model_type": "vlm"} + tenant_config_manager.get_model_config.return_value = cfg + with mock.patch.object(mgs, "tenant_config_manager", tenant_config_manager): + assert mgs._fetch_slot_config("t1", None, "vlm", "vlm") is cfg + tenant_config_manager.get_model_config.assert_called_once_with(key="vlm_config_key", tenant_id="t1") + + +# _fetch_voice_config + + +def test_fetch_voice_config_from_tenant_config(): + tenant_config_manager = mock.MagicMock() + cfg = {"model_type": "stt"} + tenant_config_manager.get_model_config.return_value = cfg + with mock.patch.object(mgs, "tenant_config_manager", tenant_config_manager): + assert mgs._fetch_voice_config("t1", "stt") is cfg + tenant_config_manager.get_model_config.assert_called_once_with("t1", "stt") + + +def test_fetch_voice_config_falls_back_to_model_records(): + tenant_config_manager = mock.MagicMock() + tenant_config_manager.get_model_config.side_effect = RuntimeError("db down") + record = {"model_type": "stt"} + with mock.patch.object(mgs, "tenant_config_manager", tenant_config_manager), mock.patch.object(mgs, "get_model_records", return_value=[record]) as mock_records: + assert mgs._fetch_voice_config("t1", "stt") is record + mock_records.assert_called_once_with({"model_type": "stt"}, "t1") + + +def test_fetch_voice_config_returns_none_when_unavailable(): + tenant_config_manager = mock.MagicMock() + tenant_config_manager.get_model_config.return_value = None + with mock.patch.object(mgs, "tenant_config_manager", tenant_config_manager), mock.patch.object(mgs, "get_model_records", return_value=[]): + assert mgs._fetch_voice_config("t1", "stt") is None + + +def test_fetch_voice_config_records_error_returns_none(): + tenant_config_manager = mock.MagicMock() + tenant_config_manager.get_model_config.return_value = None + with mock.patch.object(mgs, "tenant_config_manager", tenant_config_manager), mock.patch.object(mgs, "get_model_records", side_effect=RuntimeError("boom")): + assert mgs._fetch_voice_config("t1", "stt") is None + + +# Public entry points + + +def test_get_vlm_adapter_from_config_delegates(): + with mock.patch.object(mgs, "get_adapter_from_config", return_value="adapter") as mock_delegate: + result = mgs.get_vlm_adapter_from_config({"a": 1}, "t1", "vlm3", temperature=0.3) + + assert result == "adapter" + mock_delegate.assert_called_once_with({"a": 1}, "vlm", "vlm3", "t1", temperature=0.3) + + +def test_get_vlm_adapter_returns_adapter(): + gateway = mock.MagicMock() + gateway.get_adapter.return_value = "adapter" + cfg = {"model_factory": "openai", "model_type": "vlm"} + with mock.patch.object(mgs, "_fetch_slot_config", return_value=cfg) as mock_fetch, mock.patch.object(mgs, "get_gateway", return_value=gateway): + result = mgs.get_vlm_adapter("t1", 5, "vlm") + + assert result == "adapter" + mock_fetch.assert_called_once_with("t1", 5, expected_type="vlm", slot_key="vlm") + gateway.get_adapter.assert_called_once() + assert gateway.get_adapter.call_args.args[0].slot == "vlm" + + +def test_get_vlm_adapter_returns_none_without_config(): + with mock.patch.object(mgs, "_fetch_slot_config", return_value=None): + assert mgs.get_vlm_adapter("t1", None) is None diff --git a/test/backend/services/test_model_health_service.py b/test/backend/services/test_model_health_service.py index 0dcd5bd9cc..1871a9ee87 100644 --- a/test/backend/services/test_model_health_service.py +++ b/test/backend/services/test_model_health_service.py @@ -68,6 +68,7 @@ def __init__(self, *args, **kwargs): # Mock services packages sys.modules['services'] = MockModule() sys.modules['services.voice_service'] = MockModule() +sys.modules['services.model_gateway_service'] = MockModule() # Define the ModelConnectStatusEnum for testing @@ -194,14 +195,13 @@ async def test_perform_connectivity_check_llm(): async def test_perform_connectivity_check_vlm(): # Setup with mock.patch("backend.services.model_health_service.MessageObserver") as mock_observer, \ - mock.patch("backend.services.model_health_service.OpenAIVLModel") as mock_model: + mock.patch("backend.services.model_health_service.build_adapter_fresh") as mock_adapter_fresh: mock_observer_instance = mock.MagicMock() mock_observer.return_value = mock_observer_instance - mock_model_instance = mock.MagicMock() - mock_model_instance.check_connectivity = mock.AsyncMock( - return_value=True) - mock_model.return_value = mock_model_instance + mock_adapter = mock.MagicMock() + mock_adapter.health_check = mock.AsyncMock(return_value=True) + mock_adapter_fresh.return_value = mock_adapter # Execute result = await _perform_connectivity_check( @@ -213,14 +213,15 @@ async def test_perform_connectivity_check_vlm(): # Assert assert result is True - mock_model.assert_called_once_with( - mock_observer_instance, - model_id="gpt-4-vision", - api_base="https://api.openai.com", - api_key="test-key", - ssl_verify=True + mock_adapter_fresh.assert_called_once_with( + {"base_url": "https://api.openai.com", "api_key": "test-key", + "ssl_verify": True}, + "vlm", "vlm", None, + model_name="gpt-4-vision", + observer=mock_observer_instance, + display_name=None, ) - mock_model_instance.check_connectivity.assert_called_once() + mock_adapter.health_check.assert_awaited_once() @pytest.mark.asyncio @@ -231,7 +232,7 @@ async def test_perform_connectivity_check_dashscope_multimodal_uses_provider_cat ]) with mock.patch.dict(sys.modules, {"services.model_provider_service": model_provider_service}), \ - mock.patch("backend.services.model_health_service.OpenAIVLModel") as mock_model: + mock.patch("backend.services.model_health_service.build_adapter_fresh") as mock_adapter_fresh: result = await _perform_connectivity_check( "qwen-image-max", "vlm2", @@ -246,7 +247,7 @@ async def test_perform_connectivity_check_dashscope_multimodal_uses_provider_cat "model_type": "vlm2", "api_key": "test-key", }) - mock_model.assert_not_called() + mock_adapter_fresh.assert_not_called() @pytest.mark.asyncio @@ -985,15 +986,14 @@ async def test_perform_connectivity_check_llm_sets_monitoring_operation(): @pytest.mark.asyncio async def test_perform_connectivity_check_vlm_sets_monitoring_operation(): with mock.patch("backend.services.model_health_service.MessageObserver") as mock_observer, \ - mock.patch("backend.services.model_health_service.OpenAIVLModel") as mock_model, \ + mock.patch("backend.services.model_health_service.build_adapter_fresh") as mock_adapter_fresh, \ mock.patch("backend.services.model_health_service.set_monitoring_operation") as mock_set_op: mock_observer_instance = mock.MagicMock() mock_observer.return_value = mock_observer_instance - mock_model_instance = mock.MagicMock() - mock_model_instance.check_connectivity = mock.AsyncMock( - return_value=True) - mock_model.return_value = mock_model_instance + mock_adapter = mock.MagicMock() + mock_adapter.health_check = mock.AsyncMock(return_value=True) + mock_adapter_fresh.return_value = mock_adapter await _perform_connectivity_check( "gpt-4-vision", "vlm", "https://api.openai.com", "test-key", diff --git a/test/backend/services/test_tool_configuration_service.py b/test/backend/services/test_tool_configuration_service.py index f0bc7af500..0cf5ae2e3c 100644 --- a/test/backend/services/test_tool_configuration_service.py +++ b/test/backend/services/test_tool_configuration_service.py @@ -415,6 +415,10 @@ def validate(self): 'get_vlm_model': MagicMock(), 'get_video_understanding_model': MagicMock(), }, + 'model_gateway_service': { + 'get_llm_adapter': MagicMock(), + 'get_vlm_adapter': MagicMock(), + }, } for service_name, attrs in services_modules.items(): service_module = types.ModuleType(f'services.{service_name}') @@ -447,6 +451,10 @@ def validate(self): 'get_vlm_model': MagicMock(), 'get_video_understanding_model': MagicMock(), }, + 'model_gateway_service': { + 'get_llm_adapter': MagicMock(), + 'get_vlm_adapter': MagicMock(), + }, } for service_name, attrs in services_modules.items(): service_module = types.ModuleType(f'services.{service_name}') @@ -577,8 +585,6 @@ def _stub_get_llm_model(tenant_id): patch('services.tenant_config_service.get_selected_knowledge_list', MagicMock()).start() patch('services.tenant_config_service.build_knowledge_name_mapping', MagicMock()).start() -patch('services.image_service.get_vlm_model', MagicMock()).start() -patch('services.image_service.get_video_understanding_model', MagicMock()).start() patch('backend.database.knowledge_db.get_knowledge_name_map_by_index_names', MagicMock()).start() # Ensure this module always uses the real consts.model instead of mocks injected by other test files. @@ -1221,7 +1227,6 @@ async def test_get_all_mcp_tools_success(self, mock_urljoin, mock_get_tools, moc mock_tools1, mock_tools2, mock_default_tools] mock_urljoin.return_value = "http://default-server.com/sse" - # 导入函数 from backend.services.tool_configuration_service import get_all_mcp_tools result = await get_all_mcp_tools("test_tenant") @@ -1802,7 +1807,7 @@ async def test_update_tool_list_success(self, mock_update_table, mock_get_langch @patch('backend.services.tool_configuration_service.get_langchain_tools') @patch('backend.services.tool_configuration_service.update_tool_table_from_scan_tool_list') async def test_update_tool_list_mcp_error(self, mock_update_table, mock_get_langchain_tools, mock_get_mcp_tools, mock_get_local_tools): - """Test MCP tool retrieval failure scenario — handled gracefully (mcp_tools = []). + """Test MCP tool retrieval failure scenario - handled gracefully (mcp_tools = []). MCP errors should not block local/langchain tool updates. When MCP fails, mcp_tools is set to an empty list and the update continues. @@ -1821,7 +1826,7 @@ async def test_update_tool_list_mcp_error(self, mock_update_table, mock_get_lang from backend.services.tool_configuration_service import update_tool_list - # Should NOT raise — MCP error is logged but not propagated + # Should NOT raise - MCP error is logged but not propagated await update_tool_list("test_tenant", "test_user") # update_tool_table is still called, but with only local + langchain tools @@ -3289,7 +3294,7 @@ class TestValidateLocalToolAnalyzeImage: """Test cases for _validate_local_tool with analyze_image tool.""" @patch('backend.services.tool_configuration_service.minio_client') - @patch('backend.services.tool_configuration_service.get_vlm_model') + @patch('backend.services.tool_configuration_service.get_vlm_adapter') @patch('backend.services.tool_configuration_service._get_tool_class_by_name') @patch('backend.services.tool_configuration_service.inspect.signature') def test_validate_local_tool_analyze_image_success(self, mock_signature, mock_get_class, mock_get_vlm_model, mock_minio_client): @@ -3315,7 +3320,7 @@ def test_validate_local_tool_analyze_image_success(self, mock_signature, mock_ge ) assert result == "analyze image result" - mock_get_vlm_model.assert_called_once_with(tenant_id="tenant1", model_id=None) + mock_get_vlm_model.assert_called_once_with("tenant1", None, slot="vlm") mock_tool_class.assert_called_once() call_kwargs = mock_tool_class.call_args.kwargs assert 'vlm_model' in call_kwargs @@ -3362,7 +3367,7 @@ class TestValidateLocalToolAnalyzeAudioVideo: @pytest.mark.parametrize("tool_name", ["analyze_audio", "analyze_video"]) @patch('backend.services.tool_configuration_service.minio_client') - @patch('backend.services.tool_configuration_service.get_video_understanding_model') + @patch('backend.services.tool_configuration_service.get_vlm_adapter') @patch('backend.services.tool_configuration_service._get_tool_class_by_name') @patch('backend.services.tool_configuration_service.inspect.signature') def test_validate_local_tool_analyze_audio_video_success( @@ -3389,7 +3394,7 @@ def test_validate_local_tool_analyze_audio_video_success( ) assert result == f"{tool_name} result" - mock_get_video_model.assert_called_once_with(tenant_id="tenant1", model_id=None) + mock_get_video_model.assert_called_once_with("tenant1", None, slot="vlm3") call_kwargs = mock_tool_class.call_args.kwargs assert call_kwargs["vlm_model"] == "mock_video_model" assert "storage_client" in call_kwargs @@ -3648,7 +3653,7 @@ class TestValidateLocalToolRAGFlowSearch: @patch('backend.services.tool_configuration_service._get_tool_class_by_name') @patch('backend.services.tool_configuration_service.inspect.signature') def test_validate_local_tool_ragflow_search_success(self, mock_signature, mock_get_class): - """Test successful ragflow_search tool validation — filters out rerank params.""" + """Test successful ragflow_search tool validation - filters out rerank params.""" mock_tool_class = Mock() mock_tool_instance = Mock() mock_tool_instance.forward.return_value = "ragflow search result" @@ -5125,7 +5130,7 @@ class TestValidateLocalToolMonitoring: @patch('backend.services.tool_configuration_service.set_monitoring_operation') @patch('backend.services.tool_configuration_service.set_monitoring_context') @patch('backend.services.tool_configuration_service.minio_client') - @patch('backend.services.tool_configuration_service.get_vlm_model') + @patch('backend.services.tool_configuration_service.get_vlm_adapter') @patch('backend.services.tool_configuration_service._get_tool_class_by_name') @patch('backend.services.tool_configuration_service.inspect.signature') def test_analyze_image_sets_monitoring_context( diff --git a/test/sdk/core/gateway/__init__.py b/test/sdk/core/gateway/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/test/sdk/core/gateway/test_model_context.py b/test/sdk/core/gateway/test_model_context.py new file mode 100644 index 0000000000..4bd521a875 --- /dev/null +++ b/test/sdk/core/gateway/test_model_context.py @@ -0,0 +1,44 @@ +"""Unit tests for modality-specific gateway construction contexts.""" + +from nexent.core.gateway.model_context import LLMContext, VLMContext + + +def test_cache_key_uses_empty_defaults(): + context = VLMContext( + model_name="qwen-vl-max", + base_url="https://api.example.com", + api_key="sk-key", + modality="vlm", + factory="openai", + ) + + assert context.cache_key() == ("", "vlm", "", "qwen-vl-max", "openai") + + +def test_cache_key_includes_tenant_and_slot(): + context = VLMContext( + model_name="qwen-vl-max", + base_url="https://api.example.com", + api_key="sk-key", + modality="vlm", + factory="openai", + tenant_id="tenant-1", + slot="vlm3", + ) + + assert context.cache_key() == ("tenant-1", "vlm", "vlm3", "qwen-vl-max", "openai") + + +def test_subclass_fields_are_independent(): + llm = LLMContext( + model_name="gpt-4o", + base_url="https://api.example.com", + api_key="sk-key", + modality="llm", + factory="openai", + temperature=0.2, + stream=True, + ) + assert llm.temperature == 0.2 + assert llm.stream is True + assert not hasattr(llm, "capabilities") \ No newline at end of file diff --git a/test/sdk/core/gateway/test_multimodal_adapter.py b/test/sdk/core/gateway/test_multimodal_adapter.py new file mode 100644 index 0000000000..60a04d8ea2 --- /dev/null +++ b/test/sdk/core/gateway/test_multimodal_adapter.py @@ -0,0 +1,82 @@ +"""Unit tests for the multimodal adapter root ABC and ModelInfo.""" + +import pytest +from nexent.core.gateway.model_context import VLMContext +from nexent.core.gateway.multimodal_adapter import ModelInfo, MultimodalAdapter + + +class _ConcreteAdapter(MultimodalAdapter): + """Minimal concrete adapter for instance-level tests.""" + + modality = "vlm" + factory = "openai" + + async def invoke(self, request): + return ("invoke", request) + + async def stream(self, request): + return ("stream", request) + + async def health_check(self): + return True + + def get_model_info(self): + return ModelInfo( + model_id=self._context.model_name, + display_name="dummy", + provider=self.factory, + capabilities={"image": True}, + ) + + +def _make_context(**overrides): + fields = { + "model_name": "dummy-model", + "base_url": "https://api.example.com", + "api_key": "sk-key", + "modality": "vlm", + "factory": "openai", + } + fields.update(overrides) + return VLMContext(**fields) + + +def test_model_info_declaration(): + info = ModelInfo( + model_id="dummy-model", + display_name="Dummy VLM", + provider="openai", + capabilities={"image": True, "audio": False}, + ) + + assert info.model_id == "dummy-model" + assert info.display_name == "Dummy VLM" + assert info.provider == "openai" + assert info.capabilities == {"image": True, "audio": False} + + +def test_init_stores_context(): + context = _make_context() + adapter = _ConcreteAdapter(context) + + assert adapter._context is context + + +@pytest.mark.asyncio +async def test_abstract_contract_methods_raise_not_implemented(): + saved_abstractmethods = MultimodalAdapter.__abstractmethods__ + MultimodalAdapter.__abstractmethods__ = frozenset() + try: + adapter = MultimodalAdapter(_make_context()) + adapter.modality = "vlm" + + with pytest.raises(NotImplementedError): + await adapter.invoke({"media": "request"}) + with pytest.raises(NotImplementedError, match="does not support streaming"): + await adapter.stream({"media": "request"}) + with pytest.raises(NotImplementedError): + await adapter.health_check() + with pytest.raises(NotImplementedError): + adapter.get_model_info() + finally: + MultimodalAdapter.__abstractmethods__ = saved_abstractmethods \ No newline at end of file diff --git a/test/sdk/core/gateway/test_multimodal_gateway.py b/test/sdk/core/gateway/test_multimodal_gateway.py new file mode 100644 index 0000000000..c15a38c4ee --- /dev/null +++ b/test/sdk/core/gateway/test_multimodal_gateway.py @@ -0,0 +1,143 @@ +"""Unit tests for MultimodalGateway caching and delegation.""" + +import pytest +from nexent.core.gateway.model_context import VLMContext +from nexent.core.gateway.multimodal_adapter import ModelInfo, MultimodalAdapter +from nexent.core.gateway.multimodal_gateway import MultimodalGateway, get_gateway +from nexent.core.gateway.registry import AdapterRegistry + + +class _FakeAdapter(MultimodalAdapter): + """Concrete adapter used to observe gateway delegation.""" + + modality = "vlm" + factory = "fake" + + async def invoke(self, request): + return ("invoke", request) + + async def stream(self, request): + return ("stream", request) + + async def health_check(self): + return True + + def get_model_info(self): + return ModelInfo( + model_id=self._context.model_name, + display_name="fake", + provider=self.factory, + capabilities={"image": True}, + ) + + +def _make_registry(): + registry = AdapterRegistry() + registry.register("fake", "vlm")(_FakeAdapter) + return registry + + +def _make_context(model_name="dummy-model"): + return VLMContext( + model_name=model_name, + base_url="https://api.example.com", + api_key="sk-key", + modality="vlm", + factory="fake", + tenant_id="tenant-1", + slot="vlm", + ) + + +def test_get_adapter_builds_and_caches_by_context(): + gateway = MultimodalGateway(_make_registry()) + context = _make_context() + + first = gateway.get_adapter(context) + second = gateway.get_adapter(context) + + assert isinstance(first, _FakeAdapter) + assert first is second + assert first._context is context + + +def test_get_adapter_builds_separate_instance_for_different_key(): + gateway = MultimodalGateway(_make_registry()) + + first = gateway.get_adapter(_make_context("model-a")) + second = gateway.get_adapter(_make_context("model-b")) + + assert first is not second + + +def test_gateway_defaults_to_process_registry(): + gateway = MultimodalGateway() + context = VLMContext( + model_name="gpt-4o", + base_url="https://api.example.com", + api_key="sk-key", + modality="vlm", + factory="openai", + ) + + adapter = gateway.get_adapter(context) + assert adapter.factory == "openai" + + +@pytest.mark.asyncio +async def test_invoke_delegates_to_adapter(gateway, context): + result = await gateway.invoke(context, {"media": "request"}) + + assert result == ("invoke", {"media": "request"}) + + +@pytest.mark.asyncio +async def test_stream_delegates_to_adapter(gateway, context): + result = gateway.stream(context, {"media": "request"}) + + assert await result == ("stream", {"media": "request"}) + + +@pytest.mark.asyncio +async def test_health_check_delegates_to_adapter(gateway, context): + assert await gateway.health_check(context) is True + + +def test_invalidate_single_context(gateway, context): + cached = gateway.get_adapter(context) + assert gateway.get_adapter(context) is cached + + gateway.invalidate(context) + assert gateway.get_adapter(context) is not cached + + +def test_invalidate_all_contexts(gateway, context): + other_context = _make_context("model-other") + first_cached = gateway.get_adapter(context) + other_cached = gateway.get_adapter(other_context) + + gateway.invalidate() + + assert gateway.get_adapter(context) is not first_cached + assert gateway.get_adapter(other_context) is not other_cached + + +def test_get_gateway_is_lazy_singleton(): + from nexent.core.gateway import multimodal_gateway as gateway_module + + gateway_module._gateway = None + first = get_gateway() + second = get_gateway() + + assert first is second + assert isinstance(first, MultimodalGateway) + + +@pytest.fixture +def gateway(): + return MultimodalGateway(_make_registry()) + + +@pytest.fixture +def context(): + return _make_context() \ No newline at end of file diff --git a/test/sdk/core/gateway/test_registry.py b/test/sdk/core/gateway/test_registry.py new file mode 100644 index 0000000000..77036f774d --- /dev/null +++ b/test/sdk/core/gateway/test_registry.py @@ -0,0 +1,38 @@ +"""Unit tests for the process-wide adapter registry.""" + +import pytest +from nexent.core.gateway.registry import AdapterRegistry, get_registry, register_adapter + + +class _DummyAdapter: + """Plain placeholder class used as a registered adapter target.""" + + +def test_register_as_decorator_and_resolve(): + registry = AdapterRegistry() + registry.register("Fake", "vlm")(_DummyAdapter) + + assert registry.resolve("fake", "vlm") is _DummyAdapter + # Resolution is case-insensitive and strips surrounding whitespace. + assert registry.resolve(" FAKE ", "vlm") is _DummyAdapter + assert registry.has("fake", "vlm") is True + assert registry.has("other", "vlm") is False + assert registry.list_adapters() == [("fake", "vlm")] + + +def test_resolve_missing_pair_raises_key_error(): + registry = AdapterRegistry() + with pytest.raises(KeyError, match="No adapter registered for factory='nope'"): + registry.resolve("nope", "vlm") + + +def test_get_registry_returns_shared_singleton(): + assert get_registry() is get_registry() + + +def test_register_adapter_module_level_alias(): + @register_adapter("dummy", "vlm") + class _DummyVLMAdapter: + pass + + assert get_registry().resolve("dummy", "vlm") is _DummyVLMAdapter \ No newline at end of file diff --git a/test/sdk/core/gateway/test_transport.py b/test/sdk/core/gateway/test_transport.py new file mode 100644 index 0000000000..4a77873d86 --- /dev/null +++ b/test/sdk/core/gateway/test_transport.py @@ -0,0 +1,97 @@ +"""Unit tests for the HTTP / WebSocket transport mixins.""" + +from unittest.mock import AsyncMock + +import pytest +from nexent.core.gateway.transport import HttpTransportMixin, WebSocketTransportMixin + + +def test_http_transport_defaults(): + transport = HttpTransportMixin(base_url="https://api.example.com", api_key="sk-key") + + assert transport.transport_type == "http" + assert transport._base_url == "https://api.example.com" + assert transport._api_key == "sk-key" + assert transport._ssl_verify is True + assert transport._timeout == 30.0 + + +def test_http_transport_custom_timeout_and_ssl(): + transport = HttpTransportMixin( + base_url="https://api.example.com", + api_key="sk-key", + ssl_verify=False, + timeout=5.5, + ) + + assert transport._ssl_verify is False + assert transport._timeout == 5.5 + + +@pytest.mark.asyncio +async def test_http_transport_connect_close_health_check(): + transport = HttpTransportMixin(base_url="https://api.example.com", api_key="sk-key") + + assert await transport.connect() is None + assert await transport.close() is None + assert await transport.health_check() is True + + +def test_ws_transport_defaults(): + transport = WebSocketTransportMixin() + + assert transport.transport_type == "websocket" + assert transport._ws_url is None + assert transport._auth_headers == {} + assert transport._ws_connection is None + + +def test_ws_transport_with_params(): + transport = WebSocketTransportMixin(ws_url="wss://example.com/ws", auth_headers={"token": "abc"}) + + assert transport._ws_url == "wss://example.com/ws" + assert transport._auth_headers == {"token": "abc"} + + +@pytest.mark.asyncio +async def test_ws_transport_connect_and_close_connection(): + transport = WebSocketTransportMixin(ws_url="wss://example.com/ws") + assert await transport.connect() is None + + connection = AsyncMock() + transport._ws_connection = connection + await transport.close() + + connection.close.assert_awaited_once() + assert transport._ws_connection is None + + +@pytest.mark.asyncio +async def test_ws_transport_close_without_connection_is_idempotent(): + transport = WebSocketTransportMixin(ws_url="wss://example.com/ws") + + await transport.close() + + assert transport._ws_connection is None + + +@pytest.mark.asyncio +async def test_ws_transport_close_clears_connection_on_error(): + transport = WebSocketTransportMixin(ws_url="wss://example.com/ws") + connection = AsyncMock() + connection.close.side_effect = RuntimeError("connection reset") + transport._ws_connection = connection + + with pytest.raises(RuntimeError, match="connection reset"): + await transport.close() + + assert transport._ws_connection is None + + +@pytest.mark.asyncio +async def test_ws_transport_health_check(): + transport_with_url = WebSocketTransportMixin(ws_url="wss://example.com/ws") + transport_without_url = WebSocketTransportMixin() + + assert await transport_with_url.health_check() is True + assert await transport_without_url.health_check() is False \ No newline at end of file diff --git a/test/sdk/core/models/test_vlm_adapter.py b/test/sdk/core/models/test_vlm_adapter.py new file mode 100644 index 0000000000..6e58dcf196 --- /dev/null +++ b/test/sdk/core/models/test_vlm_adapter.py @@ -0,0 +1,501 @@ +"""Tests for OpenAIVLMAdapter - the VLM protocol that lives on the adapter.""" + +import asyncio +import base64 +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from nexent.core.gateway.modality import OpenAIVLMAdapter + + +@pytest.fixture() +def vlm_adapter(): + """Return an OpenAIVLMAdapter with a mocked _model.""" + adapter = OpenAIVLMAdapter.__new__(OpenAIVLMAdapter) + inner = MagicMock() + inner.model_id = "dummy-model" + inner.client.chat.completions.create = MagicMock() + adapter._model = inner + return adapter + + +# Tests for check_connectivity + + +@pytest.mark.asyncio +async def test_check_connectivity_success(vlm_adapter): + """check_connectivity should return True when no exception is raised.""" + with patch.object( + asyncio, + "to_thread", + new_callable=AsyncMock, + return_value=None, + ) as mock_to_thread: + result = await vlm_adapter.check_connectivity() + + assert result is True + mock_to_thread.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_check_connectivity_failure(vlm_adapter): + """check_connectivity should return False when to_thread raises.""" + with patch.object( + asyncio, + "to_thread", + new_callable=AsyncMock, + side_effect=Exception("connection error"), + ): + result = await vlm_adapter.check_connectivity() + assert result is False + + +@pytest.mark.asyncio +async def test_check_connectivity_uses_fallback_url(vlm_adapter): + """check_connectivity should use fallback remote URL when local image missing.""" + + async def mock_to_thread_func(*args, **kwargs): + return None + + with patch.object(vlm_adapter, "encode_image", return_value=""), \ + patch.object(asyncio, "to_thread", new_callable=AsyncMock, + side_effect=mock_to_thread_func): + import os.path + + with patch.object(os.path, "exists", return_value=False): + result = await vlm_adapter.check_connectivity() + + assert result is True + + +@pytest.mark.asyncio +async def test_check_connectivity_jpg_to_jpeg_conversion(vlm_adapter): + """check_connectivity should convert jpg to jpeg format for MIME type.""" + import os.path + + def mock_exists(path): + return "git-flow" in str(path) + + def mock_splitext(path): + if "git-flow" in str(path): + return ("", ".jpg") + return ("", "") + + async def mock_to_thread_func(*args, **kwargs): + return None + + with patch.object(os.path, "exists", side_effect=mock_exists), \ + patch.object(os.path, "splitext", side_effect=mock_splitext), \ + patch.object(vlm_adapter, "encode_image", return_value="fakebase64"), \ + patch.object(asyncio, "to_thread", new_callable=AsyncMock, + side_effect=mock_to_thread_func): + result = await vlm_adapter.check_connectivity() + + assert result is True + + +# Tests for encode_image + + +def test_encode_image_with_file_path(vlm_adapter, tmp_path): + test_image = tmp_path / "test.png" + test_image.write_bytes(b"fake image data") + + result = vlm_adapter.encode_image(str(test_image)) + + expected = base64.b64encode(b"fake image data").decode('utf-8') + assert result == expected + + +def test_encode_image_with_binary_io(vlm_adapter): + mock_file = MagicMock() + mock_file.read.return_value = b"binary image data" + + result = vlm_adapter.encode_image(mock_file) + + expected = base64.b64encode(b"binary image data").decode('utf-8') + assert result == expected + + +# Tests for prepare_image_message + + +def test_prepare_image_message_with_png_file(vlm_adapter, tmp_path): + test_image = tmp_path / "test.png" + test_image.write_bytes(b"fake png data") + + messages = vlm_adapter.prepare_image_message(str(test_image)) + + assert len(messages) == 2 + assert messages[0]["role"] == "system" + assert messages[1]["role"] == "user" + assert "data:image/png;base64," in messages[1]["content"][0]["image_url"]["url"] + + +def test_prepare_image_message_with_jpg_file(vlm_adapter, tmp_path): + test_image = tmp_path / "test.jpg" + test_image.write_bytes(b"fake jpg data") + + messages = vlm_adapter.prepare_image_message(str(test_image)) + + assert "data:image/jpeg;base64," in messages[1]["content"][0]["image_url"]["url"] + + +def test_prepare_image_message_with_jpeg_file(vlm_adapter, tmp_path): + test_image = tmp_path / "test.jpeg" + test_image.write_bytes(b"fake jpeg data") + + messages = vlm_adapter.prepare_image_message(str(test_image)) + + assert "data:image/jpeg;base64," in messages[1]["content"][0]["image_url"]["url"] + + +def test_prepare_image_message_with_gif_file(vlm_adapter, tmp_path): + test_image = tmp_path / "test.gif" + test_image.write_bytes(b"fake gif data") + + messages = vlm_adapter.prepare_image_message(str(test_image)) + + assert "data:image/gif;base64," in messages[1]["content"][0]["image_url"]["url"] + + +def test_prepare_image_message_with_webp_file(vlm_adapter, tmp_path): + test_image = tmp_path / "test.webp" + test_image.write_bytes(b"fake webp data") + + messages = vlm_adapter.prepare_image_message(str(test_image)) + + assert "data:image/webp;base64," in messages[1]["content"][0]["image_url"]["url"] + + +def test_prepare_image_message_with_binary_io(vlm_adapter): + mock_file = MagicMock() + mock_file.read.return_value = b"binary data" + + messages = vlm_adapter.prepare_image_message(mock_file) + + assert "data:image/jpeg;base64," in messages[1]["content"][0]["image_url"]["url"] + + +def test_prepare_image_message_custom_system_prompt(vlm_adapter, tmp_path): + test_image = tmp_path / "test.png" + test_image.write_bytes(b"fake png data") + + custom_prompt = "What is in this image?" + messages = vlm_adapter.prepare_image_message(str(test_image), system_prompt=custom_prompt) + + assert messages[0]["content"][0]["text"] == custom_prompt + + +# Tests for analyze_image / analyze_audio / analyze_video + + +def test_analyze_image_calls_prepare_image_message(vlm_adapter, tmp_path): + """analyze_image should call prepare_image_message and delegate to _model.""" + test_image = tmp_path / "test.png" + test_image.write_bytes(b"fake png data") + + with patch.object(vlm_adapter, "prepare_image_message", + return_value=[{"role": "user", "content": "test"}]) as mock_prepare: + custom_prompt = "Describe this image" + vlm_adapter.analyze_image(str(test_image), system_prompt=custom_prompt, stream=False) + + mock_prepare.assert_called_once_with(str(test_image), custom_prompt) + vlm_adapter._model.assert_called_once() + # ensure the prepared messages were forwarded to _model + _, kwargs = vlm_adapter._model.call_args + assert kwargs["messages"] == [{"role": "user", "content": "test"}] + + +def test_prepare_media_message_audio(vlm_adapter): + audio_stream = MagicMock() + audio_stream.read.return_value = b"audio bytes" + + messages = vlm_adapter.prepare_media_message( + audio_stream, + media_type="audio", + content_type="audio/mpeg", + system_prompt="Listen carefully", + ) + + assert messages[0]["content"][0]["type"] == "audio_url" + assert messages[0]["content"][0]["audio_url"]["url"].startswith("data:audio/mpeg;base64,") + assert messages[0]["content"][1] == {"type": "text", "text": "Listen carefully"} + + +def test_prepare_media_message_video(vlm_adapter): + video_stream = MagicMock() + video_stream.read.return_value = b"video bytes" + + messages = vlm_adapter.prepare_media_message( + video_stream, + media_type="video", + content_type="video/mp4", + system_prompt="Watch carefully", + ) + + assert messages[0]["content"][0]["type"] == "video_url" + assert messages[0]["content"][0]["video_url"]["url"].startswith("data:video/mp4;base64,") + assert messages[0]["content"][0]["video_url"]["max_frames"] == 16 + assert messages[0]["content"][0]["video_url"]["fps"] == 1 + assert messages[0]["content"][1] == {"type": "text", "text": "Watch carefully"} + + +def test_analyze_audio_calls_prepare_media_message(vlm_adapter): + with patch.object(vlm_adapter, "prepare_media_message", + return_value=[{"role": "user", "content": "test"}]) as mock_prepare: + vlm_adapter.analyze_audio("audio.mp3", system_prompt="Analyze", content_type="audio/mpeg") + + mock_prepare.assert_called_once_with("audio.mp3", "audio", "audio/mpeg", "Analyze") + vlm_adapter._model.assert_called_once() + + +def test_analyze_video_calls_prepare_media_message(vlm_adapter): + with patch.object(vlm_adapter, "prepare_media_message", + return_value=[{"role": "user", "content": "test"}]) as mock_prepare: + vlm_adapter.analyze_video("video.mp4", system_prompt="Analyze", content_type="video/mp4") + + mock_prepare.assert_called_once_with("video.mp4", "video", "video/mp4", "Analyze") + vlm_adapter._model.assert_called_once() + + +def test_invoke_sync_dispatches_by_media_type(vlm_adapter): + """invoke_sync routes to the adapter's own analyze_* via _METHOD_MAP.""" + with patch.object(vlm_adapter, "analyze_image", return_value="img-result") as mock_analyze: + from nexent.core.gateway.modality import VLMRequest + result = vlm_adapter.invoke_sync( + VLMRequest(media_type="image", media_input=b"bytes", prompt="p", stream=False) + ) + assert result == "img-result" + mock_analyze.assert_called_once_with(b"bytes", system_prompt="p", stream=False) + + +# Regression: _build_model must not leak frequency_penalty onto the wire. + + +def test_build_model_does_not_send_frequency_penalty(): + """Real _build_model (offline OpenAIModel construction) must keep + frequency_penalty out of self.kwargs - smolagents merges self.kwargs into + every chat.completions.create, so leaking it would silently send + frequency_penalty=0.5 to the VLM API. The original OpenAIVLModel set it + only as a dead instance attribute (never forwarded, never read). + """ + from nexent.core.gateway.model_context import VLMContext + from nexent.core.utils.observer import MessageObserver + + adapter = OpenAIVLMAdapter(VLMContext( + modality="vlm", + factory="openai", + model_name="qwen-vl-max", + display_name="vlm-test", + base_url="https://dashscope.aliyuncs.com/compatible-mode/v1", + api_key="sk-fake", + ssl_verify=True, + observer=MessageObserver(), + )) + adapter._build_model() # offline - no network call + + inner = adapter._model + # frequency_penalty must NOT ride along in the smolagents model-defaults + # dict that _prepare_completion_kwargs merges into the wire request. + assert "frequency_penalty" not in getattr(inner, "kwargs", {}), inner.kwargs + # the dead instance attribute is still set for getattr parity. + assert inner.frequency_penalty == 0.5 + # sampling defaults the old OpenAIVLModel forwarded must still be set. + assert inner.temperature == 0.7 + assert inner.top_p == 0.7 + assert inner.max_tokens == 512 + + +# get_model_info: the adapter owns the (provider, model) -> capability mapping, +# replacing analyze_audio_tool's getattr + URL sniffing on the wrapped model. + + +def _make_vlm(model_name, base_url, **ctx_overrides): + from nexent.core.gateway.model_context import VLMContext + + return OpenAIVLMAdapter(VLMContext( + modality="vlm", + factory="openai", + model_name=model_name, + display_name="vlm-test", + base_url=base_url, + api_key="sk-fake", + ssl_verify=True, + **ctx_overrides, + )) + + +def test_model_info_siliconflow_non_omni_disables_audio(): + """SiliconFlow non-omni VLMs report audio=False - callers read + get_model_info() instead of sniffing client_kwargs / model_id.""" + adapter = _make_vlm( + "Qwen/Qwen3-VL-32B-Instruct", + "https://api.siliconflow.cn/v1", + ) + info = adapter.get_model_info() + assert info.capabilities["audio"] is False + assert info.capabilities["image"] is True + assert info.capabilities["video"] is True + + +def test_model_info_siliconflow_omni_keeps_audio(): + adapter = _make_vlm( + "Qwen/Qwen3-Omni-7B", + "https://api.siliconflow.cn/v1", + ) + assert adapter.get_model_info().capabilities["audio"] is True + + +def test_model_info_non_siliconflow_keeps_audio(): + adapter = _make_vlm( + "qwen-vl-max", + "https://dashscope.aliyuncs.com/compatible-mode/v1", + ) + assert adapter.get_model_info().capabilities["audio"] is True + + +def test_model_info_explicit_capability_overrides_heuristic(): + """An explicit audio=True in context.capabilities wins over the + SiliconFlow heuristic - the config author declared the capability.""" + from nexent.core.gateway.model_context import VLMContext + + adapter = OpenAIVLMAdapter(VLMContext( + modality="vlm", + factory="openai", + model_name="Qwen/Qwen3-VL-32B-Instruct", + display_name="vlm-test", + base_url="https://api.siliconflow.cn/v1", + api_key="sk-fake", + ssl_verify=True, + capabilities={"audio": True}, + )) + assert adapter.get_model_info().capabilities["audio"] is True + + +# Coverage gaps: unknown image extensions, media type validation, lazy model +# construction, async invoke / health_check delegation, and invoke_sync kwargs. + +def test_prepare_image_message_unknown_extension_defaults_to_jpeg(vlm_adapter, tmp_path): + """Unknown extensions fall back to the default jpeg MIME type.""" + test_image = tmp_path / "test.bmp" + test_image.write_bytes(b"fake bmp data") + + messages = vlm_adapter.prepare_image_message(str(test_image)) + + assert "data:image/jpeg;base64," in messages[1]["content"][0]["image_url"]["url"] + + +def test_prepare_media_message_unsupported_type_raises(vlm_adapter): + """prepare_media_message rejects media types other than audio/video.""" + with pytest.raises(ValueError, match="Unsupported media type: text"): + vlm_adapter.prepare_media_message(b"data", "text", "text/plain", "prompt") + + +def _fresh_vlm_adapter(): + """Return a real OpenAIVLMAdapter whose wrapped model is not built yet.""" + from nexent.core.gateway.model_context import VLMContext + + return OpenAIVLMAdapter(VLMContext( + modality="vlm", + factory="openai", + model_name="qwen-vl-max", + display_name="vlm-test", + base_url="https://dashscope.aliyuncs.com/compatible-mode/v1", + api_key="sk-fake", + ssl_verify=True, + )) + + +def _build_model_side_effect(adapter): + def _build(): + adapter._model = MagicMock() + return _build + + +def test_analyze_image_builds_model_lazily(): + """analyze_image builds the wrapped model when it is missing.""" + adapter = _fresh_vlm_adapter() + + with patch.object(adapter, "_build_model", side_effect=_build_model_side_effect(adapter)) as mock_build, \ + patch.object(adapter, "prepare_image_message", return_value=[{"role": "user", "content": "x"}]): + adapter.analyze_image("dummy.png", system_prompt="Analyze") + + mock_build.assert_called_once() + adapter._model.assert_called_once() + + +def test_analyze_audio_builds_model_lazily(): + """analyze_audio builds the wrapped model when it is missing.""" + adapter = _fresh_vlm_adapter() + + with patch.object(adapter, "_build_model", side_effect=_build_model_side_effect(adapter)) as mock_build, \ + patch.object(adapter, "prepare_media_message", return_value=[{"role": "user", "content": "x"}]): + adapter.analyze_audio("dummy.mp3", system_prompt="Analyze", content_type="audio/mpeg") + + mock_build.assert_called_once() + adapter._model.assert_called_once() + + +def test_analyze_video_builds_model_lazily(): + """analyze_video builds the wrapped model when it is missing.""" + adapter = _fresh_vlm_adapter() + + with patch.object(adapter, "_build_model", side_effect=_build_model_side_effect(adapter)) as mock_build, \ + patch.object(adapter, "prepare_media_message", return_value=[{"role": "user", "content": "x"}]): + adapter.analyze_video("dummy.mp4", system_prompt="Analyze", content_type="video/mp4") + + mock_build.assert_called_once() + adapter._model.assert_called_once() + + +@pytest.mark.asyncio +async def test_check_connectivity_builds_model_and_uses_local_asset(): + """check_connectivity lazily builds the model and probes with local asset.""" + import os.path + + adapter = _fresh_vlm_adapter() + + with patch.object(adapter, "_build_model", side_effect=_build_model_side_effect(adapter)) as mock_build, \ + patch.object(os.path, "exists", return_value=True), \ + patch.object(adapter, "encode_image", return_value="fakebase64"), \ + patch.object(asyncio, "to_thread", new_callable=AsyncMock, return_value=None): + result = await adapter.check_connectivity() + + assert result is True + mock_build.assert_called_once() + + +@pytest.mark.asyncio +async def test_invoke_dispatches_to_invoke_sync(vlm_adapter): + """invoke forwards the request to invoke_sync.""" + from nexent.core.gateway.modality import VLMRequest + + request = VLMRequest(media_type="image", media_input=b"bytes", stream=False) + with patch.object(vlm_adapter, "invoke_sync", return_value="result") as mock_invoke_sync: + result = await vlm_adapter.invoke(request) + + assert result == "result" + mock_invoke_sync.assert_called_once_with(request) + + +def test_invoke_sync_without_prompt_merges_kwargs(vlm_adapter): + """invoke_sync skips system_prompt when empty and merges request kwargs.""" + from nexent.core.gateway.modality import VLMRequest + + request = VLMRequest(media_type="audio", media_input=b"bytes", stream=False, kwargs={"max_tokens": 5}) + with patch.object(vlm_adapter, "analyze_audio", return_value="audio-result") as mock_analyze: + result = vlm_adapter.invoke_sync(request) + + assert result == "audio-result" + mock_analyze.assert_called_once_with(b"bytes", stream=False, max_tokens=5) + + +@pytest.mark.asyncio +async def test_health_check_delegates_to_check_connectivity(vlm_adapter): + """health_check delegates to check_connectivity.""" + with patch.object(vlm_adapter, "check_connectivity", new_callable=AsyncMock, return_value=False) as mock_check: + result = await vlm_adapter.health_check() + + assert result is False + mock_check.assert_awaited_once() diff --git a/test/sdk/core/tools/test_analyze_audio_video_tool.py b/test/sdk/core/tools/test_analyze_audio_video_tool.py index 539ddb9a22..60f49b4de5 100644 --- a/test/sdk/core/tools/test_analyze_audio_video_tool.py +++ b/test/sdk/core/tools/test_analyze_audio_video_tool.py @@ -36,7 +36,7 @@ def _fake_get_prompt(template_type, language=None, **_): return {"system_prompt": "Analyze audio for {{ query }}"} monkeypatch.setattr(analyze_audio_tool, "get_prompt_template", _fake_get_prompt) - mock_vlm_model.analyze_audio.return_value = SimpleNamespace(content="audio result") + mock_vlm_model.invoke_sync.return_value = SimpleNamespace(content="audio result") tool = AnalyzeAudioTool( observer=observer_en, vlm_model=mock_vlm_model, @@ -47,10 +47,11 @@ def _fake_get_prompt(template_type, language=None, **_): assert result == "audio result" assert calls == [("analyze_audio", "en")] - mock_vlm_model.analyze_audio.assert_called_once() - call_kwargs = mock_vlm_model.analyze_audio.call_args.kwargs - assert hasattr(call_kwargs["audio_input"], "read") - assert call_kwargs["content_type"].startswith("audio/") + mock_vlm_model.invoke_sync.assert_called_once() + request = mock_vlm_model.invoke_sync.call_args.args[0] + assert request.media_type == "audio" + assert hasattr(request.media_input, "read") + assert request.kwargs["content_type"].startswith("audio/") def test_analyze_audio_schema_uses_single_url(): assert "audio_url" in AnalyzeAudioTool.inputs @@ -64,7 +65,7 @@ def test_analyze_audio_accepts_legacy_url_list(observer_en, mock_vlm_model, mock "get_prompt_template", lambda template_type, language=None, **_: {"system_prompt": "Analyze audio for {{ query }}"}, ) - mock_vlm_model.analyze_audio.return_value = SimpleNamespace(content="audio result") + mock_vlm_model.invoke_sync.return_value = SimpleNamespace(content="audio result") tool = AnalyzeAudioTool( observer=observer_en, vlm_model=mock_vlm_model, @@ -77,10 +78,9 @@ def test_analyze_audio_accepts_legacy_url_list(observer_en, mock_vlm_model, mock def test_analyze_audio_rejects_siliconflow_non_omni_model(observer_en, mock_storage_client): - vlm_model = SimpleNamespace( - model_id="Qwen/Qwen3-VL-32B-Instruct", - client_kwargs={"base_url": "https://api.siliconflow.cn/v1"}, - ) + vlm_model = MagicMock() + vlm_model.get_model_info.return_value = SimpleNamespace( + capabilities={"audio": False}) tool = AnalyzeAudioTool( observer=observer_en, vlm_model=vlm_model, @@ -101,7 +101,7 @@ def _fake_get_prompt(template_type, language=None, **_): return {"system_prompt": "Analyze video for {{ query }}"} monkeypatch.setattr(analyze_video_tool, "get_prompt_template", _fake_get_prompt) - mock_vlm_model.analyze_video.return_value = SimpleNamespace(content="video result") + mock_vlm_model.invoke_sync.return_value = SimpleNamespace(content="video result") tool = AnalyzeVideoTool( observer=observer_en, vlm_model=mock_vlm_model, @@ -112,10 +112,11 @@ def _fake_get_prompt(template_type, language=None, **_): assert result == "video result" assert calls == [("analyze_video", "en")] - mock_vlm_model.analyze_video.assert_called_once() - call_kwargs = mock_vlm_model.analyze_video.call_args.kwargs - assert hasattr(call_kwargs["video_input"], "read") - assert call_kwargs["content_type"].startswith("video/") + mock_vlm_model.invoke_sync.assert_called_once() + request = mock_vlm_model.invoke_sync.call_args.args[0] + assert request.media_type == "video" + assert hasattr(request.media_input, "read") + assert request.kwargs["content_type"].startswith("video/") def test_analyze_video_schema_uses_single_url(): assert "video_url" in AnalyzeVideoTool.inputs @@ -129,7 +130,7 @@ def test_analyze_video_accepts_legacy_url_list(observer_en, mock_vlm_model, mock "get_prompt_template", lambda template_type, language=None, **_: {"system_prompt": "Analyze video for {{ query }}"}, ) - mock_vlm_model.analyze_video.return_value = SimpleNamespace(content="video result") + mock_vlm_model.invoke_sync.return_value = SimpleNamespace(content="video result") tool = AnalyzeVideoTool( observer=observer_en, vlm_model=mock_vlm_model, diff --git a/test/sdk/core/tools/test_analyze_image_tool.py b/test/sdk/core/tools/test_analyze_image_tool.py index ca39df7c56..38dff08664 100644 --- a/test/sdk/core/tools/test_analyze_image_tool.py +++ b/test/sdk/core/tools/test_analyze_image_tool.py @@ -65,7 +65,7 @@ class TestAnalyzeImageTool: def test_forward_impl_success_with_multiple_images( self, tool, mock_vlm_model, mock_prompt_loader ): - mock_vlm_model.analyze_image.side_effect = [ + mock_vlm_model.invoke_sync.side_effect = [ SimpleNamespace(content="First image analysis"), SimpleNamespace(content="Second image analysis"), ] @@ -73,9 +73,9 @@ def test_forward_impl_success_with_multiple_images( result = tool._forward_impl([b"img1", b"img2"], "What is shown?") assert result == ["First image analysis", "Second image analysis"] - assert mock_vlm_model.analyze_image.call_count == 2 - for call in mock_vlm_model.analyze_image.call_args_list: - assert hasattr(call.kwargs["image_input"], "read") + assert mock_vlm_model.invoke_sync.call_count == 2 + for call in mock_vlm_model.invoke_sync.call_args_list: + assert hasattr(call.args[0].media_input, "read") assert mock_prompt_loader == [("analyze_image", "en")] def test_forward_impl_zh_observer_messages( @@ -86,7 +86,7 @@ def test_forward_impl_zh_observer_messages( vlm_model=mock_vlm_model, storage_client=mock_storage_client, ) - mock_vlm_model.analyze_image.return_value = SimpleNamespace( + mock_vlm_model.invoke_sync.return_value = SimpleNamespace( content="描述") result = tool._forward_impl([b"img"], "问题") @@ -111,7 +111,7 @@ def test_forward_impl_validates_inputs( def test_forward_impl_wraps_model_errors( self, tool, mock_vlm_model, mock_prompt_loader ): - mock_vlm_model.analyze_image.side_effect = Exception("model failed") + mock_vlm_model.invoke_sync.side_effect = Exception("model failed") with pytest.raises( Exception, @@ -119,7 +119,7 @@ def test_forward_impl_wraps_model_errors( ): tool._forward_impl([b"img"], "question") - mock_vlm_model.analyze_image.assert_called_once() + mock_vlm_model.invoke_sync.assert_called_once() class TestAnalyzeImageToolEdgeCases: @@ -159,7 +159,7 @@ def test_forward_impl_observer_none_uses_english(self, mock_vlm_model, mock_stor vlm_model=mock_vlm_model, storage_client=mock_storage_client, ) - mock_vlm_model.analyze_image.return_value = SimpleNamespace( + mock_vlm_model.invoke_sync.return_value = SimpleNamespace( content="Analysis result") result = tool._forward_impl([b"img"], "question") @@ -168,14 +168,14 @@ def test_forward_impl_observer_none_uses_english(self, mock_vlm_model, mock_stor def test_forward_impl_single_image_success(self, tool, mock_vlm_model, mock_prompt_loader): """Test successful analysis with a single image.""" - mock_vlm_model.analyze_image.return_value = SimpleNamespace( + mock_vlm_model.invoke_sync.return_value = SimpleNamespace( content="Single image description") result = tool._forward_impl( [b"single_image"], "What is in this image?") assert result == ["Single image description"] - mock_vlm_model.analyze_image.assert_called_once() + mock_vlm_model.invoke_sync.assert_called_once() def test_is_chinese_property_english(self, observer_en, mock_vlm_model, mock_storage_client): """Test that _is_chinese is False when observer lang is English.""" @@ -310,14 +310,14 @@ def test_observer_add_message_not_called_when_none(self, mock_vlm_model, mock_st vlm_model=mock_vlm_model, storage_client=mock_storage_client, ) - mock_vlm_model.analyze_image.return_value = SimpleNamespace( + mock_vlm_model.invoke_sync.return_value = SimpleNamespace( content="Result") # Should not raise any exception result = tool._forward_impl([b"img"], "question") assert result == ["Result"] - mock_vlm_model.analyze_image.assert_called_once() + mock_vlm_model.invoke_sync.assert_called_once() def test_tool_name_and_description(self, tool): """Test that tool name and description are set correctly."""