From 7543dd2e9df28bb1e4cfb45429414b0ab1a46ee4 Mon Sep 17 00:00:00 2001 From: LIghtJUNction Date: Sat, 11 Apr 2026 18:27:15 +0800 Subject: [PATCH] refactor: remove TUI module and fix type annotations - Remove TUI platform adapter and web API routes - Add DiscordEmbed, DiscordButton, DiscordReference, DiscordView to ComponentType enum - Fix Platform.terminate() and Platform.get_client() with proper implementations - Fix AstrBotMessage.raw_message type from object to Any - Add Any type annotation to CONFIG_METADATA_2 --- astrbot/core/config/default.py | 2 +- astrbot/core/message/components.py | 5 + astrbot/core/platform/astrbot_message.py | 3 +- astrbot/core/platform/manager.py | 6 - astrbot/core/platform/platform.py | 4 +- astrbot/core/platform/sources/tui/__init__.py | 5 - .../core/platform/sources/tui/tui_adapter.py | 223 ------ .../core/platform/sources/tui/tui_event.py | 203 ----- .../platform/sources/tui/tui_queue_mgr.py | 164 ----- astrbot/dashboard/routes/__init__.py | 2 - astrbot/dashboard/routes/config.py | 8 +- astrbot/dashboard/routes/tui_chat.py | 691 ------------------ astrbot/dashboard/server.py | 2 - 13 files changed, 15 insertions(+), 1303 deletions(-) delete mode 100644 astrbot/core/platform/sources/tui/__init__.py delete mode 100644 astrbot/core/platform/sources/tui/tui_adapter.py delete mode 100644 astrbot/core/platform/sources/tui/tui_event.py delete mode 100644 astrbot/core/platform/sources/tui/tui_queue_mgr.py delete mode 100644 astrbot/dashboard/routes/tui_chat.py diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index 7ec8cccf8..8f5cb8998 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -317,7 +317,7 @@ WebUI 的配置文件在 `CONFIG_METADATA_3` 中。 未来将会逐步淘汰此配置元数据。 """ -CONFIG_METADATA_2 = { +CONFIG_METADATA_2: Any = { "platform_group": { "metadata": { "platform": { diff --git a/astrbot/core/message/components.py b/astrbot/core/message/components.py index 5cc2d135a..e39c231b6 100644 --- a/astrbot/core/message/components.py +++ b/astrbot/core/message/components.py @@ -73,6 +73,11 @@ class ComponentType(str, Enum): Location = "Location" # TODO Music = "Music" Json = "Json" + # Discord-specific component types + DiscordEmbed = "DiscordEmbed" + DiscordButton = "DiscordButton" + DiscordReference = "DiscordReference" + DiscordView = "DiscordView" Unknown = "Unknown" diff --git a/astrbot/core/platform/astrbot_message.py b/astrbot/core/platform/astrbot_message.py index 8d682d728..14a894649 100644 --- a/astrbot/core/platform/astrbot_message.py +++ b/astrbot/core/platform/astrbot_message.py @@ -1,5 +1,6 @@ import time from dataclasses import dataclass +from typing import Any from astrbot.core.message.components import BaseMessageComponent @@ -55,7 +56,7 @@ class AstrBotMessage: sender: MessageMember # 发送者 message: list[BaseMessageComponent] # 消息链使用 Nakuru 的消息链格式 message_str: str # 最直观的纯文本消息字符串 - raw_message: object + raw_message: Any timestamp: int # 消息时间戳 def __init__(self) -> None: diff --git a/astrbot/core/platform/manager.py b/astrbot/core/platform/manager.py index f7bfababf..b21a34ee8 100644 --- a/astrbot/core/platform/manager.py +++ b/astrbot/core/platform/manager.py @@ -10,7 +10,6 @@ from astrbot.core.utils.webhook_utils import ensure_platform_webhook_config from .platform import Platform, PlatformStatus from .register import platform_cls_map -from .sources.tui.tui_adapter import TUIAdapter from .sources.webchat.webchat_adapter import WebChatAdapter PLATFORM_ADAPTER_MODULES: dict[str, str] = { @@ -118,11 +117,6 @@ class PlatformManager: self.platform_insts.append(webchat_inst) self._start_platform_task("webchat", webchat_inst) - # TUI - tui_inst = TUIAdapter({}, self.settings, self.event_queue) - self.platform_insts.append(tui_inst) - self._start_platform_task("tui", tui_inst) - async def load_platform(self, platform_config: dict) -> None: """实例化一个平台""" # 动态导入 diff --git a/astrbot/core/platform/platform.py b/astrbot/core/platform/platform.py index eb233a666..9924ebdf4 100644 --- a/astrbot/core/platform/platform.py +++ b/astrbot/core/platform/platform.py @@ -123,6 +123,7 @@ class Platform(abc.ABC): async def terminate(self) -> None: """终止一个平台的运行实例。""" + self._status = PlatformStatus.STOPPED @abc.abstractmethod def meta(self) -> PlatformMetadata: @@ -144,8 +145,9 @@ class Platform(abc.ABC): """提交一个事件到事件队列。""" self._event_queue.put_nowait(event) - def get_client(self) -> object: + def get_client(self) -> object | None: """获取平台的客户端对象。""" + return None async def webhook_callback(self, request: Any) -> Any: """统一 Webhook 回调入口。 diff --git a/astrbot/core/platform/sources/tui/__init__.py b/astrbot/core/platform/sources/tui/__init__.py deleted file mode 100644 index 6ac1a858a..000000000 --- a/astrbot/core/platform/sources/tui/__init__.py +++ /dev/null @@ -1,5 +0,0 @@ -from .tui_adapter import TUIAdapter -from .tui_event import TUIMessageEvent -from .tui_queue_mgr import TUIQueueMgr, tui_queue_mgr - -__all__ = ["TUIAdapter", "TUIMessageEvent", "TUIQueueMgr", "tui_queue_mgr"] diff --git a/astrbot/core/platform/sources/tui/tui_adapter.py b/astrbot/core/platform/sources/tui/tui_adapter.py deleted file mode 100644 index 64ee3ad11..000000000 --- a/astrbot/core/platform/sources/tui/tui_adapter.py +++ /dev/null @@ -1,223 +0,0 @@ -import asyncio -import os -import time -from collections.abc import Callable, Coroutine -from typing import Any - -from astrbot import logger -from astrbot.core import db_helper -from astrbot.core.db.po import PlatformMessageHistory -from astrbot.core.message.message_event_result import MessageChain -from astrbot.core.platform import ( - AstrBotMessage, - MessageMember, - MessageType, - Platform, - PlatformMetadata, -) -from astrbot.core.platform.astr_message_event import MessageSesion -from astrbot.core.platform.register import register_platform_adapter -from astrbot.core.platform.sources.webchat.message_parts_helper import ( - message_chain_to_storage_message_parts, - parse_webchat_message_parts, -) -from astrbot.core.utils.astrbot_path import get_astrbot_data_path - -from .tui_event import TUIMessageEvent -from .tui_queue_mgr import TUIQueueMgr, tui_queue_mgr - - -def _extract_conversation_id(session_id: str) -> str: - """Extract raw TUI conversation id from event/session id.""" - if session_id.startswith("tui!"): - parts = session_id.split("!", 2) - if len(parts) == 3: - return parts[2] - return session_id - - -class QueueListener: - def __init__( - self, - tui_queue_mgr: TUIQueueMgr, - callback: Callable, - stop_event: asyncio.Event, - ) -> None: - self.tui_queue_mgr = tui_queue_mgr - self.callback = callback - self.stop_event = stop_event - - async def run(self) -> None: - """Register callback and keep adapter task alive.""" - self.tui_queue_mgr.set_listener(self.callback) - try: - await self.stop_event.wait() - finally: - await self.tui_queue_mgr.clear_listener() - - -@register_platform_adapter("tui", "tui") -class TUIAdapter(Platform): - def __init__( - self, - platform_config: dict, - platform_settings: dict, - event_queue: asyncio.Queue, - ) -> None: - super().__init__(platform_config, event_queue) - self.settings = platform_settings - self.imgs_dir = os.path.join(get_astrbot_data_path(), "tui", "imgs") - self.attachments_dir = os.path.join(get_astrbot_data_path(), "attachments") - os.makedirs(self.imgs_dir, exist_ok=True) - os.makedirs(self.attachments_dir, exist_ok=True) - self.metadata = PlatformMetadata( - name="tui", - description="tui", - id="tui", - support_proactive_message=True, - ) - self._shutdown_event = asyncio.Event() - self._tui_queue_mgr = tui_queue_mgr - - async def send_by_session( - self, - session: MessageSesion, - message_chain: MessageChain, - ) -> None: - conversation_id = _extract_conversation_id(session.session_id) - active_request_ids = self._tui_queue_mgr.list_back_request_ids(conversation_id) - stream_request_ids = [ - req_id for req_id in active_request_ids if not req_id.startswith("ws_sub_") - ] - target_request_ids = stream_request_ids or active_request_ids - if not target_request_ids: - try: - await self._save_proactive_message(conversation_id, message_chain) - except Exception as e: - logger.error( - f"[TUIAdapter] Failed to save proactive message: {e}", - exc_info=True, - ) - await super().send_by_session(session, message_chain) - return - for request_id in target_request_ids: - await TUIMessageEvent._send( - request_id, - message_chain, - session.session_id, - streaming=True, - emit_complete=True, - ) - if not stream_request_ids: - try: - await self._save_proactive_message(conversation_id, message_chain) - except Exception as e: - logger.error( - f"[TUIAdapter] Failed to save proactive message: {e}", - exc_info=True, - ) - await super().send_by_session(session, message_chain) - - async def _save_proactive_message( - self, - conversation_id: str, - message_chain: MessageChain, - ) -> None: - message_parts = await message_chain_to_storage_message_parts( - message_chain, - insert_attachment=db_helper.insert_attachment, - attachments_dir=self.attachments_dir, - ) - if not message_parts: - return - await db_helper.insert_platform_message_history( - platform_id="tui", - user_id=conversation_id, - content={"type": "bot", "message": message_parts}, - sender_id="bot", - sender_name="bot", - ) - - async def _get_message_history( - self, - message_id: int, - ) -> PlatformMessageHistory | None: - return await db_helper.get_platform_message_history_by_id(message_id) - - async def _parse_message_parts( - self, - message_parts: list, - depth: int = 0, - max_depth: int = 1, - ) -> tuple[list, list[str]]: - """Parse message parts list, return message components and plain text lists.""" - - async def get_reply_parts( - message_id: Any, - ) -> tuple[list[dict], str | None, str | None] | None: - history = await self._get_message_history(message_id) - if not history or not history.content: - return None - reply_parts = history.content.get("message", []) - if not isinstance(reply_parts, list): - return None - return (reply_parts, history.sender_id, history.sender_name) - - components, text_parts, _ = await parse_webchat_message_parts( - message_parts, - strict=False, - include_empty_plain=True, - verify_media_path_exists=False, - reply_history_getter=get_reply_parts, - current_depth=depth, - max_reply_depth=max_depth, - cast_reply_id_to_str=False, - ) - return (components, text_parts) - - async def convert_message(self, data: tuple) -> AstrBotMessage: - username, cid, payload = data - abm = AstrBotMessage() - abm.self_id = "tui" - abm.sender = MessageMember(username, username) - abm.type = MessageType.FRIEND_MESSAGE - abm.session_id = f"tui!{username}!{cid}" - abm.message_id = payload.get("message_id") - message_parts = payload.get("message", []) - abm.message, message_str_parts = await self._parse_message_parts(message_parts) - logger.debug(f"TUIAdapter: {abm.message}") - abm.timestamp = int(time.time()) - abm.message_str = "".join(message_str_parts) - abm.raw_message = data - return abm - - def run(self) -> Coroutine[Any, Any, None]: - async def callback(data: tuple) -> None: - abm = await self.convert_message(data) - await self.handle_msg(abm) - - bot = QueueListener(self._tui_queue_mgr, callback, self._shutdown_event) - return bot.run() - - def meta(self) -> PlatformMetadata: - return self.metadata - - async def handle_msg(self, message: AstrBotMessage) -> None: - message_event = TUIMessageEvent( - message_str=message.message_str, - message_obj=message, - platform_meta=self.meta(), - session_id=message.session_id, - ) - _, _, payload = message.raw_message - message_event.set_extra("selected_provider", payload.get("selected_provider")) - message_event.set_extra("selected_model", payload.get("selected_model")) - message_event.set_extra( - "enable_streaming", - payload.get("enable_streaming", True), - ) - message_event.set_extra("action_type", payload.get("action_type")) - self.commit_event(message_event) - - async def terminate(self) -> None: - self._shutdown_event.set() diff --git a/astrbot/core/platform/sources/tui/tui_event.py b/astrbot/core/platform/sources/tui/tui_event.py deleted file mode 100644 index 4e2d19952..000000000 --- a/astrbot/core/platform/sources/tui/tui_event.py +++ /dev/null @@ -1,203 +0,0 @@ -import base64 -import json -import os -import shutil -import uuid - -import aiofiles - -from astrbot.api import logger -from astrbot.api.event import AstrMessageEvent, MessageChain -from astrbot.api.message_components import File, Image, Json, Plain, Record -from astrbot.core.utils.astrbot_path import get_astrbot_data_path - -from .tui_queue_mgr import tui_queue_mgr - -attachments_dir = os.path.join(get_astrbot_data_path(), "attachments") - - -def _extract_conversation_id(session_id: str) -> str: - """Extract raw TUI conversation id from event/session id.""" - if session_id.startswith("tui!"): - parts = session_id.split("!", 2) - if len(parts) == 3: - return parts[2] - return session_id - - -class TUIMessageEvent(AstrMessageEvent): - def __init__(self, message_str, message_obj, platform_meta, session_id) -> None: - super().__init__(message_str, message_obj, platform_meta, session_id) - os.makedirs(attachments_dir, exist_ok=True) - - @staticmethod - async def _send( - message_id: str, - message: MessageChain | None, - session_id: str, - streaming: bool = False, - emit_complete: bool = False, - ) -> str | None: - request_id = str(message_id) - conversation_id = _extract_conversation_id(session_id) - tui_back_queue = tui_queue_mgr.get_or_create_back_queue( - request_id, - conversation_id, - ) - if not message: - await tui_back_queue.put( - { - "type": "end", - "data": "", - "streaming": False, - "message_id": message_id, - }, - ) - return None - - data = "" - for comp in message.chain: - if isinstance(comp, Plain): - data = comp.text - await tui_back_queue.put( - { - "type": "plain", - "data": data, - "streaming": streaming, - "chain_type": message.type, - "message_id": message_id, - }, - ) - elif isinstance(comp, Json): - await tui_back_queue.put( - { - "type": "plain", - "data": json.dumps(comp.data, ensure_ascii=False), - "streaming": streaming, - "chain_type": message.type, - "message_id": message_id, - }, - ) - elif isinstance(comp, Image): - filename = f"{uuid.uuid4()!s}.jpg" - path = os.path.join(attachments_dir, filename) - image_base64 = await comp.convert_to_base64() - async with aiofiles.open(path, "wb") as f: - await f.write(base64.b64decode(image_base64)) - data = f"[IMAGE]{filename}" - await tui_back_queue.put( - { - "type": "image", - "data": data, - "streaming": streaming, - "message_id": message_id, - }, - ) - elif isinstance(comp, Record): - filename = f"{uuid.uuid4()!s}.wav" - path = os.path.join(attachments_dir, filename) - record_base64 = await comp.convert_to_base64() - async with aiofiles.open(path, "wb") as f: - await f.write(base64.b64decode(record_base64)) - data = f"[RECORD]{filename}" - await tui_back_queue.put( - { - "type": "record", - "data": data, - "streaming": streaming, - "message_id": message_id, - }, - ) - elif isinstance(comp, File): - file_path = await comp.get_file() - original_name = comp.name or os.path.basename(file_path) - ext = os.path.splitext(original_name)[1] or "" - filename = f"{uuid.uuid4()!s}{ext}" - dest_path = os.path.join(attachments_dir, filename) - shutil.copy2(file_path, dest_path) - data = f"[FILE]{filename}" - await tui_back_queue.put( - { - "type": "file", - "data": data, - "streaming": streaming, - "message_id": message_id, - }, - ) - else: - logger.debug(f"TUI ignores: {comp.type}") - - if emit_complete: - await tui_back_queue.put( - { - "type": "complete", - "data": data, - "streaming": streaming, - "chain_type": message.type, - "message_id": message_id, - }, - ) - - return data - - async def send(self, message: MessageChain | None) -> None: - message_id = self.message_obj.message_id - await TUIMessageEvent._send(message_id, message, session_id=self.session_id) - await super().send(MessageChain([])) - - async def send_streaming(self, generator, use_fallback: bool = False) -> None: - final_data = "" - reasoning_content = "" - message_id = self.message_obj.message_id - request_id = str(message_id) - conversation_id = _extract_conversation_id(self.session_id) - tui_back_queue = tui_queue_mgr.get_or_create_back_queue( - request_id, - conversation_id, - ) - async for chain in generator: - if chain.type == "audio_chunk": - audio_b64 = "" - text = None - - if chain.chain and isinstance(chain.chain[0], Plain): - audio_b64 = chain.chain[0].text - - if len(chain.chain) > 1 and isinstance(chain.chain[1], Json): - text = chain.chain[1].data.get("text") - - payload = { - "type": "audio_chunk", - "data": audio_b64, - "streaming": True, - "message_id": message_id, - } - if text: - payload["text"] = text - - await tui_back_queue.put(payload) - continue - - r = await TUIMessageEvent._send( - message_id=message_id, - message=chain, - session_id=self.session_id, - streaming=True, - ) - if not r: - continue - if chain.type == "reasoning": - reasoning_content += chain.get_plain_text() - else: - final_data += r - - await tui_back_queue.put( - { - "type": "complete", - "data": final_data, - "reasoning": reasoning_content, - "streaming": True, - "message_id": message_id, - }, - ) - await super().send_streaming(generator, use_fallback) diff --git a/astrbot/core/platform/sources/tui/tui_queue_mgr.py b/astrbot/core/platform/sources/tui/tui_queue_mgr.py deleted file mode 100644 index 5f14e9c46..000000000 --- a/astrbot/core/platform/sources/tui/tui_queue_mgr.py +++ /dev/null @@ -1,164 +0,0 @@ -import asyncio -from collections.abc import Awaitable, Callable - -from astrbot import logger - - -class TUIQueueMgr: - def __init__(self, queue_maxsize: int = 128, back_queue_maxsize: int = 512) -> None: - self.queues: dict[str, asyncio.Queue] = {} - """Conversation ID to asyncio.Queue mapping""" - self.back_queues: dict[str, asyncio.Queue] = {} - """Request ID to asyncio.Queue mapping for responses""" - self._conversation_back_requests: dict[str, set[str]] = {} - self._request_conversation: dict[str, str] = {} - self._queue_close_events: dict[str, asyncio.Event] = {} - self._listener_tasks: dict[str, asyncio.Task] = {} - self._listener_callback: Callable[[tuple], Awaitable[None]] | None = None - self.queue_maxsize = queue_maxsize - self.back_queue_maxsize = back_queue_maxsize - - 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(maxsize=self.queue_maxsize) - self._queue_close_events[conversation_id] = asyncio.Event() - self._start_listener_if_needed(conversation_id) - return self.queues[conversation_id] - - def get_or_create_back_queue( - self, - request_id: str, - conversation_id: str | None = None, - ) -> asyncio.Queue: - """Get or create a back queue for the given request ID""" - if request_id not in self.back_queues: - self.back_queues[request_id] = asyncio.Queue( - maxsize=self.back_queue_maxsize, - ) - if conversation_id: - self._request_conversation[request_id] = conversation_id - if conversation_id not in self._conversation_back_requests: - self._conversation_back_requests[conversation_id] = set() - self._conversation_back_requests[conversation_id].add(request_id) - return self.back_queues[request_id] - - def remove_back_queue(self, request_id: str) -> None: - """Remove back queue for the given request ID""" - self.back_queues.pop(request_id, None) - conversation_id = self._request_conversation.pop(request_id, None) - if conversation_id: - request_ids = self._conversation_back_requests.get(conversation_id) - if request_ids is not None: - request_ids.discard(request_id) - if not request_ids: - self._conversation_back_requests.pop(conversation_id, None) - - def remove_queues(self, conversation_id: str) -> None: - """Remove queues for the given conversation ID""" - for request_id in list( - self._conversation_back_requests.get(conversation_id, set()), - ): - self.remove_back_queue(request_id) - self._conversation_back_requests.pop(conversation_id, None) - self.remove_queue(conversation_id) - - def remove_queue(self, conversation_id: str) -> None: - """Remove input queue and listener for the given conversation ID""" - self.queues.pop(conversation_id, None) - - close_event = self._queue_close_events.pop(conversation_id, None) - if close_event is not None: - close_event.set() - - task = self._listener_tasks.pop(conversation_id, None) - if task is not None: - task.cancel() - - def list_back_request_ids(self, conversation_id: str) -> list[str]: - """List active back-queue request IDs for a conversation.""" - return list(self._conversation_back_requests.get(conversation_id, set())) - - def has_queue(self, conversation_id: str) -> bool: - """Check if a queue exists for the given conversation ID""" - return conversation_id in self.queues - - def set_listener( - self, - callback: Callable[[tuple], Awaitable[None]], - ) -> None: - self._listener_callback = callback - for conversation_id in list(self.queues.keys()): - self._start_listener_if_needed(conversation_id) - - async def clear_listener(self) -> None: - self._listener_callback = None - for close_event in list(self._queue_close_events.values()): - close_event.set() - self._queue_close_events.clear() - - listener_tasks = list(self._listener_tasks.values()) - for task in listener_tasks: - task.cancel() - if listener_tasks: - await asyncio.gather(*listener_tasks, return_exceptions=True) - self._listener_tasks.clear() - - def _start_listener_if_needed(self, conversation_id: str) -> None: - if self._listener_callback is None: - return - if conversation_id in self._listener_tasks: - task = self._listener_tasks[conversation_id] - if not task.done(): - return - queue = self.queues.get(conversation_id) - close_event = self._queue_close_events.get(conversation_id) - if queue is None or close_event is None: - return - task = asyncio.create_task( - self._listen_to_queue(conversation_id, queue, close_event), - name=f"tui_listener_{conversation_id}", - ) - self._listener_tasks[conversation_id] = task - task.add_done_callback( - lambda _: self._listener_tasks.pop(conversation_id, None), - ) - logger.debug(f"Started listener for TUI conversation: {conversation_id}") - - async def _listen_to_queue( - self, - conversation_id: str, - queue: asyncio.Queue, - close_event: asyncio.Event, - ) -> None: - while True: - get_task = asyncio.create_task(queue.get()) - close_task = asyncio.create_task(close_event.wait()) - try: - done, pending = await asyncio.wait( - {get_task, close_task}, - return_when=asyncio.FIRST_COMPLETED, - ) - for task in pending: - task.cancel() - if close_task in done: - break - data = get_task.result() - if self._listener_callback is None: - continue - try: - await self._listener_callback(data) - except Exception as e: - logger.error( - f"Error processing message from TUI conversation {conversation_id}: {e}", - ) - except asyncio.CancelledError: - break - finally: - if not get_task.done(): - get_task.cancel() - if not close_task.done(): - close_task.cancel() - - -tui_queue_mgr = TUIQueueMgr() diff --git a/astrbot/dashboard/routes/__init__.py b/astrbot/dashboard/routes/__init__.py index 86bdc0fa7..4af4e4178 100644 --- a/astrbot/dashboard/routes/__init__.py +++ b/astrbot/dashboard/routes/__init__.py @@ -23,7 +23,6 @@ from .static_file import StaticFileRoute from .subagent import SubAgentRoute from .t2i import T2iRoute from .tools import ToolsRoute -from .tui_chat import TUIChatRoute from .update import UpdateRoute __all__ = [ @@ -52,7 +51,6 @@ __all__ = [ "StaticFileRoute", "SubAgentRoute", "T2iRoute", - "TUIChatRoute", "ToolsRoute", "UpdateRoute", ] diff --git a/astrbot/dashboard/routes/config.py b/astrbot/dashboard/routes/config.py index 9b32bac56..64940e579 100644 --- a/astrbot/dashboard/routes/config.py +++ b/astrbot/dashboard/routes/config.py @@ -1289,7 +1289,7 @@ class ConfigRoute(Route): async def _get_astrbot_config(self): config = self.config - metadata = copy.deepcopy(CONFIG_METADATA_2) + metadata: Any = copy.deepcopy(CONFIG_METADATA_2) _pg: Any = metadata["platform_group"] _pg_meta: Any = _pg["metadata"] _platform_meta: Any = _pg_meta["platform"] @@ -1304,8 +1304,8 @@ class ConfigRoute(Route): _pg2: Any = metadata["platform_group"] _pg_meta2: Any = _pg2["metadata"] _platform_tmpl: Any = _pg_meta2["platform"] - platform_default_tmpl = _platform_tmpl["config_template"] - platform_i18n_translations = {} + platform_default_tmpl: Any = _platform_tmpl["config_template"] + platform_i18n_translations: dict[str, Any] = {} logo_registration_tasks = [] for platform in platform_registry: if platform.default_config_tmpl: @@ -1326,7 +1326,7 @@ class ConfigRoute(Route): await asyncio.gather(*logo_registration_tasks, return_exceptions=True) _provider_tmpl: Any = metadata["provider_group"] _provider_tmpl2: Any = _provider_tmpl["metadata"]["provider"] - provider_default_tmpl = _provider_tmpl2["config_template"] + provider_default_tmpl: Any = _provider_tmpl2["config_template"] for provider in provider_registry: if provider.default_config_tmpl: provider_default_tmpl[provider.type] = provider.default_config_tmpl diff --git a/astrbot/dashboard/routes/tui_chat.py b/astrbot/dashboard/routes/tui_chat.py deleted file mode 100644 index 6449e0ad3..000000000 --- a/astrbot/dashboard/routes/tui_chat.py +++ /dev/null @@ -1,691 +0,0 @@ -import asyncio -import json -import os -import uuid -from contextlib import asynccontextmanager -from pathlib import Path -from typing import Any - -import anyio -from quart import g, make_response, request, send_file - -from astrbot.core import logger -from astrbot.core.core_lifecycle import AstrBotCoreLifecycle -from astrbot.core.db import BaseDatabase -from astrbot.core.platform.message_type import MessageType -from astrbot.core.platform.sources.tui.tui_queue_mgr import tui_queue_mgr -from astrbot.core.platform.sources.webchat.message_parts_helper import ( - build_webchat_message_parts, - create_attachment_part_from_existing_file, - strip_message_parts_path_fields, - webchat_message_parts_have_content, -) -from astrbot.core.utils.active_event_registry import active_event_registry -from astrbot.core.utils.astrbot_path import get_astrbot_data_path -from astrbot.core.utils.datetime_utils import to_utc_isoformat - -from .route import Response, Route, RouteContext - - -@asynccontextmanager -async def track_conversation(convs: dict, conv_id: str): - convs[conv_id] = True - try: - yield - finally: - convs.pop(conv_id, None) - - -async def _poll_tui_stream_result(back_queue, username: str): - try: - result = await asyncio.wait_for(back_queue.get(), timeout=1) - except asyncio.TimeoutError: - return (None, False) - except asyncio.CancelledError: - logger.debug(f"[TUI] User {username} disconnected.") - return (None, True) - except Exception as e: - logger.error(f"TUI stream error: {e}") - return (None, False) - return (result, False) - - -def _resolve_path(path: str) -> Path: - return Path(path).resolve(strict=False) - - -class TUIChatRoute(Route): - def __init__( - self, - context: RouteContext, - db: BaseDatabase, - core_lifecycle: AstrBotCoreLifecycle, - ) -> None: - super().__init__(context) - self.routes = { - "/tui/chat": ("POST", self.chat), - "/tui/new_session": ("GET", self.new_session), - "/tui/sessions": ("GET", self.get_sessions), - "/tui/get_session": ("GET", self.get_session), - "/tui/stop": ("POST", self.stop_session), - "/tui/delete_session": ("GET", self.delete_tui_session), - "/tui/batch_delete_sessions": ("POST", self.batch_delete_sessions), - "/tui/update_session_display_name": ( - "POST", - self.update_session_display_name, - ), - "/tui/get_file": ("GET", self.get_file), - "/tui/get_attachment": ("GET", self.get_attachment), - "/tui/post_file": ("POST", self.post_file), - } - self.core_lifecycle = core_lifecycle - self.register_routes() - self.attachments_dir = os.path.join(get_astrbot_data_path(), "attachments") - os.makedirs(self.attachments_dir, exist_ok=True) - self.supported_imgs = ["jpg", "jpeg", "png", "gif", "webp"] - self.conv_mgr = core_lifecycle.conversation_manager - self.platform_history_mgr = core_lifecycle.platform_message_history_manager - self.db = db - self.umop_config_router = core_lifecycle.umop_config_router - self.running_convs: dict[str, bool] = {} - - async def get_file(self): - filename = request.args.get("filename") - if not filename: - return Response().error("Missing key: filename").to_json() - try: - file_path = os.path.join(self.attachments_dir, os.path.basename(filename)) - resolved_file_path = _resolve_path(file_path) - resolved_base_dir = _resolve_path(self.attachments_dir) - if not await anyio.Path(resolved_file_path).exists(): - return Response().error("File not found").to_json() - try: - resolved_file_path.relative_to(resolved_base_dir) - except ValueError: - return Response().error("Invalid file path").to_json() - filename_ext = os.path.splitext(filename)[1].lower() - if filename_ext == ".wav": - return await send_file(str(resolved_file_path), mimetype="audio/wav") - if filename_ext[1:] in self.supported_imgs: - return await send_file(str(resolved_file_path), mimetype="image/jpeg") - return await send_file(str(resolved_file_path)) - except (FileNotFoundError, OSError): - return Response().error("File access error").to_json() - - async def get_attachment(self): - """Get attachment file by attachment_id.""" - attachment_id = request.args.get("attachment_id") - if not attachment_id: - return Response().error("Missing key: attachment_id").to_json() - try: - attachment = await self.db.get_attachment_by_id(attachment_id) - if not attachment: - return Response().error("Attachment not found").to_json() - file_path = attachment.path - resolved_file_path = _resolve_path(file_path) - return await send_file( - str(resolved_file_path), - mimetype=attachment.mime_type, - ) - except (FileNotFoundError, OSError): - return Response().error("File access error").to_json() - - async def post_file(self): - """Upload a file and create an attachment record, return attachment_id.""" - post_data = await request.files - if "file" not in post_data: - return Response().error("Missing key: file").to_json() - file = post_data["file"] - filename = file.filename or f"{uuid.uuid4()!s}" - content_type = file.content_type or "application/octet-stream" - if content_type.startswith("image"): - attach_type = "image" - elif content_type.startswith("audio"): - attach_type = "record" - elif content_type.startswith("video"): - attach_type = "video" - else: - attach_type = "file" - path = os.path.join(self.attachments_dir, filename) - await file.save(path) - attachment = await self.db.insert_attachment( - path=path, - type=attach_type, - mime_type=content_type, - ) - if not attachment: - return Response().error("Failed to create attachment").to_json() - filename = os.path.basename(attachment.path) - return ( - Response() - .ok( - data={ - "attachment_id": attachment.attachment_id, - "filename": filename, - "type": attach_type, - }, - ) - .to_json() - ) - - async def _build_user_message_parts(self, message: str | list) -> list[dict]: - """Build user message parts list.""" - return await build_webchat_message_parts( - message, - get_attachment_by_id=self.db.get_attachment_by_id, - strict=False, - ) - - async def _create_attachment_from_file( - self, - filename: str, - attach_type: str, - ) -> dict | None: - """Create attachment from local file and return message part.""" - return await create_attachment_part_from_existing_file( - filename, - attach_type=attach_type, - insert_attachment=self.db.insert_attachment, - attachments_dir=self.attachments_dir, - ) - - async def _save_bot_message( - self, - tui_conv_id: str, - text: str, - media_parts: list, - reasoning: str, - agent_stats: dict, - refs: dict, - ): - """Save bot message to history, return saved record.""" - bot_message_parts = [] - bot_message_parts.extend(media_parts) - if text: - bot_message_parts.append({"type": "plain", "text": text}) - new_his: dict[str, Any] = {"type": "bot", "message": bot_message_parts} - if reasoning: - new_his["reasoning"] = reasoning - if agent_stats: - new_his["agent_stats"] = agent_stats - if refs: - new_his["refs"] = refs - mgr = self.platform_history_mgr - assert mgr is not None - record = await mgr.insert( - platform_id="tui", - user_id=tui_conv_id, - content=new_his, - sender_id="bot", - sender_name="bot", - ) - return record - - async def chat(self, post_data: dict | None = None): - username = g.get("username", "guest") - if post_data is None: - post_data = await request.json - if post_data is None: - return Response().error("Missing JSON body").to_json() - if "message" not in post_data and "files" not in post_data: - return Response().error("Missing key: message or files").to_json() - if "session_id" not in post_data and "conversation_id" not in post_data: - return ( - Response().error("Missing key: session_id or conversation_id").to_json() - ) - message = post_data["message"] - session_id = post_data.get("session_id", post_data.get("conversation_id")) - selected_provider = post_data.get("selected_provider") - selected_model = post_data.get("selected_model") - enable_streaming = post_data.get("enable_streaming", True) - if not session_id: - return Response().error("session_id is empty").to_json() - tui_conv_id = session_id - message_parts = await self._build_user_message_parts(message) - if not webchat_message_parts_have_content(message_parts): - return ( - Response() - .error("Message content is empty (reply only is not allowed)") - .to_json() - ) - message_id = str(uuid.uuid4()) - back_queue = tui_queue_mgr.get_or_create_back_queue(message_id, tui_conv_id) - - async def stream(): - client_disconnected = False - accumulated_parts = [] - accumulated_text = "" - accumulated_reasoning = "" - tool_calls = {} - agent_stats = {} - refs = {} - try: - session_info = { - "type": "session_id", - "data": None, - "session_id": tui_conv_id, - } - yield f"data: {json.dumps(session_info, ensure_ascii=False)}\n\n" - async with track_conversation(self.running_convs, tui_conv_id): - while True: - result, should_break = await _poll_tui_stream_result( - back_queue, - username, - ) - if should_break: - client_disconnected = True - break - if not result: - continue - if ( - "message_id" in result - and result["message_id"] != message_id - ): - logger.warning("TUI stream message_id mismatch") - continue - result_text = result["data"] - msg_type = result.get("type") - streaming = result.get("streaming", False) - chain_type = result.get("chain_type") - if chain_type == "agent_stats": - stats_info = { - "type": "agent_stats", - "data": json.loads(result_text), - } - yield f"data: {json.dumps(stats_info, ensure_ascii=False)}\n\n" - agent_stats = stats_info["data"] - continue - try: - if not client_disconnected: - yield f"data: {json.dumps(result, ensure_ascii=False)}\n\n" - except Exception as e: - if not client_disconnected: - logger.debug(f"[TUI] User {username} disconnected. {e}") - client_disconnected = True - try: - if not client_disconnected: - await asyncio.sleep(0.05) - except asyncio.CancelledError: - logger.debug(f"[TUI] User {username} disconnected.") - client_disconnected = True - if msg_type == "plain": - chain_type = result.get("chain_type") - if chain_type == "tool_call": - tool_call = json.loads(result_text) - tool_calls[tool_call.get("id")] = tool_call - if accumulated_text: - accumulated_parts.append( - {"type": "plain", "text": accumulated_text}, - ) - accumulated_text = "" - elif chain_type == "tool_call_result": - tcr = json.loads(result_text) - tc_id = tcr.get("id") - if tc_id in tool_calls: - tool_calls[tc_id]["result"] = tcr.get("result") - tool_calls[tc_id]["finished_ts"] = tcr.get("ts") - accumulated_parts.append( - { - "type": "tool_call", - "tool_calls": [tool_calls[tc_id]], - }, - ) - tool_calls.pop(tc_id, None) - elif chain_type == "reasoning": - accumulated_reasoning += result_text - elif streaming: - accumulated_text += result_text - else: - accumulated_text = result_text - elif msg_type == "image": - filename = result_text.replace("[IMAGE]", "") - part = await self._create_attachment_from_file( - filename, - "image", - ) - if part: - accumulated_parts.append(part) - elif msg_type == "record": - filename = result_text.replace("[RECORD]", "") - part = await self._create_attachment_from_file( - filename, - "record", - ) - if part: - accumulated_parts.append(part) - elif msg_type == "file": - filename = result_text.replace("[FILE]", "") - part = await self._create_attachment_from_file( - filename, - "file", - ) - if part: - accumulated_parts.append(part) - if msg_type == "end": - break - elif (streaming and msg_type == "complete") or not streaming: - if ( - chain_type == "tool_call" - or chain_type == "tool_call_result" - ): - continue - saved_record = await self._save_bot_message( - tui_conv_id, - accumulated_text, - accumulated_parts, - accumulated_reasoning, - agent_stats, - refs, - ) - if saved_record and (not client_disconnected): - saved_info = { - "type": "message_saved", - "data": { - "id": saved_record.id, - "created_at": to_utc_isoformat( - saved_record.created_at, - ), - }, - } - try: - yield f"data: {json.dumps(saved_info, ensure_ascii=False)}\n\n" - except Exception: - pass - accumulated_parts = [] - accumulated_text = "" - accumulated_reasoning = "" - agent_stats = {} - refs = {} - except BaseException as e: - logger.exception(f"TUI stream unexpected error: {e}", exc_info=True) - finally: - tui_queue_mgr.remove_back_queue(message_id) - - chat_queue = tui_queue_mgr.get_or_create_queue(tui_conv_id) - await chat_queue.put( - ( - username, - tui_conv_id, - { - "message": message_parts, - "selected_provider": selected_provider, - "selected_model": selected_model, - "enable_streaming": enable_streaming, - "message_id": message_id, - }, - ), - ) - message_parts_for_storage = strip_message_parts_path_fields(message_parts) - mgr = self.platform_history_mgr - assert mgr is not None - await mgr.insert( - platform_id="tui", - user_id=tui_conv_id, - content={"type": "user", "message": message_parts_for_storage}, - sender_id=username, - sender_name=username, - ) - response = await make_response( - stream(), - { - "Content-Type": "text/event-stream", - "Cache-Control": "no-cache", - "Transfer-Encoding": "chunked", - "Connection": "keep-alive", - }, - ) - response.timeout = None - return response - - async def stop_session(self): - """Stop active agent runs for a session.""" - post_data = await request.json - if post_data is None: - return Response().error("Missing JSON body").to_json() - session_id = post_data.get("session_id") - if not session_id: - return Response().error("Missing key: session_id").to_json() - username = g.get("username", "guest") - session = await self.db.get_platform_session_by_id(session_id) - if not session: - return Response().error(f"Session {session_id} not found").to_json() - if session.creator != username: - return Response().error("Permission denied").to_json() - message_type = ( - MessageType.GROUP_MESSAGE.value - if session.is_group - else MessageType.FRIEND_MESSAGE.value - ) - umo = f"{session.platform_id}:{message_type}:{session.platform_id}!{username}!{session_id}" - stopped_count = active_event_registry.request_agent_stop_all(umo) - return Response().ok(data={"stopped_count": stopped_count}).to_json() - - async def _delete_session_internal(self, session, username: str) -> None: - """Delete a single session and all its related data.""" - session_id = session.session_id - message_type = "GroupMessage" if session.is_group else "FriendMessage" - unified_msg_origin = f"{session.platform_id}:{message_type}:{session.platform_id}!{username}!{session_id}" - conv_mgr = self.conv_mgr - assert conv_mgr is not None - await conv_mgr.delete_conversations_by_user_id(unified_msg_origin) - mgr = self.platform_history_mgr - assert mgr is not None - history_list = await mgr.get( - platform_id=session.platform_id, - user_id=session_id, - page=1, - page_size=100000, - ) - attachment_ids = self._extract_attachment_ids(history_list) - if attachment_ids: - await self._delete_attachments(attachment_ids) - await mgr.delete( - platform_id=session.platform_id, - user_id=session_id, - offset_sec=99999999, - ) - try: - router = self.umop_config_router - if router is None: - logger.warning( - "UMOP config router not available during session cleanup", - ) - else: - await router.delete_route(unified_msg_origin) - except ValueError: - logger.warning( - "Failed to delete UMO route %s during session cleanup.", - unified_msg_origin, - ) - if session.platform_id == "tui": - tui_queue_mgr.remove_queues(session_id) - await self.db.delete_platform_session(session_id) - - async def delete_tui_session(self): - """Delete a Platform session and all its related data.""" - session_id = request.args.get("session_id") - if not session_id: - return Response().error("Missing key: session_id").to_json() - username = g.get("username", "guest") - session = await self.db.get_platform_session_by_id(session_id) - if not session: - return Response().error(f"Session {session_id} not found").to_json() - if session.creator != username: - return Response().error("Permission denied").to_json() - await self._delete_session_internal(session, username) - return Response().ok().to_json() - - async def batch_delete_sessions(self): - """Batch delete multiple Platform sessions.""" - post_data = await request.json - if post_data is None: - return Response().error("Missing JSON body").to_json() - if not isinstance(post_data, dict): - return Response().error("Invalid JSON body: expected object").to_json() - session_ids = post_data.get("session_ids") - if not session_ids or not isinstance(session_ids, list): - return Response().error("Missing or invalid key: session_ids").to_json() - username = g.get("username", "guest") - sessions = await self.db.get_platform_sessions_by_ids(session_ids) - sessions_by_id = {session.session_id: session for session in sessions} - deleted_count = 0 - failed_items = [] - for sid in session_ids: - session = sessions_by_id.get(sid) - if not session: - failed_items.append({"session_id": sid, "reason": "not found"}) - continue - if session.creator != username: - failed_items.append({"session_id": sid, "reason": "permission denied"}) - continue - try: - await self._delete_session_internal(session, username) - deleted_count += 1 - sessions_by_id.pop(sid, None) - except Exception: - logger.warning("Failed to delete session %s", sid) - failed_items.append({"session_id": sid, "reason": "internal_error"}) - return ( - Response() - .ok( - data={ - "deleted_count": deleted_count, - "failed_count": len(failed_items), - "failed_items": failed_items, - }, - ) - .to_json() - ) - - def _extract_attachment_ids(self, history_list) -> list[str]: - """Extract all attachment_ids from message history.""" - attachment_ids = [] - for history in history_list: - content = history.content - if not content or "message" not in content: - continue - message_parts = content.get("message", []) - for part in message_parts: - if isinstance(part, dict) and "attachment_id" in part: - attachment_ids.append(part["attachment_id"]) - return attachment_ids - - async def _delete_attachments(self, attachment_ids: list[str]) -> None: - """Delete attachments including DB records and disk files.""" - try: - attachments = await self.db.get_attachments(attachment_ids) - for attachment in attachments: - if not await anyio.Path(attachment.path).exists(): - continue - try: - await anyio.Path(attachment.path).unlink() - except OSError as e: - logger.warning( - f"Failed to delete attachment file {attachment.path}: {e}", - ) - except Exception as e: - logger.warning(f"Failed to get attachments: {e}") - try: - await self.db.delete_attachments(attachment_ids) - except Exception as e: - logger.warning(f"Failed to delete attachments: {e}") - - async def new_session(self): - """Create a new Platform session for TUI.""" - username = g.get("username", "guest") - session = await self.db.create_platform_session( - creator=username, - platform_id="tui", - is_group=0, - ) - return ( - Response() - .ok( - data={ - "session_id": session.session_id, - "platform_id": session.platform_id, - }, - ) - .to_json() - ) - - async def get_sessions(self): - """Get all Platform sessions for the current user filtered by TUI platform.""" - username = g.get("username", "guest") - platform_id = request.args.get("platform_id", "tui") - sessions, _ = await self.db.get_platform_sessions_by_creator_paginated( - creator=username, - platform_id=platform_id, - page=1, - page_size=100, - exclude_project_sessions=True, - ) - sessions_data = [] - for item in sessions: - session = item["session"] - sessions_data.append( - { - "session_id": session.session_id, - "platform_id": session.platform_id, - "creator": session.creator, - "display_name": session.display_name, - "is_group": session.is_group, - "created_at": to_utc_isoformat(session.created_at), - "updated_at": to_utc_isoformat(session.updated_at), - }, - ) - return Response().ok(data=sessions_data).to_json() - - async def get_session(self): - """Get session information and message history by session_id.""" - session_id = request.args.get("session_id") - if not session_id: - return Response().error("Missing key: session_id").to_json() - session = await self.db.get_platform_session_by_id(session_id) - platform_id = session.platform_id if session else "tui" - username = g.get("username", "guest") - project_info = await self.db.get_project_by_session( - session_id=session_id, - creator=username, - ) - mgr = self.platform_history_mgr - assert mgr is not None - history_ls = await mgr.get( - platform_id=platform_id, - user_id=session_id, - page=1, - page_size=1000, - ) - history_res = [history.model_dump() for history in history_ls] - response_data: dict[str, Any] = { - "history": history_res, - "is_running": self.running_convs.get(session_id, False), - } - if project_info: - response_data["project"] = { - "project_id": project_info.project_id, - "title": project_info.title, - "emoji": project_info.emoji, - } - return Response().ok(data=response_data).to_json() - - async def update_session_display_name(self): - """Update a Platform session's display name.""" - post_data = await request.json - session_id = post_data.get("session_id") - display_name = post_data.get("display_name") - if not session_id: - return Response().error("Missing key: session_id").to_json() - if display_name is None: - return Response().error("Missing key: display_name").to_json() - username = g.get("username", "guest") - session = await self.db.get_platform_session_by_id(session_id) - if not session: - return Response().error(f"Session {session_id} not found").to_json() - if session.creator != username: - return Response().error("Permission denied").to_json() - await self.db.update_platform_session( - session_id=session_id, - display_name=display_name, - ) - return Response().ok().to_json() diff --git a/astrbot/dashboard/server.py b/astrbot/dashboard/server.py index 082e04174..64ba59cdf 100644 --- a/astrbot/dashboard/server.py +++ b/astrbot/dashboard/server.py @@ -58,7 +58,6 @@ from .routes import ( SubAgentRoute, T2iRoute, ToolsRoute, - TUIChatRoute, UpdateRoute, ) from .routes.api_key import ALL_OPEN_API_SCOPES @@ -336,7 +335,6 @@ class AstrBotDashboard: self.platform_route = PlatformRoute(self.context, self.core_lifecycle) self.backup_route = BackupRoute(self.context, db, self.core_lifecycle) self.live_chat_route = LiveChatRoute(self.context, db, self.core_lifecycle) - self.tui_chat_route = TUIChatRoute(self.context, db, self.core_lifecycle) self.app.add_url_rule( "/api/plug/",