chore: stage all dev changes

This commit is contained in:
LIghtJUNction
2026-03-30 00:48:07 +08:00
parent 4a0ab861d8
commit 62d7df1da7
104 changed files with 1145 additions and 1117 deletions
+9
View File
@@ -0,0 +1,9 @@
{
"exa": {
"command": "npx",
"args": ["exa-mcp-server"],
"env": {
"EXA_API_KEY": "0093b9e9-7579-444a-8de2-b6e696b1b413"
}
}
}
+2 -2
View File
@@ -1,6 +1,6 @@
from typing import TYPE_CHECKING, Protocol, runtime_checkable
from ..message import Message
from astrbot.core.agent.message import Message
if TYPE_CHECKING:
from astrbot import logger
@@ -15,7 +15,7 @@ else:
if TYPE_CHECKING:
from astrbot.core.provider.provider import Provider
from ..context.truncator import ContextTruncator
from astrbot.core.agent.context.truncator import ContextTruncator
@runtime_checkable
+1 -1
View File
@@ -1,6 +1,6 @@
from astrbot import logger
from astrbot.core.agent.message import Message
from ..message import Message
from .compressor import LLMSummaryCompressor, TruncateByTurnsCompressor
from .config import ContextConfig
from .token_counter import EstimateTokenCounter
+7 -1
View File
@@ -1,7 +1,13 @@
import json
from typing import Protocol, runtime_checkable
from ..message import AudioURLPart, ImageURLPart, Message, TextPart, ThinkPart
from astrbot.core.agent.message import (
AudioURLPart,
ImageURLPart,
Message,
TextPart,
ThinkPart,
)
@runtime_checkable
+1 -1
View File
@@ -1,4 +1,4 @@
from ..message import Message
from astrbot.core.agent.message import Message
class ContextTruncator:
@@ -12,6 +12,11 @@ from dashscope.app.application_response import ApplicationResponse
import astrbot.core.message.components as Comp
from astrbot.core import logger, sp
from astrbot.core.agent.hooks import BaseAgentRunHooks
from astrbot.core.agent.response import AgentResponseData
from astrbot.core.agent.run_context import ContextWrapper, TContext
from astrbot.core.agent.runners.base import AgentResponse, AgentState, BaseAgentRunner
from astrbot.core.agent.tool_executor import BaseFunctionToolExecutor
from astrbot.core.message.message_event_result import MessageChain
from astrbot.core.provider.entities import (
LLMResponse,
@@ -19,12 +24,6 @@ from astrbot.core.provider.entities import (
)
from astrbot.core.provider.provider import Provider
from ...hooks import BaseAgentRunHooks
from ...response import AgentResponseData
from ...run_context import ContextWrapper, TContext
from ...tool_executor import BaseFunctionToolExecutor
from ..base import AgentResponse, AgentState, BaseAgentRunner
if sys.version_info >= (3, 12):
from typing import override
else:
@@ -12,6 +12,11 @@ from uuid import uuid4
import astrbot.core.message.components as Comp
from astrbot import logger
from astrbot.core import sp
from astrbot.core.agent.hooks import BaseAgentRunHooks
from astrbot.core.agent.response import AgentResponseData
from astrbot.core.agent.run_context import ContextWrapper, TContext
from astrbot.core.agent.runners.base import AgentResponse, AgentState, BaseAgentRunner
from astrbot.core.agent.tool_executor import BaseFunctionToolExecutor
from astrbot.core.message.message_event_result import MessageChain
from astrbot.core.provider.entities import (
LLMResponse,
@@ -20,11 +25,6 @@ from astrbot.core.provider.entities import (
from astrbot.core.provider.provider import Provider
from astrbot.core.utils.config_number import coerce_int_config
from ...hooks import BaseAgentRunHooks
from ...response import AgentResponseData
from ...run_context import ContextWrapper, TContext
from ...tool_executor import BaseFunctionToolExecutor
from ..base import AgentResponse, AgentState, BaseAgentRunner
from .constants import DEERFLOW_SESSION_PREFIX, DEERFLOW_THREAD_ID_KEY
from .deerflow_api_client import DeerFlowAPIClient
from .deerflow_content_mapper import (
@@ -1,5 +1,6 @@
import codecs
import json
import types
from collections.abc import AsyncGenerator
from typing import Any
@@ -136,7 +137,7 @@ class DeerFlowAPIClient:
self,
exc_type: type[BaseException] | None,
exc: BaseException | None,
tb: object | None,
tb: types.TracebackType | None,
) -> None:
await self.close()
+3 -6
View File
@@ -708,13 +708,10 @@ async def _process_quote_message(
except BaseException as exc:
logger.error("处理引用图片失败: %s", exc)
finally:
if (
compress_path
and compress_path != path
and os.path.exists(compress_path)
):
if compress_path and compress_path != path:
try:
os.remove(compress_path)
if await asyncio.to_thread(os.path.exists, compress_path):
await asyncio.to_thread(os.remove, compress_path)
except Exception as exc:
logger.warning("Fail to remove temporary compressed image: %s", exc)
+1 -1
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
import abc
from typing import TYPE_CHECKING
from ..olayer import (
from astrbot.core.computer.olayer import (
BrowserComponent,
FileSystemComponent,
PythonComponent,
+28 -17
View File
@@ -10,11 +10,15 @@ import sys
from dataclasses import dataclass, field
from typing import Any
from astrbot.core.computer.olayer import (
FileSystemComponent,
PythonComponent,
ShellComponent,
)
from astrbot.core.utils.astrbot_path import (
get_astrbot_temp_path,
)
from ..olayer import FileSystemComponent, PythonComponent, ShellComponent
from .base import ComputerBooter
@@ -36,6 +40,16 @@ def _decode_shell_output(output: bytes | None) -> str:
return output.decode("utf-8", errors="replace")
def _write_file_sync(path: str, content: str, mode: str, encoding: str) -> None:
with open(path, mode, encoding=encoding) as f:
f.write(content)
def _read_file_sync(path: str, encoding: str) -> str:
with open(path, encoding=encoding) as f:
return f.read()
@dataclass
class BwrapConfig:
workspace_dir: str
@@ -210,17 +224,15 @@ class HostBackedFileSystemComponent(FileSystemComponent):
self, path: str, content: str = "", mode: int = 0o644
) -> dict[str, Any]:
p = self._safe_path(path)
os.makedirs(os.path.dirname(p), exist_ok=True)
with open(p, "w", encoding="utf-8") as f:
f.write(content)
os.chmod(p, mode)
await asyncio.to_thread(os.makedirs, os.path.dirname(p), exist_ok=True)
await asyncio.to_thread(_write_file_sync, p, content, "w", "utf-8")
await asyncio.to_thread(os.chmod, p, mode)
return {"success": True, "path": p}
async def read_file(self, path: str, encoding: str = "utf-8") -> dict[str, Any]:
p = self._safe_path(path)
try:
with open(p, encoding=encoding) as f:
content = f.read()
content = await asyncio.to_thread(_read_file_sync, p, encoding)
return {"success": True, "content": content}
except Exception as e:
return {"success": False, "error": str(e)}
@@ -229,10 +241,9 @@ class HostBackedFileSystemComponent(FileSystemComponent):
self, path: str, content: str, mode: str = "w", encoding: str = "utf-8"
) -> dict[str, Any]:
p = self._safe_path(path)
os.makedirs(os.path.dirname(p), exist_ok=True)
await asyncio.to_thread(os.makedirs, os.path.dirname(p), exist_ok=True)
try:
with open(p, mode, encoding=encoding) as f:
f.write(content)
await asyncio.to_thread(_write_file_sync, p, content, mode, encoding)
return {"success": True}
except Exception as e:
return {"success": False, "error": str(e)}
@@ -240,10 +251,10 @@ class HostBackedFileSystemComponent(FileSystemComponent):
async def delete_file(self, path: str) -> dict[str, Any]:
p = self._safe_path(path)
try:
if os.path.isdir(p):
shutil.rmtree(p)
if await asyncio.to_thread(os.path.isdir, p):
await asyncio.to_thread(shutil.rmtree, p)
else:
os.remove(p)
await asyncio.to_thread(os.remove, p)
return {"success": True}
except Exception as e:
return {"success": False, "error": str(e)}
@@ -292,10 +303,10 @@ class BwrapBooter(ComputerBooter):
workspace_dir = os.path.join(
get_astrbot_temp_path(), f"sandbox_workspace_{session_id}"
)
os.makedirs(workspace_dir, exist_ok=True)
await asyncio.to_thread(os.makedirs, workspace_dir, exist_ok=True)
self.config = BwrapConfig(
workspace_dir=os.path.abspath(workspace_dir),
workspace_dir=await asyncio.to_thread(os.path.abspath, workspace_dir),
rw_binds=self._rw_binds,
ro_binds=self._ro_binds,
)
@@ -320,8 +331,8 @@ class BwrapBooter(ComputerBooter):
)
async def shutdown(self) -> None:
if self.config and os.path.exists(self.config.workspace_dir):
shutil.rmtree(self.config.workspace_dir, ignore_errors=True)
if self.config and await asyncio.to_thread(os.path.exists, self.config.workspace_dir):
await asyncio.to_thread(shutil.rmtree, self.config.workspace_dir, ignore_errors=True)
async def upload_file(self, path: str, file_name: str) -> dict:
if not self._fs or not self.config:
+5 -1
View File
@@ -10,13 +10,17 @@ from dataclasses import dataclass
from typing import Any
from astrbot.api import logger
from astrbot.core.computer.olayer import (
FileSystemComponent,
PythonComponent,
ShellComponent,
)
from astrbot.core.utils.astrbot_path import (
get_astrbot_data_path,
get_astrbot_root,
get_astrbot_temp_path,
)
from ..olayer import FileSystemComponent, PythonComponent, ShellComponent
from .base import ComputerBooter
_BLOCKED_COMMAND_PATTERNS = [
+6 -1
View File
@@ -10,7 +10,12 @@ from astrbot.api import logger
if TYPE_CHECKING:
from astrbot.core.agent.tool import FunctionTool
from ..olayer import FileSystemComponent, PythonComponent, ShellComponent
from astrbot.core.computer.olayer import (
FileSystemComponent,
PythonComponent,
ShellComponent,
)
from .base import ComputerBooter
+12 -14
View File
@@ -39,11 +39,12 @@ class NeoPythonComponent(PythonComponent):
self,
code: str,
kernel_id: str | None = None,
timeout: int = 30,
timeout_sec: int = 30,
silent: bool = False,
) -> dict[str, Any]:
_ = kernel_id # Bay runtime does not expose kernel_id in current SDK.
result = await self._sandbox.python.exec(code, timeout=timeout)
with anyio.fail_after(timeout_sec):
result = await self._sandbox.python.exec(code)
payload = _maybe_model_dump(result)
output_text = payload.get("output", "") or ""
@@ -81,7 +82,7 @@ class NeoShellComponent(ShellComponent):
command: str,
cwd: str | None = None,
env: dict[str, str] | None = None,
timeout: int | None = 30,
timeout_sec: int | None = 30,
shell: bool = True,
background: bool = False,
) -> dict[str, Any]:
@@ -103,11 +104,8 @@ class NeoShellComponent(ShellComponent):
if background:
run_command = f"nohup sh -lc {shlex.quote(run_command)} >/tmp/astrbot_bg.log 2>&1 & echo $!"
result = await self._sandbox.shell.exec(
run_command,
timeout=timeout or 30,
cwd=cwd,
)
with anyio.fail_after(timeout_sec or 30):
result = await self._sandbox.shell.exec(run_command, cwd=cwd)
payload = _maybe_model_dump(result)
stdout = payload.get("output", "") or ""
@@ -198,7 +196,7 @@ class NeoBrowserComponent(BrowserComponent):
async def exec(
self,
cmd: str,
timeout: int = 30,
timeout_sec: int = 30,
description: str | None = None,
tags: str | None = None,
learn: bool = False,
@@ -206,7 +204,7 @@ class NeoBrowserComponent(BrowserComponent):
) -> dict[str, Any]:
result = await self._sandbox.browser.exec(
cmd,
timeout=timeout,
timeout_sec=timeout_sec,
description=description,
tags=tags,
learn=learn,
@@ -217,7 +215,7 @@ class NeoBrowserComponent(BrowserComponent):
async def exec_batch(
self,
commands: list[str],
timeout: int = 60,
timeout_sec: int = 60,
stop_on_error: bool = True,
description: str | None = None,
tags: str | None = None,
@@ -226,7 +224,7 @@ class NeoBrowserComponent(BrowserComponent):
) -> dict[str, Any]:
result = await self._sandbox.browser.exec_batch(
commands,
timeout=timeout,
timeout_sec=timeout_sec,
stop_on_error=stop_on_error,
description=description,
tags=tags,
@@ -238,7 +236,7 @@ class NeoBrowserComponent(BrowserComponent):
async def run_skill(
self,
skill_key: str,
timeout: int = 60,
timeout_sec: int = 60,
stop_on_error: bool = True,
include_trace: bool = False,
description: str | None = None,
@@ -246,7 +244,7 @@ class NeoBrowserComponent(BrowserComponent):
) -> dict[str, Any]:
result = await self._sandbox.browser.run_skill(
skill_key=skill_key,
timeout=timeout,
timeout_sec=timeout_sec,
stop_on_error=stop_on_error,
include_trace=include_trace,
description=description,
+1 -1
View File
@@ -6,8 +6,8 @@ from astrbot.api import FunctionTool
from astrbot.core.agent.run_context import ContextWrapper
from astrbot.core.agent.tool import ToolExecResult
from astrbot.core.astr_agent_context import AstrAgentContext
from astrbot.core.computer.computer_client import get_booter
from ..computer_client import get_booter
from .permissions import check_admin_permission
+1 -1
View File
@@ -9,10 +9,10 @@ from astrbot.api.event import MessageChain
from astrbot.core.agent.run_context import ContextWrapper
from astrbot.core.agent.tool import ToolExecResult
from astrbot.core.astr_agent_context import AstrAgentContext
from astrbot.core.computer.computer_client import get_booter
from astrbot.core.message.components import File
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..computer_client import get_booter
from .permissions import check_admin_permission
# @dataclass
+1 -1
View File
@@ -7,9 +7,9 @@ from astrbot.api import FunctionTool
from astrbot.core.agent.run_context import ContextWrapper
from astrbot.core.agent.tool import ToolExecResult
from astrbot.core.astr_agent_context import AstrAgentContext
from astrbot.core.computer.computer_client import get_booter
from astrbot.core.skills.neo_skill_sync import NeoSkillSyncManager
from ..computer_client import get_booter
from .permissions import check_admin_permission
+22 -15
View File
@@ -17,6 +17,7 @@ Behavior:
from __future__ import annotations
import asyncio
import json
import os
import shlex
@@ -140,12 +141,13 @@ class ExecuteShellTool(FunctionTool):
else:
exec_env = session_env
# Determine timeout from config (fall back to 30) — use the cast wrapper's context
# Determine timeout from config (fall back to 30)
config = astr_ctx.context.context.get_config(umo=session_id)
provider_settings: dict = {}
if isinstance(config, dict):
provider_settings = config.get("provider_settings") or {}
try:
timeout = int(
config.get("provider_settings", {}).get("tool_call_timeout", 30)
)
timeout = int(provider_settings.get("tool_call_timeout", 30))
except (ValueError, TypeError):
timeout = 30
@@ -180,18 +182,21 @@ class ExecuteShellTool(FunctionTool):
parts = shlex.split(cd_part)
# cd with no args -> home
if len(parts) == 1:
target = os.path.expanduser("~")
target = await asyncio.to_thread(os.path.expanduser, "~")
else:
target_raw = parts[1]
# expand ~ and variables
target_raw = os.path.expanduser(target_raw)
target = (
target_raw
if os.path.isabs(target_raw)
else os.path.normpath(os.path.join(session_cwd, target_raw))
)
target_raw = await asyncio.to_thread(os.path.expanduser, target_raw)
if await asyncio.to_thread(os.path.isabs, target_raw):
target = target_raw
else:
target = await asyncio.to_thread(
os.path.normpath, os.path.join(session_cwd, target_raw)
)
if not os.path.exists(target) or not os.path.isdir(target):
target_exists = await asyncio.to_thread(os.path.exists, target)
target_isdir = await asyncio.to_thread(os.path.isdir, target)
if not target_exists or not target_isdir:
result = {
"success": False,
"exit_code": -1,
@@ -225,7 +230,8 @@ class ExecuteShellTool(FunctionTool):
# Background execution: spawn process and return pid immediately.
if background:
# Start background process; do not wait. Use shell to support pipes/redirects.
popen = subprocess.Popen(
popen = await asyncio.to_thread(
subprocess.Popen,
["/bin/sh", "-c", command_to_run],
cwd=session_cwd,
env=exec_env,
@@ -241,7 +247,8 @@ class ExecuteShellTool(FunctionTool):
return json.dumps(result)
# Foreground execution: run to completion, capture output.
completed = subprocess.run(
completed = await asyncio.to_thread(
subprocess.run,
["/bin/sh", "-c", command_to_run],
cwd=session_cwd,
env=exec_env,
@@ -268,7 +275,7 @@ class ExecuteShellTool(FunctionTool):
{
"success": False,
"exit_code": -1,
"stdout": getattr(e, "output", "") or "",
"stdout": e.stdout or "",
"stderr": f"Command timed out after {timeout} seconds",
"cwd": session_cwd,
}
+1 -1
View File
@@ -7,10 +7,10 @@ from sqlalchemy.ext.asyncio import AsyncSession
from astrbot.api import logger, sp
from astrbot.core.config import AstrBotConfig
from astrbot.core.config.default import DB_PATH
from astrbot.core.db import BaseDatabase
from astrbot.core.db.po import ConversationV2, PlatformMessageHistory
from astrbot.core.platform.astr_message_event import MessageSesion
from .. import BaseDatabase
from .shared_preferences_v3 import sp as sp_v3
from .sqlite_v3 import SQLiteDatabase as SQLiteV3DatabaseV3
+1 -1
View File
@@ -4,9 +4,9 @@ import uuid
import numpy as np
from astrbot import logger
from astrbot.core.db.vec_db.base import BaseVecDB, Result
from astrbot.core.provider.provider import EmbeddingProvider, RerankProvider
from ..base import BaseVecDB, Result
from .document_storage import DocumentStorage
from .embedding_storage import EmbeddingStorage
@@ -10,12 +10,11 @@ from astrbot import logger
from astrbot.core.db.vec_db.base import Result
from astrbot.core.db.vec_db.faiss_impl import FaissVecDB
from astrbot.core.knowledge_base.kb_db_sqlite import KBSQLiteDatabase
from astrbot.core.knowledge_base.kb_helper import KBHelper
from astrbot.core.knowledge_base.retrieval.rank_fusion import RankFusion
from astrbot.core.knowledge_base.retrieval.sparse_retriever import SparseRetriever
from astrbot.core.provider.provider import RerankProvider
from ..kb_helper import KBHelper
@dataclass
class RetrievalResult:
@@ -2,10 +2,10 @@ from collections.abc import AsyncGenerator
from astrbot.core import logger
from astrbot.core.message.message_event_result import MessageEventResult
from astrbot.core.pipeline.context import PipelineContext
from astrbot.core.pipeline.stage import Stage, register_stage
from astrbot.core.platform.astr_message_event import AstrMessageEvent
from ..context import PipelineContext
from ..stage import Stage, register_stage
from .strategies.strategy import StrategySelector
@@ -5,11 +5,10 @@ from collections.abc import AsyncGenerator
from astrbot.core import logger
from astrbot.core.message.components import Image, Plain, Record
from astrbot.core.pipeline.context import PipelineContext
from astrbot.core.pipeline.stage import Stage, register_stage
from astrbot.core.platform.astr_message_event import AstrMessageEvent
from ..context import PipelineContext
from ..stage import Stage, register_stage
@register_stage
class PreProcessStage(Stage):
@@ -6,13 +6,12 @@ from typing import Any
from astrbot.core import logger
from astrbot.core.message.message_event_result import MessageEventResult
from astrbot.core.pipeline.context import PipelineContext, call_event_hook, call_handler
from astrbot.core.pipeline.process_stage.stage import Stage
from astrbot.core.platform.astr_message_event import AstrMessageEvent
from astrbot.core.star.star import star_map
from astrbot.core.star.star_handler import EventType, StarHandlerMetadata
from ...context import PipelineContext, call_event_hook, call_handler
from ..stage import Stage
class StarRequestSubStage(Stage):
async def initialize(self, ctx: PipelineContext) -> None:
+2 -2
View File
@@ -1,11 +1,11 @@
from collections.abc import AsyncGenerator
from astrbot.core.pipeline.context import PipelineContext
from astrbot.core.pipeline.stage import Stage, register_stage
from astrbot.core.platform.astr_message_event import AstrMessageEvent
from astrbot.core.provider.entities import ProviderRequest
from astrbot.core.star.star_handler import StarHandlerMetadata
from ..context import PipelineContext
from ..stage import Stage, register_stage
from .method.agent_request import AgentRequestSubStage
from .method.star_request import StarRequestSubStage
+2 -3
View File
@@ -8,13 +8,12 @@ import astrbot.core.message.components as Comp
from astrbot.core import logger
from astrbot.core.message.components import BaseMessageComponent, ComponentType
from astrbot.core.message.message_event_result import MessageChain, ResultContentType
from astrbot.core.pipeline.context import PipelineContext, call_event_hook
from astrbot.core.pipeline.stage import Stage, register_stage
from astrbot.core.platform.astr_message_event import AstrMessageEvent
from astrbot.core.star.star_handler import EventType
from astrbot.core.utils.path_util import path_Mapping
from ..context import PipelineContext, call_event_hook
from ..stage import Stage, register_stage
@register_stage
class RespondStage(Stage):
@@ -8,15 +8,14 @@ from astrbot.core import file_token_service, html_renderer, logger
from astrbot.core.message.components import At, Image, Json, Node, Plain, Record, Reply
from astrbot.core.message.message_event_result import ResultContentType
from astrbot.core.pipeline.content_safety_check.stage import ContentSafetyCheckStage
from astrbot.core.pipeline.context import PipelineContext
from astrbot.core.pipeline.stage import Stage, register_stage, registered_stages
from astrbot.core.platform.astr_message_event import AstrMessageEvent
from astrbot.core.platform.message_type import MessageType
from astrbot.core.star.session_llm_manager import SessionServiceManager
from astrbot.core.star.star import star_map
from astrbot.core.star.star_handler import EventType, star_handlers_registry
from ..context import PipelineContext
from ..stage import Stage, register_stage, registered_stages
@register_stage
class ResultDecorateStage(Stage):
@@ -1,12 +1,11 @@
from collections.abc import AsyncGenerator
from astrbot.core import logger
from astrbot.core.pipeline.context import PipelineContext
from astrbot.core.pipeline.stage import Stage, register_stage
from astrbot.core.platform.astr_message_event import AstrMessageEvent
from astrbot.core.star.session_llm_manager import SessionServiceManager
from ..context import PipelineContext
from ..stage import Stage, register_stage
@register_stage
class SessionStatusCheckStage(Stage):
+2 -3
View File
@@ -3,6 +3,8 @@ from collections.abc import AsyncGenerator, Callable
from astrbot import logger
from astrbot.core.message.components import At, AtAll, Reply
from astrbot.core.message.message_event_result import MessageChain, MessageEventResult
from astrbot.core.pipeline.context import PipelineContext
from astrbot.core.pipeline.stage import Stage, register_stage
from astrbot.core.platform.astr_message_event import AstrMessageEvent
from astrbot.core.platform.message_type import MessageType
from astrbot.core.star.filter.command_group import CommandGroupFilter
@@ -11,9 +13,6 @@ from astrbot.core.star.session_plugin_manager import SessionPluginManager
from astrbot.core.star.star import star_map
from astrbot.core.star.star_handler import EventType, star_handlers_registry
from ..context import PipelineContext
from ..stage import Stage, register_stage
UNIQUE_SESSION_ID_BUILDERS: dict[str, Callable[[AstrMessageEvent], str | None]] = {
"aiocqhttp": lambda e: f"{e.get_sender_id()}_{e.get_group_id()}",
"slack": lambda e: f"{e.get_sender_id()}_{e.get_group_id()}",
@@ -1,12 +1,11 @@
from collections.abc import AsyncGenerator
from astrbot.core import logger
from astrbot.core.pipeline.context import PipelineContext
from astrbot.core.pipeline.stage import Stage, register_stage
from astrbot.core.platform.astr_message_event import AstrMessageEvent
from astrbot.core.platform.message_type import MessageType
from ..context import PipelineContext
from ..stage import Stage, register_stage
@register_stage
class WhitelistCheckStage(Stage):
@@ -87,8 +87,9 @@ class AiocqhttpMessageEvent(AstrMessageEvent):
new_seg.file = new_path
modified = True
except Exception as upload_err:
raise f"NapCat 文件流式上传失败: {upload_err}"
# 上传失败,保留原文件路径,但继续后续 segments 处理
raise RuntimeError(
f"NapCat 文件流式上传失败: {upload_err}"
) from upload_err
new_chain.chain.append(new_seg)
if not modified:
return False
@@ -112,45 +113,56 @@ class AiocqhttpMessageEvent(AstrMessageEvent):
file_path = path
path = Path(file_path)
if not path.exists():
if not await asyncio.to_thread(path.exists):
raise FileNotFoundError(f"文件不存在: {file_path}")
# 第一次遍历:计算文件总大小和 SHA256 哈希
hasher = hashlib.sha256()
total_size = 0
with open(path, "rb") as f:
while True:
chunk = f.read(CHUNK_SIZE)
if not chunk:
break
hasher.update(chunk)
total_size += len(chunk)
sha256_hash = hasher.hexdigest()
def _read_all_and_hash():
hasher = hashlib.sha256()
total_size = 0
with open(path, "rb") as f:
while True:
chunk = f.read(CHUNK_SIZE)
if not chunk:
break
hasher.update(chunk)
total_size += len(chunk)
return hasher.hexdigest(), total_size
sha256_hash, total_size = await asyncio.to_thread(_read_all_and_hash)
total_chunks = (total_size + CHUNK_SIZE - 1) // CHUNK_SIZE
# 第二次遍历:逐块上传
stream_id = str(uuid.uuid4())
with open(path, "rb") as f:
for i in range(total_chunks):
chunk = f.read(CHUNK_SIZE)
if not chunk:
break
chunk_b64 = base64.b64encode(chunk).decode("utf-8")
params = {
"stream_id": stream_id,
"chunk_data": chunk_b64,
"chunk_index": i,
"total_chunks": total_chunks,
"file_size": total_size,
"expected_sha256": sha256_hash,
"filename": path.name,
"file_retention": FILE_RETENTION_MS, # 单位为毫秒
}
resp = await bot.call_action("upload_file_stream", **params)
if not cls._is_upload_success_response(
resp, expected_statuses=("chunk_received", "file_complete")
):
raise OSError(f"上传分片 {i} 失败: {resp}")
async def _read_chunk(file_pos: int) -> bytes:
def _read_chunk_sync(file_pos: int) -> bytes:
with open(path, "rb") as f:
f.seek(file_pos)
return f.read(CHUNK_SIZE)
return await asyncio.to_thread(_read_chunk_sync, file_pos)
for i in range(total_chunks):
chunk = await _read_chunk(i * CHUNK_SIZE)
if not chunk:
break
chunk_b64 = base64.b64encode(chunk).decode("utf-8")
params = {
"stream_id": stream_id,
"chunk_data": chunk_b64,
"chunk_index": i,
"total_chunks": total_chunks,
"file_size": total_size,
"expected_sha256": sha256_hash,
"filename": path.name,
"file_retention": FILE_RETENTION_MS, # 单位为毫秒
}
resp = await bot.call_action("upload_file_stream", **params)
if not cls._is_upload_success_response(
resp, expected_statuses=("chunk_received", "file_complete")
):
raise OSError(f"上传分片 {i} 失败: {resp}")
# 发送完成信号
complete_params = {"stream_id": stream_id, "is_complete": True}
@@ -29,8 +29,8 @@ from astrbot.api.platform import (
PlatformMetadata,
)
from astrbot.core.platform.astr_message_event import MessageSesion
from astrbot.core.platform.register import register_platform_adapter
from ...register import register_platform_adapter
from .aiocqhttp_message_event import AiocqhttpMessageEvent
@@ -22,6 +22,7 @@ from astrbot.api.platform import (
)
from astrbot.core import sp
from astrbot.core.platform.astr_message_event import MessageSesion
from astrbot.core.platform.register import register_platform_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from astrbot.core.utils.io import download_file
from astrbot.core.utils.media_utils import (
@@ -31,7 +32,6 @@ from astrbot.core.utils.media_utils import (
get_media_duration,
)
from ...register import register_platform_adapter
from .dingtalk_event import DingtalkMessageEvent
@@ -84,7 +84,7 @@ class KookBaseDataClass(BaseModel):
def to_dict(
self,
mode: Literal["json", "python"] | str = "python",
mode: Literal["json", "python"] = "python",
by_alias=True,
exclude_none=True,
exclude_unset=False,
@@ -26,10 +26,10 @@ from astrbot.api.platform import (
PlatformMetadata,
)
from astrbot.core.platform.astr_message_event import MessageSesion
from astrbot.core.platform.register import register_platform_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from astrbot.core.utils.webhook_utils import log_webhook_info
from ...register import register_platform_adapter
from .lark_event import LarkMessageEvent
from .server import LarkWebhookServer
@@ -146,7 +146,7 @@ class LarkMessageEvent(AstrMessageEvent):
Returns:
成功返回file_key,失败返回None
"""
if not path or not os.path.exists(path):
if not path or not await asyncio.to_thread(os.path.exists, path):
logger.error(f"[Lark] 文件不存在: {path}")
return None
@@ -155,36 +155,38 @@ class LarkMessageEvent(AstrMessageEvent):
return None
try:
with open(path, "rb") as file_obj:
body_builder = (
CreateFileRequestBody.builder()
.file_type(file_type)
.file_name(os.path.basename(path))
.file(file_obj)
)
if duration is not None:
body_builder.duration(duration)
# Read file content in a thread to avoid blocking the event loop
def _read_file() -> bytes:
with open(path, "rb") as f:
return f.read()
request = (
CreateFileRequest.builder()
.request_body(body_builder.build())
.build()
)
response = await lark_client.im.v1.file.acreate(request)
file_bytes = await asyncio.to_thread(_read_file)
if not response.success():
logger.error(
f"[Lark] 无法上传文件({response.code}): {response.msg}"
)
return None
body_builder = (
CreateFileRequestBody.builder()
.file_type(file_type)
.file_name(os.path.basename(path))
.file(BytesIO(file_bytes))
)
if duration is not None:
body_builder.duration(duration)
if response.data is None:
logger.error("[Lark] 上传文件成功但未返回数据(data is None)")
return None
request = (
CreateFileRequest.builder().request_body(body_builder.build()).build()
)
response = await lark_client.im.v1.file.acreate(request)
file_key = response.data.file_key
logger.debug(f"[Lark] 文件上传成功: {file_key}")
return file_key
if not response.success():
logger.error(f"[Lark] 无法上传文件({response.code}): {response.msg}")
return None
if response.data is None:
logger.error("[Lark] 上传文件成功但未返回数据(data is None)")
return None
file_key = response.data.file_key
logger.debug(f"[Lark] 文件上传成功: {file_key}")
return file_key
except Exception as e:
logger.error(f"[Lark] 无法打开或上传文件: {e}")
@@ -217,8 +219,12 @@ class LarkMessageEvent(AstrMessageEvent):
temp_dir,
f"lark_image_{uuid.uuid4().hex[:8]}.jpg",
)
with open(file_path, "wb") as f:
f.write(BytesIO(image_data).getvalue())
def _write_image():
with open(file_path, "wb") as f:
f.write(BytesIO(image_data).getvalue())
await asyncio.to_thread(_write_image)
else:
file_path = comp.file if comp.file else ""
@@ -227,7 +233,13 @@ class LarkMessageEvent(AstrMessageEvent):
logger.error("[Lark] 图片路径为空,无法上传")
continue
try:
image_file = open(file_path, "rb")
def _open_image():
return open(file_path, "rb")
image_file = await asyncio.to_thread(
lambda: open(file_path, "rb")
)
except Exception as e:
logger.error(f"[Lark] 无法打开图片文件: {e}")
continue
@@ -634,7 +646,9 @@ class LarkMessageEvent(AstrMessageEvent):
logger.error(f"[Lark] 无法获取音频文件路径: {e}")
return
if not original_audio_path or not os.path.exists(original_audio_path):
if not original_audio_path or not await asyncio.to_thread(
os.path.exists, original_audio_path
):
logger.error(f"[Lark] 音频文件不存在: {original_audio_path}")
return
@@ -664,9 +678,11 @@ class LarkMessageEvent(AstrMessageEvent):
)
# 清理转换后的临时音频文件
if converted_audio_path and os.path.exists(converted_audio_path):
if converted_audio_path and await asyncio.to_thread(
os.path.exists, converted_audio_path
):
try:
os.remove(converted_audio_path)
await asyncio.to_thread(os.remove, converted_audio_path)
logger.debug(f"[Lark] 已删除转换后的音频文件: {converted_audio_path}")
except Exception as e:
logger.warning(f"[Lark] 删除转换后的音频文件失败: {e}")
@@ -707,7 +723,9 @@ class LarkMessageEvent(AstrMessageEvent):
logger.error(f"[Lark] 无法获取视频文件路径: {e}")
return
if not original_video_path or not os.path.exists(original_video_path):
if not original_video_path or not await asyncio.to_thread(
os.path.exists, original_video_path
):
logger.error(f"[Lark] 视频文件不存在: {original_video_path}")
return
@@ -737,9 +755,11 @@ class LarkMessageEvent(AstrMessageEvent):
)
# 清理转换后的临时视频文件
if converted_video_path and os.path.exists(converted_video_path):
if converted_video_path and await asyncio.to_thread(
os.path.exists, converted_video_path
):
try:
os.remove(converted_video_path)
await asyncio.to_thread(os.remove, converted_video_path)
logger.debug(f"[Lark] 已删除转换后的视频文件: {converted_video_path}")
except Exception as e:
logger.warning(f"[Lark] 删除转换后的视频文件失败: {e}")
@@ -17,10 +17,10 @@ from astrbot.api.platform import (
PlatformMetadata,
)
from astrbot.core.platform.astr_message_event import MessageSesion
from astrbot.core.platform.register import register_platform_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from astrbot.core.utils.webhook_utils import log_webhook_info
from ...register import register_platform_adapter
from .line_api import LineAPIClient
from .line_event import LineMessageEvent
@@ -24,8 +24,8 @@ from astrbot.api.platform import (
)
from astrbot.core.message.components import BaseMessageComponent
from astrbot.core.platform.astr_message_event import MessageSesion
from astrbot.core.platform.register import register_platform_adapter
from ...register import register_platform_adapter
from .qqofficial_message_event import QQOfficialMessageEvent
# remove logger handler
@@ -10,10 +10,12 @@ from astrbot import logger
from astrbot.api.event import MessageChain
from astrbot.api.platform import AstrBotMessage, 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.qqofficial.qqofficial_platform_adapter import (
QQOfficialPlatformAdapter,
)
from astrbot.core.utils.webhook_utils import log_webhook_info
from ...register import register_platform_adapter
from ..qqofficial.qqofficial_platform_adapter import QQOfficialPlatformAdapter
from .qo_webhook_event import QQOfficialWebhookMessageEvent
from .qo_webhook_server import QQOfficialWebhook
@@ -1,8 +1,9 @@
from botpy import Client
from astrbot.api.platform import AstrBotMessage, PlatformMetadata
from ..qqofficial.qqofficial_message_event import QQOfficialMessageEvent
from astrbot.core.platform.sources.qqofficial.qqofficial_message_event import (
QQOfficialMessageEvent,
)
class QQOfficialWebhookMessageEvent(QQOfficialMessageEvent):
@@ -20,9 +20,9 @@ from astrbot.api.platform import (
PlatformMetadata,
)
from astrbot.core.platform.astr_message_event import MessageSesion
from astrbot.core.platform.register import register_platform_adapter
from astrbot.core.utils.webhook_utils import log_webhook_info
from ...register import register_platform_adapter
from .client import SlackSocketClient, SlackWebhookClient
from .slack_event import SlackMessageEvent
@@ -174,8 +174,11 @@ class TelegramPlatformAdapter(Platform):
error_callback=self._on_polling_error
)
logger.info("Telegram Platform Adapter is running.")
while self.application.updater.running and not self._terminating:
await asyncio.sleep(1)
# Wait for termination or polling to stop.
_termination_event = asyncio.Event()
if self._terminating:
_termination_event.set()
await _termination_event.wait()
if not self._terminating:
logger.warning(
@@ -361,33 +364,36 @@ class TelegramPlatformAdapter(Platform):
logger.warning("Received an update without a message.")
return None
# Assign to local variable so type checker can infer non-None
msg = update.message
def _apply_caption() -> None:
if update.message.caption:
message.message_str = update.message.caption
if msg.caption:
message.message_str = msg.caption
message.message.append(Comp.Plain(message.message_str))
if update.message.caption and update.message.caption_entities:
for entity in update.message.caption_entities:
if msg.caption and msg.caption_entities:
for entity in msg.caption_entities:
if entity.type == "mention":
name = update.message.caption[
name = msg.caption[
entity.offset + 1 : entity.offset + entity.length
]
message.message.append(Comp.At(qq=name, name=name))
message = AstrBotMessage()
message.session_id = str(update.message.chat.id)
message.session_id = str(msg.chat.id)
# 获得是群聊还是私聊
if update.message.chat.type == ChatType.PRIVATE:
if msg.chat.type == ChatType.PRIVATE:
message.type = MessageType.FRIEND_MESSAGE
else:
message.type = MessageType.GROUP_MESSAGE
message.group_id = str(update.message.chat.id)
if update.message.is_topic_message and update.message.message_thread_id:
message.group_id = str(msg.chat.id)
if msg.is_topic_message and msg.message_thread_id:
# Telegram Topic Group: include thread id to isolate per-topic sessions.
message.group_id += "#" + str(update.message.message_thread_id)
message.group_id += "#" + str(msg.message_thread_id)
message.session_id = message.group_id
message.message_id = str(update.message.message_id)
_from_user = update.message.from_user
message.message_id = str(msg.message_id)
_from_user = msg.from_user
if not _from_user:
logger.warning("[Telegram] Received a message without a from_user.")
return None
@@ -22,9 +22,9 @@ from astrbot.api.platform import (
PlatformMetadata,
)
from astrbot.core.platform.astr_message_event import MessageSesion
from astrbot.core.platform.register import register_platform_adapter
from astrbot.core.utils.webhook_utils import log_webhook_info
from ...register import register_platform_adapter
from .wecomai_api import (
WecomAIBotAPIClient,
WecomAIBotMessageParser,
@@ -555,7 +555,7 @@ class WeixinOCAdapter(Platform):
item_type: int,
file_name: str,
) -> dict[str, Any]:
raw_bytes = media_path.read_bytes()
raw_bytes = await asyncio.to_thread(media_path.read_bytes)
raw_size = len(raw_bytes)
raw_md5 = hashlib.md5(raw_bytes).hexdigest()
file_key = uuid.uuid4().hex
@@ -761,7 +761,9 @@ class WeixinOCAdapter(Platform):
if not path:
return None
media_path = Path(path)
if not media_path.exists() or not media_path.is_file():
path_exists = await asyncio.to_thread(media_path.exists)
path_is_file = await asyncio.to_thread(media_path.is_file)
if not path_exists or not path_is_file:
return None
return media_path
@@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
import base64
import hashlib
import json
@@ -113,7 +114,7 @@ class WeixinOCClient:
aes_key_hex: str,
media_path: Path,
) -> str:
raw_data = media_path.read_bytes()
raw_data = await asyncio.to_thread(media_path.read_bytes)
logger.debug(
"weixin_oc(%s): prepare CDN upload file=%s size=%s md5=%s filekey=%s",
self.adapter_id,
@@ -15,14 +15,13 @@ from astrbot.api.provider import Provider
from astrbot.core.agent.message import ContentPart, ImageURLPart, TextPart
from astrbot.core.provider.entities import LLMResponse, TokenUsage
from astrbot.core.provider.func_tool_manager import ToolSet
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.io import download_image_by_url
from astrbot.core.utils.network_utils import (
is_connection_error,
log_connection_failure,
)
from ..register import register_provider_adapter
@register_provider_adapter(
"anthropic_chat_completion",
@@ -12,12 +12,11 @@ from httpx import AsyncClient, Timeout
from astrbot import logger
from astrbot.core.config.default import VERSION
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
TEMP_DIR = Path(get_astrbot_temp_path()) / "azure_tts"
TEMP_DIR.mkdir(parents=True, exist_ok=True)
AZURE_TTS_SUBSCRIPTION_KEY_PATTERN = r"^(?:[a-zA-Z0-9]{32}|[a-zA-Z0-9]{84})$"
@@ -3,10 +3,9 @@ import os
import aiohttp
from astrbot import logger
from ..entities import ProviderType, RerankResult
from ..provider import RerankProvider
from ..register import register_provider_adapter
from astrbot.core.provider.entities import ProviderType, RerankResult
from astrbot.core.provider.provider import RerankProvider
from astrbot.core.provider.register import register_provider_adapter
class BailianRerankError(Exception):
@@ -16,12 +16,11 @@ except (
): # pragma: no cover - older dashscope versions without Qwen TTS support
MultiModalConversation = None
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
@register_provider_adapter(
"dashscope_tts",
@@ -7,12 +7,11 @@ import anyio
import edge_tts
from astrbot.core import logger
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
"""
edge_tts 方式,能够免费、快速生成语音,使用需要先安装edge-tts库
```
@@ -9,12 +9,11 @@ from httpx import AsyncClient
from pydantic import BaseModel, conint
from astrbot import logger
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
class ServeReferenceAudio(BaseModel):
audio: bytes
@@ -5,10 +5,9 @@ from google.genai import types
from google.genai.errors import APIError
from astrbot import logger
from ..entities import ProviderType
from ..provider import EmbeddingProvider
from ..register import register_provider_adapter
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import EmbeddingProvider
from astrbot.core.provider.register import register_provider_adapter
@register_provider_adapter(
@@ -18,11 +18,10 @@ from astrbot.core.agent.message import ContentPart, ImageURLPart, TextPart
from astrbot.core.message.message_event_result import MessageChain
from astrbot.core.provider.entities import LLMResponse, TokenUsage
from astrbot.core.provider.func_tool_manager import ToolSet
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.io import download_image_by_url
from astrbot.core.utils.network_utils import is_connection_error, log_connection_failure
from ..register import register_provider_adapter
class SuppressNonTextPartsWarning(logging.Filter):
"""过滤 Gemini SDK 中的非文本部分警告"""
@@ -6,12 +6,11 @@ from google import genai
from google.genai import types
from astrbot import logger
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
@register_provider_adapter(
"gemini_tts",
+2 -1
View File
@@ -1,4 +1,5 @@
from ..register import register_provider_adapter
from astrbot.core.provider.register import register_provider_adapter
from .openai_source import ProviderOpenAIOfficial
@@ -6,12 +6,11 @@ import aiofiles
import aiohttp
from astrbot import logger
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
@register_provider_adapter(
provider_type_name="gsv_tts_selfhost",
@@ -5,12 +5,11 @@ import uuid
import aiofiles
import aiohttp
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
@register_provider_adapter(
"gsvi_tts_api",
@@ -1,4 +1,5 @@
from ..register import register_provider_adapter
from astrbot.core.provider.register import register_provider_adapter
from .anthropic_source import ProviderAnthropic
KIMI_CODE_API_BASE = "https://api.kimi.com/coding"
@@ -1,6 +1,7 @@
from ..entities import ProviderType
from ..provider import STTProvider
from ..register import register_provider_adapter
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import STTProvider
from astrbot.core.provider.register import register_provider_adapter
from .mimo_api_common import (
DEFAULT_MIMO_API_BASE,
DEFAULT_MIMO_STT_MODEL,
@@ -1,9 +1,10 @@
import base64
import uuid
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from .mimo_api_common import (
DEFAULT_MIMO_API_BASE,
DEFAULT_MIMO_TTS_MODEL,
@@ -7,12 +7,11 @@ import aiofiles
import aiohttp
from astrbot.api import logger
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
@register_provider_adapter(
"minimax_tts_api",
@@ -1,7 +1,8 @@
from collections.abc import MutableMapping
from typing import cast
from ..register import register_provider_adapter
from astrbot.core.provider.register import register_provider_adapter
from .openai_source import ProviderOpenAIOfficial
@@ -2,10 +2,9 @@ import httpx
from openai import AsyncOpenAI
from astrbot import logger
from ..entities import ProviderType
from ..provider import EmbeddingProvider
from ..register import register_provider_adapter
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import EmbeddingProvider
from astrbot.core.provider.register import register_provider_adapter
@register_provider_adapter(
@@ -6,12 +6,11 @@ import httpx
from openai import NOT_GIVEN, AsyncOpenAI
from astrbot import logger
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
@register_provider_adapter(
"openai_tts_api",
@@ -1,7 +1,8 @@
from collections.abc import MutableMapping
from typing import cast
from ..register import register_provider_adapter
from astrbot.core.provider.register import register_provider_adapter
from .openai_source import ProviderOpenAIOfficial
@@ -13,14 +13,13 @@ from funasr_onnx import SenseVoiceSmall
from funasr_onnx.utils.postprocess_utils import rich_transcription_postprocess
from astrbot.core import logger
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import STTProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from astrbot.core.utils.io import download_file
from astrbot.core.utils.tencent_record_helper import tencent_silk_to_wav
from ..entities import ProviderType
from ..provider import STTProvider
from ..register import register_provider_adapter
@register_provider_adapter(
"sensevoice_stt_selfhost",
@@ -1,10 +1,9 @@
import aiohttp
from astrbot import logger
from ..entities import ProviderType, RerankResult
from ..provider import RerankProvider
from ..register import register_provider_adapter
from astrbot.core.provider.entities import ProviderType, RerankResult
from astrbot.core.provider.provider import RerankProvider
from astrbot.core.provider.register import register_provider_adapter
@register_provider_adapter(
@@ -8,12 +8,11 @@ import aiohttp
import anyio
from astrbot import logger
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import TTSProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from ..entities import ProviderType
from ..provider import TTSProvider
from ..register import register_provider_adapter
@register_provider_adapter(
"volcengine_tts",
@@ -5,6 +5,9 @@ import anyio
from openai import NOT_GIVEN, AsyncOpenAI
from astrbot.core import logger
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import STTProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from astrbot.core.utils.io import download_file
from astrbot.core.utils.media_utils import convert_audio_to_wav
@@ -13,10 +16,6 @@ from astrbot.core.utils.tencent_record_helper import (
tencent_silk_to_wav,
)
from ..entities import ProviderType
from ..provider import STTProvider
from ..register import register_provider_adapter
def _open_file_rb(path: str):
return open(path, "rb")
@@ -7,14 +7,13 @@ import anyio
import whisper
from astrbot.core import logger
from astrbot.core.provider.entities import ProviderType
from astrbot.core.provider.provider import STTProvider
from astrbot.core.provider.register import register_provider_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from astrbot.core.utils.io import download_file
from astrbot.core.utils.tencent_record_helper import tencent_silk_to_wav
from ..entities import ProviderType
from ..provider import STTProvider
from ..register import register_provider_adapter
@register_provider_adapter(
"openai_whisper_selfhost",
+2 -1
View File
@@ -1,4 +1,5 @@
from ..register import register_provider_adapter
from astrbot.core.provider.register import register_provider_adapter
from .openai_source import ProviderOpenAIOfficial
+13 -10
View File
@@ -389,19 +389,22 @@ async def compress_image(
if len(data) < min_file_size_bytes:
return url_or_path
else:
local_path = Path(url_or_path)
if not local_path.exists():
return url_or_path
if local_path.stat().st_size < min_file_size_bytes:
return url_or_path
with local_path.open("rb") as f:
data = f.read()
if not data:
return url_or_path
def _read_local_path():
lp = Path(url_or_path)
if not lp.exists():
return None
if lp.stat().st_size < min_file_size_bytes:
return None
with lp.open("rb") as f:
return f.read()
data = await asyncio.to_thread(_read_local_path)
if not data:
return url_or_path
temp_dir = Path(get_astrbot_temp_path())
temp_dir.mkdir(parents=True, exist_ok=True)
await asyncio.to_thread(temp_dir.mkdir, parents=True, exist_ok=True)
# Offload the blocking image processing task to a thread.
return await asyncio.to_thread(
+12 -12
View File
@@ -65,7 +65,7 @@ class ApiKeyRoute(Route):
async def list_api_keys(self):
keys = await self.db.list_api_keys()
return (
Response().ok(data=[self._serialize_api_key(key) for key in keys]).__dict__
Response().ok(data=[self._serialize_api_key(key) for key in keys]).to_json()
)
async def create_api_key(self):
@@ -84,9 +84,9 @@ class ApiKeyRoute(Route):
]
normalized_scopes = list(dict.fromkeys(normalized_scopes))
if not normalized_scopes:
return Response().error("At least one valid scope is required").__dict__
return Response().error("At least one valid scope is required").to_json()
else:
return Response().error("Invalid scopes").__dict__
return Response().error("Invalid scopes").to_json()
expires_at = None
expires_in_days = post_data.get("expires_in_days")
@@ -94,10 +94,10 @@ class ApiKeyRoute(Route):
try:
expires_in_days_int = int(expires_in_days)
except (TypeError, ValueError):
return Response().error("expires_in_days must be an integer").__dict__
return Response().error("expires_in_days must be an integer").to_json()
if expires_in_days_int <= 0:
return (
Response().error("expires_in_days must be greater than 0").__dict__
Response().error("expires_in_days must be greater than 0").to_json()
)
expires_at = datetime.now(timezone.utc) + timedelta(
days=expires_in_days_int
@@ -119,26 +119,26 @@ class ApiKeyRoute(Route):
payload = self._serialize_api_key(api_key)
payload["api_key"] = raw_key
return Response().ok(data=payload).__dict__
return Response().ok(data=payload).to_json()
async def revoke_api_key(self):
post_data = await request.json or {}
key_id = post_data.get("key_id")
if not key_id:
return Response().error("Missing key: key_id").__dict__
return Response().error("Missing key: key_id").to_json()
success = await self.db.revoke_api_key(key_id)
if not success:
return Response().error("API key not found").__dict__
return Response().ok().__dict__
return Response().error("API key not found").to_json()
return Response().ok().to_json()
async def delete_api_key(self):
post_data = await request.json or {}
key_id = post_data.get("key_id")
if not key_id:
return Response().error("Missing key: key_id").__dict__
return Response().error("Missing key: key_id").to_json()
success = await self.db.delete_api_key(key_id)
if not success:
return Response().error("API key not found").__dict__
return Response().ok().__dict__
return Response().error("API key not found").to_json()
return Response().ok().to_json()
+10 -13
View File
@@ -41,15 +41,14 @@ class AuthRoute(Route):
# Security: Require non-empty credentials
if not input_username or not input_password:
return Response().error("用户名和密码不能为空").__dict__
return Response().error("用户名和密码不能为空").to_json()
# Check if password has been configured via CLI
if not self._is_password_set(stored_password_hash):
await asyncio.sleep(3)
return (
Response()
.error("管理员密码未设置,请先运行 'astrbot conf admin' 命令设置密码")
.__dict__
.error("管理员密码未设置,请先运行 'astrbot conf admin' 命令设置密码").to_json()
)
# Normal login flow - credentials must match stored admin account
@@ -65,50 +64,48 @@ class AuthRoute(Route):
"username": stored_username,
"change_pwd_hint": False,
},
)
.__dict__
).to_json()
)
# Security: Don't reveal whether it's username or password error
await asyncio.sleep(3)
return Response().error("用户名或密码错误").__dict__
return Response().error("用户名或密码错误").to_json()
async def edit_account(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
stored_password_hash = self.config["dashboard"]["password"]
post_data = await request.json
if not self._matches_dashboard_password(stored_password_hash, post_data):
return Response().error("原密码错误").__dict__
return Response().error("原密码错误").to_json()
new_pwd = post_data.get("new_password", None)
new_username = post_data.get("new_username", None)
if not new_pwd and not new_username:
return Response().error("新用户名和新密码不能同时为空").__dict__
return Response().error("新用户名和新密码不能同时为空").to_json()
# Verify password confirmation
if new_pwd:
confirm_pwd = post_data.get("confirm_password", None)
if confirm_pwd != new_pwd:
return Response().error("两次输入的新密码不一致").__dict__
return Response().error("两次输入的新密码不一致").to_json()
# Hash the new password before storing to ensure backend and CLI use the same format
try:
new_hash = hash_dashboard_password_secure(new_pwd)
except Exception as e:
return Response().error(f"Failed to hash new password: {e}").__dict__
return Response().error(f"Failed to hash new password: {e}").to_json()
self.config["dashboard"]["password"] = new_hash
if new_username:
self.config["dashboard"]["username"] = new_username
self.config.save_config()
return Response().ok(None, "修改成功").__dict__
return Response().ok(None, "修改成功").to_json()
def generate_jwt(self, username):
payload = {
+36 -40
View File
@@ -97,7 +97,7 @@ class ChatRoute(Route):
async def get_file(self):
filename = request.args.get("filename")
if not filename:
return Response().error("Missing key: filename").__dict__
return Response().error("Missing key: filename").to_json()
try:
file_path = os.path.join(self.attachments_dir, os.path.basename(filename))
@@ -116,7 +116,7 @@ class ChatRoute(Route):
try:
resolved_file_path.relative_to(resolved_base_dir)
except ValueError:
return Response().error("Invalid file path").__dict__
return Response().error("Invalid file path").to_json()
filename_ext = os.path.splitext(filename)[1].lower()
if filename_ext == ".wav":
@@ -126,18 +126,18 @@ class ChatRoute(Route):
return await send_file(str(resolved_file_path))
except (FileNotFoundError, OSError):
return Response().error("File access error").__dict__
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").__dict__
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").__dict__
return Response().error("Attachment not found").to_json()
file_path = attachment.path
resolved_file_path = _resolve_path(file_path)
@@ -147,13 +147,13 @@ class ChatRoute(Route):
)
except (FileNotFoundError, OSError):
return Response().error("File access error").__dict__
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").__dict__
return Response().error("Missing key: file").to_json()
file = post_data["file"]
filename = file.filename or f"{uuid.uuid4()!s}"
@@ -180,7 +180,7 @@ class ChatRoute(Route):
)
if not attachment:
return Response().error("Failed to create attachment").__dict__
return Response().error("Failed to create attachment").to_json()
filename = os.path.basename(attachment.path)
@@ -192,8 +192,7 @@ class ChatRoute(Route):
"filename": filename,
"type": attach_type,
}
)
.__dict__
).to_json()
)
async def _build_user_message_parts(self, message: str | list) -> list[dict]:
@@ -313,13 +312,13 @@ class ChatRoute(Route):
if post_data is None:
post_data = await request.json
if post_data is None:
return Response().error("Missing JSON body").__dict__
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").__dict__
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").__dict__
Response().error("Missing key: session_id or conversation_id").to_json()
)
message = post_data["message"]
@@ -329,7 +328,7 @@ class ChatRoute(Route):
enable_streaming = post_data.get("enable_streaming", True)
if not session_id:
return Response().error("session_id is empty").__dict__
return Response().error("session_id is empty").to_json()
webchat_conv_id = session_id
@@ -338,8 +337,7 @@ class ChatRoute(Route):
if not webchat_message_parts_have_content(message_parts):
return (
Response()
.error("Message content is empty (reply only is not allowed)")
.__dict__
.error("Message content is empty (reply only is not allowed)").to_json()
)
message_id = str(uuid.uuid4())
@@ -573,18 +571,18 @@ class ChatRoute(Route):
"""Stop active agent runs for a session."""
post_data = await request.json
if post_data is None:
return Response().error("Missing JSON body").__dict__
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").__dict__
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").__dict__
return Response().error(f"Session {session_id} not found").to_json()
if session.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
message_type = (
MessageType.GROUP_MESSAGE.value
@@ -597,7 +595,7 @@ class ChatRoute(Route):
)
stopped_count = active_event_registry.request_agent_stop_all(umo)
return Response().ok(data={"stopped_count": stopped_count}).__dict__
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."""
@@ -647,30 +645,30 @@ class ChatRoute(Route):
"""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").__dict__
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").__dict__
return Response().error(f"Session {session_id} not found").to_json()
if session.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
await self._delete_session_internal(session, username)
return Response().ok().__dict__
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").__dict__
return Response().error("Missing JSON body").to_json()
if not isinstance(post_data, dict):
return Response().error("Invalid JSON body: expected object").__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").__dict__
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)
@@ -703,8 +701,7 @@ class ChatRoute(Route):
"failed_count": len(failed_items),
"failed_items": failed_items,
}
)
.__dict__
).to_json()
)
def _extract_attachment_ids(self, history_list) -> list[str]:
@@ -763,8 +760,7 @@ class ChatRoute(Route):
"session_id": session.session_id,
"platform_id": session.platform_id,
}
)
.__dict__
).to_json()
)
async def get_sessions(self):
@@ -799,13 +795,13 @@ class ChatRoute(Route):
}
)
return Response().ok(data=sessions_data).__dict__
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").__dict__
return Response().error("Missing key: session_id").to_json()
# 获取会话信息以确定 platform_id
session = await self.db.get_platform_session_by_id(session_id)
@@ -840,7 +836,7 @@ class ChatRoute(Route):
"emoji": project_info.emoji,
}
return Response().ok(data=response_data).__dict__
return Response().ok(data=response_data).to_json()
async def update_session_display_name(self):
"""Update a Platform session's display name."""
@@ -850,18 +846,18 @@ class ChatRoute(Route):
display_name = post_data.get("display_name")
if not session_id:
return Response().error("Missing key: session_id").__dict__
return Response().error("Missing key: session_id").to_json()
if display_name is None:
return Response().error("Missing key: display_name").__dict__
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").__dict__
return Response().error(f"Session {session_id} not found").to_json()
if session.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
# 更新 display_name
await self.db.update_platform_session(
@@ -869,4 +865,4 @@ class ChatRoute(Route):
display_name=display_name,
)
return Response().ok().__dict__
return Response().ok().to_json()
+30 -32
View File
@@ -35,7 +35,7 @@ class ChatUIProjectRoute(Route):
description = post_data.get("description")
if not title:
return Response().error("Missing key: title").__dict__
return Response().error("Missing key: title").to_json()
project = await self.db.create_chatui_project(
creator=username,
@@ -55,8 +55,7 @@ class ChatUIProjectRoute(Route):
"created_at": to_utc_isoformat(project.created_at),
"updated_at": to_utc_isoformat(project.updated_at),
}
)
.__dict__
).to_json()
)
async def list_projects(self):
@@ -77,23 +76,23 @@ class ChatUIProjectRoute(Route):
for project in projects
]
return Response().ok(data=projects_data).__dict__
return Response().ok(data=projects_data).to_json()
async def get_project(self):
"""Get a specific ChatUI project."""
project_id = request.args.get("project_id")
if not project_id:
return Response().error("Missing key: project_id").__dict__
return Response().error("Missing key: project_id").to_json()
username = g.get("username", "guest")
project = await self.db.get_chatui_project_by_id(project_id)
if not project:
return Response().error(f"Project {project_id} not found").__dict__
return Response().error(f"Project {project_id} not found").to_json()
# Verify ownership
if project.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
return (
Response()
@@ -106,8 +105,7 @@ class ChatUIProjectRoute(Route):
"created_at": to_utc_isoformat(project.created_at),
"updated_at": to_utc_isoformat(project.updated_at),
}
)
.__dict__
).to_json()
)
async def update_chatui_project(self):
@@ -120,16 +118,16 @@ class ChatUIProjectRoute(Route):
description = post_data.get("description")
if not project_id:
return Response().error("Missing key: project_id").__dict__
return Response().error("Missing key: project_id").to_json()
username = g.get("username", "guest")
# Verify ownership
project = await self.db.get_chatui_project_by_id(project_id)
if not project:
return Response().error(f"Project {project_id} not found").__dict__
return Response().error(f"Project {project_id} not found").to_json()
if project.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
await self.db.update_chatui_project(
project_id=project_id,
@@ -138,26 +136,26 @@ class ChatUIProjectRoute(Route):
description=description,
)
return Response().ok().__dict__
return Response().ok().to_json()
async def delete_project(self):
"""Delete a ChatUI project."""
project_id = request.args.get("project_id")
if not project_id:
return Response().error("Missing key: project_id").__dict__
return Response().error("Missing key: project_id").to_json()
username = g.get("username", "guest")
# Verify ownership
project = await self.db.get_chatui_project_by_id(project_id)
if not project:
return Response().error(f"Project {project_id} not found").__dict__
return Response().error(f"Project {project_id} not found").to_json()
if project.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
await self.db.delete_chatui_project(project_id)
return Response().ok().__dict__
return Response().ok().to_json()
async def add_session_to_project(self):
"""Add a session to a project."""
@@ -167,29 +165,29 @@ class ChatUIProjectRoute(Route):
project_id = post_data.get("project_id")
if not session_id:
return Response().error("Missing key: session_id").__dict__
return Response().error("Missing key: session_id").to_json()
if not project_id:
return Response().error("Missing key: project_id").__dict__
return Response().error("Missing key: project_id").to_json()
username = g.get("username", "guest")
# Verify project ownership
project = await self.db.get_chatui_project_by_id(project_id)
if not project:
return Response().error(f"Project {project_id} not found").__dict__
return Response().error(f"Project {project_id} not found").to_json()
if project.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
# Verify session ownership
session = await self.db.get_platform_session_by_id(session_id)
if not session:
return Response().error(f"Session {session_id} not found").__dict__
return Response().error(f"Session {session_id} not found").to_json()
if session.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
await self.db.add_session_to_project(session_id, project_id)
return Response().ok().__dict__
return Response().ok().to_json()
async def remove_session_from_project(self):
"""Remove a session from its project."""
@@ -198,35 +196,35 @@ class ChatUIProjectRoute(Route):
session_id = post_data.get("session_id")
if not session_id:
return Response().error("Missing key: session_id").__dict__
return Response().error("Missing key: session_id").to_json()
username = g.get("username", "guest")
# Verify session ownership
session = await self.db.get_platform_session_by_id(session_id)
if not session:
return Response().error(f"Session {session_id} not found").__dict__
return Response().error(f"Session {session_id} not found").to_json()
if session.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
await self.db.remove_session_from_project(session_id)
return Response().ok().__dict__
return Response().ok().to_json()
async def get_project_sessions(self):
"""Get all sessions in a project."""
project_id = request.args.get("project_id")
if not project_id:
return Response().error("Missing key: project_id").__dict__
return Response().error("Missing key: project_id").to_json()
username = g.get("username", "guest")
# Verify project ownership
project = await self.db.get_chatui_project_by_id(project_id)
if not project:
return Response().error(f"Project {project_id} not found").__dict__
return Response().error(f"Project {project_id} not found").to_json()
if project.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
sessions = await self.db.get_project_sessions(project_id)
@@ -243,4 +241,4 @@ class ChatUIProjectRoute(Route):
for session in sessions
]
return Response().ok(data=sessions_data).__dict__
return Response().ok(data=sessions_data).to_json()
+11 -11
View File
@@ -36,11 +36,11 @@ class CommandRoute(Route):
"disabled": len([cmd for cmd in commands if not cmd["enabled"]]),
"conflicts": len([cmd for cmd in commands if cmd.get("has_conflict")]),
}
return Response().ok({"items": commands, "summary": summary}).__dict__
return Response().ok({"items": commands, "summary": summary}).to_json()
async def get_conflicts(self):
conflicts = await list_command_conflicts()
return Response().ok(conflicts).__dict__
return Response().ok(conflicts).to_json()
async def toggle_command(self):
data = await request.get_json()
@@ -48,7 +48,7 @@ class CommandRoute(Route):
enabled = data.get("enabled")
if handler_full_name is None or enabled is None:
return Response().error("handler_full_name 与 enabled 均为必填。").__dict__
return Response().error("handler_full_name 与 enabled 均为必填。").to_json()
if isinstance(enabled, str):
enabled = enabled.lower() in ("1", "true", "yes", "on")
@@ -56,10 +56,10 @@ class CommandRoute(Route):
try:
await toggle_command_service(handler_full_name, bool(enabled))
except ValueError as exc:
return Response().error(str(exc)).__dict__
return Response().error(str(exc)).to_json()
payload = await _get_command_payload(handler_full_name)
return Response().ok(payload).__dict__
return Response().ok(payload).to_json()
async def rename_command(self):
data = await request.get_json()
@@ -68,15 +68,15 @@ class CommandRoute(Route):
aliases = data.get("aliases")
if not handler_full_name or not new_name:
return Response().error("handler_full_name 与 new_name 均为必填。").__dict__
return Response().error("handler_full_name 与 new_name 均为必填。").to_json()
try:
await rename_command_service(handler_full_name, new_name, aliases=aliases)
except ValueError as exc:
return Response().error(str(exc)).__dict__
return Response().error(str(exc)).to_json()
payload = await _get_command_payload(handler_full_name)
return Response().ok(payload).__dict__
return Response().ok(payload).to_json()
async def update_permission(self):
data = await request.get_json()
@@ -85,16 +85,16 @@ class CommandRoute(Route):
if not handler_full_name or not permission:
return (
Response().error("handler_full_name 与 permission 均为必填。").__dict__
Response().error("handler_full_name 与 permission 均为必填。").to_json()
)
try:
await update_command_permission_service(handler_full_name, permission)
except ValueError as exc:
return Response().error(str(exc)).__dict__
return Response().error(str(exc)).to_json()
payload = await _get_command_payload(handler_full_name)
return Response().ok(payload).__dict__
return Response().ok(payload).to_json()
async def _get_command_payload(handler_full_name: str):
+25 -27
View File
@@ -81,7 +81,7 @@ class ConversationRoute(Route):
)
except Exception as e:
logger.error(f"数据库查询出错: {e!s}\n{traceback.format_exc()}")
return Response().error(f"数据库查询出错: {e!s}").__dict__
return Response().error(f"数据库查询出错: {e!s}").to_json()
# 计算总页数
total_pages = (
@@ -97,12 +97,12 @@ class ConversationRoute(Route):
"total_pages": total_pages,
},
}
return Response().ok(result).__dict__
return Response().ok(result).to_json()
except Exception as e:
error_msg = f"获取对话列表失败: {e!s}\n{traceback.format_exc()}"
logger.error(error_msg)
return Response().error(f"获取对话列表失败: {e!s}").__dict__
return Response().error(f"获取对话列表失败: {e!s}").to_json()
async def get_conv_detail(self):
"""获取指定对话详情(通过POST请求)"""
@@ -112,14 +112,14 @@ class ConversationRoute(Route):
cid = data.get("cid")
if not user_id or not cid:
return Response().error("缺少必要参数: user_id 和 cid").__dict__
return Response().error("缺少必要参数: user_id 和 cid").to_json()
conversation = await self.conv_mgr.get_conversation(
unified_msg_origin=user_id,
conversation_id=cid,
)
if not conversation:
return Response().error("对话不存在").__dict__
return Response().error("对话不存在").to_json()
return (
Response()
@@ -133,13 +133,12 @@ class ConversationRoute(Route):
"created_at": conversation.created_at,
"updated_at": conversation.updated_at,
},
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"获取对话详情失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"获取对话详情失败: {e!s}").__dict__
return Response().error(f"获取对话详情失败: {e!s}").to_json()
async def upd_conv(self):
"""更新对话信息(标题和角色ID)"""
@@ -150,13 +149,13 @@ class ConversationRoute(Route):
title = data.get("title")
if not user_id or not cid:
return Response().error("缺少必要参数: user_id 和 cid").__dict__
return Response().error("缺少必要参数: user_id 和 cid").to_json()
conversation = await self.conv_mgr.get_conversation(
unified_msg_origin=user_id,
conversation_id=cid,
)
if not conversation:
return Response().error("对话不存在").__dict__
return Response().error("对话不存在").to_json()
persona_id = data.get("persona_id", conversation.persona_id)
@@ -167,11 +166,11 @@ class ConversationRoute(Route):
title=title,
persona_id=persona_id,
)
return Response().ok({"message": "对话信息更新成功"}).__dict__
return Response().ok({"message": "对话信息更新成功"}).to_json()
except Exception as e:
logger.error(f"更新对话信息失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"更新对话信息失败: {e!s}").__dict__
return Response().error(f"更新对话信息失败: {e!s}").to_json()
async def del_conv(self):
"""删除对话"""
@@ -184,7 +183,7 @@ class ConversationRoute(Route):
conversations = data.get("conversations", [])
if not conversations:
return (
Response().error("批量删除时conversations参数不能为空").__dict__
Response().error("批量删除时conversations参数不能为空").to_json()
)
deleted_count = 0
@@ -222,25 +221,24 @@ class ConversationRoute(Route):
"failed_count": len(failed_items),
"failed_items": failed_items,
},
)
.__dict__
).to_json()
)
# 单个删除
user_id = data.get("user_id")
cid = data.get("cid")
if not user_id or not cid:
return Response().error("缺少必要参数: user_id 和 cid").__dict__
return Response().error("缺少必要参数: user_id 和 cid").to_json()
await self.core_lifecycle.conversation_manager.delete_conversation(
unified_msg_origin=user_id,
conversation_id=cid,
)
return Response().ok({"message": "对话删除成功"}).__dict__
return Response().ok({"message": "对话删除成功"}).to_json()
except Exception as e:
logger.error(f"删除对话失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"删除对话失败: {e!s}").__dict__
return Response().error(f"删除对话失败: {e!s}").to_json()
async def update_history(self):
"""更新对话历史内容"""
@@ -251,10 +249,10 @@ class ConversationRoute(Route):
history = data.get("history")
if not user_id or not cid:
return Response().error("缺少必要参数: user_id 和 cid").__dict__
return Response().error("缺少必要参数: user_id 和 cid").to_json()
if history is None:
return Response().error("缺少必要参数: history").__dict__
return Response().error("缺少必要参数: history").to_json()
# 历史记录必须是合法的 JSON 字符串
try:
@@ -265,7 +263,7 @@ class ConversationRoute(Route):
json.loads(history)
except json.JSONDecodeError:
return (
Response().error("history 必须是有效的 JSON 字符串或数组").__dict__
Response().error("history 必须是有效的 JSON 字符串或数组").to_json()
)
conversation = await self.conv_mgr.get_conversation(
@@ -273,7 +271,7 @@ class ConversationRoute(Route):
conversation_id=cid,
)
if not conversation:
return Response().error("对话不存在").__dict__
return Response().error("对话不存在").to_json()
history = json.loads(history) if isinstance(history, str) else history
@@ -283,11 +281,11 @@ class ConversationRoute(Route):
history=history,
)
return Response().ok({"message": "对话历史更新成功"}).__dict__
return Response().ok({"message": "对话历史更新成功"}).to_json()
except Exception as e:
logger.error(f"更新对话历史失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"更新对话历史失败: {e!s}").__dict__
return Response().error(f"更新对话历史失败: {e!s}").to_json()
async def export_conversations(self):
"""批量导出对话为 JSONL 格式"""
@@ -296,7 +294,7 @@ class ConversationRoute(Route):
conversations_to_export = data.get("conversations", [])
if not conversations_to_export:
return Response().error("导出列表不能为空").__dict__
return Response().error("导出列表不能为空").to_json()
# 收集所有对话的内容
jsonl_lines = []
@@ -351,7 +349,7 @@ class ConversationRoute(Route):
)
if exported_count == 0:
return Response().error("没有成功导出任何对话").__dict__
return Response().error("没有成功导出任何对话").to_json()
# 创建 JSONL 内容
jsonl_content = "\n".join(jsonl_lines)
@@ -374,4 +372,4 @@ class ConversationRoute(Route):
except Exception as e:
logger.error(f"批量导出对话失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"批量导出对话失败: {e!s}").__dict__
return Response().error(f"批量导出对话失败: {e!s}").to_json()
+19 -20
View File
@@ -42,27 +42,27 @@ class CronRoute(Route):
cron_mgr = self.core_lifecycle.cron_manager
if cron_mgr is None:
return jsonify(
Response().error("Cron manager not initialized").__dict__
Response().error("Cron manager not initialized").to_json()
)
job_type = request.args.get("type")
jobs = await cron_mgr.list_jobs(job_type)
data = [self._serialize_job(j) for j in jobs]
return jsonify(Response().ok(data=data).__dict__)
return jsonify(Response().ok(data=data).to_json())
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"Failed to list jobs: {e!s}").__dict__)
return jsonify(Response().error(f"Failed to list jobs: {e!s}").to_json())
async def create_job(self):
try:
cron_mgr = self.core_lifecycle.cron_manager
if cron_mgr is None:
return jsonify(
Response().error("Cron manager not initialized").__dict__
Response().error("Cron manager not initialized").to_json()
)
payload = await request.json
if not isinstance(payload, dict):
return jsonify(Response().error("Invalid payload").__dict__)
return jsonify(Response().error("Invalid payload").to_json())
name = payload.get("name") or "active_agent_task"
cron_expression = payload.get("cron_expression")
@@ -76,16 +76,15 @@ class CronRoute(Route):
run_at = payload.get("run_at")
if not session:
return jsonify(Response().error("session is required").__dict__)
return jsonify(Response().error("session is required").to_json())
if run_once and not run_at:
return jsonify(
Response().error("run_at is required when run_once=true").__dict__
Response().error("run_at is required when run_once=true").to_json()
)
if (not run_once) and not cron_expression:
return jsonify(
Response()
.error("cron_expression is required when run_once=false")
.__dict__
.error("cron_expression is required when run_once=false").to_json()
)
if run_once and cron_expression:
cron_expression = None # ignore cron when run_once specified
@@ -95,7 +94,7 @@ class CronRoute(Route):
run_at_dt = datetime.fromisoformat(str(run_at))
except Exception:
return jsonify(
Response().error("run_at must be ISO datetime").__dict__
Response().error("run_at must be ISO datetime").to_json()
)
job_payload = {
@@ -118,22 +117,22 @@ class CronRoute(Route):
run_at=run_at_dt,
)
return jsonify(Response().ok(data=self._serialize_job(job)).__dict__)
return jsonify(Response().ok(data=self._serialize_job(job)).to_json())
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"Failed to create job: {e!s}").__dict__)
return jsonify(Response().error(f"Failed to create job: {e!s}").to_json())
async def update_job(self, job_id: str):
try:
cron_mgr = self.core_lifecycle.cron_manager
if cron_mgr is None:
return jsonify(
Response().error("Cron manager not initialized").__dict__
Response().error("Cron manager not initialized").to_json()
)
payload = await request.json
if not isinstance(payload, dict):
return jsonify(Response().error("Invalid payload").__dict__)
return jsonify(Response().error("Invalid payload").to_json())
updates = {
"name": payload.get("name"),
@@ -154,21 +153,21 @@ class CronRoute(Route):
job = await cron_mgr.update_job(job_id, **updates)
if not job:
return jsonify(Response().error("Job not found").__dict__)
return jsonify(Response().ok(data=self._serialize_job(job)).__dict__)
return jsonify(Response().error("Job not found").to_json())
return jsonify(Response().ok(data=self._serialize_job(job)).to_json())
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"Failed to update job: {e!s}").__dict__)
return jsonify(Response().error(f"Failed to update job: {e!s}").to_json())
async def delete_job(self, job_id: str):
try:
cron_mgr = self.core_lifecycle.cron_manager
if cron_mgr is None:
return jsonify(
Response().error("Cron manager not initialized").__dict__
Response().error("Cron manager not initialized").to_json()
)
await cron_mgr.delete_job(job_id)
return jsonify(Response().ok(message="deleted").__dict__)
return jsonify(Response().ok(message="deleted").to_json())
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"Failed to delete job: {e!s}").__dict__)
return jsonify(Response().error(f"Failed to delete job: {e!s}").to_json())
+7 -8
View File
@@ -110,35 +110,34 @@ class LogRoute(Route):
data={
"logs": logs,
},
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"获取日志历史失败: {e}")
return Response().error(f"获取日志历史失败: {e}").__dict__
return Response().error(f"获取日志历史失败: {e}").to_json()
async def get_trace_settings(self):
"""获取 Trace 设置"""
try:
trace_enable = self.config.get("trace_enable", True)
return Response().ok(data={"trace_enable": trace_enable}).__dict__
return Response().ok(data={"trace_enable": trace_enable}).to_json()
except Exception as e:
logger.error(f"获取 Trace 设置失败: {e}")
return Response().error(f"获取 Trace 设置失败: {e}").__dict__
return Response().error(f"获取 Trace 设置失败: {e}").to_json()
async def update_trace_settings(self):
"""更新 Trace 设置"""
try:
data = await request.json
if data is None:
return Response().error("请求数据为空").__dict__
return Response().error("请求数据为空").to_json()
trace_enable = data.get("trace_enable")
if trace_enable is not None:
self.config["trace_enable"] = bool(trace_enable)
self.config.save_config()
return Response().ok(message="Trace 设置已更新").__dict__
return Response().ok(message="Trace 设置已更新").to_json()
except Exception as e:
logger.error(f"更新 Trace 设置失败: {e}")
return Response().error(f"更新 Trace 设置失败: {e}").__dict__
return Response().error(f"更新 Trace 设置失败: {e}").to_json()
+17 -20
View File
@@ -150,9 +150,9 @@ class OpenApiRoute(Route):
post_data.get("username")
)
if username_err:
return Response().error(username_err).__dict__
return Response().error(username_err).to_json()
if not effective_username:
return Response().error("Invalid username").__dict__
return Response().error("Invalid username").to_json()
raw_session_id = post_data.get("session_id", post_data.get("conversation_id"))
session_id = str(raw_session_id).strip() if raw_session_id is not None else ""
@@ -164,11 +164,11 @@ class OpenApiRoute(Route):
session_id,
)
if ensure_session_err:
return Response().error(ensure_session_err).__dict__
return Response().error(ensure_session_err).to_json()
config_id, resolve_err = self._resolve_chat_config_id(post_data)
if resolve_err:
return Response().error(resolve_err).__dict__
return Response().error(resolve_err).to_json()
original_username = g.get("username", "guest")
g.username = effective_username
@@ -191,8 +191,7 @@ class OpenApiRoute(Route):
)
return (
Response()
.error(f"Failed to update chat config route: {e}")
.__dict__
.error(f"Failed to update chat config route: {e}").to_json()
)
try:
return await self.chat_route.chat(post_data=post_data)
@@ -574,7 +573,7 @@ class OpenApiRoute(Route):
request.args.get("username")
)
if username_err:
return Response().error(username_err).__dict__
return Response().error(username_err).to_json()
assert username is not None # for type checker
@@ -582,7 +581,7 @@ class OpenApiRoute(Route):
page = int(request.args.get("page", 1))
page_size = int(request.args.get("page_size", 20))
except ValueError:
return Response().error("page and page_size must be integers").__dict__
return Response().error("page and page_size must be integers").to_json()
if page < 1:
page = 1
@@ -628,13 +627,12 @@ class OpenApiRoute(Route):
"page_size": page_size,
"total": total,
}
)
.__dict__
).to_json()
)
async def get_chat_configs(self):
conf_list = self._get_chat_config_list()
return Response().ok(data={"configs": conf_list}).__dict__
return Response().ok(data={"configs": conf_list}).to_json()
async def _build_message_chain_from_payload(
self,
@@ -652,14 +650,14 @@ class OpenApiRoute(Route):
umo = post_data.get("umo")
if message_payload is None:
return Response().error("Missing key: message").__dict__
return Response().error("Missing key: message").to_json()
if not umo:
return Response().error("Missing key: umo").__dict__
return Response().error("Missing key: umo").to_json()
try:
session = MessageSesion.from_str(str(umo))
except Exception as e:
return Response().error(f"Invalid umo: {e}").__dict__
return Response().error(f"Invalid umo: {e}").to_json()
platform_id = session.platform_name
platform_inst = next(
@@ -673,8 +671,7 @@ class OpenApiRoute(Route):
if not platform_inst:
return (
Response()
.error(f"Bot not found or not running for platform: {platform_id}")
.__dict__
.error(f"Bot not found or not running for platform: {platform_id}").to_json()
)
try:
@@ -682,12 +679,12 @@ class OpenApiRoute(Route):
message_payload
)
await platform_inst.send_by_session(session, message_chain)
return Response().ok().__dict__
return Response().ok().to_json()
except ValueError as e:
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
except Exception as e:
logger.error(f"Open API send_message failed: {e}", exc_info=True)
return Response().error(f"Failed to send message: {e}").__dict__
return Response().error(f"Failed to send message: {e}").to_json()
async def get_bots(self):
bot_ids = []
@@ -699,4 +696,4 @@ class OpenApiRoute(Route):
and platform_id not in bot_ids
):
bot_ids.append(platform_id)
return Response().ok(data={"bot_ids": bot_ids}).__dict__
return Response().ok(data={"bot_ids": bot_ids}).to_json()
+54 -65
View File
@@ -71,12 +71,11 @@ class PersonaRoute(Route):
}
for persona in personas
],
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"获取人格列表失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"获取人格列表失败: {e!s}").__dict__
return Response().error(f"获取人格列表失败: {e!s}").to_json()
async def get_persona_detail(self):
"""获取指定人格的详细信息"""
@@ -85,11 +84,11 @@ class PersonaRoute(Route):
persona_id = data.get("persona_id")
if not persona_id:
return Response().error("缺少必要参数: persona_id").__dict__
return Response().error("缺少必要参数: persona_id").to_json()
persona = await self.persona_mgr.get_persona(persona_id)
if not persona:
return Response().error("人格不存在").__dict__
return Response().error("人格不存在").to_json()
return (
Response()
@@ -110,12 +109,11 @@ class PersonaRoute(Route):
if persona.updated_at
else None,
},
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"获取人格详情失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"获取人格详情失败: {e!s}").__dict__
return Response().error(f"获取人格详情失败: {e!s}").to_json()
async def create_persona(self):
"""创建新人格"""
@@ -131,22 +129,21 @@ class PersonaRoute(Route):
sort_order = data.get("sort_order", 0)
if not persona_id:
return Response().error("人格ID不能为空").__dict__
return Response().error("人格ID不能为空").to_json()
if not system_prompt:
return Response().error("系统提示词不能为空").__dict__
return Response().error("系统提示词不能为空").to_json()
if custom_error_message is not None:
if not isinstance(custom_error_message, str):
return Response().error("自定义报错回复信息必须是字符串").__dict__
return Response().error("自定义报错回复信息必须是字符串").to_json()
custom_error_message = custom_error_message.strip() or None
# 验证 begin_dialogs 格式
if begin_dialogs and len(begin_dialogs) % 2 != 0:
return (
Response()
.error("预设对话数量必须为偶数(用户和助手轮流对话)")
.__dict__
.error("预设对话数量必须为偶数(用户和助手轮流对话)").to_json()
)
persona = await self.persona_mgr.create_persona(
@@ -182,14 +179,13 @@ class PersonaRoute(Route):
else None,
},
},
)
.__dict__
).to_json()
)
except ValueError as e:
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
except Exception as e:
logger.error(f"创建人格失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"创建人格失败: {e!s}").__dict__
return Response().error(f"创建人格失败: {e!s}").to_json()
async def update_persona(self):
"""更新人格信息"""
@@ -206,13 +202,13 @@ class PersonaRoute(Route):
custom_error_message = data.get("custom_error_message")
if not persona_id:
return Response().error("缺少必要参数: persona_id").__dict__
return Response().error("缺少必要参数: persona_id").to_json()
if has_custom_error_message:
if custom_error_message is not None and not isinstance(
custom_error_message, str
):
return Response().error("自定义报错回复信息必须是字符串").__dict__
return Response().error("自定义报错回复信息必须是字符串").to_json()
if isinstance(custom_error_message, str):
custom_error_message = custom_error_message.strip() or None
@@ -220,8 +216,7 @@ class PersonaRoute(Route):
if begin_dialogs is not None and len(begin_dialogs) % 2 != 0:
return (
Response()
.error("预设对话数量必须为偶数(用户和助手轮流对话)")
.__dict__
.error("预设对话数量必须为偶数(用户和助手轮流对话)").to_json()
)
update_kwargs = {
@@ -238,12 +233,12 @@ class PersonaRoute(Route):
await self.persona_mgr.update_persona(**update_kwargs)
return Response().ok({"message": "人格更新成功"}).__dict__
return Response().ok({"message": "人格更新成功"}).to_json()
except ValueError as e:
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
except Exception as e:
logger.error(f"更新人格失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"更新人格失败: {e!s}").__dict__
return Response().error(f"更新人格失败: {e!s}").to_json()
async def delete_persona(self):
"""删除人格"""
@@ -252,16 +247,16 @@ class PersonaRoute(Route):
persona_id = data.get("persona_id")
if not persona_id:
return Response().error("缺少必要参数: persona_id").__dict__
return Response().error("缺少必要参数: persona_id").to_json()
await self.persona_mgr.delete_persona(persona_id)
return Response().ok({"message": "人格删除成功"}).__dict__
return Response().ok({"message": "人格删除成功"}).to_json()
except ValueError as e:
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
except Exception as e:
logger.error(f"删除人格失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"删除人格失败: {e!s}").__dict__
return Response().error(f"删除人格失败: {e!s}").to_json()
async def clone_persona(self):
"""克隆人格"""
@@ -271,10 +266,10 @@ class PersonaRoute(Route):
new_persona_id = data.get("new_persona_id", "").strip()
if not source_persona_id:
return Response().error("缺少必要参数: source_persona_id").__dict__
return Response().error("缺少必要参数: source_persona_id").to_json()
if not new_persona_id:
return Response().error("新人格ID不能为空").__dict__
return Response().error("新人格ID不能为空").to_json()
persona = await self.persona_mgr.clone_persona(
source_persona_id=source_persona_id,
@@ -303,14 +298,13 @@ class PersonaRoute(Route):
else None,
},
},
)
.__dict__
).to_json()
)
except ValueError as e:
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
except Exception as e:
logger.error(f"克隆人格失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"克隆人格失败: {e!s}").__dict__
return Response().error(f"克隆人格失败: {e!s}").to_json()
async def move_persona(self):
"""移动人格到指定文件夹"""
@@ -320,16 +314,16 @@ class PersonaRoute(Route):
folder_id = data.get("folder_id") # None 表示移动到根目录
if not persona_id:
return Response().error("缺少必要参数: persona_id").__dict__
return Response().error("缺少必要参数: persona_id").to_json()
await self.persona_mgr.move_persona_to_folder(persona_id, folder_id)
return Response().ok({"message": "人格移动成功"}).__dict__
return Response().ok({"message": "人格移动成功"}).to_json()
except ValueError as e:
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
except Exception as e:
logger.error(f"移动人格失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"移动人格失败: {e!s}").__dict__
return Response().error(f"移动人格失败: {e!s}").to_json()
# ====
# Folder Routes
@@ -362,21 +356,20 @@ class PersonaRoute(Route):
}
for folder in folders
],
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"获取文件夹列表失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"获取文件夹列表失败: {e!s}").__dict__
return Response().error(f"获取文件夹列表失败: {e!s}").to_json()
async def get_folder_tree(self):
"""获取文件夹树形结构"""
try:
tree = await self.persona_mgr.get_folder_tree()
return Response().ok(tree).__dict__
return Response().ok(tree).to_json()
except Exception as e:
logger.error(f"获取文件夹树失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"获取文件夹树失败: {e!s}").__dict__
return Response().error(f"获取文件夹树失败: {e!s}").to_json()
async def get_folder_detail(self):
"""获取指定文件夹的详细信息"""
@@ -385,11 +378,11 @@ class PersonaRoute(Route):
folder_id = data.get("folder_id")
if not folder_id:
return Response().error("缺少必要参数: folder_id").__dict__
return Response().error("缺少必要参数: folder_id").to_json()
folder = await self.persona_mgr.get_folder(folder_id)
if not folder:
return Response().error("文件夹不存在").__dict__
return Response().error("文件夹不存在").to_json()
return (
Response()
@@ -407,12 +400,11 @@ class PersonaRoute(Route):
if folder.updated_at
else None,
},
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"获取文件夹详情失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"获取文件夹详情失败: {e!s}").__dict__
return Response().error(f"获取文件夹详情失败: {e!s}").to_json()
async def create_folder(self):
"""创建文件夹"""
@@ -424,7 +416,7 @@ class PersonaRoute(Route):
sort_order = data.get("sort_order", 0)
if not name:
return Response().error("文件夹名称不能为空").__dict__
return Response().error("文件夹名称不能为空").to_json()
folder = await self.persona_mgr.create_folder(
name=name,
@@ -452,12 +444,11 @@ class PersonaRoute(Route):
else None,
},
},
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"创建文件夹失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"创建文件夹失败: {e!s}").__dict__
return Response().error(f"创建文件夹失败: {e!s}").to_json()
async def update_folder(self):
"""更新文件夹信息"""
@@ -470,7 +461,7 @@ class PersonaRoute(Route):
sort_order = data.get("sort_order")
if not folder_id:
return Response().error("缺少必要参数: folder_id").__dict__
return Response().error("缺少必要参数: folder_id").to_json()
await self.persona_mgr.update_folder(
folder_id=folder_id,
@@ -480,10 +471,10 @@ class PersonaRoute(Route):
sort_order=sort_order,
)
return Response().ok({"message": "文件夹更新成功"}).__dict__
return Response().ok({"message": "文件夹更新成功"}).to_json()
except Exception as e:
logger.error(f"更新文件夹失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"更新文件夹失败: {e!s}").__dict__
return Response().error(f"更新文件夹失败: {e!s}").to_json()
async def delete_folder(self):
"""删除文件夹"""
@@ -492,14 +483,14 @@ class PersonaRoute(Route):
folder_id = data.get("folder_id")
if not folder_id:
return Response().error("缺少必要参数: folder_id").__dict__
return Response().error("缺少必要参数: folder_id").to_json()
await self.persona_mgr.delete_folder(folder_id)
return Response().ok({"message": "文件夹删除成功"}).__dict__
return Response().ok({"message": "文件夹删除成功"}).to_json()
except Exception as e:
logger.error(f"删除文件夹失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"删除文件夹失败: {e!s}").__dict__
return Response().error(f"删除文件夹失败: {e!s}").to_json()
async def reorder_items(self):
"""批量更新排序顺序
@@ -519,26 +510,24 @@ class PersonaRoute(Route):
items = data.get("items", [])
if not items:
return Response().error("items 不能为空").__dict__
return Response().error("items 不能为空").to_json()
# 验证每个 item 的格式
for item in items:
if not all(k in item for k in ("id", "type", "sort_order")):
return (
Response()
.error("每个 item 必须包含 id, type, sort_order 字段")
.__dict__
.error("每个 item 必须包含 id, type, sort_order 字段").to_json()
)
if item["type"] not in ("persona", "folder"):
return (
Response()
.error("type 字段必须是 'persona' 或 'folder'")
.__dict__
.error("type 字段必须是 'persona' 或 'folder'").to_json()
)
await self.persona_mgr.batch_update_sort_order(items)
return Response().ok({"message": "排序更新成功"}).__dict__
return Response().ok({"message": "排序更新成功"}).to_json()
except Exception as e:
logger.error(f"更新排序失败: {e!s}\n{traceback.format_exc()}")
return Response().error(f"更新排序失败: {e!s}").__dict__
return Response().error(f"更新排序失败: {e!s}").to_json()
+5 -5
View File
@@ -56,7 +56,7 @@ class PlatformRoute(Route):
if not platform_adapter:
logger.warning(f"未找到 webhook_uuid 为 {webhook_uuid} 的平台")
return Response().error("未找到对应平台").__dict__, 404
return Response().error("未找到对应平台").to_json(), 404
# 调用平台适配器的 webhook_callback 方法
try:
@@ -66,10 +66,10 @@ class PlatformRoute(Route):
logger.error(
f"平台 {platform_adapter.meta().name} 未实现 webhook_callback 方法"
)
return Response().error("平台未支持统一 Webhook 模式").__dict__, 500
return Response().error("平台未支持统一 Webhook 模式").to_json(), 500
except Exception as e:
logger.error(f"处理 webhook 回调时发生错误: {e}", exc_info=True)
return Response().error("处理回调失败").__dict__, 500
return Response().error("处理回调失败").to_json(), 500
def _find_platform_by_uuid(self, webhook_uuid: str) -> Platform | None:
"""根据 webhook_uuid 查找对应的平台适配器
@@ -94,7 +94,7 @@ class PlatformRoute(Route):
"""
try:
stats = self.platform_manager.get_all_stats()
return Response().ok(stats).__dict__
return Response().ok(stats).to_json()
except Exception as e:
logger.error(f"获取平台统计信息失败: {e}", exc_info=True)
return Response().error(f"获取统计信息失败: {e}").__dict__, 500
return Response().error(f"获取统计信息失败: {e}").to_json(), 500
+60 -74
View File
@@ -114,45 +114,42 @@ class PluginRoute(Route):
"message": message,
"astrbot_version": version_spec,
}
)
.__dict__
).to_json()
)
except Exception as e:
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def reload_failed_plugins(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
try:
data = await request.get_json()
dir_name = data.get("dir_name") # 这里拿的是目录名,不是插件名
if not dir_name:
return Response().error("缺少插件目录名").__dict__
return Response().error("缺少插件目录名").to_json()
# 调用 star_manager.py 中的函数
# 注意:传入的是目录名
success, err = await self.plugin_manager.reload_failed_plugin(dir_name)
if success:
return Response().ok(None, f"插件 {dir_name} 重载成功。").__dict__
return Response().ok(None, f"插件 {dir_name} 重载成功。").to_json()
else:
return Response().error(f"重载失败: {err}").__dict__
return Response().error(f"重载失败: {err}").to_json()
except Exception as e:
logger.error(f"/api/plugin/reload-failed: {traceback.format_exc()}")
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def reload_plugins(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
data = await request.get_json()
@@ -160,11 +157,11 @@ class PluginRoute(Route):
try:
success, message = await self.plugin_manager.reload(plugin_name)
if not success:
return Response().error(message or "插件重载失败").__dict__
return Response().ok(None, "重载成功。").__dict__
return Response().error(message or "插件重载失败").to_json()
return Response().ok(None, "重载成功。").to_json()
except Exception as e:
logger.error(f"/api/plugin/reload: {traceback.format_exc()}")
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def get_online_plugins(self):
custom = request.args.get("custom_registry")
@@ -181,7 +178,7 @@ class PluginRoute(Route):
cached_data = await self._load_plugin_cache(source.cache_file)
if cached_data:
logger.debug("缓存MD5匹配,使用缓存的插件市场数据")
return Response().ok(cached_data).__dict__
return Response().ok(cached_data).to_json()
# 尝试获取远程数据
remote_data = None
@@ -221,7 +218,7 @@ class PluginRoute(Route):
remote_data,
current_md5,
)
return Response().ok(remote_data).__dict__
return Response().ok(remote_data).to_json()
logger.error(f"请求 {url} 失败,状态码:{response.status}")
except Exception as e:
logger.error(f"请求 {url} 失败,错误:{e}")
@@ -232,9 +229,9 @@ class PluginRoute(Route):
if cached_data:
logger.warning("远程插件市场数据获取失败,使用缓存数据")
return Response().ok(cached_data, "使用缓存数据,可能不是最新版本").__dict__
return Response().ok(cached_data, "使用缓存数据,可能不是最新版本").to_json()
return Response().error("获取插件列表失败,且没有可用的缓存数据").__dict__
return Response().error("获取插件列表失败,且没有可用的缓存数据").to_json()
def _build_registry_source(self, custom_url: str | None) -> RegistrySource:
"""构建注册表源信息"""
@@ -440,13 +437,12 @@ class PluginRoute(Route):
_plugin_resp.append(_t)
return (
Response()
.ok(_plugin_resp, message=self.plugin_manager.failed_plugin_info)
.__dict__
.ok(_plugin_resp, message=self.plugin_manager.failed_plugin_info).to_json()
)
async def get_failed_plugins(self):
"""专门获取加载失败的插件列表(字典格式)"""
return Response().ok(self.plugin_manager.failed_plugin_dict).__dict__
return Response().ok(self.plugin_manager.failed_plugin_dict).to_json()
async def get_plugin_handlers_info(self, handler_full_names: list[str]):
"""解析插件行为"""
@@ -513,8 +509,7 @@ class PluginRoute(Route):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
post_data = await request.get_json()
@@ -534,7 +529,7 @@ class PluginRoute(Route):
)
# self.core_lifecycle.restart()
logger.info(f"安装插件 {repo_url} 成功。")
return Response().ok(plugin_info, "安装成功。").__dict__
return Response().ok(plugin_info, "安装成功。").to_json()
except PluginVersionIncompatibleError as e:
return {
"status": "warning",
@@ -546,14 +541,13 @@ class PluginRoute(Route):
}
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def install_plugin_upload(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
try:
@@ -575,7 +569,7 @@ class PluginRoute(Route):
)
# self.core_lifecycle.restart()
logger.info(f"安装插件 {file.filename} 成功")
return Response().ok(plugin_info, "安装成功。").__dict__
return Response().ok(plugin_info, "安装成功。").to_json()
except PluginVersionIncompatibleError as e:
return {
"status": "warning",
@@ -587,14 +581,13 @@ class PluginRoute(Route):
}
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def uninstall_plugin(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
post_data = await request.get_json()
@@ -609,17 +602,16 @@ class PluginRoute(Route):
delete_data=delete_data,
)
logger.info(f"卸载插件 {plugin_name} 成功")
return Response().ok(None, "卸载成功").__dict__
return Response().ok(None, "卸载成功").to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def uninstall_failed_plugin(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
post_data = await request.get_json()
@@ -627,7 +619,7 @@ class PluginRoute(Route):
delete_config = post_data.get("delete_config", False)
delete_data = post_data.get("delete_data", False)
if not dir_name:
return Response().error("缺少失败插件目录名").__dict__
return Response().error("缺少失败插件目录名").to_json()
try:
logger.info(f"正在卸载失败插件 {dir_name}")
@@ -637,17 +629,16 @@ class PluginRoute(Route):
delete_data=delete_data,
)
logger.info(f"卸载失败插件 {dir_name} 成功")
return Response().ok(None, "卸载成功").__dict__
return Response().ok(None, "卸载成功").to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def update_plugin(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
post_data = await request.get_json()
@@ -659,17 +650,16 @@ class PluginRoute(Route):
# self.core_lifecycle.restart()
await self.plugin_manager.reload(plugin_name)
logger.info(f"更新插件 {plugin_name} 成功。")
return Response().ok(None, "更新成功。").__dict__
return Response().ok(None, "更新成功。").to_json()
except Exception as e:
logger.error(f"/api/plugin/update: {traceback.format_exc()}")
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def update_all_plugins(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
post_data = await request.get_json()
@@ -677,7 +667,7 @@ class PluginRoute(Route):
proxy: str = post_data.get("proxy", "")
if not isinstance(plugin_names, list) or not plugin_names:
return Response().error("插件列表不能为空").__dict__
return Response().error("插件列表不能为空").to_json()
results = []
sem = asyncio.Semaphore(PLUGIN_UPDATE_CONCURRENCY)
@@ -715,14 +705,13 @@ class PluginRoute(Route):
else f"批量更新完成,其中 {len(failed)}/{len(results)} 个插件失败。"
)
return Response().ok({"results": results}, message).__dict__
return Response().ok({"results": results}, message).to_json()
async def off_plugin(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
post_data = await request.get_json()
@@ -730,17 +719,16 @@ class PluginRoute(Route):
try:
await self.plugin_manager.turn_off_plugin(plugin_name)
logger.info(f"停用插件 {plugin_name} 。")
return Response().ok(None, "停用成功。").__dict__
return Response().ok(None, "停用成功。").to_json()
except Exception as e:
logger.error(f"/api/plugin/off: {traceback.format_exc()}")
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def on_plugin(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
post_data = await request.get_json()
@@ -748,10 +736,10 @@ class PluginRoute(Route):
try:
await self.plugin_manager.turn_on_plugin(plugin_name)
logger.info(f"启用插件 {plugin_name} 。")
return Response().ok(None, "启用成功。").__dict__
return Response().ok(None, "启用成功。").to_json()
except Exception as e:
logger.error(f"/api/plugin/on: {traceback.format_exc()}")
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def get_plugin_readme(self):
plugin_name = request.args.get("name")
@@ -759,7 +747,7 @@ class PluginRoute(Route):
if not plugin_name:
logger.warning("插件名称为空")
return Response().error("插件名称不能为空").__dict__
return Response().error("插件名称不能为空").to_json()
plugin_obj = None
for plugin in self.plugin_manager.context.get_all_stars():
@@ -769,11 +757,11 @@ class PluginRoute(Route):
if not plugin_obj:
logger.warning(f"插件 {plugin_name} 不存在")
return Response().error(f"插件 {plugin_name} 不存在").__dict__
return Response().error(f"插件 {plugin_name} 不存在").to_json()
if not plugin_obj.root_dir_name:
logger.warning(f"插件 {plugin_name} 目录不存在")
return Response().error(f"插件 {plugin_name} 目录不存在").__dict__
return Response().error(f"插件 {plugin_name} 目录不存在").to_json()
if plugin_obj.reserved:
plugin_dir = os.path.join(
@@ -788,13 +776,13 @@ class PluginRoute(Route):
if not await anyio.Path(plugin_dir).is_dir():
logger.warning(f"无法找到插件目录: {plugin_dir}")
return Response().error(f"无法找到插件 {plugin_name} 的目录").__dict__
return Response().error(f"无法找到插件 {plugin_name} 的目录").to_json()
readme_path = os.path.join(plugin_dir, "README.md")
if not await anyio.Path(readme_path).is_file():
logger.warning(f"插件 {plugin_name} 没有README文件")
return Response().error(f"插件 {plugin_name} 没有README文件").__dict__
return Response().error(f"插件 {plugin_name} 没有README文件").to_json()
try:
async with await anyio.open_file(readme_path, encoding="utf-8") as f:
@@ -802,12 +790,11 @@ class PluginRoute(Route):
return (
Response()
.ok({"content": readme_content}, "成功获取README内容")
.__dict__
.ok({"content": readme_content}, "成功获取README内容").to_json()
)
except Exception as e:
logger.error(f"/api/plugin/readme: {traceback.format_exc()}")
return Response().error(f"读取README文件失败: {e!s}").__dict__
return Response().error(f"读取README文件失败: {e!s}").to_json()
async def get_plugin_changelog(self):
"""获取插件更新日志
@@ -819,7 +806,7 @@ class PluginRoute(Route):
if not plugin_name:
logger.warning("插件名称为空")
return Response().error("插件名称不能为空").__dict__
return Response().error("插件名称不能为空").to_json()
# 查找插件
plugin_obj = None
@@ -830,11 +817,11 @@ class PluginRoute(Route):
if not plugin_obj:
logger.warning(f"插件 {plugin_name} 不存在")
return Response().error(f"插件 {plugin_name} 不存在").__dict__
return Response().error(f"插件 {plugin_name} 不存在").to_json()
if not plugin_obj.root_dir_name:
logger.warning(f"插件 {plugin_name} 目录不存在")
return Response().error(f"插件 {plugin_name} 目录不存在").__dict__
return Response().error(f"插件 {plugin_name} 目录不存在").to_json()
if plugin_obj.reserved:
plugin_dir = os.path.join(
@@ -849,7 +836,7 @@ class PluginRoute(Route):
if not await anyio.Path(plugin_dir).is_dir():
logger.warning(f"无法找到插件目录: {plugin_dir}")
return Response().error(f"无法找到插件 {plugin_name} 的目录").__dict__
return Response().error(f"无法找到插件 {plugin_name} 的目录").to_json()
# 尝试多种可能的文件名
changelog_names = ["CHANGELOG.md", "changelog.md", "CHANGELOG", "changelog"]
@@ -863,21 +850,20 @@ class PluginRoute(Route):
changelog_content = await f.read()
return (
Response()
.ok({"content": changelog_content}, "成功获取更新日志")
.__dict__
.ok({"content": changelog_content}, "成功获取更新日志").to_json()
)
except Exception as e:
logger.error(f"/api/plugin/changelog: {traceback.format_exc()}")
return Response().error(f"读取更新日志失败: {e!s}").__dict__
return Response().error(f"读取更新日志失败: {e!s}").to_json()
# 没有找到 changelog 文件,返回 ok 但 content 为 null
logger.warning(f"插件 {plugin_name} 没有更新日志文件")
return Response().ok({"content": None}, "该插件没有更新日志文件").__dict__
return Response().ok({"content": None}, "该插件没有更新日志文件").to_json()
async def get_custom_source(self):
"""获取自定义插件源"""
sources = await sp.global_get("custom_plugin_sources", [])
return Response().ok(sources).__dict__
return Response().ok(sources).to_json()
async def save_custom_source(self):
"""保存自定义插件源"""
@@ -885,10 +871,10 @@ class PluginRoute(Route):
data = await request.get_json()
sources = data.get("sources", [])
if not isinstance(sources, list):
return Response().error("sources fields must be a list").__dict__
return Response().error("sources fields must be a list").to_json()
await sp.global_put("custom_plugin_sources", sources)
return Response().ok(None, "保存成功").__dict__
return Response().ok(None, "保存成功").to_json()
except Exception as e:
logger.error(f"/api/plugin/source/save: {traceback.format_exc()}")
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
+1 -1
View File
@@ -59,7 +59,7 @@ def runtime_status_response(
core_lifecycle,
include_failure_details=include_failure_details,
),
).__dict__
).to_json()
)
response.status_code = status_code
return response
+50 -62
View File
@@ -241,12 +241,11 @@ class SessionManagementRoute(Route):
"available_kbs": available_kbs,
"available_rule_keys": AVAILABLE_SESSION_RULE_KEYS,
}
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"获取规则列表失败: {e!s}")
return Response().error(f"获取规则列表失败: {e!s}").__dict__
return Response().error(f"获取规则列表失败: {e!s}").to_json()
async def update_session_rule(self):
"""更新某个 umo 的自定义规则
@@ -265,11 +264,11 @@ class SessionManagementRoute(Route):
rule_value = data.get("rule_value")
if not umo:
return Response().error("缺少必要参数: umo").__dict__
return Response().error("缺少必要参数: umo").to_json()
if not rule_key:
return Response().error("缺少必要参数: rule_key").__dict__
return Response().error("缺少必要参数: rule_key").to_json()
if rule_key not in AVAILABLE_SESSION_RULE_KEYS:
return Response().error(f"不支持的规则键: {rule_key}").__dict__
return Response().error(f"不支持的规则键: {rule_key}").to_json()
if rule_key == "session_plugin_config":
rule_value = {
@@ -281,12 +280,11 @@ class SessionManagementRoute(Route):
return (
Response()
.ok({"message": f"规则 {rule_key} 已更新", "umo": umo})
.__dict__
.ok({"message": f"规则 {rule_key} 已更新", "umo": umo}).to_json()
)
except Exception as e:
logger.error(f"更新会话规则失败: {e!s}")
return Response().error(f"更新会话规则失败: {e!s}").__dict__
return Response().error(f"更新会话规则失败: {e!s}").to_json()
async def delete_session_rule(self):
"""删除某个 umo 的自定义规则
@@ -303,25 +301,24 @@ class SessionManagementRoute(Route):
rule_key = data.get("rule_key")
if not umo:
return Response().error("缺少必要参数: umo").__dict__
return Response().error("缺少必要参数: umo").to_json()
if rule_key:
# 删除单个规则
if rule_key not in AVAILABLE_SESSION_RULE_KEYS:
return Response().error(f"不支持的规则键: {rule_key}").__dict__
return Response().error(f"不支持的规则键: {rule_key}").to_json()
await sp.session_remove(umo, rule_key)
return (
Response()
.ok({"message": f"规则 {rule_key} 已删除", "umo": umo})
.__dict__
.ok({"message": f"规则 {rule_key} 已删除", "umo": umo}).to_json()
)
else:
# 删除该 umo 的所有规则
await sp.clear_async("umo", umo)
return Response().ok({"message": "所有规则已删除", "umo": umo}).__dict__
return Response().ok({"message": "所有规则已删除", "umo": umo}).to_json()
except Exception as e:
logger.error(f"删除会话规则失败: {e!s}")
return Response().error(f"删除会话规则失败: {e!s}").__dict__
return Response().error(f"删除会话规则失败: {e!s}").to_json()
async def batch_delete_session_rule(self):
"""批量删除多个 umo 的自定义规则
@@ -347,10 +344,10 @@ class SessionManagementRoute(Route):
# 如果是自定义分组
if scope == "custom_group":
if not group_id:
return Response().error("请指定分组 ID").__dict__
return Response().error("请指定分组 ID").to_json()
groups = self._get_groups()
if group_id not in groups:
return Response().error(f"分组 '{group_id}' 不存在").__dict__
return Response().error(f"分组 '{group_id}' 不存在").to_json()
umos = groups[group_id].get("umos", [])
else:
async with self.db_helper.get_db() as session:
@@ -376,13 +373,13 @@ class SessionManagementRoute(Route):
umos = all_umos
if not umos:
return Response().error("缺少必要参数: umos 或有效的 scope").__dict__
return Response().error("缺少必要参数: umos 或有效的 scope").to_json()
if not isinstance(umos, list):
return Response().error("参数 umos 必须是数组").__dict__
return Response().error("参数 umos 必须是数组").to_json()
if rule_key and rule_key not in AVAILABLE_SESSION_RULE_KEYS:
return Response().error(f"不支持的规则键: {rule_key}").__dict__
return Response().error(f"不支持的规则键: {rule_key}").to_json()
# 批量删除
success_count = 0
@@ -411,8 +408,7 @@ class SessionManagementRoute(Route):
"success_count": success_count,
"failed_umos": failed_umos,
}
)
.__dict__
).to_json()
)
else:
return (
@@ -422,12 +418,11 @@ class SessionManagementRoute(Route):
"message": message,
"success_count": success_count,
}
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"批量删除会话规则失败: {e!s}")
return Response().error(f"批量删除会话规则失败: {e!s}").__dict__
return Response().error(f"批量删除会话规则失败: {e!s}").to_json()
async def list_umos(self):
"""列出所有有对话记录的 umo,从 Conversations 表中找
@@ -445,10 +440,10 @@ class SessionManagementRoute(Route):
)
umos = [row[0] for row in result.fetchall()]
return Response().ok({"umos": umos}).__dict__
return Response().ok({"umos": umos}).to_json()
except Exception as e:
logger.error(f"获取 UMO 列表失败: {e!s}")
return Response().error(f"获取 UMO 列表失败: {e!s}").__dict__
return Response().error(f"获取 UMO 列表失败: {e!s}").to_json()
async def list_all_umos_with_status(self):
"""获取所有有对话记录的 UMO 及其服务状态(支持分页、搜索、筛选)
@@ -598,12 +593,11 @@ class SessionManagementRoute(Route):
"available_tts_providers": available_tts_providers,
"available_stt_providers": available_stt_providers,
}
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"获取会话状态列表失败: {e!s}")
return Response().error(f"获取会话状态列表失败: {e!s}").__dict__
return Response().error(f"获取会话状态列表失败: {e!s}").to_json()
async def batch_update_service(self):
"""批量更新多个 UMO 的服务状态 (LLM/TTS/Session)
@@ -629,17 +623,17 @@ class SessionManagementRoute(Route):
# 如果没有任何修改
if llm_enabled is None and tts_enabled is None and session_enabled is None:
return Response().error("至少需要指定一个要修改的状态").__dict__
return Response().error("至少需要指定一个要修改的状态").to_json()
# 如果指定了 scope,获取符合条件的所有 umo
if scope and not umos:
# 如果是自定义分组
if scope == "custom_group":
if not group_id:
return Response().error("请指定分组 ID").__dict__
return Response().error("请指定分组 ID").to_json()
groups = self._get_groups()
if group_id not in groups:
return Response().error(f"分组 '{group_id}' 不存在").__dict__
return Response().error(f"分组 '{group_id}' 不存在").to_json()
umos = groups[group_id].get("umos", [])
else:
async with self.db_helper.get_db() as session:
@@ -665,7 +659,7 @@ class SessionManagementRoute(Route):
umos = all_umos
if not umos:
return Response().error("没有找到符合条件的会话").__dict__
return Response().error("没有找到符合条件的会话").to_json()
# 批量更新
success_count = 0
@@ -716,12 +710,11 @@ class SessionManagementRoute(Route):
"failed_count": len(failed_umos),
"failed_umos": failed_umos,
}
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"批量更新服务状态失败: {e!s}")
return Response().error(f"批量更新服务状态失败: {e!s}").__dict__
return Response().error(f"批量更新服务状态失败: {e!s}").to_json()
async def batch_update_provider(self):
"""批量更新多个 UMO 的 Provider 配置
@@ -744,8 +737,7 @@ class SessionManagementRoute(Route):
if not provider_type or not provider_id:
return (
Response()
.error("缺少必要参数: provider_type, provider_id")
.__dict__
.error("缺少必要参数: provider_type, provider_id").to_json()
)
# 转换 provider_type
@@ -757,8 +749,7 @@ class SessionManagementRoute(Route):
if provider_type not in provider_type_map:
return (
Response()
.error(f"不支持的 provider_type: {provider_type}")
.__dict__
.error(f"不支持的 provider_type: {provider_type}").to_json()
)
provider_type_enum = provider_type_map[provider_type]
@@ -769,10 +760,10 @@ class SessionManagementRoute(Route):
# 如果是自定义分组
if scope == "custom_group":
if not group_id:
return Response().error("请指定分组 ID").__dict__
return Response().error("请指定分组 ID").to_json()
groups = self._get_groups()
if group_id not in groups:
return Response().error(f"分组 '{group_id}' 不存在").__dict__
return Response().error(f"分组 '{group_id}' 不存在").to_json()
umos = groups[group_id].get("umos", [])
else:
async with self.db_helper.get_db() as session:
@@ -798,7 +789,7 @@ class SessionManagementRoute(Route):
umos = all_umos
if not umos:
return Response().error("没有找到符合条件的会话").__dict__
return Response().error("没有找到符合条件的会话").to_json()
# 批量更新
success_count = 0
@@ -826,12 +817,11 @@ class SessionManagementRoute(Route):
"failed_count": len(failed_umos),
"failed_umos": failed_umos,
}
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"批量更新 Provider 失败: {e!s}")
return Response().error(f"批量更新 Provider 失败: {e!s}").__dict__
return Response().error(f"批量更新 Provider 失败: {e!s}").to_json()
# ==================== 分组管理 API ====================
@@ -858,10 +848,10 @@ class SessionManagementRoute(Route):
"umo_count": len(group_data.get("umos", [])),
}
)
return Response().ok({"groups": groups_list}).__dict__
return Response().ok({"groups": groups_list}).to_json()
except Exception as e:
logger.error(f"获取分组列表失败: {e!s}")
return Response().error(f"获取分组列表失败: {e!s}").__dict__
return Response().error(f"获取分组列表失败: {e!s}").to_json()
async def create_group(self):
"""创建新分组"""
@@ -871,7 +861,7 @@ class SessionManagementRoute(Route):
umos = data.get("umos", [])
if not name:
return Response().error("分组名称不能为空").__dict__
return Response().error("分组名称不能为空").to_json()
groups = self._get_groups()
@@ -899,12 +889,11 @@ class SessionManagementRoute(Route):
"umo_count": len(umos),
},
}
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"创建分组失败: {e!s}")
return Response().error(f"创建分组失败: {e!s}").__dict__
return Response().error(f"创建分组失败: {e!s}").to_json()
async def update_group(self):
"""更新分组(改名、增删成员)"""
@@ -917,12 +906,12 @@ class SessionManagementRoute(Route):
remove_umos = data.get("remove_umos", [])
if not group_id:
return Response().error("分组 ID 不能为空").__dict__
return Response().error("分组 ID 不能为空").to_json()
groups = self._get_groups()
if group_id not in groups:
return Response().error(f"分组 '{group_id}' 不存在").__dict__
return Response().error(f"分组 '{group_id}' 不存在").to_json()
group = groups[group_id]
@@ -956,12 +945,11 @@ class SessionManagementRoute(Route):
"umo_count": len(group["umos"]),
},
}
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(f"更新分组失败: {e!s}")
return Response().error(f"更新分组失败: {e!s}").__dict__
return Response().error(f"更新分组失败: {e!s}").to_json()
async def delete_group(self):
"""删除分组"""
@@ -970,19 +958,19 @@ class SessionManagementRoute(Route):
group_id = data.get("id")
if not group_id:
return Response().error("分组 ID 不能为空").__dict__
return Response().error("分组 ID 不能为空").to_json()
groups = self._get_groups()
if group_id not in groups:
return Response().error(f"分组 '{group_id}' 不存在").__dict__
return Response().error(f"分组 '{group_id}' 不存在").to_json()
group_name = groups[group_id].get("name", group_id)
del groups[group_id]
self._save_groups(groups)
return Response().ok({"message": f"分组 '{group_name}' 已删除"}).__dict__
return Response().ok({"message": f"分组 '{group_name}' 已删除"}).to_json()
except Exception as e:
logger.error(f"删除分组失败: {e!s}")
return Response().error(f"删除分组失败: {e!s}").__dict__
return Response().error(f"删除分组失败: {e!s}").to_json()
+55 -71
View File
@@ -3,6 +3,7 @@ import re
import shutil
import traceback
from collections.abc import Awaitable, Callable
from dataclasses import asdict
from pathlib import Path
from typing import Any
@@ -125,10 +126,10 @@ class SkillsRoute(Route):
except ValueError as e:
# Config not ready — expected when Neo isn't set up yet
logger.debug("[Neo] %s", e)
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def get_skills(self):
try:
@@ -144,23 +145,21 @@ class SkillsRoute(Route):
Response()
.ok(
{
"skills": [skill.__dict__ for skill in skills],
"skills": [asdict(skill) for skill in skills],
"runtime": runtime,
"sandbox_cache": skill_mgr.get_sandbox_skills_cache_status(),
}
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def upload_skill(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
temp_path = None
@@ -168,10 +167,10 @@ class SkillsRoute(Route):
files = await request.files
file = files.get("file")
if not file:
return Response().error("Missing file").__dict__
return Response().error("Missing file").to_json()
filename = os.path.basename(file.filename or "skill.zip")
if not filename.lower().endswith(".zip"):
return Response().error("Only .zip files are supported").__dict__
return Response().error("Only .zip files are supported").to_json()
temp_dir = get_astrbot_temp_path()
os.makedirs(temp_dir, exist_ok=True)
@@ -201,12 +200,11 @@ class SkillsRoute(Route):
return (
Response()
.ok({"name": skill_name}, "Skill uploaded successfully.")
.__dict__
.ok({"name": skill_name}, "Skill uploaded successfully.").to_json()
)
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
finally:
if temp_path and await anyio.Path(temp_path).exists():
try:
@@ -219,8 +217,7 @@ class SkillsRoute(Route):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
try:
@@ -228,7 +225,7 @@ class SkillsRoute(Route):
file_list = files.getlist("files")
if not file_list:
return Response().error("No files provided").__dict__
return Response().error("No files provided").to_json()
succeeded = []
failed = []
@@ -323,8 +320,7 @@ class SkillsRoute(Route):
"skipped": skipped,
},
message,
)
.__dict__
).to_json()
)
if failed_count == 0 and success_count == 0:
message = f"All {total} file(s) were skipped."
@@ -338,8 +334,7 @@ class SkillsRoute(Route):
"skipped": skipped,
},
message,
)
.__dict__
).to_json()
)
if success_count == 0 and skipped_count == 0:
message = f"Upload failed for all {total} file(s)."
@@ -350,7 +345,7 @@ class SkillsRoute(Route):
"failed": failed,
"skipped": skipped,
}
return resp.__dict__
return resp.to_json()
message = f"Partial success: {success_count}/{total} skill(s) uploaded."
return (
@@ -363,21 +358,20 @@ class SkillsRoute(Route):
"skipped": skipped,
},
message,
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def download_skill(self):
try:
name = str(request.args.get("name") or "").strip()
if not name:
return Response().error("Missing skill name").__dict__
return Response().error("Missing skill name").to_json()
if not _SKILL_NAME_RE.match(name):
return Response().error("Invalid skill name").__dict__
return Response().error("Invalid skill name").to_json()
skill_mgr = SkillManager()
if skill_mgr.is_sandbox_only_skill(name):
@@ -385,14 +379,13 @@ class SkillsRoute(Route):
Response()
.error(
"Sandbox preset skill cannot be downloaded from local skill files."
)
.__dict__
).to_json()
)
skill_dir = Path(skill_mgr.skills_root) / name
skill_md = skill_dir / "SKILL.md"
if not skill_dir.is_dir() or not skill_md.exists():
return Response().error("Local skill not found").__dict__
return Response().error("Local skill not found").to_json()
export_dir = Path(get_astrbot_temp_path()) / "skill_exports"
export_dir.mkdir(parents=True, exist_ok=True)
@@ -416,48 +409,46 @@ class SkillsRoute(Route):
)
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def update_skill(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
try:
data = await request.get_json()
name = data.get("name")
active = data.get("active", True)
if not name:
return Response().error("Missing skill name").__dict__
return Response().error("Missing skill name").to_json()
SkillManager().set_skill_active(name, bool(active))
return Response().ok({"name": name, "active": bool(active)}).__dict__
return Response().ok({"name": name, "active": bool(active)}).to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def delete_skill(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
try:
data = await request.get_json()
name = data.get("name")
if not name:
return Response().error("Missing skill name").__dict__
return Response().error("Missing skill name").to_json()
SkillManager().delete_skill(name)
try:
await sync_skills_to_active_sandboxes()
except Exception:
logger.warning("Failed to sync deleted skills to active sandboxes.")
return Response().ok({"name": name}).__dict__
return Response().ok({"name": name}).to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
async def get_neo_candidates(self):
logger.info("[Neo] GET /skills/neo/candidates requested.")
@@ -476,7 +467,7 @@ class SkillsRoute(Route):
result = _to_jsonable(candidates)
total = result.get("total", "?") if isinstance(result, dict) else "?"
logger.info(f"[Neo] Candidates fetched: total={total}")
return Response().ok(result).__dict__
return Response().ok(result).to_json()
return await self._with_neo_client(_do)
@@ -499,7 +490,7 @@ class SkillsRoute(Route):
result = _to_jsonable(releases)
total = result.get("total", "?") if isinstance(result, dict) else "?"
logger.info(f"[Neo] Releases fetched: total={total}")
return Response().ok(result).__dict__
return Response().ok(result).to_json()
return await self._with_neo_client(_do)
@@ -507,12 +498,12 @@ class SkillsRoute(Route):
logger.info("[Neo] GET /skills/neo/payload requested.")
payload_ref = request.args.get("payload_ref", "")
if not payload_ref:
return Response().error("Missing payload_ref").__dict__
return Response().error("Missing payload_ref").to_json()
async def _do(client):
payload = await client.skills.get_payload(payload_ref)
logger.info(f"[Neo] Payload fetched: ref={payload_ref}")
return Response().ok(_to_jsonable(payload)).__dict__
return Response().ok(_to_jsonable(payload)).to_json()
return await self._with_neo_client(_do)
@@ -520,15 +511,14 @@ class SkillsRoute(Route):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
logger.info("[Neo] POST /skills/neo/evaluate requested.")
data = await request.get_json()
candidate_id = data.get("candidate_id")
passed_value = data.get("passed")
if not candidate_id or passed_value is None:
return Response().error("Missing candidate_id or passed").__dict__
return Response().error("Missing candidate_id or passed").to_json()
passed = _to_bool(passed_value, False)
async def _do(client):
@@ -542,7 +532,7 @@ class SkillsRoute(Route):
logger.info(
f"[Neo] Candidate evaluated: id={candidate_id}, passed={passed}"
)
return Response().ok(_to_jsonable(result)).__dict__
return Response().ok(_to_jsonable(result)).to_json()
return await self._with_neo_client(_do)
@@ -550,8 +540,7 @@ class SkillsRoute(Route):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
logger.info("[Neo] POST /skills/neo/promote requested.")
data = await request.get_json()
@@ -559,9 +548,9 @@ class SkillsRoute(Route):
stage = data.get("stage", "canary")
sync_to_local = _to_bool(data.get("sync_to_local"), True)
if not candidate_id:
return Response().error("Missing candidate_id").__dict__
return Response().error("Missing candidate_id").to_json()
if stage not in {"canary", "stable"}:
return Response().error("Invalid stage, must be canary/stable").__dict__
return Response().error("Invalid stage, must be canary/stable").to_json()
async def _do(client):
sync_mgr = NeoSkillSyncManager()
@@ -590,7 +579,7 @@ class SkillsRoute(Route):
"release": release_json,
"rollback": result.get("rollback"),
}
return resp.__dict__
return resp.to_json()
# Try to push latest local skills to all active sandboxes.
if not did_sync_to_local:
@@ -599,7 +588,7 @@ class SkillsRoute(Route):
except Exception:
logger.warning("Failed to sync skills to active sandboxes.")
return Response().ok({"release": release_json, "sync": sync_json}).__dict__
return Response().ok({"release": release_json, "sync": sync_json}).to_json()
return await self._with_neo_client(_do)
@@ -607,19 +596,18 @@ class SkillsRoute(Route):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
logger.info("[Neo] POST /skills/neo/rollback requested.")
data = await request.get_json()
release_id = data.get("release_id")
if not release_id:
return Response().error("Missing release_id").__dict__
return Response().error("Missing release_id").to_json()
async def _do(client):
result = await client.skills.rollback_release(release_id)
logger.info(f"[Neo] Release rolled back: id={release_id}")
return Response().ok(_to_jsonable(result)).__dict__
return Response().ok(_to_jsonable(result)).to_json()
return await self._with_neo_client(_do)
@@ -627,8 +615,7 @@ class SkillsRoute(Route):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
logger.info("[Neo] POST /skills/neo/sync requested.")
data = await request.get_json()
@@ -636,7 +623,7 @@ class SkillsRoute(Route):
skill_key = data.get("skill_key")
require_stable = _to_bool(data.get("require_stable"), True)
if not release_id and not skill_key:
return Response().error("Missing release_id or skill_key").__dict__
return Response().error("Missing release_id or skill_key").to_json()
async def _do(client):
sync_mgr = NeoSkillSyncManager()
@@ -662,8 +649,7 @@ class SkillsRoute(Route):
"map_path": result.map_path,
"synced_at": result.synced_at,
}
)
.__dict__
).to_json()
)
return await self._with_neo_client(_do)
@@ -672,20 +658,19 @@ class SkillsRoute(Route):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
logger.info("[Neo] POST /skills/neo/delete-candidate requested.")
data = await request.get_json()
candidate_id = data.get("candidate_id")
reason = data.get("reason")
if not candidate_id:
return Response().error("Missing candidate_id").__dict__
return Response().error("Missing candidate_id").to_json()
async def _do(client):
result = await self._delete_neo_candidate(client, candidate_id, reason)
logger.info(f"[Neo] Candidate deleted: id={candidate_id}")
return Response().ok(_to_jsonable(result)).__dict__
return Response().ok(_to_jsonable(result)).to_json()
return await self._with_neo_client(_do)
@@ -693,19 +678,18 @@ class SkillsRoute(Route):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
logger.info("[Neo] POST /skills/neo/delete-release requested.")
data = await request.get_json()
release_id = data.get("release_id")
reason = data.get("reason")
if not release_id:
return Response().error("Missing release_id").__dict__
return Response().error("Missing release_id").to_json()
async def _do(client):
result = await self._delete_neo_release(client, release_id, reason)
logger.info(f"[Neo] Release deleted: id={release_id}")
return Response().ok(_to_jsonable(result)).__dict__
return Response().ok(_to_jsonable(result)).to_json()
return await self._with_neo_client(_do)
+32 -35
View File
@@ -4,6 +4,7 @@ import re
import threading
import time
import traceback
from dataclasses import asdict
from functools import cmp_to_key
from pathlib import Path
@@ -66,12 +67,11 @@ class StatRoute(Route):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
await self.core_lifecycle.restart()
return Response().ok().__dict__
return Response().ok().to_json()
def _get_running_time_components(self, total_seconds: int):
"""将总秒数转换为时分秒组件"""
@@ -94,8 +94,7 @@ class StatRoute(Route):
"change_pwd_hint": self.is_default_cred(),
"need_migration": need_migration,
},
)
.__dict__
).to_json()
)
async def get_start_time(self):
@@ -104,19 +103,19 @@ class StatRoute(Route):
self.core_lifecycle,
include_failure_details=False,
)
return Response().ok({"start_time": self.core_lifecycle.start_time}).__dict__
return Response().ok({"start_time": self.core_lifecycle.start_time}).to_json()
async def get_runtime_status(self):
return Response().ok(build_runtime_status_data(self.core_lifecycle)).__dict__
return Response().ok(build_runtime_status_data(self.core_lifecycle)).to_json()
async def get_storage_status(self):
try:
status = await asyncio.to_thread(self.storage_cleaner.get_status)
return Response().ok(status).__dict__
return Response().ok(status).to_json()
except Exception:
logger.error("获取存储占用失败", exc_info=True)
return (
Response().error("获取存储占用失败, 请查看后端日志了解详情。").__dict__
Response().error("获取存储占用失败, 请查看后端日志了解详情。").to_json()
)
async def cleanup_storage(self):
@@ -127,12 +126,12 @@ class StatRoute(Route):
target = str(data.get("target", "all"))
result = await asyncio.to_thread(self.storage_cleaner.cleanup, target)
return Response().ok(result).__dict__
return Response().ok(result).to_json()
except ValueError as e:
return Response().error(str(e)).__dict__
return Response().error(str(e)).to_json()
except Exception:
logger.error("清理存储失败", exc_info=True)
return Response().error("清理存储失败, 请查看后端日志了解详情。").__dict__
return Response().error("清理存储失败, 请查看后端日志了解详情。").to_json()
async def get_stat(self):
if not is_runtime_request_ready(self.core_lifecycle):
@@ -156,7 +155,7 @@ class StatRoute(Route):
idx += 1
message_time_based_stats.append([bucket_end, cnt])
stat_dict = stat.__dict__
stat_dict = asdict(stat)
cpu_percent = psutil.cpu_percent(interval=0.5)
thread_count = threading.active_count()
@@ -200,10 +199,10 @@ class StatRoute(Route):
},
)
return Response().ok(stat_dict).__dict__
return Response().ok(stat_dict).to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(e.__str__()).__dict__
return Response().error(e.__str__()).to_json()
async def test_ghproxy_connection(self):
"""测试 GitHub 代理连接是否可用。"""
@@ -212,7 +211,7 @@ class StatRoute(Route):
proxy_url: str = data.get("proxy_url")
if not proxy_url:
return Response().error("proxy_url is required").__dict__
return Response().error("proxy_url is required").to_json()
proxy_url = proxy_url.rstrip("/")
@@ -232,28 +231,28 @@ class StatRoute(Route):
ret = {
"latency": round((end_time - start_time) * 1000, 2),
}
return Response().ok(data=ret).__dict__
return Response().ok(data=ret).to_json()
return (
Response().error(f"Failed. Status code: {response.status}").__dict__
Response().error(f"Failed. Status code: {response.status}").to_json()
)
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Error: {e!s}").__dict__
return Response().error(f"Error: {e!s}").to_json()
async def get_changelog(self):
"""获取指定版本的更新日志"""
try:
version = request.args.get("version")
if not version:
return Response().error("version parameter is required").__dict__
return Response().error("version parameter is required").to_json()
version = version.lstrip("v")
# 防止路径遍历攻击
if not re.match(r"^[a-zA-Z0-9._-]+$", version):
return Response().error("Invalid version format").__dict__
return Response().error("Invalid version format").to_json()
if ".." in version or "/" in version or "\\" in version:
return Response().error("Invalid version format").__dict__
return Response().error("Invalid version format").to_json()
filename = f"v{version}.md"
project_path = get_astrbot_path()
@@ -267,28 +266,26 @@ class StatRoute(Route):
logger.warning(
f"Path traversal attempt detected: {version} -> {changelog_path}",
)
return Response().error("Invalid version format").__dict__
return Response().error("Invalid version format").to_json()
if not await anyio.Path(changelog_path).exists():
return (
Response()
.error(f"Changelog for version {version} not found")
.__dict__
.error(f"Changelog for version {version} not found").to_json()
)
if not await anyio.Path(changelog_path).is_file():
return (
Response()
.error(f"Changelog for version {version} not found")
.__dict__
.error(f"Changelog for version {version} not found").to_json()
)
async with await anyio.open_file(changelog_path, encoding="utf-8") as f:
content = await f.read()
return Response().ok({"content": content, "version": version}).__dict__
return Response().ok({"content": content, "version": version}).to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Error: {e!s}").__dict__
return Response().error(f"Error: {e!s}").to_json()
async def list_changelog_versions(self):
"""获取所有可用的更新日志版本列表"""
@@ -297,7 +294,7 @@ class StatRoute(Route):
changelogs_dir = os.path.join(project_path, "changelogs")
if not await anyio.Path(changelogs_dir).exists():
return Response().ok({"versions": []}).__dict__
return Response().ok({"versions": []}).to_json()
versions = []
for filename in os.listdir(changelogs_dir):
@@ -316,10 +313,10 @@ class StatRoute(Route):
),
)
return Response().ok({"versions": versions}).__dict__
return Response().ok({"versions": versions}).to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Error: {e!s}").__dict__
return Response().error(f"Error: {e!s}").to_json()
async def get_first_notice(self):
"""读取项目根目录 FIRST_NOTICE.md 内容。"""
@@ -351,9 +348,9 @@ class StatRoute(Route):
continue
content = notice_path.read_text(encoding="utf-8")
if content.strip():
return Response().ok({"content": content}).__dict__
return Response().ok({"content": content}).to_json()
return Response().ok({"content": None}).__dict__
return Response().ok({"content": None}).to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Error: {e!s}").__dict__
return Response().error(f"Error: {e!s}").to_json()
+7 -7
View File
@@ -59,16 +59,16 @@ class SubAgentRoute(Route):
if isinstance(a, dict):
a.setdefault("provider_id", None)
a.setdefault("persona_id", None)
return jsonify(Response().ok(data=data).__dict__)
return jsonify(Response().ok(data=data).to_json())
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"获取 subagent 配置失败: {e!s}").__dict__)
return jsonify(Response().error(f"获取 subagent 配置失败: {e!s}").to_json())
async def update_config(self):
try:
data = await request.json
if not isinstance(data, dict):
return jsonify(Response().error("配置必须为 JSON 对象").__dict__)
return jsonify(Response().error("配置必须为 JSON 对象").to_json())
cfg = self.core_lifecycle.astrbot_config
cfg["subagent_orchestrator"] = data
@@ -82,10 +82,10 @@ class SubAgentRoute(Route):
if orch is not None:
await orch.reload_from_config(data)
return jsonify(Response().ok(message="保存成功").__dict__)
return jsonify(Response().ok(message="保存成功").to_json())
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"保存 subagent 配置失败: {e!s}").__dict__)
return jsonify(Response().error(f"保存 subagent 配置失败: {e!s}").to_json())
async def get_available_tools(self):
"""Return all registered tools (name/description/parameters/active/origin).
@@ -111,7 +111,7 @@ class SubAgentRoute(Route):
"handler_module_path": tool.handler_module_path,
}
)
return jsonify(Response().ok(data=tools_dict).__dict__)
return jsonify(Response().ok(data=tools_dict).to_json())
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"获取可用工具失败: {e!s}").__dict__)
return jsonify(Response().error(f"获取可用工具失败: {e!s}").to_json())
+50 -68
View File
@@ -112,10 +112,10 @@ class ToolsRoute(Route):
servers.append(server_info)
return Response().ok(servers).__dict__
return Response().ok(servers).to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Failed to get MCP server list: {e!s}").__dict__
return Response().error(f"Failed to get MCP server list: {e!s}").to_json()
async def add_mcp_server(self):
try:
@@ -125,7 +125,7 @@ class ToolsRoute(Route):
# 检查必填字段
if not name:
return Response().error("Server name cannot be empty").__dict__
return Response().error("Server name cannot be empty").to_json()
# 移除特殊字段并检查配置是否有效
has_valid_config = False
@@ -140,7 +140,7 @@ class ToolsRoute(Route):
server_data["mcpServers"]
)
except ValueError as e:
return Response().error(f"{e!s}").__dict__
return Response().error(f"{e!s}").to_json()
else:
server_config[key] = value
has_valid_config = True
@@ -148,20 +148,19 @@ class ToolsRoute(Route):
if not has_valid_config:
return (
Response()
.error("A valid server configuration is required")
.__dict__
.error("A valid server configuration is required").to_json()
)
config = self.tool_mgr.load_mcp_config()
if name in config["mcpServers"]:
return Response().error(f"Server {name} already exists").__dict__
return Response().error(f"Server {name} already exists").to_json()
try:
await self.tool_mgr.test_mcp_server_connection(server_config)
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"MCP connection test failed: {e!s}").__dict__
return Response().error(f"MCP connection test failed: {e!s}").to_json()
config["mcpServers"][name] = server_config
@@ -177,23 +176,22 @@ class ToolsRoute(Route):
err_msg = f"Timed out while enabling MCP server {name}."
if not rollback_ok:
err_msg += " Configuration rollback failed. Please check the config manually."
return Response().error(err_msg).__dict__
return Response().error(err_msg).to_json()
except Exception as e:
logger.error(traceback.format_exc())
rollback_ok = self._rollback_mcp_server(name)
err_msg = f"Failed to enable MCP server {name}: {e!s}"
if not rollback_ok:
err_msg += " Configuration rollback failed. Please check the config manually."
return Response().error(err_msg).__dict__
return Response().error(err_msg).to_json()
return (
Response()
.ok(None, f"Successfully added MCP server {name}")
.__dict__
.ok(None, f"Successfully added MCP server {name}").to_json()
)
return Response().error("Failed to save configuration").__dict__
return Response().error("Failed to save configuration").to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Failed to add MCP server: {e!s}").__dict__
return Response().error(f"Failed to add MCP server: {e!s}").to_json()
async def update_mcp_server(self):
try:
@@ -203,17 +201,17 @@ class ToolsRoute(Route):
old_name = server_data.get("oldName") or name
if not name:
return Response().error("Server name cannot be empty").__dict__
return Response().error("Server name cannot be empty").to_json()
config = self.tool_mgr.load_mcp_config()
if old_name not in config["mcpServers"]:
return Response().error(f"Server {old_name} does not exist").__dict__
return Response().error(f"Server {old_name} does not exist").to_json()
is_rename = name != old_name
if name in config["mcpServers"] and is_rename:
return Response().error(f"Server {name} already exists").__dict__
return Response().error(f"Server {name} already exists").to_json()
# 获取活动状态
old_config = config["mcpServers"][old_name]
@@ -244,7 +242,7 @@ class ToolsRoute(Route):
server_data["mcpServers"]
)
except ValueError as e:
return Response().error(f"{e!s}").__dict__
return Response().error(f"{e!s}").to_json()
else:
server_config[key] = value
only_update_active = False
@@ -279,8 +277,7 @@ class ToolsRoute(Route):
Response()
.error(
f"Timed out while disabling MCP server {old_name} before enabling: {e!s}"
)
.__dict__
).to_json()
)
except Exception as e:
logger.error(traceback.format_exc())
@@ -288,8 +285,7 @@ class ToolsRoute(Route):
Response()
.error(
f"Failed to disable MCP server {old_name} before enabling: {e!s}"
)
.__dict__
).to_json()
)
try:
await self.tool_mgr.enable_mcp_server(
@@ -300,15 +296,13 @@ class ToolsRoute(Route):
except TimeoutError:
return (
Response()
.error(f"Timed out while enabling MCP server {name}.")
.__dict__
.error(f"Timed out while enabling MCP server {name}.").to_json()
)
except Exception as e:
logger.error(traceback.format_exc())
return (
Response()
.error(f"Failed to enable MCP server {name}: {e!s}")
.__dict__
.error(f"Failed to enable MCP server {name}: {e!s}").to_json()
)
# 如果要停用服务器
elif old_name in self.tool_mgr.mcp_server_runtime_view:
@@ -317,26 +311,23 @@ class ToolsRoute(Route):
except TimeoutError:
return (
Response()
.error(f"Timed out while disabling MCP server {old_name}.")
.__dict__
.error(f"Timed out while disabling MCP server {old_name}.").to_json()
)
except Exception as e:
logger.error(traceback.format_exc())
return (
Response()
.error(f"Failed to disable MCP server {old_name}: {e!s}")
.__dict__
.error(f"Failed to disable MCP server {old_name}: {e!s}").to_json()
)
return (
Response()
.ok(None, f"Successfully updated MCP server {name}")
.__dict__
.ok(None, f"Successfully updated MCP server {name}").to_json()
)
return Response().error("Failed to save configuration").__dict__
return Response().error("Failed to save configuration").to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Failed to update MCP server: {e!s}").__dict__
return Response().error(f"Failed to update MCP server: {e!s}").to_json()
async def delete_mcp_server(self):
try:
@@ -344,12 +335,12 @@ class ToolsRoute(Route):
name = server_data.get("name", "")
if not name:
return Response().error("Server name cannot be empty").__dict__
return Response().error("Server name cannot be empty").to_json()
config = self.tool_mgr.load_mcp_config()
if name not in config["mcpServers"]:
return Response().error(f"Server {name} does not exist").__dict__
return Response().error(f"Server {name} does not exist").to_json()
del config["mcpServers"][name]
@@ -360,25 +351,22 @@ class ToolsRoute(Route):
except TimeoutError:
return (
Response()
.error(f"Timed out while disabling MCP server {name}.")
.__dict__
.error(f"Timed out while disabling MCP server {name}.").to_json()
)
except Exception as e:
logger.error(traceback.format_exc())
return (
Response()
.error(f"Failed to disable MCP server {name}: {e!s}")
.__dict__
.error(f"Failed to disable MCP server {name}: {e!s}").to_json()
)
return (
Response()
.ok(None, f"Successfully deleted MCP server {name}")
.__dict__
.ok(None, f"Successfully deleted MCP server {name}").to_json()
)
return Response().error("Failed to save configuration").__dict__
return Response().error("Failed to save configuration").to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Failed to delete MCP server: {e!s}").__dict__
return Response().error(f"Failed to delete MCP server: {e!s}").to_json()
async def test_mcp_connection(self):
"""Test MCP server connection."""
@@ -387,7 +375,7 @@ class ToolsRoute(Route):
config = server_data.get("mcp_server_config", None)
if not isinstance(config, dict) or not config:
return Response().error("Invalid MCP server configuration").__dict__
return Response().error("Invalid MCP server configuration").to_json()
if "mcpServers" in config:
mcp_servers = config["mcpServers"]
@@ -396,36 +384,32 @@ class ToolsRoute(Route):
Response()
.error(
"Only one MCP server configuration can be tested at a time"
)
.__dict__
).to_json()
)
try:
config = _extract_mcp_server_config(mcp_servers)
except EmptyMcpServersError:
return (
Response()
.error("MCP server configuration cannot be empty")
.__dict__
.error("MCP server configuration cannot be empty").to_json()
)
except ValueError as e:
return Response().error(f"{e!s}").__dict__
return Response().error(f"{e!s}").to_json()
elif not config:
return (
Response()
.error("MCP server configuration cannot be empty")
.__dict__
.error("MCP server configuration cannot be empty").to_json()
)
tools_name = await self.tool_mgr.test_mcp_server_connection(config)
return (
Response()
.ok(data=tools_name, message="🎉 MCP server is available!")
.__dict__
.ok(data=tools_name, message="🎉 MCP server is available!").to_json()
)
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Failed to test MCP connection: {e!s}").__dict__
return Response().error(f"Failed to test MCP connection: {e!s}").to_json()
async def get_tool_list(self):
"""Get all registered tools."""
@@ -462,10 +446,10 @@ class ToolsRoute(Route):
"source": source,
}
tools_dict.append(tool_info)
return Response().ok(data=tools_dict).__dict__
return Response().ok(data=tools_dict).to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Failed to get tool list: {e!s}").__dict__
return Response().error(f"Failed to get tool list: {e!s}").to_json()
async def toggle_tool(self):
"""Activate or deactivate a specified tool."""
@@ -477,34 +461,32 @@ class ToolsRoute(Route):
if not tool_name or action is None:
return (
Response()
.error("Missing required parameters: name or activate")
.__dict__
.error("Missing required parameters: name or activate").to_json()
)
# Internal tools cannot be toggled by users
for t in self.tool_mgr.func_list:
if t.name == tool_name and getattr(t, "source", "") == "internal":
return Response().error("内置工具不支持手动启用/停用").__dict__
return Response().error("内置工具不支持手动启用/停用").to_json()
if action:
try:
ok = self.tool_mgr.activate_llm_tool(tool_name, star_map=star_map)
except ValueError as e:
return Response().error(f"Failed to activate tool: {e!s}").__dict__
return Response().error(f"Failed to activate tool: {e!s}").to_json()
else:
ok = self.tool_mgr.deactivate_llm_tool(tool_name)
if ok:
return Response().ok(None, "Operation successful.").__dict__
return Response().ok(None, "Operation successful.").to_json()
return (
Response()
.error(f"Tool {tool_name} does not exist or the operation failed.")
.__dict__
.error(f"Tool {tool_name} does not exist or the operation failed.").to_json()
)
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Failed to operate tool: {e!s}").__dict__
return Response().error(f"Failed to operate tool: {e!s}").to_json()
async def sync_provider(self):
"""Sync MCP provider configuration."""
@@ -517,10 +499,10 @@ class ToolsRoute(Route):
await self.tool_mgr.sync_modelscope_mcp_servers(access_token)
case _:
return (
Response().error(f"Unknown provider: {provider_name}").__dict__
Response().error(f"Unknown provider: {provider_name}").to_json()
)
return Response().ok(message="Sync completed").__dict__
return Response().ok(message="Sync completed").to_json()
except Exception as e:
logger.error(traceback.format_exc())
return Response().error(f"Sync failed: {e!s}").__dict__
return Response().error(f"Sync failed: {e!s}").to_json()
+37 -41
View File
@@ -95,7 +95,7 @@ class TUIChatRoute(Route):
async def get_file(self):
filename = request.args.get("filename")
if not filename:
return Response().error("Missing key: filename").__dict__
return Response().error("Missing key: filename").to_json()
try:
file_path = os.path.join(self.attachments_dir, os.path.basename(filename))
@@ -103,12 +103,12 @@ class TUIChatRoute(Route):
resolved_base_dir = _resolve_path(self.attachments_dir)
if not await anyio.Path(resolved_file_path).exists():
return Response().error("File not found").__dict__
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").__dict__
return Response().error("Invalid file path").to_json()
filename_ext = os.path.splitext(filename)[1].lower()
if filename_ext == ".wav":
@@ -118,18 +118,18 @@ class TUIChatRoute(Route):
return await send_file(str(resolved_file_path))
except (FileNotFoundError, OSError):
return Response().error("File access error").__dict__
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").__dict__
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").__dict__
return Response().error("Attachment not found").to_json()
file_path = attachment.path
resolved_file_path = _resolve_path(file_path)
@@ -139,13 +139,13 @@ class TUIChatRoute(Route):
)
except (FileNotFoundError, OSError):
return Response().error("File access error").__dict__
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").__dict__
return Response().error("Missing key: file").to_json()
file = post_data["file"]
filename = file.filename or f"{uuid.uuid4()!s}"
@@ -170,7 +170,7 @@ class TUIChatRoute(Route):
)
if not attachment:
return Response().error("Failed to create attachment").__dict__
return Response().error("Failed to create attachment").to_json()
filename = os.path.basename(attachment.path)
@@ -182,8 +182,7 @@ class TUIChatRoute(Route):
"filename": filename,
"type": attach_type,
}
)
.__dict__
).to_json()
)
async def _build_user_message_parts(self, message: str | list) -> list[dict]:
@@ -243,13 +242,13 @@ class TUIChatRoute(Route):
if post_data is None:
post_data = await request.json
if post_data is None:
return Response().error("Missing JSON body").__dict__
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").__dict__
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").__dict__
Response().error("Missing key: session_id or conversation_id").to_json()
)
message = post_data["message"]
@@ -259,7 +258,7 @@ class TUIChatRoute(Route):
enable_streaming = post_data.get("enable_streaming", True)
if not session_id:
return Response().error("session_id is empty").__dict__
return Response().error("session_id is empty").to_json()
tui_conv_id = session_id
@@ -267,8 +266,7 @@ class TUIChatRoute(Route):
if not webchat_message_parts_have_content(message_parts):
return (
Response()
.error("Message content is empty (reply only is not allowed)")
.__dict__
.error("Message content is empty (reply only is not allowed)").to_json()
)
message_id = str(uuid.uuid4())
@@ -476,18 +474,18 @@ class TUIChatRoute(Route):
"""Stop active agent runs for a session."""
post_data = await request.json
if post_data is None:
return Response().error("Missing JSON body").__dict__
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").__dict__
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").__dict__
return Response().error(f"Session {session_id} not found").to_json()
if session.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
message_type = (
MessageType.GROUP_MESSAGE.value
@@ -500,7 +498,7 @@ class TUIChatRoute(Route):
)
stopped_count = active_event_registry.request_agent_stop_all(umo)
return Response().ok(data={"stopped_count": stopped_count}).__dict__
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."""
@@ -543,30 +541,30 @@ class TUIChatRoute(Route):
"""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").__dict__
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").__dict__
return Response().error(f"Session {session_id} not found").to_json()
if session.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
await self._delete_session_internal(session, username)
return Response().ok().__dict__
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").__dict__
return Response().error("Missing JSON body").to_json()
if not isinstance(post_data, dict):
return Response().error("Invalid JSON body: expected object").__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").__dict__
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)
@@ -599,8 +597,7 @@ class TUIChatRoute(Route):
"failed_count": len(failed_items),
"failed_items": failed_items,
}
)
.__dict__
).to_json()
)
def _extract_attachment_ids(self, history_list) -> list[str]:
@@ -654,8 +651,7 @@ class TUIChatRoute(Route):
"session_id": session.session_id,
"platform_id": session.platform_id,
}
)
.__dict__
).to_json()
)
async def get_sessions(self):
@@ -688,13 +684,13 @@ class TUIChatRoute(Route):
}
)
return Response().ok(data=sessions_data).__dict__
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").__dict__
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"
@@ -725,7 +721,7 @@ class TUIChatRoute(Route):
"emoji": project_info.emoji,
}
return Response().ok(data=response_data).__dict__
return Response().ok(data=response_data).to_json()
async def update_session_display_name(self):
"""Update a Platform session's display name."""
@@ -735,21 +731,21 @@ class TUIChatRoute(Route):
display_name = post_data.get("display_name")
if not session_id:
return Response().error("Missing key: session_id").__dict__
return Response().error("Missing key: session_id").to_json()
if display_name is None:
return Response().error("Missing key: display_name").__dict__
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").__dict__
return Response().error(f"Session {session_id} not found").to_json()
if session.creator != username:
return Response().error("Permission denied").__dict__
return Response().error("Permission denied").to_json()
await self.db.update_platform_session(
session_id=session_id,
display_name=display_name,
)
return Response().ok().__dict__
return Response().ok().to_json()
+18 -22
View File
@@ -37,7 +37,7 @@ class UpdateRoute(Route):
async def do_migration(self):
need_migration = await check_migration_needed_v4(self.core_lifecycle.db)
if not need_migration:
return Response().ok(None, "不需要进行迁移。").__dict__
return Response().ok(None, "不需要进行迁移。").to_json()
try:
data = await request.json
pim = data.get("platform_id_map", {})
@@ -46,10 +46,10 @@ class UpdateRoute(Route):
pim,
self.core_lifecycle.astrbot_config,
)
return Response().ok(None, "迁移成功。").__dict__
return Response().ok(None, "迁移成功。").to_json()
except Exception as e:
logger.error(f"迁移失败: {traceback.format_exc()}")
return Response().error(f"迁移失败: {e!s}").__dict__
return Response().error(f"迁移失败: {e!s}").to_json()
async def check_update(self):
type_ = request.args.get("type", None)
@@ -59,8 +59,7 @@ class UpdateRoute(Route):
if type_ == "dashboard":
return (
Response()
.ok({"has_new_version": dv != f"v{VERSION}", "current_version": dv})
.__dict__
.ok({"has_new_version": dv != f"v{VERSION}", "current_version": dv}).to_json()
)
ret = await self.astrbot_updator.check_update(None, None, False)
return Response(
@@ -72,18 +71,18 @@ class UpdateRoute(Route):
"dashboard_version": dv,
"dashboard_has_new_version": bool(dv and dv != f"v{VERSION}"),
},
).__dict__
).to_json()
except Exception as e:
logger.warning(f"检查更新失败: {e!s} (不影响除项目更新外的正常使用)")
return Response().error(e.__str__()).__dict__
return Response().error(e.__str__()).to_json()
async def get_releases(self):
try:
ret = await self.astrbot_updator.get_releases()
return Response().ok(ret).__dict__
return Response().ok(ret).to_json()
except Exception as e:
logger.error(f"/api/update/releases: {traceback.format_exc()}")
return Response().error(e.__str__()).__dict__
return Response().error(e.__str__()).to_json()
async def update_project(self):
data = await request.json
@@ -122,19 +121,17 @@ class UpdateRoute(Route):
await self.core_lifecycle.restart()
ret = (
Response()
.ok(None, "更新成功,AstrBot 将在 2 秒内全量重启以应用新的代码。")
.__dict__
.ok(None, "更新成功,AstrBot 将在 2 秒内全量重启以应用新的代码。").to_json()
)
return ret, 200, CLEAR_SITE_DATA_HEADERS
ret = (
Response()
.ok(None, "更新成功,AstrBot 将在下次启动时应用新的代码。")
.__dict__
.ok(None, "更新成功,AstrBot 将在下次启动时应用新的代码。").to_json()
)
return ret, 200, CLEAR_SITE_DATA_HEADERS
except Exception as e:
logger.error(f"/api/update_project: {traceback.format_exc()}")
return Response().error(e.__str__()).__dict__
return Response().error(e.__str__()).to_json()
async def update_dashboard(self):
try:
@@ -142,29 +139,28 @@ class UpdateRoute(Route):
await download_dashboard(version=f"v{VERSION}", latest=False)
except Exception as e:
logger.error(f"下载管理面板文件失败: {e}。")
return Response().error(f"下载管理面板文件失败: {e}").__dict__
ret = Response().ok(None, "更新成功。刷新页面即可应用新版本面板。").__dict__
return Response().error(f"下载管理面板文件失败: {e}").to_json()
ret = Response().ok(None, "更新成功。刷新页面即可应用新版本面板。").to_json()
return ret, 200, CLEAR_SITE_DATA_HEADERS
except Exception as e:
logger.error(f"/api/update_dashboard: {traceback.format_exc()}")
return Response().error(e.__str__()).__dict__
return Response().error(e.__str__()).to_json()
async def install_pip_package(self):
if DEMO_MODE:
return (
Response()
.error("You are not permitted to do this operation in demo mode")
.__dict__
.error("You are not permitted to do this operation in demo mode").to_json()
)
data = await request.json
package = data.get("package", "")
mirror = data.get("mirror", None)
if not package:
return Response().error("缺少参数 package 或不合法。").__dict__
return Response().error("缺少参数 package 或不合法。").to_json()
try:
await pip_installer.install(package, mirror=mirror)
return Response().ok(None, "安装成功。").__dict__
return Response().ok(None, "安装成功。").to_json()
except Exception as e:
logger.error(f"/api/update_pip: {traceback.format_exc()}")
return Response().error(e.__str__()).__dict__
return Response().error(e.__str__()).to_json()
+3 -3
View File
@@ -418,7 +418,7 @@ class AstrBotDashboard:
if request.path.startswith("/api/v1"):
raw_key = self._extract_raw_api_key()
if not raw_key:
r = jsonify(Response().error("Missing API key").__dict__)
r = jsonify(Response().error("Missing API key").to_json())
r.status_code = 401
return r
key_hash = hashlib.pbkdf2_hmac(
@@ -429,7 +429,7 @@ class AstrBotDashboard:
).hex()
api_key = await self.db.get_active_api_key_by_hash(key_hash)
if not api_key:
r = jsonify(Response().error("Invalid API key").__dict__)
r = jsonify(Response().error("Invalid API key").to_json())
r.status_code = 401
return r
@@ -439,7 +439,7 @@ class AstrBotDashboard:
scopes = list(ALL_OPEN_API_SCOPES)
required_scope = self._get_required_open_api_scope(request.path)
if required_scope and "*" not in scopes and required_scope not in scopes:
r = jsonify(Response().error("Insufficient API key scope").__dict__)
r = jsonify(Response().error("Insufficient API key scope").to_json())
r.status_code = 403
return r
+2
View File
@@ -11,6 +11,8 @@
rel="stylesheet"
href="https://fonts.googleapis.com/css2?family=Outfit&family=Poppins:wght@400;500;600;700&family=Roboto:wght@400;500;700&display=swap"
/>
<!-- MDI Icons Font (CDN) -->
<link rel="stylesheet" href="https://cdn.jsdelivr.net/npm/@mdi/font@7.4.47/css/materialdesignicons.min.css" />
<!-- VAD (Voice Activity Detection) Libraries -->
<script src="https://cdn.jsdelivr.net/npm/onnxruntime-web@1.22.0/dist/ort.wasm.min.js"></script>
<script src="https://cdn.jsdelivr.net/npm/@ricky0123/vad-web@0.0.29/dist/bundle.min.js"></script>
+23
View File
@@ -0,0 +1,23 @@
/* MDI Icons Font - served from public/fonts/ */
@font-face {
font-family: "Material Design Icons";
src: url("../fonts/materialdesignicons-webfont.eot?v=7.4.47");
src: url("../fonts/materialdesignicons-webfont.eot?#iefix&v=7.4.47") format("embedded-opentype"),
url("../fonts/materialdesignicons-webfont.woff2?v=7.4.47") format("woff2"),
url("../fonts/materialdesignicons-webfont.woff?v=7.4.47") format("woff"),
url("../fonts/materialdesignicons-webfont.ttf?v=7.4.47") format("truetype");
font-weight: normal;
font-style: normal;
font-display: block;
}
.mdi:before,
.mdi-set {
display: inline-block;
font: normal normal normal 24px/1 "Material Design Icons";
font-size: inherit;
text-rendering: auto;
line-height: inherit;
-webkit-font-smoothing: antialiased;
-moz-osx-font-smoothing: grayscale;
}
@@ -39,6 +39,7 @@
"subtitle": "自定义主题主色与辅助色。修改后立即生效,并保存在浏览器本地。",
"customize": {
"title": "主题颜色",
"colors": "自定义颜色",
"preset": "配色方案",
"primary": "主色",
"secondary": "辅助色",

Some files were not shown because too many files have changed in this diff Show More