diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index 327191db6..641a19bcf 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -1,6 +1,7 @@ """如需修改配置,请在 `data/cmd_config.json` 中修改或者在管理面板中可视化修改。""" import os +from typing import Any, TypedDict from astrbot.core.utils.astrbot_path import get_astrbot_data_path @@ -61,7 +62,8 @@ DEFAULT_CONFIG = { "ignore_bot_self_message": False, "ignore_at_all": False, }, - "provider": [], + "provider_sources": [], # provider sources + "provider": [], # models from provider_sources "provider_settings": { "enable": True, "default_provider_id": "", @@ -171,6 +173,22 @@ DEFAULT_CONFIG = { } +class ChatProviderTemplate(TypedDict): + id: str + provider_source_id: str + model: str + modalities: list + custom_extra_body: dict[str, Any] + + +CHAT_PROVIDER_TEMPLATE = { + "id": "", + "provide_source_id": "", + "model": "", + "modalities": [], + "custom_extra_body": {}, +} + """ AstrBot v3 时代的配置元数据,目前仅承担以下功能: @@ -844,6 +862,7 @@ CONFIG_METADATA_2 = { "metadata": { "provider": { "type": "list", + # provider sources templates "config_template": { "OpenAI": { "id": "openai", @@ -854,107 +873,10 @@ CONFIG_METADATA_2 = { "key": [], "api_base": "https://api.openai.com/v1", "timeout": 120, - "model_config": {"model": "gpt-4o-mini", "temperature": 0.4}, "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], - "hint": "也兼容所有与 OpenAI API 兼容的服务。", }, - "Azure OpenAI": { - "id": "azure", - "provider": "azure", - "type": "openai_chat_completion", - "provider_type": "chat_completion", - "enable": True, - "api_version": "2024-05-01-preview", - "key": [], - "api_base": "", - "timeout": 120, - "model_config": {"model": "gpt-4o-mini", "temperature": 0.4}, - "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], - }, - "xAI": { - "id": "xai", - "provider": "xai", - "type": "openai_chat_completion", - "provider_type": "chat_completion", - "enable": True, - "key": [], - "api_base": "https://api.x.ai/v1", - "timeout": 120, - "model_config": {"model": "grok-2-latest", "temperature": 0.4}, - "custom_headers": {}, - "custom_extra_body": {}, - "xai_native_search": False, - "modalities": ["text", "image", "tool_use"], - }, - "Anthropic": { - "hint": "注意Claude系列模型的温度调节范围为0到1.0,超出可能导致报错", - "id": "claude", - "provider": "anthropic", - "type": "anthropic_chat_completion", - "provider_type": "chat_completion", - "enable": True, - "key": [], - "api_base": "https://api.anthropic.com/v1", - "timeout": 120, - "model_config": { - "model": "claude-3-5-sonnet-latest", - "max_tokens": 4096, - "temperature": 0.2, - }, - "modalities": ["text", "image", "tool_use"], - }, - "Ollama": { - "hint": "启用前请确保已正确安装并运行 Ollama 服务端,Ollama默认不带鉴权,无需修改key", - "id": "ollama_default", - "provider": "ollama", - "type": "openai_chat_completion", - "provider_type": "chat_completion", - "enable": True, - "key": ["ollama"], # ollama 的 key 默认是 ollama - "api_base": "http://localhost:11434/v1", - "model_config": {"model": "llama3.1-8b", "temperature": 0.4}, - "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], - }, - "LM Studio": { - "id": "lm_studio", - "provider": "lm_studio", - "type": "openai_chat_completion", - "provider_type": "chat_completion", - "enable": True, - "key": ["lmstudio"], - "api_base": "http://localhost:1234/v1", - "model_config": { - "model": "llama-3.1-8b", - }, - "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], - }, - "Gemini(OpenAI兼容)": { - "id": "gemini_default", - "provider": "google", - "type": "openai_chat_completion", - "provider_type": "chat_completion", - "enable": True, - "key": [], - "api_base": "https://generativelanguage.googleapis.com/v1beta/openai/", - "timeout": 120, - "model_config": { - "model": "gemini-3-flash-preview", - "temperature": 0.4, - }, - "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], - }, - "Gemini": { - "id": "gemini_default", + "Google Gemini": { + "id": "google_gemini", "provider": "google", "type": "googlegenai_chat_completion", "provider_type": "chat_completion", @@ -962,10 +884,6 @@ CONFIG_METADATA_2 = { "key": [], "api_base": "https://generativelanguage.googleapis.com/", "timeout": 120, - "model_config": { - "model": "gemini-3-flash-preview", - "temperature": 0.4, - }, "gm_resp_image_modal": False, "gm_native_search": False, "gm_native_coderunner": False, @@ -977,10 +895,42 @@ CONFIG_METADATA_2 = { "dangerous_content": "BLOCK_MEDIUM_AND_ABOVE", }, "gm_thinking_config": {"budget": 0, "level": "HIGH"}, - "modalities": ["text", "image", "tool_use"], + }, + "Anthropic": { + "id": "anthropic", + "provider": "anthropic", + "type": "anthropic_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "api_base": "https://api.anthropic.com/v1", + "timeout": 120, + }, + "Moonshot": { + "id": "moonshot", + "provider": "moonshot", + "type": "openai_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "timeout": 120, + "api_base": "https://api.moonshot.cn/v1", + "custom_headers": {}, + }, + "xAI": { + "id": "xai", + "provider": "xai", + "type": "openai_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "api_base": "https://api.x.ai/v1", + "timeout": 120, + "custom_headers": {}, + "xai_native_search": False, }, "DeepSeek": { - "id": "deepseek_default", + "id": "deepseek", "provider": "deepseek", "type": "openai_chat_completion", "provider_type": "chat_completion", @@ -988,13 +938,75 @@ CONFIG_METADATA_2 = { "key": [], "api_base": "https://api.deepseek.com/v1", "timeout": 120, - "model_config": {"model": "deepseek-chat", "temperature": 0.4}, "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "tool_use"], + }, + "Zhipu": { + "id": "zhipu", + "provider": "zhipu", + "type": "zhipu_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "timeout": 120, + "api_base": "https://open.bigmodel.cn/api/paas/v4/", + "custom_headers": {}, + }, + "Azure OpenAI": { + "id": "azure_openai", + "provider": "azure", + "type": "openai_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "api_version": "2024-05-01-preview", + "key": [], + "api_base": "", + "timeout": 120, + "custom_headers": {}, + }, + "Ollama": { + "id": "ollama", + "provider": "ollama", + "type": "openai_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": ["ollama"], # ollama 的 key 默认是 ollama + "api_base": "http://127.0.0.1:11434/v1", + "custom_headers": {}, + }, + "LM Studio": { + "id": "lm_studio", + "provider": "lm_studio", + "type": "openai_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": ["lmstudio"], + "api_base": "http://127.0.0.1:1234/v1", + "custom_headers": {}, + }, + "ModelStack": { + "id": "modelstack", + "provider": "modelstack", + "type": "openai_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "api_base": "https://modelstack.app/v1", + "timeout": 120, + "custom_headers": {}, + }, + "Gemini_OpenAI_API": { + "id": "google_gemini_openai", + "provider": "google", + "type": "openai_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "api_base": "https://generativelanguage.googleapis.com/v1beta/openai/", + "timeout": 120, + "custom_headers": {}, }, "Groq": { - "id": "groq_default", + "id": "groq", "provider": "groq", "type": "groq_chat_completion", "provider_type": "chat_completion", @@ -1002,13 +1014,7 @@ CONFIG_METADATA_2 = { "key": [], "api_base": "https://api.groq.com/openai/v1", "timeout": 120, - "model_config": { - "model": "openai/gpt-oss-20b", - "temperature": 0.4, - }, "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "tool_use"], }, "302.AI": { "id": "302ai", @@ -1019,12 +1025,9 @@ CONFIG_METADATA_2 = { "key": [], "api_base": "https://api.302.ai/v1", "timeout": 120, - "model_config": {"model": "gpt-4.1-mini", "temperature": 0.4}, "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], }, - "硅基流动": { + "SiliconFlow": { "id": "siliconflow", "provider": "siliconflow", "type": "openai_chat_completion", @@ -1033,15 +1036,9 @@ CONFIG_METADATA_2 = { "key": [], "timeout": 120, "api_base": "https://api.siliconflow.cn/v1", - "model_config": { - "model": "deepseek-ai/DeepSeek-V3", - "temperature": 0.4, - }, "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], }, - "PPIO派欧云": { + "PPIO": { "id": "ppio", "provider": "ppio", "type": "openai_chat_completion", @@ -1050,14 +1047,9 @@ CONFIG_METADATA_2 = { "key": [], "api_base": "https://api.ppinfra.com/v3/openai", "timeout": 120, - "model_config": { - "model": "deepseek/deepseek-r1", - "temperature": 0.4, - }, "custom_headers": {}, - "custom_extra_body": {}, }, - "小马算力": { + "TokenPony": { "id": "tokenpony", "provider": "tokenpony", "type": "openai_chat_completion", @@ -1066,14 +1058,9 @@ CONFIG_METADATA_2 = { "key": [], "api_base": "https://api.tokenpony.cn/v1", "timeout": 120, - "model_config": { - "model": "kimi-k2-instruct-0905", - "temperature": 0.7, - }, "custom_headers": {}, - "custom_extra_body": {}, }, - "优云智算": { + "Compshare": { "id": "compshare", "provider": "compshare", "type": "openai_chat_completion", @@ -1082,42 +1069,18 @@ CONFIG_METADATA_2 = { "key": [], "api_base": "https://api.modelverse.cn/v1", "timeout": 120, - "model_config": { - "model": "moonshotai/Kimi-K2-Instruct", - }, "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], }, - "Kimi": { - "id": "moonshot", - "provider": "moonshot", + "ModelScope": { + "id": "modelscope", + "provider": "modelscope", "type": "openai_chat_completion", "provider_type": "chat_completion", "enable": True, "key": [], "timeout": 120, - "api_base": "https://api.moonshot.cn/v1", - "model_config": {"model": "moonshot-v1-8k", "temperature": 0.4}, + "api_base": "https://api-inference.modelscope.cn/v1", "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], - }, - "智谱 AI": { - "id": "zhipu_default", - "provider": "zhipu", - "type": "zhipu_chat_completion", - "provider_type": "chat_completion", - "enable": True, - "key": [], - "timeout": 120, - "api_base": "https://open.bigmodel.cn/api/paas/v4/", - "model_config": { - "model": "glm-4-flash", - }, - "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], }, "Dify": { "id": "dify_app_default", @@ -1132,7 +1095,6 @@ CONFIG_METADATA_2 = { "dify_query_input_key": "astrbot_text_query", "variables": {}, "timeout": 60, - "hint": "请确保你在 AstrBot 里设置的 APP 类型和 Dify 里面创建的应用的类型一致!", }, "Coze": { "id": "coze", @@ -1163,20 +1125,6 @@ CONFIG_METADATA_2 = { "variables": {}, "timeout": 60, }, - "ModelScope": { - "id": "modelscope", - "provider": "modelscope", - "type": "openai_chat_completion", - "provider_type": "chat_completion", - "enable": True, - "key": [], - "timeout": 120, - "api_base": "https://api-inference.modelscope.cn/v1", - "model_config": {"model": "Qwen/Qwen3-32B", "temperature": 0.4}, - "custom_headers": {}, - "custom_extra_body": {}, - "modalities": ["text", "image", "tool_use"], - }, "FastGPT": { "id": "fastgpt", "provider": "fastgpt", @@ -1200,7 +1148,6 @@ CONFIG_METADATA_2 = { "model": "whisper-1", }, "Whisper(Local)": { - "hint": "启用前请 pip 安装 openai-whisper 库(N卡用户大约下载 2GB,主要是 torch 和 cuda,CPU 用户大约下载 1 GB),并且安装 ffmpeg。否则将无法正常转文字。", "provider": "openai", "type": "openai_whisper_selfhost", "provider_type": "speech_to_text", @@ -1209,7 +1156,6 @@ CONFIG_METADATA_2 = { "model": "tiny", }, "SenseVoice(Local)": { - "hint": "启用前请 pip 安装 funasr、funasr_onnx、torchaudio、torch、modelscope、jieba 库(默认使用CPU,大约下载 1 GB),并且安装 ffmpeg。否则将无法正常转文字。", "type": "sensevoice_stt_selfhost", "provider": "sensevoice", "provider_type": "speech_to_text", @@ -1231,7 +1177,6 @@ CONFIG_METADATA_2 = { "timeout": "20", }, "Edge TTS": { - "hint": "提示:使用这个服务前需要安装有 ffmpeg,并且可以直接在终端调用 ffmpeg 指令。", "id": "edge_tts", "provider": "microsoft", "type": "edge_tts", @@ -1447,6 +1392,10 @@ CONFIG_METADATA_2 = { }, }, "items": { + "provider_source_id": { + "invisible": True, + "type": "string", + }, "xai_native_search": { "description": "启用原生搜索功能", "type": "bool", @@ -2015,7 +1964,6 @@ CONFIG_METADATA_2 = { "id": { "description": "ID", "type": "string", - "hint": "模型提供商名字。", }, "type": { "description": "模型提供商种类", @@ -2035,29 +1983,15 @@ CONFIG_METADATA_2 = { "description": "API Key", "type": "list", "items": {"type": "string"}, - "hint": "提供商 API Key。", }, "api_base": { "description": "API Base URL", "type": "string", - "hint": "API Base URL 请在模型提供商处获得。如出现 404 报错,尝试在地址末尾加上 /v1", }, - "model_config": { - "description": "模型配置", - "type": "object", - "items": { - "model": { - "description": "模型名称", - "type": "string", - "hint": "模型名称,如 gpt-4o-mini, deepseek-chat。", - }, - "max_tokens": { - "description": "模型最大输出长度(tokens)", - "type": "int", - }, - "temperature": {"description": "温度", "type": "float"}, - "top_p": {"description": "Top P值", "type": "float"}, - }, + "model": { + "description": "模型 ID", + "type": "string", + "hint": "模型名称,如 gpt-4o-mini, deepseek-chat。", }, "dify_api_key": { "description": "API Key", diff --git a/astrbot/core/core_lifecycle.py b/astrbot/core/core_lifecycle.py index 5a8672837..823fdf260 100644 --- a/astrbot/core/core_lifecycle.py +++ b/astrbot/core/core_lifecycle.py @@ -33,6 +33,7 @@ from astrbot.core.star.context import Context from astrbot.core.star.star_handler import EventType, star_handlers_registry, star_map from astrbot.core.umop_config_router import UmopConfigRouter from astrbot.core.updator import AstrBotUpdator +from astrbot.core.utils.llm_metadata import update_llm_metadata from astrbot.core.utils.migra_helper import migra from . import astrbot_config, html_renderer @@ -185,6 +186,8 @@ class AstrBotCoreLifecycle: # 初始化关闭控制面板的事件 self.dashboard_shutdown_event = asyncio.Event() + asyncio.create_task(update_llm_metadata()) + def _load(self) -> None: """加载事件总线和任务并初始化.""" # 创建一个异步任务来执行事件总线的 dispatch() 方法 diff --git a/astrbot/core/provider/manager.py b/astrbot/core/provider/manager.py index be8edc282..0dff2c8ed 100644 --- a/astrbot/core/provider/manager.py +++ b/astrbot/core/provider/manager.py @@ -1,4 +1,5 @@ import asyncio +import copy import traceback from typing import Protocol, runtime_checkable @@ -32,10 +33,12 @@ class ProviderManager: persona_mgr: PersonaManager, ): self.reload_lock = asyncio.Lock() + self.resource_lock = asyncio.Lock() self.persona_mgr = persona_mgr self.acm = acm config = acm.confs["default"] self.providers_config: list = config["provider"] + self.provider_sources_config: list = config.get("provider_sources", []) self.provider_settings: dict = config["provider_settings"] self.provider_stt_settings: dict = config.get("provider_stt_settings", {}) self.provider_tts_settings: dict = config.get("provider_tts_settings", {}) @@ -148,6 +151,7 @@ class ProviderManager: """ provider = None + provider_id = None if umo: provider_id = sp.get( f"provider_perf_{provider_type.value}", @@ -185,6 +189,12 @@ class ProviderManager: ) else: raise ValueError(f"Unknown provider type: {provider_type}") + + if not provider and provider_id: + logger.warning( + f"没有找到 ID 为 {provider_id} 的提供商,这可能是由于您修改了提供商(模型)ID 导致的。" + ) + return provider async def initialize(self): @@ -251,7 +261,136 @@ class ProviderManager: # 初始化 MCP Client 连接 asyncio.create_task(self.llm_tools.init_mcp_clients(), name="init_mcp_clients") + def dynamic_import_provider(self, type: str): + """动态导入提供商适配器模块 + + Args: + type (str): 提供商请求类型。 + + Raises: + ImportError: 如果提供商类型未知或无法导入对应模块,则抛出异常。 + """ + match type: + case "openai_chat_completion": + from .sources.openai_source import ( + ProviderOpenAIOfficial as ProviderOpenAIOfficial, + ) + case "zhipu_chat_completion": + from .sources.zhipu_source import ProviderZhipu as ProviderZhipu + case "groq_chat_completion": + from .sources.groq_source import ProviderGroq as ProviderGroq + case "anthropic_chat_completion": + from .sources.anthropic_source import ( + ProviderAnthropic as ProviderAnthropic, + ) + case "googlegenai_chat_completion": + from .sources.gemini_source import ( + ProviderGoogleGenAI as ProviderGoogleGenAI, + ) + case "sensevoice_stt_selfhost": + from .sources.sensevoice_selfhosted_source import ( + ProviderSenseVoiceSTTSelfHost as ProviderSenseVoiceSTTSelfHost, + ) + case "openai_whisper_api": + from .sources.whisper_api_source import ( + ProviderOpenAIWhisperAPI as ProviderOpenAIWhisperAPI, + ) + case "openai_whisper_selfhost": + from .sources.whisper_selfhosted_source import ( + ProviderOpenAIWhisperSelfHost as ProviderOpenAIWhisperSelfHost, + ) + case "xinference_stt": + from .sources.xinference_stt_provider import ( + ProviderXinferenceSTT as ProviderXinferenceSTT, + ) + case "openai_tts_api": + from .sources.openai_tts_api_source import ( + ProviderOpenAITTSAPI as ProviderOpenAITTSAPI, + ) + case "edge_tts": + from .sources.edge_tts_source import ( + ProviderEdgeTTS as ProviderEdgeTTS, + ) + case "gsv_tts_selfhost": + from .sources.gsv_selfhosted_source import ( + ProviderGSVTTS as ProviderGSVTTS, + ) + case "gsvi_tts_api": + from .sources.gsvi_tts_source import ( + ProviderGSVITTS as ProviderGSVITTS, + ) + case "fishaudio_tts_api": + from .sources.fishaudio_tts_api_source import ( + ProviderFishAudioTTSAPI as ProviderFishAudioTTSAPI, + ) + case "dashscope_tts": + from .sources.dashscope_tts import ( + ProviderDashscopeTTSAPI as ProviderDashscopeTTSAPI, + ) + case "azure_tts": + from .sources.azure_tts_source import ( + AzureTTSProvider as AzureTTSProvider, + ) + case "minimax_tts_api": + from .sources.minimax_tts_api_source import ( + ProviderMiniMaxTTSAPI as ProviderMiniMaxTTSAPI, + ) + case "volcengine_tts": + from .sources.volcengine_tts import ( + ProviderVolcengineTTS as ProviderVolcengineTTS, + ) + case "gemini_tts": + from .sources.gemini_tts_source import ( + ProviderGeminiTTSAPI as ProviderGeminiTTSAPI, + ) + case "openai_embedding": + from .sources.openai_embedding_source import ( + OpenAIEmbeddingProvider as OpenAIEmbeddingProvider, + ) + case "gemini_embedding": + from .sources.gemini_embedding_source import ( + GeminiEmbeddingProvider as GeminiEmbeddingProvider, + ) + case "vllm_rerank": + from .sources.vllm_rerank_source import ( + VLLMRerankProvider as VLLMRerankProvider, + ) + case "xinference_rerank": + from .sources.xinference_rerank_source import ( + XinferenceRerankProvider as XinferenceRerankProvider, + ) + case "bailian_rerank": + from .sources.bailian_rerank_source import ( + BailianRerankProvider as BailianRerankProvider, + ) + + def get_merged_provider_config(self, provider_config: dict) -> dict: + """获取 provider 配置和 provider_source 配置合并后的结果 + + Returns: + dict: 合并后的 provider 配置,key 为 provider id,value 为合并后的配置字典 + """ + pc = copy.deepcopy(provider_config) + provider_source_id = pc.get("provider_source_id", "") + if provider_source_id: + provider_source = None + for ps in self.provider_sources_config: + if ps.get("id") == provider_source_id: + provider_source = ps + break + + if provider_source: + # 合并配置,provider 的配置优先级更高 + merged_config = {**provider_source, **pc} + # 保持 id 为 provider 的 id,而不是 source 的 id + merged_config["id"] = pc["id"] + pc = merged_config + return pc + async def load_provider(self, provider_config: dict): + # 如果 provider_source_id 存在且不为空,则从 provider_sources 中找到对应的配置并合并 + provider_config = self.get_merged_provider_config(provider_config) + if not provider_config["enable"]: logger.info(f"Provider {provider_config['id']} is disabled, skipping") return @@ -264,99 +403,7 @@ class ProviderManager: # 动态导入 try: - match provider_config["type"]: - case "openai_chat_completion": - from .sources.openai_source import ( - ProviderOpenAIOfficial as ProviderOpenAIOfficial, - ) - case "zhipu_chat_completion": - from .sources.zhipu_source import ProviderZhipu as ProviderZhipu - case "groq_chat_completion": - from .sources.groq_source import ProviderGroq as ProviderGroq - case "anthropic_chat_completion": - from .sources.anthropic_source import ( - ProviderAnthropic as ProviderAnthropic, - ) - case "googlegenai_chat_completion": - from .sources.gemini_source import ( - ProviderGoogleGenAI as ProviderGoogleGenAI, - ) - case "sensevoice_stt_selfhost": - from .sources.sensevoice_selfhosted_source import ( - ProviderSenseVoiceSTTSelfHost as ProviderSenseVoiceSTTSelfHost, - ) - case "openai_whisper_api": - from .sources.whisper_api_source import ( - ProviderOpenAIWhisperAPI as ProviderOpenAIWhisperAPI, - ) - case "openai_whisper_selfhost": - from .sources.whisper_selfhosted_source import ( - ProviderOpenAIWhisperSelfHost as ProviderOpenAIWhisperSelfHost, - ) - case "xinference_stt": - from .sources.xinference_stt_provider import ( - ProviderXinferenceSTT as ProviderXinferenceSTT, - ) - case "openai_tts_api": - from .sources.openai_tts_api_source import ( - ProviderOpenAITTSAPI as ProviderOpenAITTSAPI, - ) - case "edge_tts": - from .sources.edge_tts_source import ( - ProviderEdgeTTS as ProviderEdgeTTS, - ) - case "gsv_tts_selfhost": - from .sources.gsv_selfhosted_source import ( - ProviderGSVTTS as ProviderGSVTTS, - ) - case "gsvi_tts_api": - from .sources.gsvi_tts_source import ( - ProviderGSVITTS as ProviderGSVITTS, - ) - case "fishaudio_tts_api": - from .sources.fishaudio_tts_api_source import ( - ProviderFishAudioTTSAPI as ProviderFishAudioTTSAPI, - ) - case "dashscope_tts": - from .sources.dashscope_tts import ( - ProviderDashscopeTTSAPI as ProviderDashscopeTTSAPI, - ) - case "azure_tts": - from .sources.azure_tts_source import ( - AzureTTSProvider as AzureTTSProvider, - ) - case "minimax_tts_api": - from .sources.minimax_tts_api_source import ( - ProviderMiniMaxTTSAPI as ProviderMiniMaxTTSAPI, - ) - case "volcengine_tts": - from .sources.volcengine_tts import ( - ProviderVolcengineTTS as ProviderVolcengineTTS, - ) - case "gemini_tts": - from .sources.gemini_tts_source import ( - ProviderGeminiTTSAPI as ProviderGeminiTTSAPI, - ) - case "openai_embedding": - from .sources.openai_embedding_source import ( - OpenAIEmbeddingProvider as OpenAIEmbeddingProvider, - ) - case "gemini_embedding": - from .sources.gemini_embedding_source import ( - GeminiEmbeddingProvider as GeminiEmbeddingProvider, - ) - case "vllm_rerank": - from .sources.vllm_rerank_source import ( - VLLMRerankProvider as VLLMRerankProvider, - ) - case "xinference_rerank": - from .sources.xinference_rerank_source import ( - XinferenceRerankProvider as XinferenceRerankProvider, - ) - case "bailian_rerank": - from .sources.bailian_rerank_source import ( - BailianRerankProvider as BailianRerankProvider, - ) + self.dynamic_import_provider(provider_config["type"]) except (ImportError, ModuleNotFoundError) as e: logger.critical( f"加载 {provider_config['type']}({provider_config['id']}) 提供商适配器失败:{e}。可能是因为有未安装的依赖。", @@ -499,6 +546,7 @@ class ProviderManager: # 和配置文件保持同步 self.providers_config = astrbot_config["provider"] + self.provider_sources_config = astrbot_config.get("provider_sources", []) config_ids = [provider["id"] for provider in self.providers_config] logger.info(f"providers in user's config: {config_ids}") for key in list(self.inst_map.keys()): @@ -570,6 +618,68 @@ class ProviderManager: ) del self.inst_map[provider_id] + async def delete_provider( + self, provider_id: str | None = None, provider_source_id: str | None = None + ): + """Delete provider and/or provider source from config and terminate the instances. Config will be saved after deletion.""" + async with self.resource_lock: + # delete from config + target_prov_ids = [] + if provider_id: + target_prov_ids.append(provider_id) + else: + for prov in self.providers_config: + if prov.get("provider_source_id") == provider_source_id: + target_prov_ids.append(prov.get("id")) + config = self.acm.default_conf + for tpid in target_prov_ids: + await self.terminate_provider(tpid) + config["provider"] = [ + prov for prov in config["provider"] if prov.get("id") != tpid + ] + config.save_config() + logger.info(f"Provider {target_prov_ids} 已从配置中删除。") + + async def update_provider(self, origin_provider_id: str, new_config: dict): + """Update provider config and reload the instance. Config will be saved after update.""" + async with self.resource_lock: + npid = new_config.get("id", None) + if not npid: + raise ValueError("New provider config must have an 'id' field") + config = self.acm.default_conf + for provider in config["provider"]: + if ( + provider.get("id", None) == npid + and provider.get("id", None) != origin_provider_id + ): + raise ValueError(f"Provider ID {npid} already exists") + # update config + for idx, provider in enumerate(config["provider"]): + if provider.get("id", None) == origin_provider_id: + config["provider"][idx] = new_config + break + else: + raise ValueError(f"Provider ID {origin_provider_id} not found") + config.save_config() + # reload instance + await self.reload(new_config) + + async def create_provider(self, new_config: dict): + """Add new provider config and load the instance. Config will be saved after addition.""" + async with self.resource_lock: + npid = new_config.get("id", None) + if not npid: + raise ValueError("New provider config must have an 'id' field") + config = self.acm.default_conf + for provider in config["provider"]: + if provider.get("id", None) == npid: + raise ValueError(f"Provider ID {npid} already exists") + # add to config + config["provider"].append(new_config) + config.save_config() + # load instance + await self.load_provider(new_config) + async def terminate(self): for provider_inst in self.provider_insts: if hasattr(provider_inst, "terminate"): diff --git a/astrbot/core/provider/sources/anthropic_source.py b/astrbot/core/provider/sources/anthropic_source.py index 7e33f40d9..0ff61e393 100644 --- a/astrbot/core/provider/sources/anthropic_source.py +++ b/astrbot/core/provider/sources/anthropic_source.py @@ -47,7 +47,7 @@ class ProviderAnthropic(Provider): base_url=self.base_url, ) - self.set_model(provider_config["model_config"]["model"]) + self.set_model(provider_config.get("model", "unknown")) def _prepare_payload(self, messages: list[dict]): """准备 Anthropic API 的请求 payload @@ -130,7 +130,11 @@ class ProviderAnthropic(Provider): if tool_list := tools.get_func_desc_anthropic_style(): payloads["tools"] = tool_list - completion = await self.client.messages.create(**payloads, stream=False) + extra_body = self.provider_config.get("custom_extra_body", {}) + + completion = await self.client.messages.create( + **payloads, stream=False, extra_body=extra_body + ) assert isinstance(completion, Message) logger.debug(f"completion: {completion}") @@ -173,11 +177,13 @@ class ProviderAnthropic(Provider): # 用于累积最终结果 final_text = "" final_tool_calls = [] - id = None usage = TokenUsage() + extra_body = self.provider_config.get("custom_extra_body", {}) - async with self.client.messages.stream(**payloads) as stream: + async with self.client.messages.stream( + **payloads, extra_body=extra_body + ) as stream: assert isinstance(stream, anthropic.AsyncMessageStream) async for event in stream: if event.type == "message_start": @@ -318,10 +324,9 @@ class ProviderAnthropic(Provider): system_prompt, new_messages = self._prepare_payload(context_query) - model_config = self.provider_config.get("model_config", {}) - model_config["model"] = model or self.get_model() + model = model or self.get_model() - payloads = {"messages": new_messages, **model_config} + payloads = {"messages": new_messages, "model": model} # Anthropic has a different way of handling system prompts if system_prompt: @@ -331,7 +336,6 @@ class ProviderAnthropic(Provider): try: llm_response = await self._query(payloads, func_tool) except Exception as e: - # logger.error(f"发生了错误。Provider 配置如下: {model_config}") raise e return llm_response @@ -373,10 +377,9 @@ class ProviderAnthropic(Provider): system_prompt, new_messages = self._prepare_payload(context_query) - model_config = self.provider_config.get("model_config", {}) - model_config["model"] = model or self.get_model() + model = model or self.get_model() - payloads = {"messages": new_messages, **model_config} + payloads = {"messages": new_messages, "model": model} # Anthropic has a different way of handling system prompts if system_prompt: diff --git a/astrbot/core/provider/sources/gemini_source.py b/astrbot/core/provider/sources/gemini_source.py index 5a56170a5..f98fd0e8f 100644 --- a/astrbot/core/provider/sources/gemini_source.py +++ b/astrbot/core/provider/sources/gemini_source.py @@ -68,7 +68,7 @@ class ProviderGoogleGenAI(Provider): self.api_base = self.api_base[:-1] self._init_client() - self.set_model(provider_config["model_config"]["model"]) + self.set_model(provider_config.get("model", "unknown")) self._init_safety_settings() def _init_client(self) -> None: @@ -689,10 +689,9 @@ class ProviderGoogleGenAI(Provider): for tcr in tool_calls_result: context_query.extend(tcr.to_openai_messages()) - model_config = self.provider_config.get("model_config", {}) - model_config["model"] = model or self.get_model() + model = model or self.get_model() - payloads = {"messages": context_query, **model_config} + payloads = {"messages": context_query, "model": model} retry = 10 keys = self.api_keys.copy() @@ -742,10 +741,9 @@ class ProviderGoogleGenAI(Provider): for tcr in tool_calls_result: context_query.extend(tcr.to_openai_messages()) - model_config = self.provider_config.get("model_config", {}) - model_config["model"] = model or self.get_model() + model = model or self.get_model() - payloads = {"messages": context_query, **model_config} + payloads = {"messages": context_query, "model": model} retry = 10 keys = self.api_keys.copy() diff --git a/astrbot/core/provider/sources/openai_source.py b/astrbot/core/provider/sources/openai_source.py index 4aeacf672..a716d0a5a 100644 --- a/astrbot/core/provider/sources/openai_source.py +++ b/astrbot/core/provider/sources/openai_source.py @@ -69,8 +69,7 @@ class ProviderOpenAIOfficial(Provider): self.client.chat.completions.create, ).parameters.keys() - model_config = provider_config.get("model_config", {}) - model = model_config.get("model", "unknown") + model = provider_config.get("model", "unknown") self.set_model(model) self.reasoning_key = "reasoning_content" @@ -375,10 +374,9 @@ class ProviderOpenAIOfficial(Provider): for tcr in tool_calls_result: context_query.extend(tcr.to_openai_messages()) - model_config = self.provider_config.get("model_config", {}) - model_config["model"] = model or self.get_model() + model = model or self.get_model() - payloads = {"messages": context_query, **model_config} + payloads = {"messages": context_query, "model": model} # xAI origin search tool inject self._maybe_inject_xai_search(payloads, **kwargs) diff --git a/astrbot/core/star/context.py b/astrbot/core/star/context.py index 2561762f1..3b666b002 100644 --- a/astrbot/core/star/context.py +++ b/astrbot/core/star/context.py @@ -267,6 +267,10 @@ class Context: ): """通过 ID 获取对应的 LLM Provider。""" prov = self.provider_manager.inst_map.get(provider_id) + if provider_id and not prov: + logger.warning( + f"没有找到 ID 为 {provider_id} 的提供商,这可能是由于您修改了提供商(模型)ID 导致的。" + ) return prov def get_all_providers(self) -> list[Provider]: @@ -296,10 +300,6 @@ class Context: provider_type=ProviderType.CHAT_COMPLETION, umo=umo, ) - if prov is None: - raise ProviderNotFoundError( - "provider not found, please choose provider first" - ) if not isinstance(prov, Provider): raise ValueError("返回的 Provider 不是 Provider 类型") return prov diff --git a/astrbot/core/utils/llm_metadata.py b/astrbot/core/utils/llm_metadata.py new file mode 100644 index 000000000..540c1efd9 --- /dev/null +++ b/astrbot/core/utils/llm_metadata.py @@ -0,0 +1,63 @@ +from typing import Literal, TypedDict + +import aiohttp + +from astrbot.core import logger + + +class LLMModalities(TypedDict): + input: list[Literal["text", "image", "audio", "video"]] + output: list[Literal["text", "image", "audio", "video"]] + + +class LLMLimit(TypedDict): + context: int + output: int + + +class LLMMetadata(TypedDict): + id: str + reasoning: bool + tool_call: bool + knowledge: str + release_date: str + modalities: LLMModalities + open_weights: bool + limit: LLMLimit + + +LLM_METADATAS: dict[str, LLMMetadata] = {} + + +async def update_llm_metadata(): + url = "https://models.dev/api.json" + try: + async with aiohttp.ClientSession() as session: + async with session.get(url) as response: + data = await response.json() + global LLM_METADATAS + models = {} + for info in data.values(): + for model in info.get("models", {}).values(): + model_id = model.get("id") + if not model_id: + continue + models[model_id] = LLMMetadata( + id=model_id, + reasoning=model.get("reasoning", False), + tool_call=model.get("tool_call", False), + knowledge=model.get("knowledge", "none"), + release_date=model.get("release_date", ""), + modalities=model.get( + "modalities", {"input": [], "output": []} + ), + open_weights=model.get("open_weights", False), + limit=model.get("limit", {"context": 0, "output": 0}), + ) + # Replace the global cache in-place so references remain valid + LLM_METADATAS.clear() + LLM_METADATAS.update(models) + logger.info(f"Successfully fetched metadata for {len(models)} LLMs.") + except Exception as e: + logger.error(f"Failed to fetch LLM metadata: {e}") + return diff --git a/astrbot/core/utils/migra_helper.py b/astrbot/core/utils/migra_helper.py index 5642d606e..b8ff677e1 100644 --- a/astrbot/core/utils/migra_helper.py +++ b/astrbot/core/utils/migra_helper.py @@ -32,6 +32,92 @@ def _migra_agent_runner_configs(conf: AstrBotConfig, ids_map: dict) -> None: logger.error(traceback.format_exc()) +def _migra_provider_to_source_structure(conf: AstrBotConfig) -> None: + """ + Migrate old provider structure to new provider-source separation. + Provider only keeps: id, provider_source_id, model, modalities, custom_extra_body + All other fields move to provider_sources. + """ + providers = conf.get("provider", []) + provider_sources = conf.get("provider_sources", []) + + # Track if any migration happened + migrated = False + + # Provider-only fields that should stay in provider + provider_only_fields = { + "id", + "provider_source_id", + "model", + "modalities", + "custom_extra_body", + "enable", + } + + # Fields that should not go to source + source_exclude_fields = provider_only_fields | {"model_config"} + + for provider in providers: + # Skip if already has provider_source_id + if provider.get("provider_source_id"): + continue + + # Skip non-chat-completion types (they don't need source separation) + provider_type = provider.get("provider_type", "") + if provider_type != "chat_completion": + # For old types without provider_type, check type field + old_type = provider.get("type", "") + if "chat_completion" not in old_type: + continue + + migrated = True + logger.info(f"Migrating provider {provider.get('id')} to new structure") + + # Extract source fields from provider + source_fields = {} + for key, value in list(provider.items()): + if key not in source_exclude_fields: + source_fields[key] = value + + # Create new provider_source + source_id = provider.get("id", "") + "_source" + new_source = {"id": source_id, **source_fields} + + # Update provider to only keep necessary fields + provider["provider_source_id"] = source_id + + # Extract model from model_config if exists + if "model_config" in provider and isinstance(provider["model_config"], dict): + model_config = provider["model_config"] + provider["model"] = model_config.get("model", "") + + # Put other model_config fields into custom_extra_body + extra_body_fields = {k: v for k, v in model_config.items() if k != "model"} + if extra_body_fields: + if "custom_extra_body" not in provider: + provider["custom_extra_body"] = {} + provider["custom_extra_body"].update(extra_body_fields) + + # Initialize new fields if not present + if "modalities" not in provider: + provider["modalities"] = [] + if "custom_extra_body" not in provider: + provider["custom_extra_body"] = {} + + # Remove fields that should be in source + keys_to_remove = [k for k in provider.keys() if k not in provider_only_fields] + for key in keys_to_remove: + del provider[key] + + # Add source to provider_sources + provider_sources.append(new_source) + + if migrated: + conf["provider_sources"] = provider_sources + conf.save_config() + logger.info("Provider-source structure migration completed") + + async def migra( db, astrbot_config_mgr, umop_config_router, acm: AstrBotConfigManager ) -> None: @@ -71,3 +157,10 @@ async def migra( for conf in acm.confs.values(): _migra_agent_runner_configs(conf, ids_map) + + # Migrate providers to new structure: extract source fields to provider_sources + try: + _migra_provider_to_source_structure(astrbot_config) + except Exception as e: + logger.error(f"Migration for provider-source structure failed: {e!s}") + logger.error(traceback.format_exc()) diff --git a/astrbot/dashboard/routes/config.py b/astrbot/dashboard/routes/config.py index 0edbe8377..abfff529b 100644 --- a/astrbot/dashboard/routes/config.py +++ b/astrbot/dashboard/routes/config.py @@ -6,7 +6,7 @@ from typing import Any from quart import request -from astrbot.core import file_token_service, logger +from astrbot.core import astrbot_config, file_token_service, logger from astrbot.core.config.astrbot_config import AstrBotConfig from astrbot.core.config.default import ( CONFIG_METADATA_2, @@ -21,6 +21,7 @@ from astrbot.core.platform.register import platform_cls_map, platform_registry from astrbot.core.provider import Provider from astrbot.core.provider.register import provider_registry from astrbot.core.star.star import star_registry +from astrbot.core.utils.llm_metadata import LLM_METADATAS from astrbot.core.utils.webhook_utils import ensure_platform_webhook_config from .route import Response, Route, RouteContext @@ -179,13 +180,149 @@ class ConfigRoute(Route): "/config/provider/new": ("POST", self.post_new_provider), "/config/provider/update": ("POST", self.post_update_provider), "/config/provider/delete": ("POST", self.post_delete_provider), + "/config/provider/template": ("GET", self.get_provider_template), "/config/provider/check_one": ("GET", self.check_one_provider_status), "/config/provider/list": ("GET", self.get_provider_config_list), "/config/provider/model_list": ("GET", self.get_provider_model_list), "/config/provider/get_embedding_dim": ("POST", self.get_embedding_dim), + "/config/provider_sources//models": ( + "GET", + self.get_provider_source_models, + ), + "/config/provider_sources//update": ( + "POST", + self.update_provider_source, + ), + "/config/provider_sources//delete": ( + "POST", + self.delete_provider_source, + ), } self.register_routes() + async def delete_provider_source(self, provider_source_id: str): + """删除 provider_source,并更新关联的 providers""" + + provider_sources = self.config.get("provider_sources", []) + target_idx = next( + ( + i + for i, ps in enumerate(provider_sources) + if ps.get("id") == provider_source_id + ), + -1, + ) + + if target_idx == -1: + return Response().error("未找到对应的 provider source").__dict__ + + # 删除 provider_source + del provider_sources[target_idx] + + # 写回配置 + self.config["provider_sources"] = provider_sources + + # 删除引用了该 provider_source 的 providers + await self.core_lifecycle.provider_manager.delete_provider( + provider_source_id=provider_source_id + ) + + try: + save_config(self.config, self.config, is_core=True) + except Exception as e: + logger.error(traceback.format_exc()) + return Response().error(str(e)).__dict__ + + return Response().ok(message="删除 provider source 成功").__dict__ + + async def update_provider_source(self, provider_source_id: str): + """更新或新增 provider_source,并重载关联的 providers""" + + post_data = await request.json + if not post_data: + return Response().error("缺少配置数据").__dict__ + + new_source_config = post_data.get("config") or post_data + original_id = provider_source_id + + if not isinstance(new_source_config, dict): + return Response().error("缺少或错误的配置数据").__dict__ + + # 确保配置中有 id 字段 + if not new_source_config.get("id"): + new_source_config["id"] = original_id + + provider_sources = self.config.get("provider_sources", []) + + for ps in provider_sources: + if ps.get("id") == new_source_config["id"] and ps.get("id") != original_id: + return ( + Response() + .error( + f"Provider source ID '{new_source_config['id']}' exists already, please try another ID.", + ) + .__dict__ + ) + + # 查找旧的 provider_source,若不存在则追加为新配置 + target_idx = next( + (i for i, ps in enumerate(provider_sources) if ps.get("id") == original_id), + -1, + ) + + old_id = original_id + if target_idx == -1: + provider_sources.append(new_source_config) + else: + old_id = provider_sources[target_idx].get("id") + provider_sources[target_idx] = new_source_config + + # 更新引用了该 provider_source 的 providers + affected_providers = [] + for provider in self.config.get("provider", []): + if provider.get("provider_source_id") == old_id: + provider["provider_source_id"] = new_source_config["id"] + affected_providers.append(provider) + + # 写回配置 + self.config["provider_sources"] = provider_sources + + try: + save_config(self.config, self.config, is_core=True) + except Exception as e: + logger.error(traceback.format_exc()) + return Response().error(str(e)).__dict__ + + # 重载受影响的 providers,使新的 source 配置生效 + reload_errors = [] + prov_mgr = self.core_lifecycle.provider_manager + for provider in affected_providers: + try: + await prov_mgr.reload(provider) + except Exception as e: + logger.error(traceback.format_exc()) + reload_errors.append(f"{provider.get('id')}: {e}") + + if reload_errors: + return ( + Response() + .error("更新成功,但部分提供商重载失败: " + ", ".join(reload_errors)) + .__dict__ + ) + + return Response().ok(message="更新 provider source 成功").__dict__ + + async def get_provider_template(self): + config_schema = { + "provider": CONFIG_METADATA_2["provider_group"]["metadata"]["provider"] + } + data = { + "config_schema": config_schema, + "providers": astrbot_config["provider"], + "provider_sources": astrbot_config["provider_sources"], + } + return Response().ok(data=data).__dict__ + async def get_uc_table(self): """获取 UMOP 配置路由表""" return Response().ok({"routing": self.ucr.umop_to_conf_id}).__dict__ @@ -433,9 +570,25 @@ class ConfigRoute(Route): return Response().error("缺少参数 provider_type").__dict__ provider_type_ls = provider_type.split(",") provider_list = [] - astrbot_config = self.core_lifecycle.astrbot_config - for provider in astrbot_config["provider"]: - if provider.get("provider_type", None) in provider_type_ls: + ps = self.core_lifecycle.provider_manager.providers_config + p_source_pt = { + psrc["id"]: psrc["provider_type"] + for psrc in self.core_lifecycle.provider_manager.provider_sources_config + } + for provider in ps: + ps_id = provider.get("provider_source_id", None) + if ( + ps_id + and ps_id in p_source_pt + and p_source_pt[ps_id] in provider_type_ls + ): + # chat + prov = self.core_lifecycle.provider_manager.get_merged_provider_config( + provider + ) + provider_list.append(prov) + elif not ps_id and provider.get("provider_type", None) in provider_type_ls: + # agent runner, embedding, etc provider_list.append(provider) return Response().ok(provider_list).__dict__ @@ -458,9 +611,18 @@ class ConfigRoute(Route): try: models = await provider.get_models() + models = models or [] + + metadata_map = {} + for model_id in models: + meta = LLM_METADATAS.get(model_id) + if meta: + metadata_map[model_id] = meta + ret = { "models": models, "provider_id": provider_id, + "model_metadata": metadata_map, } return Response().ok(ret).__dict__ except Exception as e: @@ -522,6 +684,100 @@ class ConfigRoute(Route): logger.error(traceback.format_exc()) return Response().error(f"获取嵌入维度失败: {e!s}").__dict__ + async def get_provider_source_models(self, provider_source_id: str): + """获取指定 provider_source 支持的模型列表 + + 本质上会临时初始化一个 Provider 实例,调用 get_models() 获取模型列表,然后销毁实例 + """ + try: + from astrbot.core.provider.register import provider_cls_map + + # 从配置中查找对应的 provider_source + provider_sources = self.config.get("provider_sources", []) + provider_source = None + for ps in provider_sources: + if ps.get("id") == provider_source_id: + provider_source = ps + break + + if not provider_source: + return ( + Response() + .error(f"未找到 ID 为 {provider_source_id} 的 provider_source") + .__dict__ + ) + + # 获取 provider 类型 + provider_type = provider_source.get("type", None) + if not provider_type: + return Response().error("provider_source 缺少 type 字段").__dict__ + + try: + self.core_lifecycle.provider_manager.dynamic_import_provider( + provider_type + ) + except ImportError as e: + logger.error(traceback.format_exc()) + return Response().error(f"动态导入提供商适配器失败: {e!s}").__dict__ + + # 获取对应的 provider 类 + if provider_type not in provider_cls_map: + return ( + Response() + .error(f"未找到适用于 {provider_type} 的提供商适配器") + .__dict__ + ) + + provider_metadata = provider_cls_map[provider_type] + cls_type = provider_metadata.cls_type + + if not cls_type: + return Response().error(f"无法找到 {provider_type} 的类").__dict__ + + # 检查是否是 Provider 类型 + if not issubclass(cls_type, Provider): + return ( + Response() + .error(f"提供商 {provider_type} 不支持获取模型列表") + .__dict__ + ) + + # 临时实例化 provider + inst = cls_type(provider_source, {}) + + # 如果有 initialize 方法,调用它 + init_fn = getattr(inst, "initialize", None) + if inspect.iscoroutinefunction(init_fn): + await init_fn() + + # 获取模型列表 + models = await inst.get_models() + models = models or [] + + metadata_map = {} + for model_id in models: + meta = LLM_METADATAS.get(model_id) + if meta: + metadata_map[model_id] = meta + + # 销毁实例(如果有 terminate 方法) + terminate_fn = getattr(inst, "terminate", None) + if inspect.iscoroutinefunction(terminate_fn): + await terminate_fn() + + logger.info( + f"获取到 provider_source {provider_source_id} 的模型列表: {models}", + ) + + return ( + Response() + .ok({"models": models, "model_metadata": metadata_map}) + .__dict__ + ) + except Exception as e: + logger.error(traceback.format_exc()) + return Response().error(f"获取模型列表失败: {e!s}").__dict__ + async def get_platform_list(self): """获取所有平台的列表""" platform_list = [] @@ -533,7 +789,15 @@ class ConfigRoute(Route): data = await request.json config = data.get("config", None) conf_id = data.get("conf_id", None) + try: + # 不更新 provider_sources, provider, platform + # 这些配置有单独的接口进行更新 + if conf_id == "default": + no_update_keys = ["provider_sources", "provider", "platform"] + for key in no_update_keys: + config[key] = self.acm.default_conf[key] + await self._save_astrbot_configs(config, conf_id) await self.core_lifecycle.reload_pipeline_scheduler(conf_id) return Response().ok(None, "保存成功~").__dict__ @@ -573,28 +837,30 @@ class ConfigRoute(Route): async def post_new_provider(self): new_provider_config = await request.json - self.config["provider"].append(new_provider_config) + try: - save_config(self.config, self.config, is_core=True) - await self.core_lifecycle.provider_manager.load_provider( - new_provider_config, + await self.core_lifecycle.provider_manager.create_provider( + new_provider_config ) except Exception as e: return Response().error(str(e)).__dict__ - return Response().ok(None, "新增服务提供商配置成功~").__dict__ + return Response().ok(None, "新增服务提供商配置成功").__dict__ async def post_update_platform(self): update_platform_config = await request.json - platform_id = update_platform_config.get("id", None) + origin_platform_id = update_platform_config.get("id", None) new_config = update_platform_config.get("config", None) - if not platform_id or not new_config: + if not origin_platform_id or not new_config: return Response().error("参数错误").__dict__ + if origin_platform_id != new_config.get("id", None): + return Response().error("机器人名称不允许修改").__dict__ + # 如果是支持统一 webhook 模式的平台,且启用了统一 webhook 模式,确保有 webhook_uuid ensure_platform_webhook_config(new_config) for i, platform in enumerate(self.config["platform"]): - if platform["id"] == platform_id: + if platform["id"] == origin_platform_id: self.config["platform"][i] = new_config break else: @@ -609,21 +875,15 @@ class ConfigRoute(Route): async def post_update_provider(self): update_provider_config = await request.json - provider_id = update_provider_config.get("id", None) + origin_provider_id = update_provider_config.get("id", None) new_config = update_provider_config.get("config", None) - if not provider_id or not new_config: + if not origin_provider_id or not new_config: return Response().error("参数错误").__dict__ - for i, provider in enumerate(self.config["provider"]): - if provider["id"] == provider_id: - self.config["provider"][i] = new_config - break - else: - return Response().error("未找到对应服务提供商").__dict__ - try: - save_config(self.config, self.config, is_core=True) - await self.core_lifecycle.provider_manager.reload(new_config) + await self.core_lifecycle.provider_manager.update_provider( + origin_provider_id, new_config + ) except Exception as e: return Response().error(str(e)).__dict__ return Response().ok(None, "更新成功,已经实时生效~").__dict__ @@ -646,19 +906,17 @@ class ConfigRoute(Route): async def post_delete_provider(self): provider_id = await request.json - provider_id = provider_id.get("id") - for i, provider in enumerate(self.config["provider"]): - if provider["id"] == provider_id: - del self.config["provider"][i] - break - else: - return Response().error("未找到对应服务提供商").__dict__ + provider_id = provider_id.get("id", "") + if not provider_id: + return Response().error("缺少参数 id").__dict__ + try: - save_config(self.config, self.config, is_core=True) - await self.core_lifecycle.provider_manager.terminate_provider(provider_id) + await self.core_lifecycle.provider_manager.delete_provider( + provider_id=provider_id + ) except Exception as e: return Response().error(str(e)).__dict__ - return Response().ok(None, "删除成功,已经实时生效~").__dict__ + return Response().ok(None, "删除成功,已经实时生效。").__dict__ async def get_llm_tools(self): """获取函数调用工具。包含了本地加载的以及 MCP 服务的工具""" diff --git a/dashboard/src/assets/images/provider_logos/modelstack.svg b/dashboard/src/assets/images/provider_logos/modelstack.svg new file mode 100644 index 000000000..940ac8799 --- /dev/null +++ b/dashboard/src/assets/images/provider_logos/modelstack.svg @@ -0,0 +1 @@ +
\ No newline at end of file diff --git a/dashboard/src/components/chat/ChatInput.vue b/dashboard/src/components/chat/ChatInput.vue index 78050d5e3..402a9e3a9 100644 --- a/dashboard/src/components/chat/ChatInput.vue +++ b/dashboard/src/components/chat/ChatInput.vue @@ -26,7 +26,9 @@ :initial-config-id="props.configId" @config-changed="handleConfigChange" /> - + + +