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
This commit is contained in:
LIghtJUNction
2026-04-11 18:27:15 +08:00
parent 8946a90afd
commit 7543dd2e9d
13 changed files with 15 additions and 1303 deletions
+1 -1
View File
@@ -317,7 +317,7 @@ WebUI 的配置文件在 `CONFIG_METADATA_3` 中。
未来将会逐步淘汰此配置元数据。
"""
CONFIG_METADATA_2 = {
CONFIG_METADATA_2: Any = {
"platform_group": {
"metadata": {
"platform": {
+5
View File
@@ -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"
+2 -1
View File
@@ -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:
-6
View File
@@ -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:
"""实例化一个平台"""
# 动态导入
+3 -1
View File
@@ -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 回调入口。
@@ -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"]
@@ -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()
@@ -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)
@@ -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()
-2
View File
@@ -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",
]
+4 -4
View File
@@ -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
-691
View File
@@ -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()
-2
View File
@@ -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/<path:subpath>",