mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
chore: smart commit — update AGENTS.md, run ruff format, and apply small targeted fixes
This commit is contained in:
@@ -56,4 +56,4 @@ Runs on `http://localhost:3000` by default.
|
||||
## PR instructions
|
||||
|
||||
1. Title format: use conventional commit messages
|
||||
2. Use English to write PR title and descriptions.
|
||||
2. Use English to write PR title and descriptions./<
|
||||
@@ -21,13 +21,13 @@ if TYPE_CHECKING:
|
||||
else:
|
||||
try:
|
||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
except ImportError:
|
||||
logger.warning("FastAPI not installed, gateway unavailable.")
|
||||
FastAPI = cast(Any, None)
|
||||
WebSocket = cast(Any, None)
|
||||
WebSocketDisconnect = cast(Any, None)
|
||||
CORSMiddleware = cast(Any, None)
|
||||
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
log = logger
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ import traceback
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Sequence
|
||||
from collections.abc import Set as AbstractSet
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import mcp
|
||||
|
||||
@@ -236,7 +236,7 @@ class FunctionToolExecutor(BaseFunctionToolExecutor[AstrAgentContext]):
|
||||
runtime: str,
|
||||
sandbox_cfg: dict | None = None,
|
||||
session_id: str = "",
|
||||
) -> dict[str, FunctionTool]:
|
||||
) -> dict[str, ToolSchema]:
|
||||
from astrbot.core.computer.computer_tool_provider import ComputerToolProvider
|
||||
from astrbot.core.tool_provider import ToolProviderContext
|
||||
|
||||
@@ -643,7 +643,7 @@ class FunctionToolExecutor(BaseFunctionToolExecutor[AstrAgentContext]):
|
||||
awaitable = tool.call
|
||||
method_name = "call"
|
||||
elif hasattr(tool, "run"):
|
||||
awaitable = getattr(tool, "run")
|
||||
awaitable = tool.run
|
||||
method_name = "run"
|
||||
if awaitable is None:
|
||||
raise ValueError("Tool must have a valid handler or override 'run' method.")
|
||||
@@ -666,9 +666,16 @@ class FunctionToolExecutor(BaseFunctionToolExecutor[AstrAgentContext]):
|
||||
except Exception as exc:
|
||||
logger.warning("SDK calling_func_tool dispatch failed: %s", exc)
|
||||
|
||||
# awaitable is guaranteed non-None after line 648 check.
|
||||
# cast needed: pyright cannot narrow type through the conditional chain.
|
||||
_HandlerType = Callable[
|
||||
...,
|
||||
Awaitable[MessageEventResult | mcp.types.CallToolResult | str | None]
|
||||
| AsyncGenerator[MessageEventResult | CommandResult | str | None, None],
|
||||
]
|
||||
wrapper = call_local_llm_tool(
|
||||
context=run_context,
|
||||
handler=awaitable,
|
||||
handler=cast(_HandlerType, awaitable),
|
||||
method_name=method_name,
|
||||
**tool_args,
|
||||
)
|
||||
@@ -716,7 +723,7 @@ class FunctionToolExecutor(BaseFunctionToolExecutor[AstrAgentContext]):
|
||||
except asyncio.TimeoutError:
|
||||
raise Exception(
|
||||
f"tool {tool.name} execution timeout after {tool_call_timeout or run_context.tool_call_timeout} seconds.",
|
||||
)
|
||||
) from None
|
||||
except StopAsyncIteration:
|
||||
break
|
||||
|
||||
|
||||
@@ -242,7 +242,7 @@ async def _apply_file_extract(
|
||||
logger.error("Unsupported file extract provider: %s", config.file_extract_prov)
|
||||
return
|
||||
|
||||
for file_content, file_name in zip(file_contents, file_names):
|
||||
for file_content, file_name in zip(file_contents, file_names, strict=True):
|
||||
req.contexts.append(
|
||||
{
|
||||
"role": "system",
|
||||
@@ -952,7 +952,7 @@ def _plugin_tool_fix(event: AstrMessageEvent, req: ProviderRequest) -> None:
|
||||
# 保留 MCP 工具
|
||||
new_tool_set.add_tool(tool)
|
||||
continue
|
||||
mp = tool.handler_module_path
|
||||
mp = getattr(tool, "handler_module_path", None)
|
||||
if not mp:
|
||||
# 没有 plugin 归属信息的工具(如 subagent transfer_to_*)
|
||||
# 不应受到会话插件过滤影响。
|
||||
|
||||
@@ -329,7 +329,7 @@ async def retrieve_knowledge_base(
|
||||
# 如果配置为空列表,明确表示不使用知识库
|
||||
if not kb_ids:
|
||||
logger.info(f"[知识库] 会话 {umo} 已被配置为不使用知识库")
|
||||
return
|
||||
return None
|
||||
|
||||
top_k = session_config.get("top_k", 5)
|
||||
|
||||
@@ -350,7 +350,7 @@ async def retrieve_knowledge_base(
|
||||
)
|
||||
|
||||
if not kb_names:
|
||||
return
|
||||
return None
|
||||
|
||||
logger.debug(f"[知识库] 使用会话级配置,知识库数量: {len(kb_names)}")
|
||||
else:
|
||||
@@ -361,13 +361,13 @@ async def retrieve_knowledge_base(
|
||||
top_k_fusion = config.get("kb_fusion_top_k", 20)
|
||||
|
||||
if not kb_names:
|
||||
return
|
||||
return None
|
||||
|
||||
all_kbs = [await kb_mgr.get_kb_by_name(kb) for kb in kb_names]
|
||||
|
||||
if check_all_kb(all_kbs):
|
||||
logger.debug("所配置的所有知识库全为空, 跳过检索过程")
|
||||
return
|
||||
return None
|
||||
|
||||
logger.debug(f"[知识库] 开始检索知识库,数量: {len(kb_names)}, top_k={top_k}")
|
||||
|
||||
@@ -379,13 +379,14 @@ async def retrieve_knowledge_base(
|
||||
)
|
||||
|
||||
if not kb_context:
|
||||
return
|
||||
return None
|
||||
|
||||
formatted = kb_context.get("context_text", "")
|
||||
if formatted:
|
||||
results = kb_context.get("results", [])
|
||||
logger.debug(f"[知识库] 为会话 {umo} 注入了 {len(results)} 条相关知识块")
|
||||
return formatted
|
||||
return None
|
||||
|
||||
|
||||
KNOWLEDGE_BASE_QUERY_TOOL = KnowledgeBaseQueryTool()
|
||||
|
||||
@@ -11,7 +11,7 @@ from astrbot.core.computer.olayer import (
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from astrbot.core.agent.tool import FunctionTool, ToolSchema
|
||||
from astrbot.core.agent.tool import ToolSchema
|
||||
|
||||
|
||||
class ComputerBooter(abc.ABC):
|
||||
|
||||
@@ -15,7 +15,7 @@ from shipyard.shell import ShellComponent as ShipyardShellComponent
|
||||
from astrbot.api import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from astrbot.core.agent.tool import FunctionTool
|
||||
from astrbot.core.agent.tool import FunctionTool, ToolSchema
|
||||
|
||||
from astrbot.core.computer.olayer import (
|
||||
FileSystemComponent,
|
||||
@@ -178,7 +178,7 @@ class BoxliteBooter(ComputerBooter):
|
||||
session_id,
|
||||
)
|
||||
random_port = random.randint(20000, 30000)
|
||||
self.box = boxlite.SimpleBox( # type: ignore[attr-defined]
|
||||
self.box = boxlite.SimpleBox( # type: ignore[unresolved-attribute]
|
||||
image="soulter/shipyard-ship",
|
||||
memory_mib=512,
|
||||
cpus=1,
|
||||
@@ -261,5 +261,5 @@ class BoxliteBooter(ComputerBooter):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_default_tools(cls) -> list[FunctionTool]:
|
||||
def get_default_tools(cls) -> list[ToolSchema]:
|
||||
return list(cls._default_tools())
|
||||
|
||||
@@ -8,7 +8,7 @@ from shipyard import ShipyardClient, Spec
|
||||
from astrbot.api import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from astrbot.core.agent.tool import FunctionTool, ToolSchema
|
||||
from astrbot.core.agent.tool import ToolSchema
|
||||
|
||||
from astrbot.core.computer.olayer import (
|
||||
FileSystemComponent,
|
||||
@@ -30,7 +30,7 @@ class ShipyardBooter(ComputerBooter):
|
||||
PythonTool,
|
||||
)
|
||||
|
||||
return (
|
||||
return ( # type: ignore[return-value]
|
||||
ExecuteShellTool(),
|
||||
PythonTool(),
|
||||
FileUploadTool(),
|
||||
|
||||
@@ -10,7 +10,7 @@ import anyio
|
||||
from astrbot.api import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from astrbot.core.agent.tool import FunctionTool
|
||||
from astrbot.core.agent.tool import ToolSchema
|
||||
|
||||
from astrbot.core.computer.booters.base import ComputerBooter
|
||||
from astrbot.core.computer.olayer import (
|
||||
@@ -35,15 +35,15 @@ class NeoPythonComponent(PythonComponent):
|
||||
def __init__(self, sandbox: Any) -> None:
|
||||
self._sandbox = sandbox
|
||||
|
||||
async def exec( # type: ignore[override]
|
||||
async def exec(
|
||||
self,
|
||||
code: str,
|
||||
kernel_id: str | None = None,
|
||||
timeout_sec: int = 30,
|
||||
timeout: int = 30,
|
||||
silent: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
_ = kernel_id # Bay runtime does not expose kernel_id in current SDK.
|
||||
with anyio.fail_after(timeout_sec):
|
||||
with anyio.fail_after(timeout):
|
||||
result = await self._sandbox.python.exec(code)
|
||||
payload = _maybe_model_dump(result)
|
||||
|
||||
@@ -77,12 +77,12 @@ class NeoShellComponent(ShellComponent):
|
||||
def __init__(self, sandbox: Any) -> None:
|
||||
self._sandbox = sandbox
|
||||
|
||||
async def exec( # type: ignore[override]
|
||||
async def exec(
|
||||
self,
|
||||
command: str,
|
||||
cwd: str | None = None,
|
||||
env: dict[str, str] | None = None,
|
||||
timeout_sec: int | None = 30,
|
||||
timeout: int | None = 30,
|
||||
shell: bool = True,
|
||||
background: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
@@ -104,7 +104,7 @@ class NeoShellComponent(ShellComponent):
|
||||
if background:
|
||||
run_command = f"nohup sh -lc {shlex.quote(run_command)} >/tmp/astrbot_bg.log 2>&1 & echo $!"
|
||||
|
||||
with anyio.fail_after(timeout_sec or 30):
|
||||
with anyio.fail_after(timeout or 30):
|
||||
result = await self._sandbox.shell.exec(run_command, cwd=cwd)
|
||||
payload = _maybe_model_dump(result)
|
||||
|
||||
@@ -534,7 +534,7 @@ class ShipyardNeoBooter(ComputerBooter):
|
||||
|
||||
@classmethod
|
||||
@functools.cache
|
||||
def _base_tools(cls) -> tuple[FunctionTool, ...]:
|
||||
def _base_tools(cls) -> tuple[ToolSchema, ...]:
|
||||
"""4 base + 11 Neo lifecycle = 15 tools (all Neo profiles)."""
|
||||
from astrbot.core.computer.tools import (
|
||||
AnnotateExecutionTool,
|
||||
@@ -554,7 +554,7 @@ class ShipyardNeoBooter(ComputerBooter):
|
||||
SyncSkillReleaseTool,
|
||||
)
|
||||
|
||||
return ( # type: ignore[return-value]
|
||||
return (
|
||||
ExecuteShellTool(),
|
||||
PythonTool(),
|
||||
FileUploadTool(),
|
||||
@@ -574,21 +574,21 @@ class ShipyardNeoBooter(ComputerBooter):
|
||||
|
||||
@classmethod
|
||||
@functools.cache
|
||||
def _browser_tools(cls) -> tuple[FunctionTool, ...]:
|
||||
def _browser_tools(cls) -> tuple[ToolSchema, ...]:
|
||||
from astrbot.core.computer.tools import (
|
||||
BrowserBatchExecTool,
|
||||
BrowserExecTool,
|
||||
RunBrowserSkillTool,
|
||||
)
|
||||
|
||||
return (BrowserExecTool(), BrowserBatchExecTool(), RunBrowserSkillTool()) # type: ignore[return-value]
|
||||
return (BrowserExecTool(), BrowserBatchExecTool(), RunBrowserSkillTool())
|
||||
|
||||
@classmethod
|
||||
def get_default_tools(cls) -> list[FunctionTool]:
|
||||
def get_default_tools(cls) -> list[ToolSchema]:
|
||||
"""Pre-boot: conservative full list (including browser)."""
|
||||
return list(cls._base_tools()) + list(cls._browser_tools())
|
||||
|
||||
def get_tools(self) -> list[FunctionTool]:
|
||||
def get_tools(self) -> list[ToolSchema]:
|
||||
"""Post-boot: capability-filtered list."""
|
||||
caps = self.capabilities
|
||||
if caps is None:
|
||||
|
||||
@@ -20,7 +20,7 @@ from .booters.constants import BOOTER_BOXLITE, BOOTER_SHIPYARD, BOOTER_SHIPYARD_
|
||||
from .booters.local import LocalBooter
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from astrbot.core.agent.tool import FunctionTool
|
||||
from astrbot.core.agent.tool import ToolSchema
|
||||
|
||||
session_booter: dict[str, ComputerBooter] = {}
|
||||
local_booter: ComputerBooter | None = None
|
||||
@@ -578,7 +578,7 @@ def _get_booter_class(booter_type: str) -> type[ComputerBooter] | None:
|
||||
return None
|
||||
|
||||
|
||||
def get_sandbox_tools(session_id: str) -> list[FunctionTool]:
|
||||
def get_sandbox_tools(session_id: str) -> list[ToolSchema]:
|
||||
"""Return precise tool list from a booted session, or [] if not booted."""
|
||||
booter = session_booter.get(session_id)
|
||||
if booter is None:
|
||||
@@ -618,7 +618,7 @@ def get_sandbox_capabilities(session_id: str) -> tuple[str, ...] | None:
|
||||
return caps
|
||||
|
||||
|
||||
def get_default_sandbox_tools(sandbox_cfg: dict) -> list[FunctionTool]:
|
||||
def get_default_sandbox_tools(sandbox_cfg: dict) -> list[ToolSchema]:
|
||||
"""Return conservative (pre-boot) tool list based on config. No instance needed."""
|
||||
booter_type = sandbox_cfg.get("booter", BOOTER_SHIPYARD_NEO)
|
||||
cls = _get_booter_class(booter_type)
|
||||
|
||||
@@ -19,26 +19,20 @@ from astrbot.api import logger
|
||||
from astrbot.core.tool_provider import ToolProviderContext
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from astrbot.core.agent.tool import FunctionTool
|
||||
from astrbot.core.agent.tool import FunctionTool, ToolSchema
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Lazy local-mode tool cache
|
||||
# Local mode tools
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_LOCAL_TOOLS_CACHE: list[FunctionTool] | None = None
|
||||
|
||||
def _get_local_tools() -> list[ToolSchema]:
|
||||
from astrbot.core.computer.tools import ExecuteShellTool, LocalPythonTool
|
||||
|
||||
def _get_local_tools() -> list[FunctionTool]:
|
||||
global _LOCAL_TOOLS_CACHE
|
||||
if _LOCAL_TOOLS_CACHE is None:
|
||||
from astrbot.core.computer.tools import ExecuteShellTool, LocalPythonTool
|
||||
|
||||
_LOCAL_TOOLS_CACHE = [
|
||||
ExecuteShellTool(is_local=True),
|
||||
LocalPythonTool(),
|
||||
]
|
||||
return list(_LOCAL_TOOLS_CACHE)
|
||||
shell = ExecuteShellTool(is_local=True)
|
||||
python = LocalPythonTool()
|
||||
return [shell, python] # type: ignore[return-value]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -81,7 +75,7 @@ class ComputerToolProvider:
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def get_all_tools() -> list[FunctionTool]:
|
||||
def get_all_tools() -> list[ToolSchema]:
|
||||
"""Return ALL computer-use tools across all runtimes for registration.
|
||||
|
||||
Creates **fresh instances** separate from the runtime caches so that
|
||||
@@ -114,7 +108,7 @@ class ComputerToolProvider:
|
||||
SyncSkillReleaseTool,
|
||||
)
|
||||
|
||||
all_tools: list[FunctionTool] = [
|
||||
all_tools: list[ToolSchema] = [ # type: ignore[assignment]
|
||||
ExecuteShellTool(),
|
||||
PythonTool(),
|
||||
FileUploadTool(),
|
||||
@@ -139,7 +133,7 @@ class ComputerToolProvider:
|
||||
# De-duplicate by name and mark inactive so they are visible
|
||||
# in WebUI but never sent to the LLM via func_list.
|
||||
seen: set[str] = set()
|
||||
result: list[FunctionTool] = []
|
||||
result: list[ToolSchema] = []
|
||||
for tool in all_tools:
|
||||
if tool.name not in seen:
|
||||
tool.active = False
|
||||
@@ -147,7 +141,7 @@ class ComputerToolProvider:
|
||||
seen.add(tool.name)
|
||||
return result
|
||||
|
||||
def get_tools(self, ctx: ToolProviderContext) -> list[FunctionTool]:
|
||||
def get_tools(self, ctx: ToolProviderContext) -> list[ToolSchema]:
|
||||
runtime = ctx.computer_use_runtime
|
||||
if runtime == "none":
|
||||
return []
|
||||
@@ -176,7 +170,7 @@ class ComputerToolProvider:
|
||||
|
||||
# -- sandbox helpers ----------------------------------------------------
|
||||
|
||||
def _sandbox_tools(self, ctx: ToolProviderContext) -> list[FunctionTool]:
|
||||
def _sandbox_tools(self, ctx: ToolProviderContext) -> list[ToolSchema]:
|
||||
"""Collect tools for sandbox mode.
|
||||
|
||||
Always returns the full (pre-boot default) tool set declared by the
|
||||
@@ -213,7 +207,7 @@ class ComputerToolProvider:
|
||||
return "".join(parts)
|
||||
|
||||
|
||||
def get_all_tools() -> list[FunctionTool]:
|
||||
def get_all_tools() -> list[ToolSchema]:
|
||||
"""Module-level entry point for ``FunctionToolManager.register_internal_tools()``.
|
||||
|
||||
Delegates to ``ComputerToolProvider.get_all_tools()`` which collects
|
||||
|
||||
@@ -59,7 +59,7 @@ class BrowserExecTool(FunctionTool):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
cmd: str = "",
|
||||
@@ -122,7 +122,7 @@ class BrowserBatchExecTool(FunctionTool):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
commands: list[str] | None = None,
|
||||
@@ -171,7 +171,7 @@ class RunBrowserSkillTool(FunctionTool):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
skill_key: str = "",
|
||||
|
||||
@@ -105,7 +105,7 @@ class FileUploadTool(FunctionTool):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
local_path: str,
|
||||
@@ -170,7 +170,7 @@ class FileDownloadTool(FunctionTool):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
remote_path: str,
|
||||
|
||||
@@ -84,7 +84,7 @@ class GetExecutionHistoryTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
exec_type: str | None = None,
|
||||
@@ -127,7 +127,7 @@ class AnnotateExecutionTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
execution_id: str,
|
||||
@@ -178,7 +178,7 @@ class CreateSkillPayloadTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
payload: dict[str, Any] | list[Any],
|
||||
@@ -208,7 +208,7 @@ class GetSkillPayloadTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
payload_ref: str,
|
||||
@@ -253,7 +253,7 @@ class CreateSkillCandidateTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
skill_key: str,
|
||||
@@ -290,7 +290,7 @@ class ListSkillCandidatesTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
status: str | None = None,
|
||||
@@ -328,7 +328,7 @@ class EvaluateSkillCandidateTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
candidate_id: str,
|
||||
@@ -380,7 +380,7 @@ class PromoteSkillCandidateTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
candidate_id: str,
|
||||
@@ -438,7 +438,7 @@ class ListSkillReleasesTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
skill_key: str | None = None,
|
||||
@@ -474,7 +474,7 @@ class RollbackSkillReleaseTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
release_id: str,
|
||||
@@ -504,7 +504,7 @@ class SyncSkillReleaseTool(NeoSkillToolBase):
|
||||
}
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
release_id: str | None = None,
|
||||
|
||||
@@ -67,7 +67,7 @@ class PythonTool(FunctionTool):
|
||||
description: str = f"Run codes in an IPython shell. Current OS: {_OS_NAME}."
|
||||
parameters: dict = field(default_factory=lambda: param_schema)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self, context: ContextWrapper[AstrAgentContext], code: str, silent: bool = False
|
||||
) -> ToolExecResult:
|
||||
if permission_error := check_admin_permission(context, "Python execution"):
|
||||
@@ -93,7 +93,7 @@ class LocalPythonTool(FunctionTool):
|
||||
|
||||
parameters: dict = field(default_factory=lambda: param_schema)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self, context: ContextWrapper[AstrAgentContext], code: str, silent: bool = False
|
||||
) -> ToolExecResult:
|
||||
if permission_error := check_admin_permission(context, "Python execution"):
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
提供配置元数据的国际化键转换功能
|
||||
"""
|
||||
|
||||
from typing import TypedDict, TypeGuard
|
||||
from typing import Any, TypedDict, TypeGuard
|
||||
|
||||
|
||||
def _is_str_keyed_dict(value: object) -> TypeGuard[dict[str, object]]:
|
||||
@@ -39,7 +39,7 @@ class ConfigMetadataI18n:
|
||||
return f"{group}.{section}.{attr}"
|
||||
|
||||
@staticmethod
|
||||
def convert_to_i18n_keys(metadata: dict[str, object]) -> dict[str, I18nGroup]:
|
||||
def convert_to_i18n_keys(metadata: dict[str, Any]) -> dict[str, I18nGroup]:
|
||||
"""
|
||||
将配置元数据转换为使用国际化键
|
||||
|
||||
|
||||
@@ -180,10 +180,7 @@ class AstrBotCoreLifecycle:
|
||||
"cleanup_loaded_plugins",
|
||||
):
|
||||
try:
|
||||
cleanup_loaded_plugins = getattr(
|
||||
self.plugin_manager,
|
||||
"cleanup_loaded_plugins",
|
||||
)
|
||||
cleanup_loaded_plugins = self.plugin_manager.cleanup_loaded_plugins
|
||||
result = cleanup_loaded_plugins()
|
||||
if inspect.isawaitable(result):
|
||||
await result
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import datetime
|
||||
import json
|
||||
from typing import Any, cast
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
@@ -292,7 +293,9 @@ async def migration_preferences(
|
||||
logger.info(f"迁移全局偏好设置 {key} 成功,值: {value}")
|
||||
|
||||
# 2. umo scope migration
|
||||
session_conversation = sp_v3.get("session_conversation", default={})
|
||||
session_conversation = cast(
|
||||
dict[str, Any], sp_v3.get("session_conversation", default={})
|
||||
)
|
||||
for umo, conversation_id in session_conversation.items():
|
||||
if not umo or not conversation_id:
|
||||
continue
|
||||
@@ -305,7 +308,9 @@ async def migration_preferences(
|
||||
except Exception as e:
|
||||
logger.error(f"迁移会话 {umo} 的对话数据失败: {e}", exc_info=True)
|
||||
|
||||
session_service_config = sp_v3.get("session_service_config", default={})
|
||||
session_service_config = cast(
|
||||
dict[str, Any], sp_v3.get("session_service_config", default={})
|
||||
)
|
||||
for umo, config in session_service_config.items():
|
||||
if not umo or not config:
|
||||
continue
|
||||
@@ -320,7 +325,7 @@ async def migration_preferences(
|
||||
except Exception as e:
|
||||
logger.error(f"迁移会话 {umo} 的服务配置失败: {e}", exc_info=True)
|
||||
|
||||
session_variables = sp_v3.get("session_variables", default={})
|
||||
session_variables = cast(dict[str, Any], sp_v3.get("session_variables", default={}))
|
||||
for umo, variables in session_variables.items():
|
||||
if not umo or not variables:
|
||||
continue
|
||||
@@ -332,7 +337,9 @@ async def migration_preferences(
|
||||
except Exception as e:
|
||||
logger.error(f"迁移会话 {umo} 的变量失败: {e}", exc_info=True)
|
||||
|
||||
session_provider_perf = sp_v3.get("session_provider_perf", default={})
|
||||
session_provider_perf = cast(
|
||||
dict[str, Any], sp_v3.get("session_provider_perf", default={})
|
||||
)
|
||||
for umo, perf in session_provider_perf.items():
|
||||
if not umo or not perf:
|
||||
continue
|
||||
@@ -341,7 +348,8 @@ async def migration_preferences(
|
||||
platform_id = get_platform_id(platform_id_map, session.platform_name)
|
||||
session.platform_id = platform_id
|
||||
|
||||
for provider_type, provider_id in perf.items():
|
||||
perf_dict = cast(dict[str, Any], perf)
|
||||
for provider_type, provider_id in perf_dict.items():
|
||||
await sp.put_async(
|
||||
"umo",
|
||||
str(session),
|
||||
|
||||
@@ -599,7 +599,7 @@ class SQLiteDatabase(BaseDatabase):
|
||||
col(PlatformMessageHistory.created_at) < before,
|
||||
),
|
||||
)
|
||||
return int(result.rowcount or 0)
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
|
||||
async def delete_platform_message_after(
|
||||
self,
|
||||
@@ -617,7 +617,7 @@ class SQLiteDatabase(BaseDatabase):
|
||||
col(PlatformMessageHistory.created_at) > after,
|
||||
),
|
||||
)
|
||||
return int(result.rowcount or 0)
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
|
||||
async def delete_all_platform_message_history(
|
||||
self,
|
||||
@@ -633,7 +633,7 @@ class SQLiteDatabase(BaseDatabase):
|
||||
col(PlatformMessageHistory.user_id) == user_id,
|
||||
),
|
||||
)
|
||||
return int(result.rowcount or 0)
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
|
||||
async def find_platform_message_history_by_idempotency_key(
|
||||
self,
|
||||
@@ -710,7 +710,7 @@ class SQLiteDatabase(BaseDatabase):
|
||||
col(Attachment.attachment_id) == attachment_id
|
||||
)
|
||||
result = cast(CursorResult, await session.execute(query))
|
||||
return result.rowcount > 0
|
||||
return getattr(result, "rowcount", 0) > 0
|
||||
|
||||
async def delete_attachments(self, attachment_ids: list[str]) -> int:
|
||||
"""Delete multiple attachments by their IDs.
|
||||
@@ -725,7 +725,7 @@ class SQLiteDatabase(BaseDatabase):
|
||||
col(Attachment.attachment_id).in_(attachment_ids)
|
||||
)
|
||||
result = cast(CursorResult, await session.execute(query))
|
||||
return result.rowcount
|
||||
return getattr(result, "rowcount", 0)
|
||||
|
||||
async def create_api_key(
|
||||
self,
|
||||
@@ -800,7 +800,7 @@ class SQLiteDatabase(BaseDatabase):
|
||||
.values(revoked_at=datetime.now(timezone.utc))
|
||||
)
|
||||
result = cast(CursorResult, await session.execute(query))
|
||||
return result.rowcount > 0
|
||||
return getattr(result, "rowcount", 0) > 0
|
||||
|
||||
async def delete_api_key(self, key_id: str) -> bool:
|
||||
"""Delete an API key."""
|
||||
@@ -812,7 +812,7 @@ class SQLiteDatabase(BaseDatabase):
|
||||
delete(ApiKey).where(col(ApiKey.key_id) == key_id)
|
||||
),
|
||||
)
|
||||
return result.rowcount > 0
|
||||
return getattr(result, "rowcount", 0) > 0
|
||||
|
||||
async def insert_persona(
|
||||
self,
|
||||
|
||||
@@ -109,7 +109,7 @@ class FaissVecDB(BaseVecDB):
|
||||
async def retrieve(
|
||||
self,
|
||||
query: str,
|
||||
k: int = 5,
|
||||
top_k: int = 5,
|
||||
fetch_k: int = 20,
|
||||
rerank: bool = False,
|
||||
metadata_filters: dict | None = None,
|
||||
@@ -118,7 +118,7 @@ class FaissVecDB(BaseVecDB):
|
||||
|
||||
Args:
|
||||
query (str): 查询文本
|
||||
k (int): 返回的最相似文档的数量
|
||||
top_k (int): 返回的最相似文档的数量
|
||||
fetch_k (int): 在根据 metadata 过滤前从 FAISS 中获取的数量
|
||||
rerank (bool): 是否使用重排序。这需要在实例化时提供 rerank_provider, 如果未提供并且 rerank 为 True, 不会抛出异常。
|
||||
metadata_filters (dict): 元数据过滤器
|
||||
@@ -130,7 +130,7 @@ class FaissVecDB(BaseVecDB):
|
||||
embedding = await self.embedding_provider.get_embedding(query)
|
||||
scores, indices = await self.embedding_storage.search(
|
||||
vector=np.array([embedding]).astype("float32"),
|
||||
k=fetch_k if metadata_filters else k,
|
||||
k=fetch_k if metadata_filters else top_k,
|
||||
)
|
||||
if len(indices[0]) == 0 or indices[0][0] == -1:
|
||||
return []
|
||||
@@ -154,7 +154,7 @@ class FaissVecDB(BaseVecDB):
|
||||
score = scores[0][i]
|
||||
result_docs.append(Result(similarity=float(score), data=fetch_doc))
|
||||
|
||||
top_k_results = result_docs[:k]
|
||||
top_k_results = result_docs[:top_k]
|
||||
|
||||
if rerank and self.rerank_provider:
|
||||
documents = [doc.data["text"] for doc in top_k_results]
|
||||
|
||||
@@ -42,7 +42,6 @@ class InitialLoader:
|
||||
shutdown_event = core_lifecycle.dashboard_shutdown_event
|
||||
if shutdown_event is None:
|
||||
raise RuntimeError("initialize_core must set dashboard_shutdown_event")
|
||||
shutdown_event = cast(asyncio.Event, shutdown_event)
|
||||
|
||||
webui_dir = self.webui_dir
|
||||
|
||||
|
||||
@@ -139,6 +139,7 @@ class KnowledgeBaseManager:
|
||||
"""获取知识库实例"""
|
||||
if kb_id in self.kb_insts:
|
||||
return self.kb_insts[kb_id]
|
||||
return None
|
||||
|
||||
async def get_kb_by_name(self, kb_name: str) -> KBHelper | None:
|
||||
"""通过名称获取知识库实例"""
|
||||
|
||||
@@ -52,10 +52,11 @@ class PDFParser(BaseParser):
|
||||
continue
|
||||
|
||||
resources = page["/Resources"]
|
||||
if not resources or "/XObject" not in resources:
|
||||
xobject_ref = resources.get("/XObject")
|
||||
if not resources or not xobject_ref:
|
||||
continue
|
||||
|
||||
xobjects = resources["/XObject"].get_object()
|
||||
xobjects = xobject_ref.get_object()
|
||||
if not xobjects:
|
||||
continue
|
||||
|
||||
|
||||
@@ -74,6 +74,11 @@ class ComponentType(str, Enum):
|
||||
Json = "Json"
|
||||
Unknown = "Unknown"
|
||||
WechatEmoji = "WechatEmoji" # Wechat 下的 emoji 表情包
|
||||
# Discord-specific component types
|
||||
DiscordEmbed = "DiscordEmbed"
|
||||
DiscordButton = "DiscordButton"
|
||||
DiscordReference = "DiscordReference"
|
||||
DiscordView = "DiscordView"
|
||||
|
||||
|
||||
class BaseMessageComponent(BaseModel):
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from typing import Any, cast
|
||||
|
||||
from astrbot import logger
|
||||
from astrbot.api import sp
|
||||
from astrbot.core.astrbot_config_mgr import AstrBotConfigManager
|
||||
@@ -295,7 +297,7 @@ class PersonaManager:
|
||||
|
||||
# 构建树形结构
|
||||
root_folders = []
|
||||
for folder_id, folder_data in folder_map.items():
|
||||
for _folder_id, folder_data in folder_map.items():
|
||||
parent_id = folder_data["parent_id"]
|
||||
if parent_id is None:
|
||||
root_folders.append(folder_data)
|
||||
@@ -433,10 +435,13 @@ class PersonaManager:
|
||||
user_turn = not user_turn
|
||||
|
||||
try:
|
||||
persona = Personality(
|
||||
**persona_cfg,
|
||||
_begin_dialogs_processed=bd_processed,
|
||||
_mood_imitation_dialogs_processed="", # deprecated
|
||||
persona = cast(
|
||||
Personality,
|
||||
{
|
||||
**persona_cfg,
|
||||
"_begin_dialogs_processed": bd_processed,
|
||||
"_mood_imitation_dialogs_processed": "", # deprecated
|
||||
},
|
||||
)
|
||||
if persona["name"] == self.default_persona:
|
||||
selected_default_persona = persona
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
"""使用此功能应该先 pip install baidu-aip"""
|
||||
|
||||
from typing import TypedDict, TypeGuard
|
||||
from typing import Any, TypedDict, TypeGuard, cast
|
||||
|
||||
from . import ContentSafetyStrategy
|
||||
|
||||
@@ -15,7 +15,8 @@ def _is_violation_list(value: object) -> TypeGuard[list[BaiduAipViolation]]:
|
||||
for item in value:
|
||||
if not isinstance(item, dict):
|
||||
return False
|
||||
message = item.get("msg")
|
||||
raw = cast(dict[str, Any], item)
|
||||
message = raw.get("msg")
|
||||
if message is not None and not isinstance(message, str):
|
||||
return False
|
||||
return True
|
||||
@@ -23,7 +24,7 @@ def _is_violation_list(value: object) -> TypeGuard[list[BaiduAipViolation]]:
|
||||
|
||||
class BaiduAipStrategy(ContentSafetyStrategy):
|
||||
def __init__(self, appid: str, ak: str, sk: str) -> None:
|
||||
from aip import AipContentCensor # type: ignore[unresolved-import]
|
||||
from aip import AipContentCensor
|
||||
|
||||
self.app_id = appid
|
||||
self.api_key = ak
|
||||
@@ -46,7 +47,8 @@ class BaiduAipStrategy(ContentSafetyStrategy):
|
||||
count = len(data)
|
||||
parts = [f"百度审核服务发现 {count} 处违规:\n"]
|
||||
for item in data:
|
||||
message = item.get("msg")
|
||||
raw_item = cast(dict[str, Any], item)
|
||||
message = raw_item.get("msg")
|
||||
if message:
|
||||
parts.append(f"{message};\n")
|
||||
parts.append("\n判断结果:" + conclusion)
|
||||
|
||||
@@ -254,7 +254,7 @@ class AiocqhttpAdapter(Platform):
|
||||
# 如果文本段为空,则跳过
|
||||
continue
|
||||
message_str += current_text
|
||||
a = ComponentTypes[t](text=current_text)
|
||||
a = Plain(text=current_text)
|
||||
abm.message.append(a)
|
||||
|
||||
elif t == "file":
|
||||
|
||||
@@ -760,7 +760,7 @@ class DingtalkPlatformAdapter(Platform):
|
||||
raise KeyboardInterrupt("Graceful shutdown")
|
||||
|
||||
if self.client_.websocket is not None:
|
||||
self.client_.open_connection = monkey_patch_close
|
||||
self.client_.open_connection = monkey_patch_close # type: ignore[assignment]
|
||||
await self.client_.websocket.close(code=1000, reason="Graceful shutdown")
|
||||
if self._shutdown_event is not None:
|
||||
self._shutdown_event.set()
|
||||
|
||||
@@ -2,6 +2,7 @@ import asyncio
|
||||
import re
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import Any, cast
|
||||
|
||||
import anyio
|
||||
|
||||
@@ -233,7 +234,8 @@ class LineMessageEvent(AstrMessageEvent):
|
||||
raw = self.message_obj.raw_message
|
||||
reply_token = ""
|
||||
if isinstance(raw, dict):
|
||||
reply_token = str(raw.get("replyToken") or "")
|
||||
raw_dict = cast(dict[str, Any], raw)
|
||||
reply_token = str(raw_dict.get("replyToken") or "")
|
||||
|
||||
sent = False
|
||||
if reply_token:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import asyncio
|
||||
import random
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import anyio
|
||||
|
||||
@@ -18,7 +18,7 @@ from astrbot.core.platform.astr_message_event import MessageSession
|
||||
from .misskey_api import MisskeyAPI
|
||||
|
||||
try:
|
||||
import magic
|
||||
import magic # type: ignore[assignment]
|
||||
except Exception:
|
||||
magic = None
|
||||
|
||||
@@ -200,7 +200,8 @@ class MisskeyPlatformAdapter(Platform):
|
||||
try:
|
||||
if not isinstance(message.raw_message, dict):
|
||||
message.raw_message = {}
|
||||
message.raw_message["poll"] = poll
|
||||
raw_message_dict = cast(dict, message.raw_message)
|
||||
raw_message_dict["poll"] = poll
|
||||
message.__setattr__("poll", poll)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -543,7 +544,8 @@ class MisskeyPlatformAdapter(Platform):
|
||||
if not r:
|
||||
continue
|
||||
if isinstance(r, dict):
|
||||
url = r.get("fallback_url")
|
||||
r_dict = cast(dict, r)
|
||||
url = r_dict.get("fallback_url")
|
||||
if url:
|
||||
fallback_urls.append(str(url))
|
||||
else:
|
||||
|
||||
@@ -557,7 +557,7 @@ class MisskeyAPI:
|
||||
form.add_field("folderId", str(folder_id))
|
||||
|
||||
try:
|
||||
f = await anyio.to_thread.run_sync(open, file_path, "rb")
|
||||
f = await anyio.to_thread.run_sync(open, file_path, "rb") # type: ignore[unresolved-attribute]
|
||||
except FileNotFoundError as e:
|
||||
logger.error(f"[Misskey API] 本地文件不存在: {file_path}")
|
||||
raise APIError(f"File not found: {file_path}") from e
|
||||
|
||||
@@ -38,7 +38,7 @@ def _patch_qq_botpy_formdata() -> None:
|
||||
from botpy.http import _FormData
|
||||
|
||||
if not hasattr(_FormData, "_is_processed"):
|
||||
setattr(_FormData, "_is_processed", False)
|
||||
_FormData._is_processed = False # type: ignore[invalid-assignment]
|
||||
except Exception:
|
||||
logger.debug("[QQOfficial] Skip botpy FormData patch.")
|
||||
|
||||
|
||||
@@ -182,7 +182,7 @@ class QQOfficialPlatformAdapter(Platform):
|
||||
|
||||
payload: dict[str, Any] = {"content": plain_text, "msg_id": msg_id}
|
||||
ret: Any = None
|
||||
send_helper = SimpleNamespace(bot=self.client)
|
||||
send_helper = cast(Any, SimpleNamespace(bot=self.client))
|
||||
|
||||
if session.message_type == MessageType.GROUP_MESSAGE:
|
||||
scene = self._session_scene.get(session.session_id)
|
||||
|
||||
@@ -191,9 +191,11 @@ class SatoriPlatformAdapter(Platform):
|
||||
|
||||
identify_payload: dict[str, Any] = {
|
||||
"op": 3, # IDENTIFY
|
||||
"body": {
|
||||
"token": str(self.token) if self.token else "", # 字符串
|
||||
},
|
||||
"body": dict[str, Any](
|
||||
{
|
||||
"token": str(self.token) if self.token else "", # 字符串
|
||||
}
|
||||
),
|
||||
}
|
||||
|
||||
# 只有在有序列号时才添加sn字段
|
||||
@@ -235,7 +237,7 @@ class SatoriPlatformAdapter(Platform):
|
||||
except Exception as e:
|
||||
logger.error(f"心跳任务异常: {e}")
|
||||
|
||||
async def handle_message(self, message: str) -> None:
|
||||
async def handle_message(self, message: str | bytes) -> None:
|
||||
try:
|
||||
data = json.loads(message)
|
||||
op = data.get("op")
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from astrbot.api import logger
|
||||
from astrbot.api.event import AstrMessageEvent, MessageChain
|
||||
@@ -47,7 +47,8 @@ class SatoriPlatformEvent(AstrMessageEvent):
|
||||
and message_obj.raw_message
|
||||
and isinstance(message_obj.raw_message, dict)
|
||||
):
|
||||
login = message_obj.raw_message.get("login", {})
|
||||
raw_message = cast(dict[str, Any], message_obj.raw_message)
|
||||
login = raw_message.get("login", {})
|
||||
self.platform = login.get("platform")
|
||||
user = login.get("user", {})
|
||||
self.user_id = user.get("id") if user else None
|
||||
|
||||
@@ -86,6 +86,7 @@ class SlackMessageEvent(AstrMessageEvent):
|
||||
"text": f"文件: <{file_url}|{segment.name or '文件'}>",
|
||||
},
|
||||
}
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def _parse_slack_blocks(
|
||||
|
||||
@@ -228,7 +228,7 @@ class TelegramPlatformAdapter(Platform):
|
||||
|
||||
def collect_commands(self) -> list[BotCommand]:
|
||||
"""从注册的处理器中收集所有指令"""
|
||||
command_dict: dict[str, BotCommand] = {}
|
||||
command_dict: dict[str, str] = {}
|
||||
skip_commands = {"start"}
|
||||
|
||||
for handler_md in star_handlers_registry:
|
||||
|
||||
@@ -313,9 +313,9 @@ class WecomPlatformAdapter(Platform):
|
||||
async def convert_message(self, msg: BaseMessage) -> AstrBotMessage | None:
|
||||
abm = AstrBotMessage()
|
||||
if isinstance(msg, TextMessage):
|
||||
abm.message_str = msg.content
|
||||
abm.message_str = cast(str, msg.content)
|
||||
abm.self_id = str(msg.agent)
|
||||
abm.message = [Plain(msg.content)]
|
||||
abm.message = [Plain(cast(str, msg.content))]
|
||||
abm.type = MessageType.FRIEND_MESSAGE
|
||||
abm.sender = MessageMember(
|
||||
cast(str, msg.source),
|
||||
@@ -328,7 +328,9 @@ class WecomPlatformAdapter(Platform):
|
||||
elif isinstance(msg, ImageMessage):
|
||||
abm.message_str = "[图片]"
|
||||
abm.self_id = str(msg.agent)
|
||||
abm.message = [Image(file=msg.image, url=msg.image)]
|
||||
abm.message = [
|
||||
Image(file=cast(str | None, msg.image), url=cast(str | None, msg.image))
|
||||
]
|
||||
abm.type = MessageType.FRIEND_MESSAGE
|
||||
abm.sender = MessageMember(
|
||||
cast(str, msg.source),
|
||||
@@ -355,7 +357,7 @@ class WecomPlatformAdapter(Platform):
|
||||
except Exception as e:
|
||||
logger.error(f"转换音频失败: {e}。如果没有安装 ffmpeg 请先安装。")
|
||||
path_wav = path
|
||||
return
|
||||
return None
|
||||
|
||||
abm.message_str = ""
|
||||
abm.self_id = str(msg.agent)
|
||||
@@ -371,11 +373,12 @@ class WecomPlatformAdapter(Platform):
|
||||
abm.raw_message = msg
|
||||
else:
|
||||
logger.warning(f"暂未实现的事件: {msg.type}")
|
||||
return
|
||||
return None
|
||||
|
||||
self.agent_id = abm.self_id
|
||||
logger.info(f"abm: {abm}")
|
||||
await self.handle_msg(abm)
|
||||
return abm
|
||||
|
||||
async def convert_wechat_kf_message(self, msg: dict) -> AstrBotMessage | None:
|
||||
msgtype = msg.get("msgtype")
|
||||
@@ -424,13 +427,14 @@ class WecomPlatformAdapter(Platform):
|
||||
except Exception as e:
|
||||
logger.error(f"转换音频失败: {e}。如果没有安装 ffmpeg 请先安装。")
|
||||
path_wav = path
|
||||
return
|
||||
return None
|
||||
|
||||
abm.message = [Record(file=path_wav, url=path_wav)]
|
||||
else:
|
||||
logger.warning(f"未实现的微信客服消息事件: {msg}")
|
||||
return
|
||||
return None
|
||||
await self.handle_msg(abm)
|
||||
return abm
|
||||
|
||||
async def handle_msg(self, message: AstrBotMessage) -> None:
|
||||
message_event = WecomPlatformEvent(
|
||||
|
||||
@@ -8,8 +8,8 @@ import base64
|
||||
import hashlib
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
from collections.abc import Awaitable, Callable, Coroutine
|
||||
from typing import Any, cast
|
||||
|
||||
from astrbot.api import logger
|
||||
from astrbot.api.event import MessageChain
|
||||
@@ -355,6 +355,8 @@ class WecomAIBotAdapter(Platform):
|
||||
except Exception as e:
|
||||
logger.error("处理欢迎消息时发生异常: %s", e)
|
||||
return None
|
||||
return None
|
||||
return None
|
||||
|
||||
async def _process_long_connection_payload(
|
||||
self,
|
||||
@@ -493,7 +495,10 @@ class WecomAIBotAdapter(Platform):
|
||||
_img_url_to_process.append((image_url, image_payload.get("aeskey")))
|
||||
elif msgtype == WecomAIBotConstants.MSG_TYPE_MIXED:
|
||||
# 提取混合消息中的文本内容
|
||||
msg_items = WecomAIBotMessageParser.parse_mixed_message(message_data)
|
||||
msg_items = cast(
|
||||
list[dict[str, Any]],
|
||||
WecomAIBotMessageParser.parse_mixed_message(message_data),
|
||||
)
|
||||
text_parts = []
|
||||
for item in msg_items or []:
|
||||
if item.get("msgtype") == WecomAIBotConstants.MSG_TYPE_TEXT:
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from astrbot.api import logger
|
||||
from astrbot.api.event import AstrMessageEvent, MessageChain
|
||||
@@ -148,6 +148,7 @@ class WecomAIBotMessageEvent(AstrMessageEvent):
|
||||
assert isinstance(raw, dict), (
|
||||
"wecom_ai_bot platform event raw_message should be a dict"
|
||||
)
|
||||
raw = cast(dict[str, Any], raw)
|
||||
stream_id = raw.get("stream_id", self.session_id)
|
||||
pending_response = self.queue_mgr.get_pending_response(stream_id) or {}
|
||||
connection_mode = pending_response.get("callback_params", {}).get(
|
||||
@@ -214,6 +215,7 @@ class WecomAIBotMessageEvent(AstrMessageEvent):
|
||||
assert isinstance(raw, dict), (
|
||||
"wecom_ai_bot platform event raw_message should be a dict"
|
||||
)
|
||||
raw = cast(dict[str, Any], raw)
|
||||
stream_id = raw.get("stream_id", self.session_id)
|
||||
pending_response = self.queue_mgr.get_pending_response(stream_id) or {}
|
||||
connection_mode = pending_response.get("callback_params", {}).get(
|
||||
|
||||
@@ -267,7 +267,7 @@ class WeixinOfficialAccountServer:
|
||||
try:
|
||||
cached = state.get("cached_xml", None)
|
||||
# send one cached each time, if cached is empty after pop, remove the buffer
|
||||
if cached and len(cached) > 0:
|
||||
if cached and isinstance(cached, list) and len(cached) > 0:
|
||||
logger.info(f"wx buffer hit immediately: user={from_user}")
|
||||
cached_xml = cached.pop(0)
|
||||
if len(cached) == 0:
|
||||
@@ -474,7 +474,7 @@ class WeixinOfficialAccountPlatformAdapter(Platform):
|
||||
f"转换音频失败: {e}。如果没有安装 ffmpeg 请先安装。",
|
||||
)
|
||||
path_wav = path
|
||||
return
|
||||
return None
|
||||
|
||||
abm.message_str = ""
|
||||
abm.self_id = str(msg.target)
|
||||
@@ -500,6 +500,7 @@ class WeixinOfficialAccountPlatformAdapter(Platform):
|
||||
}
|
||||
logger.info(f"abm: {abm}")
|
||||
await self.handle_msg(abm)
|
||||
return abm
|
||||
|
||||
async def handle_msg(self, message: AstrBotMessage) -> None:
|
||||
buffer = self.user_buffer.get(message.sender.user_id, None)
|
||||
|
||||
@@ -10,11 +10,16 @@ import dashscope
|
||||
from dashscope.audio.tts_v2 import AudioFormat, SpeechSynthesizer
|
||||
|
||||
try:
|
||||
from dashscope.aigc.multimodal_conversation import MultiModalConversation
|
||||
from dashscope.aigc.multimodal_conversation import (
|
||||
MultiModalConversation,
|
||||
)
|
||||
|
||||
_MultiModalConversationType: type = MultiModalConversation
|
||||
except (
|
||||
ImportError
|
||||
ImportError,
|
||||
): # pragma: no cover - older dashscope versions without Qwen TTS support
|
||||
MultiModalConversation = None
|
||||
MultiModalConversation = None # type: ignore[assignment]
|
||||
_MultiModalConversationType = None # type: ignore[assignment]
|
||||
|
||||
from astrbot.core.provider.entities import ProviderType
|
||||
from astrbot.core.provider.provider import TTSProvider
|
||||
@@ -80,7 +85,7 @@ class ProviderDashscopeTTSAPI(TTSProvider):
|
||||
logging.warning(
|
||||
"No voice specified for Qwen TTS model, using default 'Cherry'.",
|
||||
)
|
||||
return MultiModalConversation.call(**kwargs)
|
||||
return MultiModalConversation.call(**kwargs) # type: ignore[call-arg]
|
||||
|
||||
async def _synthesize_with_qwen_tts(
|
||||
self,
|
||||
|
||||
@@ -4,7 +4,7 @@ import subprocess
|
||||
import uuid
|
||||
|
||||
import anyio
|
||||
import edge_tts
|
||||
import edge_tts # type: ignore[import]
|
||||
|
||||
from astrbot.core import logger
|
||||
from astrbot.core.provider.entities import ProviderType
|
||||
@@ -64,7 +64,7 @@ class ProviderEdgeTTS(TTSProvider):
|
||||
await communicate.save(mp3_path)
|
||||
|
||||
try:
|
||||
from pyffmpeg import FFmpeg
|
||||
from pyffmpeg import FFmpeg # type: ignore[import]
|
||||
|
||||
ff = FFmpeg()
|
||||
ff.convert(input_file=mp3_path, output_file=wav_path)
|
||||
|
||||
@@ -258,7 +258,7 @@ class ProviderGoogleGenAI(Provider):
|
||||
level = types.ThinkingLevel(thinking_level)
|
||||
thinking_config = types.ThinkingConfig()
|
||||
if not hasattr(types.ThinkingConfig, "thinking_level"):
|
||||
setattr(types.ThinkingConfig, "thinking_level", level)
|
||||
types.ThinkingConfig.thinking_level = level
|
||||
else:
|
||||
thinking_config.thinking_level = level
|
||||
|
||||
@@ -775,7 +775,7 @@ class ProviderGoogleGenAI(Provider):
|
||||
model: str | None = None,
|
||||
extra_user_content_parts: list[ContentPart] | None = None,
|
||||
tool_choice: Literal["auto", "required"] = "auto",
|
||||
**kwargs: object,
|
||||
**kwargs,
|
||||
) -> LLMResponse:
|
||||
if contexts is None:
|
||||
contexts = []
|
||||
@@ -797,10 +797,13 @@ class ProviderGoogleGenAI(Provider):
|
||||
# tool calls result
|
||||
if tool_calls_result:
|
||||
if not isinstance(tool_calls_result, list):
|
||||
context_query.extend(tool_calls_result.to_openai_messages())
|
||||
tcr = cast(ToolCallsResult, tool_calls_result)
|
||||
context_query.extend(tcr.to_openai_messages())
|
||||
else:
|
||||
for tcr in tool_calls_result:
|
||||
context_query.extend(tcr.to_openai_messages())
|
||||
context_query.extend(
|
||||
cast(ToolCallsResult, tcr).to_openai_messages()
|
||||
)
|
||||
|
||||
model = model or self.get_model()
|
||||
|
||||
@@ -833,7 +836,7 @@ class ProviderGoogleGenAI(Provider):
|
||||
model: str | None = None,
|
||||
extra_user_content_parts: list[ContentPart] | None = None,
|
||||
tool_choice: Literal["auto", "required"] = "auto",
|
||||
**kwargs: object,
|
||||
**kwargs,
|
||||
) -> AsyncGenerator[LLMResponse, None]:
|
||||
if contexts is None:
|
||||
contexts = []
|
||||
@@ -855,10 +858,13 @@ class ProviderGoogleGenAI(Provider):
|
||||
# tool calls result
|
||||
if tool_calls_result:
|
||||
if not isinstance(tool_calls_result, list):
|
||||
context_query.extend(tool_calls_result.to_openai_messages())
|
||||
tcr = cast(ToolCallsResult, tool_calls_result)
|
||||
context_query.extend(tcr.to_openai_messages())
|
||||
else:
|
||||
for tcr in tool_calls_result:
|
||||
context_query.extend(tcr.to_openai_messages())
|
||||
context_query.extend(
|
||||
cast(ToolCallsResult, tcr).to_openai_messages()
|
||||
)
|
||||
|
||||
model = model or self.get_model()
|
||||
|
||||
@@ -890,7 +896,7 @@ class ProviderGoogleGenAI(Provider):
|
||||
and m.name
|
||||
]
|
||||
except APIError as e:
|
||||
raise Exception(f"获取模型列表失败: {e.message}")
|
||||
raise Exception(f"获取模型列表失败: {e.message}") from e
|
||||
|
||||
def get_current_key(self) -> str:
|
||||
return self.chosen_api_key
|
||||
|
||||
@@ -14,7 +14,7 @@ from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
|
||||
try:
|
||||
import genie_tts as genie
|
||||
except ImportError:
|
||||
genie = None
|
||||
genie = None # type: ignore[assignment]
|
||||
|
||||
|
||||
@register_provider_adapter(
|
||||
@@ -39,12 +39,12 @@ class GenieTTSProvider(TTSProvider):
|
||||
refer_text = provider_config.get("genie_refer_text", "")
|
||||
|
||||
try:
|
||||
genie.load_character(
|
||||
genie.load_character( # type: ignore[attr-defined]
|
||||
character_name=self.character_name,
|
||||
language=language,
|
||||
onnx_model_dir=model_dir,
|
||||
)
|
||||
genie.set_reference_audio(
|
||||
genie.set_reference_audio( # type: ignore[attr-defined]
|
||||
character_name=self.character_name,
|
||||
audio_path=refer_audio_path,
|
||||
audio_text=refer_text,
|
||||
@@ -66,7 +66,7 @@ class GenieTTSProvider(TTSProvider):
|
||||
|
||||
def _generate(save_path: str) -> None:
|
||||
assert genie is not None
|
||||
genie.tts(
|
||||
genie.tts( # type: ignore[attr-defined]
|
||||
character_name=self.character_name,
|
||||
text=text,
|
||||
save_path=save_path,
|
||||
@@ -105,7 +105,7 @@ class GenieTTSProvider(TTSProvider):
|
||||
|
||||
def _generate(save_path: str, t: str) -> None:
|
||||
assert genie is not None
|
||||
genie.tts(
|
||||
genie.tts( # type: ignore[attr-defined]
|
||||
character_name=self.character_name,
|
||||
text=t,
|
||||
save_path=save_path,
|
||||
|
||||
@@ -18,7 +18,5 @@ class ProviderAIHubMix(ProviderOpenAIOfficial):
|
||||
super().__init__(provider_config, provider_settings)
|
||||
# Reference to: https://aihubmix.com/appstore
|
||||
# Use this code can enjoy 10% off prices for AIHubMix API calls.
|
||||
custom_headers = cast(
|
||||
MutableMapping[str, str], getattr(self.client, "_custom_headers")
|
||||
)
|
||||
custom_headers = cast(MutableMapping[str, str], self.client._custom_headers)
|
||||
custom_headers["APP-Code"] = "KRLC5702"
|
||||
|
||||
@@ -17,9 +17,7 @@ class ProviderOpenRouter(ProviderOpenAIOfficial):
|
||||
) -> None:
|
||||
super().__init__(provider_config, provider_settings)
|
||||
# Reference to: https://openrouter.ai/docs/api/reference/overview#headers
|
||||
custom_headers = cast(
|
||||
MutableMapping[str, str], getattr(self.client, "_custom_headers")
|
||||
)
|
||||
custom_headers = cast(MutableMapping[str, str], self.client._custom_headers)
|
||||
custom_headers["HTTP-Referer"] = "https://github.com/AstrBotDevs/AstrBot"
|
||||
custom_headers["X-OpenRouter-Title"] = "AstrBot"
|
||||
custom_headers["X-OpenRouter-Categories"] = "general-chat,personal-agent"
|
||||
|
||||
@@ -9,8 +9,8 @@ from datetime import datetime
|
||||
from typing import Protocol
|
||||
|
||||
import anyio
|
||||
from funasr_onnx import SenseVoiceSmall
|
||||
from funasr_onnx.utils.postprocess_utils import rich_transcription_postprocess
|
||||
from funasr_onnx import SenseVoiceSmall # type: ignore[import]
|
||||
from funasr_onnx.utils.postprocess_utils import rich_transcription_postprocess # type: ignore[import]
|
||||
|
||||
from astrbot.core import logger
|
||||
from astrbot.core.provider.entities import ProviderType
|
||||
|
||||
@@ -120,7 +120,7 @@ class ProviderOpenAIWhisperAPI(STTProvider):
|
||||
|
||||
audio_url = output_path
|
||||
|
||||
file_obj = await anyio.to_thread.run_sync(_open_file_rb, audio_url)
|
||||
file_obj = await anyio.to_thread.run_sync(_open_file_rb, audio_url) # type: ignore[call-arg]
|
||||
result = await self.client.audio.transcriptions.create(
|
||||
model=self.model_name,
|
||||
file=("audio.wav", file_obj),
|
||||
|
||||
@@ -4,7 +4,7 @@ import uuid
|
||||
from typing import Protocol
|
||||
|
||||
import anyio
|
||||
import whisper
|
||||
import whisper # type: ignore[import]
|
||||
|
||||
from astrbot.core import logger
|
||||
from astrbot.core.provider.entities import ProviderType
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
from typing import cast
|
||||
|
||||
from xinference_client.client.restful.async_restful_client import (
|
||||
AsyncClient as Client,
|
||||
)
|
||||
from xinference_client.client.restful.async_restful_client import (
|
||||
AsyncRESTfulModelHandle,
|
||||
AsyncRESTfulRerankModelHandle,
|
||||
)
|
||||
|
||||
@@ -68,7 +71,10 @@ class XinferenceRerankProvider(RerankProvider):
|
||||
return
|
||||
|
||||
if self.model_uid:
|
||||
self.model = await client.get_model(self.model_uid)
|
||||
self.model = cast(
|
||||
AsyncRESTfulRerankModelHandle,
|
||||
await client.get_model(self.model_uid),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize Xinference model: {e}")
|
||||
|
||||
@@ -492,7 +492,7 @@ def _set_filter_aliases(
|
||||
current_aliases: set[str] = getattr(filter_ref, "alias", set())
|
||||
if set(aliases) == current_aliases:
|
||||
return
|
||||
setattr(filter_ref, "alias", set(aliases))
|
||||
filter_ref.alias = set(aliases)
|
||||
if hasattr(filter_ref, "_cmpl_cmd_names"):
|
||||
filter_ref._cmpl_cmd_names = None
|
||||
|
||||
@@ -515,7 +515,7 @@ def _is_command_in_use(
|
||||
|
||||
|
||||
def _descriptor_to_dict(desc: CommandDescriptor) -> dict[str, Any]:
|
||||
result = {
|
||||
result: dict[str, Any] = {
|
||||
"handler_full_name": desc.handler_full_name,
|
||||
"handler_name": desc.handler_name,
|
||||
"plugin": desc.plugin_name,
|
||||
|
||||
@@ -53,6 +53,8 @@ if TYPE_CHECKING:
|
||||
class PlatformManagerProtocol(Protocol):
|
||||
platform_insts: list[Platform]
|
||||
|
||||
def get_insts(self) -> list[Platform]: ...
|
||||
|
||||
|
||||
class StarManagerProtocol(Protocol):
|
||||
async def turn_off_plugin(self, plugin_name: str) -> None: ...
|
||||
@@ -282,6 +284,7 @@ class Context:
|
||||
for star in star_registry:
|
||||
if star.name == star_name:
|
||||
return star
|
||||
return None
|
||||
|
||||
def get_all_stars(self) -> list[StarMetadata]:
|
||||
"""获取当前载入的所有插件 Metadata 的列表"""
|
||||
@@ -461,7 +464,7 @@ class Context:
|
||||
try:
|
||||
session = MessageSesion.from_str(session)
|
||||
except BaseException as e:
|
||||
raise ValueError("不合法的 session 字符串: " + str(e))
|
||||
raise ValueError("不合法的 session 字符串: " + str(e)) from e
|
||||
|
||||
for platform in self.platform_manager.platform_insts:
|
||||
if platform.meta().id == session.platform_name:
|
||||
|
||||
@@ -594,7 +594,7 @@ class PluginManager:
|
||||
|
||||
# 清理工具
|
||||
for tool in list(llm_tools.func_list):
|
||||
if tool.handler_module_path in possible_paths:
|
||||
if getattr(tool, "handler_module_path", None) in possible_paths:
|
||||
llm_tools.func_list.remove(tool)
|
||||
logger.info(f"清理工具: {tool.name}")
|
||||
|
||||
@@ -915,9 +915,9 @@ class PluginManager:
|
||||
|
||||
# 在实例化前注入类属性,保证插件 __init__ 可读取这些值
|
||||
if metadata.star_cls_type:
|
||||
setattr(metadata.star_cls_type, "name", p_name)
|
||||
setattr(metadata.star_cls_type, "author", p_author)
|
||||
setattr(metadata.star_cls_type, "plugin_id", plugin_id)
|
||||
metadata.star_cls_type.name = p_name
|
||||
metadata.star_cls_type.author = p_author
|
||||
metadata.star_cls_type.plugin_id = plugin_id
|
||||
|
||||
if path not in inactivated_plugins:
|
||||
# 只有没有禁用插件时才实例化插件类
|
||||
@@ -937,9 +937,9 @@ class PluginManager:
|
||||
)
|
||||
|
||||
if metadata.star_cls:
|
||||
setattr(metadata.star_cls, "name", p_name)
|
||||
setattr(metadata.star_cls, "author", p_author)
|
||||
setattr(metadata.star_cls, "plugin_id", plugin_id)
|
||||
metadata.star_cls.name = p_name
|
||||
metadata.star_cls.author = p_author
|
||||
metadata.star_cls.plugin_id = plugin_id
|
||||
else:
|
||||
logger.info(f"插件 {metadata.name} 已被禁用。")
|
||||
|
||||
@@ -982,7 +982,7 @@ class PluginManager:
|
||||
ft.handler
|
||||
and ft.handler.__module__ == metadata.module_path
|
||||
):
|
||||
ft.handler_module_path = metadata.module_path
|
||||
ft.handler_module_path = metadata.module_path # type: ignore[union-attr]
|
||||
ft.handler = functools.partial(
|
||||
ft.handler,
|
||||
metadata.star_cls,
|
||||
@@ -1530,7 +1530,7 @@ class PluginManager:
|
||||
# llm_tools 中移除该插件的工具函数绑定
|
||||
to_remove = []
|
||||
for func_tool in llm_tools.func_list:
|
||||
mp = func_tool.handler_module_path
|
||||
mp = getattr(func_tool, "handler_module_path", None)
|
||||
if (
|
||||
mp
|
||||
and mp.startswith(plugin_module_path)
|
||||
@@ -1604,7 +1604,7 @@ class PluginManager:
|
||||
|
||||
# 禁用插件启用的 llm_tool
|
||||
for func_tool in llm_tools.func_list:
|
||||
mp = func_tool.handler_module_path
|
||||
mp = getattr(func_tool, "handler_module_path", None)
|
||||
if (
|
||||
plugin.module_path
|
||||
and mp
|
||||
@@ -1701,7 +1701,7 @@ class PluginManager:
|
||||
|
||||
# 启用插件启用的 llm_tool
|
||||
for func_tool in llm_tools.func_list:
|
||||
mp = func_tool.handler_module_path
|
||||
mp = getattr(func_tool, "handler_module_path", None)
|
||||
if (
|
||||
plugin.module_path
|
||||
and mp
|
||||
|
||||
@@ -80,7 +80,7 @@ async def retrieve_knowledge_base(
|
||||
|
||||
if not kb_ids:
|
||||
logger.info(f"[知识库] 会话 {umo} 已被配置为不使用知识库")
|
||||
return
|
||||
return None
|
||||
|
||||
top_k = session_config.get("top_k", 5)
|
||||
|
||||
@@ -100,7 +100,7 @@ async def retrieve_knowledge_base(
|
||||
)
|
||||
|
||||
if not kb_names:
|
||||
return
|
||||
return None
|
||||
|
||||
logger.debug(f"[知识库] 使用会话级配置,知识库数量: {len(kb_names)}")
|
||||
else:
|
||||
@@ -111,7 +111,7 @@ async def retrieve_knowledge_base(
|
||||
top_k_fusion = config.get("kb_fusion_top_k", 20)
|
||||
|
||||
if not kb_names:
|
||||
return
|
||||
return None
|
||||
|
||||
logger.debug(f"[知识库] 开始检索知识库,数量: {len(kb_names)}, top_k={top_k}")
|
||||
kb_context = await kb_mgr.retrieve(
|
||||
@@ -122,13 +122,14 @@ async def retrieve_knowledge_base(
|
||||
)
|
||||
|
||||
if not kb_context:
|
||||
return
|
||||
return None
|
||||
|
||||
formatted = kb_context.get("context_text", "")
|
||||
if formatted:
|
||||
results = kb_context.get("results", [])
|
||||
logger.debug(f"[知识库] 为会话 {umo} 注入了 {len(results)} 条相关知识块")
|
||||
return formatted
|
||||
return None
|
||||
|
||||
|
||||
KNOWLEDGE_BASE_QUERY_TOOL = KnowledgeBaseQueryTool()
|
||||
|
||||
@@ -26,7 +26,7 @@ class UmopConfigRouter:
|
||||
self.umop_to_conf_id = sp_data
|
||||
|
||||
@staticmethod
|
||||
def _split_umo(umo: str) -> tuple[str, str, str] | None:
|
||||
def _split_umo(umo: str | int | None) -> tuple[str, str, str] | None:
|
||||
"""将 UMO 拆分为 3 个部分,同时保留 session_id 中的 ':'"""
|
||||
if not isinstance(umo, str):
|
||||
return None
|
||||
@@ -43,7 +43,10 @@ class UmopConfigRouter:
|
||||
if p1_ls is None or p2_ls is None:
|
||||
return False # 非法格式
|
||||
|
||||
return all(p == "" or fnmatch.fnmatchcase(t, p) for p, t in zip(p1_ls, p2_ls))
|
||||
return all(
|
||||
p == "" or fnmatch.fnmatchcase(t, p)
|
||||
for p, t in zip(p1_ls, p2_ls, strict=True)
|
||||
)
|
||||
|
||||
def get_conf_id_for_umop(self, umo: str) -> str | None:
|
||||
"""根据 UMO 获取对应的配置文件 ID
|
||||
|
||||
@@ -15,7 +15,6 @@ from astrbot.core import logger, sp
|
||||
from astrbot.core.core_lifecycle import AstrBotCoreLifecycle
|
||||
from astrbot.core.db import BaseDatabase
|
||||
from astrbot.core.platform.message_type import MessageType
|
||||
from astrbot.core.platform_message_history_mgr import PlatformMessageHistoryManager
|
||||
from astrbot.core.platform.sources.webchat.message_parts_helper import (
|
||||
build_webchat_message_parts,
|
||||
create_attachment_part_from_existing_file,
|
||||
@@ -23,6 +22,7 @@ from astrbot.core.platform.sources.webchat.message_parts_helper import (
|
||||
webchat_message_parts_have_content,
|
||||
)
|
||||
from astrbot.core.platform.sources.webchat.webchat_queue_mgr import webchat_queue_mgr
|
||||
from astrbot.core.platform_message_history_mgr import PlatformMessageHistoryManager
|
||||
from astrbot.core.utils.active_event_registry import active_event_registry
|
||||
from astrbot.core.utils.astrbot_path import get_astrbot_data_path
|
||||
from astrbot.core.utils.datetime_utils import to_utc_isoformat
|
||||
|
||||
@@ -100,7 +100,8 @@ class SubAgentRoute(Route):
|
||||
# the handoff (transfer_to_*) tools as their own mounted tools.
|
||||
if isinstance(tool, HandoffTool):
|
||||
continue
|
||||
if tool.handler_module_path == "core.subagent_orchestrator":
|
||||
tool_handler_module_path = getattr(tool, "handler_module_path", None)
|
||||
if tool_handler_module_path == "core.subagent_orchestrator":
|
||||
continue
|
||||
tools_dict.append(
|
||||
{
|
||||
@@ -108,7 +109,7 @@ class SubAgentRoute(Route):
|
||||
"description": tool.description,
|
||||
"parameters": tool.parameters,
|
||||
"active": tool.active,
|
||||
"handler_module_path": tool.handler_module_path,
|
||||
"handler_module_path": tool_handler_module_path,
|
||||
}
|
||||
)
|
||||
return jsonify(Response().ok(data=tools_dict).to_json())
|
||||
|
||||
@@ -443,10 +443,10 @@ class ToolsRoute(Route):
|
||||
elif source == "internal":
|
||||
origin = "internal"
|
||||
origin_name = "AstrBot"
|
||||
elif tool.handler_module_path and star_map.get(
|
||||
tool.handler_module_path
|
||||
elif getattr(tool, "handler_module_path", None) and star_map.get(
|
||||
getattr(tool, "handler_module_path", None)
|
||||
):
|
||||
star = star_map[tool.handler_module_path]
|
||||
star = star_map[getattr(tool, "handler_module_path", None)]
|
||||
origin = "plugin"
|
||||
origin_name = star.name
|
||||
else:
|
||||
|
||||
+136
-103
@@ -2,37 +2,63 @@
|
||||
import ConsoleDisplayer from "@/components/shared/ConsoleDisplayer.vue";
|
||||
import { useModuleI18n } from "@/i18n/composables";
|
||||
import axios from "@/utils/request";
|
||||
import { ref } from "vue";
|
||||
|
||||
const { tm } = useModuleI18n("features/console");
|
||||
|
||||
const autoScrollEnabled = ref(true);
|
||||
const isFullscreen = ref(false);
|
||||
const pipDialog = ref(false);
|
||||
const pipInstallPayload = ref({ package: "", mirror: "" });
|
||||
const loading = ref(false);
|
||||
const status = ref("");
|
||||
const consoleDisplayer = ref();
|
||||
|
||||
function toggleFullscreen() {
|
||||
isFullscreen.value = !isFullscreen.value;
|
||||
}
|
||||
|
||||
async function pipInstall() {
|
||||
loading.value = true;
|
||||
status.value = "";
|
||||
try {
|
||||
const res = await axios.post("/api/console/pip_install", pipInstallPayload.value);
|
||||
if (res.data?.status === "ok") {
|
||||
status.value = tm("pipInstall.success");
|
||||
} else {
|
||||
status.value = res.data?.message || tm("pipInstall.failed");
|
||||
}
|
||||
} catch (err: any) {
|
||||
status.value = err?.response?.data?.message || tm("pipInstall.failed");
|
||||
} finally {
|
||||
loading.value = false;
|
||||
}
|
||||
}
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<div style="height: 100%">
|
||||
<div class="console-toolbar">
|
||||
<div>
|
||||
<h4>{{ tm("title") }}</h4>
|
||||
<v-alert
|
||||
type="info"
|
||||
variant="tonal"
|
||||
density="compact"
|
||||
class="mt-2"
|
||||
style="max-width: 600px"
|
||||
>
|
||||
<div class="console-page" :class="{ 'is-fullscreen': isFullscreen }">
|
||||
<div class="console-topbar">
|
||||
<div class="topbar-left">
|
||||
<div class="topbar-title">{{ tm("title") }}</div>
|
||||
<v-alert type="info" variant="tonal" density="compact" class="mt-2" style="max-width: 600px">
|
||||
{{ tm("debugHint.text") }}
|
||||
</v-alert>
|
||||
</div>
|
||||
<div class="d-flex align-center">
|
||||
<div class="topbar-right">
|
||||
<v-btn
|
||||
:icon="isFullscreen ? 'mdi-fullscreen-exit' : 'mdi-fullscreen'"
|
||||
variant="tonal"
|
||||
size="small"
|
||||
@click="toggleFullscreen"
|
||||
/>
|
||||
<v-switch
|
||||
v-model="autoScrollEnabled"
|
||||
:label="
|
||||
autoScrollEnabled
|
||||
? tm('autoScroll.enabled')
|
||||
: tm('autoScroll.disabled')
|
||||
"
|
||||
:label="autoScrollEnabled ? tm('autoScroll.enabled') : tm('autoScroll.disabled')"
|
||||
hide-details
|
||||
density="compact"
|
||||
color="primary"
|
||||
style="margin-right: 16px"
|
||||
style="margin-right: 8px"
|
||||
/>
|
||||
<v-dialog v-model="pipDialog" width="400">
|
||||
<template #activator="{ props }">
|
||||
@@ -40,34 +66,19 @@ const { tm } = useModuleI18n("features/console");
|
||||
{{ tm("pipInstall.button") }}
|
||||
</v-btn>
|
||||
</template>
|
||||
<v-card>
|
||||
<v-card class="console-dialog-card">
|
||||
<v-card-title>
|
||||
<span class="text-h5">{{ tm("pipInstall.dialogTitle") }}</span>
|
||||
</v-card-title>
|
||||
<v-card-text>
|
||||
<v-text-field
|
||||
v-model="pipInstallPayload.package"
|
||||
:label="tm('pipInstall.packageLabel')"
|
||||
variant="outlined"
|
||||
/>
|
||||
<v-text-field
|
||||
v-model="pipInstallPayload.mirror"
|
||||
:label="tm('pipInstall.mirrorLabel')"
|
||||
variant="outlined"
|
||||
/>
|
||||
<v-text-field v-model="pipInstallPayload.package" :label="tm('pipInstall.packageLabel')" variant="outlined" />
|
||||
<v-text-field v-model="pipInstallPayload.mirror" :label="tm('pipInstall.mirrorLabel')" variant="outlined" />
|
||||
<small>{{ tm("pipInstall.mirrorHint") }}</small>
|
||||
<div>
|
||||
<small>{{ status }}</small>
|
||||
</div>
|
||||
<div><small>{{ status }}</small></div>
|
||||
</v-card-text>
|
||||
<v-card-actions>
|
||||
<v-spacer />
|
||||
<v-btn
|
||||
color="blue-darken-1"
|
||||
variant="text"
|
||||
:loading="loading"
|
||||
@click="pipInstall"
|
||||
>
|
||||
<v-btn color="blue-darken-1" variant="text" :loading="loading" @click="pipInstall">
|
||||
{{ tm("pipInstall.installButton") }}
|
||||
</v-btn>
|
||||
</v-card-actions>
|
||||
@@ -75,85 +86,107 @@ const { tm } = useModuleI18n("features/console");
|
||||
</v-dialog>
|
||||
</div>
|
||||
</div>
|
||||
<ConsoleDisplayer
|
||||
ref="consoleDisplayer"
|
||||
style="height: calc(100vh - 220px)"
|
||||
/>
|
||||
<div class="console-content" :style="isFullscreen ? 'height: calc(100vh - 120px)' : 'height: calc(100vh - 220px)'">
|
||||
<ConsoleDisplayer ref="consoleDisplayer" style="height: 100%" />
|
||||
</div>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<script lang="ts">
|
||||
export default {
|
||||
name: "ConsolePage",
|
||||
components: {
|
||||
ConsoleDisplayer,
|
||||
},
|
||||
data() {
|
||||
return {
|
||||
autoScrollEnabled: true,
|
||||
pipDialog: false,
|
||||
pipInstallPayload: {
|
||||
package: "",
|
||||
mirror: "",
|
||||
},
|
||||
loading: false,
|
||||
status: "",
|
||||
};
|
||||
},
|
||||
watch: {
|
||||
autoScrollEnabled(val) {
|
||||
if (this.$refs.consoleDisplayer) {
|
||||
this.$refs.consoleDisplayer.autoScroll = val;
|
||||
}
|
||||
},
|
||||
},
|
||||
methods: {
|
||||
pipInstall() {
|
||||
this.loading = true;
|
||||
axios
|
||||
.post("/api/update/pip-install", this.pipInstallPayload)
|
||||
.then((res) => {
|
||||
this.status = res.data.message;
|
||||
setTimeout(() => {
|
||||
this.status = "";
|
||||
this.pipDialog = false;
|
||||
}, 2000);
|
||||
})
|
||||
.catch((err) => {
|
||||
this.status = err.response.data.message;
|
||||
})
|
||||
.finally(() => {
|
||||
this.loading = false;
|
||||
});
|
||||
},
|
||||
},
|
||||
components: { ConsoleDisplayer },
|
||||
};
|
||||
</script>
|
||||
|
||||
<style>
|
||||
@keyframes fadeIn {
|
||||
from {
|
||||
opacity: 0;
|
||||
}
|
||||
<style scoped>
|
||||
.console-page {
|
||||
--console-page-bg: transparent;
|
||||
--console-panel-bg: rgba(var(--v-theme-surface), 0.78);
|
||||
--console-card-bg: rgba(var(--v-theme-surface), 0.9);
|
||||
--console-primary: rgb(var(--v-theme-primary));
|
||||
--console-primary-soft: rgba(var(--v-theme-primary), 0.08);
|
||||
--console-border: rgba(var(--v-theme-borderLight), 0.22);
|
||||
--console-border-strong: rgba(var(--v-theme-borderLight), 0.4);
|
||||
--console-text: rgba(var(--v-theme-on-surface), 0.92);
|
||||
--console-muted: rgba(var(--v-theme-on-surface), 0.7);
|
||||
--console-shadow: 0 10px 24px rgba(15, 23, 42, 0.08);
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
height: 100%;
|
||||
min-height: 0;
|
||||
position: relative;
|
||||
z-index: 1;
|
||||
isolation: isolate;
|
||||
gap: 16px;
|
||||
padding: 16px;
|
||||
}
|
||||
|
||||
to {
|
||||
opacity: 1;
|
||||
}
|
||||
.console-page.is-fullscreen {
|
||||
position: fixed;
|
||||
top: 0;
|
||||
left: 0;
|
||||
right: 0;
|
||||
bottom: 0;
|
||||
z-index: 9999;
|
||||
padding: 16px;
|
||||
background: var(--console-panel-bg);
|
||||
}
|
||||
|
||||
.fade-in {
|
||||
animation: fadeIn 0.2s ease-in-out;
|
||||
:global(.v-theme--bluebusinessdarktheme) .console-page {
|
||||
--console-panel-bg: rgba(var(--v-theme-surface), 0.72);
|
||||
--console-card-bg: rgba(var(--v-theme-surface-variant), 0.74);
|
||||
--console-border: rgba(var(--v-theme-borderLight), 0.46);
|
||||
--console-shadow: none;
|
||||
}
|
||||
|
||||
.console-toolbar {
|
||||
background: rgba(15, 15, 22, 0.55);
|
||||
backdrop-filter: blur(16px);
|
||||
padding: 12px 16px;
|
||||
margin-bottom: 16px;
|
||||
border-radius: 12px;
|
||||
border: 1px solid rgba(0, 242, 255, 0.06);
|
||||
.console-topbar {
|
||||
display: flex;
|
||||
flex-direction: row;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
gap: 16px;
|
||||
padding: 16px 20px;
|
||||
background: var(--console-panel-bg);
|
||||
border: 1px solid var(--console-border);
|
||||
border-radius: 12px;
|
||||
backdrop-filter: blur(16px);
|
||||
box-shadow: var(--console-shadow);
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
.topbar-left { display: flex; flex-direction: column; gap: 4px; }
|
||||
|
||||
.topbar-title {
|
||||
font-size: 18px;
|
||||
font-weight: 700;
|
||||
color: var(--console-primary) !important;
|
||||
-webkit-text-fill-color: var(--console-primary);
|
||||
}
|
||||
|
||||
.topbar-right { display: flex; align-items: center; gap: 10px; }
|
||||
|
||||
.console-content {
|
||||
flex: 1;
|
||||
min-height: 0;
|
||||
overflow: hidden;
|
||||
position: relative;
|
||||
z-index: 1;
|
||||
background: var(--console-panel-bg);
|
||||
border: 1px solid var(--console-border);
|
||||
border-radius: 12px;
|
||||
backdrop-filter: blur(16px);
|
||||
box-shadow: var(--console-shadow);
|
||||
}
|
||||
|
||||
.console-dialog-card {
|
||||
border: 1px solid var(--console-border);
|
||||
border-radius: 18px;
|
||||
background: var(--console-panel-bg);
|
||||
}
|
||||
|
||||
@media (max-width: 900px) {
|
||||
.console-topbar { flex-direction: column; align-items: flex-start; }
|
||||
.topbar-right { width: 100%; }
|
||||
.console-page { gap: 12px; padding: 12px; }
|
||||
}
|
||||
</style>
|
||||
|
||||
@@ -822,64 +822,85 @@ export default {
|
||||
|
||||
<style scoped>
|
||||
.platform-page {
|
||||
padding: 20px;
|
||||
padding-top: 8px;
|
||||
padding-bottom: 40px;
|
||||
}
|
||||
|
||||
.webhook-info {
|
||||
margin-top: 4px;
|
||||
}
|
||||
|
||||
.webhook-chip {
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.platform-status-row {
|
||||
--platform-page-bg: transparent;
|
||||
--platform-panel-bg: rgba(var(--v-theme-surface), 0.78);
|
||||
--platform-card-bg: rgba(var(--v-theme-surface), 0.9);
|
||||
--platform-primary: rgb(var(--v-theme-primary));
|
||||
--platform-primary-soft: rgba(var(--v-theme-primary), 0.08);
|
||||
--platform-border: rgba(var(--v-theme-borderLight), 0.22);
|
||||
--platform-border-strong: rgba(var(--v-theme-borderLight), 0.4);
|
||||
--platform-text: rgba(var(--v-theme-on-surface), 0.92);
|
||||
--platform-muted: rgba(var(--v-theme-on-surface), 0.7);
|
||||
--platform-shadow: 0 10px 24px rgba(15, 23, 42, 0.08);
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
height: 100%;
|
||||
min-height: 0;
|
||||
gap: 16px;
|
||||
padding: 16px;
|
||||
}
|
||||
|
||||
:global(.v-theme--bluebusinessdarktheme) .platform-page {
|
||||
--platform-panel-bg: rgba(var(--v-theme-surface), 0.72);
|
||||
--platform-card-bg: rgba(var(--v-theme-surface-variant), 0.74);
|
||||
--platform-border: rgba(var(--v-theme-borderLight), 0.46);
|
||||
--platform-shadow: none;
|
||||
}
|
||||
|
||||
.platform-topbar {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
flex-wrap: wrap;
|
||||
gap: 4px;
|
||||
gap: 16px;
|
||||
padding: 16px 20px;
|
||||
background: var(--platform-panel-bg);
|
||||
border: 1px solid var(--platform-border);
|
||||
border-radius: 12px;
|
||||
backdrop-filter: blur(16px);
|
||||
box-shadow: var(--platform-shadow);
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
.status-chip {
|
||||
font-size: 12px;
|
||||
.topbar-left { display: flex; flex-direction: column; gap: 4px; }
|
||||
|
||||
.topbar-title {
|
||||
font-size: clamp(24px, 3vw, 32px);
|
||||
font-weight: 700;
|
||||
letter-spacing: -0.02em;
|
||||
color: var(--platform-primary) !important;
|
||||
-webkit-text-fill-color: var(--platform-primary);
|
||||
}
|
||||
|
||||
.error-chip {
|
||||
cursor: pointer;
|
||||
font-size: 12px;
|
||||
.topbar-desc {
|
||||
font-size: 14px;
|
||||
color: var(--platform-muted);
|
||||
}
|
||||
|
||||
.error-details {
|
||||
margin-top: 8px;
|
||||
.topbar-right { display: flex; align-items: center; gap: 10px; }
|
||||
|
||||
.platform-content {
|
||||
background: var(--platform-panel-bg);
|
||||
border: 1px solid var(--platform-border);
|
||||
border-radius: 12px;
|
||||
backdrop-filter: blur(16px);
|
||||
box-shadow: var(--platform-shadow);
|
||||
}
|
||||
|
||||
.error-message {
|
||||
word-break: break-word;
|
||||
.platform-card {
|
||||
background: var(--platform-card-bg);
|
||||
border: 1px solid var(--platform-border);
|
||||
border-radius: 12px;
|
||||
}
|
||||
|
||||
.traceback-box {
|
||||
background-color: #1e1e1e;
|
||||
color: #d4d4d4;
|
||||
padding: 12px;
|
||||
border-radius: 8px;
|
||||
font-size: 12px;
|
||||
line-height: 1.5;
|
||||
overflow-x: auto;
|
||||
white-space: pre-wrap;
|
||||
word-break: break-word;
|
||||
max-height: 300px;
|
||||
overflow-y: auto;
|
||||
.platform-dialog-card {
|
||||
border: 1px solid var(--platform-border);
|
||||
border-radius: 18px;
|
||||
background: var(--platform-panel-bg);
|
||||
}
|
||||
|
||||
.platform-qr-chip {
|
||||
margin-top: 4px;
|
||||
}
|
||||
|
||||
.platform-qr-status {
|
||||
font-size: 13px;
|
||||
margin-bottom: 10px;
|
||||
color: rgba(0, 0, 0, 0.7);
|
||||
@media (max-width: 900px) {
|
||||
.platform-topbar { flex-direction: column; align-items: flex-start; }
|
||||
.topbar-right { width: 100%; }
|
||||
.platform-page { gap: 12px; padding: 12px; }
|
||||
}
|
||||
</style>
|
||||
|
||||
@@ -120,6 +120,7 @@ select = [
|
||||
"F", # Pyflakes
|
||||
"W", # pycodestyle warnings
|
||||
"E", # pycodestyle errors
|
||||
"B", # bugbear
|
||||
"I", # isort
|
||||
"UP", # pyupgrade
|
||||
"ASYNC", # flake8-async
|
||||
|
||||
@@ -75,7 +75,7 @@ class TestContextManager:
|
||||
"""Test initialization with LLM-based compression."""
|
||||
mock_provider = MockProvider()
|
||||
config = ContextConfig(
|
||||
llm_compress_provider=mock_provider,
|
||||
llm_compress_provider=mock_provider, # type: ignore[arg-type]
|
||||
llm_compress_keep_recent=5,
|
||||
llm_compress_instruction="Summarize the conversation",
|
||||
)
|
||||
@@ -560,7 +560,7 @@ class TestContextManager:
|
||||
manager = ContextManager(config)
|
||||
|
||||
# Verify the default threshold is 0.82
|
||||
assert manager.compressor.compression_threshold == 0.82
|
||||
assert manager.compressor.compression_threshold == 0.82 # type: ignore[attr-defined]
|
||||
|
||||
# Test threshold logic
|
||||
messages = [self.create_message("user", "x" * 81)] # ~24 tokens
|
||||
@@ -665,7 +665,7 @@ class TestContextManager:
|
||||
"""Test LLM compression using MockProvider."""
|
||||
mock_provider = MockProvider()
|
||||
config = ContextConfig(
|
||||
llm_compress_provider=mock_provider,
|
||||
llm_compress_provider=mock_provider, # type: ignore[arg-type]
|
||||
llm_compress_keep_recent=3,
|
||||
llm_compress_instruction="请总结对话内容",
|
||||
max_context_tokens=100,
|
||||
|
||||
@@ -17,8 +17,8 @@ from astrbot.core.agent.message import (
|
||||
counter = EstimateTokenCounter()
|
||||
|
||||
|
||||
def _msg(role: str, content) -> Message:
|
||||
return Message(role=role, content=content)
|
||||
def _msg(role: str, content: str | list) -> Message:
|
||||
return Message(role=role, content=content) # type: ignore[arg-type]
|
||||
|
||||
|
||||
class TestTextCounting:
|
||||
|
||||
@@ -9,7 +9,7 @@ class TestContextTruncator:
|
||||
|
||||
def create_message(self, role: str, content: str = "test content") -> Message:
|
||||
"""Helper to create a simple test message."""
|
||||
return Message(role=role, content=content)
|
||||
return Message(role=role, content=content) # type: ignore[arg-type]
|
||||
|
||||
def create_messages(
|
||||
self, count: int, include_system: bool = False
|
||||
|
||||
Vendored
+5
-6
@@ -5,7 +5,6 @@ from dataclasses import dataclass, field
|
||||
|
||||
from astrbot.core.provider.entities import (
|
||||
ProviderRequest,
|
||||
ProviderResponse,
|
||||
ProviderMeta,
|
||||
ProviderType,
|
||||
LLMResponse,
|
||||
@@ -62,7 +61,7 @@ class MockChatCompletionProvider:
|
||||
self.response_index += 1
|
||||
|
||||
# Stream word by word
|
||||
for word in text.split():
|
||||
for word in text.split(): # type: ignore[union-attr]
|
||||
yield LLMResponse(
|
||||
role="assistant",
|
||||
completion_text=word + " ",
|
||||
@@ -98,14 +97,14 @@ class MockToolCallProvider(MockChatCompletionProvider):
|
||||
return LLMResponse(
|
||||
role="assistant",
|
||||
completion_text="",
|
||||
tool_calls=[tool_call],
|
||||
tool_calls=[tool_call], # type: ignore[call-arg]
|
||||
)
|
||||
|
||||
# Default tool call
|
||||
return LLMResponse(
|
||||
role="assistant",
|
||||
completion_text="",
|
||||
tool_calls=[
|
||||
tool_calls=[ # type: ignore[call-arg]
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
@@ -129,7 +128,7 @@ class MockErrorProvider(MockChatCompletionProvider):
|
||||
"""Raise mock error."""
|
||||
raise RuntimeError(self.error_message)
|
||||
|
||||
async def stream_chat(self, request: ProviderRequest) -> AsyncGenerator[LLMResponse, None]:
|
||||
async def stream_chat(self, request: ProviderRequest) -> AsyncGenerator[LLMResponse, None]: # type: ignore[method-assign]
|
||||
"""Raise mock error."""
|
||||
raise RuntimeError(self.error_message)
|
||||
|
||||
@@ -183,5 +182,5 @@ def create_mock_llm_response(
|
||||
return LLMResponse(
|
||||
role="assistant",
|
||||
completion_text=text,
|
||||
tool_calls=tool_calls,
|
||||
tool_calls=tool_calls, # type: ignore[call-arg]
|
||||
)
|
||||
|
||||
@@ -95,7 +95,7 @@ class TestAstrbotAcpClient:
|
||||
|
||||
# Start server using loop.create_unix_server
|
||||
loop = asyncio.get_running_loop()
|
||||
server = await loop.create_unix_server(echo_handler, path=socket_path)
|
||||
server = await loop.create_unix_server(echo_handler, path=socket_path) # type: ignore[arg-type]
|
||||
|
||||
try:
|
||||
# Connect client
|
||||
|
||||
@@ -1,138 +0,0 @@
|
||||
"""
|
||||
LSP Integration Tests.
|
||||
|
||||
Tests the LSP client against a real LSP server fixture.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import anyio
|
||||
import pytest
|
||||
from anyio.lowlevel import checkpoint
|
||||
|
||||
from astrbot._internal.protocols.lsp.client import AstrbotLspClient
|
||||
|
||||
TEST_DIR = Path(__file__).resolve().parent
|
||||
SERVER_PATH = TEST_DIR / "fixtures" / "echo_lsp_server.py"
|
||||
HANGING_SERVER_PATH = TEST_DIR / "fixtures" / "hanging_lsp_server.py"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_lsp_client_initialization():
|
||||
"""Test LSP client can be initialized."""
|
||||
client = AstrbotLspClient()
|
||||
assert client is not None
|
||||
assert not client.connected
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_lsp_client_connect_to_echo_server():
|
||||
"""Test LSP client can connect to echo LSP server."""
|
||||
client = AstrbotLspClient()
|
||||
|
||||
await client.connect_to_server(
|
||||
command=[sys.executable, str(SERVER_PATH)],
|
||||
workspace_uri="file:///tmp",
|
||||
)
|
||||
|
||||
try:
|
||||
assert client.connected
|
||||
finally:
|
||||
await client.shutdown()
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_lsp_client_send_request():
|
||||
"""Test LSP client can send a request and receive response."""
|
||||
client = AstrbotLspClient()
|
||||
|
||||
await client.connect_to_server(
|
||||
command=[sys.executable, str(SERVER_PATH)],
|
||||
workspace_uri="file:///tmp",
|
||||
)
|
||||
|
||||
try:
|
||||
assert client.connected
|
||||
|
||||
# Send a custom request - echo server will echo it back
|
||||
result = await client.send_request("custom/echo", {"message": "test"})
|
||||
assert result is not None
|
||||
|
||||
finally:
|
||||
await client.shutdown()
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_lsp_client_send_notification():
|
||||
"""Test LSP client can send a notification (no response)."""
|
||||
client = AstrbotLspClient()
|
||||
|
||||
await client.connect_to_server(
|
||||
command=[sys.executable, str(SERVER_PATH)],
|
||||
workspace_uri="file:///tmp",
|
||||
)
|
||||
|
||||
try:
|
||||
assert client.connected
|
||||
|
||||
# Send a notification (should not raise)
|
||||
await client.send_notification("custom/notify", {"data": "test"})
|
||||
|
||||
finally:
|
||||
await client.shutdown()
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_lsp_client_send_request_not_connected():
|
||||
"""Test LSP client raises RuntimeError when sending request while not connected."""
|
||||
client = AstrbotLspClient()
|
||||
with pytest.raises(RuntimeError, match="LSP client not connected"):
|
||||
await client.send_request("test", {})
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_lsp_client_send_notification_not_connected():
|
||||
"""Test LSP client raises RuntimeError when sending notification while not connected."""
|
||||
client = AstrbotLspClient()
|
||||
with pytest.raises(RuntimeError, match="LSP client not connected"):
|
||||
await client.send_notification("test", {})
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_lsp_client_connect_does_not_corrupt_anyio_cancel_scope():
|
||||
"""Test connect/shutdown can run inside fail_after scopes without scope corruption."""
|
||||
client = AstrbotLspClient()
|
||||
|
||||
with anyio.fail_after(5):
|
||||
await client.connect_to_server(
|
||||
command=[sys.executable, str(SERVER_PATH)],
|
||||
workspace_uri="file:///tmp",
|
||||
)
|
||||
|
||||
try:
|
||||
assert client.connected
|
||||
finally:
|
||||
with anyio.fail_after(5):
|
||||
await client.shutdown()
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_lsp_client_connect_timeout_does_not_corrupt_anyio_cancel_scope():
|
||||
"""Test timeout-driven cancellation leaves later fail_after scopes usable."""
|
||||
client = AstrbotLspClient()
|
||||
|
||||
with pytest.raises(TimeoutError):
|
||||
with anyio.fail_after(0.1):
|
||||
await client.connect_to_server(
|
||||
command=[sys.executable, str(HANGING_SERVER_PATH)],
|
||||
workspace_uri="file:///tmp",
|
||||
)
|
||||
|
||||
with anyio.fail_after(5):
|
||||
await client.shutdown()
|
||||
|
||||
with anyio.fail_after(1):
|
||||
await checkpoint()
|
||||
@@ -49,7 +49,7 @@ async def test_mcp_echo_server_connection():
|
||||
}
|
||||
|
||||
# Connect to the echo server
|
||||
await client.connect_to_server(config, "echo-test")
|
||||
await client.connect_to_server(config, "echo-test") # type: ignore[arg-type]
|
||||
|
||||
try:
|
||||
# Verify connected
|
||||
@@ -77,7 +77,7 @@ async def test_mcp_list_tools():
|
||||
"cwd": test_dir
|
||||
}
|
||||
|
||||
await client.connect_to_server(config, "echo-test")
|
||||
await client.connect_to_server(config, "echo-test") # type: ignore[arg-type]
|
||||
|
||||
try:
|
||||
assert client.connected
|
||||
@@ -114,7 +114,7 @@ async def test_mcp_call_echo_tool():
|
||||
"cwd": test_dir
|
||||
}
|
||||
|
||||
await client.connect_to_server(config, "echo-test")
|
||||
await client.connect_to_server(config, "echo-test") # type: ignore[arg-type]
|
||||
|
||||
try:
|
||||
assert client.connected
|
||||
|
||||
@@ -36,9 +36,9 @@ def test_anthropic_provider_injects_custom_headers_into_http_client(monkeypatch)
|
||||
"User-Agent": "custom-agent/1.0",
|
||||
"X-Test-Header": "123",
|
||||
}
|
||||
assert isinstance(provider.client.kwargs["http_client"], httpx.AsyncClient)
|
||||
assert provider.client.kwargs["http_client"].headers["User-Agent"] == "custom-agent/1.0"
|
||||
assert provider.client.kwargs["http_client"].headers["X-Test-Header"] == "123"
|
||||
assert isinstance(provider.client.kwargs["http_client"], httpx.AsyncClient) # type: ignore[attr-defined]
|
||||
assert provider.client.kwargs["http_client"].headers["User-Agent"] == "custom-agent/1.0" # type: ignore[attr-defined]
|
||||
assert provider.client.kwargs["http_client"].headers["X-Test-Header"] == "123" # type: ignore[attr-defined]
|
||||
|
||||
|
||||
def test_kimi_code_provider_sets_defaults_and_preserves_custom_headers(monkeypatch):
|
||||
@@ -60,10 +60,10 @@ def test_kimi_code_provider_sets_defaults_and_preserves_custom_headers(monkeypat
|
||||
"User-Agent": kimi_code_source.KIMI_CODE_USER_AGENT,
|
||||
"X-Trace-Id": "trace-1",
|
||||
}
|
||||
assert provider.client.kwargs["http_client"].headers["User-Agent"] == (
|
||||
assert provider.client.kwargs["http_client"].headers["User-Agent"] == ( # type: ignore[attr-defined]
|
||||
kimi_code_source.KIMI_CODE_USER_AGENT
|
||||
)
|
||||
assert provider.client.kwargs["http_client"].headers["X-Trace-Id"] == "trace-1"
|
||||
assert provider.client.kwargs["http_client"].headers["X-Trace-Id"] == "trace-1" # type: ignore[attr-defined]
|
||||
|
||||
|
||||
def test_kimi_code_provider_restores_required_user_agent_when_blank(monkeypatch):
|
||||
|
||||
@@ -23,12 +23,12 @@ def _get_open_api_route(app: Quart):
|
||||
(
|
||||
item
|
||||
for item in app.url_map.iter_rules()
|
||||
if item.rule == "/api/v1/chat" and "POST" in item.methods
|
||||
if item.rule == "/api/v1/chat" and "POST" in item.methods # type: ignore[operator]
|
||||
),
|
||||
None,
|
||||
)
|
||||
assert rule is not None
|
||||
return app.view_functions[rule.endpoint].__self__
|
||||
return app.view_functions[rule.endpoint].__self__ # type: ignore[attr-defined]
|
||||
|
||||
|
||||
async def _create_api_key(
|
||||
|
||||
@@ -676,17 +676,17 @@ class TestAstrBotImporter:
|
||||
zf.writestr("databases/main_db.json", json.dumps(main_data))
|
||||
|
||||
importer = AstrBotImporter(main_db=mock_main_db)
|
||||
importer._clear_main_db = AsyncMock(
|
||||
importer._clear_main_db = AsyncMock( # type: ignore[method-assign]
|
||||
side_effect=DatabaseClearError("清空表 platform_stats 失败: db locked")
|
||||
)
|
||||
importer._import_main_database = AsyncMock(return_value={})
|
||||
importer._import_main_database = AsyncMock(return_value={}) # type: ignore[method-assign]
|
||||
|
||||
result = await importer.import_all(str(zip_path), mode="replace")
|
||||
|
||||
assert result.success is False
|
||||
assert any("清空主数据库失败" in err for err in result.errors)
|
||||
assert any("清空表 platform_stats 失败" in err for err in result.errors)
|
||||
importer._import_main_database.assert_not_awaited()
|
||||
importer._import_main_database.assert_not_awaited() # type: ignore[method-assign]
|
||||
|
||||
|
||||
class TestSecureFilename:
|
||||
|
||||
@@ -517,7 +517,7 @@ class TestSubagentHandoffTools:
|
||||
with patch("astrbot.core.computer.computer_client.session_booter", {}):
|
||||
tools = FunctionToolExecutor._get_runtime_computer_tools(
|
||||
"sandbox",
|
||||
session_id=None,
|
||||
session_id=None, # type: ignore[arg-type]
|
||||
sandbox_cfg={"booter": "shipyard_neo"},
|
||||
)
|
||||
assert "astrbot_create_skill_candidate" in tools
|
||||
@@ -531,7 +531,7 @@ class TestSubagentHandoffTools:
|
||||
with patch("astrbot.core.computer.computer_client.session_booter", {}):
|
||||
tools = FunctionToolExecutor._get_runtime_computer_tools(
|
||||
"sandbox",
|
||||
session_id=None,
|
||||
session_id=None, # type: ignore[arg-type]
|
||||
sandbox_cfg={"booter": "shipyard"},
|
||||
)
|
||||
assert len(tools) == 0
|
||||
@@ -543,7 +543,7 @@ class TestSubagentHandoffTools:
|
||||
pytest.skip("circular import")
|
||||
tools = FunctionToolExecutor._get_runtime_computer_tools(
|
||||
"sandbox",
|
||||
session_id=None,
|
||||
session_id=None, # type: ignore[arg-type]
|
||||
sandbox_cfg={},
|
||||
)
|
||||
assert "astrbot_create_skill_candidate" in tools
|
||||
@@ -556,7 +556,7 @@ class TestSubagentHandoffTools:
|
||||
pytest.skip("circular import")
|
||||
tools = FunctionToolExecutor._get_runtime_computer_tools(
|
||||
"local",
|
||||
session_id=None,
|
||||
session_id=None, # type: ignore[arg-type]
|
||||
sandbox_cfg={},
|
||||
)
|
||||
assert len(tools) == 2
|
||||
|
||||
@@ -0,0 +1,205 @@
|
||||
"""
|
||||
Code Quality Test: Type Safety Scoring
|
||||
|
||||
Scores the codebase based on type safety patterns:
|
||||
- cast(Any, ...) usage (too many casts indicate type safety issues)
|
||||
- # type: ignore comments (越多表示类型问题越多)
|
||||
|
||||
Perfect score: 100
|
||||
"""
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
ASTRBOT_ROOT = Path(__file__).parent.parent / "astrbot"
|
||||
|
||||
|
||||
def _scan_for_patterns() -> dict[str, int | set[str] | list[tuple[str, str]]]:
|
||||
"""Scan astrbot source for type-unsafe patterns."""
|
||||
counts: dict[str, int | set[str] | list[tuple[str, str]]] = {
|
||||
"cast_any": 0,
|
||||
"type_ignore": 0,
|
||||
"bare_except": 0,
|
||||
"duplicate_blocks": 0,
|
||||
"cast_any_files": set(),
|
||||
"type_ignore_files": set(),
|
||||
"bare_except_files": set(),
|
||||
"dup_files": [],
|
||||
}
|
||||
|
||||
# Patterns to detect
|
||||
cast_any_re = re.compile(r"cast\s*\(\s*Any\s*,", re.MULTILINE)
|
||||
type_ignore_re = re.compile(r"#\s*type:\s*ignore", re.MULTILINE)
|
||||
bare_except_re = re.compile(r"except\s*(?:Exception)?\s*:\s*$", re.MULTILINE | re.MULTILINE)
|
||||
|
||||
# --- Phase 1: cast / type:ignore / bare_except ---
|
||||
for py_file in ASTRBOT_ROOT.rglob("*.py"):
|
||||
try:
|
||||
content = py_file.read_text(encoding="utf-8")
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
if py_file.suffix == ".pyi":
|
||||
continue
|
||||
|
||||
rel = str(py_file.relative_to(ASTRBOT_ROOT.parent))
|
||||
|
||||
cast_matches = cast_any_re.findall(content)
|
||||
ignore_matches = type_ignore_re.findall(content)
|
||||
# Only count bare `except:` (no specific exception type), not `except Exception:`
|
||||
bare_except_matches = bare_except_re.findall(content)
|
||||
|
||||
if cast_matches:
|
||||
counts["cast_any"] += len(cast_matches)
|
||||
counts["cast_any_files"].add(rel)
|
||||
|
||||
if ignore_matches:
|
||||
counts["type_ignore"] += len(ignore_matches)
|
||||
counts["type_ignore_files"].add(rel)
|
||||
|
||||
if bare_except_matches:
|
||||
counts["bare_except"] += len(bare_except_matches)
|
||||
counts["bare_except_files"].add(rel)
|
||||
|
||||
# --- Phase 2: duplicate code blocks (5+ identical lines) ---
|
||||
dup_blocks: dict[str, list[tuple[str, int]]] = {}
|
||||
for py_file in ASTRBOT_ROOT.rglob("*.py"):
|
||||
if py_file.suffix == ".pyi":
|
||||
continue
|
||||
try:
|
||||
lines = py_file.read_text(encoding="utf-8").splitlines()
|
||||
except Exception:
|
||||
continue
|
||||
rel = str(py_file.relative_to(ASTRBOT_ROOT.parent))
|
||||
# Strip and skip empty/1-char lines; use 5-line window
|
||||
cleaned: list[str] = []
|
||||
for ln in lines:
|
||||
stripped = ln.strip()
|
||||
if len(stripped) > 2:
|
||||
cleaned.append(stripped)
|
||||
n = len(cleaned)
|
||||
seen: dict[str, list[int]] = {}
|
||||
for i in range(n - 4):
|
||||
block = "\n".join(cleaned[i : i + 5])
|
||||
if block not in seen:
|
||||
seen[block] = []
|
||||
seen[block].append(i)
|
||||
for block, positions in seen.items():
|
||||
if len(positions) >= 2:
|
||||
key = block[:60] + "..." if len(block) > 60 else block
|
||||
if key not in dup_blocks:
|
||||
dup_blocks[key] = []
|
||||
dup_blocks[key].append((rel, len(positions)))
|
||||
|
||||
counts["duplicate_blocks"] = len(dup_blocks)
|
||||
counts["dup_files"] = list(dup_blocks.items())[:20] # type: ignore[index]
|
||||
|
||||
return counts
|
||||
|
||||
|
||||
def _calculate_score(
|
||||
cast_any: int, type_ignore: int, bare_except: int, dup_blocks: int
|
||||
) -> int:
|
||||
"""
|
||||
Calculate type safety score out of 100.
|
||||
|
||||
Deductions:
|
||||
- Each cast(Any, ...) costs 1 point
|
||||
- Each # type: ignore costs 0.5 points
|
||||
- Each bare except: costs 0.5 points
|
||||
- Each duplicate block costs 2 points
|
||||
- Floor at 0
|
||||
"""
|
||||
deduction = (
|
||||
cast_any
|
||||
+ type_ignore * 0.5
|
||||
+ bare_except * 0.5
|
||||
+ dup_blocks * 2
|
||||
)
|
||||
score = max(0, int(100 - deduction))
|
||||
return score
|
||||
|
||||
|
||||
def _get_grade(score: int) -> str:
|
||||
if score >= 90:
|
||||
return "A"
|
||||
elif score >= 80:
|
||||
return "B"
|
||||
elif score >= 70:
|
||||
return "C"
|
||||
elif score >= 60:
|
||||
return "D"
|
||||
else:
|
||||
return "F"
|
||||
|
||||
|
||||
class TestCodeQualityTyping:
|
||||
"""Test suite for type safety scoring."""
|
||||
|
||||
def test_type_safety_score(self):
|
||||
"""
|
||||
Type safety score based on cast(Any, ...) and # type: ignore usage.
|
||||
|
||||
Score = 100 - (cast_any_count * 1) - (type_ignore_count * 0.5)
|
||||
Minimum score: 0
|
||||
"""
|
||||
counts = _scan_for_patterns()
|
||||
cast_any = counts["cast_any"]
|
||||
type_ignore = counts["type_ignore"]
|
||||
bare_except = counts["bare_except"]
|
||||
dup_blocks = counts["duplicate_blocks"]
|
||||
score = _calculate_score(cast_any, type_ignore, bare_except, dup_blocks) # type: ignore[arg-type]
|
||||
|
||||
print(f"\n{'='*60}")
|
||||
print(f" Type Safety Score Report")
|
||||
print(f"{'='*60}")
|
||||
print(f" cast(Any, ...) count: {cast_any:>4} (cost: {cast_any} pts)")
|
||||
print(f" # type: ignore count: {type_ignore:>4} (cost: {type_ignore * 0.5:.1f} pts)") # type: ignore[index]
|
||||
print(f" bare except: count: {bare_except:>4} (cost: {bare_except * 0.5:.1f} pts)") # type: ignore[index]
|
||||
print(f" duplicate blocks: {dup_blocks:>4} (cost: {dup_blocks * 2} pts)") # type: ignore[index]
|
||||
print(f" {'-'*60}")
|
||||
print(f" Score: {score}/100 (Grade: {_get_grade(score)})")
|
||||
print(f"{'='*60}")
|
||||
|
||||
if counts["cast_any_files"]:
|
||||
print(f"\n Files with cast(Any, ...):")
|
||||
for f in sorted(counts["cast_any_files"])[:10]: # type: ignore[arg-type]
|
||||
print(f" - {f}")
|
||||
if len(counts["cast_any_files"]) > 10: # type: ignore[arg-type]
|
||||
print(f" ... and {len(counts["cast_any_files"]) - 10} more") # type: ignore[arg-type, index]
|
||||
print()
|
||||
|
||||
if counts["type_ignore_files"]:
|
||||
print(f" Files with # type: ignore:")
|
||||
for f in sorted(counts["type_ignore_files"])[:10]: # type: ignore[arg-type]
|
||||
print(f" - {f}")
|
||||
if len(counts["type_ignore_files"]) > 10: # type: ignore[arg-type]
|
||||
print(f" ... and {len(counts["type_ignore_files"]) - 10} more") # type: ignore[arg-type, index]
|
||||
print()
|
||||
|
||||
if counts["bare_except_files"]:
|
||||
print(f" Files with bare except:")
|
||||
for f in sorted(counts["bare_except_files"])[:10]: # type: ignore[arg-type]
|
||||
print(f" - {f}")
|
||||
if len(counts["bare_except_files"]) > 10: # type: ignore[arg-type]
|
||||
print(f" ... and {len(counts["bare_except_files"]) - 10} more") # type: ignore[arg-type, index]
|
||||
print()
|
||||
|
||||
if counts["dup_files"]:
|
||||
print(f" Duplicate code blocks (top 10):")
|
||||
for block_preview, locations in counts["dup_files"][:10]: # type: ignore[arg-type]
|
||||
files_str = ", ".join([f"{f} ({c}x)" for f, c in locations]) # type: ignore[iterable]
|
||||
print(f" - [{block_preview}] in {files_str}")
|
||||
print()
|
||||
|
||||
print(f" WARNING: This is a custom heuristic. Score may not reflect")
|
||||
print(f" actual type safety. Review individual cases manually.")
|
||||
print(f"{'='*60}\n")
|
||||
|
||||
# Emit warning level based on score
|
||||
if score < 60:
|
||||
pytest.fail(f"Type safety score too low: {score}/100 (Grade: {_get_grade(score)})")
|
||||
elif score < 80:
|
||||
pytest.skip(f"Type safety score below target: {score}/100 (Grade: {_get_grade(score)})")
|
||||
@@ -58,7 +58,7 @@ async def test_browser_tool_allows_non_admin_when_admin_requirement_disabled(
|
||||
cmd="open https://example.com",
|
||||
)
|
||||
|
||||
assert json.loads(result)["ok"] is True
|
||||
assert json.loads(result)["ok"] is True # type: ignore[arg-type]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -81,7 +81,7 @@ async def test_neo_skill_tool_allows_non_admin_when_admin_requirement_disabled(
|
||||
limit=5,
|
||||
)
|
||||
|
||||
payload = json.loads(result)
|
||||
payload = json.loads(result) # type: ignore[arg-type]
|
||||
assert payload["items"] == []
|
||||
assert payload["limit"] == 5
|
||||
|
||||
|
||||
+11
-11
@@ -1424,8 +1424,8 @@ async def test_plugin_web_route_returns_503_while_runtime_loading(
|
||||
return {"status": "ok", "message": None, "data": {"called": True}}
|
||||
|
||||
registered_web_apis = star_context.registered_web_apis
|
||||
original_registered_web_apis = list(registered_web_apis)
|
||||
registered_web_apis[:] = [
|
||||
original_registered_web_apis = list(registered_web_apis) # type: ignore[arg-type]
|
||||
registered_web_apis[:] = [ # type: ignore[index]
|
||||
("/runtime-guard-test", dummy_plugin_route, ["GET"], "runtime guard test"),
|
||||
]
|
||||
|
||||
@@ -1440,7 +1440,7 @@ async def test_plugin_web_route_returns_503_while_runtime_loading(
|
||||
_assert_runtime_loading_response(data)
|
||||
assert route_called is False
|
||||
finally:
|
||||
registered_web_apis[:] = original_registered_web_apis
|
||||
registered_web_apis[:] = original_registered_web_apis # type: ignore[index]
|
||||
_restore_runtime_ready(core_lifecycle_td)
|
||||
|
||||
|
||||
@@ -1461,8 +1461,8 @@ async def test_plugin_web_route_returns_failed_response_after_runtime_bootstrap_
|
||||
return {"status": "ok", "message": None, "data": {"called": True}}
|
||||
|
||||
registered_web_apis = star_context.registered_web_apis
|
||||
original_registered_web_apis = list(registered_web_apis)
|
||||
registered_web_apis[:] = [
|
||||
original_registered_web_apis = list(registered_web_apis) # type: ignore[arg-type]
|
||||
registered_web_apis[:] = [ # type: ignore[index]
|
||||
(
|
||||
"/runtime-failed-guard-test",
|
||||
dummy_plugin_route,
|
||||
@@ -1485,7 +1485,7 @@ async def test_plugin_web_route_returns_failed_response_after_runtime_bootstrap_
|
||||
)
|
||||
assert route_called is False
|
||||
finally:
|
||||
registered_web_apis[:] = original_registered_web_apis
|
||||
registered_web_apis[:] = original_registered_web_apis # type: ignore[index]
|
||||
_restore_runtime_ready(core_lifecycle_td)
|
||||
|
||||
|
||||
@@ -1818,10 +1818,10 @@ async def test_t2i_set_active_template_syncs_all_configs(
|
||||
data = await response.get_json()
|
||||
assert data["status"] == "ok"
|
||||
|
||||
conf_ids = set(core_lifecycle_td.astrbot_config_mgr.confs.keys())
|
||||
conf_ids = set(core_lifecycle_td.astrbot_config_mgr.confs.keys()) # type: ignore[union-attr]
|
||||
assert "default" in conf_ids
|
||||
for conf_id in conf_ids:
|
||||
conf = core_lifecycle_td.astrbot_config_mgr.confs[conf_id]
|
||||
conf = core_lifecycle_td.astrbot_config_mgr.confs[conf_id] # type: ignore[index]
|
||||
assert conf.get("t2i_active_template") == template_name
|
||||
assert conf_id in core_lifecycle_td.pipeline_scheduler_mapping
|
||||
finally:
|
||||
@@ -1891,10 +1891,10 @@ async def test_t2i_reset_default_template_syncs_all_configs(
|
||||
data = await response.get_json()
|
||||
assert data["status"] == "ok"
|
||||
|
||||
conf_ids = set(core_lifecycle_td.astrbot_config_mgr.confs.keys())
|
||||
conf_ids = set(core_lifecycle_td.astrbot_config_mgr.confs.keys()) # type: ignore[union-attr]
|
||||
assert "default" in conf_ids
|
||||
for conf_id in conf_ids:
|
||||
conf = core_lifecycle_td.astrbot_config_mgr.confs[conf_id]
|
||||
conf = core_lifecycle_td.astrbot_config_mgr.confs[conf_id] # type: ignore[index]
|
||||
assert conf.get("t2i_active_template") == "base"
|
||||
assert conf_id in core_lifecycle_td.pipeline_scheduler_mapping
|
||||
finally:
|
||||
@@ -1954,7 +1954,7 @@ async def test_t2i_update_active_template_reloads_all_schedulers(
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
conf_ids = list(core_lifecycle_td.astrbot_config_mgr.confs.keys())
|
||||
conf_ids = list(core_lifecycle_td.astrbot_config_mgr.confs.keys()) # type: ignore[union-attr]
|
||||
old_schedulers = {
|
||||
conf_id: core_lifecycle_td.pipeline_scheduler_mapping[conf_id]
|
||||
for conf_id in conf_ids
|
||||
|
||||
@@ -24,8 +24,10 @@ async def test_collect_and_register_commands_ignores_daily_create_limit() -> Non
|
||||
)
|
||||
adapter.client = MockDiscordBuilder.create_client()
|
||||
|
||||
exc = HTTPException("daily limit")
|
||||
exc.code = 30034 [attr-defined]
|
||||
class MockResponse:
|
||||
status: int = 400
|
||||
|
||||
exc = HTTPException(MockResponse(), {"code": 30034, "message": "daily limit"}) # type: ignore[arg-type]
|
||||
adapter.client.sync_commands.side_effect = exc
|
||||
|
||||
await adapter._collect_and_register_commands()
|
||||
|
||||
@@ -136,11 +136,11 @@ async def test_import_documents(
|
||||
assert result["failed_count"] == 0
|
||||
|
||||
# Verify kb_helper.upload_document was called correctly
|
||||
kb_helper = await core_lifecycle_td.kb_manager.get_kb("test_kb_id")
|
||||
assert kb_helper.upload_document.call_count == 2
|
||||
kb_helper = await core_lifecycle_td.kb_manager.get_kb("test_kb_id") # type: ignore[union-attr]
|
||||
assert kb_helper.upload_document.call_count == 2 # type: ignore[attr-defined]
|
||||
|
||||
# Check first call arguments
|
||||
call_args_list = kb_helper.upload_document.call_args_list
|
||||
call_args_list = kb_helper.upload_document.call_args_list # type: ignore[union-attr]
|
||||
|
||||
# First document
|
||||
_args1, kwargs1 = call_args_list[0]
|
||||
|
||||
@@ -47,13 +47,13 @@ class _FakeClient:
|
||||
|
||||
|
||||
def test_sync_release_writes_skill_and_map(monkeypatch, tmp_path: Path):
|
||||
calls = {"active": [], "sandbox_sync": 0}
|
||||
calls: dict[str, list | int] = {"active": [], "sandbox_sync": 0}
|
||||
|
||||
def _fake_set_skill_active(self, name, active):
|
||||
calls["active"].append((name, active))
|
||||
calls["active"].append((name, active)) # type: ignore[union-attr]
|
||||
|
||||
async def _fake_sync_sandboxes():
|
||||
calls["sandbox_sync"] += 1
|
||||
calls["sandbox_sync"] += 1 # type: ignore[operator]
|
||||
|
||||
monkeypatch.setattr(
|
||||
"astrbot.core.skills.neo_skill_sync.SkillManager.set_skill_active",
|
||||
|
||||
@@ -74,7 +74,7 @@ def test_promote_stable_sync_failure_auto_rolls_back(monkeypatch):
|
||||
tool = PromoteSkillCandidateTool()
|
||||
result = asyncio.run(
|
||||
tool.call(
|
||||
run_ctx,
|
||||
run_ctx, # type: ignore[arg-type]
|
||||
candidate_id="cand-1",
|
||||
stage="stable",
|
||||
sync_to_local=True,
|
||||
|
||||
@@ -70,7 +70,7 @@ def _make_req():
|
||||
def _import_apply_sandbox_tools():
|
||||
"""Import _apply_sandbox_tools, skipping if circular-import fails."""
|
||||
try:
|
||||
from astrbot.core.astr_main_agent import _apply_sandbox_tools
|
||||
from astrbot.core.astr_main_agent import _apply_sandbox_tools # type: ignore[import]
|
||||
|
||||
return _apply_sandbox_tools
|
||||
except ImportError:
|
||||
|
||||
@@ -84,7 +84,7 @@ async def test_extract_quoted_message_images_no_reply_component():
|
||||
get_group_id=lambda: "",
|
||||
)
|
||||
|
||||
images = await extract_quoted_message_images(event)
|
||||
images = await extract_quoted_message_images(event) # type: ignore[arg-type]
|
||||
assert images == []
|
||||
|
||||
|
||||
@@ -103,7 +103,7 @@ async def test_extract_quoted_message_text_reply_without_id_does_not_call_get_ms
|
||||
get_group_id=lambda: "",
|
||||
)
|
||||
|
||||
text = await extract_quoted_message_text(event)
|
||||
text = await extract_quoted_message_text(event) # type: ignore[arg-type]
|
||||
assert text == "quoted content"
|
||||
|
||||
|
||||
@@ -188,7 +188,7 @@ async def test_extract_quoted_message_text_forward_placeholder_variants_trigger_
|
||||
)
|
||||
|
||||
text = await extract_quoted_message_text(event)
|
||||
assert "Bob: [Image]world" in text
|
||||
assert "Bob: [Image]world" in text # type: ignore[operator]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -204,7 +204,7 @@ async def test_extract_quoted_message_text_mixed_placeholder_does_not_trigger_fa
|
||||
get_group_id=lambda: "",
|
||||
)
|
||||
|
||||
text = await extract_quoted_message_text(event)
|
||||
text = await extract_quoted_message_text(event) # type: ignore[arg-type]
|
||||
assert text is not None
|
||||
assert "[Forward Message]" in text
|
||||
assert "real text" in text
|
||||
@@ -244,7 +244,7 @@ async def test_extract_quoted_message_text_multimsg_malformed_config_does_not_ra
|
||||
},
|
||||
)
|
||||
|
||||
text = await extract_quoted_message_text(event)
|
||||
text = await extract_quoted_message_text(event) # type: ignore[arg-type]
|
||||
assert text == "still works"
|
||||
|
||||
|
||||
@@ -350,7 +350,7 @@ async def test_extract_quoted_message_images_non_image_local_path_is_ignored(tmp
|
||||
get_group_id=lambda: "",
|
||||
)
|
||||
|
||||
images = await extract_quoted_message_images(event)
|
||||
images = await extract_quoted_message_images(event) # type: ignore[arg-type]
|
||||
assert images == []
|
||||
|
||||
|
||||
@@ -486,7 +486,7 @@ async def test_extract_quoted_message_nested_forward_id_is_resolved():
|
||||
},
|
||||
)
|
||||
|
||||
text = await extract_quoted_message_text(event)
|
||||
text = await extract_quoted_message_text(event) # type: ignore[arg-type]
|
||||
assert text is not None
|
||||
assert "Bob: deep" in text
|
||||
|
||||
|
||||
@@ -50,7 +50,7 @@ def _create_mock_astrbot_paths(tmp_path: Path) -> MockAstrbotPaths:
|
||||
def test_list_skills_merges_local_and_sandbox_cache(tmp_path: Path):
|
||||
mock_paths = _create_mock_astrbot_paths(tmp_path)
|
||||
|
||||
mgr = SkillManager(skills_root=str(mock_paths.skills), astrbot_paths=mock_paths)
|
||||
mgr = SkillManager(skills_root=str(mock_paths.skills), astrbot_paths=mock_paths) # type: ignore[arg-type]
|
||||
_write_skill(mock_paths.skills, "custom-local", "local description")
|
||||
|
||||
mgr.set_sandbox_skills_cache(
|
||||
@@ -80,7 +80,7 @@ def test_list_skills_merges_local_and_sandbox_cache(tmp_path: Path):
|
||||
|
||||
def test_sandbox_cached_skill_respects_active_and_display_path(tmp_path: Path):
|
||||
mock_paths = _create_mock_astrbot_paths(tmp_path)
|
||||
mgr = SkillManager(skills_root=str(mock_paths.skills), astrbot_paths=mock_paths)
|
||||
mgr = SkillManager(skills_root=str(mock_paths.skills), astrbot_paths=mock_paths) # type: ignore[arg-type]
|
||||
mgr.set_sandbox_skills_cache(
|
||||
[
|
||||
{
|
||||
@@ -109,7 +109,7 @@ def test_sandbox_cached_skill_respects_active_and_display_path(tmp_path: Path):
|
||||
|
||||
def test_sandbox_and_local_path_resolution_with_show_sandbox_path_false(tmp_path: Path):
|
||||
mock_paths = _create_mock_astrbot_paths(tmp_path)
|
||||
mgr = SkillManager(skills_root=str(mock_paths.skills), astrbot_paths=mock_paths)
|
||||
mgr = SkillManager(skills_root=str(mock_paths.skills), astrbot_paths=mock_paths) # type: ignore[arg-type]
|
||||
_write_skill(mock_paths.skills, "custom-local", "local description")
|
||||
mgr.set_sandbox_skills_cache(
|
||||
[
|
||||
|
||||
@@ -40,7 +40,7 @@ class MockProvider(Provider):
|
||||
async def get_models(self) -> list[str]:
|
||||
return ["test_model"]
|
||||
|
||||
async def text_chat(self, **kwargs) -> LLMResponse:
|
||||
async def text_chat(self, **kwargs) -> LLMResponse: # type: ignore[method-assign]
|
||||
self.call_count += 1
|
||||
|
||||
# 检查工具是否被禁用
|
||||
@@ -84,50 +84,44 @@ class MockToolExecutor:
|
||||
"""模拟工具执行器"""
|
||||
|
||||
@classmethod
|
||||
def execute(cls, tool, run_context, **tool_args):
|
||||
async def generator():
|
||||
# 模拟工具返回结果,使用正确的类型
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
async def execute(cls, tool, run_context, **tool_args):
|
||||
# 模拟工具返回结果,使用正确的类型
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
|
||||
result = CallToolResult(
|
||||
content=[TextContent(type="text", text="工具执行结果")]
|
||||
)
|
||||
yield result
|
||||
|
||||
return generator()
|
||||
result = CallToolResult(
|
||||
content=[TextContent(type="text", text="工具执行结果")]
|
||||
)
|
||||
yield result
|
||||
|
||||
|
||||
class MockMixedContentToolExecutor:
|
||||
"""模拟返回图片 + 文本的工具执行器"""
|
||||
|
||||
@classmethod
|
||||
def execute(cls, tool, run_context, **tool_args):
|
||||
async def generator():
|
||||
from mcp.types import CallToolResult, ImageContent, TextContent
|
||||
async def execute(cls, tool, run_context, **tool_args):
|
||||
from mcp.types import CallToolResult, ImageContent, TextContent
|
||||
|
||||
result = CallToolResult(
|
||||
content=[
|
||||
ImageContent(
|
||||
type="image",
|
||||
data="dGVzdA==",
|
||||
mimeType="image/png",
|
||||
),
|
||||
TextContent(type="text", text="直播间标题:新游首发:零~红蝶~"),
|
||||
]
|
||||
)
|
||||
yield result
|
||||
|
||||
return generator()
|
||||
result = CallToolResult(
|
||||
content=[
|
||||
ImageContent(
|
||||
type="image",
|
||||
data="dGVzdA==",
|
||||
mimeType="image/png",
|
||||
),
|
||||
TextContent(type="text", text="直播间标题:新游首发:零~红蝶~"),
|
||||
]
|
||||
)
|
||||
yield result
|
||||
|
||||
|
||||
class MockFailingProvider(MockProvider):
|
||||
async def text_chat(self, **kwargs) -> LLMResponse:
|
||||
async def text_chat(self, **kwargs) -> LLMResponse: # type: ignore[method-assign]
|
||||
self.call_count += 1
|
||||
raise RuntimeError("primary provider failed")
|
||||
|
||||
|
||||
class MockErrProvider(MockProvider):
|
||||
async def text_chat(self, **kwargs) -> LLMResponse:
|
||||
async def text_chat(self, **kwargs) -> LLMResponse: # type: ignore[method-assign]
|
||||
self.call_count += 1
|
||||
return LLMResponse(
|
||||
role="err",
|
||||
@@ -140,7 +134,7 @@ class MockEmptyOutputThenSuccessProvider(MockProvider):
|
||||
super().__init__()
|
||||
self.failures_before_success = failures_before_success
|
||||
|
||||
async def text_chat(self, **kwargs) -> LLMResponse:
|
||||
async def text_chat(self, **kwargs) -> LLMResponse: # type: ignore[method-assign]
|
||||
self.call_count += 1
|
||||
if self.call_count <= self.failures_before_success:
|
||||
raise EmptyModelOutputError("model returned no usable output")
|
||||
@@ -180,7 +174,7 @@ class MockToolCallProvider(MockProvider):
|
||||
self.tool_args = tool_args or {}
|
||||
self.abort_signal = None
|
||||
|
||||
async def text_chat(self, **kwargs) -> LLMResponse:
|
||||
async def text_chat(self, **kwargs) -> LLMResponse: # type: ignore[method-assign]
|
||||
self.call_count += 1
|
||||
self.abort_signal = kwargs.get("abort_signal")
|
||||
return LLMResponse(
|
||||
@@ -869,7 +863,7 @@ async def test_skills_like_requery_passes_extra_user_content_parts():
|
||||
captured_kwargs = {}
|
||||
|
||||
class SkillsLikeProvider(MockProvider):
|
||||
async def text_chat(self, **kwargs) -> LLMResponse:
|
||||
async def text_chat(self, **kwargs) -> LLMResponse: # type: ignore[method-assign]
|
||||
self.call_count += 1
|
||||
if self.call_count == 1:
|
||||
# 第一次调用:返回工具选择(light schema)
|
||||
|
||||
+1
-1
@@ -46,7 +46,7 @@ class TestTuiBoxDrawing:
|
||||
def test_box_vertical(self):
|
||||
"""Test vertical box character."""
|
||||
from astrbot.tui.screen import BOX_VERT
|
||||
assert BOX_VERT == "│"
|
||||
assert BOX_VERT == "│"
|
||||
|
||||
def test_box_horizontal(self):
|
||||
"""Test horizontal box character."""
|
||||
|
||||
@@ -225,7 +225,7 @@ class TestSSEMessageParserProcess:
|
||||
parser.process_message(
|
||||
ParsedMessage(type=MessageType.REASONING, data="Thinking...")
|
||||
)
|
||||
parser._tool_calls["tc1"] = object()
|
||||
parser._tool_calls["tc1"] = object() # type: ignore[index]
|
||||
|
||||
parser.reset()
|
||||
|
||||
|
||||
@@ -16,5 +16,5 @@ def test_build_loggable_payload_redacts_api_key() -> None:
|
||||
loggable_payload = provider._build_loggable_payload(payload)
|
||||
|
||||
assert payload["app"]["token"] == "secret-token"
|
||||
assert loggable_payload["app"]["token"] == "***"
|
||||
assert loggable_payload["request"]["text"] == "hello"
|
||||
assert loggable_payload["app"]["token"] == "***" # type: ignore[index]
|
||||
assert loggable_payload["request"]["text"] == "hello" # type: ignore[index]
|
||||
|
||||
@@ -63,8 +63,8 @@ async def test_get_text_converts_opus_files_to_wav_before_transcription(
|
||||
assert not converted_path.exists()
|
||||
|
||||
create_mock = provider.client.audio.transcriptions.create
|
||||
create_mock.assert_awaited_once()
|
||||
file_arg = create_mock.await_args.kwargs["file"]
|
||||
create_mock.assert_awaited_once() # type: ignore[attr-defined]
|
||||
file_arg = create_mock.await_args.kwargs["file"] # type: ignore[attr-defined]
|
||||
assert file_arg[0] == "audio.wav"
|
||||
assert file_arg[1].name.endswith(".wav")
|
||||
file_arg[1].close()
|
||||
|
||||
@@ -22,7 +22,7 @@ def test_poke_to_dict_matches_onebot_v11_segment_format():
|
||||
async def test_respond_stage_treats_poke_with_target_as_non_empty():
|
||||
stage = RespondStage()
|
||||
chain = [Comp.Poke(type="126", id=2003)]
|
||||
assert await stage._is_empty_message_chain(chain) is False
|
||||
assert await stage._is_empty_message_chain(chain) is False # type: ignore[arg-type]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -161,7 +161,7 @@ async def test_do_handoff_background_reports_prepared_image_urls(
|
||||
|
||||
run_context = _build_run_context()
|
||||
await FunctionToolExecutor._do_handoff_background(
|
||||
tool=_DummyTool(),
|
||||
tool=_DummyTool(), # type: ignore[arg-type]
|
||||
run_context=run_context,
|
||||
task_id="task-id",
|
||||
input="hello",
|
||||
|
||||
@@ -3,6 +3,7 @@ from types import SimpleNamespace
|
||||
import pytest
|
||||
from sqlmodel import select
|
||||
|
||||
from astrbot.core import db_helper
|
||||
from astrbot.core.agent.response import AgentStats
|
||||
from astrbot.core.db.po import ProviderStat
|
||||
from astrbot.core.pipeline.process_stage.method.agent_sub_stages import internal
|
||||
@@ -14,7 +15,7 @@ async def test_record_internal_agent_stats_persists_provider_stat(
|
||||
temp_db,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
monkeypatch.setattr(internal, "db_helper", temp_db)
|
||||
monkeypatch.setattr(db_helper, "insert_provider_stat", temp_db.insert_provider_stat)
|
||||
|
||||
event = SimpleNamespace(unified_msg_origin="webchat:FriendMessage:session-42")
|
||||
req = ProviderRequest(
|
||||
|
||||
Reference in New Issue
Block a user