mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
fix: clean lint suppressions and async route errors
This commit is contained in:
@@ -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}")
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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}"
|
||||
)
|
||||
|
||||
@@ -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,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"
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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__)
|
||||
|
||||
@@ -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__
|
||||
# 检查重排序模型可用性
|
||||
|
||||
@@ -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}")
|
||||
|
||||
|
||||
@@ -97,7 +97,7 @@ class LogRoute(Route):
|
||||
},
|
||||
),
|
||||
)
|
||||
response.timeout = None # type: ignore
|
||||
setattr(response, "timeout", None)
|
||||
return response
|
||||
|
||||
async def log_history(self):
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user