diff --git a/.github/ISSUE_TEMPLATE/PLUGIN_PUBLISH.md b/.github/ISSUE_TEMPLATE/PLUGIN_PUBLISH.md index 73f5009ca..0358a5b27 100644 --- a/.github/ISSUE_TEMPLATE/PLUGIN_PUBLISH.md +++ b/.github/ISSUE_TEMPLATE/PLUGIN_PUBLISH.md @@ -17,6 +17,7 @@ assignees: '' { "name": "插件名", "desc": "插件介绍", + "author": "作者名", "repo": "插件仓库链接", "tags": [], "social_link": "" diff --git a/.github/dependabot.yml b/.github/dependabot.yml new file mode 100644 index 000000000..be006de9a --- /dev/null +++ b/.github/dependabot.yml @@ -0,0 +1,13 @@ +# Keep GitHub Actions up to date with GitHub's Dependabot... +# https://docs.github.com/en/code-security/dependabot/working-with-dependabot/keeping-your-actions-up-to-date-with-dependabot +# https://docs.github.com/en/code-security/dependabot/dependabot-version-updates/configuration-options-for-the-dependabot.yml-file#package-ecosystem +version: 2 +updates: + - package-ecosystem: github-actions + directory: / + groups: + github-actions: + patterns: + - "*" # Group all Actions updates into a single larger pull request + schedule: + interval: weekly diff --git a/.github/workflows/auto_release.yml b/.github/workflows/auto_release.yml index 854a69f16..46d914a9a 100644 --- a/.github/workflows/auto_release.yml +++ b/.github/workflows/auto_release.yml @@ -73,7 +73,7 @@ jobs: uses: actions/checkout@v4 - name: Set up Python - uses: actions/setup-python@v4 + uses: actions/setup-python@v5 with: python-version: '3.10' diff --git a/.github/workflows/coverage_test.yml b/.github/workflows/coverage_test.yml index 30e9237ed..e9c94d679 100644 --- a/.github/workflows/coverage_test.yml +++ b/.github/workflows/coverage_test.yml @@ -1,6 +1,6 @@ name: Run tests and upload coverage -on: +on: push: branches: - master @@ -8,6 +8,7 @@ on: - 'README.md' - 'changelogs/**' - 'dashboard/**' + pull_request: workflow_dispatch: jobs: @@ -21,25 +22,24 @@ jobs: fetch-depth: 0 - name: Set up Python - uses: actions/setup-python@v4 + uses: actions/setup-python@v5 - name: Install dependencies run: | python -m pip install --upgrade pip - pip install -r requirements.txt - pip install pytest pytest-cov pytest-asyncio + pip install pytest pytest-asyncio pytest-cov + pip install --editable . - name: Run tests run: | - mkdir data - mkdir data/plugins - mkdir data/config - mkdir data/temp + mkdir -p data/plugins + mkdir -p data/config + mkdir -p data/temp export TESTING=true export ZHIPU_API_KEY=${{ secrets.OPENAI_API_KEY }} - PYTHONPATH=./ pytest --cov=. tests/ -v -o log_cli=true -o log_level=DEBUG + pytest --cov=. -v -o log_cli=true -o log_level=DEBUG - name: Upload results to Codecov - uses: codecov/codecov-action@v4 + uses: codecov/codecov-action@v5 with: - token: ${{ secrets.CODECOV_TOKEN }} \ No newline at end of file + token: ${{ secrets.CODECOV_TOKEN }} diff --git a/.github/workflows/docker-image.yml b/.github/workflows/docker-image.yml index e0e235c09..c0610a3c1 100644 --- a/.github/workflows/docker-image.yml +++ b/.github/workflows/docker-image.yml @@ -12,7 +12,7 @@ jobs: steps: - name: Pull The Codes - uses: actions/checkout@v3 + uses: actions/checkout@v4 with: fetch-depth: 0 # Must be 0 so we can fetch tags diff --git a/.github/workflows/stale.yml b/.github/workflows/stale.yml index 310e250b6..283e99989 100644 --- a/.github/workflows/stale.yml +++ b/.github/workflows/stale.yml @@ -18,7 +18,7 @@ jobs: pull-requests: write steps: - - uses: actions/stale@v5 + - uses: actions/stale@v9 with: repo-token: ${{ secrets.GITHUB_TOKEN }} stale-issue-message: 'Stale issue message' diff --git a/astrbot/core/__init__.py b/astrbot/core/__init__.py index 104a9edb6..16f108ece 100644 --- a/astrbot/core/__init__.py +++ b/astrbot/core/__init__.py @@ -28,5 +28,3 @@ pip_installer = PipInstaller( astrbot_config.get("pip_install_arg", ""), astrbot_config.get("pypi_index_url", None), ) -web_chat_queue = asyncio.Queue(maxsize=32) -web_chat_back_queue = asyncio.Queue(maxsize=32) diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index 4a75a11bc..b3982cb13 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -6,7 +6,7 @@ import os from astrbot.core.utils.astrbot_path import get_astrbot_data_path -VERSION = "3.5.18" +VERSION = "3.5.19" DB_PATH = os.path.join(get_astrbot_data_path(), "data_v3.db") # 默认配置 @@ -64,7 +64,7 @@ DEFAULT_CONFIG = { "streaming_response": False, "show_tool_use_status": False, "streaming_segmented": False, - "separate_provider": False, + "separate_provider": True, }, "provider_stt_settings": { "enable": False, @@ -724,16 +724,16 @@ CONFIG_METADATA_2 = { "model": "deepseek-chat", }, }, - "智谱 AI": { - "id": "zhipu_default", - "type": "zhipu_chat_completion", + "302.AI": { + "id": "302ai", + "type": "openai_chat_completion", "provider_type": "chat_completion", "enable": True, "key": [], + "api_base": "https://api.302.ai/v1", "timeout": 120, - "api_base": "https://open.bigmodel.cn/api/paas/v4/", "model_config": { - "model": "glm-4-flash", + "model": "gpt-4.1-mini", }, }, "硅基流动": { @@ -748,6 +748,18 @@ CONFIG_METADATA_2 = { "model": "deepseek-ai/DeepSeek-V3", }, }, + "PPIO派欧云": { + "id": "ppio", + "type": "openai_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "api_base": "https://api.ppinfra.com/v3/openai", + "timeout": 120, + "model_config": { + "model": "deepseek/deepseek-r1", + }, + }, "Kimi": { "id": "moonshot", "type": "openai_chat_completion", @@ -760,16 +772,16 @@ CONFIG_METADATA_2 = { "model": "moonshot-v1-8k", }, }, - "PPIO派欧云": { - "id": "ppio", - "type": "openai_chat_completion", + "智谱 AI": { + "id": "zhipu_default", + "type": "zhipu_chat_completion", "provider_type": "chat_completion", "enable": True, "key": [], - "api_base": "https://api.ppinfra.com/v3/openai", "timeout": 120, + "api_base": "https://open.bigmodel.cn/api/paas/v4/", "model_config": { - "model": "deepseek/deepseek-r1", + "model": "glm-4-flash", }, }, "Dify": { diff --git a/astrbot/core/conversation_mgr.py b/astrbot/core/conversation_mgr.py index b0f5c136d..b665488e4 100644 --- a/astrbot/core/conversation_mgr.py +++ b/astrbot/core/conversation_mgr.py @@ -88,7 +88,10 @@ class ConversationManager: return self.session_conversations.get(unified_msg_origin, None) async def get_conversation( - self, unified_msg_origin: str, conversation_id: str + self, + unified_msg_origin: str, + conversation_id: str, + create_if_not_exists: bool = False, ) -> Conversation: """获取会话的对话 @@ -98,6 +101,13 @@ class ConversationManager: Returns: conversation (Conversation): 对话对象 """ + conv = self.db.get_conversation_by_user_id(unified_msg_origin, conversation_id) + if not conv and create_if_not_exists: + # 如果对话不存在且需要创建,则新建一个对话 + conversation_id = await self.new_conversation(unified_msg_origin) + return self.db.get_conversation_by_user_id( + unified_msg_origin, conversation_id + ) return self.db.get_conversation_by_user_id(unified_msg_origin, conversation_id) async def get_conversations(self, unified_msg_origin: str) -> List[Conversation]: diff --git a/astrbot/core/pipeline/context.py b/astrbot/core/pipeline/context.py index d98f7c341..0b9d9e533 100644 --- a/astrbot/core/pipeline/context.py +++ b/astrbot/core/pipeline/context.py @@ -23,7 +23,12 @@ class PipelineContext: event: AstrMessageEvent, hook_type: EventType, *args, - ): + ) -> bool: + """调用事件钩子函数 + + Returns: + bool: 如果事件被终止,返回 True + """ platform_id = event.get_platform_id() handlers = star_handlers_registry.get_handlers_by_event_type( hook_type, platform_id=platform_id @@ -41,7 +46,8 @@ class PipelineContext: logger.info( f"{star_map[handler.handler_module_path].name} - {handler.handler_name} 终止了事件传播。" ) - return + + return event.is_stopped() async def call_handler( self, diff --git a/astrbot/core/pipeline/process_stage/agent_runner/tool_loop_agent.py b/astrbot/core/pipeline/process_stage/agent_runner/tool_loop_agent.py index 3163e02e4..c2961ded5 100644 --- a/astrbot/core/pipeline/process_stage/agent_runner/tool_loop_agent.py +++ b/astrbot/core/pipeline/process_stage/agent_runner/tool_loop_agent.py @@ -106,7 +106,6 @@ class ToolLoopAgent(BaseAgentRunner): # 处理 LLM 响应 llm_resp = llm_resp_result - logger.debug(f"LLMResp: {llm_resp}") if llm_resp.role == "err": # 如果 LLM 响应错误,转换到错误状态 @@ -127,9 +126,10 @@ class ToolLoopAgent(BaseAgentRunner): self._transition_state(AgentState.DONE) # 执行事件钩子 - await self.pipeline_ctx.call_event_hook( + if await self.pipeline_ctx.call_event_hook( self.event, EventType.OnLLMResponseEvent, llm_resp - ) + ): + return # 返回 LLM 结果 if llm_resp.result_chain: @@ -218,7 +218,9 @@ class ToolLoopAgent(BaseAgentRunner): content="返回了图片(已直接发送给用户)", ) ) - yield MessageChain().base64_image(res.content[0].data) + yield MessageChain(type="tool_direct_result").base64_image( + res.content[0].data + ) elif isinstance(res.content[0], EmbeddedResource): resource = res.content[0].resource if isinstance(resource, TextResourceContents): @@ -242,7 +244,9 @@ class ToolLoopAgent(BaseAgentRunner): content="返回了图片(已直接发送给用户)", ) ) - yield MessageChain().base64_image(res.content[0].data) + yield MessageChain(type="tool_direct_result").base64_image( + res.content[0].data + ) else: tool_call_result_blocks.append( ToolCallMessageSegment( @@ -275,7 +279,9 @@ class ToolLoopAgent(BaseAgentRunner): self._transition_state(AgentState.DONE) if res := self.event.get_result(): if res.chain: - yield MessageChain(chain=res.chain) + yield MessageChain( + chain=res.chain, type="tool_direct_result" + ) self.event.clear_result() except Exception as e: diff --git a/astrbot/core/pipeline/process_stage/method/llm_request.py b/astrbot/core/pipeline/process_stage/method/llm_request.py index 0e64d733b..770bd65e8 100644 --- a/astrbot/core/pipeline/process_stage/method/llm_request.py +++ b/astrbot/core/pipeline/process_stage/method/llm_request.py @@ -24,8 +24,8 @@ from astrbot.core.provider.entities import ( ) from astrbot.core.star.session_llm_manager import SessionServiceManager from astrbot.core.star.star_handler import EventType -from astrbot.core import web_chat_back_queue from ..agent_runner.tool_loop_agent import ToolLoopAgent +from astrbot.core.provider import Provider class LLMRequestSubStage(Stage): @@ -53,22 +53,35 @@ class LLMRequestSubStage(Stage): self.conv_manager = ctx.plugin_manager.context.conversation_manager + def _select_provider(self, event: AstrMessageEvent) -> Provider | None: + """选择使用的 LLM 提供商""" + sel_provider = event.get_extra("selected_provider") + _ctx = self.ctx.plugin_manager.context + if sel_provider and isinstance(sel_provider, str): + provider = _ctx.get_provider_by_id(sel_provider) + if not provider: + logger.error(f"未找到指定的提供商: {sel_provider}。") + return provider + + return _ctx.get_using_provider(umo=event.unified_msg_origin) + async def process( self, event: AstrMessageEvent, _nested: bool = False ) -> Union[None, AsyncGenerator[None, None]]: - req: ProviderRequest = None + req: ProviderRequest | None = None if not self.ctx.astrbot_config["provider_settings"]["enable"]: logger.debug("未启用 LLM 能力,跳过处理。") return + # 检查会话级别的LLM启停状态 if not SessionServiceManager.should_process_llm_request(event): logger.debug(f"会话 {event.unified_msg_origin} 禁用了 LLM,跳过处理。") return - umo = event.unified_msg_origin - provider = self.ctx.plugin_manager.context.get_using_provider(umo=umo) + + provider = self._select_provider(event) if provider is None: return @@ -83,6 +96,8 @@ class LLMRequestSubStage(Stage): else: req = ProviderRequest(prompt="", image_urls=[]) + if sel_model := event.get_extra("selected_model"): + req.model = sel_model if self.provider_wake_prefix: if not event.message_str.startswith(self.provider_wake_prefix): return @@ -121,7 +136,8 @@ class LLMRequestSubStage(Stage): return # 执行请求 LLM 前事件钩子。 - await self.ctx.call_event_hook(event, EventType.OnLLMRequestEvent, req) + if await self.ctx.call_event_hook(event, EventType.OnLLMRequestEvent, req): + return if isinstance(req.contexts, str): req.contexts = json.loads(req.contexts) @@ -174,13 +190,24 @@ class LLMRequestSubStage(Stage): step_idx += 1 try: async for resp in tool_loop_agent.step(): + if event.is_stopped(): + return if resp.type == "tool_call_result": - continue # 跳过工具调用结果 + msg_chain = resp.data["chain"] + if msg_chain.type == "tool_direct_result": + # tool_direct_result 用于标记 llm tool 需要直接发送给用户的内容 + resp.data["chain"].type = "tool_call_result" + await event.send(resp.data["chain"]) + continue + # 对于其他情况,暂时先不处理 if resp.type == "tool_call": if self.streaming_response: # 用来标记流式响应需要分节 yield MessageChain(chain=[], type="break") - if self.show_tool_use or event.get_platform_name() == "webchat": + if ( + self.show_tool_use + or event.get_platform_name() == "webchat" + ): resp.data["chain"].type = "tool_call" await event.send(resp.data["chain"]) continue @@ -249,11 +276,13 @@ class LLMRequestSubStage(Stage): # 异步处理 WebChat 特殊情况 if event.get_platform_name() == "webchat": - asyncio.create_task(self._handle_webchat(event, req)) + asyncio.create_task(self._handle_webchat(event, req, provider)) await self._save_to_history(event, req, tool_loop_agent.get_final_llm_resp()) - async def _handle_webchat(self, event: AstrMessageEvent, req: ProviderRequest): + async def _handle_webchat( + self, event: AstrMessageEvent, req: ProviderRequest, prov: Provider + ): """处理 WebChat 平台的特殊情况,包括第一次 LLM 对话时总结对话内容生成 title""" # 检查会话级别的LLM启停状态,防止标题生成功能绕过会话级别限制 if not SessionServiceManager.should_process_llm_request(event): @@ -268,17 +297,16 @@ class LLMRequestSubStage(Stage): latest_pair = messages[-2:] if not latest_pair: return - provider = self.ctx.plugin_manager.context.get_using_provider() cleaned_text = "User: " + latest_pair[0].get("content", "").strip() logger.debug(f"WebChat 对话标题生成请求,清理后的文本: {cleaned_text}") - llm_resp = await provider.text_chat( + llm_resp = await prov.text_chat( system_prompt="You are expert in summarizing user's query.", prompt=( f"Please summarize the following query of user:\n" f"{cleaned_text}\n" "Only output the summary within 10 words, DO NOT INCLUDE any other text." "You must use the same language as the user." - "If you think the dialog is too short to summarize, only output a special mark: `None`" + "If you think the dialog is too short to summarize, only output a special mark: ``" ), ) if llm_resp and llm_resp.completion_text: @@ -286,7 +314,7 @@ class LLMRequestSubStage(Stage): f"WebChat 对话标题生成响应: {llm_resp.completion_text.strip()}" ) title = llm_resp.completion_text.strip() - if not title or "None" == title: + if not title or "" in title: return await self.conv_manager.update_conversation_title( event.unified_msg_origin, title=title @@ -302,13 +330,6 @@ class LLMRequestSubStage(Stage): cid=cid, title=title, ) - web_chat_back_queue.put_nowait( - { - "type": "update_title", - "cid": cid, - "data": title, - } - ) async def _save_to_history( self, @@ -340,7 +361,6 @@ class LLMRequestSubStage(Stage): await self.conv_manager.update_conversation( event.unified_msg_origin, req.conversation.cid, history=messages ) - logger.debug(f"messages persisted: {messages}") def fix_messages(self, messages: list[dict]) -> list[dict]: """验证并且修复上下文""" diff --git a/astrbot/core/pipeline/waking_check/stage.py b/astrbot/core/pipeline/waking_check/stage.py index 52759a00c..3797751bf 100644 --- a/astrbot/core/pipeline/waking_check/stage.py +++ b/astrbot/core/pipeline/waking_check/stage.py @@ -165,7 +165,7 @@ class WakingCheckStage(Stage): "parsed_params" ) - event.clear_extra() + event._extras.pop("parsed_params", None) # 根据会话配置过滤插件处理器 activated_handlers = SessionPluginManager.filter_handlers_by_session(event, activated_handlers) diff --git a/astrbot/core/platform/sources/webchat/webchat_adapter.py b/astrbot/core/platform/sources/webchat/webchat_adapter.py index fa384ed99..aaac8e289 100644 --- a/astrbot/core/platform/sources/webchat/webchat_adapter.py +++ b/astrbot/core/platform/sources/webchat/webchat_adapter.py @@ -2,7 +2,7 @@ import time import asyncio import uuid import os -from typing import Awaitable, Any +from typing import Awaitable, Any, Callable from astrbot.core.platform import ( Platform, AstrBotMessage, @@ -13,7 +13,7 @@ from astrbot.core.platform import ( from astrbot.core.message.message_event_result import MessageChain from astrbot.core.message.components import Plain, Image, Record # noqa: F403 from astrbot import logger -from astrbot.core import web_chat_queue +from .webchat_queue_mgr import webchat_queue_mgr, WebChatQueueMgr from .webchat_event import WebChatMessageEvent from astrbot.core.platform.astr_message_event import MessageSesion from ...register import register_platform_adapter @@ -21,14 +21,46 @@ from astrbot.core.utils.astrbot_path import get_astrbot_data_path class QueueListener: - def __init__(self, queue: asyncio.Queue, callback: callable) -> None: - self.queue = queue + def __init__(self, webchat_queue_mgr: WebChatQueueMgr, callback: Callable) -> None: + self.webchat_queue_mgr = webchat_queue_mgr self.callback = callback + self.running_tasks = set() + + async def listen_to_queue(self, conversation_id: str): + """Listen to a specific conversation queue""" + queue = self.webchat_queue_mgr.get_or_create_queue(conversation_id) + while True: + try: + data = await queue.get() + await self.callback(data) + except Exception as e: + logger.error( + f"Error processing message from conversation {conversation_id}: {e}" + ) + break async def run(self): + """Monitor for new conversation queues and start listeners""" + monitored_conversations = set() + while True: - data = await self.queue.get() - await self.callback(data) + # Check for new conversations + current_conversations = set(self.webchat_queue_mgr.queues.keys()) + new_conversations = current_conversations - monitored_conversations + + # Start listeners for new conversations + for conversation_id in new_conversations: + task = asyncio.create_task(self.listen_to_queue(conversation_id)) + self.running_tasks.add(task) + task.add_done_callback(self.running_tasks.discard) + monitored_conversations.add(conversation_id) + logger.debug(f"Started listener for conversation: {conversation_id}") + + # Clean up monitored conversations that no longer exist + removed_conversations = monitored_conversations - current_conversations + monitored_conversations -= removed_conversations + + await asyncio.sleep(1) # Check for new conversations every second @register_platform_adapter("webchat", "webchat") @@ -45,7 +77,7 @@ class WebChatAdapter(Platform): os.makedirs(self.imgs_dir, exist_ok=True) self.metadata = PlatformMetadata( - name="webchat", description="webchat", id=self.config.get("id") + name="webchat", description="webchat", id=self.config.get("id", "") ) async def send_by_session( @@ -105,7 +137,7 @@ class WebChatAdapter(Platform): abm = await self.convert_message(data) await self.handle_msg(abm) - bot = QueueListener(web_chat_queue, callback) + bot = QueueListener(webchat_queue_mgr, callback) return bot.run() def meta(self) -> PlatformMetadata: @@ -119,6 +151,10 @@ class WebChatAdapter(Platform): session_id=message.session_id, ) + _, _, payload = message.raw_message # type: ignore + message_event.set_extra("selected_provider", payload.get("selected_provider")) + message_event.set_extra("selected_model", payload.get("selected_model")) + self.commit_event(message_event) async def terminate(self): diff --git a/astrbot/core/platform/sources/webchat/webchat_event.py b/astrbot/core/platform/sources/webchat/webchat_event.py index 111027a5c..c4e5d63c0 100644 --- a/astrbot/core/platform/sources/webchat/webchat_event.py +++ b/astrbot/core/platform/sources/webchat/webchat_event.py @@ -5,8 +5,8 @@ from astrbot.api import logger from astrbot.api.event import AstrMessageEvent, MessageChain from astrbot.api.message_components import Plain, Image, Record from astrbot.core.utils.io import download_image_by_url -from astrbot.core import web_chat_back_queue from astrbot.core.utils.astrbot_path import get_astrbot_data_path +from .webchat_queue_mgr import webchat_queue_mgr imgs_dir = os.path.join(get_astrbot_data_path(), "webchat", "imgs") @@ -18,13 +18,14 @@ class WebChatMessageEvent(AstrMessageEvent): @staticmethod async def _send(message: MessageChain, session_id: str, streaming: bool = False): + cid = session_id.split("!")[-1] + web_chat_back_queue = webchat_queue_mgr.get_or_create_back_queue(cid) if not message: await web_chat_back_queue.put( {"type": "end", "data": "", "streaming": False} ) return "" - cid = session_id.split("!")[-1] data = "" for comp in message.chain: if isinstance(comp, Plain): @@ -98,18 +99,22 @@ class WebChatMessageEvent(AstrMessageEvent): async def send(self, message: MessageChain): await WebChatMessageEvent._send(message, session_id=self.session_id) + cid = self.session_id.split("!")[-1] + web_chat_back_queue = webchat_queue_mgr.get_or_create_back_queue(cid) await web_chat_back_queue.put( { "type": "end", "data": "", "streaming": False, - "cid": self.session_id.split("!")[-1], + "cid": cid, } ) await super().send(message) async def send_streaming(self, generator, use_fallback: bool = False): final_data = "" + cid = self.session_id.split("!")[-1] + web_chat_back_queue = webchat_queue_mgr.get_or_create_back_queue(cid) async for chain in generator: if chain.type == "break" and final_data: # 分割符 @@ -118,7 +123,7 @@ class WebChatMessageEvent(AstrMessageEvent): "type": "end", "data": final_data, "streaming": True, - "cid": self.session_id.split("!")[-1], + "cid": cid, } ) final_data = "" @@ -132,7 +137,7 @@ class WebChatMessageEvent(AstrMessageEvent): "type": "end", "data": final_data, "streaming": True, - "cid": self.session_id.split("!")[-1], + "cid": cid, } ) await super().send_streaming(generator, use_fallback) diff --git a/astrbot/core/platform/sources/webchat/webchat_queue_mgr.py b/astrbot/core/platform/sources/webchat/webchat_queue_mgr.py new file mode 100644 index 000000000..96e172212 --- /dev/null +++ b/astrbot/core/platform/sources/webchat/webchat_queue_mgr.py @@ -0,0 +1,33 @@ +import asyncio + +class WebChatQueueMgr: + def __init__(self) -> None: + self.queues = {} + """Conversation ID to asyncio.Queue mapping""" + self.back_queues = {} + """Conversation ID to asyncio.Queue mapping for responses""" + + def get_or_create_queue(self, conversation_id: str) -> asyncio.Queue: + """Get or create a queue for the given conversation ID""" + if conversation_id not in self.queues: + self.queues[conversation_id] = asyncio.Queue() + return self.queues[conversation_id] + + def get_or_create_back_queue(self, conversation_id: str) -> asyncio.Queue: + """Get or create a back queue for the given conversation ID""" + if conversation_id not in self.back_queues: + self.back_queues[conversation_id] = asyncio.Queue() + return self.back_queues[conversation_id] + + def remove_queues(self, conversation_id: str): + """Remove queues for the given conversation ID""" + if conversation_id in self.queues: + del self.queues[conversation_id] + if conversation_id in self.back_queues: + del self.back_queues[conversation_id] + + def has_queue(self, conversation_id: str) -> bool: + """Check if a queue exists for the given conversation ID""" + return conversation_id in self.queues + +webchat_queue_mgr = WebChatQueueMgr() diff --git a/astrbot/core/platform/sources/wechatpadpro/wechatpadpro_adapter.py b/astrbot/core/platform/sources/wechatpadpro/wechatpadpro_adapter.py index 58e3c9b19..7d5984416 100644 --- a/astrbot/core/platform/sources/wechatpadpro/wechatpadpro_adapter.py +++ b/astrbot/core/platform/sources/wechatpadpro/wechatpadpro_adapter.py @@ -210,6 +210,16 @@ class WeChatPadProAdapter(Platform): logger.error(traceback.format_exc()) return False + def _extract_auth_key(self, data): + """Helper method to extract auth_key from response data.""" + if isinstance(data, dict): + auth_keys = data.get("authKeys") # 新接口 + if isinstance(auth_keys, list) and auth_keys: + return auth_keys[0] + elif isinstance(data, list) and data: # 旧接口 + return data[0] + return None + async def generate_auth_key(self): """ 生成授权码。 @@ -218,28 +228,26 @@ class WeChatPadProAdapter(Platform): params = {"key": self.admin_key} payload = {"Count": 1, "Days": 365} # 生成一个有效期365天的授权码 + self.auth_key = None # Reset auth_key before generating a new one + async with aiohttp.ClientSession() as session: try: async with session.post(url, params=params, json=payload) as response: + if response.status != 200: + logger.error(f"生成授权码失败: {response.status}, {await response.text()}") + return + response_data = await response.json() - # 修正成功判断条件和授权码提取路径 - if response.status == 200 and response_data.get("Code") == 200: - # 授权码在 Data 字段的列表中 - if ( - response_data.get("Data") - and isinstance(response_data["Data"], list) - and len(response_data["Data"]) > 0 - ): - self.auth_key = response_data["Data"][0] - logger.info(f"成功获取授权码 {self.auth_key[:8]}...") + if response_data.get("Code") == 200: + if data := response_data.get("Data"): + self.auth_key = self._extract_auth_key(data) + + if self.auth_key: + logger.info("成功获取授权码") else: - logger.error( - f"生成授权码成功但未找到授权码: {response_data}" - ) + logger.error(f"生成授权码成功但未找到授权码: {response_data}") else: - logger.error( - f"生成授权码失败: {response.status}, {response_data}" - ) + logger.error(f"生成授权码失败: {response_data}") except aiohttp.ClientConnectorError as e: logger.error(f"连接到 WeChatPadPro 服务失败: {e}") except Exception as e: diff --git a/astrbot/core/provider/entities.py b/astrbot/core/provider/entities.py index abb01960c..2d120d7f6 100644 --- a/astrbot/core/provider/entities.py +++ b/astrbot/core/provider/entities.py @@ -110,6 +110,9 @@ class ProviderRequest: tool_calls_result: list[ToolCallsResult] | ToolCallsResult | None = None """附加的上次请求后工具调用的结果。参考: https://platform.openai.com/docs/guides/function-calling#handling-function-calls""" + model: str | None = None + """模型名称,为 None 时使用提供商的默认模型""" + def __repr__(self): return f"ProviderRequest(prompt={self.prompt}, session_id={self.session_id}, image_urls={self.image_urls}, func_tool={self.func_tool}, contexts={self._print_friendly_context()}, system_prompt={self.system_prompt.strip()}, tool_calls_result={self.tool_calls_result})" diff --git a/astrbot/core/provider/manager.py b/astrbot/core/provider/manager.py index 2abe59d65..05747c3ff 100644 --- a/astrbot/core/provider/manager.py +++ b/astrbot/core/provider/manager.py @@ -40,11 +40,13 @@ class ProviderManager: begin_dialogs = [] user_turn = True for dialog in begin_dialogs: - bd_processed.append({ - "role": "user" if user_turn else "assistant", - "content": dialog, - "_no_save": None, # 不持久化到 db - }) + bd_processed.append( + { + "role": "user" if user_turn else "assistant", + "content": dialog, + "_no_save": None, # 不持久化到 db + } + ) user_turn = not user_turn if mood_imitation_dialogs: if len(mood_imitation_dialogs) % 2 != 0: @@ -93,15 +95,15 @@ class ProviderManager: """加载的 Text To Speech Provider 的实例""" self.embedding_provider_insts: List[Provider] = [] """加载的 Embedding Provider 的实例""" - self.inst_map = {} + self.inst_map: dict[str, Provider] = {} """Provider 实例映射. key: provider_id, value: Provider 实例""" self.llm_tools = llm_tools - self.curr_provider_inst: Provider = None + self.curr_provider_inst: Provider | None = None """默认的 Provider 实例""" - self.curr_stt_provider_inst: STTProvider = None + self.curr_stt_provider_inst: STTProvider | None = None """默认的 Speech To Text Provider 实例""" - self.curr_tts_provider_inst: TTSProvider = None + self.curr_tts_provider_inst: TTSProvider | None = None """默认的 Text To Speech Provider 实例""" self.db_helper = db_helper @@ -145,21 +147,24 @@ class ProviderManager: await self.load_provider(provider_config) # 设置默认提供商 - self.curr_provider_inst = self.inst_map.get( - self.provider_settings.get("default_provider_id") + selected_provider_id = sp.get( + "curr_provider", self.provider_settings.get("default_provider_id") ) + selected_stt_provider_id = sp.get( + "curr_provider_stt", self.provider_stt_settings.get("provider_id") + ) + selected_tts_provider_id = sp.get( + "curr_provider_tts", self.provider_tts_settings.get("provider_id") + ) + self.curr_provider_inst = self.inst_map.get(selected_provider_id) if not self.curr_provider_inst and self.provider_insts: self.curr_provider_inst = self.provider_insts[0] - self.curr_stt_provider_inst = self.inst_map.get( - self.provider_stt_settings.get("provider_id") - ) + self.curr_stt_provider_inst = self.inst_map.get(selected_stt_provider_id) if not self.curr_stt_provider_inst and self.stt_provider_insts: self.curr_stt_provider_inst = self.stt_provider_insts[0] - self.curr_tts_provider_inst = self.inst_map.get( - self.provider_tts_settings.get("provider_id") - ) + self.curr_tts_provider_inst = self.inst_map.get(selected_tts_provider_id) if not self.curr_tts_provider_inst and self.tts_provider_insts: self.curr_tts_provider_inst = self.tts_provider_insts[0] @@ -417,7 +422,7 @@ class ProviderManager: self.curr_tts_provider_inst = None if getattr(self.inst_map[provider_id], "terminate", None): - await self.inst_map[provider_id].terminate() + await self.inst_map[provider_id].terminate() # type: ignore logger.info( f"{provider_id} 提供商适配器已终止({len(self.provider_insts)}, {len(self.stt_provider_insts)}, {len(self.tts_provider_insts)})" @@ -427,6 +432,6 @@ class ProviderManager: async def terminate(self): for provider_inst in self.provider_insts: if hasattr(provider_inst, "terminate"): - await provider_inst.terminate() + await provider_inst.terminate() # type: ignore # 清理 MCP Client 连接 await self.llm_tools.mcp_service_queue.put({"type": "terminate"}) diff --git a/astrbot/core/provider/provider.py b/astrbot/core/provider/provider.py index 1ecca3537..98e8fab85 100644 --- a/astrbot/core/provider/provider.py +++ b/astrbot/core/provider/provider.py @@ -88,6 +88,7 @@ class Provider(AbstractProvider): contexts: list = None, system_prompt: str = None, tool_calls_result: ToolCallsResult | list[ToolCallsResult] = None, + model: str | None = None, **kwargs, ) -> LLMResponse: """获得 LLM 的文本对话结果。会使用当前的模型进行对话。 @@ -116,6 +117,7 @@ class Provider(AbstractProvider): contexts: list = None, system_prompt: str = None, tool_calls_result: ToolCallsResult | list[ToolCallsResult] = None, + model: str | None = None, **kwargs, ) -> AsyncGenerator[LLMResponse, None]: """获得 LLM 的流式文本对话结果。会使用当前的模型进行对话。在生成的最后会返回一次完整的结果。 diff --git a/astrbot/core/provider/sources/anthropic_source.py b/astrbot/core/provider/sources/anthropic_source.py index a53250fb7..aaff177e5 100644 --- a/astrbot/core/provider/sources/anthropic_source.py +++ b/astrbot/core/provider/sources/anthropic_source.py @@ -235,6 +235,7 @@ class ProviderAnthropic(Provider): contexts=None, system_prompt=None, tool_calls_result=None, + model=None, **kwargs, ) -> LLMResponse: if contexts is None: @@ -259,7 +260,7 @@ class ProviderAnthropic(Provider): system_prompt, new_messages = self._prepare_payload(context_query) model_config = self.provider_config.get("model_config", {}) - model_config["model"] = self.get_model() + model_config["model"] = model or self.get_model() payloads = {"messages": new_messages, **model_config} @@ -285,6 +286,7 @@ class ProviderAnthropic(Provider): contexts=..., system_prompt=None, tool_calls_result=None, + model=None, **kwargs, ): if contexts is None: @@ -300,12 +302,16 @@ class ProviderAnthropic(Provider): # tool calls result if tool_calls_result: - context_query.extend(tool_calls_result.to_openai_messages()) + if not isinstance(tool_calls_result, list): + context_query.extend(tool_calls_result.to_openai_messages()) + else: + for tcr in tool_calls_result: + context_query.extend(tcr.to_openai_messages()) system_prompt, new_messages = self._prepare_payload(context_query) model_config = self.provider_config.get("model_config", {}) - model_config["model"] = self.get_model() + model_config["model"] = model or self.get_model() payloads = {"messages": new_messages, **model_config} diff --git a/astrbot/core/provider/sources/azure_tts_source.py b/astrbot/core/provider/sources/azure_tts_source.py index c35c7ec6c..6ddf452d4 100644 --- a/astrbot/core/provider/sources/azure_tts_source.py +++ b/astrbot/core/provider/sources/azure_tts_source.py @@ -19,6 +19,7 @@ from ..register import register_provider_adapter TEMP_DIR = Path("data/temp/azure_tts") TEMP_DIR.mkdir(parents=True, exist_ok=True) + class OTTSProvider: def __init__(self, config: Dict): self.skey = config["OTTS_SKEY"] @@ -70,12 +71,12 @@ class OTTSProvider: "style": voice_params["style"], "role": voice_params["role"], "rate": voice_params["rate"], - "volume": voice_params["volume"] + "volume": voice_params["volume"], }, headers={ "User-Agent": f"AstrBot/{VERSION}", - "UAK": "AstrBot/AzureTTS" - } + "UAK": "AstrBot/AzureTTS", + }, ) response.raise_for_status() file_path.parent.mkdir(parents=True, exist_ok=True) @@ -88,14 +89,19 @@ class OTTSProvider: raise RuntimeError(f"OTTS请求失败: {str(e)}") from e await asyncio.sleep(0.5 * (attempt + 1)) + class AzureNativeProvider(TTSProvider): def __init__(self, provider_config: dict, provider_settings: dict): super().__init__(provider_config, provider_settings) - self.subscription_key = provider_config.get("azure_tts_subscription_key", "").strip() + self.subscription_key = provider_config.get( + "azure_tts_subscription_key", "" + ).strip() if not re.fullmatch(r"^[a-zA-Z0-9]{32}$", self.subscription_key): raise ValueError("无效的Azure订阅密钥") self.region = provider_config.get("azure_tts_region", "eastus").strip() - self.endpoint = f"https://{self.region}.tts.speech.microsoft.com/cognitiveservices/v1" + self.endpoint = ( + f"https://{self.region}.tts.speech.microsoft.com/cognitiveservices/v1" + ) self.client = None self.token = None self.token_expire = 0 @@ -104,15 +110,17 @@ class AzureNativeProvider(TTSProvider): "style": provider_config.get("azure_tts_style", "cheerful"), "role": provider_config.get("azure_tts_role", "Boy"), "rate": provider_config.get("azure_tts_rate", "1"), - "volume": provider_config.get("azure_tts_volume", "100") + "volume": provider_config.get("azure_tts_volume", "100"), } async def __aenter__(self): - self.client = AsyncClient(headers={ - "User-Agent": f"AstrBot/{VERSION}", - "Content-Type": "application/ssml+xml", - "X-Microsoft-OutputFormat": "riff-48khz-16bit-mono-pcm" - }) + self.client = AsyncClient( + headers={ + "User-Agent": f"AstrBot/{VERSION}", + "Content-Type": "application/ssml+xml", + "X-Microsoft-OutputFormat": "riff-48khz-16bit-mono-pcm", + } + ) return self async def __aexit__(self, exc_type, exc_val, exc_tb): @@ -120,10 +128,11 @@ class AzureNativeProvider(TTSProvider): await self.client.aclose() async def _refresh_token(self): - token_url = f"https://{self.region}.api.cognitive.microsoft.com/sts/v1.0/issuetoken" + token_url = ( + f"https://{self.region}.api.cognitive.microsoft.com/sts/v1.0/issuetoken" + ) response = await self.client.post( - token_url, - headers={"Ocp-Apim-Subscription-Key": self.subscription_key} + token_url, headers={"Ocp-Apim-Subscription-Key": self.subscription_key} ) response.raise_for_status() self.token = response.text @@ -150,8 +159,8 @@ class AzureNativeProvider(TTSProvider): content=ssml, headers={ "Authorization": f"Bearer {self.token}", - "User-Agent": f"AstrBot/{VERSION}" - } + "User-Agent": f"AstrBot/{VERSION}", + }, ) response.raise_for_status() file_path.parent.mkdir(parents=True, exist_ok=True) @@ -160,6 +169,7 @@ class AzureNativeProvider(TTSProvider): f.write(chunk) return str(file_path.resolve()) + @register_provider_adapter("azure_tts", "Azure TTS", ProviderType.TEXT_TO_SPEECH) class AzureTTSProvider(TTSProvider): def __init__(self, provider_config: dict, provider_settings: dict): @@ -183,7 +193,7 @@ class AzureTTSProvider(TTSProvider): error_msg = ( f"JSON解析失败,请检查格式(错误位置:行 {e.lineno} 列 {e.colno})\n" f"错误详情: {e.msg}\n" - f"错误上下文: {json_str[max(0, e.pos-30):e.pos+30]}" + f"错误上下文: {json_str[max(0, e.pos - 30) : e.pos + 30]}" ) raise ValueError(error_msg) from e except KeyError as e: @@ -202,8 +212,8 @@ class AzureTTSProvider(TTSProvider): "style": self.provider_config.get("azure_tts_style"), "role": self.provider_config.get("azure_tts_role"), "rate": self.provider_config.get("azure_tts_rate"), - "volume": self.provider_config.get("azure_tts_volume") - } + "volume": self.provider_config.get("azure_tts_volume"), + }, ) else: async with self.provider as provider: diff --git a/astrbot/core/provider/sources/dashscope_source.py b/astrbot/core/provider/sources/dashscope_source.py index 3498f8346..46b12726b 100644 --- a/astrbot/core/provider/sources/dashscope_source.py +++ b/astrbot/core/provider/sources/dashscope_source.py @@ -67,6 +67,7 @@ class ProviderDashscope(ProviderOpenAIOfficial): func_tool: FuncCall = None, contexts: List = None, system_prompt: str = None, + model=None, **kwargs, ) -> LLMResponse: if contexts is None: @@ -163,6 +164,7 @@ class ProviderDashscope(ProviderOpenAIOfficial): contexts=..., system_prompt=None, tool_calls_result=None, + model=None, **kwargs, ): # raise NotImplementedError("This method is not implemented yet.") diff --git a/astrbot/core/provider/sources/dify_source.py b/astrbot/core/provider/sources/dify_source.py index 81c910d66..e6f345741 100644 --- a/astrbot/core/provider/sources/dify_source.py +++ b/astrbot/core/provider/sources/dify_source.py @@ -18,7 +18,7 @@ class ProviderDify(Provider): self, provider_config, provider_settings, - default_persona = None, + default_persona=None, ) -> None: super().__init__( provider_config, @@ -60,12 +60,14 @@ class ProviderDify(Provider): func_tool: FuncCall = None, contexts: List = None, system_prompt: str = None, + tool_calls_result=None, + model=None, **kwargs, ) -> LLMResponse: if image_urls is None: image_urls = [] result = "" - session_id = session_id or kwargs.get("user") # 1734 + session_id = session_id or kwargs.get("user") or "unknown" # 1734 conversation_id = self.conversation_ids.get(session_id, "") files_payload = [] @@ -197,6 +199,7 @@ class ProviderDify(Provider): contexts=..., system_prompt=None, tool_calls_result=None, + model=None, **kwargs, ): # raise NotImplementedError("This method is not implemented yet.") diff --git a/astrbot/core/provider/sources/gemini_source.py b/astrbot/core/provider/sources/gemini_source.py index e1d1f11bd..56526c121 100644 --- a/astrbot/core/provider/sources/gemini_source.py +++ b/astrbot/core/provider/sources/gemini_source.py @@ -14,7 +14,7 @@ import astrbot.core.message.components as Comp from astrbot import logger from astrbot.api.provider import Provider from astrbot.core.message.message_event_result import MessageChain -from astrbot.core.provider.entities import LLMResponse, ToolCallsResult +from astrbot.core.provider.entities import LLMResponse from astrbot.core.provider.func_tool_manager import FuncCall from astrbot.core.utils.io import download_image_by_url @@ -259,10 +259,12 @@ class ProviderGoogleGenAI(Provider): contents.append(content_cls(parts=part)) gemini_contents: list[types.Content] = [] - native_tool_enabled = any([ - self.provider_config.get("gm_native_coderunner", False), - self.provider_config.get("gm_native_search", False), - ]) + native_tool_enabled = any( + [ + self.provider_config.get("gm_native_coderunner", False), + self.provider_config.get("gm_native_search", False), + ] + ) for message in payloads["messages"]: role, content = message["role"], message.get("content") @@ -505,6 +507,7 @@ class ProviderGoogleGenAI(Provider): contexts=None, system_prompt=None, tool_calls_result=None, + model=None, **kwargs, ) -> LLMResponse: if contexts is None: @@ -527,7 +530,7 @@ class ProviderGoogleGenAI(Provider): context_query.extend(tcr.to_openai_messages()) model_config = self.provider_config.get("model_config", {}) - model_config["model"] = self.get_model() + model_config["model"] = model or self.get_model() payloads = {"messages": context_query, **model_config} @@ -544,13 +547,14 @@ class ProviderGoogleGenAI(Provider): async def text_chat_stream( self, - prompt: str, - session_id: str = None, - image_urls: list[str] = None, - func_tool: FuncCall = None, - contexts: str = None, - system_prompt: str = None, - tool_calls_result: ToolCallsResult = None, + prompt, + session_id=None, + image_urls=None, + func_tool=None, + contexts=None, + system_prompt=None, + tool_calls_result=None, + model=None, **kwargs, ) -> AsyncGenerator[LLMResponse, None]: if contexts is None: @@ -566,10 +570,14 @@ class ProviderGoogleGenAI(Provider): # tool calls result if tool_calls_result: - context_query.extend(tool_calls_result.to_openai_messages()) + if not isinstance(tool_calls_result, list): + context_query.extend(tool_calls_result.to_openai_messages()) + else: + for tcr in tool_calls_result: + context_query.extend(tcr.to_openai_messages()) model_config = self.provider_config.get("model_config", {}) - model_config["model"] = self.get_model() + model_config["model"] = model or self.get_model() payloads = {"messages": context_query, **model_config} @@ -628,10 +636,12 @@ class ProviderGoogleGenAI(Provider): if not image_data: logger.warning(f"图片 {image_url} 得到的结果为空,将忽略。") continue - user_content["content"].append({ - "type": "image_url", - "image_url": {"url": image_data}, - }) + user_content["content"].append( + { + "type": "image_url", + "image_url": {"url": image_data}, + } + ) return user_content else: return {"role": "user", "content": text} diff --git a/astrbot/core/provider/sources/openai_source.py b/astrbot/core/provider/sources/openai_source.py index ef6131d8c..ec1624776 100644 --- a/astrbot/core/provider/sources/openai_source.py +++ b/astrbot/core/provider/sources/openai_source.py @@ -30,7 +30,7 @@ class ProviderOpenAIOfficial(Provider): self, provider_config, provider_settings, - default_persona = None, + default_persona=None, ) -> None: super().__init__( provider_config, @@ -222,6 +222,7 @@ class ProviderOpenAIOfficial(Provider): contexts: list | None = None, system_prompt: str | None = None, tool_calls_result: ToolCallsResult | list[ToolCallsResult] | None = None, + model: str | None = None, **kwargs, ) -> tuple: """准备聊天所需的有效载荷和上下文""" @@ -245,7 +246,7 @@ class ProviderOpenAIOfficial(Provider): context_query.extend(tcr.to_openai_messages()) model_config = self.provider_config.get("model_config", {}) - model_config["model"] = self.get_model() + model_config["model"] = model or self.get_model() payloads = {"messages": context_query, **model_config} @@ -346,6 +347,7 @@ class ProviderOpenAIOfficial(Provider): contexts=None, system_prompt=None, tool_calls_result=None, + model=None, **kwargs, ) -> LLMResponse: payloads, context_query = await self._prepare_chat_payload( @@ -354,6 +356,7 @@ class ProviderOpenAIOfficial(Provider): contexts, system_prompt, tool_calls_result, + model=model, **kwargs, ) @@ -413,6 +416,7 @@ class ProviderOpenAIOfficial(Provider): contexts=[], system_prompt=None, tool_calls_result=None, + model=None, **kwargs, ) -> AsyncGenerator[LLMResponse, None]: """流式对话,与服务商交互并逐步返回结果""" @@ -422,6 +426,7 @@ class ProviderOpenAIOfficial(Provider): contexts, system_prompt, tool_calls_result, + model=model, **kwargs, ) @@ -482,7 +487,7 @@ class ProviderOpenAIOfficial(Provider): if flag: flag = False # 删除 image 后,下一条(LLM 响应)也要删除 continue - if isinstance(context["content"], list): + if "content" in context and isinstance(context["content"], list): flag = True # continue new_content = [] @@ -526,7 +531,10 @@ class ProviderOpenAIOfficial(Provider): logger.warning(f"图片 {image_url} 得到的结果为空,将忽略。") continue user_content["content"].append( - {"type": "image_url", "image_url": {"url": image_data}} + { + "type": "image_url", + "image_url": {"url": image_data}, + } ) return user_content else: diff --git a/astrbot/core/provider/sources/volcengine_tts.py b/astrbot/core/provider/sources/volcengine_tts.py index dca0196b1..12e7ed9cd 100644 --- a/astrbot/core/provider/sources/volcengine_tts.py +++ b/astrbot/core/provider/sources/volcengine_tts.py @@ -5,12 +5,12 @@ import os import traceback import asyncio import aiohttp -import requests from ..provider import TTSProvider from ..entities import ProviderType from ..register import register_provider_adapter from astrbot import logger + @register_provider_adapter( "volcengine_tts", "火山引擎 TTS", provider_type=ProviderType.TEXT_TO_SPEECH ) @@ -22,7 +22,9 @@ class ProviderVolcengineTTS(TTSProvider): self.cluster = provider_config.get("volcengine_cluster", "") self.voice_type = provider_config.get("volcengine_voice_type", "") self.speed_ratio = provider_config.get("volcengine_speed_ratio", 1.0) - self.api_base = provider_config.get("api_base", f"https://openspeech.bytedance.com/api/v1/tts") + self.api_base = provider_config.get( + "api_base", "https://openspeech.bytedance.com/api/v1/tts" + ) self.timeout = provider_config.get("timeout", 20) def _build_request_payload(self, text: str) -> dict: @@ -30,11 +32,9 @@ class ProviderVolcengineTTS(TTSProvider): "app": { "appid": self.appid, "token": self.api_key, - "cluster": self.cluster - }, - "user": { - "uid": str(uuid.uuid4()) + "cluster": self.cluster, }, + "user": {"uid": str(uuid.uuid4())}, "audio": { "voice_type": self.voice_type, "encoding": "mp3", @@ -48,60 +48,61 @@ class ProviderVolcengineTTS(TTSProvider): "text_type": "plain", "operation": "query", "with_frontend": 1, - "frontend_type": "unitTson" - } + "frontend_type": "unitTson", + }, } async def get_audio(self, text: str) -> str: """异步方法获取语音文件路径""" headers = { "Content-Type": "application/json", - "Authorization": f"Bearer; {self.api_key}" + "Authorization": f"Bearer; {self.api_key}", } - + payload = self._build_request_payload(text) - + logger.debug(f"请求头: {headers}") logger.debug(f"请求 URL: {self.api_base}") logger.debug(f"请求体: {json.dumps(payload, ensure_ascii=False)[:100]}...") - + try: async with aiohttp.ClientSession() as session: async with session.post( self.api_base, - data=json.dumps(payload), + data=json.dumps(payload), headers=headers, - timeout=self.timeout + timeout=self.timeout, ) as response: logger.debug(f"响应状态码: {response.status}") - + response_text = await response.text() logger.debug(f"响应内容: {response_text[:200]}...") - + if response.status == 200: resp_data = json.loads(response_text) - + if "data" in resp_data: audio_data = base64.b64decode(resp_data["data"]) - + os.makedirs("data/temp", exist_ok=True) - + file_path = f"data/temp/volcengine_tts_{uuid.uuid4()}.mp3" - + loop = asyncio.get_running_loop() await loop.run_in_executor( - None, - lambda: open(file_path, "wb").write(audio_data) + None, lambda: open(file_path, "wb").write(audio_data) ) - + return file_path else: error_msg = resp_data.get("message", "未知错误") raise Exception(f"火山引擎 TTS API 返回错误: {error_msg}") else: - raise Exception(f"火山引擎 TTS API 请求失败: {response.status}, {response_text}") - + raise Exception( + f"火山引擎 TTS API 请求失败: {response.status}, {response_text}" + ) + except Exception as e: error_details = traceback.format_exc() logger.debug(f"火山引擎 TTS 异常详情: {error_details}") - raise Exception(f"火山引擎 TTS 异常: {str(e)}") \ No newline at end of file + raise Exception(f"火山引擎 TTS 异常: {str(e)}") diff --git a/astrbot/core/provider/sources/zhipu_source.py b/astrbot/core/provider/sources/zhipu_source.py index 428dee8f4..cf52e95fc 100644 --- a/astrbot/core/provider/sources/zhipu_source.py +++ b/astrbot/core/provider/sources/zhipu_source.py @@ -28,6 +28,7 @@ class ProviderZhipu(ProviderOpenAIOfficial): func_tool: FuncCall = None, contexts=None, system_prompt=None, + model=None, **kwargs, ) -> LLMResponse: if contexts is None: @@ -38,7 +39,7 @@ class ProviderZhipu(ProviderOpenAIOfficial): context_query = [*contexts, new_record] model_cfgs: dict = self.provider_config.get("model_config", {}) - model = self.get_model() + model = model or self.get_model() # glm-4v-flash 只支持一张图片 if model.lower() == "glm-4v-flash" and image_urls and len(context_query) > 1: logger.debug("glm-4v-flash 只支持一张图片,将只保留最后一张图片") diff --git a/astrbot/core/star/__init__.py b/astrbot/core/star/__init__.py index f871c67c8..3337e4c25 100644 --- a/astrbot/core/star/__init__.py +++ b/astrbot/core/star/__init__.py @@ -1,4 +1,4 @@ -from .star import StarMetadata +from .star import StarMetadata, star_map from .star_manager import PluginManager from .context import Context from astrbot.core.provider import Provider @@ -14,12 +14,22 @@ class Star(CommandParserMixin): StarTools.initialize(context) self.context = context - async def text_to_image(self, text: str, return_url=True) -> str: + def __init_subclass__(cls, **kwargs): + super().__init_subclass__(**kwargs) + metadata = StarMetadata( + star_cls_type=cls, + module_path=cls.__module__, + ) + star_map[cls.__module__] = metadata + + @staticmethod + async def text_to_image(text: str, return_url=True) -> str: """将文本转换为图片""" return await html_renderer.render_t2i(text, return_url=return_url) + @staticmethod async def html_render( - self, tmpl: str, data: dict, return_url=True, options: dict = None + tmpl: str, data: dict, return_url=True, options: dict = None ) -> str: """渲染 HTML""" return await html_renderer.render_custom_template( diff --git a/astrbot/core/star/filter/platform_adapter_type.py b/astrbot/core/star/filter/platform_adapter_type.py index 0926cc337..fffaf8553 100644 --- a/astrbot/core/star/filter/platform_adapter_type.py +++ b/astrbot/core/star/filter/platform_adapter_type.py @@ -8,22 +8,48 @@ from typing import Union class PlatformAdapterType(enum.Flag): AIOCQHTTP = enum.auto() QQOFFICIAL = enum.auto() - VCHAT = enum.auto() GEWECHAT = enum.auto() TELEGRAM = enum.auto() WECOM = enum.auto() LARK = enum.auto() - ALL = AIOCQHTTP | QQOFFICIAL | VCHAT | GEWECHAT | TELEGRAM | WECOM | LARK + WECHATPADPRO = enum.auto() + DINGTALK = enum.auto() + DISCORD = enum.auto() + SLACK = enum.auto() + KOOK = enum.auto() + VOCECHAT = enum.auto() + WEIXIN_OFFICIAL_ACCOUNT = enum.auto() + ALL = ( + AIOCQHTTP + | QQOFFICIAL + | GEWECHAT + | TELEGRAM + | WECOM + | LARK + | WECHATPADPRO + | DINGTALK + | DISCORD + | SLACK + | KOOK + | VOCECHAT + | WEIXIN_OFFICIAL_ACCOUNT + ) ADAPTER_NAME_2_TYPE = { "aiocqhttp": PlatformAdapterType.AIOCQHTTP, "qq_official": PlatformAdapterType.QQOFFICIAL, - "vchat": PlatformAdapterType.VCHAT, "gewechat": PlatformAdapterType.GEWECHAT, "telegram": PlatformAdapterType.TELEGRAM, "wecom": PlatformAdapterType.WECOM, "lark": PlatformAdapterType.LARK, + "dingtalk": PlatformAdapterType.DINGTALK, + "discord": PlatformAdapterType.DISCORD, + "slack": PlatformAdapterType.SLACK, + "kook": PlatformAdapterType.KOOK, + "wechatpadpro": PlatformAdapterType.WECHATPADPRO, + "vocechat": PlatformAdapterType.VOCECHAT, + "weixin_official_account": PlatformAdapterType.WEIXIN_OFFICIAL_ACCOUNT, } diff --git a/astrbot/core/star/register/star.py b/astrbot/core/star/register/star.py index 01ff9adaa..7e5f89fd2 100644 --- a/astrbot/core/star/register/star.py +++ b/astrbot/core/star/register/star.py @@ -1,9 +1,15 @@ -from ..star import star_registry, StarMetadata, star_map +import warnings + +_warned_register_star = False def register_star(name: str, author: str, desc: str, version: str, repo: str = None): """注册一个插件(Star)。 + [DEPRECATED] 该装饰器已废弃,将在未来版本中移除。 + 在 v3.5.19 版本之后(不含),您不需要使用该装饰器来装饰插件类, + AstrBot 会自动识别继承自 Star 的类并将其作为插件类加载。 + Args: name: 插件名称。 author: 作者。 @@ -21,18 +27,16 @@ def register_star(name: str, author: str, desc: str, version: str, repo: str = N 帮助信息会被自动提取。使用 `/plugin <插件名> 可以查看帮助信息。` """ - def decorator(cls): - star_metadata = StarMetadata( - name=name, - author=author, - desc=desc, - version=version, - repo=repo, - star_cls_type=cls, - module_path=cls.__module__, + global _warned_register_star + if not _warned_register_star: + _warned_register_star = True + warnings.warn( + "The 'register_star' decorator is deprecated and will be removed in a future version.", + DeprecationWarning, + stacklevel=2, ) - star_registry.append(star_metadata) - star_map[cls.__module__] = star_metadata + + def decorator(cls): return cls return decorator diff --git a/astrbot/core/star/star.py b/astrbot/core/star/star.py index 10cf90c8b..bc8f3d404 100644 --- a/astrbot/core/star/star.py +++ b/astrbot/core/star/star.py @@ -1,12 +1,12 @@ from __future__ import annotations -from types import ModuleType -from typing import List, Dict from dataclasses import dataclass, field +from types import ModuleType + from astrbot.core.config import AstrBotConfig -star_registry: List[StarMetadata] = [] -star_map: Dict[str, StarMetadata] = {} +star_registry: list[StarMetadata] = [] +star_map: dict[str, StarMetadata] = {} """key 是模块路径,__module__""" @@ -18,22 +18,27 @@ class StarMetadata: 当 activated 为 False 时,star_cls 可能为 None,请不要在插件未激活时调用 star_cls 的方法。 """ - name: str - author: str # 插件作者 - desc: str # 插件简介 - version: str # 插件版本 - repo: str = None # 插件仓库地址 + name: str | None = None + """插件名""" + author: str | None = None + """插件作者""" + desc: str | None = None + """插件简介""" + version: str | None = None + """插件版本""" + repo: str | None = None + """插件仓库地址""" - star_cls_type: type = None + star_cls_type: type | None = None """插件的类对象的类型""" - module_path: str = None + module_path: str | None = None """插件的模块路径""" - star_cls: object = None + star_cls: object | None = None """插件的类对象""" - module: ModuleType = None + module: ModuleType | None = None """插件的模块对象""" - root_dir_name: str = None + root_dir_name: str | None = None """插件的目录名称""" reserved: bool = False """是否是 AstrBot 的保留插件""" @@ -41,13 +46,13 @@ class StarMetadata: activated: bool = True """是否被激活""" - config: AstrBotConfig = None + config: AstrBotConfig | None = None """插件配置""" - star_handler_full_names: List[str] = field(default_factory=list) + star_handler_full_names: list[str] = field(default_factory=list) """注册的 Handler 的全名列表""" - supported_platforms: Dict[str, bool] = field(default_factory=dict) + supported_platforms: dict[str, bool] = field(default_factory=dict) """插件支持的平台ID字典,key为平台ID,value为是否支持""" def __str__(self) -> str: diff --git a/astrbot/core/star/star_manager.py b/astrbot/core/star/star_manager.py index 3dd4cd1cf..b8365ed61 100644 --- a/astrbot/core/star/star_manager.py +++ b/astrbot/core/star/star_manager.py @@ -11,7 +11,6 @@ import os import sys import traceback from types import ModuleType -from typing import List import yaml @@ -119,7 +118,8 @@ class PluginManager: reloaded_plugins.add(plugin_name) break - def _get_classes(self, arg: ModuleType): + @staticmethod + def _get_classes(arg: ModuleType): """获取指定模块(可以理解为一个 python 文件)下所有的类""" classes = [] clsmembers = inspect.getmembers(arg, inspect.isclass) @@ -129,7 +129,8 @@ class PluginManager: break return classes - def _get_modules(self, path): + @staticmethod + def _get_modules(path): modules = [] dirs = os.listdir(path) @@ -155,7 +156,7 @@ class PluginManager: ) return modules - def _get_plugin_modules(self) -> List[dict]: + def _get_plugin_modules(self) -> list[dict]: plugins = [] if os.path.exists(self.plugin_store_path): plugins.extend(self._get_modules(self.plugin_store_path)) @@ -189,7 +190,8 @@ class PluginManager: except Exception as e: logger.error(f"更新插件 {p} 的依赖失败。Code: {str(e)}") - def _load_plugin_metadata(self, plugin_path: str, plugin_obj=None) -> StarMetadata: + @staticmethod + def _load_plugin_metadata(plugin_path: str, plugin_obj=None) -> StarMetadata: """v3.4.0 以前的方式载入插件元数据 先寻找 metadata.yaml 文件,如果不存在,则使用插件对象的 info() 函数获取元数据。 @@ -228,8 +230,9 @@ class PluginManager: return metadata + @staticmethod def _get_plugin_related_modules( - self, plugin_root_dir: str, is_reserved: bool = False + plugin_root_dir: str, is_reserved: bool = False ) -> list[str]: """获取与指定插件相关的所有已加载模块名 @@ -435,7 +438,7 @@ class PluginManager: ) if path in star_map: - # 通过装饰器的方式注册插件 + # 通过__init__subclass__注册插件 metadata = star_map[path] try: @@ -504,6 +507,8 @@ class PluginManager: if func_tool.name in inactivated_llm_tools: func_tool.active = False + star_registry.append(metadata) + else: # v3.4.0 以前的方式注册插件 logger.debug( @@ -775,7 +780,8 @@ class PluginManager: plugin.activated = False - async def _terminate_plugin(self, star_metadata: StarMetadata): + @staticmethod + async def _terminate_plugin(star_metadata: StarMetadata): """终止插件,调用插件的 terminate() 和 __del__() 方法""" logger.info(f"正在终止插件 {star_metadata.name} ...") diff --git a/astrbot/core/utils/tencent_record_helper.py b/astrbot/core/utils/tencent_record_helper.py index 9d0552c1e..2c97a01ed 100644 --- a/astrbot/core/utils/tencent_record_helper.py +++ b/astrbot/core/utils/tencent_record_helper.py @@ -117,7 +117,7 @@ async def audio_to_tencent_silk_base64(audio_path: str) -> tuple[str, float]: try: import pilk except ImportError as e: - raise Exception("未安装 pysilk,请执行: pip install pysilk") from e + raise Exception("未安装 pilk: pip install pilk") from e temp_dir = os.path.join(get_astrbot_data_path(), "temp") os.makedirs(temp_dir, exist_ok=True) diff --git a/astrbot/dashboard/routes/chat.py b/astrbot/dashboard/routes/chat.py index 270c92b44..b704a8888 100644 --- a/astrbot/dashboard/routes/chat.py +++ b/astrbot/dashboard/routes/chat.py @@ -2,7 +2,7 @@ import uuid import json import os from .route import Route, Response, RouteContext -from astrbot.core import web_chat_queue, web_chat_back_queue +from astrbot.core.platform.sources.webchat.webchat_queue_mgr import webchat_queue_mgr from quart import request, Response as QuartResponse, g, make_response from astrbot.core.db import BaseDatabase import asyncio @@ -21,7 +21,6 @@ class ChatRoute(Route): super().__init__(context) self.routes = { "/chat/send": ("POST", self.chat), - "/chat/listen": ("GET", self.listener), "/chat/new_conversation": ("GET", self.new_conversation), "/chat/conversations": ("GET", self.get_conversations), "/chat/get_conversation": ("GET", self.get_conversation), @@ -40,9 +39,6 @@ class ChatRoute(Route): self.supported_imgs = ["jpg", "jpeg", "png", "gif", "webp"] - self.curr_user_cid = {} - self.curr_chat_sse = {} - async def status(self): has_llm_enabled = ( self.core_lifecycle.provider_manager.curr_provider_inst is not None @@ -124,6 +120,8 @@ class ChatRoute(Route): conversation_id = post_data["conversation_id"] image_url = post_data.get("image_url") audio_url = post_data.get("audio_url") + selected_provider = post_data.get("selected_provider") + selected_model = post_data.get("selected_model") if not message and not image_url and not audio_url: return ( Response() @@ -133,21 +131,10 @@ class ChatRoute(Route): if not conversation_id: return Response().error("conversation_id is empty").__dict__ - self.curr_user_cid[username] = conversation_id + # Get conversation-specific queues + back_queue = webchat_queue_mgr.get_or_create_back_queue(conversation_id) - await web_chat_queue.put( - ( - username, - conversation_id, - { - "message": message, - "image_url": image_url, # list - "audio_url": audio_url, - }, - ) - ) - - # 持久化 + # append user message conversation = self.db.get_conversation_by_user_id(username, conversation_id) try: history = json.loads(conversation.history) @@ -164,30 +151,12 @@ class ChatRoute(Route): username, conversation_id, history=json.dumps(history) ) - return Response().ok().__dict__ - - async def listener(self): - """一直保持长连接""" - - username = g.get("username", "guest") - - if username in self.curr_chat_sse: - return Response().error("Already connected").__dict__ - - self.curr_chat_sse[username] = None - - heartbeat = json.dumps({"type": "heartbeat", "data": "ping"}) - async def stream(): try: - yield f"data: {heartbeat}\n\n" # 心跳包 while True: try: - result = await asyncio.wait_for( - web_chat_back_queue.get(), timeout=10 - ) # 设置超时时间为5秒 + result = await asyncio.wait_for(back_queue.get(), timeout=10) except asyncio.TimeoutError: - yield f"data: {heartbeat}\n\n" # 心跳包 continue if not result: @@ -197,19 +166,16 @@ class ChatRoute(Route): type = result.get("type") cid = result.get("cid") streaming = result.get("streaming", False) - if cid != self.curr_user_cid.get(username): - # 丢弃 - continue + chain_type = result.get("chain_type") yield f"data: {json.dumps(result, ensure_ascii=False)}\n\n" await asyncio.sleep(0.05) if streaming and type != "end": - continue - - if type == "update_title": + # If the result is still streaming, we continue to wait for more data continue if result_text: + # append bot message conversation = self.db.get_conversation_by_user_id( username, cid ) @@ -222,11 +188,31 @@ class ChatRoute(Route): self.db.update_conversation( username, cid, history=json.dumps(history) ) + if chain_type not in ["tool_call", "tool_call_result"]: + # If the result is not a tool call or tool call result, + # we can break the loop and end the stream + break + except BaseException as _: logger.debug(f"用户 {username} 断开聊天长连接。") - self.curr_chat_sse.pop(username) return + # Put message to conversation-specific queue + chat_queue = webchat_queue_mgr.get_or_create_queue(conversation_id) + await chat_queue.put( + ( + username, + conversation_id, + { + "message": message, + "image_url": image_url, # list + "audio_url": audio_url, + "selected_provider": selected_provider, + "selected_model": selected_model, + }, + ) + ) + response = await make_response( stream(), { @@ -236,7 +222,6 @@ class ChatRoute(Route): "Connection": "keep-alive", }, ) - response.timeout = None return response async def delete_conversation(self): @@ -245,6 +230,8 @@ class ChatRoute(Route): if not conversation_id: return Response().error("Missing key: conversation_id").__dict__ + # Clean up queues when deleting conversation + webchat_queue_mgr.remove_queues(conversation_id) self.db.delete_conversation(username, conversation_id) return Response().ok().__dict__ @@ -279,6 +266,4 @@ class ChatRoute(Route): conversation = self.db.get_conversation_by_user_id(username, conversation_id) - self.curr_user_cid[username] = conversation_id - return Response().ok(data=conversation).__dict__ diff --git a/astrbot/dashboard/routes/config.py b/astrbot/dashboard/routes/config.py index c225c762a..1dbe4de4a 100644 --- a/astrbot/dashboard/routes/config.py +++ b/astrbot/dashboard/routes/config.py @@ -9,6 +9,7 @@ from astrbot.core.platform.register import platform_registry from astrbot.core.provider.register import provider_registry from astrbot.core.star.star import star_registry from astrbot.core import logger +from astrbot.core.provider import Provider import asyncio @@ -166,8 +167,9 @@ class ConfigRoute(Route): "/config/provider/update": ("POST", self.post_update_provider), "/config/provider/delete": ("POST", self.post_delete_provider), "/config/llmtools": ("GET", self.get_llm_tools), - "/config/provider/check_status": ("GET", self.check_all_providers_status), + "/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_session_seperate": ( "GET", lambda: Response() @@ -256,33 +258,37 @@ class ConfigRoute(Route): ) return status_info - async def check_all_providers_status(self): - """ - API 接口: 检查所有 LLM Providers 的状态 - """ - logger.info("API call received: /config/provider/check_status") + def _error_response(self, message: str, status_code: int = 500, log_fn=logger.error): + log_fn(message) + # 记录更详细的traceback信息,但只在是严重错误时 + if status_code == 500: + log_fn(traceback.format_exc()) + return Response().error(message, status_code=status_code).__dict__ + + async def check_one_provider_status(self): + """API: check a single LLM Provider's status by id""" + provider_id = request.args.get("id") + if not provider_id: + return self._error_response("Missing provider_id parameter", 400, logger.warning) + + logger.info(f"API call: /config/provider/check_one id={provider_id}") try: - all_providers: typing.List = ( - self.core_lifecycle.star_context.get_all_providers() + all_providers = self.core_lifecycle.star_context.get_all_providers() + # replace manual loop with next(filter(...)) + target = next( + (p for p in all_providers if p.provider_config.get("id") == provider_id), + None ) - logger.debug(f"Found {len(all_providers)} providers to check.") + if not target: + return self._error_response(f"Provider with id '{provider_id}' not found", 404, logger.warning) - if not all_providers: - logger.info("No providers found to check.") - return Response().ok([]).__dict__ + result = await self._test_single_provider(target) + return Response().ok(result).__dict__ - tasks = [self._test_single_provider(p) for p in all_providers] - logger.debug(f"Created {len(tasks)} tasks for concurrent provider checks.") - - results = await asyncio.gather(*tasks) - logger.info(f"Provider status check completed. Results: {results}") - - return Response().ok(results).__dict__ except Exception as e: - logger.error(f"Critical error in check_all_providers_status: {str(e)}") - logger.error(traceback.format_exc()) - return ( - Response().error(f"检查 Provider 状态时发生严重错误: {str(e)}").__dict__ + return self._error_response( + f"Critical error checking provider {provider_id}: {e}", + 500 ) async def get_configs(self): @@ -319,6 +325,28 @@ class ConfigRoute(Route): provider_list.append(provider) return Response().ok(provider_list).__dict__ + async def get_provider_model_list(self): + """获取指定提供商的模型列表""" + provider_id = request.args.get("provider_id", None) + if not provider_id: + return Response().error("缺少参数 provider_id").__dict__ + + prov_mgr = self.core_lifecycle.provider_manager + provider: Provider | None = prov_mgr.inst_map.get(provider_id, None) + if not provider: + return Response().error(f"未找到 ID 为 {provider_id} 的提供商").__dict__ + + try: + models = await provider.get_models() + ret = { + "models": models, + "provider_id": provider_id, + } + return Response().ok(ret).__dict__ + except Exception as e: + logger.error(traceback.format_exc()) + return Response().error(str(e)).__dict__ + async def post_astrbot_configs(self): post_configs = await request.json try: diff --git a/astrbot/dashboard/routes/conversation.py b/astrbot/dashboard/routes/conversation.py index aa8b0af36..d73e6186a 100644 --- a/astrbot/dashboard/routes/conversation.py +++ b/astrbot/dashboard/routes/conversation.py @@ -29,6 +29,7 @@ class ConversationRoute(Route): ), } self.db_helper = db_helper + self.core_lifecycle = core_lifecycle self.register_routes() async def list_conversations(self): @@ -165,11 +166,9 @@ class ConversationRoute(Route): if not user_id or not cid: return Response().error("缺少必要参数: user_id 和 cid").__dict__ - conversation = self.db_helper.get_conversation_by_user_id(user_id, cid) - if not conversation: - return Response().error("对话不存在").__dict__ - self.db_helper.delete_conversation(user_id, cid) - + self.core_lifecycle.conversation_manager.delete_conversation( + unified_msg_origin=user_id, conversation_id=cid + ) return Response().ok({"message": "对话删除成功"}).__dict__ except Exception as e: diff --git a/changelogs/v3.5.19.md b/changelogs/v3.5.19.md new file mode 100644 index 000000000..cb821cef6 --- /dev/null +++ b/changelogs/v3.5.19.md @@ -0,0 +1,10 @@ +# What's Changed + +1. 修复: 通过 provider 指令设置提供商,重启后失效 +2. 新增: WebChat 支持直接选择提供商和模型 +3. 优化: WebUI 视觉效果、WebChat 视觉效果 +4. 优化: WebUI 测试提供商功能 +5. 优化: 修复潜在的 README XSS 注入问题 +6. 修复: WechatPadPro 授权码提取逻辑以适配上游新版本,并提高安全性 +7. 修复: Gemini 下,多轮工具调用时可能报错的问题 +8. 其他修复与优化 \ No newline at end of file diff --git a/dashboard/package.json b/dashboard/package.json index 7a5dd44a5..4f7ca6753 100644 --- a/dashboard/package.json +++ b/dashboard/package.json @@ -26,6 +26,7 @@ "js-md5": "^0.8.3", "lodash": "4.17.21", "marked": "^15.0.7", + "markdown-it": "^14.1.0", "pinia": "2.1.6", "remixicon": "3.5.0", "vee-validate": "4.11.3", diff --git a/dashboard/src/components/chat/ProviderModelSelector.vue b/dashboard/src/components/chat/ProviderModelSelector.vue new file mode 100644 index 000000000..7509b5295 --- /dev/null +++ b/dashboard/src/components/chat/ProviderModelSelector.vue @@ -0,0 +1,353 @@ + + + + + diff --git a/dashboard/src/components/shared/ExtensionCard.vue b/dashboard/src/components/shared/ExtensionCard.vue index 35666b740..8a0075173 100644 --- a/dashboard/src/components/shared/ExtensionCard.vue +++ b/dashboard/src/components/shared/ExtensionCard.vue @@ -49,6 +49,11 @@ const reloadExtension = () => { }; const $confirm = inject("$confirm"); + +const installExtension = async () => { + emit('install', props.extension); +}; + const uninstallExtension = async () => { if (typeof $confirm !== "function") { console.error(tm("card.errors.confirmNotRegistered")); @@ -117,6 +122,10 @@ const viewReadme = () => { {{ extension.handlers?.length }}{{ tm("card.status.handlersCount") }} + + {{ tag === 'danger' ? tm('tags.danger') : tag }} +
@@ -139,7 +148,7 @@ const viewReadme = () => { + @click="installExtension"> @@ -200,6 +209,7 @@ const viewReadme = () => { + diff --git a/dashboard/src/components/shared/ItemCardGrid.vue b/dashboard/src/components/shared/ItemCardGrid.vue index 71841d2ab..5176c186b 100644 --- a/dashboard/src/components/shared/ItemCardGrid.vue +++ b/dashboard/src/components/shared/ItemCardGrid.vue @@ -9,10 +9,10 @@ - +
- {{ getItemTitle(item) }} + {{ getItemTitle(item) }}