fix: clean lint suppressions and async route errors

This commit is contained in:
LIghtJUNction
2026-03-18 18:06:42 +08:00
parent 26d6d1b36f
commit 14ec513b0d
32 changed files with 198 additions and 157 deletions
@@ -51,7 +51,7 @@ class PersonaCommands:
return lines
async def persona(self, message: AstrMessageEvent) -> None:
l = message.message_str.split(" ") # noqa: E741
parts = message.message_str.split(" ")
umo = message.unified_msg_origin
curr_persona_name = "无"
@@ -103,7 +103,7 @@ class PersonaCommands:
curr_cid_title = conv.title if conv.title else "新对话"
curr_cid_title += f"({cid[:4]})"
if len(l) == 1:
if len(parts) == 1:
message.set_result(
MessageEventResult()
.message(
@@ -122,7 +122,7 @@ class PersonaCommands:
)
.use_t2i(False),
)
elif l[1] == "list":
elif parts[1] == "list":
# 获取文件夹树和所有人格
folder_tree = await self.context.persona_manager.get_folder_tree()
all_personas = self.context.persona_manager.personas
@@ -149,11 +149,11 @@ class PersonaCommands:
msg = "\n".join(lines)
message.set_result(MessageEventResult().message(msg).use_t2i(False))
elif l[1] == "view":
if len(l) == 2:
elif parts[1] == "view":
if len(parts) == 2:
message.set_result(MessageEventResult().message("请输入人格情景名"))
return
ps = l[2].strip()
ps = parts[2].strip()
if persona := next(
builtins.filter(
lambda persona: persona["name"] == ps,
@@ -166,7 +166,7 @@ class PersonaCommands:
else:
msg = f"人格{ps}不存在"
message.set_result(MessageEventResult().message(msg))
elif l[1] == "unset":
elif parts[1] == "unset":
if not cid:
message.set_result(
MessageEventResult().message("当前没有对话,无法取消人格。"),
@@ -178,7 +178,7 @@ class PersonaCommands:
)
message.set_result(MessageEventResult().message("取消人格成功。"))
else:
ps = "".join(l[1:]).strip()
ps = "".join(parts[1:]).strip()
if not cid:
message.set_result(
MessageEventResult().message(
@@ -4,7 +4,6 @@ from astrbot.core import DEMO_MODE, logger
from astrbot.core.star.filter.command import CommandFilter
from astrbot.core.star.filter.command_group import CommandGroupFilter
from astrbot.core.star.star_handler import StarHandlerMetadata, star_handlers_registry
from astrbot.core.star.star_manager import PluginManager
class PluginCommands:
@@ -40,7 +39,10 @@ class PluginCommands:
MessageEventResult().message("/plugin off <插件名> 禁用插件。"),
)
return
await self.context._star_manager.turn_off_plugin(plugin_name) # type: ignore
if self.context._star_manager is None:
event.set_result(MessageEventResult().message("插件管理器未初始化。"))
return
await self.context._star_manager.turn_off_plugin(plugin_name)
event.set_result(MessageEventResult().message(f"插件 {plugin_name} 已禁用。"))
async def plugin_on(self, event: AstrMessageEvent, plugin_name: str = "") -> None:
@@ -53,7 +55,10 @@ class PluginCommands:
MessageEventResult().message("/plugin on <插件名> 启用插件。"),
)
return
await self.context._star_manager.turn_on_plugin(plugin_name) # type: ignore
if self.context._star_manager is None:
event.set_result(MessageEventResult().message("插件管理器未初始化。"))
return
await self.context._star_manager.turn_on_plugin(plugin_name)
event.set_result(MessageEventResult().message(f"插件 {plugin_name} 已启用。"))
async def plugin_get(self, event: AstrMessageEvent, plugin_repo: str = "") -> None:
@@ -68,9 +73,9 @@ class PluginCommands:
return
logger.info(f"准备从 {plugin_repo} 安装插件。")
if self.context._star_manager:
star_mgr: PluginManager = self.context._star_manager
star_mgr = self.context._star_manager
try:
await star_mgr.install_plugin(plugin_repo) # type: ignore
await star_mgr.install_plugin(plugin_repo)
event.set_result(MessageEventResult().message("安装插件成功。"))
except Exception as e:
logger.error(f"安装插件失败: {e}")
+15 -1
View File
@@ -22,7 +22,7 @@ from astrbot.core.utils.requirements_utils import (
from astrbot.core.utils.shared_preferences import SharedPreferences
from astrbot.core.utils.t2i.renderer import HtmlRenderer
from .log import LogBroker, LogManager # noqa
from .log import LogBroker, LogManager
from .utils.astrbot_path import get_astrbot_data_path
# 初始化数据存储文件夹
@@ -47,3 +47,17 @@ pip_installer = PipInstaller(
astrbot_config.get("pip_install_arg", ""),
astrbot_config.get("pypi_index_url", None),
)
__all__ = [
"AstrBotConfig",
"DEMO_MODE",
"astrbot_config",
"t2i_base_url",
"html_renderer",
"logger",
"LogBroker",
"LogManager",
"db_helper",
"sp",
"file_token_service",
"pip_installer",
]
@@ -263,7 +263,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
if has_stream_output:
return
except Exception as exc: # noqa: BLE001
except Exception as exc:
last_exception = exc
logger.warning(
"Chat Model %s request error: %s",
+2 -2
View File
@@ -152,7 +152,7 @@ class FunctionToolExecutor(BaseFunctionToolExecutor[AstrAgentContext]):
task_id=task_id,
**tool_args,
)
except Exception as e: # noqa: BLE001
except Exception as e:
logger.error(
f"Background task {task_id} failed: {e!s}",
exc_info=True,
@@ -399,7 +399,7 @@ class FunctionToolExecutor(BaseFunctionToolExecutor[AstrAgentContext]):
task_id=task_id,
**tool_args,
)
except Exception as e: # noqa: BLE001
except Exception as e:
logger.error(
f"Background handoff {task_id} ({tool.name}) failed: {e!s}",
exc_info=True,
+5 -5
View File
@@ -192,7 +192,7 @@ async def _apply_kb(
req.system_prompt += (
f"\n\n[Related Knowledge Base Results]:\n{kb_result}"
)
except Exception as exc: # noqa: BLE001
except Exception as exc:
logger.error("Error occurred while retrieving knowledge base: %s", exc)
else:
if req.func_tool is None:
@@ -451,7 +451,7 @@ async def _ensure_img_caption(
TextPart(text=f"<image_caption>{caption}</image_caption>")
)
req.image_urls = []
except Exception as exc: # noqa: BLE001
except Exception as exc:
logger.error("处理图片描述失败: %s", exc)
@@ -563,7 +563,7 @@ def _append_system_reminders(
try:
now = datetime.datetime.now(zoneinfo.ZoneInfo(timezone))
current_time = now.strftime("%Y-%m-%d %H:%M (%Z)")
except Exception as exc: # noqa: BLE001
except Exception as exc:
logger.error("时区设置错误: %s, 使用本地时区", exc)
if not current_time:
current_time = (
@@ -1011,7 +1011,7 @@ async def build_main_agent(
req.image_urls.append(image_ref)
fallback_quoted_image_count += 1
_append_quoted_image_attachment(req, image_ref)
except Exception as exc: # noqa: BLE001
except Exception as exc:
logger.warning(
"Failed to resolve fallback quoted images for umo=%s, reply_id=%s: %s",
event.unified_msg_origin,
@@ -1032,7 +1032,7 @@ async def build_main_agent(
if config.file_extract_enabled:
try:
await _apply_file_extract(event, req, config)
except Exception as exc: # noqa: BLE001
except Exception as exc:
logger.error("Error occurred while applying file extract: %s", exc)
if not req.prompt and not req.image_urls:
+2 -2
View File
@@ -101,7 +101,7 @@ class LocalShellComponent(ShellComponent):
# `command` is intentionally executed through the current shell so
# local computer-use behavior matches existing tool semantics.
# Safety relies on `_is_safe_command()` and the allowed-root checks.
proc = subprocess.Popen( # noqa: S602 # nosemgrep: python.lang.security.audit.dangerous-subprocess-use-audit
proc = subprocess.Popen( # nosemgrep: python.lang.security.audit.dangerous-subprocess-use-audit
command,
shell=shell,
cwd=working_dir,
@@ -113,7 +113,7 @@ class LocalShellComponent(ShellComponent):
# `command` is intentionally executed through the current shell so
# local computer-use behavior matches existing tool semantics.
# Safety relies on `_is_safe_command()` and the allowed-root checks.
result = subprocess.run( # noqa: S602 # nosemgrep: python.lang.security.audit.dangerous-subprocess-use-audit
result = subprocess.run( # nosemgrep: python.lang.security.audit.dangerous-subprocess-use-audit
command,
shell=shell,
cwd=working_dir,
+2 -1
View File
@@ -256,7 +256,8 @@ class AstrBotCoreLifecycle:
# 把插件中注册的所有协程函数注册到事件总线中并执行
extra_tasks = []
for task in self.star_context._register_tasks:
extra_tasks.append(asyncio.create_task(task, name=task.__name__)) # type: ignore
task_name = getattr(task, "__name__", task.__class__.__name__)
extra_tasks.append(asyncio.create_task(task, name=task_name))
tasks_ = [event_bus_task, *(extra_tasks if extra_tasks else [])]
if cron_task:
+2 -2
View File
@@ -205,7 +205,7 @@ class CronJobManager:
await self._run_active_agent_job(job, start_time=start_time)
else:
raise ValueError(f"Unknown cron job type: {job.job_type}")
except Exception as e: # noqa: BLE001
except Exception as e:
status = "failed"
last_error = str(e)
logger.error(f"Cron job {job_id} failed: {e!s}", exc_info=True)
@@ -286,7 +286,7 @@ class CronJobManager:
if isinstance(session_str, MessageSession)
else MessageSession.from_str(session_str)
)
except Exception as e: # noqa: BLE001
except Exception as e:
logger.error(f"Invalid session for cron job: {e}")
return
+25 -12
View File
@@ -8,7 +8,6 @@ from pathlib import Path
import aiofiles
from astrbot.core import logger
from astrbot.core.db.vec_db.base import BaseVecDB
from astrbot.core.db.vec_db.faiss_impl.vec_db import FaissVecDB
from astrbot.core.provider.manager import ProviderManager
from astrbot.core.provider.provider import (
@@ -106,7 +105,7 @@ Text chunk to process:
class KBHelper:
vec_db: BaseVecDB
vec_db: FaissVecDB | None
kb: KnowledgeBase
def __init__(
@@ -126,6 +125,7 @@ class KBHelper:
self.kb_dir = Path(self.kb_root_dir) / self.kb.kb_id
self.kb_medias_dir = Path(self.kb_dir) / "medias" / self.kb.kb_id
self.kb_files_dir = Path(self.kb_dir) / "files" / self.kb.kb_id
self.vec_db = None
self.kb_medias_dir.mkdir(parents=True, exist_ok=True)
self.kb_files_dir.mkdir(parents=True, exist_ok=True)
@@ -133,16 +133,25 @@ class KBHelper:
async def initialize(self) -> None:
await self._ensure_vec_db()
def _get_vec_db(self) -> FaissVecDB:
if self.vec_db is None:
raise ValueError("Vector database is not initialized")
return self.vec_db
async def get_ep(self) -> EmbeddingProvider:
if not self.kb.embedding_provider_id:
raise ValueError(f"知识库 {self.kb.kb_name} 未配置 Embedding Provider")
ep: EmbeddingProvider = await self.prov_mgr.get_provider_by_id(
self.kb.embedding_provider_id,
) # type: ignore
)
if not ep:
raise ValueError(
f"无法找到 ID 为 {self.kb.embedding_provider_id} 的 Embedding Provider",
)
if not isinstance(ep, EmbeddingProvider):
raise ValueError(
f"Provider {self.kb.embedding_provider_id} is not an Embedding Provider",
)
return ep
async def get_rp(self) -> RerankProvider | None:
@@ -150,11 +159,15 @@ class KBHelper:
return None
rp: RerankProvider = await self.prov_mgr.get_provider_by_id(
self.kb.rerank_provider_id,
) # type: ignore
)
if not rp:
raise ValueError(
f"无法找到 ID 为 {self.kb.rerank_provider_id} 的 Rerank Provider",
)
if not isinstance(rp, RerankProvider):
raise ValueError(
f"Provider {self.kb.rerank_provider_id} is not a Rerank Provider",
)
return rp
async def _ensure_vec_db(self) -> FaissVecDB:
@@ -297,7 +310,7 @@ class KBHelper:
if progress_callback:
await progress_callback("embedding", current, total)
await self.vec_db.insert_batch(
await self._get_vec_db().insert_batch(
contents=contents,
metadatas=metadatas,
batch_size=batch_size,
@@ -327,7 +340,7 @@ class KBHelper:
await session.refresh(doc)
vec_db: FaissVecDB = self.vec_db # type: ignore
vec_db = self._get_vec_db()
await self.kb_db.update_kb_stats(kb_id=self.kb.kb_id, vec_db=vec_db)
await self.refresh_kb()
await self.refresh_document(doc_id)
@@ -364,21 +377,21 @@ class KBHelper:
"""删除单个文档及其相关数据"""
await self.kb_db.delete_document_by_id(
doc_id=doc_id,
vec_db=self.vec_db, # type: ignore
vec_db=self._get_vec_db(),
)
await self.kb_db.update_kb_stats(
kb_id=self.kb.kb_id,
vec_db=self.vec_db, # type: ignore
vec_db=self._get_vec_db(),
)
await self.refresh_kb()
async def delete_chunk(self, chunk_id: str, doc_id: str) -> None:
"""删除单个文本块及其相关数据"""
vec_db: FaissVecDB = self.vec_db # type: ignore
vec_db = self._get_vec_db()
await vec_db.delete(chunk_id)
await self.kb_db.update_kb_stats(
kb_id=self.kb.kb_id,
vec_db=self.vec_db, # type: ignore
vec_db=self._get_vec_db(),
)
await self.refresh_kb()
await self.refresh_document(doc_id)
@@ -409,7 +422,7 @@ class KBHelper:
limit: int = 100,
) -> list[dict]:
"""获取文档的所有块及其元数据"""
vec_db: FaissVecDB = self.vec_db # type: ignore
vec_db = self._get_vec_db()
chunks = await vec_db.document_storage.get_documents(
metadata_filters={"kb_doc_id": doc_id},
offset=offset,
@@ -432,7 +445,7 @@ class KBHelper:
async def get_chunk_count_by_doc_id(self, doc_id: str) -> int:
"""获取文档的块数量"""
vec_db: FaissVecDB = self.vec_db # type: ignore
vec_db = self._get_vec_db()
count = await vec_db.count_documents(metadata_filter={"kb_doc_id": doc_id})
return count
+2 -1
View File
@@ -27,7 +27,8 @@ from astrbot.core.utils.metrics import Metric
from astrbot.core.utils.trace import TraceSpan
from .astrbot_message import AstrBotMessage, Group
from .message_session import MessageSesion, MessageSession # noqa
from .message_session import MessageSesion as MessageSesion
from .message_session import MessageSession
from .platform_metadata import PlatformMetadata
+22 -59
View File
@@ -2,6 +2,7 @@ import asyncio
import traceback
from asyncio import Queue
from dataclasses import dataclass
from importlib import import_module
from astrbot.core import logger
from astrbot.core.config.astrbot_config import AstrBotConfig
@@ -12,6 +13,24 @@ from .platform import Platform, PlatformStatus
from .register import platform_cls_map
from .sources.webchat.webchat_adapter import WebChatAdapter
PLATFORM_ADAPTER_MODULES: dict[str, str] = {
"aiocqhttp": ".sources.aiocqhttp.aiocqhttp_platform_adapter",
"qq_official": ".sources.qqofficial.qqofficial_platform_adapter",
"qq_official_webhook": ".sources.qqofficial_webhook.qo_webhook_adapter",
"lark": ".sources.lark.lark_adapter",
"dingtalk": ".sources.dingtalk.dingtalk_adapter",
"telegram": ".sources.telegram.tg_adapter",
"wecom": ".sources.wecom.wecom_adapter",
"wecom_ai_bot": ".sources.wecom_ai_bot.wecomai_adapter",
"weixin_official_account": ".sources.weixin_official_account.weixin_offacc_adapter",
"discord": ".sources.discord.discord_platform_adapter",
"misskey": ".sources.misskey.misskey_adapter",
"slack": ".sources.slack.slack_adapter",
"satori": ".sources.satori.satori_adapter",
"line": ".sources.line.line_adapter",
"kook": ".sources.kook.kook_adapter",
}
@dataclass
class PlatformTasks:
@@ -125,65 +144,9 @@ class PlatformManager:
logger.info(
f"载入 {platform_config['type']}({platform_config['id']}) 平台适配器 ...",
)
match platform_config["type"]:
case "aiocqhttp":
from .sources.aiocqhttp.aiocqhttp_platform_adapter import (
AiocqhttpAdapter, # noqa: F401
)
case "qq_official":
from .sources.qqofficial.qqofficial_platform_adapter import (
QQOfficialPlatformAdapter, # noqa: F401
)
case "qq_official_webhook":
from .sources.qqofficial_webhook.qo_webhook_adapter import (
QQOfficialWebhookPlatformAdapter, # noqa: F401
)
case "lark":
from .sources.lark.lark_adapter import (
LarkPlatformAdapter, # noqa: F401
)
case "dingtalk":
from .sources.dingtalk.dingtalk_adapter import (
DingtalkPlatformAdapter, # noqa: F401
)
case "telegram":
from .sources.telegram.tg_adapter import (
TelegramPlatformAdapter, # noqa: F401
)
case "wecom":
from .sources.wecom.wecom_adapter import (
WecomPlatformAdapter, # noqa: F401
)
case "wecom_ai_bot":
from .sources.wecom_ai_bot.wecomai_adapter import (
WecomAIBotAdapter, # noqa: F401
)
case "weixin_official_account":
from .sources.weixin_official_account.weixin_offacc_adapter import (
WeixinOfficialAccountPlatformAdapter, # noqa: F401
)
case "discord":
from .sources.discord.discord_platform_adapter import (
DiscordPlatformAdapter, # noqa: F401
)
case "misskey":
from .sources.misskey.misskey_adapter import (
MisskeyPlatformAdapter, # noqa: F401
)
case "slack":
from .sources.slack.slack_adapter import SlackAdapter # noqa: F401
case "satori":
from .sources.satori.satori_adapter import (
SatoriPlatformAdapter, # noqa: F401
)
case "line":
from .sources.line.line_adapter import (
LinePlatformAdapter, # noqa: F401
)
case "kook":
from .sources.kook.kook_adapter import (
KookPlatformAdapter, # noqa: F401
)
module_path = PLATFORM_ADAPTER_MODULES.get(platform_config["type"])
if module_path is not None:
import_module(module_path, package=__package__)
except (ImportError, ModuleNotFoundError) as e:
logger.error(
f"加载平台适配器 {platform_config['type']} 失败,原因:{e}。请检查依赖库是否安装。提示:可以在 管理面板->平台日志->安装Pip库 中安装依赖库。",
@@ -293,7 +293,7 @@ class OrderMessage(BaseModel):
class KookMessageSignal(IntEnum):
"""KOOK WebSocket 信令类型
ws文档: https://developer.kookapp.cn/doc/websocket""" # noqa: W291
ws文档: https://developer.kookapp.cn/doc/websocket"""
MESSAGE = 0
"""server->client 消息(s包含聊天和通知消息)"""
@@ -436,8 +436,8 @@ class KookWebsocketEvent(KookBaseDataClass):
] = Field(None, validation_alias="d", serialization_alias="d")
"""数据事件主体,对应原字段是'd'"""
sn: int | None = None
"""消息序号 , 用来确定消息顺序和ws重连时使用
详见ws连接流程文档: https://developer.kookapp.cn/doc/websocket#%E8%BF%9E%E6%8E%A5%E6%B5%81%E7%A8%8B""" # noqa: W291
"""消息序号 , 用来确定消息顺序和ws重连时使用
详见ws连接流程文档: https://developer.kookapp.cn/doc/websocket#%E8%BF%9E%E6%8E%A5%E6%B5%81%E7%A8%8B"""
@model_validator(mode="before")
@classmethod
@@ -4,7 +4,7 @@ import time
import uuid
from collections.abc import Callable, Coroutine
from pathlib import Path
from typing import Any
from typing import Any, cast
from astrbot import logger
from astrbot.core import db_helper
@@ -239,7 +239,7 @@ class WebChatAdapter(Platform):
session_id=message.session_id,
)
_, _, payload = message.raw_message # type: ignore
_, _, payload = cast(tuple[Any, Any, dict[str, Any]], message.raw_message)
message_event.set_extra("selected_provider", payload.get("selected_provider"))
message_event.set_extra("selected_model", payload.get("selected_model"))
message_event.set_extra(
+1 -1
View File
@@ -578,7 +578,7 @@ class FunctionToolManager:
"""安全清理单个 MCP 客户端,避免清理异常中断主流程。"""
try:
await mcp_client.cleanup()
except Exception as cleanup_exc: # noqa: BLE001 - only log here
except Exception as cleanup_exc: # only log here
logger.error(
f"Failed to cleanup MCP client resources {name}: {cleanup_exc}"
)
+11 -5
View File
@@ -2,7 +2,7 @@ import asyncio
import copy
import os
import traceback
from collections.abc import Callable
from collections.abc import Awaitable, Callable
from typing import Protocol, runtime_checkable
from astrbot.core import astrbot_config, logger, sp
@@ -28,6 +28,11 @@ class HasInitialize(Protocol):
async def initialize(self) -> None: ...
@runtime_checkable
class SupportsTerminate(Protocol):
def terminate(self) -> Awaitable[object]: ...
class ProviderManager:
def __init__(
self,
@@ -739,8 +744,9 @@ class ProviderManager:
if self.inst_map[provider_id] == self.curr_tts_provider_inst:
self.curr_tts_provider_inst = None
if getattr(self.inst_map[provider_id], "terminate", None):
await self.inst_map[provider_id].terminate() # type: ignore
inst = self.inst_map[provider_id]
if isinstance(inst, SupportsTerminate):
await inst.terminate()
logger.info(
f"{provider_id} 提供商适配器已终止({len(self.provider_insts)}, {len(self.stt_provider_insts)}, {len(self.tts_provider_insts)})",
@@ -820,8 +826,8 @@ class ProviderManager:
pass
for provider_inst in self.provider_insts:
if hasattr(provider_inst, "terminate"):
await provider_inst.terminate() # type: ignore
if isinstance(provider_inst, SupportsTerminate):
await provider_inst.terminate()
try:
await self.llm_tools.disable_mcp_server()
except Exception:
+2 -2
View File
@@ -2,7 +2,7 @@ import abc
import asyncio
import os
from collections.abc import AsyncGenerator
from typing import TypeAlias, Union
from typing import TypeAlias, Union, cast
import aiofiles
import anyio
@@ -157,7 +157,7 @@ class Provider(AbstractProvider):
"""
if False: # pragma: no cover - make this an async generator for typing
yield None # type: ignore
yield cast(LLMResponse, None)
raise NotImplementedError()
async def pop_record(self, context: list) -> None:
@@ -127,7 +127,7 @@ class ProviderAnthropic(Provider):
if "tool_calls" in message and isinstance(message["tool_calls"], list):
for tool_call in message["tool_calls"]:
blocks.append( # noqa: PERF401
blocks.append(
{
"type": "tool_use",
"name": tool_call["function"]["name"],
@@ -1,3 +1,6 @@
from collections.abc import MutableMapping
from typing import cast
from ..register import register_provider_adapter
from .openai_source import ProviderOpenAIOfficial
@@ -14,4 +17,7 @@ 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.
self.client._custom_headers["APP-Code"] = "KRLC5702" # type: ignore
custom_headers = cast(
MutableMapping[str, str], getattr(self.client, "_custom_headers")
)
custom_headers["APP-Code"] = "KRLC5702"
@@ -1,3 +1,6 @@
from collections.abc import MutableMapping
from typing import cast
from ..register import register_provider_adapter
from .openai_source import ProviderOpenAIOfficial
@@ -13,10 +16,9 @@ class ProviderOpenRouter(ProviderOpenAIOfficial):
) -> None:
super().__init__(provider_config, provider_settings)
# Reference to: https://openrouter.ai/docs/api/reference/overview#headers
self.client._custom_headers["HTTP-Referer"] = ( # type: ignore
"https://github.com/AstrBotDevs/AstrBot"
)
self.client._custom_headers["X-OpenRouter-Title"] = "AstrBot" # type: ignore
self.client._custom_headers["X-OpenRouter-Categories"] = (
"general-chat,personal-agent" # type: ignore
custom_headers = cast(
MutableMapping[str, str], getattr(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 -3
View File
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging
from asyncio import Queue
from collections.abc import Awaitable, Callable
from collections.abc import Awaitable, Callable, Coroutine
from typing import TYPE_CHECKING, Any, Protocol
from deprecated import deprecated
@@ -53,14 +53,20 @@ class PlatformManagerProtocol(Protocol):
platform_insts: list[Platform]
class StarManagerProtocol(Protocol):
async def turn_off_plugin(self, plugin_name: str) -> None: ...
async def turn_on_plugin(self, plugin_name: str) -> None: ...
async def install_plugin(self, repo_url: str, proxy: str = "") -> dict | None: ...
class Context:
"""暴露给插件的接口上下文。"""
registered_web_apis: list = []
# 向后兼容的变量
_register_tasks: list[Awaitable] = []
_star_manager = None
_register_tasks: list[Coroutine[object, object, object]] = []
_star_manager: StarManagerProtocol | None = None
def __init__(
self,
+1 -1
View File
@@ -197,7 +197,7 @@ class StarHandlerRegistry(Generic[T]):
return len(self._handlers)
star_handlers_registry = StarHandlerRegistry() # type: ignore
star_handlers_registry: StarHandlerRegistry = StarHandlerRegistry()
class EventType(enum.Enum):
+6 -3
View File
@@ -160,7 +160,7 @@ class PluginManager:
self.updator = PluginUpdator()
self.context = context
self.context._star_manager = self # type: ignore
self.context._star_manager = self
StarTools.initialize(context)
self.config = config
@@ -910,6 +910,9 @@ class PluginManager:
assert metadata.module_path is not None, (
f"插件 {metadata.name} 的模块路径为空。"
)
assert metadata.star_cls is not None, (
f"插件 {metadata.name} 的实例为空。"
)
# 绑定 handler
related_handlers = (
@@ -920,7 +923,7 @@ class PluginManager:
for handler in related_handlers:
handler.handler = functools.partial(
handler.handler,
metadata.star_cls, # type: ignore
metadata.star_cls,
)
# 绑定 llm_tool handler
for func_tool in llm_tools.func_list:
@@ -942,7 +945,7 @@ class PluginManager:
ft.handler_module_path = metadata.module_path
ft.handler = functools.partial(
ft.handler,
metadata.star_cls, # type: ignore
metadata.star_cls,
)
if ft.name in inactivated_llm_tools:
ft.active = False
+17 -4
View File
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any
from astrbot import logger
from astrbot.core.agent.agent import Agent
from astrbot.core.agent.handoff import HandoffTool
from astrbot.core.agent.tool import FunctionTool
from astrbot.core.provider.func_tool_manager import FunctionToolManager
if TYPE_CHECKING:
@@ -60,7 +61,7 @@ class SubAgentOrchestrator:
provider_id = item.get("provider_id")
if provider_id is not None:
provider_id = str(provider_id).strip() or None
tools = item.get("tools", [])
tools: list[str | FunctionTool] | None = item.get("tools", [])
begin_dialogs = None
if persona_data:
@@ -70,7 +71,11 @@ class SubAgentOrchestrator:
begin_dialogs = copy.deepcopy(
persona_data.get("_begin_dialogs_processed")
)
tools = persona_data.get("tools")
persona_tools = persona_data.get("tools")
if isinstance(persona_tools, list):
tools = [str(t).strip() for t in persona_tools if str(t).strip()]
else:
tools = None
if public_description == "" and prompt:
public_description = prompt[:120]
if tools is None:
@@ -78,12 +83,20 @@ class SubAgentOrchestrator:
elif not isinstance(tools, list):
tools = []
else:
tools = [str(t).strip() for t in tools if str(t).strip()]
tools = [
t if isinstance(t, FunctionTool) else str(t).strip()
for t in tools
if (
isinstance(t, FunctionTool)
or (isinstance(t, str) and t.strip())
or (not isinstance(t, FunctionTool) and str(t).strip())
)
]
agent = Agent[AstrAgentContext](
name=name,
instructions=instructions,
tools=tools, # type: ignore
tools=tools,
)
agent.begin_dialogs = begin_dialogs
# The tool description should be a short description for the main LLM,
+1 -1
View File
@@ -20,7 +20,7 @@ async def persist_agent_history(
history = []
try:
history = json.loads(req.conversation.history or "[]")
except Exception as exc: # noqa: BLE001
except Exception as exc:
logger.warning("Failed to parse conversation history: %s", exc)
history.append({"role": "user", "content": "Output your last task result below."})
history.append({"role": "assistant", "content": summary_note})
+2 -1
View File
@@ -72,6 +72,7 @@ class ApiKeyRoute(Route):
post_data = await request.json or {}
name = str(post_data.get("name", "")).strip() or "Untitled API Key"
normalized_scopes: list[str]
scopes = post_data.get("scopes")
if scopes is None:
normalized_scopes = list(ALL_OPEN_API_SCOPES)
@@ -111,7 +112,7 @@ class ApiKeyRoute(Route):
name=name,
key_hash=key_hash,
key_prefix=key_prefix,
scopes=normalized_scopes, # type: ignore
scopes=normalized_scopes,
created_by=created_by,
expires_at=expires_at,
)
+4 -4
View File
@@ -48,7 +48,7 @@ class CronRoute(Route):
jobs = await cron_mgr.list_jobs(job_type)
data = [self._serialize_job(j) for j in jobs]
return jsonify(Response().ok(data=data).__dict__)
except Exception as e: # noqa: BLE001
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"Failed to list jobs: {e!s}").__dict__)
@@ -119,7 +119,7 @@ class CronRoute(Route):
)
return jsonify(Response().ok(data=self._serialize_job(job)).__dict__)
except Exception as e: # noqa: BLE001
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"Failed to create job: {e!s}").__dict__)
@@ -156,7 +156,7 @@ class CronRoute(Route):
if not job:
return jsonify(Response().error("Job not found").__dict__)
return jsonify(Response().ok(data=self._serialize_job(job)).__dict__)
except Exception as e: # noqa: BLE001
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"Failed to update job: {e!s}").__dict__)
@@ -169,6 +169,6 @@ class CronRoute(Route):
)
await cron_mgr.delete_job(job_id)
return jsonify(Response().ok(message="deleted").__dict__)
except Exception as e: # noqa: BLE001
except Exception as e:
logger.error(traceback.format_exc())
return jsonify(Response().error(f"Failed to delete job: {e!s}").__dict__)
+8 -6
View File
@@ -366,7 +366,7 @@ class KnowledgeBaseRoute(Route):
return Response().error("缺少参数 embedding_provider_id").__dict__
prv = await kb_manager.provider_manager.get_provider_by_id(
embedding_provider_id,
) # type: ignore
)
if not prv or not isinstance(prv, EmbeddingProvider):
return (
Response().error(f"嵌入模型不存在或类型错误({type(prv)})").__dict__
@@ -381,11 +381,13 @@ class KnowledgeBaseRoute(Route):
return Response().error(f"测试嵌入模型失败: {e!s}").__dict__
# pre-check rerank
if rerank_provider_id:
rerank_prv: RerankProvider = (
await kb_manager.provider_manager.get_provider_by_id(
rerank_provider_id,
)
) # type: ignore
rerank_prv = await kb_manager.provider_manager.get_provider_by_id(
rerank_provider_id,
)
if rerank_prv is not None and not isinstance(
rerank_prv, RerankProvider
):
return Response().error("重排序模型类型错误").__dict__
if not rerank_prv:
return Response().error("重排序模型不存在").__dict__
# 检查重排序模型可用性
+2 -2
View File
@@ -95,7 +95,7 @@ class LiveChatSession:
logger.error(f"[Live Chat] 组装 WAV 文件失败: {e}", exc_info=True)
return None, 0.0
def cleanup(self) -> None:
async def cleanup(self) -> None:
"""清理临时文件"""
if self.temp_audio_path and await anyio.Path(self.temp_audio_path).exists():
try:
@@ -179,7 +179,7 @@ class LiveChatRoute(Route):
# 清理会话
if session_id in self.sessions:
await self._cleanup_chat_subscriptions(live_session)
live_session.cleanup()
await live_session.cleanup()
del self.sessions[session_id]
logger.info(f"[Live Chat] WebSocket 连接关闭: {username}")
+1 -1
View File
@@ -97,7 +97,7 @@ class LogRoute(Route):
},
),
)
response.timeout = None # type: ignore
setattr(response, "timeout", None)
return response
async def log_history(self):
+14 -10
View File
@@ -167,7 +167,7 @@ class PluginRoute(Route):
if not force_refresh:
# 先检查MD5是否匹配,如果匹配则使用缓存
if await self._is_cache_valid(source):
cached_data = self._load_plugin_cache(source.cache_file)
cached_data = await self._load_plugin_cache(source.cache_file)
if cached_data:
logger.debug("缓存MD5匹配,使用缓存的插件市场数据")
return Response().ok(cached_data).__dict__
@@ -205,7 +205,7 @@ class PluginRoute(Route):
)
# 获取最新的MD5并保存到缓存
current_md5 = await self._fetch_remote_md5(source.md5_url)
self._save_plugin_cache(
await self._save_plugin_cache(
source.cache_file,
remote_data,
current_md5,
@@ -217,7 +217,7 @@ class PluginRoute(Route):
# 如果远程获取失败,尝试使用缓存数据
if not cached_data:
cached_data = self._load_plugin_cache(source.cache_file)
cached_data = await self._load_plugin_cache(source.cache_file)
if cached_data:
logger.warning("远程插件市场数据获取失败,使用缓存数据")
@@ -249,14 +249,14 @@ class PluginRoute(Route):
]
return RegistrySource(urls=urls, cache_file=cache_file, md5_url=md5_url)
def _load_cached_md5(self, cache_file: str) -> str | None:
async def _load_cached_md5(self, cache_file: str) -> str | None:
"""从缓存文件中加载MD5"""
if not await anyio.Path(cache_file).exists():
return None
try:
async with await anyio.open_file(cache_file, encoding="utf-8") as f:
cache_data = json.load(f)
cache_data = json.loads(await f.read())
return cache_data.get("md5")
except Exception as e:
logger.warning(f"加载缓存MD5失败: {e}")
@@ -288,7 +288,7 @@ class PluginRoute(Route):
async def _is_cache_valid(self, source: RegistrySource) -> bool:
"""检查缓存是否有效(基于MD5)"""
try:
cached_md5 = self._load_cached_md5(source.cache_file)
cached_md5 = await self._load_cached_md5(source.cache_file)
if not cached_md5:
logger.debug("缓存文件中没有MD5信息")
return False
@@ -308,12 +308,12 @@ class PluginRoute(Route):
logger.warning(f"检查缓存有效性失败: {e}")
return False
def _load_plugin_cache(self, cache_file: str):
async def _load_plugin_cache(self, cache_file: str):
"""加载本地缓存的插件市场数据"""
try:
if await anyio.Path(cache_file).exists():
async with await anyio.open_file(cache_file, encoding="utf-8") as f:
cache_data = json.load(f)
cache_data = json.loads(await f.read())
# 检查缓存是否有效
if "data" in cache_data and "timestamp" in cache_data:
logger.debug(
@@ -324,7 +324,9 @@ class PluginRoute(Route):
logger.warning(f"加载插件市场缓存失败: {e}")
return None
def _save_plugin_cache(self, cache_file: str, data, md5: str | None = None) -> None:
async def _save_plugin_cache(
self, cache_file: str, data, md5: str | None = None
) -> None:
"""保存插件市场数据到本地缓存"""
try:
# 确保目录存在
@@ -337,7 +339,9 @@ class PluginRoute(Route):
}
async with await anyio.open_file(cache_file, "w", encoding="utf-8") as f:
json.dump(cache_data, f, ensure_ascii=False, indent=2)
await f.write(
json.dumps(cache_data, ensure_ascii=False, indent=2),
)
logger.debug(f"插件市场数据已缓存到: {cache_file}, MD5: {md5}")
except Exception as e:
logger.warning(f"保存插件市场缓存失败: {e}")
+2 -1
View File
@@ -1,6 +1,7 @@
import base64
import traceback
from io import BytesIO
from typing import cast
from astrbot.api import logger
from astrbot.core.db.vec_db.faiss_impl import FaissVecDB
@@ -81,7 +82,7 @@ async def generate_tsne_visualization(
index.reconstruct(i, vectors[i])
# 获取查询向量
vec_db: FaissVecDB = kb_helper.vec_db # type: ignore
vec_db = cast(FaissVecDB, kb_helper.vec_db)
embedding_provider = vec_db.embedding_provider
query_embedding = await embedding_provider.get_embedding(query)
query_vector = np.array([query_embedding], dtype=np.float32)