Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 46 additions & 11 deletions astrbot/core/provider/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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):
Expand All @@ -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()
Expand Down Expand Up @@ -583,51 +585,82 @@ 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", "<unknown>")
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)

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", "<unknown>"),
)
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
Expand Down Expand Up @@ -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:
Expand All @@ -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()
Expand Down
85 changes: 85 additions & 0 deletions tests/unit/test_provider_manager_config_validation.py
Original file line number Diff line number Diff line change
@@ -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 == {}
Loading