From b1654880394ea119788580b0824c980c16d2348b Mon Sep 17 00:00:00 2001 From: EterUltimate <1831303476@qq.com> Date: Sun, 26 Jul 2026 19:00:17 +0800 Subject: [PATCH] fix: reject providers missing adapter type --- astrbot/core/provider/manager.py | 57 ++++++++++--- ...test_provider_manager_config_validation.py | 85 +++++++++++++++++++ 2 files changed, 131 insertions(+), 11 deletions(-) create mode 100644 tests/unit/test_provider_manager_config_validation.py diff --git a/astrbot/core/provider/manager.py b/astrbot/core/provider/manager.py index 7c2ac7312c..39b6a24244 100644 --- a/astrbot/core/provider/manager.py +++ b/astrbot/core/provider/manager.py @@ -3,14 +3,13 @@ import os import traceback from collections.abc import Callable -from typing import Protocol, runtime_checkable +from typing import TYPE_CHECKING, Protocol, runtime_checkable from astrbot.core import astrbot_config, logger, sp from astrbot.core.astrbot_config_mgr import AstrBotConfigManager from astrbot.core.db import BaseDatabase from astrbot.core.utils.error_redaction import safe_error -from ..persona_mgr import PersonaManager from .entities import ProviderType from .provider import ( EmbeddingProvider, @@ -22,6 +21,9 @@ ) from .register import llm_tools, provider_cls_map +if TYPE_CHECKING: + from ..persona_mgr import PersonaManager + @runtime_checkable class HasInitialize(Protocol): @@ -33,7 +35,7 @@ def __init__( self, acm: AstrBotConfigManager, db_helper: BaseDatabase, - persona_mgr: PersonaManager, + persona_mgr: "PersonaManager", ) -> None: self.reload_lock = asyncio.Lock() self.resource_lock = asyncio.Lock() @@ -583,6 +585,29 @@ def _resolve_env_key_list(self, provider_config: dict) -> dict: provider_config["key"] = resolved_keys return provider_config + def _require_loadable_provider_config(self, provider_config: dict) -> None: + merged_config = self.get_merged_provider_config(provider_config) + if not merged_config.get("enable", True): + return + if merged_config.get("provider_type", "") == "agent_runner": + return + + provider_type = merged_config.get("type") + if isinstance(provider_type, str) and provider_type.strip(): + return + + provider_id = merged_config.get("id", "") + provider_source_id = provider_config.get("provider_source_id") + if provider_source_id: + raise ValueError( + f"Provider {provider_id} references provider source " + f"{provider_source_id}, but the merged config is missing a valid " + "'type' field" + ) + raise ValueError( + f"Provider {provider_id} config is missing a valid 'type' field" + ) + async def load_provider(self, provider_config: dict) -> None: # 如果 provider_source_id 存在且不为空,则从 provider_sources 中找到对应的配置并合并 provider_config = self.get_merged_provider_config(provider_config) @@ -590,44 +615,52 @@ async def load_provider(self, provider_config: dict) -> None: if provider_config.get("provider_type", "") == "chat_completion": provider_config = self._resolve_env_key_list(provider_config) - if not provider_config["enable"]: + if not provider_config.get("enable", True): logger.info(f"Provider {provider_config['id']} is disabled, skipping") return if provider_config.get("provider_type", "") == "agent_runner": return + provider_type = provider_config.get("type") + if not isinstance(provider_type, str) or not provider_type.strip(): + logger.error( + "Provider %s has no valid adapter type. Skipped.", + provider_config.get("id", ""), + ) + return + logger.info( "Loading model %s(%s) ...", - provider_config["type"], + provider_type, provider_config["id"], ) # 动态导入 try: - self.dynamic_import_provider(provider_config["type"]) + self.dynamic_import_provider(provider_type) except (ImportError, ModuleNotFoundError) as e: logger.critical( - f"Failed to load provider adapter {provider_config['type']}" + f"Failed to load provider adapter {provider_type}" f"({provider_config['id']}): {e}. A dependency may be missing.", exc_info=True, ) return except Exception as e: logger.critical( - f"Failed to load provider adapter {provider_config['type']}" + f"Failed to load provider adapter {provider_type}" f"({provider_config['id']}): {e}. Unknown cause.", exc_info=True, ) return - if provider_config["type"] not in provider_cls_map: + if provider_type not in provider_cls_map: logger.error( - f"Provider adapter not found: {provider_config['type']}({provider_config['id']}). Skipped.", + f"Provider adapter not found: {provider_type}({provider_config['id']}). Skipped.", exc_info=True, ) return - provider_metadata = provider_cls_map[provider_config["type"]] + provider_metadata = provider_cls_map[provider_type] try: # 按任务实例化提供商 cls_type = provider_metadata.cls_type @@ -868,6 +901,7 @@ async def update_provider(self, origin_provider_id: str, new_config: dict) -> No and provider.get("id", None) != origin_provider_id ): raise ValueError(f"Provider ID {npid} already exists") + self._require_loadable_provider_config(new_config) # update config for idx, provider in enumerate(config["provider"]): if provider.get("id", None) == origin_provider_id: @@ -889,6 +923,7 @@ async def create_provider(self, new_config: dict) -> None: for provider in config["provider"]: if provider.get("id", None) == npid: raise ValueError(f"Provider ID {npid} already exists") + self._require_loadable_provider_config(new_config) # add to config config["provider"].append(new_config) config.save_config() diff --git a/tests/unit/test_provider_manager_config_validation.py b/tests/unit/test_provider_manager_config_validation.py new file mode 100644 index 0000000000..733f3e6051 --- /dev/null +++ b/tests/unit/test_provider_manager_config_validation.py @@ -0,0 +1,85 @@ +from types import SimpleNamespace + +import pytest + +from astrbot.core.provider.manager import ProviderManager + + +class SaveCountingConfig(dict[str, object]): + def __init__(self, initial: dict[str, object]) -> None: + super().__init__(initial) + self.save_count = 0 + + def save_config(self) -> None: + self.save_count += 1 + + +class FakeConfigManager: + def __init__(self, config: SaveCountingConfig) -> None: + self.confs = {"default": config} + + @property + def default_conf(self) -> SaveCountingConfig: + return self.confs["default"] + + +def make_manager(config: SaveCountingConfig) -> ProviderManager: + return ProviderManager( + FakeConfigManager(config), + db_helper=SimpleNamespace(), + persona_mgr=SimpleNamespace(default_persona="default"), + ) + + +@pytest.mark.asyncio +async def test_create_provider_rejects_missing_merged_type_before_saving() -> None: + config = SaveCountingConfig( + { + "provider_sources": [ + { + "id": "deepseek_1", + "provider_type": "chat_completion", + "enable": True, + } + ], + "provider": [], + "provider_settings": {}, + } + ) + manager = make_manager(config) + + with pytest.raises(ValueError, match="missing a valid 'type' field"): + await manager.create_provider( + { + "id": "deepseek_1/deepseek-v4-flash", + "provider_source_id": "deepseek_1", + "provider_type": "chat_completion", + "model": "deepseek-v4-flash", + "enable": True, + } + ) + + assert config["provider"] == [] + assert config.save_count == 0 + + +@pytest.mark.asyncio +async def test_load_provider_skips_missing_type_without_keyerror() -> None: + config = SaveCountingConfig( + { + "provider_sources": [], + "provider": [], + "provider_settings": {}, + } + ) + manager = make_manager(config) + + await manager.load_provider( + { + "id": "missing-type", + "provider_type": "chat_completion", + "enable": True, + } + ) + + assert manager.inst_map == {}