mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
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:
@@ -317,7 +317,7 @@ WebUI 的配置文件在 `CONFIG_METADATA_3` 中。
|
||||
|
||||
未来将会逐步淘汰此配置元数据。
|
||||
"""
|
||||
CONFIG_METADATA_2 = {
|
||||
CONFIG_METADATA_2: Any = {
|
||||
"platform_group": {
|
||||
"metadata": {
|
||||
"platform": {
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
"""实例化一个平台"""
|
||||
# 动态导入
|
||||
|
||||
@@ -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()
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
@@ -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>",
|
||||
|
||||
Reference in New Issue
Block a user