diff --git a/README.md b/README.md index 6774d4305..214f5ec69 100644 --- a/README.md +++ b/README.md @@ -86,7 +86,7 @@ astrbot run > AstrBot requires Python 3.12 or later. The `--python 3.12` option ensures that `uv` creates the tool environment with Python 3.12. > [!NOTE] -> For macOS user: due to macOS security checks, the first run of the `astrbot` command may take longer (about 10-20s). +> For macOS users: due to macOS security checks, the first run of the `astrbot` command may take longer (about 10-20s). Update `astrbot`: @@ -101,7 +101,7 @@ uv tool upgrade astrbot --python 3.12 For users familiar with containers and looking for a more stable, production-ready deployment method, we recommend deploying AstrBot with Docker / Docker Compose. -Please refer to the official documentation: [Deploy AstrBot with Docker](https://astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot). +Please refer to the official documentation: [Deploy AstrBot with Docker](https://docs.astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot). ### Deploy on RainYun @@ -139,7 +139,7 @@ yay -S astrbot-git **More deployment methods** -If you need panel-based management or deeper customization, see [BT-Panel Deployment](https://astrbot.app/deploy/astrbot/btpanel.html) for BT Panel app-store setup, [1Panel Deployment](https://astrbot.app/deploy/astrbot/1panel.html) for 1Panel app-market deployment, [CasaOS Deployment](https://astrbot.app/deploy/astrbot/casaos.html) for NAS/home-server visual deployment, and [Manual Deployment](https://astrbot.app/deploy/astrbot/cli.html) for fully custom source-based installation with `uv`. +If you need panel-based management or deeper customization, see [BT-Panel Deployment](https://docs.astrbot.app/deploy/astrbot/btpanel.html) for BT Panel app-store setup, [1Panel Deployment](https://docs.astrbot.app/deploy/astrbot/1panel.html) for 1Panel app-market deployment, [CasaOS Deployment](https://docs.astrbot.app/deploy/astrbot/casaos.html) for NAS/home-server visual deployment, and [Manual Deployment](https://docs.astrbot.app/deploy/astrbot/cli.html) for fully custom source-based installation with `uv`. ## Supported Messaging Platforms diff --git a/README_fr.md b/README_fr.md index c7dbeac3c..658a047cd 100644 --- a/README_fr.md +++ b/README_fr.md @@ -100,7 +100,7 @@ uv tool upgrade astrbot --python 3.12 Pour les utilisateurs familiers avec les conteneurs et qui souhaitent une méthode plus stable et adaptée à la production, nous recommandons de déployer AstrBot avec Docker / Docker Compose. -Veuillez consulter la documentation officielle [Déployer AstrBot avec Docker](https://astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot). +Veuillez consulter la documentation officielle [Déployer AstrBot avec Docker](https://docs.astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot). ### Déployer sur RainYun @@ -138,7 +138,7 @@ yay -S astrbot-git **Autres méthodes de déploiement** -Si vous avez besoin d'une gestion par panneau ou d'une personnalisation plus poussée, consultez [Déploiement BT-Panel](https://astrbot.app/deploy/astrbot/btpanel.html) pour une installation via BT Panel, [Déploiement 1Panel](https://astrbot.app/deploy/astrbot/1panel.html) pour le marketplace 1Panel, [Déploiement CasaOS](https://astrbot.app/deploy/astrbot/casaos.html) pour un déploiement visuel sur NAS/serveur domestique, et [Déploiement manuel](https://astrbot.app/deploy/astrbot/cli.html) pour une installation complète depuis les sources avec `uv`. +Si vous avez besoin d'une gestion par panneau ou d'une personnalisation plus poussée, consultez [Déploiement BT-Panel](https://docs.astrbot.app/deploy/astrbot/btpanel.html) pour une installation via BT Panel, [Déploiement 1Panel](https://docs.astrbot.app/deploy/astrbot/1panel.html) pour le marketplace 1Panel, [Déploiement CasaOS](https://docs.astrbot.app/deploy/astrbot/casaos.html) pour un déploiement visuel sur NAS/serveur domestique, et [Déploiement manuel](https://docs.astrbot.app/deploy/astrbot/cli.html) pour une installation complète depuis les sources avec `uv`. ## Plateformes de messagerie prises en charge diff --git a/README_ja.md b/README_ja.md index 1e91a0ebc..d806b8a97 100644 --- a/README_ja.md +++ b/README_ja.md @@ -100,7 +100,7 @@ uv tool upgrade astrbot --python 3.12 コンテナ運用に慣れており、より安定した本番向けのデプロイ方法を求めるユーザーには、Docker / Docker Compose での AstrBot デプロイをおすすめします。 -公式ドキュメント [Docker を使用した AstrBot のデプロイ](https://astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot) をご参照ください。 +公式ドキュメント [Docker を使用した AstrBot のデプロイ](https://docs.astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot) をご参照ください。 ### 雨云でのデプロイ @@ -138,7 +138,7 @@ yay -S astrbot-git **その他のデプロイ方法** -パネル操作での導入やより高度なカスタマイズが必要な場合は、[宝塔パネルデプロイ](https://astrbot.app/deploy/astrbot/btpanel.html)(BT Panel 経由の導入)、[1Panel デプロイ](https://astrbot.app/deploy/astrbot/1panel.html)(1Panel アプリマーケット経由)、[CasaOS デプロイ](https://astrbot.app/deploy/astrbot/casaos.html)(NAS / ホームサーバー向け可視化導入)、[手動デプロイ](https://astrbot.app/deploy/astrbot/cli.html)(`uv` とソースベースのフルカスタム導入)を参照してください。 +パネル操作での導入やより高度なカスタマイズが必要な場合は、[宝塔パネルデプロイ](https://docs.astrbot.app/deploy/astrbot/btpanel.html)(BT Panel 経由の導入)、[1Panel デプロイ](https://docs.astrbot.app/deploy/astrbot/1panel.html)(1Panel アプリマーケット経由)、[CasaOS デプロイ](https://docs.astrbot.app/deploy/astrbot/casaos.html)(NAS / ホームサーバー向け可視化導入)、[手動デプロイ](https://docs.astrbot.app/deploy/astrbot/cli.html)(`uv` とソースベースのフルカスタム導入)を参照してください。 ## サポートされているメッセージプラットフォーム diff --git a/README_ru.md b/README_ru.md index a1fb90678..d40a21357 100644 --- a/README_ru.md +++ b/README_ru.md @@ -100,7 +100,7 @@ uv tool upgrade astrbot --python 3.12 Для пользователей, знакомых с контейнерами и которым нужен более стабильный и подходящий для production способ, мы рекомендуем разворачивать AstrBot через Docker / Docker Compose. -См. официальную документацию [Развёртывание AstrBot с Docker](https://astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot). +См. официальную документацию [Развёртывание AstrBot с Docker](https://docs.astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot). ### Развёртывание на RainYun @@ -138,7 +138,7 @@ yay -S astrbot-git **Другие способы развёртывания** -Если вам нужна панельная установка или более глубокая кастомизация, смотрите [Развёртывание BT-Panel](https://astrbot.app/deploy/astrbot/btpanel.html) (установка через BT Panel), [Развёртывание 1Panel](https://astrbot.app/deploy/astrbot/1panel.html) (развёртывание через маркетплейс 1Panel), [Развёртывание CasaOS](https://astrbot.app/deploy/astrbot/casaos.html) (визуальный вариант для NAS и домашних серверов) и [Ручное развёртывание](https://astrbot.app/deploy/astrbot/cli.html) (полностью настраиваемая установка из исходников через `uv`). +Если вам нужна панельная установка или более глубокая кастомизация, смотрите [Развёртывание BT-Panel](https://docs.astrbot.app/deploy/astrbot/btpanel.html) (установка через BT Panel), [Развёртывание 1Panel](https://docs.astrbot.app/deploy/astrbot/1panel.html) (развёртывание через маркетплейс 1Panel), [Развёртывание CasaOS](https://docs.astrbot.app/deploy/astrbot/casaos.html) (визуальный вариант для NAS и домашних серверов) и [Ручное развёртывание](https://docs.astrbot.app/deploy/astrbot/cli.html) (полностью настраиваемая установка из исходников через `uv`). ## Поддерживаемые платформы обмена сообщениями diff --git a/README_zh-TW.md b/README_zh-TW.md index da5f3abf4..6edc905e3 100644 --- a/README_zh-TW.md +++ b/README_zh-TW.md @@ -32,7 +32,7 @@ 文件 | Blog | 路線圖 | -問題回報 +問題回報 | Email @@ -100,7 +100,7 @@ uv tool upgrade astrbot --python 3.12 對於熟悉容器、希望獲得更穩定且更適合正式環境部署方式的使用者,我們推薦使用 Docker / Docker Compose 部署 AstrBot。 -請參考官方文件 [使用 Docker 部署 AstrBot](https://astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot)。 +請參考官方文件 [使用 Docker 部署 AstrBot](https://docs.astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot)。 ### 在雨雲上部署 @@ -138,7 +138,7 @@ yay -S astrbot-git **更多部署方式** -若你需要面板化或更高自訂程度的部署,可參考 [寶塔面板](https://astrbot.app/deploy/astrbot/btpanel.html)(BT Panel 應用商店安裝)、[1Panel](https://astrbot.app/deploy/astrbot/1panel.html)(1Panel 應用商店安裝)、[CasaOS](https://astrbot.app/deploy/astrbot/casaos.html)(NAS / 家用伺服器可視化部署)與 [手動部署](https://astrbot.app/deploy/astrbot/cli.html)(基於原始碼與 `uv` 的完整自訂安裝)。 +若你需要面板化或更高自訂程度的部署,可參考 [寶塔面板](https://docs.astrbot.app/deploy/astrbot/btpanel.html)(BT Panel 應用商店安裝)、[1Panel](https://docs.astrbot.app/deploy/astrbot/1panel.html)(1Panel 應用商店安裝)、[CasaOS](https://docs.astrbot.app/deploy/astrbot/casaos.html)(NAS / 家用伺服器可視化部署)與 [手動部署](https://docs.astrbot.app/deploy/astrbot/cli.html)(基於原始碼與 `uv` 的完整自訂安裝)。 ## 支援的訊息平台 @@ -160,7 +160,7 @@ yay -S astrbot-git | KOOK | 官方維護 | | Misskey | 官方維護 | | Mattermost | 官方維護 | -| Whatsapp(即將支援) | 官方維護 | +| WhatsApp(即將支援) | 官方維護 | | [Matrix](https://github.com/stevessr/astrbot_plugin_matrix_adapter) | 社群維護 | | [Rocket.Chat](https://github.com/NET-Homeless/astrbot_plugin_rocket_chat_adapter) | 社群維護 | | [VoceChat](https://github.com/HikariFroya/astrbot_plugin_vocechat) | 社群維護 | diff --git a/README_zh.md b/README_zh.md index 44ffa13a9..f926227ab 100644 --- a/README_zh.md +++ b/README_zh.md @@ -31,12 +31,12 @@ 文档 | 博客 | 路线图 | -问题提交 +问题提交 | Email -AstrBot 是一个开源的一站式 Agentic 个人和群聊助手,可在 QQ、Telegram、企业微信、飞书、钉钉、Slack、等数十款主流即时通讯软件上部署,此外还内置类似 OpenWebUI 的轻量化 ChatUI,为个人、开发者和团队打造可靠、可扩展的对话式智能基础设施。无论是个人 AI 伙伴、智能客服、自动化助手,还是企业知识库,AstrBot 都能在你的即时通讯软件平台的工作流中快速构建 AI 应用。 +AstrBot 是一个开源的一站式 Agentic 个人和群聊助手,可在 QQ、Telegram、企业微信、飞书、钉钉、Slack 等数十款主流即时通讯软件上部署,此外还内置类似 OpenWebUI 的轻量化 ChatUI,为个人、开发者和团队打造可靠、可扩展的对话式智能基础设施。无论是个人 AI 伙伴、智能客服、自动化助手,还是企业知识库,AstrBot 都能在你的即时通讯软件平台的工作流中快速构建 AI 应用。 ![landingpage](https://github.com/user-attachments/assets/45fc5699-cddf-4e21-af35-13040706f6c0) @@ -103,7 +103,7 @@ uv tool upgrade astrbot --python 3.12 对于熟悉容器、希望获得更稳定且更适合生产环境部署方式的用户,我们推荐使用 Docker / Docker Compose 部署 AstrBot。 -请参考官方文档 [使用 Docker 部署 AstrBot](https://astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot)。 +请参考官方文档 [使用 Docker 部署 AstrBot](https://docs.astrbot.app/deploy/astrbot/docker.html#%E4%BD%BF%E7%94%A8-docker-%E9%83%A8%E7%BD%B2-astrbot)。 ### 在 雨云 上部署 @@ -141,7 +141,7 @@ yay -S astrbot-git **更多部署方式** -若你需要面板化或更高自定义部署,可参考 [宝塔面板](https://astrbot.app/deploy/astrbot/btpanel.html)(BT Panel 应用商店安装)、[1Panel](https://astrbot.app/deploy/astrbot/1panel.html)(1Panel 应用商店安装)、[CasaOS](https://astrbot.app/deploy/astrbot/casaos.html)(NAS / 家庭服务器可视化部署)和 [手动部署](https://astrbot.app/deploy/astrbot/cli.html)(基于源码与 `uv` 的完整自定义安装)。 +若你需要面板化或更高自定义部署,可参考 [宝塔面板](https://docs.astrbot.app/deploy/astrbot/btpanel.html)(BT Panel 应用商店安装)、[1Panel](https://docs.astrbot.app/deploy/astrbot/1panel.html)(1Panel 应用商店安装)、[CasaOS](https://docs.astrbot.app/deploy/astrbot/casaos.html)(NAS / 家庭服务器可视化部署)和 [手动部署](https://docs.astrbot.app/deploy/astrbot/cli.html)(基于源码与 `uv` 的完整自定义安装)。 ## 支持的消息平台 @@ -163,7 +163,7 @@ yay -S astrbot-git | **KOOK** | 官方维护 | | **Misskey** | 官方维护 | | **Mattermost** | 官方维护 | -| **Whatsapp (将支持)** | 官方维护 | +| **WhatsApp(将支持)** | 官方维护 | | [**Matrix**](https://github.com/stevessr/astrbot_plugin_matrix_adapter) | 社区维护 | | [**Rocket.Chat**](https://github.com/NET-Homeless/astrbot_plugin_rocket_chat_adapter) | 社区维护 | | [**VoceChat**](https://github.com/HikariFroya/astrbot_plugin_vocechat) | 社区维护 | diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index 8f5cb8998..76a2b846a 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -1180,7 +1180,7 @@ CONFIG_METADATA_2: Any = { "provider_type": "chat_completion", "enable": True, "key": [], - "api_base": "https://api.kimi.com/coding/", + "api_base": "https://api.kimi.com/coding", "timeout": 120, "proxy": "", "custom_headers": {"User-Agent": "claude-code/0.1.0"}, @@ -1210,6 +1210,19 @@ CONFIG_METADATA_2: Any = { "proxy": "", "custom_headers": {}, }, + "MiniMax Token Plan": { + "id": "minimax-token-plan", + "provider": "minimax-token-plan", + "type": "minimax_token_plan", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "api_base": "https://api.minimaxi.com/anthropic", + "timeout": 120, + "proxy": "", + "custom_headers": {"User-Agent": "claude-code/0.1.0"}, + "anth_thinking_config": {"type": "", "budget": 0, "effort": ""}, + }, "xAI": { "id": "xai", "provider": "xai", diff --git a/astrbot/core/cron/manager.py b/astrbot/core/cron/manager.py index c7597018f..59898d03c 100644 --- a/astrbot/core/cron/manager.py +++ b/astrbot/core/cron/manager.py @@ -23,6 +23,12 @@ if TYPE_CHECKING: from astrbot.core.star.context import Context +class CronJobSchedulingError(Exception): + """Raised when a cron job fails to be scheduled.""" + + pass + + class CronJobManager: """Central scheduler for BasicCronJob and ActiveAgentCronJob.""" @@ -60,7 +66,10 @@ class CronJobManager: job.job_id, ) continue - self._schedule_job(job) + try: + self._schedule_job(job) + except CronJobSchedulingError: + continue # Error already logged in _schedule_job async def add_basic_job( self, @@ -201,12 +210,15 @@ class CronJobManager: next_run_time=self._get_next_run_time(job.job_id), ), ) - except Exception as e: - logger.error(f"Failed to schedule cron job {job.job_id}: {e!s}") + except (ValueError, TypeError) as e: + logger.exception("Failed to schedule cron job %s", job.job_id) + raise CronJobSchedulingError(str(e)) from e def _get_next_run_time(self, job_id: str): aps_job = self.scheduler.get_job(job_id) - return aps_job.next_run_time if aps_job else None + if not aps_job or aps_job.next_run_time is None: + return None + return aps_job.next_run_time.astimezone(timezone.utc) async def _run_job(self, job_id: str) -> None: job = await self.db.get_cron_job(job_id) diff --git a/astrbot/core/db/vec_db/faiss_impl/document_storage.py b/astrbot/core/db/vec_db/faiss_impl/document_storage.py index 93b298411..8eb0f48af 100644 --- a/astrbot/core/db/vec_db/faiss_impl/document_storage.py +++ b/astrbot/core/db/vec_db/faiss_impl/document_storage.py @@ -3,8 +3,9 @@ import os from collections.abc import AsyncIterator from contextlib import asynccontextmanager from datetime import datetime +from pathlib import Path -from sqlalchemy import Column, Text +from sqlalchemy import Column, Text, bindparam from sqlalchemy.ext.asyncio import ( AsyncEngine, AsyncSession, @@ -14,6 +15,14 @@ from sqlalchemy.ext.asyncio import ( from sqlmodel import Field, MetaData, SQLModel, col, func, select, text from astrbot.core import logger +from astrbot.core.knowledge_base.retrieval.tokenizer import ( + build_fts5_or_query, + load_stopwords, + to_fts5_search_text, +) + +FTS_TABLE_NAME = "documents_fts" +FTS_REBUILD_BATCH_SIZE = 1000 class BaseDocModel(SQLModel, table=False): @@ -47,6 +56,10 @@ class DocumentStorage: os.path.dirname(__file__), "sqlite_init.sql", ) + self.fts5_available = False + self._fts_contentless_delete = False + self._fts_index_ready = False + self._stopwords: set[str] | None = None async def initialize(self) -> None: """Initialize the SQLite database and create the documents table if it doesn't exist.""" @@ -84,8 +97,49 @@ class DocumentStorage: except BaseException: pass + await self._initialize_fts5(conn) await conn.commit() + async def _initialize_fts5(self, executor) -> None: + try: + try: + await executor.execute( + text( + f""" + CREATE VIRTUAL TABLE IF NOT EXISTS {FTS_TABLE_NAME} + USING fts5( + search_text, + content='', + contentless_delete=1, + tokenize='unicode61' + ) + """, + ), + ) + self._fts_contentless_delete = True + except Exception: + await executor.execute( + text( + f""" + CREATE VIRTUAL TABLE IF NOT EXISTS {FTS_TABLE_NAME} + USING fts5( + search_text, + content='', + tokenize='unicode61' + ) + """, + ), + ) + self._fts_contentless_delete = False + self.fts5_available = True + except Exception as e: + self.fts5_available = False + self._fts_contentless_delete = False + logger.warning( + f"SQLite FTS5 is unavailable for document storage {self.db_path}; " + f"falling back to in-memory BM25 sparse retrieval: {e}", + ) + async def connect(self) -> None: """Connect to the SQLite database.""" if self.engine is None: @@ -108,6 +162,18 @@ class DocumentStorage: async with self.async_session_maker() as session: yield session + @property + def stopwords(self) -> set[str]: + if self._stopwords is None: + stopwords_path = ( + Path(__file__).parents[3] + / "knowledge_base" + / "retrieval" + / "hit_stopwords.txt" + ) + self._stopwords = load_stopwords(stopwords_path) + return self._stopwords + async def get_documents( self, metadata_filters: dict, @@ -181,6 +247,8 @@ class DocumentStorage: session.add(document) await session.flush() # Flush to get the ID assert document.id is not None, "Inserted document ID was not generated." + if document.id is not None: + await self._insert_fts_row(session, int(document.id), text) return document.id async def insert_documents_batch( @@ -224,6 +292,7 @@ class DocumentStorage: "Inserted document ID was not generated." ) document_ids.append(document.id) + await self._insert_fts_rows_batch(session, documents, texts) return document_ids async def delete_document_by_doc_id(self, doc_id: str) -> None: @@ -241,6 +310,8 @@ class DocumentStorage: document = result.scalar_one_or_none() if document: + if document.id is not None: + await self._delete_fts_row(session, int(document.id), document.text) await session.delete(document) async def get_document_by_doc_id(self, doc_id: str): @@ -280,9 +351,13 @@ class DocumentStorage: document = result.scalar_one_or_none() if document: + if document.id is not None: + await self._delete_fts_row(session, int(document.id), document.text) document.text = new_text document.updated_at = datetime.now() session.add(document) + if document.id is not None: + await self._insert_fts_row(session, int(document.id), new_text) async def delete_documents(self, metadata_filters: dict) -> None: """Delete documents by their metadata filters. @@ -308,6 +383,7 @@ class DocumentStorage: result = await session.execute(query) documents = result.scalars().all() + await self._delete_fts_rows_batch(session, documents) for doc in documents: await session.delete(doc) @@ -338,6 +414,286 @@ class DocumentStorage: count = result.scalar_one_or_none() return count if count is not None else 0 + async def ensure_fts_index(self) -> bool: + """Ensure the FTS5 sparse index exists and matches the documents table.""" + if not self.fts5_available: + return False + if self._fts_index_ready: + return True + + assert self.engine is not None, "Database connection is not initialized." + + async with self.get_session() as session: + doc_count = await self._count_documents_in_session(session) + fts_count = await self._count_fts_rows(session) + if doc_count == fts_count: + self._fts_index_ready = True + return True + + logger.info( + f"Rebuilding FTS5 sparse index for {self.db_path}: " + f"documents={doc_count}, fts_rows={fts_count}", + ) + await self.rebuild_fts_index() + return self.fts5_available + + async def rebuild_fts_index(self) -> None: + """Rebuild the contentless FTS5 sparse index from documents.""" + if not self.fts5_available: + return + + assert self.engine is not None, "Database connection is not initialized." + + async with self.get_session() as session, session.begin(): + await session.execute(text(f"DROP TABLE IF EXISTS {FTS_TABLE_NAME}")) + await self._initialize_fts5(session) + if not self.fts5_available: + return + + last_id = 0 + while True: + query = ( + select(Document) + .where(col(Document.id) > last_id) + .order_by(col(Document.id)) + .limit(FTS_REBUILD_BATCH_SIZE) + ) + result = await session.execute(query) + documents = result.scalars().all() + if not documents: + break + + await self._insert_fts_rows_batch( + session, + documents, + [doc.text for doc in documents], + ) + last_id = int(documents[-1].id or last_id) + + self._fts_index_ready = True + + async def search_sparse( + self, + query_tokens: list[str], + limit: int, + ) -> list[dict] | None: + """Search chunks using the FTS5 sparse index. + + Returns None when FTS5 is unavailable so callers can fall back to another + sparse retrieval implementation. + """ + if limit <= 0: + return [] + if not await self.ensure_fts_index(): + return None + + match_query = build_fts5_or_query(query_tokens) + if not match_query: + return [] + + async with self.get_session() as session: + try: + result = await session.execute( + text( + f""" + SELECT + d.id AS id, + d.doc_id AS doc_id, + d.text AS text, + d.metadata AS metadata, + d.created_at AS created_at, + d.updated_at AS updated_at, + bm25({FTS_TABLE_NAME}) AS score + FROM {FTS_TABLE_NAME} + JOIN documents d ON d.id = {FTS_TABLE_NAME}.rowid + WHERE {FTS_TABLE_NAME} MATCH :query + ORDER BY score ASC, d.id ASC + LIMIT :limit + """, + ), + {"query": match_query, "limit": int(limit)}, + ) + except Exception as e: + logger.warning( + f"FTS5 sparse search failed for {self.db_path}; " + f"falling back to in-memory BM25: {e}", + ) + self.fts5_available = False + return None + + rows = result.mappings().all() + return [ + { + "id": row["id"], + "doc_id": row["doc_id"], + "text": row["text"], + "metadata": row["metadata"], + "created_at": row["created_at"], + "updated_at": row["updated_at"], + "score": float(row["score"]), + } + for row in rows + ] + + async def _count_documents_in_session(self, session: AsyncSession) -> int: + result = await session.execute(select(func.count(col(Document.id)))) + count = result.scalar_one_or_none() + return int(count or 0) + + async def _count_fts_rows(self, session: AsyncSession) -> int: + result = await session.execute( + text(f"SELECT count(*) FROM {FTS_TABLE_NAME}"), + ) + count = result.scalar_one_or_none() + return int(count or 0) + + async def _insert_fts_row( + self, + session: AsyncSession, + rowid: int, + content: str, + ) -> None: + if not self.fts5_available: + return + + search_text = to_fts5_search_text(content, self.stopwords) + await session.execute( + text( + f""" + INSERT INTO {FTS_TABLE_NAME}(rowid, search_text) + VALUES (:rowid, :search_text) + """, + ), + {"rowid": rowid, "search_text": search_text}, + ) + + async def _insert_fts_rows_batch( + self, + session: AsyncSession, + documents: list[Document], + contents: list[str], + ) -> None: + if not self.fts5_available: + return + + fts_params = [ + { + "rowid": int(doc.id), + "search_text": to_fts5_search_text(content, self.stopwords), + } + for doc, content in zip(documents, contents) + if doc.id is not None + ] + if not fts_params: + return + + await session.execute( + text( + f""" + INSERT INTO {FTS_TABLE_NAME}(rowid, search_text) + VALUES (:rowid, :search_text) + """, + ), + fts_params, + ) + + async def _delete_fts_row( + self, + session: AsyncSession, + rowid: int, + content: str, + ) -> None: + if not self.fts5_available: + return + + if self._fts_contentless_delete: + await session.execute( + text(f"DELETE FROM {FTS_TABLE_NAME} WHERE rowid = :rowid"), + {"rowid": rowid}, + ) + return + + if not await self._fts_row_exists(session, rowid): + return + + search_text = to_fts5_search_text(content, self.stopwords) + await session.execute( + text( + f""" + INSERT INTO {FTS_TABLE_NAME}({FTS_TABLE_NAME}, rowid, search_text) + VALUES ('delete', :rowid, :search_text) + """, + ), + {"rowid": rowid, "search_text": search_text}, + ) + + async def _delete_fts_rows_batch( + self, + session: AsyncSession, + documents: list[Document], + ) -> None: + if not self.fts5_available: + return + + docs_with_ids = [doc for doc in documents if doc.id is not None] + if not docs_with_ids: + return + + if self._fts_contentless_delete: + await session.execute( + text(f"DELETE FROM {FTS_TABLE_NAME} WHERE rowid = :rowid"), + [{"rowid": int(doc.id)} for doc in docs_with_ids if doc.id is not None], + ) + return + + existing_rowids = await self._existing_fts_rowids( + session, + [int(doc.id) for doc in docs_with_ids if doc.id is not None], + ) + fts_params = [ + { + "rowid": int(doc.id), + "search_text": to_fts5_search_text(doc.text, self.stopwords), + } + for doc in docs_with_ids + if doc.id is not None and int(doc.id) in existing_rowids + ] + if not fts_params: + return + + await session.execute( + text( + f""" + INSERT INTO {FTS_TABLE_NAME}({FTS_TABLE_NAME}, rowid, search_text) + VALUES ('delete', :rowid, :search_text) + """, + ), + fts_params, + ) + + async def _fts_row_exists(self, session: AsyncSession, rowid: int) -> bool: + result = await session.execute( + text(f"SELECT 1 FROM {FTS_TABLE_NAME} WHERE rowid = :rowid LIMIT 1"), + {"rowid": rowid}, + ) + return result.scalar_one_or_none() is not None + + async def _existing_fts_rowids( + self, + session: AsyncSession, + rowids: list[int], + ) -> set[int]: + if not rowids: + return set() + + result = await session.execute( + text( + f"SELECT rowid FROM {FTS_TABLE_NAME} WHERE rowid IN :rowids" + ).bindparams(bindparam("rowids", expanding=True)), + {"rowids": rowids}, + ) + return {int(row[0]) for row in result.fetchall()} + async def get_user_ids(self) -> list[str]: """Retrieve all user IDs from the documents table. diff --git a/astrbot/core/knowledge_base/retrieval/__init__.py b/astrbot/core/knowledge_base/retrieval/__init__.py index f5d196cb9..b7c88075d 100644 --- a/astrbot/core/knowledge_base/retrieval/__init__.py +++ b/astrbot/core/knowledge_base/retrieval/__init__.py @@ -1,8 +1,11 @@ """检索模块""" -from .manager import RetrievalManager, RetrievalResult -from .rank_fusion import FusedResult, RankFusion -from .sparse_retriever import SparseResult, SparseRetriever +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from .manager import RetrievalManager, RetrievalResult + from .rank_fusion import FusedResult, RankFusion + from .sparse_retriever import SparseResult, SparseRetriever __all__ = [ "FusedResult", @@ -12,3 +15,31 @@ __all__ = [ "SparseResult", "SparseRetriever", ] + + +def __getattr__(name: str): + if name in {"RetrievalManager", "RetrievalResult"}: + from .manager import RetrievalManager, RetrievalResult + + return { + "RetrievalManager": RetrievalManager, + "RetrievalResult": RetrievalResult, + }[name] + + if name in {"FusedResult", "RankFusion"}: + from .rank_fusion import FusedResult, RankFusion + + return { + "FusedResult": FusedResult, + "RankFusion": RankFusion, + }[name] + + if name in {"SparseResult", "SparseRetriever"}: + from .sparse_retriever import SparseResult, SparseRetriever + + return { + "SparseResult": SparseResult, + "SparseRetriever": SparseRetriever, + }[name] + + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/astrbot/core/knowledge_base/retrieval/sparse_retriever.py b/astrbot/core/knowledge_base/retrieval/sparse_retriever.py index 9d485dab2..f06eb5090 100644 --- a/astrbot/core/knowledge_base/retrieval/sparse_retriever.py +++ b/astrbot/core/knowledge_base/retrieval/sparse_retriever.py @@ -8,10 +8,13 @@ import os from dataclasses import dataclass from typing import TYPE_CHECKING -import jieba from rank_bm25 import BM25Okapi from astrbot.core.knowledge_base.kb_db_sqlite import KBSQLiteDatabase +from astrbot.core.knowledge_base.retrieval.tokenizer import ( + load_stopwords, + tokenize_text, +) if TYPE_CHECKING: from astrbot.core.db.vec_db.faiss_impl import FaissVecDB @@ -47,13 +50,9 @@ class SparseRetriever: self.kb_db = kb_db self._index_cache = {} # 缓存 BM25 索引 - with open( + self.hit_stopwords = load_stopwords( os.path.join(os.path.dirname(__file__), "hit_stopwords.txt"), - encoding="utf-8", - ) as f: - self.hit_stopwords = { - word.strip() for word in set(f.read().splitlines()) if word.strip() - } + ) async def retrieve( self, @@ -72,7 +71,52 @@ class SparseRetriever: List[SparseResult]: 检索结果列表 """ - # 1. 获取所有相关块 + fts_results = [] + fallback_kb_ids = [] + query_tokens = tokenize_text(query, self.hit_stopwords) + for kb_id in kb_ids: + vec_db: FaissVecDB | None = kb_options.get(kb_id, {}).get("vec_db") + if not vec_db: + continue + top_k_sparse = kb_options.get(kb_id, {}).get("top_k_sparse", 50) + result = await vec_db.document_storage.search_sparse( + query_tokens=query_tokens, + limit=top_k_sparse, + ) + if result is None: + fallback_kb_ids.append(kb_id) + continue + + for doc in result: + chunk_md = json.loads(doc["metadata"]) + fts_results.append( + SparseResult( + chunk_id=doc["doc_id"], + chunk_index=chunk_md["chunk_index"], + doc_id=chunk_md["kb_doc_id"], + kb_id=kb_id, + content=doc["text"], + score=-float(doc["score"]), + ), + ) + + fallback_results = [] + if fallback_kb_ids: + fallback_results = await self._retrieve_with_bm25( + query=query, + kb_ids=fallback_kb_ids, + kb_options=kb_options, + ) + results = fts_results + fallback_results + results.sort(key=lambda x: x.score, reverse=True) + return results + + async def _retrieve_with_bm25( + self, + query: str, + kb_ids: list[str], + kb_options: dict, + ) -> list[SparseResult]: top_k_sparse = 0 chunks = [] for kb_id in kb_ids: @@ -103,20 +147,13 @@ class SparseRetriever: # 2. 准备文档和索引 corpus = [chunk["text"] for chunk in chunks] - tokenized_corpus = [list(jieba.cut(doc)) for doc in corpus] - tokenized_corpus = [ - [word for word in doc if word not in self.hit_stopwords] - for doc in tokenized_corpus - ] + tokenized_corpus = [tokenize_text(doc, self.hit_stopwords) for doc in corpus] # 3. 构建 BM25 索引 bm25 = BM25Okapi(tokenized_corpus) # 4. 执行检索 - tokenized_query = list(jieba.cut(query)) - tokenized_query = [ - word for word in tokenized_query if word not in self.hit_stopwords - ] + tokenized_query = tokenize_text(query, self.hit_stopwords) scores = bm25.get_scores(tokenized_query) # 5. 排序并返回 Top-K diff --git a/astrbot/core/knowledge_base/retrieval/tokenizer.py b/astrbot/core/knowledge_base/retrieval/tokenizer.py new file mode 100644 index 000000000..1f4f07bb9 --- /dev/null +++ b/astrbot/core/knowledge_base/retrieval/tokenizer.py @@ -0,0 +1,39 @@ +"""Tokenization helpers shared by sparse retrieval indexes.""" + +import re +from pathlib import Path +from re import Pattern + +import jieba + +_TERM_PATTERN: Pattern[str] = re.compile(r"\w", re.UNICODE) + + +def load_stopwords(path: Path | str) -> set[str]: + with Path(path).open(encoding="utf-8") as f: + return {word.strip() for word in set(f.read().splitlines()) if word.strip()} + + +def tokenize_text(text: str, stopwords: set[str]) -> list[str]: + tokens = [] + for token in jieba.cut(text or ""): + token = token.strip() + if not token or token in stopwords: + continue + if not _TERM_PATTERN.search(token): + continue + tokens.append(token) + return tokens + + +def to_fts5_search_text(text: str, stopwords: set[str]) -> str: + return " ".join(tokenize_text(text, stopwords)) + + +def quote_fts5_token(token: str) -> str: + return '"' + token.replace('"', '""') + '"' + + +def build_fts5_or_query(tokens: list[str]) -> str: + quoted_tokens = [quote_fts5_token(token) for token in tokens if token] + return " OR ".join(quoted_tokens) diff --git a/astrbot/core/pipeline/rate_limit_check/stage.py b/astrbot/core/pipeline/rate_limit_check/stage.py index 49ad2f56e..7d79d7b62 100644 --- a/astrbot/core/pipeline/rate_limit_check/stage.py +++ b/astrbot/core/pipeline/rate_limit_check/stage.py @@ -61,6 +61,8 @@ class RateLimitStage(Stage): timestamps = self.event_timestamps[session_id] self._remove_expired_timestamps(timestamps, now) + if self.rate_limit_count <= 0: + break if len(timestamps) < self.rate_limit_count: timestamps.append(now) break diff --git a/astrbot/core/platform/sources/telegram/tg_adapter.py b/astrbot/core/platform/sources/telegram/tg_adapter.py index 20607465e..65ced22cd 100644 --- a/astrbot/core/platform/sources/telegram/tg_adapter.py +++ b/astrbot/core/platform/sources/telegram/tg_adapter.py @@ -4,6 +4,7 @@ import re import sys import uuid +from apscheduler.events import EVENT_JOB_ERROR from apscheduler.schedulers.asyncio import AsyncIOScheduler from telegram import BotCommand, Update from telegram.constants import ChatType @@ -80,6 +81,15 @@ class TelegramPlatformAdapter(Platform): self.client = self.application.bot logger.debug(f"Telegram base url: {self.client.base_url}") self.scheduler = AsyncIOScheduler() + self.scheduler.add_listener( + lambda ev: logger.error( + "Scheduled job %s raised: %s", + ev.job_id, + ev.exception, + exc_info=ev.exception, + ), + EVENT_JOB_ERROR, + ) self._terminating = False raw_delay = self.config.get("telegram_polling_restart_delay", 5.0) try: @@ -513,22 +523,32 @@ class TelegramPlatformAdapter(Platform): logger.info( f"Processing media group {media_group_id}, total {len(updates_and_contexts)} items", ) - first_update, first_context = updates_and_contexts[0] - abm = await self.convert_message(first_update, first_context) - if not abm: - logger.warning( - f"Failed to convert the first message of media group {media_group_id}", + + try: + first_update, first_context = updates_and_contexts[0] + abm = await self.convert_message(first_update, first_context) + + if not abm: + logger.warning( + f"Failed to convert the first message of media group {media_group_id}" + ) + return + + for update, context in updates_and_contexts[1:]: + extra = await self.convert_message(update, context, get_reply=False) + if not extra: + continue + + abm.message.extend(extra.message) + logger.debug( + f"Added {len(extra.message)} components to media group {media_group_id}" + ) + + await self.handle_msg(abm) + except Exception: + logger.error( + f"Failed to process media group {media_group_id}", exc_info=True ) - return - for update, context in updates_and_contexts[1:]: - extra = await self.convert_message(update, context, get_reply=False) - if not extra: - continue - abm.message.extend(extra.message) - logger.debug( - f"Added {len(extra.message)} components to media group {media_group_id}", - ) - await self.handle_msg(abm) async def handle_msg(self, message: AstrBotMessage) -> None: message_event = TelegramPlatformEvent( diff --git a/astrbot/core/platform/sources/weixin_oc/weixin_oc_adapter.py b/astrbot/core/platform/sources/weixin_oc/weixin_oc_adapter.py index 0aee360d0..01994dd69 100644 --- a/astrbot/core/platform/sources/weixin_oc/weixin_oc_adapter.py +++ b/astrbot/core/platform/sources/weixin_oc/weixin_oc_adapter.py @@ -7,6 +7,8 @@ import io import sys import time import uuid +from collections import deque +from collections.abc import Mapping from dataclasses import dataclass, field from pathlib import Path from typing import Any @@ -107,6 +109,7 @@ class WeixinOCAdapter(Platform): self._sync_buf = "" self._qr_expired_count = 0 self._context_tokens: dict[str, str] = {} + self._context_tokens_dirty = False self._typing_states: dict[str, TypingSessionState] = {} self._last_inbound_error = "" self._typing_keepalive_interval_s = max( @@ -444,12 +447,31 @@ class WeixinOCAdapter(Platform): saved_base = str(self.config.get("weixin_oc_base_url", "")).strip() if saved_base: self.base_url = saved_base.rstrip("/") + raw_context_tokens = self.config.get("weixin_oc_context_tokens", {}) + if isinstance(raw_context_tokens, dict): + self._context_tokens = self._normalize_context_tokens(raw_context_tokens) + + def _normalize_context_tokens( + self, raw_context_tokens: Mapping[object, object] + ) -> dict[str, str]: + normalized_context_tokens: dict[str, str] = {} + for user_id, context_token in raw_context_tokens.items(): + normalized_user_id = str(user_id).strip() + normalized_context_token = str(context_token).strip() + if not normalized_user_id or not normalized_context_token: + continue + normalized_context_tokens[normalized_user_id] = normalized_context_token + return normalized_context_tokens async def _save_account_state(self) -> None: + normalized_context_tokens = self._normalize_context_tokens(self._context_tokens) self.config["weixin_oc_token"] = self.token or "" self.config["weixin_oc_account_id"] = self.account_id or "" self.config["weixin_oc_sync_buf"] = self._sync_buf self.config["weixin_oc_base_url"] = self.base_url + self.config["weixin_oc_context_tokens"] = normalized_context_tokens + + for platform in astrbot_config.get("platform", []): if not isinstance(platform, dict): continue @@ -461,9 +483,11 @@ class WeixinOCAdapter(Platform): platform["weixin_oc_account_id"] = self.account_id or "" platform["weixin_oc_sync_buf"] = self._sync_buf platform["weixin_oc_base_url"] = self.base_url + platform["weixin_oc_context_tokens"] = normalized_context_tokens break self._sync_client_state() astrbot_config.save_config() + self._context_tokens_dirty = False def _is_login_session_valid( self, @@ -1043,15 +1067,21 @@ class WeixinOCAdapter(Platform): self._last_inbound_error, ) return + + should_save_state = self._context_tokens_dirty if data.get("get_updates_buf"): self._sync_buf = str(data.get("get_updates_buf")) - await self._save_account_state() + should_save_state = True + + for msg in data.get("msgs", []) if isinstance(data.get("msgs"), list) else []: if self._shutdown_event.is_set(): return if not isinstance(msg, dict): continue await self._handle_inbound_message(msg) + if should_save_state: + await self._save_account_state() def _message_chain_to_text(self, message_chain: MessageChain) -> str: text = "" diff --git a/astrbot/core/provider/manager.py b/astrbot/core/provider/manager.py index 94bdbcd8d..d38427125 100644 --- a/astrbot/core/provider/manager.py +++ b/astrbot/core/provider/manager.py @@ -378,6 +378,10 @@ class ProviderManager: ) case "longcat_chat_completion": from .sources.longcat_source import ProviderLongCat as ProviderLongCat + case "minimax_token_plan": + from .sources.minimax_token_plan_source import ( + ProviderMiniMaxTokenPlan as ProviderMiniMaxTokenPlan, + ) case "zhipu_chat_completion": from .sources.zhipu_source import ProviderZhipu as ProviderZhipu case "groq_chat_completion": diff --git a/astrbot/core/provider/sources/anthropic_source.py b/astrbot/core/provider/sources/anthropic_source.py index 1ee910166..64e6b3d4c 100644 --- a/astrbot/core/provider/sources/anthropic_source.py +++ b/astrbot/core/provider/sources/anthropic_source.py @@ -304,7 +304,7 @@ class ProviderAnthropic(Provider): extra_body = self.provider_config.get("custom_extra_body", {}) if "max_tokens" not in payloads: - payloads["max_tokens"] = 1024 + payloads["max_tokens"] = 65536 self._apply_thinking_config(payloads) try: @@ -402,7 +402,7 @@ class ProviderAnthropic(Provider): reasoning_signature = "" if "max_tokens" not in payloads: - payloads["max_tokens"] = 1024 + payloads["max_tokens"] = 65536 self._apply_thinking_config(payloads) async with self.client.messages.stream( diff --git a/astrbot/core/provider/sources/minimax_token_plan_source.py b/astrbot/core/provider/sources/minimax_token_plan_source.py new file mode 100644 index 000000000..5a578424f --- /dev/null +++ b/astrbot/core/provider/sources/minimax_token_plan_source.py @@ -0,0 +1,55 @@ +from astrbot.core.provider.sources.anthropic_source import ProviderAnthropic + +from ..register import register_provider_adapter + +MINIMAX_TOKEN_PLAN_MODELS = [ + "MiniMax-M2.7", + "MiniMax-M2.5", + "MiniMax-M2.1", + "MiniMax-M2", +] + + +@register_provider_adapter( + "minimax_token_plan", + "MiniMax Token Plan Provider Adapter", +) +class ProviderMiniMaxTokenPlan(ProviderAnthropic): + """MiniMax Token Plan provider. + + The Token Plan API does not support the /models endpoint, so get_models() + returns a hard-coded model list. This is a Token Plan API limitation. + See https://github.com/AstrBotDevs/AstrBot/issues/7585 for details. + """ + + def __init__( + self, + provider_config, + provider_settings, + ) -> None: + # Keep api_base fixed; Token Plan users do not need to configure it. + provider_config["api_base"] = "https://api.minimaxi.com/anthropic" + # MiniMax Token Plan requires the Authorization: Bearer header. + key = provider_config.get("key", "") + actual_key = key[0] if isinstance(key, list) else key + provider_config.setdefault("custom_headers", {})["Authorization"] = ( + f"Bearer {actual_key}" + ) + + super().__init__( + provider_config, + provider_settings, + ) + + configured_model = provider_config.get("model", "MiniMax-M2.7") + if configured_model not in MINIMAX_TOKEN_PLAN_MODELS: + raise ValueError( + f"Unsupported model: {configured_model!r}. " + f"Supported models: {', '.join(MINIMAX_TOKEN_PLAN_MODELS)}" + ) + + self.set_model(configured_model) + + async def get_models(self) -> list[str]: + """Return the hard-coded known model list because Token Plan cannot fetch it dynamically.""" + return MINIMAX_TOKEN_PLAN_MODELS.copy() diff --git a/astrbot/core/star/updator.py b/astrbot/core/star/updator.py index 3abfa240b..7e548c3de 100644 --- a/astrbot/core/star/updator.py +++ b/astrbot/core/star/updator.py @@ -10,8 +10,8 @@ from astrbot.core.utils.io import on_error, remove_dir class PluginUpdator(RepoZipUpdator): - def __init__(self, repo_mirror: str = "") -> None: - super().__init__(repo_mirror) + def __init__(self, repo_mirror: str = "", verify: str | bool | None = None) -> None: + super().__init__(repo_mirror, verify=verify) self.plugin_store_path = get_astrbot_plugin_path() def get_plugin_store_path(self) -> str: diff --git a/astrbot/core/updator.py b/astrbot/core/updator.py index c6175d733..c42ed9f13 100644 --- a/astrbot/core/updator.py +++ b/astrbot/core/updator.py @@ -7,7 +7,6 @@ import psutil from astrbot.core import logger from astrbot.core.config.default import VERSION from astrbot.core.utils.astrbot_path import get_astrbot_path -from astrbot.core.utils.io import download_file from .zip_updator import ReleaseInfo, RepoZipUpdator @@ -18,8 +17,8 @@ class AstrBotUpdator(RepoZipUpdator): 功能包括检查更新、下载更新文件、解压缩更新文件等 """ - def __init__(self, repo_mirror: str = "") -> None: - super().__init__(repo_mirror) + def __init__(self, repo_mirror: str = "", verify: str | bool | None = None) -> None: + super().__init__(repo_mirror, verify=verify) self.MAIN_PATH = get_astrbot_path() self.ASTRBOT_RELEASE_API = "https://api.soulter.top/releases" @@ -182,8 +181,8 @@ class AstrBotUpdator(RepoZipUpdator): file_url = f"{proxy}/{file_url}" try: - await download_file(file_url, "temp.zip") - logger.info("下载 AstrBot Core 更新文件完成,正在执行解压...") + await self._download_file(file_url, "temp.zip") + logger.info("下载 AstrBot Core 更新文件完成,正在执行解压...") self.unzip_file("temp.zip", self.MAIN_PATH) except BaseException as e: raise e diff --git a/astrbot/core/zip_updator.py b/astrbot/core/zip_updator.py index bc15e45ea..41b2a0882 100644 --- a/astrbot/core/zip_updator.py +++ b/astrbot/core/zip_updator.py @@ -1,15 +1,15 @@ import os import re import shutil -import ssl import zipfile +from pathlib import Path from typing import NoReturn -import aiohttp import certifi +import httpx from astrbot.core import logger -from astrbot.core.utils.io import download_file, on_error +from astrbot.core.utils.io import on_error from astrbot.core.utils.version_comparator import VersionComparator @@ -33,36 +33,53 @@ class ReleaseInfo: class RepoZipUpdator: - def __init__(self, repo_mirror: str = "") -> None: + def __init__(self, repo_mirror: str = "", verify: str | bool | None = None) -> None: self.repo_mirror = repo_mirror self.rm_on_error = on_error + self.httpx_verify = certifi.where() if verify is None else verify + + def _create_httpx_client(self, timeout: float = 30.0) -> httpx.AsyncClient: + return httpx.AsyncClient( + follow_redirects=True, + timeout=timeout, + trust_env=True, + verify=self.httpx_verify, + ) + + @staticmethod + def _truncate_response_body(body: str, max_len: int = 1000) -> str: + if len(body) <= max_len: + return body + return body[:max_len] + "...[truncated]" + + async def _download_file( + self, url: str, path: str, timeout: float = 1800.0 + ) -> None: + target_path = Path(path) + target_path.parent.mkdir(parents=True, exist_ok=True) + + try: + async with self._create_httpx_client(timeout=timeout) as client: + async with client.stream("GET", url) as response: + response.raise_for_status() + with target_path.open("wb") as file: + async for chunk in response.aiter_bytes(8192): + file.write(chunk) + except Exception as e: + logger.error(f"下载文件失败: {url} -> {target_path}, 错误: {e}") + if self.rm_on_error and target_path.exists(): + target_path.unlink() + raise async def fetch_release_info(self, url: str, latest: bool = True) -> list: """请求版本信息。 返回一个列表,每个元素是一个字典,包含版本号、发布时间、更新内容、commit hash等信息。 """ try: - ssl_context = ssl.create_default_context( - cafile=certifi.where(), - ) # 新增:创建基于 certifi 的 SSL 上下文 - connector = aiohttp.TCPConnector( - ssl=ssl_context, - ) # 新增:使用 TCPConnector 指定 SSL 上下文 - async with ( - aiohttp.ClientSession( - trust_env=True, - connector=connector, - ) as session, - session.get(url) as response, - ): - # 检查 HTTP 状态码 - if response.status != 200: - text = await response.text() - logger.error( - f"请求 {url} 失败,状态码: {response.status}, 内容: {text}", - ) - raise Exception(f"请求失败,状态码: {response.status}") - result = await response.json() + async with self._create_httpx_client() as client: + response = await client.get(url) + response.raise_for_status() + result = response.json() if not result: return [] # if latest: @@ -80,9 +97,17 @@ class RepoZipUpdator: "zipball_url": release["zipball_url"], }, ) + except httpx.HTTPStatusError as e: + response_body = "" + if e.response is not None: + response_body = self._truncate_response_body(e.response.text) + logger.error( + f"请求 {url} 失败,状态码: {e.response.status_code}, 内容: {response_body}", + ) + raise Exception("解析版本信息失败") from e except Exception as e: logger.error(f"解析版本信息时发生异常: {e}") - raise Exception("解析版本信息失败") + raise Exception("解析版本信息失败") from e return ret def github_api_release_parser(self, releases: list) -> list: @@ -209,7 +234,7 @@ class RepoZipUpdator: f"检查到设置了镜像站,将使用镜像站下载 {author}/{repo} 仓库源码: {release_url}", ) - await download_file(release_url, target_path + ".zip") + await self._download_file(release_url, target_path + ".zip") def parse_github_url(self, url: str): """使用正则表达式解析 GitHub 仓库 URL,支持 `.git` 后缀和 `tree/branch` 结构 diff --git a/astrbot/dashboard/routes/cron.py b/astrbot/dashboard/routes/cron.py index 1644078f8..721cd8ee7 100644 --- a/astrbot/dashboard/routes/cron.py +++ b/astrbot/dashboard/routes/cron.py @@ -1,5 +1,5 @@ import traceback -from datetime import datetime +from datetime import datetime, timezone from quart import jsonify, request @@ -28,8 +28,12 @@ class CronRoute(Route): def _serialize_job(self, job) -> dict: data = job.model_dump() if hasattr(job, "model_dump") else job.__dict__ for k in ["created_at", "updated_at", "last_run_at", "next_run_time"]: - if isinstance(data.get(k), datetime): - data[k] = data[k].isoformat() + v = data.get(k) + if isinstance(v, datetime): + # Attach UTC + if v.tzinfo is None: + v = v.replace(tzinfo=timezone.utc) + data[k] = v.isoformat() # expose note explicitly for UI (prefer payload.note then description) payload = data.get("payload") or {} data["note"] = payload.get("note") or data.get("description") or "" diff --git a/dashboard/src/components/shared/ConfigItemRenderer.vue b/dashboard/src/components/shared/ConfigItemRenderer.vue index d1564a8a7..13324e009 100644 --- a/dashboard/src/components/shared/ConfigItemRenderer.vue +++ b/dashboard/src/components/shared/ConfigItemRenderer.vue @@ -203,7 +203,9 @@ @update:model-value="(val) => emitUpdate(toNumber(val))" /> None: + return None + + def json(self): + return self._payload + + +class _FakeStreamResponse: + def __init__(self, payload: bytes): + self._payload = payload + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb) -> None: + return None + + def raise_for_status(self) -> None: + return None + + async def aiter_bytes(self, chunk_size: int = 8192): + for start in range(0, len(self._payload), chunk_size): + yield self._payload[start : start + chunk_size] + + +class _FakeFailingStreamResponse: + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb) -> None: + return None + + def raise_for_status(self) -> None: + return None + + async def aiter_bytes(self, chunk_size: int = 8192): # noqa: ARG002 + yield b"partial" + raise RuntimeError("stream interrupted") + + +class _FakeStatusErrorResponse: + def __init__(self, status_code: int, body: str, url: str): + self._status_code = status_code + self._body = body + self._url = url + + def raise_for_status(self) -> None: + request = httpx.Request("GET", self._url) + response = httpx.Response( + self._status_code, + text=self._body, + request=request, + ) + raise httpx.HTTPStatusError( + "status error", + request=request, + response=response, + ) + + +@dataclass +class _FakeAsyncClientState: + json_payload: list[dict] = field(default_factory=list) + stream_payload: bytes = b"" + init_kwargs: dict | None = None + requested_urls: list[str] = field(default_factory=list) + stream_urls: list[str] = field(default_factory=list) + + +class _FakeStatusErrorAsyncClient: + def __init__(self, response: _FakeStatusErrorResponse): + self._response = response + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb) -> None: + return None + + async def get(self, url: str): + return self._response + + +class _FakeFailingStreamAsyncClient: + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb) -> None: + return None + + def stream(self, method: str, url: str): # noqa: ARG002 + return _FakeFailingStreamResponse() + + +def _build_fake_httpx_module(state: _FakeAsyncClientState) -> SimpleNamespace: + class _FakeAsyncClient: + def __init__(self, **kwargs): + state.init_kwargs = kwargs + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb) -> None: + return None + + async def get(self, url: str): + state.requested_urls.append(url) + return _FakeJSONResponse(state.json_payload) + + def stream(self, method: str, url: str): + assert method == "GET" + state.stream_urls.append(url) + return _FakeStreamResponse(state.stream_payload) + + return SimpleNamespace( + AsyncClient=_FakeAsyncClient, + HTTPStatusError=httpx.HTTPStatusError, + ) + + +@pytest.fixture +def fake_async_client_state() -> _FakeAsyncClientState: + return _FakeAsyncClientState() + + +@pytest.mark.asyncio +async def test_fetch_release_info_uses_httpx_client_with_env_proxy_support( + monkeypatch: pytest.MonkeyPatch, + fake_async_client_state: _FakeAsyncClientState, +) -> None: + import astrbot.core.zip_updator as zip_updator_module + + fake_async_client_state.json_payload = [ + { + "name": "AstrBot v4.23.2", + "published_at": "2026-04-16T00:00:00Z", + "body": "fix updater socks proxy support", + "tag_name": "v4.23.2", + "zipball_url": "https://example.com/astrbot.zip", + } + ] + + monkeypatch.setattr( + zip_updator_module, + "aiohttp", + SimpleNamespace( + ClientSession=lambda *args, **kwargs: (_ for _ in ()).throw( + AssertionError( + "fetch_release_info should not use aiohttp.ClientSession" + ) + ) + ), + raising=False, + ) + monkeypatch.setattr( + zip_updator_module, + "httpx", + _build_fake_httpx_module(fake_async_client_state), + raising=False, + ) + + release_info = await RepoZipUpdator().fetch_release_info( + "https://api.soulter.top/releases" + ) + + assert release_info == [ + { + "version": "AstrBot v4.23.2", + "published_at": "2026-04-16T00:00:00Z", + "body": "fix updater socks proxy support", + "tag_name": "v4.23.2", + "zipball_url": "https://example.com/astrbot.zip", + } + ] + assert fake_async_client_state.requested_urls == ["https://api.soulter.top/releases"] + assert fake_async_client_state.init_kwargs is not None + assert fake_async_client_state.init_kwargs["follow_redirects"] is True + assert fake_async_client_state.init_kwargs["timeout"] == 30.0 + assert fake_async_client_state.init_kwargs["trust_env"] is True + assert fake_async_client_state.init_kwargs["verify"] == certifi.where() + + +@pytest.mark.asyncio +async def test_download_from_repo_url_uses_httpx_stream_for_zip_download( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + fake_async_client_state: _FakeAsyncClientState, +) -> None: + import astrbot.core.zip_updator as zip_updator_module + + fake_async_client_state.stream_payload = b"zip-data" + + async def fake_fetch_release_info(self, url: str, latest: bool = True): # noqa: ARG001 + return [ + { + "version": "AstrBot v4.23.2", + "published_at": "2026-04-16T00:00:00Z", + "body": "fix updater socks proxy support", + "tag_name": "v4.23.2", + "zipball_url": "https://example.com/archive.zip", + } + ] + + monkeypatch.setattr(RepoZipUpdator, "fetch_release_info", fake_fetch_release_info) + monkeypatch.setattr( + zip_updator_module, + "download_file", + lambda *args, **kwargs: (_ for _ in ()).throw( + AssertionError("download_from_repo_url should not use aiohttp download_file") + ), + raising=False, + ) + monkeypatch.setattr( + zip_updator_module, + "httpx", + _build_fake_httpx_module(fake_async_client_state), + raising=False, + ) + + target_path = tmp_path / "AstrBot" + await RepoZipUpdator().download_from_repo_url( + str(target_path), + "https://github.com/AstrBotDevs/AstrBot", + ) + + assert (tmp_path / "AstrBot.zip").read_bytes() == b"zip-data" + assert fake_async_client_state.stream_urls == ["https://example.com/archive.zip"] + assert fake_async_client_state.init_kwargs is not None + assert fake_async_client_state.init_kwargs["follow_redirects"] is True + assert fake_async_client_state.init_kwargs["timeout"] == 1800.0 + assert fake_async_client_state.init_kwargs["trust_env"] is True + assert fake_async_client_state.init_kwargs["verify"] == certifi.where() + + +def test_create_httpx_client_uses_custom_verify_setting( + monkeypatch: pytest.MonkeyPatch, + fake_async_client_state: _FakeAsyncClientState, +) -> None: + import astrbot.core.zip_updator as zip_updator_module + + custom_verify = "/tmp/custom-ca.pem" + + monkeypatch.setattr( + zip_updator_module, + "httpx", + _build_fake_httpx_module(fake_async_client_state), + raising=False, + ) + + RepoZipUpdator(verify=custom_verify)._create_httpx_client(timeout=45.0) + + assert fake_async_client_state.init_kwargs is not None + assert fake_async_client_state.init_kwargs["follow_redirects"] is True + assert fake_async_client_state.init_kwargs["timeout"] == 45.0 + assert fake_async_client_state.init_kwargs["trust_env"] is True + assert fake_async_client_state.init_kwargs["verify"] == custom_verify + + +@pytest.mark.asyncio +async def test_fetch_release_info_logs_status_code_and_truncated_body_on_http_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import astrbot.core.zip_updator as zip_updator_module + + url = "https://api.soulter.top/releases" + body = "x" * 1005 + log_messages: list[str] = [] + + monkeypatch.setattr( + RepoZipUpdator, + "_create_httpx_client", + staticmethod( + lambda timeout=30.0: _FakeStatusErrorAsyncClient( # noqa: ARG005 + _FakeStatusErrorResponse(502, body, url) + ) + ), + ) + monkeypatch.setattr( + zip_updator_module.logger, + "error", + lambda message: log_messages.append(message), + ) + + with pytest.raises(Exception, match="解析版本信息失败"): + await RepoZipUpdator().fetch_release_info(url) + + assert any("状态码: 502" in message for message in log_messages) + assert any("内容: " in message for message in log_messages) + assert any("...[truncated]" in message for message in log_messages) + + +@pytest.mark.asyncio +async def test_download_file_removes_partial_file_when_stream_fails( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + monkeypatch.setattr( + RepoZipUpdator, + "_create_httpx_client", + staticmethod( + lambda timeout=30.0: _FakeFailingStreamAsyncClient() # noqa: ARG005 + ), + ) + + target_path = tmp_path / "partial.zip" + + with pytest.raises(RuntimeError, match="stream interrupted"): + await RepoZipUpdator()._download_file( + "https://example.com/archive.zip", + str(target_path), + ) + + assert not target_path.exists() + + +@pytest.mark.asyncio +async def test_download_file_logs_url_and_target_path_on_failure( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + import astrbot.core.zip_updator as zip_updator_module + + url = "https://example.com/archive.zip" + target_path = tmp_path / "logged-partial.zip" + log_messages: list[str] = [] + + monkeypatch.setattr( + RepoZipUpdator, + "_create_httpx_client", + staticmethod( + lambda timeout=30.0: _FakeFailingStreamAsyncClient() # noqa: ARG005 + ), + ) + monkeypatch.setattr( + zip_updator_module.logger, + "error", + lambda message: log_messages.append(message), + ) + + with pytest.raises(RuntimeError, match="stream interrupted"): + await RepoZipUpdator()._download_file(url, str(target_path)) + + assert any(url in message for message in log_messages) + assert any(str(target_path) in message for message in log_messages) diff --git a/tests/unit/test_cron_manager.py b/tests/unit/test_cron_manager.py index b111384ac..9dd3fc34d 100644 --- a/tests/unit/test_cron_manager.py +++ b/tests/unit/test_cron_manager.py @@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from astrbot.core.cron.manager import CronJobManager +from astrbot.core.cron.manager import CronJobManager, CronJobSchedulingError from astrbot.core.db.po import CronJob @@ -190,24 +190,25 @@ class TestAddActiveJob: @pytest.mark.asyncio async def test_add_active_job_run_once(self, cron_manager, mock_db, sample_cron_job): - """Test adding a run-once active job.""" + """Test adding a run-once active job with an invalid returned job.""" sample_cron_job.job_type = "active_agent" sample_cron_job.run_once = True mock_db.create_cron_job.return_value = sample_cron_job run_at = datetime.now(timezone.utc) + timedelta(days=30) - result = await cron_manager.add_active_job( - name="Test Run Once Job", - cron_expression=None, - payload={"session": "test:group:123"}, - run_once=True, - run_at=run_at, - ) + with pytest.raises(CronJobSchedulingError, match="Invalid isoformat string"): + await cron_manager.add_active_job( + name="Test Run Once Job", + cron_expression=None, + payload={"session": "test:group:123"}, + run_once=True, + run_at=run_at, + ) - assert result == sample_cron_job call_kwargs = mock_db.create_cron_job.call_args.kwargs assert call_kwargs["run_once"] is True + assert call_kwargs["payload"]["run_at"] == run_at.isoformat() class TestUpdateJob: diff --git a/tests/unit/test_document_storage_fts.py b/tests/unit/test_document_storage_fts.py new file mode 100644 index 000000000..753c371ef --- /dev/null +++ b/tests/unit/test_document_storage_fts.py @@ -0,0 +1,75 @@ +import pytest + +from astrbot.core.db.vec_db.faiss_impl.document_storage import DocumentStorage + + +@pytest.mark.asyncio +async def test_document_storage_fts_insert_search_and_delete(tmp_path): + storage = DocumentStorage(str(tmp_path / "doc.db")) + await storage.initialize() + + assert storage.fts5_available is True + + await storage.insert_documents_batch( + doc_ids=["chunk-1", "chunk-2"], + texts=["AstrBot 知识库召回性能优化", "FAISS 向量检索"], + metadatas=[ + {"kb_doc_id": "doc-1", "kb_id": "kb-1", "chunk_index": 0}, + {"kb_doc_id": "doc-1", "kb_id": "kb-1", "chunk_index": 1}, + ], + ) + + results = await storage.search_sparse(["知识库"], limit=10) + + assert results is not None + assert [result["doc_id"] for result in results] == ["chunk-1"] + + await storage.delete_document_by_doc_id("chunk-1") + results = await storage.search_sparse(["知识库"], limit=10) + + assert results == [] + + await storage.close() + + +@pytest.mark.asyncio +async def test_document_storage_fts_rebuilds_existing_documents(tmp_path): + storage = DocumentStorage(str(tmp_path / "doc.db")) + await storage.initialize() + + storage.fts5_available = False + await storage.insert_document( + doc_id="legacy-chunk", + text="legacy 知识库 文本", + metadata={"kb_doc_id": "doc-1", "kb_id": "kb-1", "chunk_index": 0}, + ) + + storage.fts5_available = True + storage._fts_index_ready = False + + results = await storage.search_sparse(["知识库"], limit=10) + + assert results is not None + assert [result["doc_id"] for result in results] == ["legacy-chunk"] + + await storage.close() + + +@pytest.mark.asyncio +async def test_document_storage_fts_delete_skips_missing_fts_row(tmp_path): + storage = DocumentStorage(str(tmp_path / "doc.db")) + await storage.initialize() + + storage.fts5_available = False + await storage.insert_document( + doc_id="legacy-chunk", + text="legacy 知识库 文本", + metadata={"kb_doc_id": "doc-1", "kb_id": "kb-1", "chunk_index": 0}, + ) + + storage.fts5_available = True + await storage.delete_document_by_doc_id("legacy-chunk") + + assert await storage.get_document_by_doc_id("legacy-chunk") is None + + await storage.close() diff --git a/tests/unit/test_sparse_retriever.py b/tests/unit/test_sparse_retriever.py new file mode 100644 index 000000000..11c491b4d --- /dev/null +++ b/tests/unit/test_sparse_retriever.py @@ -0,0 +1,93 @@ +import json +from types import SimpleNamespace + +import pytest + +from astrbot.core.knowledge_base.retrieval.sparse_retriever import SparseRetriever + + +def make_doc(chunk_id: str, text: str, chunk_index: int = 0) -> dict: + return { + "doc_id": chunk_id, + "text": text, + "metadata": json.dumps( + { + "chunk_index": chunk_index, + "kb_doc_id": f"doc-{chunk_index}", + "kb_id": "kb-1", + }, + ), + } + + +class FTSStorage: + def __init__(self): + self.search_sparse_calls = 0 + self.get_documents_calls = 0 + + async def search_sparse(self, query_tokens: list[str], limit: int): + self.search_sparse_calls += 1 + assert query_tokens == ["apple"] + assert limit == 1 + return [ + { + **make_doc("chunk-1", "apple banana", 0), + "score": -1.0, + }, + ] + + async def get_documents(self, *args, **kwargs): + self.get_documents_calls += 1 + return [] + + +class FallbackStorage: + def __init__(self): + self.search_sparse_calls = 0 + self.get_documents_calls = 0 + + async def search_sparse(self, query_tokens: list[str], limit: int): + self.search_sparse_calls += 1 + return None + + async def get_documents(self, metadata_filters: dict, limit: int | None, offset): + self.get_documents_calls += 1 + return [ + make_doc("chunk-1", "apple banana", 0), + make_doc("chunk-2", "orange pear", 1), + make_doc("chunk-3", "grape melon", 2), + ] + + +@pytest.mark.asyncio +async def test_sparse_retriever_uses_fts5_when_available(): + storage = FTSStorage() + vec_db = SimpleNamespace(document_storage=storage) + retriever = SparseRetriever(kb_db=None) + + results = await retriever.retrieve( + query="apple", + kb_ids=["kb-1"], + kb_options={"kb-1": {"vec_db": vec_db, "top_k_sparse": 1}}, + ) + + assert [result.chunk_id for result in results] == ["chunk-1"] + assert storage.search_sparse_calls == 1 + assert storage.get_documents_calls == 0 + + +@pytest.mark.asyncio +async def test_sparse_retriever_falls_back_to_bm25_when_fts5_is_unavailable(): + storage = FallbackStorage() + vec_db = SimpleNamespace(document_storage=storage) + retriever = SparseRetriever(kb_db=None) + + results = await retriever.retrieve( + query="apple", + kb_ids=["kb-1"], + kb_options={"kb-1": {"vec_db": vec_db, "top_k_sparse": 1}}, + ) + + assert [result.chunk_id for result in results] == ["chunk-1"] + assert storage.search_sparse_calls == 1 + assert storage.get_documents_calls == 1