diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 0000000..e8c87a1 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,39 @@ +# Repository Guidelines + +## Project Structure & Module Organization + +This is an AstrBot plugin. `main.py` centralizes AstrBot command registration. `src/` has four top-level packages only: `application/` for use cases and ports; `domain/` for Setu, fortune, and access-control rules; `infrastructure/` for AstrBot adapters, providers, persistence, permissions, and sending; and `shared/` for helpers and config models. Tests live in `tests/`, WebUI pages in `pages/`, and plugin skills in `skills/`. + +## Agent Skills & Tooling + +Use `$skill-astrbot-dev` for AstrBot structure, decorators/hooks, lifecycle, config schema, message flow, platform adapters, and LLM tools. If docs and source disagree, trust source. Use `$github` for issues, PRs, CI, and advanced repository queries through `gh`. + +## Build, Test, and Development Commands + +- `python -m pip install -e ".[dev]"`: install the plugin with dev tools. +- `python -m pip install -U astrbot`: refresh the AstrBot SDK for local signatures. +- `PYTHONPATH=/path/to/data/plugins python -m pytest`: run the full test suite. +- `python -m pytest tests/domain`: run focused domain tests. +- `python -m ruff check .`: lint Python code. +- `python -m ruff format .`: format Python files. +- `python -m py_compile main.py src/**/*.py tests/**/*.py`: syntax check. + +## Ruff Tooling + +Ruff settings come from repository config. The project targets Python 3.10, uses 100-character lines, and sorts `astrbot` and `astrbot_plugin_setu` as first-party imports. If the parent cache is not writable, use `RUFF_CACHE_DIR=.ruff_cache`. + +## Coding Style & Naming Conventions + +Use 4-space indentation, type hints, and `from __future__ import annotations` in new Python modules. Use `snake_case` for modules/functions, `PascalCase` for classes, and `UPPER_SNAKE_CASE` for constants. Handlers, hooks, and tool functions should be `async def`. Keep `main.py` focused on registration and forwarding; keep reusable orchestration in `application/`. + +## Testing Guidelines + +Use pytest and pytest-asyncio. Name files `test_*.py`, classes `TestFeatureName`, and methods `test_behavior`. Reuse fixtures from `tests/conftest.py` for AstrBot config, events, providers, and temp data directories. Add focused unit tests for domain and infrastructure changes. + +## Commit & Pull Request Guidelines + +Recent history mostly follows Conventional Commits, for example `feat(safety): ...`, `fix: ...`, `refactor: ...`, `docs: ...`, and `chore: ...`. Keep subjects imperative and scoped when useful. PRs should explain motivation, summarize core file changes, identify breaking changes, and include verification output or screenshots. + +## Security & Configuration Tips + +Do not commit secrets, tokens, local AstrBot runtime data, or downloaded image caches. Session overrides are runtime data; keep fixtures under `tests/`. Keep `_conf_schema.json`, `src/shared/config/models.py`, `metadata.yaml`, `requirements.txt`, and README examples in sync when adding settings or dependencies. diff --git a/CHANGELOG.md b/CHANGELOG.md index 2830628..f0a24eb 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,22 @@ # Changelog +## [2.0.0] - 2026-05-17 + +### Changed +- **统一命令入口**:合并运势相关命令,减少主入口分发重复 +- **发送链路日志增强**:补全 sender / provider / fallback / NapCat stream 的关键日志 +- **配置加载路径修正**:插件初始化改为读取插件自身配置,不再误用 AstrBot 全局配置 + +### Fixed +- **修复图片反代未生效**:Lolicon / Atri 现在都会使用插件配置中的 proxy,并对 Pixiv 图片 URL 做兜底改写 +- **修复空 aspect_ratio 导致启动失败**:兼容配置文件中的空字符串,避免 Pydantic 校验报错 +- **修复旧临时文件路径直发失败**:普通发送前先将本地图片物化为稳定载荷,降低 OneBot 发送阶段失败率 + +### Added +- **统一 provider 诊断日志**:补充抓取、下载、缓存命中和失败汇总日志,便于定位“拿不到图”的根因 + +## [1.3.1] - 2026-04-19 + ## [1.3.0] - 2026-04-10 ### Added @@ -233,4 +250,3 @@ ### Added - 瑟瑟诞生!初始版本发布 - diff --git a/README.md b/README.md index b2983d0..5654f4c 100644 --- a/README.md +++ b/README.md @@ -8,11 +8,11 @@ Moe Counter -**一个支持多平台、可自定义、带防审核机制的随机色图插件,支持多 API、HTML 卡片包装、LLM 工具调用。** +**一个支持多平台、可自定义、带防审核机制的随机色图插件,支持多 API、会话级配置、LLM 工具调用。** [![License: APGL](https://img.shields.io/badge/License-APGL-blue.svg)](https://opensource.org/licenses/agpl-3.0) ![Python Version](https://img.shields.io/badge/Python-3.10%2B-blue) -![AstrBot](https://img.shields.io/badge/AstrBot-%E2%89%A54.10.4-green) +![AstrBot](https://img.shields.io/badge/AstrBot-%E2%89%A54.24.0-green) ![Platform](https://img.shields.io/badge/Platform-Windows%20%7C%20Linux-lightgrey) @@ -55,11 +55,12 @@ ## ✨ 功能特性 - 🎨 **多 API 支持** - Lolicon、SexNyan、自定义 API 等 +- 🧩 **分层架构** - `application` / `domain` / `infrastructure` / `shared` 清晰分离 - 🖼️ **HTML 卡片包装** - 防止平台审核,支持自定义样式 - 🤖 **LLM 工具调用** - 可通过大模型自动获取色图 - 🏷️ **标签搜索** - 支持多标签、中文标签、模糊匹配 - 🔄 **多种发送模式** - 直接发送、合并转发、文件封装 -- 🛡️ **防审核机制** - 图片混淆、延迟撤回、Docx 封装 +- 🛡️ **防审核机制** - HTML 卡片 fallback、NapCat 流式上传、延迟撤回、Docx 封装 - ⚡ **性能优化** - 并发下载、磁盘缓存、自动补图、httpx、分段下载 - 🌐 **多平台适配** - 兼容 AstrBot 支持的所有平台 @@ -90,11 +91,11 @@ | 配置项 | 类型 | 说明 | 可选值 | 默认值 | |--------|------|------|--------|--------| -| `api_type` | 字符串 | API 类型 | `lolicon` / `sexnyan` / `custom` / `all` | `lolicon` | +| `api_type` | 字符串 | API 类型 | `lolicon` / `atri` / `sexnyan` / `custom` / `all` | `lolicon` | | `send_mode` | 字符串 | 发送模式 | `auto` / `image` / `forward` | `auto` | | `content_mode` | 字符串 | 内容模式 | `sfw` / `r18` / `mix` | `sfw` | | `max_count` | 整数 | 单次最大图片数 | 1-20 | `10` | -| `cache_enabled` | 布尔值 | 是否启用图片磁盘缓存 | `true` / `false` | `true` | +| `cache_enabled` | 布尔值 | 是否复用本地发送缓存 | `true` / `false` | `true` | | `exclude_ai` | 布尔值 | 是否排除 AI 生成图片 | `true` / `false` | `false` | ### HTML 卡片与防审核配置 @@ -102,6 +103,7 @@ | 配置项 | 类型 | 说明 | 可选值 | 默认值 | |--------|------|------|--------|--------| | `html_card_strategy` | 字符串 | HTML 卡片策略 | `never` / `fallback` / `always` | `fallback` | +| `napcat_stream_mode` | 字符串 | NapCat 流式上传策略 | `disabled` / `fallback` / `always` | `fallback` | | `auto_revoke_r18` | 布尔值 | R18 图片是否自动撤回 | `true` / `false` | `false` | | `r18_docx_mode` | 布尔值 | R18 是否使用 Docx 封装 | `true` / `false` | `true` | @@ -115,7 +117,7 @@ | `download_concurrent_limit` | 整数 | 并发下载限制 | 1-50 | `10` | | `download_timeout_seconds` | 整数 | 下载超时时间(秒) | 5-300 | `30` | -### 访问控制配置(1.3.0+) +### 访问控制配置(2.0.0) > 安全配置统一位于 `safety` 下,且 `setu` / `fortune` 完全独立。 @@ -154,7 +156,7 @@ - `safety.setu_access_control_mode` -> `safety.setu_user_access_control_mode` + `safety.setu_group_access_control_mode` - `safety.fortune_access_control_mode` -> `safety.fortune_user_access_control_mode` + `safety.fortune_group_access_control_mode` -> 旧键当前仍兼容读取,建议在 WebUI 中显式配置新键,后续版本更稳。 +> 旧键当前仍兼容读取,但 `v2.0.0` 起主实现已完全转到新键。 ### 配置示例 @@ -211,12 +213,38 @@ /setu 白丝 萝莉 /setu 3 白丝 /setu 4 白丝 萝莉 -/setu_mode r18 +/session_config set setu.content_mode r18 ``` - 数量范围支持中文数字 - 标签支持空格、逗号、顿号分隔 -- `/setu_mode` 可切换内容模式 +- `/session_config` 统一管理当前会话的覆盖配置 + +### v2.0.0 更新摘要 + +- 运势管理命令收敛为更少的统一入口,旧命令保留兼容别名 +- provider / sender / fallback 日志补全,方便定位“有 URL 但发不出去”或“provider 无结果”的问题 +- 反代配置读取、provider 重建、插件配置来源修正,减少 WebUI 改配置后不生效的问题 + +### 会话配置命令(管理员设置) + +会话覆盖配置会写入插件数据目录下的 `session_overrides.json`,不会修改全局 WebUI 配置。 +也可以在插件 WebUI 的 `sessionConfig` 页面集中管理所有群聊/私聊会话覆盖。 + +``` +/session_config get +/session_config get setu.content_mode +/session_config get json +/session_config set setu.content_mode r18 +/session_config set setu.r18_docx true +/session_config set setu.auto_revoke false +/session_config set setu.send_mode forward +/session_config set fortune.tags 白丝,猫耳 +/session_config clear setu.send_mode +/session_config clear +``` + +可用配置项:`setu.content_mode`、`setu.r18_docx`、`setu.auto_revoke`、`setu.send_mode`、`fortune.tags`、`fortune.content_mode`。 ### 黑白名单管理命令(管理员) @@ -235,10 +263,10 @@ | 命令 | 说明 | 示例 | |------|------|------| -| `/拉黑运势用户 @用户` | 将用户加入运势黑名单(必须AT) | `/拉黑运势用户 @小明` | -| `/解除运势拉黑 @用户` | 将用户从运势黑名单移除(必须AT) | `/解除运势拉黑 @小明` | -| `/信任运势用户 @用户` | 将用户加入运势白名单(必须AT) | `/信任运势用户 @小明` | -| `/取消运势信任 @用户` | 将用户从运势白名单移除(必须AT) | `/取消运势信任 @小明` | +| `/运势用户 拉黑 @用户` | 将用户加入运势黑名单(必须AT) | `/运势用户 拉黑 @小明` | +| `/运势用户 解黑 @用户` | 将用户从运势黑名单移除(必须AT) | `/运势用户 解黑 @小明` | +| `/运势用户 信任 @用户` | 将用户加入运势白名单(必须AT) | `/运势用户 信任 @小明` | +| `/运势用户 取消信任 @用户` | 将用户从运势白名单移除(必须AT) | `/运势用户 取消信任 @小明` | #### 群组功能开关 @@ -246,14 +274,23 @@ |------|------|------| | `/开启色图` | 在本群开启色图功能(移出色图群组黑名单) | `/开启色图` | | `/关闭色图` | 在本群关闭色图功能(加入色图群组黑名单) | `/关闭色图` | -| `/开启运势` | 在本群开启运势功能(移出运势群组黑名单) | `/开启运势` | -| `/关闭运势` | 在本群关闭运势功能(加入运势群组黑名单) | `/关闭运势` | +| `/运势开关 开` | 在本群开启运势功能(移出运势群组黑名单) | `/运势开关 开` | +| `/运势开关 关` | 在本群关闭运势功能(加入运势群组黑名单) | `/运势开关 关` | + +#### 运势刷新 + +| 命令 | 说明 | 示例 | +|------|------|------| +| `/运势刷新` | 刷新自己的今日运势 | `/运势刷新` | +| `/运势刷新 本群` | 刷新当前群今日运势 | `/运势刷新 本群` | +| `/运势刷新 全局` | 刷新全局今日运势 | `/运势刷新 全局` | **注意:** - 以上命令仅限管理员或超级管理员使用。 - 用户类命令必须通过 `@` 指定目标用户,不支持直接输入用户 ID。 - 群组命令仅作用于当前群。 - 白名单/黑名单会自动保持互斥(同一功能下不会同时存在)。 +- 旧命令 `/开启运势`、`/关闭运势`、`/拉黑运势用户`、`/解除运势拉黑`、`/信任运势用户`、`/取消运势信任`、`/刷新今日运势`、`/刷新本群今日运势`、`/刷新全局今日运势` 仍可继续使用。 #### 安全配置示例(`safety`) @@ -289,11 +326,14 @@ | 工具名 | 作用 | 参数 | 权限 | |---|---|---|---| | `get_setu_image` | 获取并发送随机图片 | `count: integer`(数量), `tags: string[]`(标签) | 普通用户可用 | -| `get_setu_content_mode` | 查看当前会话生效的内容模式 | 无 | 普通用户可用 | -| `set_setu_content_mode` | 设置当前会话内容模式 | `mode: string`,可选 `sfw/r18/mix/clear` | 管理员/超级管理员 | -| `set_setu_r18_docx_mode` | 设置当前会话 R18 Docx 封装开关 | `enabled: boolean`(部分场景支持 clear 语义) | 管理员/超级管理员 | -| `set_setu_auto_revoke` | 设置当前会话 R18 自动撤回开关 | `enabled: boolean`(部分场景支持 clear 语义) | 管理员/超级管理员 | -| `set_setu_send_mode` | 设置当前会话发送模式 | `mode: string`,可选 `image/forward/auto/clear` | 管理员/超级管理员 | + +##### 会话配置工具 + +| 工具名 | 作用 | 参数 | 权限 | +|---|---|---|---| +| `get_session_config` | 查看当前会话全部配置或单个 key 的生效值 | `key?: string` | 普通用户可用 | +| `set_session_config` | 设置当前会话一个覆盖配置 | `key: string`, `value: string` | 管理员/超级管理员 | +| `clear_session_config` | 清除当前会话一个覆盖配置,或清空全部覆盖 | `key?: string` | 管理员/超级管理员 | ##### 今日运势工具 @@ -303,13 +343,11 @@ | `refresh_my_fortune` | 刷新"我的"今日运势 | 无 | 管理员 | | `refresh_group_fortune` | 刷新当前群今日运势 | 无 | 管理员 | | `refresh_all_fortune` | 刷新全局今日运势 | 无 | 超级管理员 | -| `get_fortune_config` | 查看当前会话运势配置 | 无 | 普通用户可用 | -| `set_fortune_config` | 设置当前会话运势配置 | `tags: string`, `mode: string(sfw/r18/mix)` | 管理员 | ##### 调用建议 -- 需要"仅查看状态"时优先调用查询类工具(如 `get_setu_content_mode`、`get_fortune_config`)。 -- 需要会话级覆写时使用 `set_*` 工具;希望回到全局配置时可使用 `clear` 语义参数(支持的工具见上表)。 +- 需要"仅查看状态"时优先调用 `get_session_config`。 +- 需要会话级覆写时使用 `set_session_config`;希望回到全局配置时使用 `clear_session_config`。 - 对于发送类工具,插件会直接把结果发送到当前会话,工具返回文本用于说明执行结果。 ### 高级用法 @@ -317,8 +355,8 @@ | 功能 | 说明 | 配置方式 | |------|------|----------| | **自定义 API** | 设置 `api_type` 为 `custom`,填写自定义 API 地址和解析规则,实现对接任意第三方色图接口 | 配置面板 | -| **图片混淆** | 如遇平台审核拦截,插件会自动尝试对图片进行字节级混淆重发 | 自动触发 | -| **磁盘缓存** | 通过 `cache_enabled` 启用图片磁盘缓存,提升多次请求同一图片的响应速度 | 配置面板 | +| **NapCat 流式上传** | 通过 `napcat_stream_mode` 控制本地图片文件的流式上传,默认普通发送失败后自动重试 | 配置面板 | +| **发送缓存** | 图片先落盘再发送;`cache_enabled` 控制是否复用未过期的本地文件,降低 original 大图内存压力 | 配置面板 | | **多 API 策略** | 支持 `all` 模式自动切换多 API,提升获取成功率 | 设置 `api_type` 为 `all` | | **标签与过滤** | 支持多标签、中文标签、AI 过滤(`exclude_ai`),可灵活组合搜索条件 | 配置面板 | @@ -334,6 +372,8 @@ | 参数 | 说明 | 推荐值 | |------|------|--------| +| `napcat_stream_mode` | NapCat/OneBot 图片传输策略:`fallback` 先普通发送,失败后流式上传;`always` 发送前先流式上传;`disabled` 不使用流式上传 | `fallback` | +| `cache_enabled` | 是否复用发送缓存。即使关闭复用,运行时仍会优先落盘发送,避免大图长期停留在内存中 | `true` | | `enable_range_download` | 启用分段下载,将大图片分多段并行下载,适合高带宽服务器 | `false`(一般)/ `true`(高带宽) | | `range_segments` | 分段数 | 2-4 | | `range_download_threshold` | 分段下载阈值(KB),大于此值才启用分段 | 512 | @@ -343,9 +383,9 @@ --- ## 未来更新 -- [ ] 更好的自定义API -- [x] 新增一个作者自己的图库内置API,该图库的更新速度会比目前的图库API更快,并且跟随了当前版本潮流! -- [x] 一些基于色图的额外插件功能 (今日运势) +- [ ] 更好的自定义 API +- [ ] 更细粒度的 provider 健康检查与降级策略 +- [ ] 继续补全 WebUI 配置与诊断能力 --- ## 📄 开源协议 diff --git a/_conf_schema.json b/_conf_schema.json index 92c5e6b..e4f228c 100644 --- a/_conf_schema.json +++ b/_conf_schema.json @@ -84,24 +84,12 @@ "hint": "发送 R18 内容后多久自动撤回,默认 30 秒。", "default": 30 }, - "url_send_mode": { - "type": "bool", - "description": "URL 发送模式", - "hint": "开启后直接发送图片 URL 而不是下载后发送。可显著降低服务器带宽和内存占用,图片由客户端直接加载。", - "default": false - }, - "url_send_verify": { - "type": "bool", - "description": "验证 URL 有效性", - "hint": "发送前检查图片 URL 是否可访问(返回 200)。建议开启以避免发送失效链接。", - "default": true - }, - "url_send_timeout": { - "type": "int", - "description": "URL 验证超时(秒)", - "hint": "验证 URL 时的 HTTP 请求超时时间。默认 5 秒。", - "default": 5, - "slider": {"min": 2, "max": 10, "step": 1} + "napcat_stream_mode": { + "type": "string", + "description": "NapCat 流式上传", + "hint": "disabled=不用流式上传,fallback=普通发送失败后流式重试,always=发送前优先流式上传。", + "default": "fallback", + "options": ["disabled", "fallback", "always"] } } }, @@ -186,14 +174,14 @@ } }, "cache": { - "description": "图片缓存", + "description": "发送缓存", "type": "object", - "hint": "简单的基于 URL 键值的本地图片缓存,带索引文件。", + "hint": "图片会先流式下载到插件数据目录,再从本地文件发送,避免 original 大图进入内存。", "items": { "enabled": { "type": "bool", "description": "启用缓存", - "hint": "通过 URL 缓存图片以减少重复下载。", + "hint": "开启后复用未过期的本地发送缓存;关闭后仍会落盘发送,但不复用旧文件。", "default": true }, "ttl_hours": { @@ -546,102 +534,5 @@ "default": 30 } } - }, - "session_configs": { - "type": "template_list", - "description": "会话级配置覆盖", - "hint": "为特定群聊或私聊设置独立配置,优先于全局配置。可通过命令 /setu_config 在会话中设置。", - "default": [], - "obvious_hint": true, - "templates": { - "session_template": { - "name": "会话配置", - "hint": "一个群聊或私聊的独立配置", - "items": { - "session_id": { - "type": "string", - "description": "会话 ID", - "hint": "群号或用户 ID", - "default": "" - }, - "session_type": { - "type": "string", - "description": "会话类型", - "hint": "group=群聊,private=私聊", - "default": "group", - "options": ["group", "private"] - }, - "content_mode": { - "type": "string", - "description": "内容模式", - "hint": "覆盖全局内容模式设置", - "default": "", - "options": ["", "sfw", "r18", "mix"] - }, - "r18_docx_mode": { - "type": "string", - "description": "R18 Docx 打包模式", - "hint": "覆盖全局 R18 Docx 打包设置", - "default": "", - "options": ["", "enabled", "disabled"] - }, - "auto_revoke_r18": { - "type": "string", - "description": "R18 自动撤回", - "hint": "覆盖全局自动撤回设置", - "default": "", - "options": ["", "enabled", "disabled"] - }, - "send_mode": { - "type": "string", - "description": "发送模式", - "hint": "覆盖全局发送模式设置", - "default": "", - "options": ["", "image", "forward", "auto"] - } - } - } - } - }, - "fortune_session_configs": { - "type": "template_list", - "description": "今日运势会话配置覆盖", - "hint": "为特定群聊或私聊设置独立的今日运势配置。可通过命令 /jrys_config 在会话中设置。", - "default": [], - "obvious_hint": true, - "templates": { - "fortune_session_template": { - "name": "今日运势会话配置", - "hint": "一个群聊或私聊的今日运势独立配置", - "items": { - "session_id": { - "type": "string", - "description": "会话 ID", - "hint": "群号或用户 ID", - "default": "" - }, - "session_type": { - "type": "string", - "description": "会话类型", - "hint": "group=群聊,private=私聊", - "default": "group", - "options": ["group", "private"] - }, - "tags": { - "type": "string", - "description": "标签", - "hint": "获取今日运势图片时的标签,如 '少女,可爱'", - "default": "" - }, - "content_mode": { - "type": "string", - "description": "内容模式", - "hint": "覆盖全局内容模式设置", - "default": "", - "options": ["", "sfw", "r18", "mix"] - } - } - } - } } } diff --git a/config/__init__.py b/config/__init__.py deleted file mode 100644 index 5eafeae..0000000 --- a/config/__init__.py +++ /dev/null @@ -1,48 +0,0 @@ -"""Setu 插件配置解析和兼容性辅助模块。 - -提供配置读取、标签解析、以及新旧配置格式的兼容性支持。 -""" - -from __future__ import annotations - -from astrbot.core import AstrBotConfig - -from .api import ApiConfigMixin -from .base import ConfigBase -from .delivery import DeliveryConfigMixin -from .helpers import parse_count -from .html_card import HtmlCardConfigMixin -from .messages import MessagesConfigMixin -from .safety import SafetyConfigMixin - - -class SetuConfig( - ConfigBase, - ApiConfigMixin, - DeliveryConfigMixin, - HtmlCardConfigMixin, - SafetyConfigMixin, - MessagesConfigMixin, -): - """Setu 插件配置包装类。 - - 支持嵌套配置和旧版配置的兼容性处理,提供类型安全的配置访问方法。 - - 通过多重继承组合各个功能模块的配置: - - ApiConfigMixin: API 相关配置 - - DeliveryConfigMixin: 发送相关配置 - - HtmlCardConfigMixin: HTML 卡片配置 - - SafetyConfigMixin: 安全和缓存配置 - - MessagesConfigMixin: 消息配置 - """ - - def __init__(self, config: AstrBotConfig): - """初始化配置包装器。 - - 参数: - config: AstrBot 配置对象 - """ - super().__init__(config) - - -__all__ = ["SetuConfig", "parse_count"] diff --git a/config/api.py b/config/api.py deleted file mode 100644 index 3f97159..0000000 --- a/config/api.py +++ /dev/null @@ -1,332 +0,0 @@ -"""API 相关配置。""" - -from __future__ import annotations - -from typing import TYPE_CHECKING, Any - -from .helpers import safe_bool, safe_int - -if TYPE_CHECKING: - from .base import ConfigBase - - -class ApiConfigMixin: - """API 配置混入类。""" - - _read: Any - - @property - def api_type(self: ConfigBase) -> str: - """API 提供商类型。 - - 返回: - 返回 API 类型:lolicon、atri、sexnyan、custom 或 all - """ - api_type = self._read( - ("setu_general", "api_type"), - ("general", "api_type"), - "api_type", - "apiType", - default="lolicon", - ) - return ( - api_type - if api_type in ("lolicon", "atri", "sexnyan", "custom", "all") - else "lolicon" - ) - - @property - def multi_api_strategy(self: ConfigBase) -> str: - """多 API 策略。 - - 返回: - 返回策略类型:round_robin(轮询)、random(随机)、failover(故障转移) - """ - strategy = self._read( - ("setu_general", "multi_api_strategy"), - ("general", "multi_api_strategy"), - "multi_api_strategy", - default="round_robin", - ) - return ( - strategy - if strategy in ("round_robin", "random", "failover") - else "round_robin" - ) - - @property - def content_mode(self: ConfigBase) -> str: - """内容模式(分级)。 - - 返回: - 返回内容模式:sfw(全年龄)、r18(成人)、mix(混合) - """ - mode = self._read( - ("setu_general", "content_mode"), - ("general", "content_mode"), - "content_mode", - "contentMode", - default="sfw", - ) - return mode if mode in ("sfw", "r18", "mix") else "sfw" - - @property - def exclude_ai(self: ConfigBase) -> bool: - """排除 AI 生成作品(仅 lolicon 生效)。 - - 返回: - 是否排除 AI 生成的图片 - """ - return safe_bool( - self._read( - ("api", "lolicon", "exclude_ai"), - "exclude_ai", - "excludeAi", - default=True, - ), - True, - ) - - @property - def image_size(self: ConfigBase) -> str: - """图片尺寸(仅 lolicon 生效)。 - - 返回: - 返回图片尺寸:original(原图)、regular、small、thumb、mini - """ - size = self._read( - ("api", "lolicon", "image_size"), "image_size", default="original" - ) - return ( - size - if size in ("original", "regular", "small", "thumb", "mini") - else "original" - ) - - @property - def proxy(self: ConfigBase) -> str: - """图片代理主机(仅 lolicon 生效)。 - - 返回: - 代理服务器地址,默认 i.pixiv.re - """ - return str( - self._read(("api", "lolicon", "proxy"), "proxy", default="i.pixiv.re") - ) - - @property - def aspect_ratio(self: ConfigBase) -> str: - """宽高比过滤(仅 lolicon 生效)。 - - 返回: - 返回宽高比:horizontal(横向)、vertical(纵向)、square(方形) - """ - ratio = self._read( - ("api", "lolicon", "aspect_ratio"), "aspect_ratio", default="" - ) - return ratio if ratio in ("horizontal", "vertical", "square") else "" - - @property - def uid(self: ConfigBase) -> list[int]: - """作者 UID 列表(仅 lolicon 生效)。 - - 返回: - 指定作者的 UID 列表,用于筛选特定作者的作品 - """ - uids = self._read(("api", "lolicon", "uid"), "uid", default=[]) - if isinstance(uids, list): - return [safe_int(uid, 0) for uid in uids if safe_int(uid, 0) > 0] - return [] - - @property - def keyword(self: ConfigBase) -> str: - """关键词过滤(仅 lolicon 生效)。 - - 返回: - 用于过滤图片的关键词 - """ - return str(self._read(("api", "lolicon", "keyword"), "keyword", default="")) - - @property - def max_replenish_rounds(self: ConfigBase) -> int: - """最大补充轮次。 - - 返回: - 部分图片下载失败时的重试轮次 - """ - return safe_int( - self._read( - ("setu_general", "max_replenish_rounds"), - ("general", "max_replenish_rounds"), - "max_replenish_rounds", - "maxReplenishRounds", - default=3, - ), - 3, - ) - - @property - def custom_api_configs(self: ConfigBase) -> list[dict[str, Any]]: - """自定义 API 配置列表。 - - 返回: - 用户自定义的 API 配置列表 - """ - configs = self._read( - ("api", "custom_api_configs"), "custom_api_configs", default=[] - ) - if isinstance(configs, list): - return configs - return [] - - def get_custom_api_config(self, name: str | None = None) -> dict[str, Any] | None: - """获取自定义 API 配置。 - - 参数: - name: 配置名称,如果不指定则返回第一个配置 - - 返回: - 配置字典或 None - """ - configs = self.custom_api_configs - if not configs: - return None - - if name: - for cfg in configs: - if cfg.get("name") == name: - return cfg - return None - return configs[0] - - @property - def custom_api(self) -> dict[str, Any]: - """自定义 API 基本信息。 - - 返回: - 包含 url、method、timeout 的字典 - """ - cfg = self.get_custom_api_config() - if cfg: - return { - "url": cfg.get("url", ""), - "method": cfg.get("method", "GET"), - "timeout": safe_int(cfg.get("timeout"), 30), - } - return {"url": "", "method": "GET", "timeout": 30} - - @property - def api_response_parser(self) -> dict[str, Any]: - """API 响应解析器配置。 - - 返回: - 包含解析器类型和 JSON 路径的字典 - """ - cfg = self.get_custom_api_config() - if cfg: - return { - "type": cfg.get("parser_type", "auto"), - "json_path": cfg.get("json_path", ""), - } - return {"type": "auto", "json_path": ""} - - # Atri 配置属性(和 Lolicon 相同结构) - @property - def atri_image_size(self: ConfigBase) -> str: - """图片尺寸(仅 atri 生效)。 - - 返回: - 返回图片尺寸:original(原图)、regular、small、thumb、mini - """ - size = self._read( - ("api", "atri", "image_size"), "image_size", default="original" - ) - return ( - size - if size in ("original", "regular", "small", "thumb", "mini") - else "original" - ) - - @property - def atri_proxy(self: ConfigBase) -> str: - """图片代理主机(仅 atri 生效)。 - - 返回: - 代理服务器地址,默认 i.pixiv.re - """ - return str(self._read(("api", "atri", "proxy"), "proxy", default="i.pixiv.re")) - - @property - def atri_aspect_ratio(self: ConfigBase) -> str: - """宽高比过滤(仅 atri 生效)。 - - 返回: - 返回宽高比:horizontal(横向)、vertical(纵向)、square(方形) - """ - ratio = self._read(("api", "atri", "aspect_ratio"), "aspect_ratio", default="") - return ratio if ratio in ("horizontal", "vertical", "square") else "" - - @property - def atri_uid(self: ConfigBase) -> list[int]: - """作者 UID 列表(仅 atri 生效)。 - - 返回: - 指定作者的 UID 列表,用于筛选特定作者的作品 - """ - uids = self._read(("api", "atri", "uid"), "uid", default=[]) - if isinstance(uids, list): - return [safe_int(uid, 0) for uid in uids if safe_int(uid, 0) > 0] - return [] - - @property - def atri_keyword(self: ConfigBase) -> str: - """关键词过滤(仅 atri 生效)。 - - 返回: - 用于过滤图片的关键词 - """ - return str(self._read(("api", "atri", "keyword"), "keyword", default="")) - - @property - def atri_exclude_ai(self: ConfigBase) -> bool: - """排除 AI 生成作品(仅 atri 生效)。 - - 返回: - 是否排除 AI 生成的图片 - """ - return safe_bool( - self._read( - ("api", "atri", "exclude_ai"), - "exclude_ai", - "excludeAi", - default=True, - ), - True, - ) - - @property - def fortune_api_type(self: ConfigBase) -> str: - """今日运势的 API 提供商类型。 - - 返回: - 返回 API 类型:inherit(继承色图配置)、lolicon、atri、sexnyan、custom - """ - api_type = self._read( - ("fortune", "api_type"), "fortune_api_type", default="inherit" - ) - if api_type in ("inherit", "lolicon", "atri", "sexnyan", "custom"): - return api_type - return "inherit" - - def get_effective_fortune_api_type(self: ConfigBase) -> str: - """获取生效的今日运势 API 类型。 - - 如果设置为 inherit,则使用色图的 api_type。 - - 返回: - 生效的 API 类型:lolicon、atri、sexnyan、custom、all - """ - fortune_api = self.fortune_api_type - if fortune_api == "inherit": - return self.api_type - return fortune_api diff --git a/config/base.py b/config/base.py deleted file mode 100644 index 206f106..0000000 --- a/config/base.py +++ /dev/null @@ -1,61 +0,0 @@ -"""配置基础类。""" - -from __future__ import annotations - -from typing import Any - -from astrbot.core import AstrBotConfig - - -class ConfigBase: - """配置基础类,提供嵌套配置读取功能。""" - - def __init__(self, config: AstrBotConfig): - """初始化配置包装器。 - - 参数: - config: AstrBot 配置对象 - """ - self._cfg = config - - def _read( - self, - nested_path: tuple[str, ...], - *legacy_keys: str, - default: Any = None, - ) -> Any: - """读取配置值,支持嵌套路径和旧版键名回退。 - - 优先尝试嵌套路径读取,如果失败则尝试旧版键名,最后返回默认值。 - - 参数: - nested_path: 嵌套配置路径,如 ("general", "api_type") - *legacy_keys: 旧版配置键名,用于兼容性回退 - default: 默认值,读取失败时返回 - - 返回: - 配置值或默认值 - """ - current: Any = self._cfg - for key in nested_path: - if isinstance(current, dict): - current = current.get(key) - elif hasattr(current, "get"): - current = current.get(key) - else: - current = None - if current is None: - break - - if current is not None: - return current - - # 尝试旧版键名 - for key in legacy_keys: - try: - value = self._cfg.get(key) - except (KeyError, AttributeError): - value = None - if value is not None: - return value - return default diff --git a/config/delivery.py b/config/delivery.py deleted file mode 100644 index 1018b56..0000000 --- a/config/delivery.py +++ /dev/null @@ -1,151 +0,0 @@ -"""发送和交付相关配置。""" - -from __future__ import annotations - -from typing import TYPE_CHECKING, Any - -from .helpers import safe_bool, safe_int - -if TYPE_CHECKING: - from .base import ConfigBase - - -class DeliveryConfigMixin: - """发送配置混入类。""" - - _read: Any - - @property - def r18_docx_mode(self: ConfigBase) -> bool: - """R18 Docx 打包模式。 - - 返回: - 是否将 R18 图片打包到 docx 文件中 - """ - return safe_bool( - self._read(("delivery", "r18_docx_mode"), "r18_docx_mode", default=True), - True, - ) - - @property - def send_mode(self: ConfigBase) -> str: - """图片发送模式。 - - 返回: - 返回发送模式:image(直接发送)、forward(合并转发)、auto(自动) - """ - mode = self._read( - ("delivery", "send_mode"), "send_mode", "sendMode", default="image" - ) - return mode if mode in ("image", "forward", "auto") else "image" - - @property - def auto_handle_send_failure(self: ConfigBase) -> bool: - """自动处理发送失败。 - - 返回: - 发送失败时是否自动尝试 HTML 卡片降级发送 - """ - return safe_bool( - self._read( - ("delivery", "auto_handle_send_failure"), - "auto_handle_send_failure", - default=True, - ), - True, - ) - - @property - def auto_revoke_r18(self: ConfigBase) -> bool: - """自动撤回 R18 内容。 - - 返回: - 是否自动撤回 R18 图片/文件 - """ - return safe_bool( - self._read( - ("delivery", "auto_revoke_r18"), - "auto_revoke_r18", - default=False, - ), - False, - ) - - @property - def auto_revoke_delay(self: ConfigBase) -> int: - """自动撤回延迟时间(秒)。 - - 返回: - R18 内容发送后多久自动撤回 - """ - return safe_int( - self._read( - ("delivery", "auto_revoke_delay"), - "auto_revoke_delay", - default=30, - ), - 30, - ) - - @property - def max_count(self: ConfigBase) -> int: - """每次请求最大图片数。 - - 返回: - 单次命令的上限 - """ - return safe_int( - self._read(("general", "max_count"), "max_count", "maxCount", default=10), - 10, - ) - - @property - def url_send_mode(self: ConfigBase) -> bool: - """URL 发送模式。 - - 返回: - 是否直接发送图片 URL 而不是下载后发送 - 开启后插件不会下载图片,而是直接发送图片链接 - 可以降低服务器带宽和内存占用 - """ - return safe_bool( - self._read( - ("delivery", "url_send_mode"), - "url_send_mode", - default=False, - ), - False, - ) - - @property - def url_send_verify(self: ConfigBase) -> bool: - """URL 发送前验证链接有效性。 - - 返回: - 是否在发送前验证 URL 是否可访问(返回 200) - 仅当 url_send_mode 为 True 时生效 - """ - return safe_bool( - self._read( - ("delivery", "url_send_verify"), - "url_send_verify", - default=True, - ), - True, - ) - - @property - def url_send_timeout(self: ConfigBase) -> int: - """URL 验证超时时间(秒)。 - - 返回: - 验证 URL 时的 HTTP 请求超时时间 - """ - return safe_int( - self._read( - ("delivery", "url_send_timeout"), - "url_send_timeout", - default=5, - ), - 5, - ) diff --git a/config/helpers.py b/config/helpers.py deleted file mode 100644 index a21914c..0000000 --- a/config/helpers.py +++ /dev/null @@ -1,75 +0,0 @@ -"""配置辅助函数。""" - -from __future__ import annotations - -from typing import Any - -from astrbot.api import logger - - -def safe_int(value: Any, default: int) -> int: - """安全地解析整数。 - - 参数: - value: 待解析的值 - default: 解析失败时返回的默认值 - - 返回: - 解析后的正整数,失败返回默认值 - """ - try: - parsed = int(value) - return parsed if parsed > 0 else default - except Exception: - return default - - -def safe_bool(value: Any, default: bool) -> bool: - """安全地解析布尔值。 - - 支持布尔值和字符串形式的布尔值(如 "true", "yes", "1" 等)。 - - 参数: - value: 待解析的值 - default: 解析失败时返回的默认值 - - 返回: - 解析后的布尔值,失败返回默认值 - """ - if isinstance(value, bool): - return value - if isinstance(value, str): - lowered = value.strip().lower() - if lowered in {"1", "true", "yes", "on"}: - return True - if lowered in {"0", "false", "no", "off"}: - return False - return default - - -def parse_count(raw: str) -> int: - """解析阿拉伯数字或中文数字字符串。 - - 支持简单数字(如 "3")和中文数字(如 "三"、"十五"、"二十三")。 - 解析失败返回 -1。 - - 参数: - raw: 待解析的数字字符串 - - 返回: - 解析后的整数,失败返回 -1 - """ - from ..utils import cn_to_an - - s = (raw or "").strip() - if not s: - return 1 - if s.isdigit(): - return int(s) - # 尝试使用 utils 中的 cn_to_an 解析复杂中文数字 - try: - result = cn_to_an(s) - return result if result > 0 else -1 - except Exception as e: - logger.debug("error parsing count: %s", e) - return -1 diff --git a/config/html_card.py b/config/html_card.py deleted file mode 100644 index 4da6136..0000000 --- a/config/html_card.py +++ /dev/null @@ -1,62 +0,0 @@ -"""HTML 卡片相关配置。""" - -from __future__ import annotations - -from typing import TYPE_CHECKING, Any - -from .helpers import safe_int - -if TYPE_CHECKING: - from .base import ConfigBase - - -class HtmlCardConfigMixin: - """HTML 卡片配置混入类。""" - - _read: Any - - @property - def html_card_strategy(self: ConfigBase) -> str: - """HTML 卡片策略。 - - 返回: - 返回策略:never(从不)、fallback(失败时降级)、always(总是) - """ - strategy = self._read( - ("html_card", "strategy"), "html_card_strategy", default="fallback" - ) - if strategy in ("never", "fallback", "always"): - return strategy - # 兼容旧版 enabled 配置 - old_enabled = self._read(("html_card", "enabled"), default=None) - if old_enabled is True: - return "fallback" - return "never" - - @property - def html_card_mode(self: ConfigBase) -> str: - """HTML 卡片模式。 - - 返回: - 返回卡片模式:single(单张)、multiple(多张) - """ - mode = self._read(("html_card", "mode"), "html_card_mode", default="single") - return mode if mode in ("single", "multiple") else "single" - - @property - def html_card_padding(self: ConfigBase) -> int: - """HTML 卡片内边距。 - - 返回: - 卡片内部图片与边框的距离 - """ - return safe_int(self._read(("html_card", "card_padding"), default=6), 6) - - @property - def html_card_gap(self: ConfigBase) -> int: - """HTML 卡片间距。 - - 返回: - 多张图片之间的间距 - """ - return safe_int(self._read(("html_card", "card_gap"), default=6), 6) diff --git a/config/messages.py b/config/messages.py deleted file mode 100644 index 78e1cc1..0000000 --- a/config/messages.py +++ /dev/null @@ -1,100 +0,0 @@ -"""消息配置相关。""" - -from __future__ import annotations - -from typing import TYPE_CHECKING, Any - -from .helpers import safe_bool - -if TYPE_CHECKING: - pass - - -class MessagesConfigMixin: - """消息配置混入类。""" - - _read: Any - - @property - def msg_fetching_enabled(self) -> bool: - """是否启用获取中提示。 - - 返回: - 开始获取图片时是否发送提示消息 - """ - return safe_bool( - self._read(("messages", "fetching", "enabled"), default=True), - True, - ) - - @property - def msg_fetching_text(self) -> str: - """获取中提示文本。 - - 返回: - 开始获取图片时显示的文本 - """ - return str( - self._read( - ("messages", "fetching", "text"), - default="正在获取图片,请稍候...", - ) - ) - - @property - def msg_found_enabled(self) -> bool: - """是否启用找到图片提示。 - - 返回: - 成功找到图片后是否发送提示消息 - """ - return safe_bool( - self._read(("messages", "found", "enabled"), default=True), - True, - ) - - @property - def msg_found_text(self) -> str: - """找到图片提示文本。 - - 返回: - 成功找到图片后显示的文本,可使用 {count} 占位符 - """ - return str( - self._read( - ("messages", "found", "text"), - default="找到 {count} 张符合要求的图片~", - ) - ) - - @property - def msg_send_failed_text(self) -> str: - """发送失败提示文本。 - - 返回: - 图片发送失败时显示的文本 - """ - return str( - self._read( - ("messages", "send_failed", "text"), - "msg_send_failed_text", - default="图片发送失败,请稍后再试。", - ) - ) - - def format_found_message(self, count: int, revoke_delay: int | None = None) -> str: - """格式化找到图片的消息。 - - 将 msg_found_text 中的 {count} 和 {revoke_delay} 占位符替换为实际值。 - - 参数: - count: 找到的图片数量 - revoke_delay: 自动撤回延迟(秒),为 None 时不替换该变量 - - 返回: - 格式化后的消息文本 - """ - result = self.msg_found_text.replace("{count}", str(count)) - if revoke_delay is not None: - result = result.replace("{revoke_delay}", str(revoke_delay)) - return result diff --git a/config/safety.py b/config/safety.py deleted file mode 100644 index 2fdb968..0000000 --- a/config/safety.py +++ /dev/null @@ -1,262 +0,0 @@ -"""安全和缓存相关配置。""" - -from __future__ import annotations - -from typing import TYPE_CHECKING, Any - -from astrbot.api import logger - -from ..constants import DEFAULT_TAG_ALIAS -from .helpers import safe_bool, safe_int - -if TYPE_CHECKING: - pass - - -class SafetyConfigMixin: - """安全和缓存配置混入类。 - - 注意:用户级黑白名单已完全分离为 setu 和 fortune 独立的配置, - 请通过 AccessControlManager 直接使用功能级黑白名单方法。 - """ - - _read: Any - - @staticmethod - def _normalize_access_mode(value: Any, default: str = "none") -> str: - """规范化访问控制模式字符串。""" - if isinstance(value, str) and value in {"none", "blacklist", "whitelist"}: - return value - return default - - @property - def setu_access_control_mode(self) -> str: - """色图访问控制模式(旧配置,兼容保留)。""" - value = self._read(("safety", "setu_access_control_mode"), default="none") - return self._normalize_access_mode(value) - - @property - def fortune_access_control_mode(self) -> str: - """运势访问控制模式(旧配置,兼容保留)。""" - value = self._read(("safety", "fortune_access_control_mode"), default="none") - return self._normalize_access_mode(value) - - @property - def setu_user_access_control_mode(self) -> str: - """色图用户访问控制模式。""" - value = self._read( - ("safety", "setu_user_access_control_mode"), - "setu_user_access_control_mode", - default=self.setu_access_control_mode, - ) - return self._normalize_access_mode(value, self.setu_access_control_mode) - - @property - def setu_group_access_control_mode(self) -> str: - """色图群组访问控制模式。""" - value = self._read( - ("safety", "setu_group_access_control_mode"), - "setu_group_access_control_mode", - default=self.setu_access_control_mode, - ) - return self._normalize_access_mode(value, self.setu_access_control_mode) - - @property - def fortune_user_access_control_mode(self) -> str: - """运势用户访问控制模式。""" - value = self._read( - ("safety", "fortune_user_access_control_mode"), - "fortune_user_access_control_mode", - default=self.fortune_access_control_mode, - ) - return self._normalize_access_mode(value, self.fortune_access_control_mode) - - @property - def fortune_group_access_control_mode(self) -> str: - """运势群组访问控制模式。""" - value = self._read( - ("safety", "fortune_group_access_control_mode"), - "fortune_group_access_control_mode", - default=self.fortune_access_control_mode, - ) - return self._normalize_access_mode(value, self.fortune_access_control_mode) - - @property - def cache_enabled(self) -> bool: - """是否启用图片缓存。 - - 返回: - 是否启用 URL 图片磁盘缓存 - """ - return safe_bool(self._read(("cache", "enabled"), default=True), True) - - @property - def cache_ttl_hours(self) -> int: - """缓存 TTL(小时)。 - - 返回: - 缓存条目的存活时间 - """ - return safe_int(self._read(("cache", "ttl_hours"), default=2), 2) - - @property - def cache_max_items(self) -> int: - """最大缓存条目数。 - - 返回: - 缓存最多保留的条目数量 - """ - return safe_int(self._read(("cache", "max_items"), default=1), 1) - - @property - def cache_cleanup_on_start(self) -> bool: - """启动时清理缓存。 - - 返回: - 启动时是否自动清理过期缓存 - """ - return safe_bool(self._read(("cache", "cleanup_on_start"), default=True), True) - - @property - def download_concurrent_limit(self) -> int: - """并发下载限制。 - - 返回: - 同时下载图片的最大并发数,适用于高带宽服务器 - """ - return safe_int( - self._read(("performance", "download_concurrent_limit"), default=10), - 10, - ) - - @property - def download_timeout_seconds(self) -> int: - """下载超时时间(秒)。 - - 返回: - 单个图片下载的最大超时时间 - """ - return safe_int( - self._read(("performance", "download_timeout_seconds"), default=30), - 30, - ) - - @property - def enable_range_download(self) -> bool: - """启用分段下载。 - - 返回: - 是否将单张大图分成多段并行下载 - """ - return safe_bool( - self._read(("performance", "enable_range_download"), default=False), - False, - ) - - @property - def range_segments(self) -> int: - """分段下载的段数。 - - 返回: - 单张图片的分段下载数,2-4 段通常效果最佳 - """ - return safe_int( - self._read(("performance", "range_segments"), default=3), - 3, - ) - - @property - def range_threshold(self) -> int: - """分段下载阈值(KB)。 - - 返回: - 图片大于此值才启用分段下载(单位 KB) - """ - return safe_int( - self._read(("performance", "range_download_threshold"), default=512), - 512, - ) - - @property - def tag_alias(self) -> dict[str, list[str]]: - """标签别名映射。 - - 返回: - 标签到别名列表的映射字典 - """ - # 从 setu_general 读取(优先)或 safety(兼容旧配置) - alias_str = self._read( - ("setu_general", "tag_alias"), - ("safety", "tag_alias"), - "tag_alias", - default="", - ) - if not alias_str or not isinstance(alias_str, str): - return DEFAULT_TAG_ALIAS - - result: dict[str, list[str]] = {} - lines = alias_str.strip().replace("\r\n", "\n").split("\n") - for line in lines: - line = line.strip() - if not line or line.startswith("#") or line.startswith(";"): - continue - - if "=" not in line: - continue - key, value = line.split("=", 1) - key = key.strip() - value = value.strip() - if not key or not value: - continue - aliases = [a.strip() for a in value.split(",") if a.strip()] - if aliases: - result[key] = aliases - - if not result: - logger.debug("tag_alias parsing returned empty map, fallback to defaults") - return DEFAULT_TAG_ALIAS.copy() - return result - - def resolve_tags(self, raw_tag: str) -> list[str]: - """解析并规范化标签字符串。 - - 将逗号或空格分隔的标签字符串解析为列表,并应用别名映射。 - - 参数: - raw_tag: 原始标签字符串 - - 返回: - 规范化后的标签列表 - """ - if not raw_tag: - return [] - - normalized = raw_tag.replace(",", ",").replace(" ", ",") - tags = [t.strip() for t in normalized.split(",") if t.strip()] - - result: list[str] = [] - for tag in tags: - canonical = self._find_canonical_tag(tag) - result.append(canonical if canonical else tag) - return result - - def _find_canonical_tag(self, tag: str) -> str | None: - """查找标签的标准名称。 - - 参数: - tag: 标签名称(可能是别名) - - 返回: - 标准标签名称,如果未找到返回 None - """ - normalized = tag.lower() - for canonical, aliases in self.tag_alias.items(): - if not isinstance(canonical, str): - continue - if normalized == canonical.lower(): - return canonical - if isinstance(aliases, list): - for alias in aliases: - if isinstance(alias, str) and normalized == alias.lower(): - return canonical - return None diff --git a/core/__init__.py b/core/__init__.py deleted file mode 100644 index 594b88b..0000000 --- a/core/__init__.py +++ /dev/null @@ -1,975 +0,0 @@ -"""Setu 插件核心逻辑。""" - -from __future__ import annotations - -import asyncio -import random -from collections.abc import AsyncGenerator -from pathlib import Path -from typing import Any - -import httpx - -import astrbot.api.message_components as Comp -from astrbot.api import logger -from astrbot.api.event import AstrMessageEvent - -from ..config import SetuConfig -from ..providers import get_provider -from ..services import ( - AccessControlManager, - ConfigManager, - DocxService, - HtmlCardRenderer, - ImageService, - UrlImageDiskCache, -) -from ..session_config import SessionConfigManager -from .rate_limiter import SessionRequestLimiter -from .revoke_manager import RevokeManager -from .revoke_tasks import RevokeTaskMixin -from .send_revoke import SendWithRevokeMixin - - -class SetuCore(RevokeTaskMixin, SendWithRevokeMixin): - """Setu 插件核心逻辑类。""" - - def __init__(self, plugin, config: SetuConfig, data_dir: Path, astrbot_config=None): - self.plugin = plugin - self._config = config - self.data_dir = data_dir - self._revoke_manager = RevokeManager(data_dir) - self._session_config = SessionConfigManager( - astrbot_config or plugin.config, data_dir - ) - self._rate_limiter = SessionRequestLimiter() - self._docx_service = DocxService() - self._revoke_tasks: set[asyncio.Task] = set() - self._cache: UrlImageDiskCache | None = None - self._image_service: ImageService | None = None - self._html_renderer: HtmlCardRenderer | None = None - self._config_manager = ConfigManager(data_dir, astrbot_config or plugin.config) - self._access_control = AccessControlManager(self._config_manager) - - async def initialize(self) -> None: - """初始化核心组件。""" - await self._revoke_manager.initialize() - await self._session_config.initialize() - await self._config_manager.initialize() - await self._restore_pending_revokes() - - if ( - self._config.html_card_strategy in ("fallback", "always") - or self._config.auto_handle_send_failure - ): - self._ensure_html_renderer() - - try: - if self._config.cache_enabled: - cache_dir = self.data_dir / "image_cache" - self._cache = UrlImageDiskCache( - cache_dir=cache_dir, - ttl_hours=self._config.cache_ttl_hours, - max_items=self._config.cache_max_items, - ) - if self._config.cache_cleanup_on_start: - await self._cache.cleanup_expired() - self._image_service = ImageService( - self._cache, - concurrent_limit=self._config.download_concurrent_limit, - timeout_seconds=self._config.download_timeout_seconds, - enable_range_download=self._config.enable_range_download, - range_segments=self._config.range_segments, - range_threshold=self._config.range_threshold, - ) - except (OSError, RuntimeError, ValueError): - logger.exception( - "SetuCore initialize failed, fallback to no-cache ImageService" - ) - self._cache = None - self._image_service = ImageService( - None, - concurrent_limit=self._config.download_concurrent_limit, - timeout_seconds=self._config.download_timeout_seconds, - ) - - def terminate(self) -> None: - """终止插件,取消所有后台任务。""" - for task in list(self._revoke_tasks): - if not task.done(): - task.cancel() - self._revoke_tasks.clear() - logger.info("[revoke] All revoke tasks cancelled") - - def _get_provider(self): - cfg = self._config - lolicon_config = None - atri_config = None - if cfg.api_type in ("lolicon", "all"): - lolicon_config = { - "image_size": cfg.image_size, - "proxy": cfg.proxy, - "aspect_ratio": cfg.aspect_ratio, - "uid": cfg.uid, - "keyword": cfg.keyword, - } - if cfg.api_type in ("atri", "all"): - atri_config = { - "image_size": cfg.atri_image_size, - "proxy": cfg.atri_proxy, - "aspect_ratio": cfg.atri_aspect_ratio, - "uid": cfg.atri_uid, - "keyword": cfg.atri_keyword, - } - return get_provider( - cfg.api_type, - custom_config=cfg.custom_api if cfg.api_type == "custom" else None, - parser_config=cfg.api_response_parser if cfg.api_type == "custom" else None, - custom_api_configs=cfg.custom_api_configs - if cfg.api_type in ("custom", "all") - else None, - multi_api_strategy=cfg.multi_api_strategy, - lolicon_config=lolicon_config, - atri_config=atri_config, - ) - - def _get_fortune_provider(self): - """获取今日运势专用的 provider。 - - 支持独立的 API 提供商配置。 - """ - cfg = self._config - # 获取生效的 fortune API 类型 - effective_api_type = cfg.get_effective_fortune_api_type() - - lolicon_config = None - atri_config = None - if effective_api_type in ("lolicon", "all"): - lolicon_config = { - "image_size": cfg.image_size, - "proxy": cfg.proxy, - "aspect_ratio": cfg.aspect_ratio, - "uid": cfg.uid, - "keyword": cfg.keyword, - } - if effective_api_type in ("atri", "all"): - atri_config = { - "image_size": cfg.atri_image_size, - "proxy": cfg.atri_proxy, - "aspect_ratio": cfg.atri_aspect_ratio, - "uid": cfg.atri_uid, - "keyword": cfg.atri_keyword, - } - return get_provider( - effective_api_type, - custom_config=cfg.custom_api if effective_api_type == "custom" else None, - parser_config=cfg.api_response_parser - if effective_api_type == "custom" - else None, - custom_api_configs=cfg.custom_api_configs - if effective_api_type in ("custom", "all") - else None, - multi_api_strategy=cfg.multi_api_strategy, - lolicon_config=lolicon_config, - atri_config=atri_config, - ) - - async def get_effective_content_mode(self, event: AstrMessageEvent) -> str: - """获取生效的内容模式(优先会话配置)。""" - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - session_mode = await self._session_config.get_session_content_mode( - session_id, is_group - ) - - # 调试日志 - logger.debug( - "[content_mode] session_id=%s, is_group=%s, session_mode=%s, global_mode=%s", - session_id, - is_group, - session_mode, - self._config.content_mode, - ) - - if session_mode: - logger.debug("[content_mode] Using session mode: %s", session_mode) - return session_mode - logger.debug("[content_mode] Using global mode: %s", self._config.content_mode) - return self._config.content_mode - - async def get_effective_r18_docx_mode(self, event: AstrMessageEvent) -> bool: - """获取生效的 R18 Docx 模式设置。""" - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - session_mode = await self._session_config.get_session_r18_docx_mode( - session_id, is_group - ) - if session_mode is not None: - return session_mode - return self._config.r18_docx_mode - - async def get_effective_auto_revoke_r18(self, event: AstrMessageEvent) -> bool: - """获取生效的自动撤回 R18 设置。""" - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - session_mode = await self._session_config.get_session_auto_revoke_r18( - session_id, is_group - ) - if session_mode is not None: - return session_mode - return self._config.auto_revoke_r18 - - async def get_effective_send_mode(self, event: AstrMessageEvent) -> str: - """获取生效的发送模式(优先会话配置)。""" - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - session_mode = await self._session_config.get_session_send_mode( - session_id, is_group - ) - if session_mode: - logger.debug("[send_mode] Using session mode: %s", session_mode) - return session_mode - logger.debug("[send_mode] Using global mode: %s", self._config.send_mode) - return self._config.send_mode - - @staticmethod - def determine_r18(content_mode: str) -> bool: - """根据内容模式确定是否为 R18。""" - if content_mode == "r18": - return True - if content_mode == "mix": - return random.random() > 0.5 - return False - - @staticmethod - def _resolve_send_mode(send_mode: str, image_count: int) -> str: - if send_mode == "auto": - return "forward" if image_count > 1 else "image" - return send_mode - - def _ensure_html_renderer(self) -> bool: - if self._html_renderer is not None: - return True - try: - template_path = Path(__file__).parent.parent / "templates" / "setu.html" - self._html_renderer = HtmlCardRenderer(template_path) - return True - except (OSError, ValueError): - logger.exception("html renderer initialize failed") - return False - - def is_group_blocked( - self, event: AstrMessageEvent, feature: str | None = None - ) -> bool: - """检查群聊或用户是否被屏蔽。 - - 完全独立的功能级黑白名单:setu 和 fortune 有各自独立的黑白名单配置。 - 用户级黑白名单也是分开的:setu 的用户黑白名单只影响色图功能, - fortune 的用户黑白名单只影响运势功能。 - - 参数: - event: 消息事件 - feature: 功能名称,可选值为 "setu" 或 "fortune" - - 返回: - 如果被屏蔽返回 True,否则返回 False - """ - try: - # 获取用户 ID - user_id = event.get_sender_id() - - # 获取群组 ID(私聊时可能为 None) - group_id = None - try: - group_id = event.message_obj.group_id - except AttributeError: - pass - - # 根据功能类型检查对应的黑白名单 - if feature == "setu": - setu_user_mode = self._config.setu_user_access_control_mode - setu_group_mode = self._config.setu_group_access_control_mode - logger.debug( - "[setu] Access control mode from SetuConfig: user_mode=%s, group_mode=%s (user=%s, group=%s)", - setu_user_mode, - setu_group_mode, - user_id, - group_id, - ) - - is_blocked, reason = self._access_control.check_setu_access( - user_id, - group_id, - user_access_control_mode=setu_user_mode, - group_access_control_mode=setu_group_mode, - ) - if is_blocked: - logger.info( - "[setu] Access denied for user=%s, group=%s: %s", - user_id, - group_id, - reason, - ) - return True - elif feature == "fortune": - fortune_user_mode = self._config.fortune_user_access_control_mode - fortune_group_mode = self._config.fortune_group_access_control_mode - logger.debug( - "[fortune] Access control mode from SetuConfig: user_mode=%s, group_mode=%s (user=%s, group=%s)", - fortune_user_mode, - fortune_group_mode, - user_id, - group_id, - ) - - is_blocked, reason = self._access_control.check_fortune_access( - user_id, - group_id, - user_access_control_mode=fortune_user_mode, - group_access_control_mode=fortune_group_mode, - ) - if is_blocked: - logger.debug( - "[fortune] Access denied for user=%s, group=%s: %s", - user_id, - group_id, - reason, - ) - return True - else: - # 失败关闭:未指定/非法 feature 时,避免静默绕过访问控制。 - logger.warning( - "[access_control] Missing or invalid feature=%s, deny by default (user=%s, group=%s)", - feature, - user_id, - group_id, - ) - return True - - except AttributeError: - logger.debug("failed to inspect user/group id for blocked check") - - return False - - def _is_forward_supported(self, event: AstrMessageEvent) -> bool: - """检查当前平台是否支持合并转发消息。 - - 目前支持的平台:OneBot v11 (aiocqhttp) - """ - # 获取平台信息(多种方式) - platform_name = None - try: - platform_type_name = "Unknown" - # 方式1:通过 event.platform.name - if hasattr(event, "platform") and event.platform: - platform_obj = event.platform - platform_type_name = type(platform_obj).__name__ - if hasattr(platform_obj, "name"): - platform_name = platform_obj.name - - # 方式2:通过 event.get_platform_name() - if not platform_name and hasattr(event, "get_platform_name"): - try: - platform_name = event.get_platform_name() - except Exception: - pass - - # 方式3:通过 context 获取 - if not platform_name and hasattr(self, "plugin") and self.plugin: - try: - ctx = getattr(self.plugin, "context", None) - if ctx and hasattr(ctx, "get_platform"): - platform = ctx.get_platform() - if platform and hasattr(platform, "name"): - platform_name = platform.name - platform_type_name = type(platform).__name__ - except Exception: - pass - - # 调试日志 - logger.debug( - "[platform] Checking forward support: platform_name=%s, platform_type=%s", - platform_name, - platform_type_name, - ) - - # 支持的平台列表(检查名称) - supported_platforms = ( - "aiocqhttp", - "onebot11", - "onebot", - "go-cqhttp", - "napcat", - "llonebot", - ) - if platform_name: - platform_name_lower = platform_name.lower() - for p in supported_platforms: - if p in platform_name_lower: - logger.debug("[platform] Forward supported: name match '%s'", p) - return True - - # 检查 platform 对象类型名 - if platform_type_name and any( - p in platform_type_name.lower() for p in ("cqhttp", "onebot", "gocq") - ): - logger.debug( - "[platform] Forward supported: type match '%s'", platform_type_name - ) - return True - - # 检查 event 是否有 bot 属性且看起来是 OneBot(有 call_action 方法) - if hasattr(event, "bot") and event.bot: - bot = event.bot - if hasattr(bot, "call_action"): - logger.debug("[platform] Forward supported: bot has call_action") - return True - - # 检查 unified_msg_origin 格式 - if hasattr(event, "unified_msg_origin"): - umo = event.unified_msg_origin - if umo and isinstance(umo, str): - if any(p in umo.lower() for p in ("aiocqhttp", "onebot", "gocq")): - logger.debug( - "[platform] Forward supported: unified_msg_origin match" - ) - return True - - except Exception as exc: - logger.debug("[platform] Error checking forward support: %s", exc) - - logger.debug("[platform] Forward not supported for platform=%s", platform_name) - return False - - @property - def session_config(self) -> SessionConfigManager: - return self._session_config - - @property - def rate_limiter(self) -> SessionRequestLimiter: - return self._rate_limiter - - @property - def config(self) -> SetuConfig: - return self._config - - @property - def access_control(self) -> AccessControlManager: - return self._access_control - - async def fetch_and_download_images( - self, num: int, tags: list[str], is_r18: bool - ) -> list[bytes]: - """获取并下载图片。""" - try: - provider = self._get_provider() - except (ValueError, RuntimeError): - logger.exception("provider initialization failed") - return [] - - if not provider: - logger.error("no provider available") - return [] - - exclude_ai = self._config.exclude_ai - max_replenish = self._config.max_replenish_rounds - - try: - img_urls = await provider.fetch_image_urls( - num=num, tags=tags, r18=is_r18, exclude_ai=exclude_ai - ) - except (RuntimeError, ConnectionError, TimeoutError) as exc: - logger.error("provider fetch failed: %s", exc) - return [] - - if not img_urls or not self._image_service: - return [] - - downloaded = await self._image_service.download_parallel(img_urls) - - # 补充机制 - round_num = 0 - while len(downloaded) < num and round_num < max_replenish: - missing = num - len(downloaded) - if len(downloaded) == len(img_urls): - break - try: - extra_urls = await provider.fetch_image_urls( - num=missing, tags=tags, r18=is_r18, exclude_ai=exclude_ai - ) - if extra_urls: - extra_downloaded = await self._image_service.download_parallel( - extra_urls - ) - downloaded.extend(extra_downloaded) - if len(extra_urls) < missing: - break - else: - break - except Exception as exc: - logger.warning("replenish round %d failed: %s", round_num + 1, exc) - round_num += 1 - return downloaded - - async def _direct_send(self, event: AstrMessageEvent, chain: list) -> bool: - """通过 context.send_message 直接发送消息,返回是否成功。""" - try: - result = event.chain_result(chain) - - # 检查图片大小,Telegram 等平台有 10MB 或 20MB 限制 - for comp in chain: - if hasattr(comp, "file") and isinstance(comp.file, bytes): - size_mb = len(comp.file) / (1024 * 1024) - if size_mb > 10: # Telegram 通常限制 10-20MB - logger.warning( - "[send] Image size %.2f MB may exceed platform limit", - size_mb, - ) - - send_result = await self.plugin.context.send_message( - event.unified_msg_origin, result - ) - - # 获取平台名称 - platform_name = getattr(event.platform, "name", "unknown") - - # 检查发送结果 - # 注意:某些平台(如 Telegram)send_message 可能返回 None 但实际发送成功 - # 只对特定平台严格检查返回值 - if send_result is None and platform_name == "aiocqhttp": - # OneBot v11 平台应该返回消息 ID,如果没有可能是失败了 - logger.warning( - "[send] send_message returned None on %s, " - "platform may not support this message type", - platform_name, - ) - return False - - # 对于其他平台(Telegram 等),不严格检查返回值 - logger.debug("[send] Direct send completed on %s", platform_name) - return True - except TimeoutError: - logger.warning( - "[send] Direct send timed out on %s", - getattr(event.platform, "name", "unknown"), - ) - return False - except Exception as exc: - logger.warning("[send] Direct send failed: %s", exc, exc_info=True) - return False - - async def _direct_send_plain(self, event: AstrMessageEvent, text: str) -> bool: - """直接发送纯文本消息,返回是否成功。""" - try: - result = event.plain_result(text) - await self.plugin.context.send_message(event.unified_msg_origin, result) - return True - except Exception as exc: - logger.warning("direct send plain failed: %s", exc) - return False - - async def send_images( - self, - event: AstrMessageEvent, - images: list[bytes], - is_r18: bool, - tags: list[str] | None = None, - ) -> AsyncGenerator[Any, None]: - """发送图片,失败时尝试 HTML 卡片降级发送(复用已下载的图片)。""" - if not images: - yield event.plain_result("运气不好,一张图都没拿到...") - return - - cfg = self._config - # 使用会话级发送模式(优先会话配置) - send_mode = await self.get_effective_send_mode(event) - actual_send_mode = self._resolve_send_mode(send_mode, len(images)) - - effective_auto_revoke = await self.get_effective_auto_revoke_r18(event) - effective_r18_docx = await self.get_effective_r18_docx_mode(event) - auto_revoke = is_r18 and effective_auto_revoke - - # R18 Docx 模式 - if is_r18 and effective_r18_docx: - docx_path = self._docx_service.create_docx_with_images(images, tags=tags) - if docx_path: - if auto_revoke: - message_id = await self._send_file_with_revoke( - event, str(docx_path), docx_path.name - ) - if message_id: - await self._schedule_revoke( - event, message_id, cfg.auto_revoke_delay - ) - if cfg.msg_found_enabled: - yield event.plain_result( - cfg.format_found_message( - len(images), cfg.auto_revoke_delay - ) - ) - return - if cfg.msg_found_enabled: - yield event.plain_result(cfg.format_found_message(len(images))) - yield event.chain_result( - [Comp.File(file=str(docx_path), name=docx_path.name)] - ) - return - yield event.plain_result("R18 Docx 封装失败,请稍后再试或联系管理员。") - return - - # 准备提示消息 - found_message = ( - cfg.format_found_message(len(images)) if cfg.msg_found_enabled else None - ) - - # 判断 HTML 卡片策略 - html_strategy = cfg.html_card_strategy - use_html_always = html_strategy == "always" - use_html_fallback = html_strategy == "fallback" or cfg.auto_handle_send_failure - - # 如果策略是 always,直接使用 HTML 卡片发送 - if use_html_always: - if found_message: - await self._direct_send_plain(event, found_message) - send_success = await self._try_html_card_fallback( - event, images, actual_send_mode, auto_revoke - ) - if not send_success: - yield event.plain_result("HTML 卡片发送失败,请检查网络或联系管理员。") - return - - # 尝试普通发送(使用直接发送以捕获异常) - send_success = False - - # 检测平台是否支持合并转发 - forward_supported = self._is_forward_supported(event) - if actual_send_mode == "forward" and not forward_supported: - logger.debug( - "[send] Platform %s does not support forward messages, " - "falling back to normal image send", - getattr(event.platform, "name", "unknown"), - ) - actual_send_mode = "image" - - if actual_send_mode == "forward": - import time - - build_start = time.monotonic() - nodes = [] - for img_data in images: - node = Comp.Node( - uin=event.get_self_id(), - name="色图", - content=[Comp.Image.fromBytes(img_data)], - ) - nodes.append(node) - build_end = time.monotonic() - logger.debug( - "[forward] Built %d nodes in %.3fs", len(nodes), build_end - build_start - ) - - if auto_revoke: - message_id = await self._send_nodes_with_revoke(event, nodes) - if message_id: - await self._schedule_revoke( - event, message_id, cfg.auto_revoke_delay - ) - if cfg.msg_found_enabled: - yield event.plain_result( - cfg.format_found_message(len(images), cfg.auto_revoke_delay) - ) - send_success = True - else: - if found_message: - await self._direct_send_plain(event, found_message) - # 使用优化的合并转发发送(绕过 AstrBot 内部处理) - send_start = time.monotonic() - send_success = await self._send_nodes_direct(event, nodes) - logger.debug( - "[forward] Send nodes completed in %.3fs", - time.monotonic() - send_start, - ) - # 如果合并转发失败(平台不支持),降级为普通图片发送 - if not send_success: - logger.info( - "[forward] Forward message not supported or failed, " - "falling back to normal image send" - ) - chain = [Comp.Image.fromBytes(img) for img in images] - send_success = await self._direct_send(event, chain) - else: - if auto_revoke: - chain = [Comp.Image.fromBytes(img) for img in images] - message_id = await self._send_with_revoke_support( - event, - chain, - bool(event.get_group_id()), - event.get_group_id() or event.get_sender_id(), - ) - if message_id: - await self._schedule_revoke( - event, message_id, cfg.auto_revoke_delay - ) - if cfg.msg_found_enabled: - yield event.plain_result( - cfg.format_found_message(len(images), cfg.auto_revoke_delay) - ) - send_success = True - else: - if found_message: - await self._direct_send_plain(event, found_message) - # 将所有图片放在一条消息链中发送,提高速度和可靠性 - chain = [Comp.Image.fromBytes(img) for img in images] - send_success = await self._direct_send(event, chain) - - # 如果普通发送失败,尝试 HTML 卡片降级(使用已下载的图片) - if not send_success and use_html_fallback: - logger.debug("Image send failed, attempting HTML card fallback") - send_success = await self._try_html_card_fallback( - event, images, actual_send_mode, auto_revoke - ) - logger.debug("HTML card fallback result: %s", send_success) - - if not send_success: - if use_html_fallback: - yield event.plain_result( - "图片发送失败,HTML 卡片降级发送也失败了,请检查网络或联系管理员。" - ) - else: - yield event.plain_result( - "图片发送失败,可尝试在插件配置中启用 HTML 卡片模式作为备选方案。" - ) - else: - # 发送成功,yield 一个标记供调用者识别 - logger.debug("Images sent successfully: count=%d", len(images)) - yield {"send_success": True, "image_count": len(images)} - - async def _try_html_card_fallback( - self, - event: AstrMessageEvent, - images: list[bytes], - send_mode: str, - auto_revoke: bool = False, - ) -> bool: - """使用 HTML 卡片包装已下载的图片发送,返回是否成功。""" - cfg = self._config - - if not self._ensure_html_renderer() or not self._html_renderer: - logger.warning("HTML renderer not available for fallback") - return False - - render_style = { - "card_padding": cfg.html_card_padding, - "card_gap": cfg.html_card_gap, - } - - # 渲染 HTML 卡片(现在直接返回字节数据) - html_image_data: list[bytes] = [] - for i, img_data in enumerate(images): - logger.debug("[html_fallback] Rendering image %d/%d", i + 1, len(images)) - rendered = await self._html_renderer.render_single_image( - context=self.plugin, - image=img_data, - style_options=render_style, - ) - if rendered: - html_image_data.append(rendered) - logger.debug( - "[html_fallback] Image %d rendered, size=%d", i + 1, len(rendered) - ) - else: - logger.warning("[html_fallback] Failed to render image %d", i + 1) - - if not html_image_data: - logger.warning("HTML card rendering produced no images") - return False - - logger.debug( - "[html_fallback] Successfully rendered %d HTML cards", len(html_image_data) - ) - - # 发送渲染后的 HTML 卡片图片 - if send_mode == "forward" and self._is_forward_supported(event): - nodes = [] - for img_data in html_image_data: - node = Comp.Node( - uin=event.get_self_id(), - name="色图", - content=[Comp.Image.fromBytes(img_data)], - ) - nodes.append(node) - - if auto_revoke: - message_id = await self._send_nodes_with_revoke(event, nodes) - if message_id: - await self._schedule_revoke( - event, message_id, cfg.auto_revoke_delay - ) - if cfg.msg_found_enabled: - await self._direct_send_plain( - event, - cfg.format_found_message( - len(html_image_data), cfg.auto_revoke_delay - ), - ) - return True - return False - - send_success = await self._send_nodes_direct(event, nodes) - # 如果合并转发失败,降级为普通发送 - if not send_success: - logger.debug( - "[forward] HTML card forward send failed, falling back to normal send" - ) - chain = [Comp.Image.fromBytes(img) for img in html_image_data] - return await self._direct_send(event, chain) - return True - else: - if auto_revoke: - chain = [Comp.Image.fromBytes(img) for img in html_image_data] - message_id = await self._send_with_revoke_support( - event, - chain, - bool(event.get_group_id()), - event.get_group_id() or event.get_sender_id(), - ) - if message_id: - await self._schedule_revoke( - event, message_id, cfg.auto_revoke_delay - ) - if cfg.msg_found_enabled: - await self._direct_send_plain( - event, - cfg.format_found_message( - len(html_image_data), cfg.auto_revoke_delay - ), - ) - return True - return False - - # 所有 HTML 卡片图片放在一条消息链中发送 - chain = [Comp.Image.fromBytes(img) for img in html_image_data] - logger.debug( - "[html_fallback] Sending %d HTML card images via direct send", - len(chain), - ) - result = await self._direct_send(event, chain) - logger.info("[html_fallback] Direct send result: %s", result) - return result - - async def verify_image_urls(self, urls: list[str]) -> list[str]: - """验证图片 URL 是否可访问。 - - 使用 HEAD 请求快速验证 URL 是否返回 200。 - 支持并发请求以提高性能。 - - 参数: - urls: 图片 URL 列表 - - 返回: - 验证通过的 URL 列表 - """ - if not urls: - return [] - - valid_urls = [] - semaphore = asyncio.Semaphore(8) # 限制并发数为8 - - async def _check_single_url(client: httpx.AsyncClient, url: str) -> str | None: - async with semaphore: - try: - response = await client.head(url, follow_redirects=True) - if response.status_code == 200: - logger.debug("[url_verify] URL valid: %s", url) - return url - else: - logger.warning( - "[url_verify] URL returned %d: %s", - response.status_code, - url, - ) - except httpx.HTTPError as exc: - logger.warning( - "[url_verify] URL verification failed: %s - %s", url, exc - ) - return None - - async with httpx.AsyncClient(timeout=self._config.url_send_timeout) as client: - tasks = [_check_single_url(client, url) for url in urls] - results = await asyncio.gather(*tasks, return_exceptions=True) - - for result in results: - if isinstance(result, str): - valid_urls.append(result) - - return valid_urls - - async def send_images_by_url( - self, - event: AstrMessageEvent, - urls: list[str], - is_r18: bool, - tags: list[str] | None = None, - ): - """通过 URL 直接发送图片(不下载)。 - - 参数: - event: 消息事件 - urls: 图片 URL 列表 - is_r18: 是否为 R18 内容 - tags: 标签列表(保留用于接口一致性) - """ - # tags 参数保留用于接口一致性,当前实现不使用 - _ = tags - if not urls: - yield event.plain_result("运气不好,没有获取到图片链接...") - return - - cfg = self._config - - # 验证 URL(如果启用) - if cfg.url_send_verify: - valid_urls = await self.verify_image_urls(urls) - if not valid_urls: - yield event.plain_result("获取的图片链接均无法访问,请稍后再试。") - return - if len(valid_urls) < len(urls): - logger.warning( - "[url_send] %d/%d URLs invalid, using valid ones", - len(urls) - len(valid_urls), - len(urls), - ) - else: - valid_urls = urls - - # 准备提示消息 - found_message = ( - cfg.format_found_message(len(valid_urls)) if cfg.msg_found_enabled else None - ) - - # 发送图片 URL - if found_message: - await self._direct_send_plain(event, found_message) - - # 构建消息链 - 使用 Comp.Image.fromURL - chain = [Comp.Image.fromURL(url) for url in valid_urls] - - effective_auto_revoke = await self.get_effective_auto_revoke_r18(event) - auto_revoke = is_r18 and effective_auto_revoke - - if auto_revoke: - message_id = await self._send_with_revoke_support( - event, - chain, - bool(event.get_group_id()), - event.get_group_id() or event.get_sender_id(), - ) - if message_id: - await self._schedule_revoke(event, message_id, cfg.auto_revoke_delay) - if cfg.msg_found_enabled: - yield event.plain_result( - cfg.format_found_message(len(valid_urls), cfg.auto_revoke_delay) - ) - else: - await self._direct_send(event, chain) - yield {"send_success": True, "image_count": len(valid_urls)} diff --git a/core/rate_limiter.py b/core/rate_limiter.py deleted file mode 100644 index 253a18c..0000000 --- a/core/rate_limiter.py +++ /dev/null @@ -1,120 +0,0 @@ -"""请求限流管理器 - 实现会话级用户并发控制。""" - -from __future__ import annotations - -import asyncio -from typing import TYPE_CHECKING - -from astrbot.api import logger - -if TYPE_CHECKING: - from astrbot.api.event import AstrMessageEvent - - -class SessionRequestLimiter: - """会话级请求限流器。 - - 确保每个用户在每个会话中同时只能有一个正在处理的请求。 - 使用字典跟踪正在处理的请求,键格式为: "{session_id}:{user_id}"。 - - Attributes: - _locks: 存储每个用户的请求锁 - _global_lock: 用于保护 _locks 字典的锁 - """ - - def __init__(self) -> None: - self._locks: dict[str, asyncio.Lock] = {} - self._global_lock = asyncio.Lock() - - @staticmethod - def _get_key(session_id: str, user_id: str) -> str: - """生成唯一的请求键。""" - return f"{session_id}:{user_id}" - - async def acquire(self, event: AstrMessageEvent) -> bool: - """尝试获取请求锁。 - - 参数: - event: 消息事件对象 - - 返回: - 如果成功获取锁返回 True,如果用户已有请求在处理返回 False - """ - session_id = event.get_session_id() - user_id = event.get_sender_id() - key = self._get_key(session_id, user_id) - - async with self._global_lock: - if key in self._locks: - # 用户已有请求在处理中 - logger.debug( - "[rate_limit] Request rejected for user %s in session %s: " - "already has a pending request", - user_id, - session_id, - ) - return False - - self._locks[key] = asyncio.Lock() - lock = self._locks[key] - - try: - # 新锁理论上不会阻塞,但这里是可取消的 await 点 - await lock.acquire() - except asyncio.CancelledError: - async with self._global_lock: - current = self._locks.get(key) - if current is lock: - del self._locks[key] - raise - - logger.debug( - "[rate_limit] Request acquired for user %s in session %s", - user_id, - session_id, - ) - return True - - async def release(self, event: AstrMessageEvent) -> None: - """释放请求锁。 - - 参数: - event: 消息事件对象 - """ - session_id = event.get_session_id() - user_id = event.get_sender_id() - key = self._get_key(session_id, user_id) - - async with self._global_lock: - if key in self._locks: - try: - self._locks[key].release() - except RuntimeError: - # 锁可能已被释放 - pass - del self._locks[key] - logger.debug( - "[rate_limit] Request released for user %s in session %s", - user_id, - session_id, - ) - - async def is_pending(self, event: AstrMessageEvent) -> bool: - """检查用户是否有请求在处理中。 - - 参数: - event: 消息事件对象 - - 返回: - 如果有请求在处理中返回 True - """ - session_id = event.get_session_id() - user_id = event.get_sender_id() - key = self._get_key(session_id, user_id) - - async with self._global_lock: - return key in self._locks - - def get_pending_count(self) -> int: - """获取当前正在处理的请求数量。""" - return len(self._locks) diff --git a/core/revoke_manager.py b/core/revoke_manager.py deleted file mode 100644 index 3f94d9d..0000000 --- a/core/revoke_manager.py +++ /dev/null @@ -1,136 +0,0 @@ -"""撤回管理器模块。""" - -from __future__ import annotations - -import asyncio -import json -import time -from pathlib import Path -from typing import Any - -from astrbot.api import logger - - -class RevokeManager: - """管理 revoke.json,用于追踪被撤回的 R18 消息。""" - - def __init__(self, data_dir: Path): - self.data_dir = data_dir / "setu" - self.revoke_file = self.data_dir / "setu_revoke.json" - self._lock = asyncio.Lock() - self._data: dict[str, Any] = {"entries": {}, "meta": {}} - - async def initialize(self) -> None: - """初始化 revoke.json 文件。""" - self.data_dir.mkdir(parents=True, exist_ok=True) - # 检查并迁移旧数据(从 data_dir/revoke.json 到 data_dir/setu/revoke.json) - await self._migrate_old_data() - await self._load() - - async def _migrate_old_data(self) -> None: - """迁移旧位置的数据文件到新位置。""" - old_revoke_file = self.data_dir.parent / "revoke.json" - if old_revoke_file.exists() and not self.revoke_file.exists(): - try: - content = old_revoke_file.read_text(encoding="utf-8") - self._data = json.loads(content) - await self._save() - old_revoke_file.unlink() - logger.info("[revoke_manager] Migrated old revoke data to new location") - except (OSError, json.JSONDecodeError) as exc: - logger.warning("[revoke_manager] Failed to migrate old data: %s", exc) - - async def _load(self) -> None: - """从文件加载撤回数据。""" - if not self.revoke_file.exists(): - self._data = {"entries": {}, "meta": {"created_at": int(time.time())}} - await self._save() - return - try: - async with self._lock: - content = self.revoke_file.read_text(encoding="utf-8") - loaded = json.loads(content) - self._data = { - "entries": loaded.get("entries", {}), - "meta": loaded.get("meta", {}), - } - except (json.JSONDecodeError, OSError): - logger.exception("Failed to load revoke.json") - self._data = {"entries": {}, "meta": {"created_at": int(time.time())}} - await self._save() - - async def _save(self) -> None: - """保存撤回数据到文件。""" - try: - self.revoke_file.write_text( - json.dumps(self._data, ensure_ascii=False, indent=2), - encoding="utf-8", - ) - except OSError: - logger.exception("Failed to save revoke.json") - - async def add_entry( - self, - message_id: str, - platform: str, - session_id: str, - is_group: bool, - revoke_time: int, - ) -> None: - """添加撤回条目。""" - async with self._lock: - self._data["entries"][message_id] = { - "platform": platform, - "session_id": session_id, - "is_group": is_group, - "revoke_time": revoke_time, - "revoked": False, - "created_at": int(time.time()), - } - await self._save() - - async def mark_revoked(self, message_id: str) -> None: - """标记消息为已撤回,并清理旧记录。""" - async with self._lock: - if message_id in self._data["entries"]: - self._data["entries"][message_id]["revoked"] = True - self._data["entries"][message_id]["revoked_at"] = int(time.time()) - # 清理已撤销的过期记录(保留最近7天的) - await self._cleanup_revoked_entries() - await self._save() - - async def _cleanup_revoked_entries(self, max_age_days: int = 7) -> int: - """清理已撤销的旧记录。 - - 参数: - max_age_days: 已撤销记录保留天数 - - 返回: - 清理的记录数量 - """ - cutoff = int(time.time()) - (max_age_days * 24 * 3600) - to_remove = [] - - for message_id, entry in self._data["entries"].items(): - # 只清理已撤销的过期记录 - if entry.get("revoked", False): - revoked_at = entry.get("revoked_at", entry.get("created_at", 0)) - if revoked_at < cutoff: - to_remove.append(message_id) - - for message_id in to_remove: - del self._data["entries"][message_id] - - if to_remove: - logger.debug("[revoke] Cleaned up %d old revoked entries", len(to_remove)) - - return len(to_remove) - - def get_pending_entries(self) -> list[dict[str, Any]]: - """获取待撤回的条目列表。""" - entries = [] - for message_id, entry in self._data["entries"].items(): - if not entry.get("revoked", False): - entry["message_id"] = message_id - entries.append(entry) - return entries diff --git a/core/revoke_tasks.py b/core/revoke_tasks.py deleted file mode 100644 index 0f9a3e3..0000000 --- a/core/revoke_tasks.py +++ /dev/null @@ -1,152 +0,0 @@ -"""撤回任务处理混入类。""" - -from __future__ import annotations - -import asyncio -import time -from typing import TYPE_CHECKING, Any - -from astrbot.api import logger - -if TYPE_CHECKING: - from astrbot.api.event import AstrMessageEvent - - from .revoke_manager import RevokeManager - - -class RevokeTaskMixin: - """撤回任务处理混入类。""" - - _revoke_manager: RevokeManager - _revoke_tasks: set[asyncio.Task] - - async def _delayed_revoke( - self, - message_id: str, - delay: int, - _platform: str, - session_id: str, - is_group: bool, - bot_id: int | None, - bot: Any | None, - ) -> None: - """后台任务:在 delay 秒后撤回消息。""" - await asyncio.sleep(delay) - - success = False - actual_bot = bot - - # 如果 bot 为 None,尝试从插件 context 获取 - if not actual_bot and hasattr(self, "plugin") and self.plugin: - try: - # 尝试获取平台的 bot 对象 - context = getattr(self.plugin, "context", None) - if context: - # 通过 session 获取适配器,再获取 bot - platform = context.get_platform() - if platform and hasattr(platform, "bot"): - actual_bot = platform.bot - logger.debug( - "[revoke] Retrieved bot from platform for message %s", - message_id, - ) - except Exception as exc: - logger.debug("[revoke] Failed to retrieve bot from context: %s", exc) - - if actual_bot: - try: - params_list = [ - {"message_id": message_id}, - { - "message_id": int(message_id) - if str(message_id).isdigit() - else message_id - }, - ] - for params in params_list: - try: - await actual_bot.call_action("delete_msg", **params) - success = True - break - except (RuntimeError, ConnectionError, TimeoutError): - continue - except Exception as exc: - logger.warning("[revoke] Background revoke failed: %s", exc) - - if success: - await self._revoke_manager.mark_revoked(message_id) - logger.info("[revoke] Successfully revoked message %s", message_id) - else: - await self._revoke_manager.mark_revoked(message_id) - if actual_bot: - logger.warning( - "[revoke] Failed to revoke message %s, marked as revoked", - message_id, - ) - else: - logger.warning( - "[revoke] Cannot revoke message %s: no bot available, marked as revoked", - message_id, - ) - - async def _schedule_revoke( - self, event: AstrMessageEvent, message_id: str | int, delay: int - ) -> None: - """调度消息在 delay 秒后撤回。""" - if not message_id: - return - - platform = event.get_platform_name() - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - revoke_time = int(time.time()) + delay - - bot = getattr(event, "bot", None) - bot_id = id(bot) if bot else None - - await self._revoke_manager.add_entry( - str(message_id), platform, session_id, is_group, revoke_time - ) - - task = asyncio.create_task( - self._delayed_revoke( - str(message_id), delay, platform, session_id, is_group, bot_id, bot - ) - ) - self._revoke_tasks.add(task) - task.add_done_callback(self._revoke_tasks.discard) - - async def _restore_pending_revokes(self) -> None: - """恢复未处理的撤回任务(插件重启后)。""" - try: - pending = self._revoke_manager.get_pending_entries() - if not pending: - return - logger.info("[revoke] Restoring %d pending revoke tasks", len(pending)) - now = int(time.time()) - for entry in pending: - message_id = entry.get("message_id") - revoke_time = entry.get("revoke_time", 0) - if revoke_time <= now: - # 已过期,直接标记为已撤销 - await self._revoke_manager.mark_revoked(message_id) - logger.debug( - "[revoke] Expired entry marked as revoked: %s", message_id - ) - else: - delay = revoke_time - now - task = asyncio.create_task( - self._delayed_revoke( - message_id, - delay, - entry.get("platform", ""), - entry.get("session_id", ""), - entry.get("is_group", False), - None, - None, - ) - ) - self._revoke_tasks.add(task) - task.add_done_callback(self._revoke_tasks.discard) - except Exception as exc: - logger.exception("[revoke] Failed to restore pending revokes: %s", exc) diff --git a/core/send_revoke.py b/core/send_revoke.py deleted file mode 100644 index a43ad7e..0000000 --- a/core/send_revoke.py +++ /dev/null @@ -1,283 +0,0 @@ -"""图片发送混入类。""" - -from __future__ import annotations - -import asyncio -from pathlib import Path -from typing import TYPE_CHECKING, Any - -import astrbot.api.message_components as Comp -from astrbot.api import logger - -if TYPE_CHECKING: - from astrbot.api.event import AstrMessageEvent - - -class SendWithRevokeMixin: - """支持撤回的消息发送混入类。""" - - @staticmethod - def _get_bot_api(event: AstrMessageEvent) -> Any | None: - """从事件中获取底层 bot API 客户端。""" - return getattr(event, "bot", None) - - async def _send_forward_messages( - self, - event: AstrMessageEvent, - messages: list[dict], - is_group: bool, - session_id: str, - ) -> dict | None: - """发送合并转发消息的共享逻辑。 - - 返回 API 结果字典,出错时返回 None。 - """ - bot = self._get_bot_api(event) - if not bot: - logger.debug("[forward] No bot API available") - return None - - try: - if is_group: - return await bot.call_action( - "send_group_forward_msg", - group_id=int(session_id) - if str(session_id).isdigit() - else session_id, - messages=messages, - ) - return await bot.call_action( - "send_private_forward_msg", - user_id=int(session_id) if str(session_id).isdigit() else session_id, - messages=messages, - ) - except Exception: - logger.exception("[forward] Failed to send forward messages") - return None - - @staticmethod - def _check_forward_result(result: dict | None) -> tuple[bool, str | None]: - """检查合并转发 API 返回结果。 - - 返回: (是否成功, message_id 或 None) - """ - if not isinstance(result, dict): - return False, None - - # 检查是否有错误返回 - if result.get("status") == "failed" or result.get("retcode") not in (0, None): - logger.debug("[forward] API returned failure: %s", result) - return False, None - - data = result.get("data", result) - message_id = data.get("message_id") if isinstance(data, dict) else None - - if not message_id: - logger.debug( - "[forward] API returned no message_id, platform may not support forward" - ) - return False, None - - return True, str(message_id) if message_id else None - - async def _send_with_revoke_support( - self, - event: AstrMessageEvent, - chain: list[Any], - is_group: bool, - session_id: str, - ) -> str | None: - """发送消息并返回 message_id 以支持撤回。""" - try: - bot = self._get_bot_api(event) - if not bot: - return None - - messages = [] - for comp in chain: - if isinstance(comp, Comp.Plain): - if comp.text.strip(): - messages.append({"type": "text", "data": {"text": comp.text}}) - elif isinstance(comp, Comp.Image): - if comp.file and comp.file.startswith("base64://"): - messages.append({"type": "image", "data": {"file": comp.file}}) - elif comp.file: - messages.append({"type": "image", "data": {"file": comp.file}}) - elif comp.url: - messages.append({"type": "image", "data": {"file": comp.url}}) - elif isinstance(comp, Comp.File): - if comp.file: - messages.append({"type": "file", "data": {"file": comp.file}}) - - if not messages: - return None - - if is_group: - result = await bot.call_action( - "send_group_msg", - group_id=int(session_id) if session_id.isdigit() else session_id, - message=messages, - ) - else: - result = await bot.call_action( - "send_private_msg", - user_id=int(session_id) if session_id.isdigit() else session_id, - message=messages, - ) - - if isinstance(result, dict): - data = result.get("data", result) - message_id = data.get("message_id") if isinstance(data, dict) else None - if message_id: - return str(message_id) - return None - except Exception: - logger.exception("[revoke] Failed to send with revoke support") - return None - - async def _send_file_with_revoke( - self, event: AstrMessageEvent, file_path: str, file_name: str - ) -> str | None: - """发送文件并返回 message_id。""" - try: - import base64 - - bot = self._get_bot_api(event) - if not bot: - return None - - is_group = bool(event.get_group_id()) - session_id = event.get_group_id() or event.get_sender_id() - - path_obj = Path(file_path) - file_size = path_obj.stat().st_size - max_size = 5 * 1024 * 1024 - - if file_size > max_size: - logger.warning("[revoke] File too large (%d bytes)", file_size) - return None - - file_data = path_obj.read_bytes() - file_b64 = base64.b64encode(file_data).decode() - - messages = [ - { - "type": "file", - "data": {"file": f"base64://{file_b64}", "name": file_name}, - } - ] - - if is_group: - result = await bot.call_action( - "send_group_msg", - group_id=int(session_id) - if str(session_id).isdigit() - else session_id, - message=messages, - ) - else: - result = await bot.call_action( - "send_private_msg", - user_id=int(session_id) - if str(session_id).isdigit() - else session_id, - message=messages, - ) - - if isinstance(result, dict): - data = result.get("data", result) - message_id = data.get("message_id") if isinstance(data, dict) else None - if message_id: - return str(message_id) - return None - except (OSError, RuntimeError): - logger.exception("[revoke] Failed to send file with revoke support") - return None - - async def _send_nodes_with_revoke( - self, event: AstrMessageEvent, nodes: list[Comp.Node] - ) -> str | None: - """发送合并转发消息并返回 message_id。优化版本:并行转换 node 到 dict。""" - import time - - start_time = time.monotonic() - - is_group = bool(event.get_group_id()) - session_id = event.get_group_id() or event.get_sender_id() - - # 并行转换所有 nodes 到 dict,提高效率 - tasks = [node.to_dict() for node in nodes] - messages = list(await asyncio.gather(*tasks)) - - dict_time = time.monotonic() - logger.debug( - "[forward] Converted %d nodes to dict in %.3fs", - len(nodes), - dict_time - start_time, - ) - - # 使用共享逻辑发送 - result = await self._send_forward_messages( - event, messages, is_group, session_id - ) - - end_time = time.monotonic() - - # 检查结果 - success, message_id = self._check_forward_result(result) - if success: - logger.debug( - "[forward] Sent %d nodes in %.3fs (dict: %.3fs, api: %.3fs)", - len(nodes), - end_time - start_time, - dict_time - start_time, - end_time - dict_time, - ) - return message_id - return None - - async def _send_nodes_direct( - self, event: AstrMessageEvent, nodes: list[Comp.Node] - ) -> bool: - """直接发送合并转发消息(无撤回支持),优化版本。 - - 注意:合并转发消息仅在 OneBot v11 等特定平台受支持。 - 对于不支持的平台,会返回 False,调用方应降级为普通发送。 - """ - import time - - start_time = time.monotonic() - - is_group = bool(event.get_group_id()) - session_id = event.get_group_id() or event.get_sender_id() - - # 并行转换所有 nodes 到 dict - tasks = [node.to_dict() for node in nodes] - messages = list(await asyncio.gather(*tasks)) - - dict_time = time.monotonic() - logger.debug( - "[forward] Converted %d nodes to dict in %.3fs", - len(nodes), - dict_time - start_time, - ) - - # 使用共享逻辑发送 - result = await self._send_forward_messages( - event, messages, is_group, session_id - ) - - end_time = time.monotonic() - - # 检查结果 - success, _ = self._check_forward_result(result) - if success: - logger.debug( - "[forward] Sent %d nodes in %.3fs (dict: %.3fs, api: %.3fs)", - len(nodes), - end_time - start_time, - dict_time - start_time, - end_time - dict_time, - ) - return True - return False diff --git a/fortune/__init__.py b/fortune/__init__.py deleted file mode 100644 index b0ab789..0000000 --- a/fortune/__init__.py +++ /dev/null @@ -1,103 +0,0 @@ -"""今日运势模块。 - -集成到 Setu 插件的今日运势功能。 -参照 Java 版本 winefox-bot FortunePlugin 实现。 -""" - -from __future__ import annotations - -from pathlib import Path -from typing import Any - -from astrbot.api import logger - -from .core import FortuneCore -from .handlers import FortuneCommandHandler -from .llm_handlers import FortuneLlmHandler -from .renderer import FortuneRenderer -from .session_config import FortuneSessionConfig - - -class FortuneManager: - """今日运势管理器。 - - 封装今日运势的所有功能,便于在主插件中集成。 - """ - - def __init__( - self, plugin, data_dir: Path, config: dict[str, Any], astrbot_config=None - ): - self.plugin = plugin - self.data_dir = data_dir - self.config = config - self._astrbot_config = astrbot_config or plugin.config - - self._core: FortuneCore | None = None - self._renderer: FortuneRenderer | None = None - self._session_config: FortuneSessionConfig | None = None - self._cmd_handler: FortuneCommandHandler | None = None - self._llm_handler: FortuneLlmHandler | None = None - - async def initialize(self) -> None: - """初始化今日运势模块。""" - logger.info("[fortune] Initializing FortuneManager...") - - # 初始化核心组件 - self._core = FortuneCore(self.data_dir, self.config) - await self._core.initialize() - - self._renderer = FortuneRenderer() - self._session_config = FortuneSessionConfig(self._astrbot_config, self.data_dir) - await self._session_config.initialize() - - # 初始化处理器 - self._cmd_handler = FortuneCommandHandler( - self.plugin._core, # SetuCore - self.plugin.config, # SetuConfig - self._core, - self._renderer, - self._session_config, - ) - - self._llm_handler = FortuneLlmHandler( - self.plugin, - self.plugin._core, - self._core, - self._session_config, - ) - - logger.info("[fortune] FortuneManager initialized successfully") - - def terminate(self) -> None: - """清理资源。""" - logger.info("[fortune] FortuneManager terminated") - - @property - def cmd_handler(self) -> FortuneCommandHandler | None: - """获取命令处理器。""" - return self._cmd_handler - - @property - def llm_handler(self) -> FortuneLlmHandler | None: - """获取 LLM 处理器。""" - return self._llm_handler - - @property - def core(self) -> FortuneCore | None: - """获取核心实例。""" - return self._core - - @property - def session_config(self) -> FortuneSessionConfig | None: - """获取会话配置。""" - return self._session_config - - -__all__ = [ - "FortuneManager", - "FortuneCore", - "FortuneRenderer", - "FortuneSessionConfig", - "FortuneCommandHandler", - "FortuneLlmHandler", -] diff --git a/fortune/core.py b/fortune/core.py deleted file mode 100644 index cb4c8c0..0000000 --- a/fortune/core.py +++ /dev/null @@ -1,664 +0,0 @@ -"""今日运势核心模块。 - -管理运势生成、数据库持久化、图片缓存等核心功能。 -参照 Java 版本 FortuneDataServiceImpl 实现。 -""" - -from __future__ import annotations - -import asyncio -import datetime -import random -from pathlib import Path -from typing import Any - -from astrbot.api import logger - -# 默认权重数组(对应0-7星) -DEFAULT_WEIGHTS = [0.1, 0.15, 0.2, 0.25, 0.15, 0.12, 0.07, 0.005] - -# 默认运势标题(对应星级0-7) -DEFAULT_TITLES = ["凶", "末吉", "末小吉", "小吉", "中吉", "吉", "大吉", "超大吉"] - -# 默认运势描述(对应星级0-7) -DEFAULT_MESSAGES = [ - "长夜再暗,火种仍在,转机终会到来。", - "微光不灭,步步向前,黎明就在眼前。", - "心怀希冀,顺流而行,好事悄然靠近。", - "逆境翻篇,机遇迎面,惊喜不期而至。", - "小吉随身,难题化易,幸运与你并肩。", - "吉星高照,所行皆坦,所愿皆如愿。", - "福泽深厚,大吉加身,一路花开有声。", - "七星同耀,奇迹频现,今日万事皆成。", -] - - -class FortuneCore: - """今日运势核心类。""" - - def __init__(self, data_dir: Path, config: dict[str, Any]): - self.data_dir = data_dir - self.config = config - self.db_path = data_dir / "fortune.db" - self.cache_dir = data_dir / "cache" - self._db_lock = asyncio.Lock() - self._cache_lock = asyncio.Lock() - self._db_inited = False - self._last_cleanup_date: str | None = None - - async def initialize(self) -> None: - """初始化数据库和缓存目录。""" - self.data_dir.mkdir(parents=True, exist_ok=True) - self.cache_dir.mkdir(parents=True, exist_ok=True) - await self._init_db() - # 启动时清理过期缓存 - await self._cleanup_expired_cache() - # 启动定时清理任务 - asyncio.create_task(self._scheduled_cleanup()) - - async def _scheduled_cleanup(self) -> None: - """定时任务,每天凌晨执行:清理缓存 + 预生成活跃用户运势。""" - while True: - try: - # 计算到明天凌晨的时间 - now = datetime.datetime.now() - tomorrow = now + datetime.timedelta(days=1) - midnight = datetime.datetime.combine(tomorrow.date(), datetime.time.min) - wait_seconds = (midnight - now).total_seconds() - - logger.debug("[fortune] Scheduled tasks in %.0f seconds", wait_seconds) - await asyncio.sleep(wait_seconds) - - # 执行清理 - await self._cleanup_expired_cache() - # 执行预生成 - await self._pregenerate_active_users() - except asyncio.CancelledError: - break - except Exception as exc: - logger.warning("[fortune] Scheduled tasks failed: %s", exc) - await asyncio.sleep(3600) # 出错后1小时再试 - - async def _pregenerate_active_users(self) -> int: - """预生成活跃用户的今日运势。 - - 只为3天内有查看记录的用户预生成,避免滥用。 - 如果今天已有运势则不覆盖。 - - 返回: - 预生成的用户数量 - """ - import aiosqlite - - today_str = datetime.date.today().isoformat() - # 计算3天前的日期 - three_days_ago = ( - datetime.date.today() - datetime.timedelta(days=3) - ).isoformat() - - pregenerated = 0 - - try: - async with self._db_lock: - async with aiosqlite.connect(str(self.db_path)) as db: - # 查询3天内有查看记录的用户(且今天还没有运势) - cursor = await db.execute( - """ - SELECT DISTINCT user_id FROM fortune_data - WHERE last_view_date >= ? - AND user_id NOT IN ( - SELECT user_id FROM fortune_data WHERE date_str = ? - ) - """, - (three_days_ago, today_str), - ) - active_users = await cursor.fetchall() - - if not active_users: - logger.debug("[fortune] No active users to pregenerate") - return 0 - - logger.info( - "[fortune] Pregenerating fortunes for %d active users", - len(active_users), - ) - - # 为每个活跃用户预生成运势 - for (user_id,) in active_users: - try: - # 再次检查今天是否已有运势(避免并发竞争) - check_cursor = await db.execute( - "SELECT 1 FROM fortune_data WHERE user_id = ? AND date_str = ?", - (user_id, today_str), - ) - if await check_cursor.fetchone(): - # 已有运势,跳过 - continue - - # 生成运势 - fortune = self._generate_fortune( - user_id, "指挥官", today_str - ) - - # 保存到数据库(使用 INSERT OR IGNORE 避免覆盖) - await db.execute( - "INSERT INTO fortune_data " - "(user_id, date_str, title, stars, desc_text, extra, theme, image_cached, img_url, last_view_date) " - "VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", - ( - user_id, - today_str, - fortune["title"], - fortune["star_count"], - fortune["description"], - fortune["extra_message"], - fortune["theme_color"], - 0, - None, - today_str, # 预生成时设置 last_view_date,保持活跃状态 - ), - ) - pregenerated += 1 - except Exception as exc: - logger.warning( - "[fortune] Failed to pregenerate for user %s: %s", - user_id, - exc, - ) - - await db.commit() - - if pregenerated > 0: - logger.info( - "[fortune] Pregenerated %d/%d fortunes for active users", - pregenerated, - len(active_users), - ) - - except Exception as exc: - logger.error("[fortune] Pregeneration failed: %s", exc) - - return pregenerated - - async def _init_db(self) -> None: - """初始化数据库表。""" - import aiosqlite - - async with self._db_lock: - async with aiosqlite.connect(str(self.db_path)) as db: - # 创建表(如果不存在) - await db.execute( - """ - CREATE TABLE IF NOT EXISTS fortune_data ( - user_id TEXT NOT NULL, - date_str TEXT NOT NULL, - title TEXT NOT NULL, - stars INTEGER NOT NULL, - desc_text TEXT NOT NULL, - extra TEXT NOT NULL, - theme TEXT NOT NULL, - image_cached INTEGER DEFAULT 0, - img_url TEXT, - last_view_date TEXT, - PRIMARY KEY (user_id, date_str) - ) - """ - ) - await db.commit() - - # 检查并添加缺失的列(版本迁移) - await self._migrate_db(db) - - self._db_inited = True - logger.info("[fortune] Database initialized") - - async def _migrate_db(self, db) -> None: - """数据库版本迁移:添加缺失的列。""" - # 获取表中所有列 - cursor = await db.execute("PRAGMA table_info(fortune_data)") - columns = await cursor.fetchall() - column_names = [col[1] for col in columns] - - # 检查并添加 img_url 列 - if "img_url" not in column_names: - logger.info("[fortune] Migrating database: adding img_url column") - await db.execute("ALTER TABLE fortune_data ADD COLUMN img_url TEXT") - await db.commit() - logger.info("[fortune] Migration completed") - - # 检查并添加 last_view_date 列(用于预生成机制) - if "last_view_date" not in column_names: - logger.info("[fortune] Migrating database: adding last_view_date column") - await db.execute("ALTER TABLE fortune_data ADD COLUMN last_view_date TEXT") - await db.commit() - logger.info("[fortune] Migration completed") - - # 检查并添加 group_id 列(用于群组运势刷新) - if "group_id" not in column_names: - logger.info("[fortune] Migrating database: adding group_id column") - await db.execute("ALTER TABLE fortune_data ADD COLUMN group_id TEXT") - await db.commit() - # 创建索引以优化按群组查询 - await db.execute( - "CREATE INDEX IF NOT EXISTS idx_fortune_group ON fortune_data(group_id, date_str)" - ) - await db.commit() - logger.info("[fortune] Migration completed") - - @staticmethod - def _get_weights() -> list[float]: - """获取权重配置(固定值)。""" - return DEFAULT_WEIGHTS - - @staticmethod - def _get_titles() -> list[str]: - """获取标题配置(固定值)。""" - return DEFAULT_TITLES - - @staticmethod - def _get_messages() -> list[str]: - """获取描述文案配置(固定值)。""" - return DEFAULT_MESSAGES - - @staticmethod - def _get_extra_message() -> str: - """获取额外消息配置(固定为空)。""" - return "" - - def _calculate_luck(self) -> int: - """根据权重计算运势星级。 - - 返回: - 星级 0-7 - """ - weights = self._get_weights() - total_weight = sum(weights) - random_value = random.random() * total_weight - current_weight = 0 - - for i, w in enumerate(weights): - current_weight += w - if random_value <= current_weight: - return i - return 0 - - def _get_theme_color(self, star_num: int) -> str: - """根据星级获取主题颜色。 - - Java版本: - - 7, 6星 -> theme-red - - 5, 4星 -> theme-gold - - 1, 0星 -> theme-gray - - 其他 -> theme-blue - """ - if star_num in (7, 6): - return "theme-red" - elif star_num in (5, 4): - return "theme-gold" - elif star_num in (1, 0): - return "theme-gray" - else: - return "theme-blue" - - def _generate_fortune( - self, user_id: str, username: str, date_str: str, need_new: bool = True - ) -> dict[str, Any]: - """生成运势数据。 - - 参数: - user_id: 用户ID - username: 用户名 - date_str: 日期字符串 - need_new: 是否生成新的运势(刷新时使用) - - 返回: - 运势数据字典 - """ - # 计算星级 - star_num = self._calculate_luck() - - # 获取配置 - titles = self._get_titles() - messages = self._get_messages() - extra = self._get_extra_message() - - # 获取对应星级的标题和描述 - title = titles[min(star_num, len(titles) - 1)] - desc = messages[min(star_num, len(messages) - 1)] - theme = self._get_theme_color(star_num) - - return { - "user_id": user_id, - "username": username or "指挥官", - "date_str": date_str, - "title": title, - "star_count": star_num, - "max_stars": 7, - "description": desc, - "extra_message": extra, - "theme_color": theme, - "image_cached": False, - "img_url": None, - } - - def _get_cache_path(self, user_id: str, date_str: str) -> Path: - """获取用户图片缓存路径。""" - return self.cache_dir / f"{user_id}_{date_str}.jpg" - - async def _cleanup_expired_cache(self) -> int: - """清理过期的缓存文件(非今天的)。 - - 使用 _last_cleanup_date 避免同一天重复清理。 - """ - today = datetime.date.today().isoformat() - - # 检查是否已经清理过今天 - if self._last_cleanup_date == today: - return 0 - - removed = 0 - removed_size = 0 - - try: - async with self._cache_lock: - for file_path in self.cache_dir.iterdir(): - if not file_path.is_file(): - continue - - # 解析文件名中的日期 (格式: {user_id}_{date}.jpg) - try: - # 提取日期部分,假设格式为 user_id_YYYY-MM-DD.jpg - file_name = file_path.stem # 去掉 .jpg - parts = file_name.rsplit( - "_", 1 - ) # 从右边分割,得到 [user_id, date] - - if len(parts) >= 2: - file_date = parts[-1] - # 验证日期格式 - if file_date != today: - # 非今天的缓存,删除 - file_size = file_path.stat().st_size - file_path.unlink() - removed += 1 - removed_size += file_size - else: - # 无法解析的文件名,也删除 - file_size = file_path.stat().st_size - file_path.unlink() - removed += 1 - removed_size += file_size - except (OSError, ValueError): - # 删除异常文件 - try: - file_size = file_path.stat().st_size - file_path.unlink() - removed += 1 - removed_size += file_size - except OSError: - pass - - # 更新最后清理日期 - if removed > 0: - logger.info( - "[fortune] Cleaned up %d expired cache files (%.2f MB)", - removed, - removed_size / 1024 / 1024, - ) - else: - logger.debug("[fortune] No expired cache files to clean") - - self._last_cleanup_date = today - - except OSError as exc: - logger.warning("[fortune] Failed to cleanup cache: %s", exc) - - return removed - - async def get_today_fortune( - self, user_id: str, username: str, group_id: str | None = None - ) -> dict[str, Any] | None: - """获取今日运势。 - - 仅查询今天的记录;如果今天尚未生成,则立即生成并保存。 - 每次获取都会更新 last_view_date。 - - 参数: - user_id: 用户ID - username: 用户名 - group_id: 群组ID(可选,用于群维度运势管理) - - 返回: - 运势数据字典 - """ - import aiosqlite - - today_str = datetime.date.today().isoformat() - - # 先检查并清理过期缓存(每天只执行一次) - await self._cleanup_expired_cache() - - async with self._db_lock: - async with aiosqlite.connect(str(self.db_path)) as db: - # 仅查询今天的记录 - cursor = await db.execute( - "SELECT title, stars, desc_text, extra, theme, image_cached, img_url " - "FROM fortune_data WHERE user_id = ? AND date_str = ?", - (user_id, today_str), - ) - row = await cursor.fetchone() - - if row: - # 今天已有运势,更新最后查看日期并返回 - await db.execute( - "UPDATE fortune_data SET last_view_date = ?, group_id = COALESCE(?, group_id) WHERE user_id = ? AND date_str = ?", - (today_str, group_id, user_id, today_str), - ) - await db.commit() - return { - "user_id": user_id, - "username": username, - "date_str": today_str, - "title": row[0], - "star_count": row[1], - "max_stars": 7, - "description": row[2], - "extra_message": row[3], - "theme_color": row[4], - "image_cached": bool(row[5]), - "img_url": row[6], - } - - # 今天没有记录,生成新运势 - fortune = self._generate_fortune(user_id, username, today_str) - - # 保存到数据库(包含 group_id 和 last_view_date) - await db.execute( - "INSERT OR REPLACE INTO fortune_data " - "(user_id, date_str, title, stars, desc_text, extra, theme, image_cached, img_url, last_view_date, group_id) " - "VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", - ( - user_id, - today_str, - fortune["title"], - fortune["star_count"], - fortune["description"], - fortune["extra_message"], - fortune["theme_color"], - 0, - None, - today_str, - group_id, - ), - ) - await db.commit() - - return fortune - - async def update_fortune_image_cache( - self, user_id: str, date_str: str, image_data: bytes, img_url: str | None = None - ) -> Path: - """更新运势图片缓存。""" - cache_path = self._get_cache_path(user_id, date_str) - - async with self._cache_lock: - await asyncio.to_thread(cache_path.write_bytes, image_data) - - # 更新数据库标记 - import aiosqlite - - async with self._db_lock: - async with aiosqlite.connect(str(self.db_path)) as db: - await db.execute( - "UPDATE fortune_data SET image_cached = 1, img_url = ? " - "WHERE user_id = ? AND date_str = ?", - (img_url, user_id, date_str), - ) - await db.commit() - - return cache_path - - async def get_cached_image(self, user_id: str, date_str: str) -> bytes | None: - """获取缓存的图片数据。""" - cache_path = self._get_cache_path(user_id, date_str) - - try: - async with self._cache_lock: - if cache_path.exists(): - return await asyncio.to_thread(cache_path.read_bytes) - except OSError: - pass - return None - - async def refresh_fortune(self, user_id: str, username: str) -> dict[str, Any]: - """刷新用户今日运势(生成新的)。""" - import aiosqlite - - date_str = datetime.date.today().isoformat() - - # 生成新运势 - fortune = self._generate_fortune(user_id, username, date_str) - - # 更新数据库(包含 last_view_date) - async with self._db_lock: - async with aiosqlite.connect(str(self.db_path)) as db: - await db.execute( - "INSERT OR REPLACE INTO fortune_data " - "(user_id, date_str, title, stars, desc_text, extra, theme, image_cached, img_url, last_view_date) " - "VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", - ( - user_id, - date_str, - fortune["title"], - fortune["star_count"], - fortune["description"], - fortune["extra_message"], - fortune["theme_color"], - 0, - None, - date_str, # 刷新时更新 last_view_date - ), - ) - await db.commit() - - # 删除旧缓存 - cache_path = self._get_cache_path(user_id, date_str) - try: - cache_path.unlink(missing_ok=True) - except OSError: - pass - - return fortune - - async def refresh_group_fortune(self, group_id: str) -> int: - """刷新指定群的所有今日运势。 - - 返回: - 被刷新的记录数量 - """ - import aiosqlite - - date_str = datetime.date.today().isoformat() - - async with self._db_lock: - async with aiosqlite.connect(str(self.db_path)) as db: - # 获取今天该群的所有记录(group_id 匹配或 NULL) - cursor = await db.execute( - "SELECT user_id FROM fortune_data WHERE date_str = ? AND (group_id = ? OR group_id IS NULL)", - (date_str, group_id), - ) - rows = await cursor.fetchall() - - # 删除这些记录 - await db.execute( - "DELETE FROM fortune_data WHERE date_str = ? AND (group_id = ? OR group_id IS NULL)", - (date_str, group_id), - ) - await db.commit() - - # 清理缓存文件(只清理该群用户的缓存) - async with self._cache_lock: - for file_path in self.cache_dir.iterdir(): - if file_path.is_file() and date_str in file_path.name: - # 检查是否属于该群用户 - user_id_from_file = ( - file_path.stem.split("_")[0] if "_" in file_path.stem else None - ) - if user_id_from_file and any( - row[0] == user_id_from_file for row in rows - ): - try: - file_path.unlink() - except OSError: - pass - - logger.info( - "[fortune] Refreshed %d fortunes for group %s on date %s", - len(rows), - group_id, - date_str, - ) - return len(rows) - - async def refresh_all_fortune(self) -> int: - """刷新全局今日运势。 - - 返回: - 被删除的记录数量 - """ - import aiosqlite - - date_str = datetime.date.today().isoformat() - - async with self._db_lock: - async with aiosqlite.connect(str(self.db_path)) as db: - cursor = await db.execute( - "SELECT COUNT(*) FROM fortune_data WHERE date_str = ?", - (date_str,), - ) - row = await cursor.fetchone() - count = row[0] if row else 0 - - await db.execute( - "DELETE FROM fortune_data WHERE date_str = ?", - (date_str,), - ) - await db.commit() - - # 清理所有今天的缓存文件 - async with self._cache_lock: - for file_path in self.cache_dir.iterdir(): - if file_path.is_file() and date_str in file_path.name: - try: - file_path.unlink() - except OSError: - pass - - logger.info("[fortune] Refreshed all %d fortunes for date %s", count, date_str) - return count - - def format_stars(self, count: int, max_count: int = 7) -> str: - """格式化星级显示。""" - filled = "★" * count - empty = "☆" * (max_count - count) - return filled + empty diff --git a/fortune/handlers.py b/fortune/handlers.py deleted file mode 100644 index 3d92a2f..0000000 --- a/fortune/handlers.py +++ /dev/null @@ -1,425 +0,0 @@ -"""今日运势命令处理器。""" - -from __future__ import annotations - -from typing import TYPE_CHECKING - -from astrbot.api import logger -from astrbot.api.event import AstrMessageEvent - -if TYPE_CHECKING: - from ..config import SetuConfig - from ..core import SetuCore - from .core import FortuneCore - from .renderer import FortuneRenderer - from .session_config import FortuneSessionConfig - - -class FortuneCommandHandler: - """今日运势命令处理器。""" - - def __init__( - self, - core: SetuCore, - config: SetuConfig, - fortune_core: FortuneCore, - fortune_renderer: FortuneRenderer, - session_config: FortuneSessionConfig, - ): - self._core = core - self._config = config - self._fortune_core = fortune_core - self._fortune_renderer = fortune_renderer - self._session_config = session_config - - @staticmethod - def _check_admin(event: AstrMessageEvent) -> bool: - """检查用户是否为管理员。""" - try: - if hasattr(event, "is_admin") and callable(getattr(event, "is_admin")): - if event.is_admin(): - return True - if hasattr(event, "is_super_user") and callable( - getattr(event, "is_super_user") - ): - if event.is_super_user(): - return True - if hasattr(event, "message_obj"): - msg_obj = event.message_obj - if hasattr(msg_obj, "sender") and hasattr(msg_obj.sender, "role"): - role = msg_obj.sender.role - if role in ("admin", "owner"): - return True - except AttributeError: - pass - return False - - async def _get_image_for_fortune(self, event: AstrMessageEvent) -> bytes | None: - """获取今日运势的图片。 - - 使用 Setu 插件的图片供应商,根据配置获取图片。 - """ - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - - # 获取生效的配置 - fortune_cfg = getattr(self._config, "fortune", {}) - global_tags = ( - fortune_cfg.get("tags", "") if isinstance(fortune_cfg, dict) else "" - ) - global_mode = ( - fortune_cfg.get("content_mode", "sfw") - if isinstance(fortune_cfg, dict) - else "sfw" - ) - - tags_str = await self._session_config.get_effective_tags( - session_id, is_group, global_tags - ) - content_mode = await self._session_config.get_effective_content_mode( - session_id, is_group, global_mode - ) - - # 解析标签 - tags = [t.strip() for t in tags_str.split(",") if t.strip()] if tags_str else [] - - # 确定 R18 - import random - - is_r18 = False - if content_mode == "r18": - is_r18 = True - elif content_mode == "mix": - is_r18 = random.random() > 0.5 - - logger.debug( - "[fortune] Getting image: tags=%s, mode=%s, is_r18=%s", - tags, - content_mode, - is_r18, - ) - - try: - # 使用 SetuCore 的 fortune_provider 获取图片 - provider = self._core._get_fortune_provider() - if not provider: - logger.error("[fortune] No provider available") - return None - - # 获取 1 张图片 URL - img_urls = await provider.fetch_image_urls( - num=1, tags=tags, r18=is_r18, exclude_ai=self._config.exclude_ai - ) - - if not img_urls: - logger.warning("[fortune] No image URLs returned") - return None - - # 下载图片 - downloaded = await self._core.fetch_and_download_images(1, tags, is_r18) - - if downloaded and len(downloaded) > 0: - return downloaded[0] - - except Exception as exc: - logger.exception("[fortune] Failed to get image: %s", exc) - - return None - - async def handle_fortune(self, event: AstrMessageEvent): - """处理 /jrys 或 /今日运势 命令。""" - # 检查全局黑白名单和运势功能级黑名单 - if self._core.is_group_blocked(event, feature="fortune"): - return - - user_id = event.get_sender_id() - username = event.get_sender_name() or user_id - group_id = event.get_group_id() - - # 获取今日运势数据 - fortune = await self._fortune_core.get_today_fortune( - user_id, username, group_id - ) - if not fortune: - yield event.plain_result("运势获取失败,请稍后重试。") - return - - # 检查是否已有缓存的渲染图片 - cached_image = await self._fortune_core.get_cached_image( - user_id, fortune["date_str"] - ) - - if cached_image: - # 使用缓存的图片 - logger.debug("[fortune] Using cached image for %s", user_id) - import astrbot.api.message_components as Comp - - yield event.chain_result([Comp.Image.fromBytes(cached_image)]) - return - - try: - # 获取图片 - image_data = await self._get_image_for_fortune(event) - - if not image_data: - # 图片获取失败,发送文字版运势 - logger.warning("[fortune] Failed to get image, sending text version") - stars = "★" * fortune["star_count"] + "☆" * (7 - fortune["star_count"]) - msg = ( - f"【今日运势】\n" - f"用户:{fortune['username']}\n" - f"日期:{fortune['date_str']}\n" - f"运势:{fortune['title']}\n" - f"星级:{stars}\n" - f"\n{fortune['description']}" - ) - yield event.plain_result(msg) - return - - # 转换为 base64 用于渲染 - import base64 - - image_base64 = base64.b64encode(image_data).decode("utf-8") - - # 获取 AstrBot 的 html_render 方法 - html_renderer = None - if ( - hasattr(self._core, "plugin") - and self._core.plugin - and hasattr(self._core.plugin, "html_render") - ): - html_renderer = self._core.plugin.html_render - - if not html_renderer: - # 没有 HTML 渲染器,发送文字版运势 - logger.warning( - "[fortune] No HTML renderer available, sending text version" - ) - stars = "★" * fortune["star_count"] + "☆" * (7 - fortune["star_count"]) - msg = ( - f"【今日运势】\n" - f"用户:{fortune['username']}\n" - f"日期:{fortune['date_str']}\n" - f"运势:{fortune['title']}\n" - f"星级:{stars}\n" - f"\n{fortune['description']}" - ) - yield event.plain_result(msg) - return - - # 渲染为图片(使用 Jinja2 模板引擎) - logger.debug("[fortune] Rendering image with html_renderer...") - rendered_image = await self._fortune_renderer.render_to_image( - fortune=fortune, image_base64=image_base64, html_renderer=html_renderer - ) - - if not rendered_image: - logger.warning( - "[fortune] Image render returned None/empty, sending text version" - ) - stars = "★" * fortune["star_count"] + "☆" * (7 - fortune["star_count"]) - msg = ( - f"【今日运势】\n" - f"用户:{fortune['username']}\n" - f"日期:{fortune['date_str']}\n" - f"运势:{fortune['title']}\n" - f"星级:{stars}\n" - f"\n{fortune['description']}" - ) - yield event.plain_result(msg) - return - - logger.debug( - "[fortune] Image rendered successfully: %d bytes", len(rendered_image) - ) - - # 保存到缓存 - await self._fortune_core.update_fortune_image_cache( - user_id, fortune["date_str"], rendered_image - ) - - # 发送图片 - import astrbot.api.message_components as Comp - - yield event.chain_result([Comp.Image.fromBytes(rendered_image)]) - - except Exception as exc: - logger.exception("[fortune] Failed to process fortune: %s", exc) - yield event.plain_result("今日运势生成失败,请稍后重试。") - - async def handle_refresh_fortune(self, event: AstrMessageEvent): - """处理 /刷新今日运势 命令(个人)。""" - if not self._check_admin(event): - # 检查配置是否允许非管理员刷新 - fortune_cfg = getattr(self._config, "fortune", {}) - allow_refresh = ( - fortune_cfg.get("allow_user_refresh", False) - if isinstance(fortune_cfg, dict) - else False - ) - if not allow_refresh: - yield event.plain_result("只有管理员可以刷新运势。") - return - - user_id = event.get_sender_id() - username = event.get_sender_name() or user_id - - try: - # 刷新运势 - await self._fortune_core.refresh_fortune(user_id, username) - yield event.plain_result( - "已刷新你的今日运势,发送 今日运势 或 jrys 查看新的运势。" - ) - except Exception as exc: - logger.exception("[fortune] Failed to refresh fortune: %s", exc) - yield event.plain_result("刷新失败,请稍后重试。") - - async def handle_refresh_group_fortune(self, event: AstrMessageEvent): - """处理 /刷新本群今日运势 命令(仅管理员)。""" - if not self._check_admin(event): - yield event.plain_result("只有管理员可以刷新群组运势。") - return - - group_id = event.get_group_id() - if not group_id: - yield event.plain_result("此命令只能在群聊中使用。") - return - - try: - # 刷新群组运势(删除该群所有今天的记录) - count = await self._fortune_core.refresh_group_fortune(str(group_id)) - yield event.plain_result( - f"已刷新本群今日运势,共 {count} 条记录将被重新生成。" - ) - except Exception as exc: - logger.exception("[fortune] Failed to refresh group fortune: %s", exc) - yield event.plain_result("刷新失败,请稍后重试。") - - async def handle_refresh_all_fortune(self, event: AstrMessageEvent): - """处理 /刷新全局今日运势 命令(仅管理员)。""" - if not self._check_admin(event): - yield event.plain_result("只有管理员可以刷新全局运势。") - return - - try: - count = await self._fortune_core.refresh_all_fortune() - yield event.plain_result( - f"已刷新全局今日运势,共 {count} 条记录将被重新生成。" - ) - except Exception as exc: - logger.exception("[fortune] Failed to refresh all fortune: %s", exc) - yield event.plain_result("刷新失败,请稍后重试。") - - async def handle_fortune_config(self, event: AstrMessageEvent, args: str = ""): - """处理 /jrys_config 命令(配置今日运势)。""" - if not self._check_admin(event): - yield event.plain_result("只有管理员可以配置今日运势。") - return - - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - - args = args.strip() - if not args: - # 显示当前配置 - async for result in self._show_fortune_config(event): - yield result - return - - # 解析命令(只将命令部分转为小写) - parts = args.split(maxsplit=1) - cmd = parts[0].lower() - value = parts[1] if len(parts) > 1 else "" - - if cmd == "tags": - # 设置标签(标签保持原样,不转小写) - if value: - success = await self._session_config.set_session_tags( - session_id, is_group, value - ) - if success: - yield event.plain_result(f"✅ 已设置今日运势标签为:{value}") - else: - yield event.plain_result("❌ 设置失败。") - else: - # 清除标签 - success = await self._session_config.clear_session_tags( - session_id, is_group - ) - if success: - yield event.plain_result("✅ 已清除今日运势标签设置。") - else: - yield event.plain_result("ℹ️ 当前没有标签设置。") - - elif cmd == "mode": - # 设置内容模式 - value_lower = value.lower() - if value_lower in ("sfw", "r18", "mix"): - success = await self._session_config.set_session_content_mode( - session_id, is_group, value_lower - ) - if success: - yield event.plain_result( - f"✅ 已设置今日运势内容模式为:{value_lower}" - ) - else: - yield event.plain_result("❌ 设置失败。") - elif value_lower == "clear": - success = await self._session_config.clear_session_content_mode( - session_id, is_group - ) - if success: - yield event.plain_result("✅ 已清除今日运势内容模式设置。") - else: - yield event.plain_result("ℹ️ 当前没有内容模式设置。") - else: - yield event.plain_result( - "❌ 无效的模式。可用:sfw(全年龄)、r18(成人)、mix(混合)、clear(清除)" - ) - else: - yield event.plain_result( - "用法:/jrys_config tags <标签> | /jrys_config mode " - ) - - async def _show_fortune_config(self, event: AstrMessageEvent): - """显示当前今日运势配置。""" - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - - session_tags = await self._session_config.get_session_tags(session_id, is_group) - session_mode = await self._session_config.get_session_content_mode( - session_id, is_group - ) - - fortune_cfg = getattr(self._config, "fortune", {}) - global_tags = ( - fortune_cfg.get("tags", "未设置") - if isinstance(fortune_cfg, dict) - else "未设置" - ) - global_mode = ( - fortune_cfg.get("content_mode", "sfw") - if isinstance(fortune_cfg, dict) - else "sfw" - ) - - effective_tags = session_tags if session_tags is not None else global_tags - effective_mode = session_mode if session_mode else global_mode - - session_type = "群聊" if is_group else "私聊" - - msg = ( - f"📋 今日运势配置({session_type})\n\n" - f"标签设置:\n" - f" 会话覆盖:{session_tags if session_tags is not None else '未设置'}\n" - f" 全局配置:{global_tags}\n" - f" 生效标签:{effective_tags}\n\n" - f"内容模式:\n" - f" 会话覆盖:{session_mode if session_mode else '未设置'}\n" - f" 全局配置:{global_mode}\n" - f" 生效模式:{effective_mode}\n\n" - f"命令:\n" - f" /jrys_config tags <标签> - 设置标签\n" - f" /jrys_config mode - 设置内容模式" - ) - yield event.plain_result(msg) diff --git a/fortune/llm_handlers.py b/fortune/llm_handlers.py deleted file mode 100644 index 59ce38e..0000000 --- a/fortune/llm_handlers.py +++ /dev/null @@ -1,441 +0,0 @@ -"""今日运势 LLM 工具处理器。""" - -from __future__ import annotations - -import base64 -from typing import TYPE_CHECKING - -import astrbot.api.message_components as Comp -from astrbot.api import logger -from astrbot.api.event import AstrMessageEvent -from astrbot.core.message.message_event_result import MessageChain - -if TYPE_CHECKING: - from ..core import SetuCore - from .core import FortuneCore - from .session_config import FortuneSessionConfig - - -class FortuneLlmHandler: - """今日运势 LLM 工具处理器。""" - - def __init__( - self, - setu_plugin, - core: SetuCore, - fortune_core: FortuneCore, - session_config: FortuneSessionConfig, - ): - self._setu_plugin = setu_plugin - self._core = core - self._fortune_core = fortune_core - self._session_config = session_config - self._renderer = None - - @staticmethod - async def _check_admin(event: AstrMessageEvent) -> bool: - """检查用户是否为管理员。""" - try: - if hasattr(event, "is_admin") and callable(getattr(event, "is_admin")): - if event.is_admin(): - return True - if hasattr(event, "is_super_user") and callable( - getattr(event, "is_super_user") - ): - if event.is_super_user(): - return True - except AttributeError: - pass - return False - - @staticmethod - async def _check_super_admin(event: AstrMessageEvent) -> bool: - """检查用户是否为超级管理员。""" - try: - if hasattr(event, "is_super_user") and callable( - getattr(event, "is_super_user") - ): - if event.is_super_user(): - return True - except AttributeError: - pass - return False - - async def _get_image_for_fortune(self, event: AstrMessageEvent) -> bytes | None: - """获取今日运势的背景图片。 - - 使用 Setu 插件的图片供应商,根据配置获取图片。 - """ - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - - # 获取生效的配置 - - config = self._setu_plugin.config - fortune_cfg = getattr(config, "fortune", {}) - global_tags = ( - fortune_cfg.get("tags", "") if isinstance(fortune_cfg, dict) else "" - ) - global_mode = ( - fortune_cfg.get("content_mode", "sfw") - if isinstance(fortune_cfg, dict) - else "sfw" - ) - - tags_str = await self._session_config.get_effective_tags( - session_id, is_group, global_tags - ) - content_mode = await self._session_config.get_effective_content_mode( - session_id, is_group, global_mode - ) - - # 解析标签 - tags = [t.strip() for t in tags_str.split(",") if t.strip()] if tags_str else [] - - # 确定 R18 - import random - - is_r18 = False - if content_mode == "r18": - is_r18 = True - elif content_mode == "mix": - is_r18 = random.random() > 0.5 - - logger.debug( - "[fortune_llm] Getting image: tags=%s, mode=%s, is_r18=%s", - tags, - content_mode, - is_r18, - ) - - try: - # 使用 SetuCore 的 provider 获取图片 - provider = self._core._get_provider() - if not provider: - logger.error("[fortune_llm] No provider available") - return None - - # 获取图片 - downloaded = await self._core.fetch_and_download_images(1, tags, is_r18) - - if downloaded and len(downloaded) > 0: - return downloaded[0] - - except Exception as exc: - logger.exception("[fortune_llm] Failed to get image: %s", exc) - - return None - - async def _send_fortune_image(self, event: AstrMessageEvent, fortune: dict) -> bool: - """生成并发送运势图片给用户。 - - 参数: - event: 消息事件 - fortune: 运势数据 - - 返回: - 是否成功发送图片 - """ - try: - # 初始化 renderer - if self._renderer is None: - from .renderer import FortuneRenderer - - self._renderer = FortuneRenderer() - - # 获取 HTML 渲染器 - html_renderer = None - if hasattr(self._setu_plugin, "html_render"): - html_renderer = self._setu_plugin.html_render - - if not html_renderer: - logger.warning("[fortune_llm] No HTML renderer available") - return False - - # 获取背景图片 - image_data = await self._get_image_for_fortune(event) - image_base64 = "" - if image_data: - image_base64 = base64.b64encode(image_data).decode("utf-8") - - # 渲染运势图片 - rendered_image = await self._renderer.render_to_image( - fortune=fortune, image_base64=image_base64, html_renderer=html_renderer - ) - - if not rendered_image: - logger.warning("[fortune_llm] Failed to render fortune image") - return False - - # 保存到缓存 - user_id = event.get_sender_id() - await self._fortune_core.update_fortune_image_cache( - user_id, fortune["date_str"], rendered_image - ) - - # 发送图片 - message_chain = MessageChain([Comp.Image.fromBytes(rendered_image)]) - await self._setu_plugin.context.send_message( - event.unified_msg_origin, - message_chain, - ) - return True - - except Exception as exc: - logger.exception("[fortune_llm] Failed to send fortune image: %s", exc) - return False - - async def llm_get_fortune(self, event: AstrMessageEvent, **kwargs) -> dict: - """LLM 工具:获取今日运势。 - - 返回用户的今日运势信息,并发送运势图片给用户。 - """ - user_id = event.get_sender_id() - username = event.get_sender_name() or user_id - group_id = event.get_group_id() - - try: - fortune = await self._fortune_core.get_today_fortune( - user_id, username, group_id - ) - - if not fortune: - return { - "success": False, - "message": "无法获取今日运势,请稍后重试。", - } - - # 检查是否已有缓存的渲染图片 - cached_image = await self._fortune_core.get_cached_image( - user_id, fortune["date_str"] - ) - - if cached_image: - # 使用缓存的图片 - logger.debug("[fortune_llm] Using cached image for %s", user_id) - message_chain = MessageChain([Comp.Image.fromBytes(cached_image)]) - await self._setu_plugin.context.send_message( - event.unified_msg_origin, - message_chain, - ) - else: - # 没有缓存,生成新图片 - logger.debug( - "[fortune_llm] No cached image, generating new one for %s", user_id - ) - await self._send_fortune_image(event, fortune) - - return { - "success": True, - "fortune": { - "title": fortune["title"], - "stars": fortune["star_count"], - "max_stars": fortune["max_stars"], - "description": fortune["description"], - "extra": fortune["extra_message"], - "date": fortune["date_str"], - }, - "message": f"今日运势:{fortune['title']}({fortune['star_count']}星)\n{fortune['description']}", - } - except Exception as exc: - logger.exception("[fortune_llm] Failed to get fortune: %s", exc) - return { - "success": False, - "message": "获取运势时出错,请稍后重试。", - } - - async def llm_refresh_fortune(self, event: AstrMessageEvent, **kwargs) -> dict: - """LLM 工具:刷新今日运势。 - - 刷新用户的今日运势(管理员或配置允许时)。 - """ - user_id = event.get_sender_id() - username = event.get_sender_name() or user_id - - # 检查权限 - is_admin = await self._check_admin(event) - if not is_admin: - return { - "success": False, - "message": "只有管理员可以刷新运势。", - } - - try: - await self._fortune_core.refresh_fortune(user_id, username) - return { - "success": True, - "message": "已成功刷新今日运势,可以重新获取查看。", - } - except Exception as exc: - logger.exception("[fortune_llm] Failed to refresh fortune: %s", exc) - return { - "success": False, - "message": "刷新运势时出错,请稍后重试。", - } - - async def llm_refresh_group_fortune( - self, event: AstrMessageEvent, **kwargs - ) -> dict: - """LLM 工具:刷新群组今日运势。 - - 刷新当前群组所有成员的今日运势(仅管理员)。 - """ - # 检查权限 - is_admin = await self._check_admin(event) - if not is_admin: - return { - "success": False, - "message": "只有管理员可以刷新群组运势。", - } - - group_id = event.get_group_id() - if not group_id: - return { - "success": False, - "message": "此操作只能在群聊中进行。", - } - - try: - count = await self._fortune_core.refresh_group_fortune(str(group_id)) - return { - "success": True, - "message": f"已成功刷新本群今日运势,共 {count} 条记录将被重新生成。", - } - except Exception as exc: - logger.exception("[fortune_llm] Failed to refresh group fortune: %s", exc) - return { - "success": False, - "message": "刷新群组运势时出错,请稍后重试。", - } - - async def llm_refresh_all_fortune(self, event: AstrMessageEvent, **kwargs) -> dict: - """LLM 工具:刷新全局今日运势。 - - 刷新所有用户的今日运势(仅超级管理员)。 - """ - # 检查权限:仅允许超级管理员 - is_super_admin = await self._check_super_admin(event) - if not is_super_admin: - return { - "success": False, - "message": "只有超级管理员可以刷新全局运势。", - } - - try: - count = await self._fortune_core.refresh_all_fortune() - return { - "success": True, - "message": f"已成功刷新全局今日运势,共 {count} 条记录将被重新生成。", - } - except Exception as exc: - logger.exception("[fortune_llm] Failed to refresh all fortune: %s", exc) - return { - "success": False, - "message": "刷新全局运势时出错,请稍后重试。", - } - - async def llm_get_fortune_config(self, event: AstrMessageEvent, **kwargs) -> dict: - """LLM 工具:获取今日运势配置。 - - 返回当前会话的今日运势配置信息。 - """ - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - - session_tags = await self._session_config.get_session_tags(session_id, is_group) - session_mode = await self._session_config.get_session_content_mode( - session_id, is_group - ) - - config = self._setu_plugin.config - fortune_cfg = getattr(config, "fortune", {}) - global_tags = ( - fortune_cfg.get("tags", "未设置") - if isinstance(fortune_cfg, dict) - else "未设置" - ) - global_mode = ( - fortune_cfg.get("content_mode", "sfw") - if isinstance(fortune_cfg, dict) - else "sfw" - ) - - effective_tags = session_tags if session_tags is not None else global_tags - effective_mode = session_mode if session_mode else global_mode - - return { - "success": True, - "config": { - "session_tags": session_tags, - "global_tags": global_tags, - "effective_tags": effective_tags, - "session_mode": session_mode, - "global_mode": global_mode, - "effective_mode": effective_mode, - }, - "message": ( - f"当前今日运势配置:\n" - f"标签:{effective_tags}\n" - f"内容模式:{effective_mode}" - ), - } - - async def llm_set_fortune_config( - self, - event: AstrMessageEvent, - tags: str | None = None, - mode: str | None = None, - **kwargs, - ) -> dict: - """LLM 工具:设置今日运势配置。 - - 设置当前会话的今日运势标签或内容模式(仅管理员)。 - - 参数: - tags: 标签字符串,如 "少女,可爱" - mode: 内容模式,可选 sfw、r18、mix - """ - # 检查权限 - is_admin = await self._check_admin(event) - if not is_admin: - return { - "success": False, - "message": "只有管理员可以配置今日运势。", - } - - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - - results = [] - - if tags is not None: - if tags == "": - # 清除标签 - await self._session_config.clear_session_tags(session_id, is_group) - results.append("已清除标签设置") - else: - await self._session_config.set_session_tags(session_id, is_group, tags) - results.append(f"已设置标签为:{tags}") - - if mode is not None: - if mode in ("sfw", "r18", "mix"): - await self._session_config.set_session_content_mode( - session_id, is_group, mode - ) - results.append(f"已设置内容模式为:{mode}") - else: - return { - "success": False, - "message": f"无效的内容模式:{mode},可选:sfw、r18、mix", - } - - if not results: - return { - "success": False, - "message": "请提供要设置的参数:tags 或 mode", - } - - return { - "success": True, - "message": ";".join(results), - } diff --git a/fortune/renderer.py b/fortune/renderer.py deleted file mode 100644 index 2497ab5..0000000 --- a/fortune/renderer.py +++ /dev/null @@ -1,265 +0,0 @@ -"""今日运势 HTML 渲染模块。""" - -from __future__ import annotations - -import base64 -from pathlib import Path -from typing import Any - -from astrbot.api import logger - - -class FortuneRenderer: - """运势卡片渲染器。""" - - def __init__(self, template_path: Path | None = None): - if template_path is None: - template_path = Path(__file__).parent.parent / "templates" / "fortune.html" - self.template_path = template_path - self._fonts_dir = self.template_path.parent / "res" / "fonts" - self._embedded_fonts_css: str | None = None - - def _get_embedded_fonts_css(self) -> str: - """生成内嵌字体的 CSS(使用 base64)。 - - 字体文件会被缓存,避免重复读取。 - """ - if self._embedded_fonts_css is not None: - return self._embedded_fonts_css - - fonts_config = [ - { - "name": "NotoSansSC-Regular", - "file": "NotoSansSC-Regular.woff2", - "format": "woff2", - "mime": "font/woff2", - }, - { - "name": "NotoSansSC-Bold", - "file": "NotoSansSC-Bold.woff2", - "format": "woff2", - "mime": "font/woff2", - }, - { - "name": "SSFangTangTi", - "file": "SSFangTangTi.woff2", - "format": "woff2", - "mime": "font/woff2", - }, - ] - - css_parts = [" /* Embedded fonts - auto-generated */"] - - for font in fonts_config: - font_path = self._fonts_dir / font["file"] - try: - if font_path.exists(): - font_data = font_path.read_bytes() - b64_data = base64.b64encode(font_data).decode("ascii") - css_parts.append(f""" @font-face {{ - font-family: '{font["name"]}'; - src: url('data:{font["mime"]};base64,{b64_data}') format('{font["format"]}'); - font-weight: normal; - font-style: normal; - font-display: swap; - }}""") - else: - logger.warning("[fortune] Font file not found: %s", font_path) - except OSError as exc: - logger.error("[fortune] Failed to read font %s: %s", font["file"], exc) - - self._embedded_fonts_css = "\n".join(css_parts) - return self._embedded_fonts_css - - @staticmethod - def _truncate_username(username: str, max_length: int = 15) -> str: - """截断用户名,超过最大长度时用...省略。 - - 参数: - username: 原始用户名 - max_length: 最大字符数,默认15 - - 返回: - 截断后的用户名 - """ - if len(username) > max_length: - return username[:max_length] + "..." - return username - - def _get_template(self) -> str: - try: - return self.template_path.read_text(encoding="utf-8") - except OSError as exc: - logger.error("[fortune] Failed to read template: %s", exc) - # 返回一个极简的备用模板 - return """ - -
-

{{ username }} 的今日运势 ({{ date_str }})

-

运势: {{ title }}

-

星级: {{ stars_display }}

-

{{ description }}

-
""" - - def render(self, fortune: dict[str, Any], image_base64: str | None = None) -> str: - """渲染运势卡片为 HTML 字符串。 - - 参数: - fortune: 运势数据 - image_base64: 背景图片的 base64 编码(可选) - - 返回: - HTML 字符串(包含内嵌字体,适用于 Playwright 截图) - """ - template = self._get_template() - - # 构建渲染数据(包含内嵌字体 CSS) - username = self._truncate_username(fortune.get("username", "用户")) - data = { - "fonts_css": self._get_embedded_fonts_css(), - "username": username, - "date_str": fortune.get("date_str", ""), - "title": fortune.get("title", "未知"), - "stars_display": self._format_stars( - fortune.get("star_count", 3), fortune.get("max_stars", 7) - ), - "description": fortune.get("description", ""), - "extra_message": fortune.get("extra_message", ""), - "theme_color": fortune.get("theme_color", "theme-gray"), - "image_base64": image_base64 or "", - } - - # 简单的模板渲染(替换 {{ variable }}) - html = template - for key, value in data.items(): - placeholder = f"{{{{ {key} }}}}" - html = html.replace(placeholder, str(value)) - - # 处理条件语句 {% if variable %}...{% endif %} - import re - - # 处理 if 语句 - def process_if(match: re.Match) -> str: - condition = match.group(1).strip() - content = match.group(2) - # 检查条件是否为真(变量存在且非空) - if condition in data and data[condition]: - return content - return "" - - html = re.sub( - r"{%\s*if\s+(\w+)\s*%}(.*?){%\s*endif\s*%}", - process_if, - html, - flags=re.DOTALL, - ) - - return html - - def _format_stars(self, count: int, max_count: int = 7) -> str: - """格式化星级显示为HTML。""" - stars_html = [] - # 填充的星星 - for _ in range(count): - stars_html.append('') - # 空星星 - for _ in range(max_count - count): - stars_html.append('') - return "".join(stars_html) - - async def render_to_image( - self, - fortune: dict[str, Any], - image_base64: str | None = None, - html_renderer=None, - ) -> bytes | None: - """渲染运势卡片为图片。 - - 使用 AstrBot 的 T2I 服务(Jinja2 模板引擎)。 - - 参数: - fortune: 运势数据 - image_base64: 背景图片 base64 - html_renderer: HTML 渲染器(AstrBot context 的 html_render 方法) - - 返回: - 图片字节数据,失败返回 None - """ - if html_renderer is None: - logger.error("[fortune] No HTML renderer available") - return None - - try: - # 读取 HTML 模板(使用 Jinja2 语法) - template = self._get_template() - - # 准备模板数据(通过 data 参数传递给 Jinja2,包含内嵌字体 CSS) - username = self._truncate_username(fortune.get("username", "用户")) - tmpl_data = { - "fonts_css": self._get_embedded_fonts_css(), - "username": username, - "date_str": fortune.get("date_str", ""), - "title": fortune.get("title", "未知"), - "stars_display": self._format_stars( - fortune.get("star_count", 3), fortune.get("max_stars", 7) - ), - "description": fortune.get("description", ""), - "extra_message": fortune.get("extra_message", ""), - "theme_color": fortune.get("theme_color", "theme-gray"), - "image_base64": image_base64 or "", - } - - # 使用 AstrBot 的 html_render 方法 - # 参考 Playwright screenshot API: https://playwright.dev/python/docs/api/class-page#page-screenshot - render_options = { - "full_page": True, - "type": "png", - "scale": "device", - } - result = await html_renderer( - tmpl=template, data=tmpl_data, return_url=False, options=render_options - ) - - # 处理返回值(可能是文件路径字符串) - if result is None: - logger.error("[fortune] html_render returned None") - return None - - if isinstance(result, str): - # 结果是文件路径 - import asyncio - from pathlib import Path - - path = Path(result) - if path.exists(): - # 检查文件大小 - file_size = path.stat().st_size - logger.debug( - "[fortune] Reading image from path: %s (%d bytes)", - result, - file_size, - ) - if file_size < 100: - logger.error( - "[fortune] Rendered image too small (%d bytes), likely failed", - file_size, - ) - return None - return await asyncio.to_thread(path.read_bytes) - else: - logger.error("[fortune] Rendered file not found: %s", result) - return None - - logger.error( - "[fortune] Unexpected return type from html_render: %s", type(result) - ) - return None - - except NotImplementedError: - logger.error( - "[fortune] HTML rendering not supported by current T2I strategy (local strategy doesn't support custom templates)" - ) - return None - except Exception as exc: - logger.exception("[fortune] Failed to render image: %s", exc) - return None diff --git a/fortune/session_config.py b/fortune/session_config.py deleted file mode 100644 index 28dd530..0000000 --- a/fortune/session_config.py +++ /dev/null @@ -1,237 +0,0 @@ -"""今日运势会话配置管理。 - -支持会话级别的独立配置,包括标签和内容模式。 -配置存储在 AstrBotConfig 的 fortune_session_configs 字段中,可在 WebUI 中管理。 -""" - -from __future__ import annotations - -import asyncio -import json -from pathlib import Path -from typing import TYPE_CHECKING, Any - -from astrbot.api import logger - -from ..session_config_base import SessionConfigBase - -if TYPE_CHECKING: - from astrbot.core import AstrBotConfig - - -class FortuneSessionConfig(SessionConfigBase): - """今日运势会话配置管理器。 - - 配置存储在 AstrBotConfig 的 fortune_session_configs 字段中,可在 WebUI 中管理。 - """ - - # 允许会话覆盖的配置项 - ALLOWED_KEYS = {"tags", "content_mode"} - - # 有效的内容模式值 - VALID_CONTENT_MODES = {"sfw", "r18", "mix"} - - def __init__(self, config: AstrBotConfig, data_dir: Path | None = None): - """初始化配置管理器。 - - 参数: - config: AstrBotConfig 配置对象 - data_dir: 插件数据目录(用于迁移旧配置) - """ - super().__init__(config, config_key="fortune_session_configs") - self._data_dir = data_dir - self._migrated = False - - async def initialize(self) -> None: - """初始化配置,迁移旧配置文件到新格式。""" - if self._data_dir and not self._migrated: - await self._migrate_old_config() - self._migrated = True - logger.info("[fortune_session] Session config initialized") - - async def _migrate_old_config(self) -> None: - """迁移旧位置的配置文件到 AstrBotConfig。""" - old_config_file = self._data_dir / "fortune_session_config.json" - - if not old_config_file.exists(): - return - - try: - content = await asyncio.to_thread( - old_config_file.read_text, encoding="utf-8" - ) - old_data = json.loads(content) - sessions = old_data.get("sessions", {}) - - if not sessions: - return - - # 转换格式到新的 template_list 格式 - new_configs = self._get_configs() - - for session_key, session_data in sessions.items(): - # 解析 session_key (格式: "group:12345" 或 "private:67890") - parsed = self._parse_session_key(session_key) - if not parsed: - continue - - session_type, session_id = parsed - - # 构建新的配置项 - new_item = { - "session_id": session_id, - "session_type": session_type, - } - - # 迁移各配置项 - if "tags" in session_data: - new_item["tags"] = session_data["tags"] - if "content_mode" in session_data: - new_item["content_mode"] = session_data["content_mode"] - - # 检查是否已存在 - self.merge_session_item(new_configs, new_item) - - # 保存迁移后的配置 - self._save_configs(new_configs) - - # 删除旧配置文件 - await asyncio.to_thread(old_config_file.unlink) - logger.info( - "[fortune_session] Migrated old config to AstrBotConfig, %d sessions", - len(new_configs), - ) - except (OSError, json.JSONDecodeError) as exc: - logger.warning("[fortune_session] Failed to migrate old config: %s", exc) - - def _get_session_config(self, session_id: str, is_group: bool) -> dict | None: - """获取会话配置项。""" - return self._find_session_config(session_id, is_group) - - def _get_session_key(self, session_id: str, is_group: bool) -> str: - """生成会话键。""" - prefix = "group" if is_group else "private" - return f"{prefix}:{session_id}" - - async def set_config( - self, session_id: str, is_group: bool, key: str, value: Any - ) -> bool: - """设置会话配置项。""" - if key not in self.ALLOWED_KEYS: - logger.warning("[fortune_session] Key %s is not allowed", key) - return False - - if key == "content_mode" and value not in self.VALID_CONTENT_MODES: - logger.warning("[fortune_session] Invalid content_mode: %s", value) - return False - - async with self._lock: - session_type = "group" if is_group else "private" - configs = self._get_configs() - - # 查找现有配置 - existing_idx = self._find_session_index(configs, session_id, is_group) - - if existing_idx is not None: - # 更新现有配置 - configs[existing_idx][key] = value - else: - # 创建新配置 - new_item = { - "session_id": session_id, - "session_type": session_type, - key: value, - } - configs.append(new_item) - - self._save_configs(configs) - - session_key = self._get_session_key(session_id, is_group) - logger.info("[fortune_session] Set %s=%s for %s", key, value, session_key) - return True - - async def get_config( - self, session_id: str, is_group: bool, key: str, default: Any = None - ) -> Any: - """获取会话配置项。""" - async with self._lock: - cfg = self._get_session_config(session_id, is_group) - if not cfg: - return default - - value = cfg.get(key) - if value is None or value == "": - return default - - return value - - async def clear_config(self, session_id: str, is_group: bool, key: str) -> bool: - """清除会话配置项。""" - async with self._lock: - configs = self._get_configs() - - idx = self._find_session_index(configs, session_id, is_group) - if idx is None: - return False - - cfg = configs[idx] - if key not in cfg: - return False - - del cfg[key] - # 如果配置项都为空,删除整个会话配置 - remaining_keys = [k for k in self.ALLOWED_KEYS if k in cfg] - if not remaining_keys: - configs.pop(idx) - self._save_configs(configs) - session_key = self._get_session_key(session_id, is_group) - logger.info("[fortune_session] Cleared %s for %s", key, session_key) - return True - - async def get_session_tags(self, session_id: str, is_group: bool) -> str | None: - """获取会话的标签配置。""" - return await self.get_config(session_id, is_group, "tags") - - async def set_session_tags( - self, session_id: str, is_group: bool, tags: str - ) -> bool: - """设置会话的标签配置。""" - return await self.set_config(session_id, is_group, "tags", tags) - - async def clear_session_tags(self, session_id: str, is_group: bool) -> bool: - """清除会话的标签配置。""" - return await self.clear_config(session_id, is_group, "tags") - - async def get_session_content_mode( - self, session_id: str, is_group: bool - ) -> str | None: - """获取会话的内容模式配置。""" - return await self.get_config(session_id, is_group, "content_mode") - - async def set_session_content_mode( - self, session_id: str, is_group: bool, mode: str - ) -> bool: - """设置会话的内容模式配置。""" - return await self.set_config(session_id, is_group, "content_mode", mode) - - async def clear_session_content_mode(self, session_id: str, is_group: bool) -> bool: - """清除会话的内容模式配置。""" - return await self.clear_config(session_id, is_group, "content_mode") - - async def get_effective_tags( - self, session_id: str, is_group: bool, global_tags: str - ) -> str: - """获取生效的标签(优先会话配置)。""" - session_tags = await self.get_session_tags(session_id, is_group) - if session_tags is not None: - return session_tags - return global_tags - - async def get_effective_content_mode( - self, session_id: str, is_group: bool, global_mode: str - ) -> str: - """获取生效的内容模式(优先会话配置)。""" - session_mode = await self.get_session_content_mode(session_id, is_group) - if session_mode: - return session_mode - return global_mode diff --git a/handlers/__init__.py b/handlers/__init__.py deleted file mode 100644 index 2e0c164..0000000 --- a/handlers/__init__.py +++ /dev/null @@ -1,8 +0,0 @@ -"""Handlers package.""" - -from __future__ import annotations - -from .command_handlers import CommandHandler -from .llm_handlers import LlmHandlers - -__all__ = ["CommandHandler", "LlmHandlers"] diff --git a/handlers/command_handlers.py b/handlers/command_handlers.py deleted file mode 100644 index 4dc20c6..0000000 --- a/handlers/command_handlers.py +++ /dev/null @@ -1,838 +0,0 @@ -"""命令处理器。""" - -from __future__ import annotations - -import asyncio -import re -from typing import TYPE_CHECKING - -from astrbot.api import logger -from astrbot.api.event import AstrMessageEvent - -from ..config import SetuConfig, parse_count -from ..constants import COMMAND_PATTERN - -if TYPE_CHECKING: - from ..core import SetuCore - - -class CommandHandler: - """命令处理器。""" - - def __init__(self, core: SetuCore, config: SetuConfig): - self._core = core - self._config = config - - def _check_admin(self, event: AstrMessageEvent) -> bool: - """检查用户是否为管理员。""" - try: - if hasattr(event, "is_admin") and callable(getattr(event, "is_admin")): - if event.is_admin(): - return True - if hasattr(event, "is_super_user") and callable( - getattr(event, "is_super_user") - ): - if event.is_super_user(): - return True - if hasattr(event, "message_obj"): - msg_obj = event.message_obj - if hasattr(msg_obj, "sender") and hasattr(msg_obj.sender, "role"): - role = msg_obj.sender.role - if role in ("admin", "owner"): - return True - except AttributeError: - pass - return False - - async def handle_random_picture(self, event: AstrMessageEvent): - """处理色图请求命令。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - # 限流检查:每个用户同时只能有一个请求 - if not await self._core.rate_limiter.acquire(event): - yield event.plain_result("你有一个请求正在处理中,请稍后再试~") - return - - try: - async for result in self._handle_random_picture_internal(event): - yield result - finally: - # 无论成功或失败,都释放锁 - await self._core.rate_limiter.release(event) - - async def _handle_random_picture_internal(self, event: AstrMessageEvent): - """处理色图请求命令的内部逻辑。""" - match = re.match(COMMAND_PATTERN, event.message_str.strip()) - if not match: - return - - # 检查全局黑白名单和色图功能级黑名单 - if self._core.is_group_blocked(event, feature="setu"): - return - - num_str = match.group(2) - num = parse_count(num_str) - max_count = self._config.max_count - - if num < 1 or num > max_count: - if num == -1: - yield event.plain_result( - f"数量解析失败,请使用数字或中文数字,图片数量必须在1-{max_count}之间" - ) - elif num > max_count: - yield event.plain_result(f"一次最多只能获取{max_count}张哦~") - else: - yield event.plain_result(f"图片数量必须在1-{max_count}之间哦~") - return - - tag_str = match.group(4).strip() - tags = self._config.resolve_tags(tag_str) - - if self._config.msg_fetching_enabled: - yield event.plain_result(self._config.msg_fetching_text) - - try: - effective_content_mode = await self._core.get_effective_content_mode(event) - is_r18 = self._core.determine_r18(effective_content_mode) - logger.debug("num = %d invoke params = %s , r18 = %s", num, tags, is_r18) - - # 检查是否使用 URL 发送模式 - if self._config.url_send_mode: - # URL 模式:获取 URL 并直接发送 - provider = self._core._get_provider() - if provider: - try: - img_urls = await asyncio.wait_for( - provider.fetch_image_urls( - num=num, - tags=tags, - r18=is_r18, - exclude_ai=self._config.exclude_ai, - ), - timeout=60.0, - ) - except asyncio.TimeoutError: - logger.warning( - "url mode fetch_image_urls timeout (>60s), tags=%s, r18=%s", - tags, - is_r18, - ) - yield event.plain_result("图片获取超时,请稍后重试。") - return - - async for result in self._core.send_images_by_url( - event, img_urls, is_r18, tags - ): - yield result - else: - yield event.plain_result("没有可用的图片源,请联系管理员。") - else: - # 正常模式:下载后发送 - downloaded = await asyncio.wait_for( - self._core.fetch_and_download_images(num, tags, is_r18), - timeout=60.0, - ) - if not downloaded: - tags_info = f"标签: {', '.join(tags)}" if tags else "" - yield event.plain_result( - f"未找到{tags_info}符合要求的图片,请尝试其他标签或检查标签拼写~" - ) - return - - async for result in self._core.send_images( - event, downloaded, is_r18, tags - ): - yield result - except asyncio.TimeoutError: - logger.warning("get_random_picture timeout (>60s)") - yield event.plain_result("获取图片超时,网络可能不稳定,请稍后再试。") - except (OSError, RuntimeError, ValueError): - logger.exception("get_random_picture failed") - yield event.plain_result("获取图片失败,网络或服务异常,请稍后再试。") - - async def handle_setu_command( - self, event: AstrMessageEvent, count: str = "1", *, tags: str = "" - ): - """处理 /setu 命令。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - # 限流检查:每个用户同时只能有一个请求 - if not await self._core.rate_limiter.acquire(event): - yield event.plain_result("你有一个请求正在处理中,请稍后再试~") - return - - try: - async for result in self._handle_setu_command_internal( - event, count, tags=tags - ): - yield result - finally: - await self._core.rate_limiter.release(event) - - async def _handle_setu_command_internal( - self, event: AstrMessageEvent, count: str = "1", *, tags: str = "" - ): - """处理 /setu 命令的内部逻辑。""" - if self._core.is_group_blocked(event, feature="setu"): - return - - max_count = self._config.max_count - num_str = count - num = parse_count(num_str) - - extra_tag = "" - if num == -1: - num = 1 - extra_tag = count - - all_tags = tags - if extra_tag: - all_tags = f"{extra_tag} {all_tags}".strip() - - if num > max_count: - yield event.plain_result(f"一次最多只能获取{max_count}张哦~") - return - - parsed_tags = self._config.resolve_tags(all_tags) - - if self._config.msg_fetching_enabled: - yield event.plain_result(self._config.msg_fetching_text) - - try: - effective_content_mode = await self._core.get_effective_content_mode(event) - is_r18 = self._core.determine_r18(effective_content_mode) - - # 检查是否使用 URL 发送模式 - if self._config.url_send_mode: - provider = self._core._get_provider() - if provider: - try: - img_urls = await asyncio.wait_for( - provider.fetch_image_urls( - num=num, - tags=parsed_tags, - r18=is_r18, - exclude_ai=self._config.exclude_ai, - ), - timeout=60.0, - ) - except asyncio.TimeoutError: - logger.warning( - "url mode fetch_image_urls timeout (>60s), tags=%s, r18=%s", - parsed_tags, - is_r18, - ) - yield event.plain_result("图片获取超时,请稍后重试。") - return - - async for result in self._core.send_images_by_url( - event, img_urls, is_r18, parsed_tags - ): - yield result - else: - yield event.plain_result("没有可用的图片源,请联系管理员。") - else: - downloaded = await asyncio.wait_for( - self._core.fetch_and_download_images(num, parsed_tags, is_r18), - timeout=60.0, - ) - if not downloaded: - tags_info = f"标签: {', '.join(parsed_tags)}" if parsed_tags else "" - yield event.plain_result( - f"未找到{tags_info}符合要求的图片,请尝试其他标签或检查标签拼写~" - ) - return - - async for result in self._core.send_images( - event, downloaded, is_r18, parsed_tags - ): - yield result - except asyncio.TimeoutError: - logger.warning("setu command timeout (>60s)") - yield event.plain_result("获取图片超时,网络可能不稳定,请稍后再试。") - except (OSError, RuntimeError, ValueError): - logger.exception("setu command failed") - yield event.plain_result("获取图片失败,网络或服务异常,请稍后再试。") - - async def _show_mode_status(self, event: AstrMessageEvent): - """显示当前会话的模式状态。""" - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - - current_mode = await self._core.session_config.get_session_content_mode( - session_id, is_group - ) - global_mode = self._core.config.content_mode - - session_docx = await self._core.session_config.get_session_r18_docx_mode( - session_id, is_group - ) - global_docx = self._core.config.r18_docx_mode - effective_docx = await self._core.get_effective_r18_docx_mode(event) - - session_revoke = await self._core.session_config.get_session_auto_revoke_r18( - session_id, is_group - ) - global_revoke = self._core.config.auto_revoke_r18 - delay = self._core.config.auto_revoke_delay - effective_revoke = await self._core.get_effective_auto_revoke_r18(event) - - session_send = await self._core.session_config.get_session_send_mode( - session_id, is_group - ) - global_send = self._core.config.send_mode - effective_send = await self._core.get_effective_send_mode(event) - - def fmt(val, is_bool=True): - if val is None: - return "未设置" - return "启用" if val else "禁用" if is_bool else str(val) - - msg = ( - f"📋 当前会话配置:\n\n" - f"1️⃣ 内容分级:\n" - f" 会话覆盖:{fmt(current_mode, False) if current_mode else '未设置'}\n" - f" 全局配置:{global_mode}\n" - f" 生效模式:{current_mode or global_mode}\n\n" - f"2️⃣ R18 Docx 打包:\n" - f" 会话覆盖:{fmt(session_docx)}\n" - f" 全局配置:{'启用' if global_docx else '禁用'}\n" - f" 生效设置:{'启用' if effective_docx else '禁用'}\n\n" - f"3️⃣ R18 自动撤回:\n" - f" 会话覆盖:{fmt(session_revoke)}\n" - f" 全局配置:{'启用' if global_revoke else '禁用'}(延迟 {delay} 秒)\n" - f" 生效设置:{'启用' if effective_revoke else '禁用'}\n\n" - f"4️⃣ 发送模式:\n" - f" 会话覆盖:{fmt(session_send, False) if session_send else '未设置'}\n" - f" 全局配置:{global_send}\n" - f" 生效模式:{effective_send}\n\n" - f"可用命令:\n" - f" /setu_config mode - 设置内容分级\n" - f" /setu_config docx - 设置 R18 Docx 模式\n" - f" /setu_config revoke - 设置自动撤回\n" - f" /setu_config send - 设置发送模式\n" - f" /setu_config show - 显示当前配置" - ) - yield event.plain_result(msg) - - async def handle_setu_config(self, event: AstrMessageEvent, args: str = ""): - """处理 /setu_config 命令(统一配置管理)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - args = args.strip() - if not args or args.lower() == "show": - async for result in self._show_mode_status(event): - yield result - return - - # 解析命令(只将命令部分转为小写) - parts = args.split(maxsplit=1) - cmd = parts[0].lower() - value = parts[1].lower() if len(parts) > 1 else "" - - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - session_type = "群聊" if is_group else "私聊" - - if cmd == "mode": - # 内容分级设置 - async for result in self._handle_config_mode( - event, session_id, is_group, session_type, value - ): - yield result - - elif cmd == "docx": - # R18 Docx 打包模式设置 - async for result in self._handle_config_docx( - event, session_id, is_group, session_type, value - ): - yield result - - elif cmd == "revoke": - # 自动撤回设置 - async for result in self._handle_config_revoke( - event, session_id, is_group, session_type, value - ): - yield result - - elif cmd == "send": - # 发送模式设置 - async for result in self._handle_config_send( - event, session_id, is_group, session_type, value - ): - yield result - - else: - yield event.plain_result( - "❌ 未知的配置项。\n" - "可用配置项:mode(内容分级)、docx(打包模式)、revoke(自动撤回)、send(发送模式)\n" - "使用 /setu_config show 查看当前配置。" - ) - - async def _handle_config_mode( - self, event, session_id, is_group, session_type, value - ): - """处理内容分级配置。""" - if value not in ("sfw", "r18", "mix", "clear"): - yield event.plain_result( - "❌ 无效的模式。\n" - "可用模式:sfw(全年龄)、r18(成人)、mix(混合)、clear(清除设置)" - ) - return - - if value == "clear": - success = await self._core.session_config.clear_session_content_mode( - session_id, is_group - ) - if success: - global_mode = self._core.config.content_mode - yield event.plain_result( - f"✅ 已清除当前{session_type}的内容分级设置,将使用全局配置。\n" - f"当前全局配置为:{global_mode}" - ) - else: - yield event.plain_result("ℹ️ 当前会话没有设置覆盖,已在使用全局配置。") - else: - success = await self._core.session_config.set_session_content_mode( - session_id, is_group, value - ) - if success: - yield event.plain_result( - f"✅ 已将当前{session_type}的内容分级设置为:{value}\n" - f"此后发送的图片将使用此模式(优先于全局配置)。" - ) - else: - yield event.plain_result("❌ 设置失败,请稍后再试。") - - async def _handle_config_docx( - self, event, session_id, is_group, session_type, value - ): - """处理 R18 Docx 打包模式配置。""" - if value not in ("on", "off", "clear", ""): - yield event.plain_result( - "❌ 无效的设置。\n可用值:on(启用)、off(禁用)、clear(清除设置)" - ) - return - - global_docx = self._core.config.r18_docx_mode - - if value == "clear" or value == "": - success = await self._core.session_config.clear_session_r18_docx_mode( - session_id, is_group - ) - if success: - yield event.plain_result( - f"✅ 已清除当前{session_type}的 R18 Docx 打包模式设置,将使用全局配置。\n" - f"当前全局配置为:{'启用' if global_docx else '禁用'}" - ) - else: - yield event.plain_result("ℹ️ 当前会话没有设置覆盖,已在使用全局配置。") - else: - enabled = value == "on" - success = await self._core.session_config.set_session_r18_docx_mode( - session_id, is_group, enabled - ) - if success: - yield event.plain_result( - f"✅ 已将当前{session_type}的 R18 Docx 打包模式设置为:{'启用' if enabled else '禁用'}\n" - f"此后发送的 R18 图片将{'打包为 DOCX 文件' if enabled else '直接发送'}(优先于全局配置)。" - ) - else: - yield event.plain_result("❌ 设置失败,请稍后再试。") - - async def _handle_config_revoke( - self, event, session_id, is_group, session_type, value - ): - """处理自动撤回配置。""" - if value not in ("on", "off", "clear", ""): - yield event.plain_result( - "❌ 无效的设置。\n可用值:on(启用)、off(禁用)、clear(清除设置)" - ) - return - - global_revoke = self._core.config.auto_revoke_r18 - delay = self._core.config.auto_revoke_delay - - if value == "clear" or value == "": - success = await self._core.session_config.clear_session_auto_revoke_r18( - session_id, is_group - ) - if success: - yield event.plain_result( - f"✅ 已清除当前{session_type}的自动撤回设置,将使用全局配置。\n" - f"当前全局配置为:{'启用' if global_revoke else '禁用'}(延迟 {delay} 秒)" - ) - else: - yield event.plain_result("ℹ️ 当前会话没有设置覆盖,已在使用全局配置。") - else: - enabled = value == "on" - success = await self._core.session_config.set_session_auto_revoke_r18( - session_id, is_group, enabled - ) - if success: - yield event.plain_result( - f"✅ 已将当前{session_type}的自动撤回设置为:{'启用' if enabled else '禁用'}\n" - f"此后发送的 R18 内容将{'在 {delay} 秒后自动撤回' if enabled else '不会自动撤回'}(优先于全局配置)。" - ) - else: - yield event.plain_result("❌ 设置失败,请稍后再试。") - - async def _handle_config_send( - self, event, session_id, is_group, session_type, value - ): - """处理发送模式配置。""" - if value not in ("image", "forward", "auto", "clear", ""): - yield event.plain_result( - "❌ 无效的发送模式。\n" - "可用模式:image(直接发送)、forward(合并转发)、auto(自动选择)、clear(清除设置)" - ) - return - - global_send = self._core.config.send_mode - - if value == "clear" or value == "": - success = await self._core.session_config.clear_session_send_mode( - session_id, is_group - ) - if success: - yield event.plain_result( - f"✅ 已清除当前{session_type}的发送模式设置,将使用全局配置。\n" - f"当前全局配置为:{global_send}" - ) - else: - yield event.plain_result("ℹ️ 当前会话没有设置覆盖,已在使用全局配置。") - else: - success = await self._core.session_config.set_session_send_mode( - session_id, is_group, value - ) - if success: - mode_desc = { - "image": "直接发送图片", - "forward": "合并转发消息", - "auto": "自动选择(单张直接发送,多张合并转发)", - } - yield event.plain_result( - f"✅ 已将当前{session_type}的发送模式设置为:{value}\n" - f"说明:{mode_desc.get(value, '')}(优先于全局配置)。" - ) - else: - yield event.plain_result("❌ 设置失败,请稍后再试。") - - # ============ 黑白名单管理命令(简化中文版)============ - - async def handle_setu_block_user(self, event: AstrMessageEvent, args: str = ""): - """处理 /拉黑色图用户 命令(AT某人加入Setu黑名单)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - target_id = self._extract_target_id(event, args) - if not target_id: - yield event.plain_result( - "❌ 请通过 AT (@) 指定要拉黑的用户。\n用法:/拉黑色图用户 @用户名" - ) - return - - sender_id = event.get_sender_id() - if target_id == str(sender_id): - yield event.plain_result("❌ 不能将自己加入黑名单。") - return - - success = self._core.access_control.add_setu_blocked_user(target_id) - if success: - yield event.plain_result(f"✅ 已将用户 `{target_id}` 加入色图功能黑名单。") - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_setu_unblock_user(self, event: AstrMessageEvent, args: str = ""): - """处理 /解除色图拉黑 命令(AT某人从Setu黑名单移除)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - target_id = self._extract_target_id(event, args) - if not target_id: - yield event.plain_result( - "❌ 请通过 AT (@) 指定要解除拉黑的用户。\n用法:/解除色图拉黑 @用户名" - ) - return - - success = self._core.access_control.remove_setu_blocked_user(target_id) - if success: - yield event.plain_result( - f"✅ 已将用户 `{target_id}` 从色图功能黑名单移除。" - ) - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_fortune_block_user(self, event: AstrMessageEvent, args: str = ""): - """处理 /拉黑运势用户 命令(AT某人加入Fortune黑名单)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - target_id = self._extract_target_id(event, args) - if not target_id: - yield event.plain_result( - "❌ 请通过 AT (@) 指定要拉黑的用户。\n用法:/拉黑运势用户 @用户名" - ) - return - - sender_id = event.get_sender_id() - if target_id == str(sender_id): - yield event.plain_result("❌ 不能将自己加入黑名单。") - return - - success = self._core.access_control.add_fortune_blocked_user(target_id) - if success: - yield event.plain_result(f"✅ 已将用户 `{target_id}` 加入运势功能黑名单。") - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_fortune_unblock_user( - self, event: AstrMessageEvent, args: str = "" - ): - """处理 /解除运势拉黑 命令(AT某人从Fortune黑名单移除)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - target_id = self._extract_target_id(event, args) - if not target_id: - yield event.plain_result( - "❌ 请通过 AT (@) 指定要解除拉黑的用户。\n用法:/解除运势拉黑 @用户名" - ) - return - - success = self._core.access_control.remove_fortune_blocked_user(target_id) - if success: - yield event.plain_result( - f"✅ 已将用户 `{target_id}` 从运势功能黑名单移除。" - ) - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_setu_trust_user(self, event: AstrMessageEvent, args: str = ""): - """处理 /信任色图用户 命令(AT某人加入Setu白名单)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - target_id = self._extract_target_id(event, args) - if not target_id: - yield event.plain_result( - "❌ 请通过 AT (@) 指定要信任的用户。\n用法:/信任色图用户 @用户名" - ) - return - - success = self._core.access_control.add_setu_whitelist_user(target_id) - if success: - yield event.plain_result(f"✅ 已将用户 `{target_id}` 加入色图功能白名单。") - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_setu_untrust_user(self, event: AstrMessageEvent, args: str = ""): - """处理 /取消色图信任 命令(AT某人从Setu白名单移除)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - target_id = self._extract_target_id(event, args) - if not target_id: - yield event.plain_result( - "❌ 请通过 AT (@) 指定要取消信任的用户。\n用法:/取消色图信任 @用户名" - ) - return - - success = self._core.access_control.remove_setu_whitelist_user(target_id) - if success: - yield event.plain_result( - f"✅ 已将用户 `{target_id}` 从色图功能白名单移除。" - ) - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_fortune_trust_user(self, event: AstrMessageEvent, args: str = ""): - """处理 /信任运势用户 命令(AT某人加入Fortune白名单)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - target_id = self._extract_target_id(event, args) - if not target_id: - yield event.plain_result( - "❌ 请通过 AT (@) 指定要信任的用户。\n用法:/信任运势用户 @用户名" - ) - return - - success = self._core.access_control.add_fortune_whitelist_user(target_id) - if success: - yield event.plain_result(f"✅ 已将用户 `{target_id}` 加入运势功能白名单。") - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_fortune_untrust_user( - self, event: AstrMessageEvent, args: str = "" - ): - """处理 /取消运势信任 命令(AT某人从Fortune白名单移除)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - target_id = self._extract_target_id(event, args) - if not target_id: - yield event.plain_result( - "❌ 请通过 AT (@) 指定要取消信任的用户。\n用法:/取消运势信任 @用户名" - ) - return - - success = self._core.access_control.remove_fortune_whitelist_user(target_id) - if success: - yield event.plain_result( - f"✅ 已将用户 `{target_id}` 从运势功能白名单移除。" - ) - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_enable_setu_group(self, event: AstrMessageEvent, args: str = ""): - """处理 /开启色图 命令(从色图黑名单移除当前群组)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - group_id = event.get_group_id() - if not group_id: - yield event.plain_result("❌ 此命令只能在群聊中使用。") - return - - gid = str(group_id) - success = self._core.access_control.remove_setu_blocked_group(gid) - if success: - yield event.plain_result("✅ 已在本群开启色图功能。") - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_disable_setu_group(self, event: AstrMessageEvent, args: str = ""): - """处理 /关闭色图 命令(将当前群组加入色图黑名单)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - group_id = event.get_group_id() - if not group_id: - yield event.plain_result("❌ 此命令只能在群聊中使用。") - return - - gid = str(group_id) - success = self._core.access_control.add_setu_blocked_group(gid) - if success: - yield event.plain_result("✅ 已在本群关闭色图功能。") - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_enable_fortune_group( - self, event: AstrMessageEvent, args: str = "" - ): - """处理 /开启运势 命令(从运势黑名单移除当前群组)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - group_id = event.get_group_id() - if not group_id: - yield event.plain_result("❌ 此命令只能在群聊中使用。") - return - - gid = str(group_id) - success = self._core.access_control.remove_fortune_blocked_group(gid) - if success: - yield event.plain_result("✅ 已在本群开启运势功能。") - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - async def handle_disable_fortune_group( - self, event: AstrMessageEvent, args: str = "" - ): - """处理 /关闭运势 命令(将当前群组加入运势黑名单)。""" - if not self._core: - yield event.plain_result("插件尚未就绪,请稍后再试。") - return - - if not self._check_admin(event): - yield event.plain_result("❌ 权限不足:此命令仅限管理员或超级管理员使用。") - return - - group_id = event.get_group_id() - if not group_id: - yield event.plain_result("❌ 此命令只能在群聊中使用。") - return - - gid = str(group_id) - success = self._core.access_control.add_fortune_blocked_group(gid) - if success: - yield event.plain_result("✅ 已在本群关闭运势功能。") - else: - yield event.plain_result("❌ 操作失败,请稍后再试。") - - def _extract_target_id(self, event: AstrMessageEvent, args: str) -> str | None: - """从消息中提取目标用户 ID(仅支持 AT)。""" - for comp in event.get_messages(): - if hasattr(comp, "qq") and comp.qq: - target = str(comp.qq) - if target not in ("all", "0"): - return target - return None diff --git a/handlers/llm_handlers.py b/handlers/llm_handlers.py deleted file mode 100644 index 9e2921d..0000000 --- a/handlers/llm_handlers.py +++ /dev/null @@ -1,544 +0,0 @@ -"""LLM 工具处理器模块。""" - -from __future__ import annotations - -import json -import re - -import mcp.types - -from astrbot.api import logger -from astrbot.api.event import AstrMessageEvent - -from .command_handlers import CommandHandler - -# 匹配 \uXXXX 格式的 Unicode 转义序列 -_UNICODE_ESCAPE_PATTERN = re.compile(r"\\u([0-9a-fA-F]{4})") - - -def _decode_unicode_escapes(text: str) -> str: - """解码字符串中的 Unicode 转义序列(如 \\u74f7\\u7b25\\u6728\\u684c -> 碧蓝档案)。 - - 只处理 \\uXXXX 格式的 Unicode 转义,避免过度解码其他转义序列(如 \\n, \\t 等)。 - """ - if not text or "\\u" not in text: - return text - try: - # 先尝试使用 json.loads 解码(处理带引号的 JSON 字符串) - if text.startswith('"') and text.endswith('"'): - return json.loads(text) - # 使用正则表达式只替换 \\uXXXX 格式的转义序列 - return _UNICODE_ESCAPE_PATTERN.sub(lambda m: chr(int(m.group(1), 16)), text) - except (ValueError, UnicodeDecodeError): - return text - - -class LlmHandlers: - """LLM 工具处理器集合。""" - - def __init__(self, plugin): - """初始化处理器。 - - 参数: - plugin: SetuPlugin 实例 - """ - self.plugin = plugin - self._cmd_handler: CommandHandler | None = None - - @property - def core(self): - """获取核心实例。""" - return self.plugin._core - - @property - def cmd_handler(self) -> CommandHandler: - """获取命令处理器(惰性初始化)。""" - if self._cmd_handler is None: - self._cmd_handler = CommandHandler(self.core, self.core.config) - return self._cmd_handler - - def _check_admin(self, event: AstrMessageEvent) -> bool: - """检查用户是否为管理员。""" - try: - if hasattr(event, "is_admin") and callable(getattr(event, "is_admin")): - if event.is_admin(): - return True - if hasattr(event, "is_super_user") and callable( - getattr(event, "is_super_user") - ): - if event.is_super_user(): - return True - if hasattr(event, "message_obj"): - msg_obj = event.message_obj - if hasattr(msg_obj, "sender") and hasattr(msg_obj.sender, "role"): - role = msg_obj.sender.role - if role in ("admin", "owner"): - return True - except AttributeError: - pass - return False - - async def _llm_get_setu_handler( - self, - event: AstrMessageEvent, - count=1, - tags: list[str] | str | dict | None = None, - ) -> mcp.types.CallToolResult: - """LLM 工具处理器:获取色图并直接发送。""" - if not self.core: - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text="插件尚未就绪,请稍后再试。" - ) - ] - ) - try: - # 处理可能被包装成字典的参数 - if isinstance(count, dict): - count = count.get("value", 1) - if isinstance(tags, dict): - tags = tags.get("value", []) - - # 构建 tags 字符串,确保所有元素为字符串,并解码 Unicode 转义 - if isinstance(tags, list): - decoded_tags = [_decode_unicode_escapes(str(tag)) for tag in tags] - tags_str = " ".join(decoded_tags) - else: - tags_str = _decode_unicode_escapes(str(tags or "")) - - # 跟踪发送结果 - sent_count = 0 - has_error = False - - async for result in self.cmd_handler.handle_setu_command( - event, count=str(count), tags=tags_str - ): - if result is not None: - # 检查内部成功标记 - if isinstance(result, dict) and result.get("send_success"): - sent_count = result.get("image_count", 0) - continue - try: - await self.plugin.context.send_message( - event.unified_msg_origin, result - ) - # 注意:这里不递增 sent_count,因为 result 可能是错误消息 - # 只有成功标记中的 image_count 才是实际发送的图片数 - except Exception as exc: - has_error = True - logger.warning("[llm_tool] Failed to send: %s", exc) - - # 根据实际结果返回不同消息 - if sent_count == 0: - if has_error: - msg = "图片发送失败" - else: - msg = "没有获取到图片或发送被阻止" - else: - msg = f"已成功发送 {sent_count} 张图片" - - return mcp.types.CallToolResult( - content=[mcp.types.TextContent(type="text", text=msg)] - ) - except (TypeError, ValueError, RuntimeError) as e: - logger.exception("LLM 工具获取色图失败") - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent(type="text", text=f"获取图片失败:{str(e)}") - ] - ) - - async def _llm_get_content_mode_handler( - self, - event: AstrMessageEvent, - ) -> mcp.types.CallToolResult: - """LLM 工具处理器:查看当前内容分级。""" - if not self.core: - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text="插件尚未就绪,请稍后再试。" - ) - ] - ) - try: - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - - # 获取会话级别的覆盖配置 - session_mode = await self.core.session_config.get_session_content_mode( - session_id, is_group - ) - # 获取全局配置 - global_mode = self.core.config.content_mode - # 获取实际生效的模式 - effective_mode = await self.core.get_effective_content_mode(event) - - session_type = "群聊" if is_group else "私聊" - - if session_mode: - msg = ( - f"当前{session_type}的内容分级设置:\n" - f"- 会话覆盖:{session_mode}\n" - f"- 全局配置:{global_mode}\n" - f"- 生效模式:{effective_mode}\n\n" - f"说明:此会话已设置独立的内容分级,优先于全局配置。" - ) - else: - msg = ( - f"当前{session_type}的内容分级设置:\n" - f"- 会话覆盖:未设置(使用全局配置)\n" - f"- 全局配置:{global_mode}\n" - f"- 生效模式:{effective_mode}\n\n" - f"说明:此会话使用全局内容分级配置。" - ) - - return mcp.types.CallToolResult( - content=[mcp.types.TextContent(type="text", text=msg)] - ) - except Exception as e: - logger.exception("LLM 工具查看内容分级失败") - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text=f"查看内容分级失败:{str(e)}" - ) - ] - ) - - async def _llm_set_content_mode_handler( - self, - event: AstrMessageEvent, - mode: str | dict | None = None, - ) -> mcp.types.CallToolResult: - """LLM 工具处理器:设置内容分级。""" - if not self.core: - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text="插件尚未就绪,请稍后再试。" - ) - ] - ) - - # 检查权限(管理员或超级管理员) - if not self._check_admin(event): - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", - text="❌ 权限不足:设置内容分级需要管理员或超级管理员权限。", - ) - ] - ) - - try: - # 处理可能被包装成字典的参数 - if isinstance(mode, dict): - mode = mode.get("value", "") - if not mode: - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", - text="❌ 请指定要设置的内容分级:sfw(全年龄)、r18(成人)、mix(混合)或 clear(清除设置)。", - ) - ] - ) - - mode = str(mode).strip().lower() - if mode not in ("sfw", "r18", "mix", "clear"): - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", - text=f"❌ 无效的模式 '{mode}'。\n可用模式:sfw(全年龄)、r18(成人)、mix(混合)、clear(清除设置)", - ) - ] - ) - - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - session_type = "群聊" if is_group else "私聊" - - if mode == "clear": - success = await self.core.session_config.clear_session_content_mode( - session_id, is_group - ) - if success: - global_mode = self.core.config.content_mode - msg = ( - f"✅ 已清除当前{session_type}的内容分级设置,将使用全局配置。\n" - f"当前全局配置为:{global_mode}" - ) - else: - msg = "ℹ️ 当前会话没有设置覆盖,已在使用全局配置。" - else: - success = await self.core.session_config.set_session_content_mode( - session_id, is_group, mode - ) - if success: - msg = ( - f"✅ 已将当前{session_type}的内容分级设置为:{mode}\n" - f"此后发送的图片将使用此模式(优先于全局配置)。" - ) - else: - msg = "❌ 设置失败,请稍后再试。" - - return mcp.types.CallToolResult( - content=[mcp.types.TextContent(type="text", text=msg)] - ) - except Exception as e: - logger.exception("LLM 工具设置内容分级失败") - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text=f"设置内容分级失败:{str(e)}" - ) - ] - ) - - async def _llm_set_r18_docx_mode_handler( - self, - event: AstrMessageEvent, - enabled: bool | dict | None = None, - ) -> mcp.types.CallToolResult: - """LLM 工具处理器:设置会话级别的 R18 Docx 模式。""" - if not self.core: - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text="插件尚未就绪,请稍后再试。" - ) - ] - ) - - if not self._check_admin(event): - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", - text="❌ 权限不足:设置 R18 Docx 模式需要管理员或超级管理员权限。", - ) - ] - ) - - try: - # 处理可能被包装成字典的参数 - if isinstance(enabled, dict): - enabled = enabled.get("value") - - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - session_type = "群聊" if is_group else "私聊" - - # 获取当前全局配置 - global_mode = self.core.config.r18_docx_mode - - if enabled == "clear" or enabled is None: - # 清除会话设置 - success = await self.core.session_config.clear_session_r18_docx_mode( - session_id, is_group - ) - if success: - msg = ( - f"✅ 已清除当前{session_type}的 R18 Docx 模式设置,将使用全局配置。\n" - f"当前全局配置为:{'启用' if global_mode else '禁用'}" - ) - else: - msg = "ℹ️ 当前会话没有设置覆盖,已在使用全局配置。" - else: - # 转换为布尔值 - bool_enabled = bool(enabled) - success = await self.core.session_config.set_session_r18_docx_mode( - session_id, is_group, bool_enabled - ) - if success: - msg = ( - f"✅ 已将当前{session_type}的 R18 Docx 模式设置为:{'启用' if bool_enabled else '禁用'}\n" - f"此后发送的 R18 图片将{'打包为 DOCX 文件' if bool_enabled else '直接发送'}(优先于全局配置)。" - ) - else: - msg = "❌ 设置失败,请稍后再试。" - - return mcp.types.CallToolResult( - content=[mcp.types.TextContent(type="text", text=msg)] - ) - except Exception as e: - logger.exception("LLM 工具设置 R18 Docx 模式失败") - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text=f"设置 R18 Docx 模式失败:{str(e)}" - ) - ] - ) - - async def _llm_set_auto_revoke_handler( - self, - event: AstrMessageEvent, - enabled: bool | dict | None = None, - ) -> mcp.types.CallToolResult: - """LLM 工具处理器:设置会话级别的自动撤回。""" - if not self.core: - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text="插件尚未就绪,请稍后再试。" - ) - ] - ) - - if not self._check_admin(event): - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", - text="❌ 权限不足:设置自动撤回需要管理员或超级管理员权限。", - ) - ] - ) - - try: - # 处理可能被包装成字典的参数 - if isinstance(enabled, dict): - enabled = enabled.get("value") - - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - session_type = "群聊" if is_group else "私聊" - - # 获取当前全局配置 - global_mode = self.core.config.auto_revoke_r18 - delay = self.core.config.auto_revoke_delay - - if enabled == "clear" or enabled is None: - # 清除会话设置 - success = await self.core.session_config.clear_session_auto_revoke_r18( - session_id, is_group - ) - if success: - msg = ( - f"✅ 已清除当前{session_type}的自动撤回设置,将使用全局配置。\n" - f"当前全局配置为:{'启用' if global_mode else '禁用'}(延迟 {delay} 秒)" - ) - else: - msg = "ℹ️ 当前会话没有设置覆盖,已在使用全局配置。" - else: - # 转换为布尔值 - bool_enabled = bool(enabled) - success = await self.core.session_config.set_session_auto_revoke_r18( - session_id, is_group, bool_enabled - ) - if success: - msg = ( - f"✅ 已将当前{session_type}的自动撤回设置为:{'启用' if bool_enabled else '禁用'}\n" - f"此后发送的 R18 内容将{'在 {delay} 秒后自动撤回' if bool_enabled else '不会自动撤回'}(优先于全局配置)。" - ) - else: - msg = "❌ 设置失败,请稍后再试。" - - return mcp.types.CallToolResult( - content=[mcp.types.TextContent(type="text", text=msg)] - ) - except Exception as e: - logger.exception("LLM 工具设置自动撤回失败") - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text=f"设置自动撤回失败:{str(e)}" - ) - ] - ) - - async def _llm_set_send_mode_handler( - self, - event: AstrMessageEvent, - mode: str | dict | None = None, - ) -> mcp.types.CallToolResult: - """LLM 工具处理器:设置会话级别的发送模式。""" - if not self.core: - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text="插件尚未就绪,请稍后再试。" - ) - ] - ) - - if not self._check_admin(event): - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", - text="❌ 权限不足:设置发送模式需要管理员或超级管理员权限。", - ) - ] - ) - - try: - # 处理可能被包装成字典的参数 - if isinstance(mode, dict): - mode = mode.get("value") - - session_id = event.get_session_id() - is_group = bool(event.get_group_id()) - session_type = "群聊" if is_group else "私聊" - - # 获取当前全局配置 - global_mode = self.core.config.send_mode - - if mode == "clear" or mode is None: - # 清除会话设置 - success = await self.core.session_config.clear_session_send_mode( - session_id, is_group - ) - if success: - msg = ( - f"✅ 已清除当前{session_type}的发送模式设置,将使用全局配置。\n" - f"当前全局配置为:{global_mode}" - ) - else: - msg = "ℹ️ 当前会话没有设置覆盖,已在使用全局配置。" - else: - mode_str = str(mode).strip().lower() - if mode_str not in ("image", "forward", "auto"): - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", - text=f"❌ 无效的发送模式 '{mode_str}'。\n" - f"可用模式:image(直接发送)、forward(合并转发)、auto(自动选择)、clear(清除设置)", - ) - ] - ) - - success = await self.core.session_config.set_session_send_mode( - session_id, is_group, mode_str - ) - if success: - mode_desc = { - "image": "直接发送图片", - "forward": "合并转发消息", - "auto": "自动选择(单张直接发送,多张合并转发)", - } - msg = ( - f"✅ 已将当前{session_type}的发送模式设置为:{mode_str}\n" - f"说明:{mode_desc.get(mode_str, '')}(优先于全局配置)。" - ) - else: - msg = "❌ 设置失败,请稍后再试。" - - return mcp.types.CallToolResult( - content=[mcp.types.TextContent(type="text", text=msg)] - ) - except Exception as e: - logger.exception("LLM 工具设置发送模式失败") - return mcp.types.CallToolResult( - content=[ - mcp.types.TextContent( - type="text", text=f"设置发送模式失败:{str(e)}" - ) - ] - ) diff --git a/llm_registry.py b/llm_registry.py deleted file mode 100644 index ff8e6a4..0000000 --- a/llm_registry.py +++ /dev/null @@ -1,382 +0,0 @@ -"""LLM 工具注册管理模块。 - -集中管理所有 LLM 工具的定义、注册和注销。 -将工具定义从 main.py 分离,便于维护和扩展。 -""" - -from __future__ import annotations - -from collections.abc import Callable -from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any - -from astrbot.core.provider.register import llm_tools - -if TYPE_CHECKING: - pass - - -@dataclass -class LlmToolDefinition: - """LLM 工具定义。 - - 包含工具的名称、描述、参数和处理函数。 - """ - - name: str - handler: Callable[..., Any] - args: list[dict[str, Any]] = field(default_factory=list) - description: str = "" - module_path: str | None = None - - -class LlmToolRegistry: - """LLM 工具注册管理器。 - - 管理工具的注册和注销,支持按模块分组管理。 - """ - - def __init__(self): - self._registered_tools: dict[str, str] = {} - # key: tool_name, value: module_path (用于追踪工具来源) - - def register_tool(self, tool_def: LlmToolDefinition, module_path: str) -> bool: - """注册单个 LLM 工具。 - - Args: - tool_def: 工具定义 - module_path: 模块路径(用于插件管理) - - Returns: - 是否成功注册 - """ - try: - llm_tools.add_func( - name=tool_def.name, - func_args=tool_def.args, - desc=tool_def.description, - handler=tool_def.handler, - ) - tool = llm_tools.get_func(tool_def.name) - if tool: - tool.handler_module_path = module_path - - self._registered_tools[tool_def.name] = module_path - return True - except (AttributeError, RuntimeError): - return False - - def register_tools( - self, tools: list[LlmToolDefinition], module_path: str - ) -> list[str]: - """批量注册 LLM 工具。 - - Args: - tools: 工具定义列表 - module_path: 模块路径 - - Returns: - 成功注册的工具名称列表 - """ - registered: list[str] = [] - for tool_def in tools: - if self.register_tool(tool_def, module_path): - registered.append(tool_def.name) - return registered - - def unregister_tool(self, name: str) -> bool: - """注销单个 LLM 工具。 - - Args: - name: 工具名称 - - Returns: - 是否成功注销 - """ - try: - llm_tools.remove_func(name) - self._registered_tools.pop(name, None) - return True - except (AttributeError, RuntimeError): - return False - - def unregister_tools(self, names: list[str]) -> list[str]: - """批量注销 LLM 工具。 - - Args: - names: 工具名称列表 - - Returns: - 成功注销的工具名称列表 - """ - unregistered: list[str] = [] - for name in names: - if self.unregister_tool(name): - unregistered.append(name) - return unregistered - - def unregister_by_module(self, module_path: str) -> list[str]: - """注销指定模块的所有工具。 - - Args: - module_path: 模块路径 - - Returns: - 成功注销的工具名称列表 - """ - tools_to_remove = [ - name for name, path in self._registered_tools.items() if path == module_path - ] - return self.unregister_tools(tools_to_remove) - - def get_registered_tools(self) -> dict[str, str]: - """获取已注册的工具映射。 - - Returns: - 工具名称到模块路径的映射 - """ - return self._registered_tools.copy() - - -# 全局注册管理器实例 -_registry = LlmToolRegistry() - - -def get_registry() -> LlmToolRegistry: - """获取全局 LLM 工具注册管理器。""" - return _registry - - -# ==================== Setu 工具定义 ==================== - - -def get_setu_tool_definitions(handler) -> list[LlmToolDefinition]: - """获取 Setu 相关的 LLM 工具定义。 - - Args: - handler: LlmHandlers 实例 - - Returns: - 工具定义列表 - """ - return [ - LlmToolDefinition( - name="get_setu_image", - handler=handler._llm_get_setu_handler, - args=[ - { - "name": "count", - "type": "integer", - "description": "Number of images.", - }, - {"name": "tags", "type": "array", "items": {"type": "string"}}, - ], - description="Fetch random anime images.", - ), - LlmToolDefinition( - name="get_setu_content_mode", - handler=handler._llm_get_content_mode_handler, - args=[], - description="Get content mode.", - ), - LlmToolDefinition( - name="set_setu_content_mode", - handler=handler._llm_set_content_mode_handler, - args=[ - { - "name": "mode", - "type": "string", - "enum": ["sfw", "r18", "mix", "clear"], - }, - ], - description="Set content mode.", - ), - LlmToolDefinition( - name="set_setu_r18_docx_mode", - handler=handler._llm_set_r18_docx_mode_handler, - args=[ - {"name": "enabled", "type": "boolean"}, - ], - description="Set R18 Docx mode.", - ), - LlmToolDefinition( - name="set_setu_auto_revoke", - handler=handler._llm_set_auto_revoke_handler, - args=[ - {"name": "enabled", "type": "boolean"}, - ], - description="Set auto-revoke.", - ), - LlmToolDefinition( - name="set_setu_send_mode", - handler=handler._llm_set_send_mode_handler, - args=[ - { - "name": "mode", - "type": "string", - "enum": ["image", "forward", "auto", "clear"], - "description": "Send mode override for current session.", - } - ], - description="Set session send mode.", - ), - ] - - -SETU_TOOL_NAMES = [ - "get_setu_image", - "get_setu_content_mode", - "set_setu_content_mode", - "set_setu_r18_docx_mode", - "set_setu_auto_revoke", - "set_setu_send_mode", -] - - -def register_setu_tools(handler, module_path: str) -> list[str]: - """注册 Setu 相关的 LLM 工具。 - - Args: - handler: LlmHandlers 实例 - module_path: 模块路径 - - Returns: - 成功注册的工具名称列表 - """ - tools = get_setu_tool_definitions(handler) - return _registry.register_tools(tools, module_path) - - -def unregister_setu_tools() -> list[str]: - """注销 Setu 相关的 LLM 工具。 - - Returns: - 成功注销的工具名称列表 - """ - return _registry.unregister_tools(SETU_TOOL_NAMES) - - -# ==================== Fortune 工具定义 ==================== - - -def get_fortune_tool_definitions(handler) -> list[LlmToolDefinition]: - """获取今日运势相关的 LLM 工具定义。 - - Args: - handler: FortuneLlmHandler 实例 - - Returns: - 工具定义列表 - """ - return [ - LlmToolDefinition( - name="get_today_fortune", - handler=handler.llm_get_fortune, - args=[], - description="Get today's fortune for the user.", - ), - LlmToolDefinition( - name="refresh_my_fortune", - handler=handler.llm_refresh_fortune, - args=[], - description="Refresh my today's fortune (admin only).", - ), - LlmToolDefinition( - name="refresh_group_fortune", - handler=handler.llm_refresh_group_fortune, - args=[], - description="Refresh today's fortune for the current group (admin only).", - ), - LlmToolDefinition( - name="refresh_all_fortune", - handler=handler.llm_refresh_all_fortune, - args=[], - description="Refresh today's fortune for all users (super admin only).", - ), - LlmToolDefinition( - name="get_fortune_config", - handler=handler.llm_get_fortune_config, - args=[], - description="Get the fortune configuration for the current session.", - ), - LlmToolDefinition( - name="set_fortune_config", - handler=handler.llm_set_fortune_config, - args=[ - { - "name": "tags", - "type": "string", - "description": "Tags for fortune images, e.g., 'girl,cute'. Leave empty to clear.", - }, - { - "name": "mode", - "type": "string", - "enum": ["sfw", "r18", "mix"], - "description": "Content mode for fortune images.", - }, - ], - description="Set the fortune configuration for the current session (admin only).", - ), - ] - - -FORTUNE_TOOL_NAMES = [ - "get_today_fortune", - "refresh_my_fortune", - "refresh_group_fortune", - "refresh_all_fortune", - "get_fortune_config", - "set_fortune_config", -] - - -def register_fortune_tools(handler, module_path: str) -> list[str]: - """注册今日运势相关的 LLM 工具。 - - Args: - handler: FortuneLlmHandler 实例 - module_path: 模块路径 - - Returns: - 成功注册的工具名称列表 - """ - tools = get_fortune_tool_definitions(handler) - return _registry.register_tools(tools, module_path) - - -def unregister_fortune_tools() -> list[str]: - """注销今日运势相关的 LLM 工具。 - - Returns: - 成功注销的工具名称列表 - """ - return _registry.unregister_tools(FORTUNE_TOOL_NAMES) - - -def unregister_all_tools() -> list[str]: - """注销所有工具(包括 Setu 和 Fortune)。 - - Returns: - 成功注销的工具名称列表 - """ - all_names = SETU_TOOL_NAMES + FORTUNE_TOOL_NAMES - return _registry.unregister_tools(all_names) - - -__all__ = [ - "LlmToolDefinition", - "LlmToolRegistry", - "get_registry", - # Setu 工具 - "get_setu_tool_definitions", - "register_setu_tools", - "unregister_setu_tools", - "SETU_TOOL_NAMES", - # Fortune 工具 - "get_fortune_tool_definitions", - "register_fortune_tools", - "unregister_fortune_tools", - "FORTUNE_TOOL_NAMES", - # 综合操作 - "unregister_all_tools", -] diff --git a/main.py b/main.py index 87512fc..b2a39ed 100644 --- a/main.py +++ b/main.py @@ -1,431 +1,373 @@ -"""Setu(随机图片)插件 - 主入口。 +"""AstrBot Setu Plugin - Main entry point. -支持多 API 提供商、可配置的内容模式、多种发送模式、 -HTML 卡片包装、LLM 工具调用、以及自定义 API 支持。 -集成今日运势功能。 +Commands are defined directly on the Star subclass so AstrBot's decorator-based +registration discovers them. Business logic is delegated to handler helpers. """ from __future__ import annotations -from pathlib import Path -from typing import AsyncGenerator, Any +from collections.abc import AsyncGenerator +from typing import Any +from astrbot.api import logger from astrbot.api.event import AstrMessageEvent, filter from astrbot.api.star import Context, Star, StarTools from astrbot.core import AstrBotConfig -from astrbot.core.message.components import Plain -from astrbot.core.provider.register import llm_tools -from .config import SetuConfig -from .constants import COMMAND_PATTERN, FORTUNE_PATTERN -from .core import SetuCore -from .fortune import FortuneManager -from .handlers import CommandHandler, LlmHandlers +from .src.infrastructure import ( + init_access_control_repo, + init_fortune_repo, + init_provider_from_config, + init_session_config_repo, +) +from .src.infrastructure.astrbot import init_config, set_plugin_context +from .src.infrastructure.astrbot.commands import ( + FortuneCommandHandler, + SessionConfigCommandHandler, + SetuCommandHandler, + register_fortune_llm_tools, + register_session_config_llm_tools, + register_setu_llm_tools, + unregister_fortune_llm_tools, + unregister_session_config_llm_tools, + unregister_setu_llm_tools, +) +from .src.infrastructure.astrbot.session_config_api import ( + register_session_config_web_apis, +) +from .src.shared.send_cache import clear_send_cache, init_send_cache + +# Regex patterns for command triggers +SETU_REGEX_PATTERN = r"^/?(来\s*(.*?)(份|个|张|点))(.*?)(?:福利|色|瑟|涩|塞)?图$" + +# Module-level handler singletons +_setu_handler: SetuCommandHandler | None = None +_fortune_handler: FortuneCommandHandler | None = None +_session_config_handler: SessionConfigCommandHandler | None = None + +FORTUNE_REFRESH_COMMANDS = { + "刷新今日运势": "self", + "刷新jrys": "self", + "flush_jrys": "self", + "刷新本群今日运势": "group", + "刷新本群jrys": "group", + "flush_group_jrys": "group", + "刷新全局今日运势": "all", + "刷新全局jrys": "all", + "flush_all_jrys": "all", +} +FORTUNE_TOGGLE_COMMANDS = {"开启运势": "enable", "关闭运势": "disable"} +FORTUNE_USER_COMMANDS = { + "拉黑运势用户": "block", + "解除运势拉黑": "unblock", + "信任运势用户": "trust", + "取消运势信任": "untrust", +} +FORTUNE_REFRESH_ARG_ALIASES = { + "": "self", + "我": "self", + "自己": "self", + "我的": "self", + "self": "self", + "me": "self", + "本群": "group", + "群": "group", + "group": "group", + "全局": "all", + "全部": "all", + "all": "all", + "global": "all", +} +FORTUNE_TOGGLE_ARG_ALIASES = { + "开": "enable", + "开启": "enable", + "on": "enable", + "关": "disable", + "关闭": "disable", + "off": "disable", +} +FORTUNE_USER_ARG_ALIASES = { + "拉黑": "block", + "黑名单": "block", + "block": "block", + "解除拉黑": "unblock", + "解黑": "unblock", + "取消拉黑": "unblock", + "unblock": "unblock", + "信任": "trust", + "白名单": "trust", + "trust": "trust", + "取消信任": "untrust", + "移除信任": "untrust", + "取消白名单": "untrust", + "untrust": "untrust", +} + + +def _get_invoked_command(event: AstrMessageEvent) -> str: + raw_message = getattr(event, "message_str", None) + if not raw_message and hasattr(event, "get_message_str"): + raw_message = event.get_message_str() + text = str(raw_message or "").strip() + if text.startswith("/"): + text = text[1:].strip() + return text.split(maxsplit=1)[0] if text else "" + + +def _resolve_fortune_refresh_target(event: AstrMessageEvent, args: str) -> str: + command = _get_invoked_command(event) + if command in FORTUNE_REFRESH_COMMANDS: + return FORTUNE_REFRESH_COMMANDS[command] + + target = FORTUNE_REFRESH_ARG_ALIASES.get((args or "").strip().lower()) + if target: + return target + raise ValueError("用法:/运势刷新 [我|本群|全局]") + + +def _resolve_fortune_toggle_action(event: AstrMessageEvent, args: str) -> str: + command = _get_invoked_command(event) + if command in FORTUNE_TOGGLE_COMMANDS: + return FORTUNE_TOGGLE_COMMANDS[command] + + action = FORTUNE_TOGGLE_ARG_ALIASES.get((args or "").strip().lower()) + if action: + return action + raise ValueError("用法:/运势开关 <开|关>") + + +def _resolve_fortune_user_action(event: AstrMessageEvent, args: str) -> tuple[str, str]: + command = _get_invoked_command(event) + if command in FORTUNE_USER_COMMANDS: + return FORTUNE_USER_COMMANDS[command], args + + normalized_args = (args or "").strip() + if not normalized_args: + raise ValueError("用法:/运势用户 <拉黑|解黑|信任|取消信任> [用户]") + + parts = normalized_args.split(maxsplit=1) + action = FORTUNE_USER_ARG_ALIASES.get(parts[0].lower()) + if not action: + raise ValueError("用法:/运势用户 <拉黑|解黑|信任|取消信任> [用户]") + target = parts[1] if len(parts) > 1 else "" + return action, target class SetuPlugin(Star): - """色图插件主类(含今日运势)。""" + """Main plugin class with command handlers on the Star subclass. + + AstrBot only discovers @filter.command/@filter.regex handlers whose + __module__ matches the plugin's main module path. Defining handlers + here ensures they are found and bound to the Star instance. + """ def __init__(self, context: Context, config: AstrBotConfig): super().__init__(context, config) self.context = context - self.config = config - self._core: SetuCore | None = None - self._plugin_data_dir: Path = StarTools.get_data_dir(self.name) - self._cmd_handler: CommandHandler | None = None - self._llm_handlers: LlmHandlers | None = None - self._fortune_manager: FortuneManager | None = None + self._plugin_config = config + register_session_config_web_apis(self.context) async def initialize(self) -> None: - """初始化插件。""" - cfg = SetuConfig(self.config) - self._core = SetuCore(self, cfg, self._plugin_data_dir, self.config) - await self._core.initialize() - - self._cmd_handler = CommandHandler(self._core, cfg) - self._llm_handlers = LlmHandlers(self) - - # 初始化今日运势模块 - fortune_cfg = getattr(cfg, "fortune", {}) - if isinstance(fortune_cfg, dict) and fortune_cfg.get("enabled", True): - self._fortune_manager = FortuneManager( - self, self._plugin_data_dir / "fortune", fortune_cfg - ) - await self._fortune_manager.initialize() - await self._register_fortune_llm_tools() - - # 注册 Setu LLM 工具 - await self._register_setu_llm_tools() - - async def _register_setu_llm_tools(self) -> None: - """注册 Setu 相关的 LLM 工具。""" - tools = [ - ( - "get_setu_image", - self._llm_handlers._llm_get_setu_handler, - [ - { - "name": "count", - "type": "integer", - "description": "Number of images.", - }, - {"name": "tags", "type": "array", "items": {"type": "string"}}, - ], - "Fetch random anime images.", - ), - ( - "get_setu_content_mode", - self._llm_handlers._llm_get_content_mode_handler, - [], - "Get content mode.", - ), - ( - "set_setu_content_mode", - self._llm_handlers._llm_set_content_mode_handler, - [ - { - "name": "mode", - "type": "string", - "enum": ["sfw", "r18", "mix", "clear"], - }, - ], - "Set content mode.", - ), - ( - "set_setu_r18_docx_mode", - self._llm_handlers._llm_set_r18_docx_mode_handler, - [ - {"name": "enabled", "type": "boolean"}, - ], - "Set R18 Docx mode.", - ), - ( - "set_setu_auto_revoke", - self._llm_handlers._llm_set_auto_revoke_handler, - [ - {"name": "enabled", "type": "boolean"}, - ], - "Set auto-revoke.", - ), - ( - "set_setu_send_mode", - self._llm_handlers._llm_set_send_mode_handler, - [ - { - "name": "mode", - "type": "string", - "enum": ["image", "forward", "auto", "clear"], - }, - ], - "Set session send mode.", - ), - ] - - for name, handler, args, desc in tools: - try: - llm_tools.add_func( - name=name, func_args=args, desc=desc, handler=handler - ) - tool = llm_tools.get_func(name) - if tool: - tool.handler_module_path = __name__ - except (AttributeError, RuntimeError): - pass - - async def _register_fortune_llm_tools(self) -> None: - """注册今日运势的 LLM 工具。""" - if not self._fortune_manager or not self._fortune_manager.llm_handler: + global _setu_handler, _fortune_handler, _session_config_handler + + raw_config = self._runtime_plugin_config() + cfg = init_config(raw_config) + set_plugin_context(self.context) + data_dir = StarTools.get_data_dir(self.name) + + init_provider_from_config(cfg) + + await init_access_control_repo(data_dir, raw_config) + await init_fortune_repo(data_dir) + await init_session_config_repo(data_dir) + await init_send_cache( + data_dir, + enabled=cfg.cache_enabled, + ttl_hours=cfg.cache_ttl_hours, + max_items=cfg.cache_max_items, + cleanup_on_start=cfg.cache_cleanup_on_start, + ) + + _setu_handler = SetuCommandHandler() + _fortune_handler = FortuneCommandHandler() + _session_config_handler = SessionConfigCommandHandler() + + try: + register_setu_llm_tools() + register_fortune_llm_tools() + register_session_config_llm_tools() + except Exception as e: + logger.warning("Failed to register LLM tools: %s", e) + + logger.info("SetuPlugin initialized successfully") + + async def terminate(self) -> None: + global _setu_handler, _fortune_handler, _session_config_handler + + from .src.infrastructure import ( + clear_fortune_repo, + clear_provider, + clear_repo, + clear_session_config_repo, + ) + + try: + unregister_setu_llm_tools() + unregister_fortune_llm_tools() + unregister_session_config_llm_tools() + except Exception: + pass + + clear_provider() + clear_repo() + clear_fortune_repo() + clear_session_config_repo() + clear_send_cache() + + _setu_handler = None + _fortune_handler = None + _session_config_handler = None + + logger.info("SetuPlugin terminated") + + def _runtime_plugin_config(self) -> dict[str, Any]: + """Return the plugin-scoped config dict passed in by AstrBot.""" + if isinstance(self._plugin_config, dict): + return dict(self._plugin_config) + items = getattr(self._plugin_config, "items", None) + if callable(items): + return dict(items()) + return dict(self._plugin_config) + + # ==================== Setu Commands ==================== + + @filter.regex(SETU_REGEX_PATTERN) + async def get_random_picture( + self, event: AstrMessageEvent + ) -> AsyncGenerator[Any, None]: + """来份色图 / 来张图 etc.""" + if _setu_handler is None: + yield event.plain_result("插件未初始化") return + async for result in _setu_handler.get_random_picture(event): + yield result - handler = self._fortune_manager.llm_handler - - tools = [ - ( - "get_today_fortune", - handler.llm_get_fortune, - [], - "Get today's fortune for the user.", - ), - ( - "refresh_my_fortune", - handler.llm_refresh_fortune, - [], - "Refresh my today's fortune (admin only).", - ), - ( - "refresh_group_fortune", - handler.llm_refresh_group_fortune, - [], - "Refresh today's fortune for the current group (admin only).", - ), - ( - "refresh_all_fortune", - handler.llm_refresh_all_fortune, - [], - "Refresh today's fortune for all users (super admin only).", - ), - ( - "get_fortune_config", - handler.llm_get_fortune_config, - [], - "Get the fortune configuration for the current session.", - ), - ( - "set_fortune_config", - handler.llm_set_fortune_config, - [ - { - "name": "tags", - "type": "string", - "description": "Tags for fortune images, e.g., 'girl,cute'. Leave empty to clear.", - }, - { - "name": "mode", - "type": "string", - "enum": ["sfw", "r18", "mix"], - "description": "Content mode for fortune images.", - }, - ], - "Set the fortune configuration for the current session (admin only).", - ), - ] - - for name, h, args, desc in tools: - try: - llm_tools.add_func(name=name, func_args=args, desc=desc, handler=h) - tool = llm_tools.get_func(name) - if tool: - tool.handler_module_path = __name__ - except (AttributeError, RuntimeError): - pass + @filter.command("setu") + async def setu_command( + self, event: AstrMessageEvent, count: str = "1", *, tags: str = "" + ) -> AsyncGenerator[Any, None]: + """/setu [count] [tags...]""" + if _setu_handler is None: + yield event.plain_result("插件未初始化") + return + async for result in _setu_handler.setu_command(event, count, tags=tags): + yield result - async def terminate(self) -> None: - """卸载插件。""" - if self._core: - self._core.terminate() - if self._fortune_manager: - self._fortune_manager.terminate() - - # 注销 LLM 工具 - for name in [ - "get_setu_image", - "get_setu_content_mode", - "set_setu_content_mode", - "set_setu_r18_docx_mode", - "set_setu_auto_revoke", - "set_setu_send_mode", - "get_today_fortune", - "refresh_my_fortune", - "refresh_group_fortune", - "refresh_all_fortune", - "get_fortune_config", - "set_fortune_config", - ]: - try: - llm_tools.remove_func(name) - except (AttributeError, RuntimeError): - pass - - def _is_explicit_slash_command(self, event: AstrMessageEvent) -> bool: - """Check whether original message text uses an explicit '/' command prefix.""" - for comp in event.get_messages(): - if isinstance(comp, Plain): - return comp.text.strip().startswith("/") - return False - - def _mark_event_handled(self, event: AstrMessageEvent, key: str) -> bool: - """Mark a plugin flow as handled for this event.""" - extra_key = f"setu_{key}_handled" - if event.get_extra(extra_key): - return False - event.set_extra(extra_key, True) - return True - - async def _run_fortune(self, event: AstrMessageEvent) -> AsyncGenerator[Any, None]: - if not self._mark_event_handled(event, "fortune"): + @filter.command("session_config") + async def session_config_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """/session_config get|set|clear""" + if _session_config_handler is None: + yield event.plain_result("插件未初始化") return - if self._fortune_manager and self._fortune_manager.cmd_handler: - async for result in self._fortune_manager.cmd_handler.handle_fortune(event): - yield result + async for result in _session_config_handler.session_config_command(event, args): + yield result - # ==================== Setu 命令 ==================== + # ==================== Fortune Commands ==================== - @filter.regex(COMMAND_PATTERN) - async def get_random_picture(self, event): - """处理色图请求命令。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_random_picture(event): - yield result + @filter.command("今日运势", alias={"jrys"}) + async def fortune_command( + self, event: AstrMessageEvent + ) -> AsyncGenerator[Any, None]: + """今日运势 / jrys""" + if _fortune_handler is None: + yield event.plain_result("插件未初始化") + return + async for result in _fortune_handler.fortune_command(event): + yield result - @filter.command("setu") - async def setu_command(self, event, count: str = "1", *, tags: str = ""): - """处理 /setu 命令。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_setu_command( - event, count, tags=tags - ): - yield result - - @filter.command("setu_mode") - async def setu_mode_command(self, event, mode: str = ""): - """处理 /setu_mode 命令。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_setu_mode(event, mode): - yield result - - # ==================== Fortune 命令 ==================== - - @filter.regex(FORTUNE_PATTERN) - async def fortune_regex(self, event): - """处理纯文本今日运势触发(今日运势, jrys)。""" - # Avoid duplicate handling when user sends /今日运势 or /jrys. - if self._is_explicit_slash_command(event): + @filter.command( + "运势刷新", + alias={ + "刷新今日运势", + "刷新jrys", + "flush_jrys", + "刷新本群今日运势", + "刷新本群jrys", + "flush_group_jrys", + "刷新全局今日运势", + "刷新全局jrys", + "flush_all_jrys", + }, + ) + async def fortune_refresh_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """运势刷新 [我|本群|全局]""" + if _fortune_handler is None: + yield event.plain_result("插件未初始化") return - async for result in self._run_fortune(event): + + try: + target = _resolve_fortune_refresh_target(event, args) + except ValueError as exc: + yield event.plain_result(str(exc)) + return + + handler_map = { + "self": _fortune_handler.refresh_fortune_command, + "group": _fortune_handler.refresh_group_fortune_command, + "all": _fortune_handler.refresh_all_fortune_command, + } + async for result in handler_map[target](event): yield result - @filter.command(command_name="今日运势", alias={"jrys"}) - async def fortune_command(self, event): - """处理今日运势命令 (/今日运势, /jrys)。""" - async for result in self._run_fortune(event): + @filter.command("运势开关", alias={"开启运势", "关闭运势"}) + async def fortune_toggle_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """运势开关 <开|关>""" + if _fortune_handler is None: + yield event.plain_result("插件未初始化") + return + + try: + action = _resolve_fortune_toggle_action(event, args) + except ValueError as exc: + yield event.plain_result(str(exc)) + return + + handler_map = { + "enable": _fortune_handler.enable_fortune_group_command, + "disable": _fortune_handler.disable_fortune_group_command, + } + async for result in handler_map[action](event, ""): yield result - @filter.command("刷新今日运势", alias={"刷新jrys", "flush_jrys"}) - async def refresh_fortune_command(self, event): - """处理 /刷新今日运势 命令。""" - if self._fortune_manager and self._fortune_manager.cmd_handler: - async for ( - result - ) in self._fortune_manager.cmd_handler.handle_refresh_fortune(event): - yield result - - @filter.command("刷新本群今日运势", alias={"刷新本群jrys", "flush_group_jrys"}) - async def refresh_group_fortune_command(self, event): - """处理 /刷新本群今日运势 命令。""" - if self._fortune_manager and self._fortune_manager.cmd_handler: - async for ( - result - ) in self._fortune_manager.cmd_handler.handle_refresh_group_fortune(event): - yield result - - @filter.command("刷新全局今日运势", alias={"刷新全局jrys", "flush_all_jrys"}) - async def refresh_all_fortune_command(self, event): - """处理 /刷新全局今日运势 命令。""" - if self._fortune_manager and self._fortune_manager.cmd_handler: - async for ( - result - ) in self._fortune_manager.cmd_handler.handle_refresh_all_fortune(event): - yield result - - @filter.command("jrys_config") - async def jrys_config_command(self, event, args: str = ""): - """处理 /jrys_config 命令。""" - if self._fortune_manager and self._fortune_manager.cmd_handler: - async for result in self._fortune_manager.cmd_handler.handle_fortune_config( - event, args - ): - yield result - - # ==================== 黑白名单管理命令(中文 - 完全独立的Setu和Fortune)==================== - - @filter.command("开启色图") - async def enable_setu_group_command(self, event, args: str = ""): - """处理 /开启色图 命令(在当前群组开启色图)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_enable_setu_group(event, args): - yield result - - @filter.command("关闭色图") - async def disable_setu_group_command(self, event, args: str = ""): - """处理 /关闭色图 命令(在当前群组关闭色图)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_disable_setu_group( - event, args - ): - yield result - - @filter.command("开启运势") - async def enable_fortune_group_command(self, event, args: str = ""): - """处理 /开启运势 命令(在当前群组开启运势)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_enable_fortune_group( - event, args - ): - yield result - - @filter.command("关闭运势") - async def disable_fortune_group_command(self, event, args: str = ""): - """处理 /关闭运势 命令(在当前群组关闭运势)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_disable_fortune_group( - event, args - ): - yield result - - # ----- 色图功能用户黑白名单 ----- - - @filter.command("拉黑色图用户") - async def block_setu_user_command(self, event, args: str = ""): - """处理 /拉黑色图用户 命令(AT某人加入色图黑名单)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_setu_block_user(event, args): - yield result - - @filter.command("解除色图拉黑") - async def unblock_setu_user_command(self, event, args: str = ""): - """处理 /解除色图拉黑 命令(AT某人从色图黑名单移除)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_setu_unblock_user(event, args): - yield result - - @filter.command("信任色图用户") - async def trust_setu_user_command(self, event, args: str = ""): - """处理 /信任色图用户 命令(AT某人加入色图白名单)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_setu_trust_user(event, args): - yield result - - @filter.command("取消色图信任") - async def untrust_setu_user_command(self, event, args: str = ""): - """处理 /取消色图信任 命令(AT某人从色图白名单移除)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_setu_untrust_user(event, args): - yield result - - # ----- 运势功能用户黑白名单 ----- - - @filter.command("拉黑运势用户") - async def block_fortune_user_command(self, event, args: str = ""): - """处理 /拉黑运势用户 命令(AT某人加入运势黑名单)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_fortune_block_user( - event, args - ): - yield result - - @filter.command("解除运势拉黑") - async def unblock_fortune_user_command(self, event, args: str = ""): - """处理 /解除运势拉黑 命令(AT某人从运势黑名单移除)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_fortune_unblock_user( - event, args - ): - yield result - - @filter.command("信任运势用户") - async def trust_fortune_user_command(self, event, args: str = ""): - """处理 /信任运势用户 命令(AT某人加入运势白名单)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_fortune_trust_user( - event, args - ): - yield result - - @filter.command("取消运势信任") - async def untrust_fortune_user_command(self, event, args: str = ""): - """处理 /取消运势信任 命令(AT某人从运势白名单移除)。""" - if self._cmd_handler: - async for result in self._cmd_handler.handle_fortune_untrust_user( - event, args - ): - yield result + @filter.command( + "运势用户", + alias={"拉黑运势用户", "解除运势拉黑", "信任运势用户", "取消运势信任"}, + ) + async def fortune_user_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """运势用户 <拉黑|解黑|信任|取消信任> [用户]""" + if _fortune_handler is None: + yield event.plain_result("插件未初始化") + return + + try: + action, target = _resolve_fortune_user_action(event, args) + except ValueError as exc: + yield event.plain_result(str(exc)) + return + + handler_map = { + "block": _fortune_handler.block_fortune_user_command, + "unblock": _fortune_handler.unblock_fortune_user_command, + "trust": _fortune_handler.trust_fortune_user_command, + "untrust": _fortune_handler.untrust_fortune_user_command, + } + async for result in handler_map[action](event, target): + yield result diff --git a/metadata.yaml b/metadata.yaml index 86866ca..efcbcb0 100644 --- a/metadata.yaml +++ b/metadata.yaml @@ -1,6 +1,6 @@ name: astrbot_plugin_setu display_name: 瑟瑟! -version: v1.3.0 +version: v2.0.0 author: FlanChanXwO desc: 随机福利图插件,支持标签与数量,以及图片分级控制,可以针对于特定平台配置概率绕过审核的发送方式。 repo: https://github.com/FlanChanXwO/astrbot_plugin_setu diff --git a/pages/sessionConfig/app.js b/pages/sessionConfig/app.js new file mode 100644 index 0000000..4612533 --- /dev/null +++ b/pages/sessionConfig/app.js @@ -0,0 +1,146 @@ +(function () { + const bridge = window.AstrBotPluginPage || null; + let toastTimer = null; + + function emptyForm() { + return { + session_id: '', + session_type: 'group', + display_name: '', + overrides: {}, + }; + } + + function apiResult(result) { + if (result && result.success === false) { + throw new Error(result.error || '请求失败'); + } + return result || {}; + } + + const store = { + loading: false, + keys: [], + sessions: [], + globalValues: {}, + form: emptyForm(), + toast: {show: false, type: 'success', message: ''}, + + async reload() { + this.loading = true; + try { + if (!bridge) throw new Error('AstrBotPluginPage bridge not available'); + await bridge.ready(); + const result = apiResult(await bridge.apiGet('session-config')); + this.keys = result.keys || []; + this.sessions = result.sessions || []; + this.globalValues = result.global || {}; + if (this.form.session_id) { + const selected = this.sessions.find((item) => item.session_id === this.form.session_id); + if (selected) this.selectSession(selected); + } + } catch (err) { + this.showToast(err.message || '加载失败', 'error'); + } finally { + this.loading = false; + } + }, + + newSession() { + this.form = emptyForm(); + }, + + selectSession(session) { + this.form = { + session_id: session.session_id || '', + session_type: session.session_type || 'group', + display_name: session.display_name || '', + overrides: {...(session.overrides || {})}, + }; + }, + + hasOverride(key) { + return Object.prototype.hasOwnProperty.call(this.form.overrides, key); + }, + + toggleOverride(key, enabled) { + if (enabled) { + this.form.overrides[key] = this.effectiveValue(key); + } else { + delete this.form.overrides[key]; + } + }, + + effectiveValue(key) { + return this.hasOverride(key) ? this.form.overrides[key] : this.globalValues[key]; + }, + + displayValue(value) { + if (value === true) return '启用'; + if (value === false) return '禁用'; + if (value === '' || value === undefined || value === null) return '空'; + return String(value); + }, + + async save() { + if (!this.form.session_id.trim()) { + this.showToast('session_id 不能为空', 'error'); + return; + } + this.loading = true; + try { + const result = apiResult(await bridge.apiPost('session-config/upsert', this.form)); + this.selectSession(result.data || this.form); + await this.reload(); + this.showToast('已保存'); + } catch (err) { + this.showToast(err.message || '保存失败', 'error'); + } finally { + this.loading = false; + } + }, + + async clearAll() { + if (!this.form.session_id) return; + this.loading = true; + try { + const result = apiResult(await bridge.apiPost('session-config/clear', this.form)); + this.selectSession(result.data || this.form); + await this.reload(); + this.showToast('已清空覆盖'); + } catch (err) { + this.showToast(err.message || '清空失败', 'error'); + } finally { + this.loading = false; + } + }, + + async deleteSession() { + if (!this.form.session_id) return; + this.loading = true; + try { + apiResult(await bridge.apiPost('session-config/delete', {session_id: this.form.session_id})); + this.form = emptyForm(); + await this.reload(); + this.showToast('已删除'); + } catch (err) { + this.showToast(err.message || '删除失败', 'error'); + } finally { + this.loading = false; + } + }, + + showToast(message, type = 'success') { + if (toastTimer) window.clearTimeout(toastTimer); + this.toast = {show: true, type, message}; + toastTimer = window.setTimeout(() => { + this.toast.show = false; + }, 2600); + }, + }; + + window.sessionConfigStore = store; + PetiteVue.createApp(store).mount('#app'); + document.body.classList.add('ready'); + store.reload(); +})(); diff --git a/pages/sessionConfig/index.html b/pages/sessionConfig/index.html new file mode 100644 index 0000000..b3a0460 --- /dev/null +++ b/pages/sessionConfig/index.html @@ -0,0 +1,250 @@ + + + + + + sessionConfig + + + +
正在加载会话配置...
+
+ +
+
+
+
+

会话

+
+ + + +
+
+
+ + + +
+
+
+
+

配置覆盖

+ +
+
+
+
+ {{ item.label }} +
{{ item.key }}
+
+ +
+
全局配置
+ {{ displayValue(globalValues[item.key]) }} +
+
+
生效值
+ {{ displayValue(effectiveValue(item.key)) }} +
+
+ + + +
+
+
+
+
+
+
+ {{ toast.message }} +
+
+ + + + diff --git a/providers/base.py b/providers/base.py deleted file mode 100644 index 4f36ef9..0000000 --- a/providers/base.py +++ /dev/null @@ -1,46 +0,0 @@ -"""图片提供商基类。""" - -from __future__ import annotations - - -class SetuImageProvider: - """色图图片提供商基类(策略模式)。""" - - async def fetch_image_urls( - self, - num: int, - tags: list[str], - r18: bool, - exclude_ai: bool = True, - ) -> list[str]: - """从 API 获取图片 URL 列表。 - - 参数: - num: 要获取的图片数量。 - tags: 搜索标签/关键词。 - r18: 是否请求 R18 内容。 - exclude_ai: 是否排除 AI 生成的作品。 - - 返回: - 图片 URL 列表。 - """ - raise NotImplementedError - - @staticmethod - def _normalize_bool(value, default: bool = True) -> bool: - """Normalize possibly-dirty boolean input from config/runtime sources.""" - if value is None: - return default - if isinstance(value, bool): - return value - if isinstance(value, (int, float)): - return bool(value) - if isinstance(value, str): - lowered = value.strip().lower() - if lowered in {"true", "1", "yes", "on"}: - return True - if lowered in {"false", "0", "no", "off"}: - return False - if lowered in {"none", "null", ""}: - return default - return default diff --git a/requirements.txt b/requirements.txt index 5da2415..a6383fc 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,3 +1,5 @@ python-docx>=1.0.0 playwright>=1.45.0 httpx[http2]>=0.28.0 +pydantic>=2.0 +aiosqlite>=0.19.0 diff --git a/services/__init__.py b/services/__init__.py deleted file mode 100644 index 50383e3..0000000 --- a/services/__init__.py +++ /dev/null @@ -1,21 +0,0 @@ -"""服务层模块。 - -整合图片下载、缓存、Docx生成、HTML渲染、配置管理等服务。 -""" - -from __future__ import annotations - -from .cache import UrlImageDiskCache -from .config_manager import AccessControlManager, ConfigManager -from .docx import DocxService -from .html import HtmlCardRenderer -from .image import ImageService - -__all__ = [ - "UrlImageDiskCache", - "ConfigManager", - "AccessControlManager", - "DocxService", - "HtmlCardRenderer", - "ImageService", -] diff --git a/services/cache.py b/services/cache.py deleted file mode 100644 index d3b0881..0000000 --- a/services/cache.py +++ /dev/null @@ -1,271 +0,0 @@ -"""图片缓存服务。""" - -from __future__ import annotations - -import asyncio -import hashlib -import json -import time -from pathlib import Path -from typing import Any - -from astrbot.api import logger - - -class UrlImageDiskCache: - """基于 URL 的简单图片磁盘缓存,支持 TTL 和数量限制。""" - - def __init__( - self, cache_dir: Path, ttl_hours: int, max_items: int, enabled: bool = True - ): - self.enabled = enabled - self.cache_dir = cache_dir - self.index_path = cache_dir / "image_cache_index.json" - self.ttl_seconds = max(1, int(ttl_hours) * 3600) - self.max_items = max(1, int(max_items)) - self._index: dict[str, dict[str, Any]] = {"entries": {}, "meta": {}} - self._lock = asyncio.Lock() - - async def initialize(self, cleanup_on_start: bool = True) -> None: - """初始化缓存。""" - if not self.enabled: - logger.info("[setu.cache] cache disabled") - return - try: - self.cache_dir.mkdir(parents=True, exist_ok=True) - await self._load_index() - if cleanup_on_start: - removed = await self.cleanup_expired() - logger.info("[setu.cache] startup cleanup removed=%d", removed) - # 检查并清理孤立缓存文件(索引中不存在的 .bin 文件) - await self._cleanup_orphaned_files() - except (OSError, json.JSONDecodeError) as exc: - logger.exception("[setu.cache] initialize failed: %s", exc) - - async def _cleanup_orphaned_files(self) -> int: - """清理孤立缓存文件(索引中不存在的 .bin 文件)。""" - removed = 0 - try: - async with self._lock: - # 获取索引中记录的文件路径 - indexed_paths = { - entry.get("path", "") - for entry in self._index.get("entries", {}).values() - } - - for file_path in self.cache_dir.iterdir(): - if not file_path.is_file(): - continue - # 处理旧的 .img 文件(迁移到 .bin) - if file_path.suffix == ".img": - try: - # 尝试迁移:读取旧文件,创建新文件,删除旧文件 - new_path = file_path.with_suffix(".bin") - if not new_path.exists(): - data = file_path.read_bytes() - new_path.write_bytes(data) - # 更新索引中对应的路径 - for entry in self._index.get("entries", {}).values(): - if entry.get("path") == str(file_path): - entry["path"] = str(new_path) - break - file_path.unlink() - removed += 1 - logger.debug( - "[setu.cache] migrated .img to .bin: %s", file_path.name - ) - except OSError as exc: - logger.warning( - "[setu.cache] failed to migrate .img file: %s", exc - ) - continue - # 只处理 .bin 文件(新的扩展名) - if file_path.suffix == ".bin": - str_path = str(file_path) - if str_path not in indexed_paths: - try: - file_path.unlink() - removed += 1 - logger.debug( - "[setu.cache] removed orphaned file: %s", - file_path.name, - ) - except OSError: - pass - # 如果有迁移,保存索引 - if removed > 0: - await self._save_index_locked() - except OSError as exc: - logger.warning("[setu.cache] failed to cleanup orphaned files: %s", exc) - - if removed > 0: - logger.info( - "[setu.cache] cleaned up %d orphaned/migrated cache files", removed - ) - return removed - - async def get(self, url: str) -> bytes | None: - """获取缓存图片。""" - if not self.enabled: - return None - key = self._url_key(url) - async with self._lock: - entry = self._index.get("entries", {}).get(key) - if not entry: - return None - - now = int(time.time()) - expires_at = int(entry.get("expires_at", 0)) - cached_path = Path(entry.get("path", "")) - if expires_at <= now or not cached_path.is_file(): - self._remove_entry_locked(key, delete_file=True) - await self._save_index_locked() - return None - - try: - data = cached_path.read_bytes() - except OSError as exc: - logger.exception( - "[setu.cache] failed to read cache file key=%s: %s", key, exc - ) - self._remove_entry_locked(key, delete_file=True) - await self._save_index_locked() - return None - - entry["last_hit"] = now - await self._save_index_locked() - logger.debug("[setu.cache] hit key=%s path=%s", key, cached_path) - return data - - async def put(self, url: str, data: bytes) -> None: - """写入图片到缓存。""" - if not self.enabled or not data: - return - key = self._url_key(url) - now = int(time.time()) - expires_at = now + self.ttl_seconds - file_path = self.cache_dir / f"{key}.bin" - - async with self._lock: - try: - file_path.write_bytes(data) - except OSError as exc: - logger.exception( - "[setu.cache] failed to write cache file=%s: %s", file_path, exc - ) - return - - entries = self._index.setdefault("entries", {}) - entries[key] = { - "url": url, - "path": str(file_path), - "created_at": now, - "expires_at": expires_at, - "last_hit": now, - "size": len(data), - } - removed = self._prune_locked(now) - self._index.setdefault("meta", {})["last_prune_removed"] = removed - self._index["meta"]["last_update_at"] = now - await self._save_index_locked() - logger.debug( - "[setu.cache] put key=%s size=%d removed=%d", key, len(data), removed - ) - - async def cleanup_expired(self) -> int: - """清理过期缓存。""" - if not self.enabled: - return 0 - now = int(time.time()) - async with self._lock: - removed = self._prune_locked(now) - self._index.setdefault("meta", {})["last_cleanup_at"] = now - self._index["meta"]["last_cleanup_removed"] = removed - await self._save_index_locked() - return removed - - async def _load_index(self) -> None: - """加载缓存索引文件。""" - if not self.index_path.is_file(): - self._index = {"entries": {}, "meta": {}} - await self._save_index_locked() - return - try: - raw = self.index_path.read_text(encoding="utf-8") - loaded = json.loads(raw) - entries = loaded.get("entries", {}) if isinstance(loaded, dict) else {} - meta = loaded.get("meta", {}) if isinstance(loaded, dict) else {} - self._index = { - "entries": entries if isinstance(entries, dict) else {}, - "meta": meta if isinstance(meta, dict) else {}, - } - except Exception as exc: - logger.exception("[setu.cache] index parse failed, reset index: %s", exc) - self._index = {"entries": {}, "meta": {}} - await self._save_index_locked() - - async def _save_index_locked(self) -> None: - """保存缓存索引(加锁)。""" - tmp_path = self.index_path.with_suffix(".json.tmp") - try: - # 确保目录存在 - self.cache_dir.mkdir(parents=True, exist_ok=True) - tmp_path.write_text( - json.dumps(self._index, ensure_ascii=False, indent=2), - encoding="utf-8", - ) - tmp_path.replace(self.index_path) - except Exception as exc: - logger.exception( - "[setu.cache] failed to save index %s: %s", self.index_path, exc - ) - try: - if tmp_path.exists(): - tmp_path.unlink() - except Exception as exc2: - logger.debug("[setu.cache] failed to remove temp index file: %s", exc2) - - def _prune_locked(self, now: int) -> int: - """清理过期和超出数量限制的缓存项。""" - entries = self._index.setdefault("entries", {}) - removed = 0 - - expired_keys = [ - k for k, v in entries.items() if int(v.get("expires_at", 0)) <= now - ] - for key in expired_keys: - self._remove_entry_locked(key, delete_file=True) - removed += 1 - - if len(entries) > self.max_items: - sorted_items = sorted( - entries.items(), - key=lambda item: int( - item[1].get("last_hit", item[1].get("created_at", 0)) - ), - ) - overflow = len(entries) - self.max_items - for key, _ in sorted_items[:overflow]: - self._remove_entry_locked(key, delete_file=True) - removed += 1 - return removed - - def _remove_entry_locked(self, key: str, delete_file: bool) -> None: - """移除缓存项。""" - entries = self._index.setdefault("entries", {}) - entry = entries.pop(key, None) - if not entry or not delete_file: - return - cached_path = Path(entry.get("path", "")) - try: - if cached_path.is_file(): - cached_path.unlink() - except Exception as exc: - logger.warning( - "[setu.cache] failed to remove cache file path=%s: %s", cached_path, exc - ) - - @staticmethod - def _url_key(url: str) -> str: - """生成 URL 的哈希 key。""" - return hashlib.sha256(url.encode("utf-8", errors="ignore")).hexdigest() diff --git a/services/config_manager.py b/services/config_manager.py deleted file mode 100644 index 67876cc..0000000 --- a/services/config_manager.py +++ /dev/null @@ -1,657 +0,0 @@ -"""配置管理服务,用于持久化黑白名单等配置到文件。""" - -from __future__ import annotations - -import json -import threading -from pathlib import Path -from typing import TYPE_CHECKING, Any - -from astrbot.api import logger - -if TYPE_CHECKING: - from astrbot.core import AstrBotConfig - - -class ConfigManager: - """配置管理器,负责读写插件配置文件。 - - 同时管理两个配置存储位置: - 1. 插件数据目录的 config.json - 运行时快速访问 - 2. AstrBot 主配置 - 供 webui 显示和编辑 - """ - - SAFETY_LIST_KEYS = ( - "setu_blocked_users", - "setu_whitelist_users", - "setu_blocked_groups", - "setu_whitelist_groups", - "fortune_blocked_users", - "fortune_whitelist_users", - "fortune_blocked_groups", - "fortune_whitelist_groups", - ) - SAFETY_MODE_KEYS = ( - "setu_user_access_control_mode", - "setu_group_access_control_mode", - "fortune_user_access_control_mode", - "fortune_group_access_control_mode", - ) - LEGACY_MODE_TO_NEW = { - "setu_access_control_mode": ( - "setu_user_access_control_mode", - "setu_group_access_control_mode", - ), - "fortune_access_control_mode": ( - "fortune_user_access_control_mode", - "fortune_group_access_control_mode", - ), - } - - def __init__( - self, plugin_data_dir: Path, astrbot_config: AstrBotConfig | None = None - ): - """初始化配置管理器。 - - 参数: - plugin_data_dir: 插件数据目录 - astrbot_config: AstrBot 配置对象,用于同步到 webui - """ - self._data_dir = plugin_data_dir - self._config_file = plugin_data_dir / "config.json" - self._cache: dict[str, Any] = {} - self._astrbot_config = astrbot_config - self._main_config_cache: dict[str, Any] | None = None - self._main_config_cache_mtime: float | None = None - self._main_config_cache_path: Path | None = None - self._main_config_cache_lock = threading.Lock() - - async def initialize(self) -> None: - """初始化,加载现有配置。""" - self._load_config() - # 启动时优先吸收 WebUI 最新值,避免本地缓存反向覆盖。 - imported = self._sync_from_astrbot_config() - if not imported: - self._sync_to_astrbot_config() - - def _load_config(self) -> None: - """从文件加载配置。""" - if not self._config_file.exists(): - self._cache = {} - return - - try: - with open(self._config_file, encoding="utf-8") as f: - self._cache = json.load(f) - except (json.JSONDecodeError, OSError) as e: - logger.warning("Failed to load config file: %s", e) - self._cache = {} - - def _save_config(self) -> bool: - """保存配置到文件。 - - 返回: - 是否保存成功 - """ - try: - self._data_dir.mkdir(parents=True, exist_ok=True) - with open(self._config_file, "w", encoding="utf-8") as f: - json.dump(self._cache, f, ensure_ascii=False, indent=2) - # 同时同步到 AstrBot 配置供 webui 显示 - self._sync_to_astrbot_config() - return True - except (OSError, TypeError) as e: - logger.error("Failed to save config file: %s", e) - return False - - def _invalidate_main_config_cache(self) -> None: - """失效主配置缓存。""" - with self._main_config_cache_lock: - self._main_config_cache = None - self._main_config_cache_mtime = None - self._main_config_cache_path = None - - def _iter_main_config_candidates(self) -> list[Path]: - """构造可能的主配置文件路径候选。""" - config_file_name = "astrbot_plugin_setu_config.json" - candidates: list[Path] = [] - # 常见结构:data/plugins/ - candidates.append(self._data_dir.parent.parent / "config" / config_file_name) - # 兜底:遍历父目录尝试 data/config 路径 - candidates.extend( - parent / "data" / "config" / config_file_name - for parent in self._data_dir.parents - ) - return candidates - - def _load_main_config_with_cache(self) -> dict[str, Any] | None: - """加载主配置并基于 mtime 进行缓存。""" - config_path: Path | None = None - for candidate in self._iter_main_config_candidates(): - if candidate.is_file(): - config_path = candidate - break - - if config_path is None: - return None - - try: - mtime = config_path.stat().st_mtime - except OSError: - return None - - with self._main_config_cache_lock: - if ( - self._main_config_cache_path == config_path - and self._main_config_cache_mtime == mtime - and self._main_config_cache is not None - ): - return self._main_config_cache - - try: - # AstrBot 主配置在 Windows 上可能带 BOM,使用 utf-8-sig 兼容读取。 - with open(config_path, encoding="utf-8-sig") as f: - data = json.load(f) - except (OSError, json.JSONDecodeError) as e: - logger.warning("Failed to load main config %s: %s", config_path, e) - return None - - if not isinstance(data, dict): - return None - - self._main_config_cache = data - self._main_config_cache_mtime = mtime - self._main_config_cache_path = config_path - return data - - def _get_safety_section(self) -> dict[str, Any] | None: - """从主配置中获取 safety 配置段(带缓存)。""" - main_cfg = self._load_main_config_with_cache() - if not main_cfg: - return None - safety = main_cfg.get("safety") - if not isinstance(safety, dict): - return None - return safety - - def _get_safety_value_from_file(self, key: str, default: Any = None) -> Any: - """从主配置文件读取单个 safety 配置值(带缓存)。""" - safety_config = self._get_safety_section() - if safety_config is None: - return default - - # 兼容旧键:请求新键时可回退到旧键。 - if key not in safety_config: - if ( - key - in ( - "setu_user_access_control_mode", - "setu_group_access_control_mode", - ) - and "setu_access_control_mode" in safety_config - ): - value = safety_config["setu_access_control_mode"] - elif ( - key - in ( - "fortune_user_access_control_mode", - "fortune_group_access_control_mode", - ) - and "fortune_access_control_mode" in safety_config - ): - value = safety_config["fortune_access_control_mode"] - else: - return default - else: - value = safety_config[key] - - if key in self.SAFETY_LIST_KEYS: - if not isinstance(value, list): - return [] - return [str(v).strip() for v in value if str(v).strip()] - - if key in self.SAFETY_MODE_KEYS: - return value if value in {"none", "blacklist", "whitelist"} else default - - return value - - def _sync_to_astrbot_config(self) -> None: - """将本地配置同步到 AstrBot 配置供 webui 显示。 - - 配置在 schema 中定义于顶层 safety 下,webui 期望从该路径读取。 - """ - if self._astrbot_config is None: - return - - try: - updated = False - - # 确保顶层 safety 存在(webui 配置路径) - if "safety" not in self._astrbot_config: - self._astrbot_config["safety"] = {} - updated = True - - safety_config = self._astrbot_config["safety"] - if not isinstance(safety_config, dict): - safety_config = {} - self._astrbot_config["safety"] = safety_config - updated = True - - # 仅同步缓存中明确存在的键,避免启动时空缓存覆盖 WebUI。 - for key in self.SAFETY_LIST_KEYS: - if key not in self._cache: - continue - value = self._cache.get(key) - if not isinstance(value, list): - value = [] - if safety_config.get(key) != value: - safety_config[key] = value - updated = True - - for key in self.SAFETY_MODE_KEYS: - if key not in self._cache: - continue - value = self._cache.get(key) - if value is not None and safety_config.get(key) != value: - safety_config[key] = value - updated = True - - if ( - updated - and hasattr(self._astrbot_config, "save_config") - and callable(getattr(self._astrbot_config, "save_config")) - ): - self._astrbot_config.save_config() - self._invalidate_main_config_cache() - logger.debug( - "[config_manager] Synced config to AstrBot safety (top-level) and saved" - ) - - except Exception as e: - logger.debug("[config_manager] Failed to sync to AstrBot config: %s", e) - - def _sync_from_astrbot_config(self) -> bool: - """从 AstrBot 配置同步到本地(webui 修改后)。 - - 从顶层 safety 同步黑白名单列表和访问控制模式到本地缓存。 - 返回是否成功导入了任意安全配置键。 - """ - if self._astrbot_config is None: - return False - - try: - imported = False - updated = False - - safety_config = self._astrbot_config.get("safety", {}) - if not isinstance(safety_config, dict): - return False - - for key in self.SAFETY_LIST_KEYS: - if key not in safety_config: - continue - imported = True - value = safety_config.get(key) - if not isinstance(value, list): - value = [] - value = [str(v).strip() for v in value if str(v).strip()] - if self._cache.get(key) != value: - self._cache[key] = value - updated = True - - valid_modes = {"none", "blacklist", "whitelist"} - # 优先读取新键 - for key in self.SAFETY_MODE_KEYS: - if key not in safety_config: - continue - imported = True - value = safety_config.get(key) - if value not in valid_modes: - continue - if self._cache.get(key) != value: - self._cache[key] = value - updated = True - - # 若新键缺失,尝试旧键映射。 - for legacy_key, mapped_keys in self.LEGACY_MODE_TO_NEW.items(): - if legacy_key not in safety_config: - continue - imported = True - value = safety_config.get(legacy_key) - if value not in valid_modes: - continue - for key in mapped_keys: - if key not in safety_config and self._cache.get(key) != value: - self._cache[key] = value - updated = True - - if updated: - self._data_dir.mkdir(parents=True, exist_ok=True) - with open(self._config_file, "w", encoding="utf-8") as f: - json.dump(self._cache, f, ensure_ascii=False, indent=2) - logger.debug( - "[config_manager] Synced config from AstrBot/webui safety (top-level)" - ) - - return imported - - except Exception as e: - logger.debug("[config_manager] Failed to sync from AstrBot config: %s", e) - return False - - def get(self, key: str, default: Any = None) -> Any: - """获取配置值。""" - if key in self.SAFETY_LIST_KEYS or key in self.SAFETY_MODE_KEYS: - file_value = self._get_safety_value_from_file(key, None) - if file_value is not None: - return file_value - - if self._astrbot_config is not None: - try: - safety_config = self._astrbot_config.get("safety", {}) - if isinstance(safety_config, dict) and key in safety_config: - return safety_config[key] - except Exception: - pass - - if key in self._cache: - return self._cache[key] - return default - - if key in self._cache: - return self._cache[key] - - return default - - def _get_access_control_mode_from_file(self, key: str, default: Any = None) -> Any: - """兼容方法:访问控制模式读取统一复用 safety 读取路径。""" - return self._get_safety_value_from_file(key, default) - - def set(self, key: str, value: Any) -> bool: - """设置配置值。""" - self._cache[key] = value - return self._save_config() - - def get_list(self, key: str) -> list[str]: - """获取列表类型的配置值。""" - value = self.get(key, []) - if isinstance(value, list): - return [str(v).strip() for v in value if str(v).strip()] - return [] - - def add_to_list(self, key: str, item: str) -> bool: - """添加项目到列表。""" - current = self.get_list(key) - item_str = str(item).strip() - - if not item_str: - return False - if item_str in current: - return True - - current.append(item_str) - self._cache[key] = current - return self._save_config() - - def remove_from_list(self, key: str, item: str) -> bool: - """从列表中移除项目。""" - current = self.get_list(key) - item_str = str(item).strip() - - if item_str not in current: - return True - - current.remove(item_str) - self._cache[key] = current - return self._save_config() - - def is_in_list(self, key: str, item: str) -> bool: - """检查项目是否在列表中。""" - current = self.get_list(key) - return str(item).strip() in current - - -class AccessControlManager: - """访问控制管理器,封装黑白名单管理逻辑。""" - - KEY_SETU_ACCESS_CONTROL_MODE = "setu_access_control_mode" - KEY_SETU_BLOCKED_USERS = "setu_blocked_users" - KEY_SETU_WHITELIST_USERS = "setu_whitelist_users" - KEY_SETU_BLOCKED_GROUPS = "setu_blocked_groups" - KEY_SETU_WHITELIST_GROUPS = "setu_whitelist_groups" - - KEY_FORTUNE_ACCESS_CONTROL_MODE = "fortune_access_control_mode" - KEY_FORTUNE_BLOCKED_USERS = "fortune_blocked_users" - KEY_FORTUNE_WHITELIST_USERS = "fortune_whitelist_users" - KEY_FORTUNE_BLOCKED_GROUPS = "fortune_blocked_groups" - KEY_FORTUNE_WHITELIST_GROUPS = "fortune_whitelist_groups" - - def __init__(self, config_manager: ConfigManager): - self._cfg = config_manager - - def add_setu_blocked_user(self, user_id: str) -> bool: - user_id = str(user_id).strip() - if not user_id: - return False - # 避免同一用户同时存在于黑白名单。 - self._cfg.remove_from_list(self.KEY_SETU_WHITELIST_USERS, user_id) - return self._cfg.add_to_list(self.KEY_SETU_BLOCKED_USERS, user_id) - - def remove_setu_blocked_user(self, user_id: str) -> bool: - return self._cfg.remove_from_list(self.KEY_SETU_BLOCKED_USERS, user_id) - - def is_setu_user_blocked(self, user_id: str) -> bool: - return self._cfg.is_in_list(self.KEY_SETU_BLOCKED_USERS, user_id) - - def get_setu_blocked_users(self) -> list[str]: - return self._cfg.get_list(self.KEY_SETU_BLOCKED_USERS) - - def add_setu_whitelist_user(self, user_id: str) -> bool: - user_id = str(user_id).strip() - if not user_id: - return False - # 被信任时自动从黑名单移除。 - self._cfg.remove_from_list(self.KEY_SETU_BLOCKED_USERS, user_id) - return self._cfg.add_to_list(self.KEY_SETU_WHITELIST_USERS, user_id) - - def remove_setu_whitelist_user(self, user_id: str) -> bool: - return self._cfg.remove_from_list(self.KEY_SETU_WHITELIST_USERS, user_id) - - def is_setu_user_whitelisted(self, user_id: str) -> bool: - return self._cfg.is_in_list(self.KEY_SETU_WHITELIST_USERS, user_id) - - def get_setu_whitelist_users(self) -> list[str]: - return self._cfg.get_list(self.KEY_SETU_WHITELIST_USERS) - - def add_setu_blocked_group(self, group_id: str) -> bool: - return self._cfg.add_to_list(self.KEY_SETU_BLOCKED_GROUPS, group_id) - - def remove_setu_blocked_group(self, group_id: str) -> bool: - return self._cfg.remove_from_list(self.KEY_SETU_BLOCKED_GROUPS, group_id) - - def is_setu_group_blocked(self, group_id: str) -> bool: - return self._cfg.is_in_list(self.KEY_SETU_BLOCKED_GROUPS, group_id) - - def get_setu_blocked_groups(self) -> list[str]: - return self._cfg.get_list(self.KEY_SETU_BLOCKED_GROUPS) - - def add_setu_whitelist_group(self, group_id: str) -> bool: - return self._cfg.add_to_list(self.KEY_SETU_WHITELIST_GROUPS, group_id) - - def remove_setu_whitelist_group(self, group_id: str) -> bool: - return self._cfg.remove_from_list(self.KEY_SETU_WHITELIST_GROUPS, group_id) - - def is_setu_group_whitelisted(self, group_id: str) -> bool: - return self._cfg.is_in_list(self.KEY_SETU_WHITELIST_GROUPS, group_id) - - def get_setu_whitelist_groups(self) -> list[str]: - return self._cfg.get_list(self.KEY_SETU_WHITELIST_GROUPS) - - def add_fortune_blocked_user(self, user_id: str) -> bool: - user_id = str(user_id).strip() - if not user_id: - return False - self._cfg.remove_from_list(self.KEY_FORTUNE_WHITELIST_USERS, user_id) - return self._cfg.add_to_list(self.KEY_FORTUNE_BLOCKED_USERS, user_id) - - def remove_fortune_blocked_user(self, user_id: str) -> bool: - return self._cfg.remove_from_list(self.KEY_FORTUNE_BLOCKED_USERS, user_id) - - def is_fortune_user_blocked(self, user_id: str) -> bool: - return self._cfg.is_in_list(self.KEY_FORTUNE_BLOCKED_USERS, user_id) - - def get_fortune_blocked_users(self) -> list[str]: - return self._cfg.get_list(self.KEY_FORTUNE_BLOCKED_USERS) - - def add_fortune_whitelist_user(self, user_id: str) -> bool: - user_id = str(user_id).strip() - if not user_id: - return False - self._cfg.remove_from_list(self.KEY_FORTUNE_BLOCKED_USERS, user_id) - return self._cfg.add_to_list(self.KEY_FORTUNE_WHITELIST_USERS, user_id) - - def remove_fortune_whitelist_user(self, user_id: str) -> bool: - return self._cfg.remove_from_list(self.KEY_FORTUNE_WHITELIST_USERS, user_id) - - def is_fortune_user_whitelisted(self, user_id: str) -> bool: - return self._cfg.is_in_list(self.KEY_FORTUNE_WHITELIST_USERS, user_id) - - def get_fortune_whitelist_users(self) -> list[str]: - return self._cfg.get_list(self.KEY_FORTUNE_WHITELIST_USERS) - - def add_fortune_blocked_group(self, group_id: str) -> bool: - return self._cfg.add_to_list(self.KEY_FORTUNE_BLOCKED_GROUPS, group_id) - - def remove_fortune_blocked_group(self, group_id: str) -> bool: - return self._cfg.remove_from_list(self.KEY_FORTUNE_BLOCKED_GROUPS, group_id) - - def is_fortune_group_blocked(self, group_id: str) -> bool: - return self._cfg.is_in_list(self.KEY_FORTUNE_BLOCKED_GROUPS, group_id) - - def get_fortune_blocked_groups(self) -> list[str]: - return self._cfg.get_list(self.KEY_FORTUNE_BLOCKED_GROUPS) - - def add_fortune_whitelist_group(self, group_id: str) -> bool: - return self._cfg.add_to_list(self.KEY_FORTUNE_WHITELIST_GROUPS, group_id) - - def remove_fortune_whitelist_group(self, group_id: str) -> bool: - return self._cfg.remove_from_list(self.KEY_FORTUNE_WHITELIST_GROUPS, group_id) - - def is_fortune_group_whitelisted(self, group_id: str) -> bool: - return self._cfg.is_in_list(self.KEY_FORTUNE_WHITELIST_GROUPS, group_id) - - def get_fortune_whitelist_groups(self) -> list[str]: - return self._cfg.get_list(self.KEY_FORTUNE_WHITELIST_GROUPS) - - def check_setu_access( - self, - user_id: str | None, - group_id: str | None, - user_access_control_mode: str = "none", - group_access_control_mode: str = "none", - ) -> tuple[bool, str]: - """检查色图功能访问权限。 - - 参数: - user_id: 用户 ID - group_id: 群组 ID - user_access_control_mode: 用户访问控制模式 (none/blacklist/whitelist) - group_access_control_mode: 群组访问控制模式 (none/blacklist/whitelist) - - 返回: - (是否被屏蔽, 屏蔽原因) - """ - if user_id is not None: - uid = str(user_id) - if user_access_control_mode == "blacklist": - is_blocked = self.is_setu_user_blocked(uid) - blocked_users = self.get_setu_blocked_users() - logger.debug( - "[check_setu_access] User blacklist mode: user=%s, blocked=%s, blocked_users=%s", - uid, - is_blocked, - blocked_users, - ) - if is_blocked: - return True, "用户被禁用" - elif ( - user_access_control_mode == "whitelist" - and not self.is_setu_user_whitelisted(uid) - ): - return True, "用户不在白名单中" - - if group_id is not None: - gid = str(group_id) - if group_access_control_mode == "blacklist" and self.is_setu_group_blocked( - gid - ): - return True, "群组被禁用" - if ( - group_access_control_mode == "whitelist" - and not self.is_setu_group_whitelisted(gid) - ): - return True, "群组不在白名单中" - - return False, "" - - def check_fortune_access( - self, - user_id: str | None, - group_id: str | None, - user_access_control_mode: str = "none", - group_access_control_mode: str = "none", - ) -> tuple[bool, str]: - """检查运势功能访问权限。 - - 参数: - user_id: 用户 ID - group_id: 群组 ID - user_access_control_mode: 用户访问控制模式 (none/blacklist/whitelist) - group_access_control_mode: 群组访问控制模式 (none/blacklist/whitelist) - - 返回: - (是否被屏蔽, 屏蔽原因) - """ - if user_id is not None: - uid = str(user_id) - if user_access_control_mode == "blacklist" and self.is_fortune_user_blocked( - uid - ): - return True, "用户被禁用" - if ( - user_access_control_mode == "whitelist" - and not self.is_fortune_user_whitelisted(uid) - ): - return True, "用户不在白名单中" - - if group_id is not None: - gid = str(group_id) - if ( - group_access_control_mode == "blacklist" - and self.is_fortune_group_blocked(gid) - ): - return True, "群组被禁用" - if ( - group_access_control_mode == "whitelist" - and not self.is_fortune_group_whitelisted(gid) - ): - return True, "群组不在白名单中" - - return False, "" - - def get_all_lists(self) -> dict[str, list[str]]: - """获取所有黑白名单列表。 - - 返回: - 包含所有列表的字典 - """ - return { - "setu_blocked_users": self.get_setu_blocked_users(), - "setu_whitelist_users": self.get_setu_whitelist_users(), - "setu_blocked_groups": self.get_setu_blocked_groups(), - "setu_whitelist_groups": self.get_setu_whitelist_groups(), - "fortune_blocked_users": self.get_fortune_blocked_users(), - "fortune_whitelist_users": self.get_fortune_whitelist_users(), - "fortune_blocked_groups": self.get_fortune_blocked_groups(), - "fortune_whitelist_groups": self.get_fortune_whitelist_groups(), - } diff --git a/services/docx.py b/services/docx.py deleted file mode 100644 index ac6c97e..0000000 --- a/services/docx.py +++ /dev/null @@ -1,116 +0,0 @@ -"""Docx 文件生成服务。""" - -from __future__ import annotations - -import io -import os -import re -from pathlib import Path - -from astrbot.api import logger - - -class DocxService: - """Docx 文件生成服务。""" - - def __init__(self): - self._has_dependency = self._check_dependency() - - async def initialize(self) -> None: - """初始化 Docx 服务(统一初始化接口占位)。 - - 实际依赖检查已在 __init__ 完成,此方法供 SetuCore 统一调用。 - """ - pass - - @staticmethod - def _check_dependency() -> bool: - """检查 python-docx 是否已安装。""" - try: - import importlib.util - - return importlib.util.find_spec("docx") is not None - except ImportError: - logger.error("python-docx 未安装,请运行: pip install python-docx") - return False - - def create_docx_with_images( - self, - images: list[bytes], - output_path: Path | None = None, - tags: list[str] | None = None, - ) -> Path | None: - """创建包含图片的 Docx 文件。""" - if not self._has_dependency: - logger.error("python-docx 未安装,无法生成 Docx 文件") - return None - - try: - from docx import Document - from docx.enum.text import WD_ALIGN_PARAGRAPH - from docx.shared import Inches - except ImportError: - logger.error("python-docx 未安装,无法生成 Docx 文件") - return None - - try: - doc = Document() - - sections = doc.sections - for section in sections: - section.top_margin = Inches(0.5) - section.bottom_margin = Inches(0.5) - section.left_margin = Inches(0.5) - section.right_margin = Inches(0.5) - - for i, img_data in enumerate(images): - img_stream = io.BytesIO(img_data) - paragraph = doc.add_paragraph() - paragraph.alignment = WD_ALIGN_PARAGRAPH.CENTER - - run = paragraph.add_run() - run.add_picture(img_stream, width=Inches(6.0)) - - if i < len(images) - 1: - doc.add_paragraph() - - if output_path is None: - from astrbot.core.utils.astrbot_path import get_astrbot_temp_path - - temp_dir = get_astrbot_temp_path() - if isinstance(temp_dir, str): - temp_dir = Path(temp_dir) - - if tags: - clean_tags = [self._sanitize_filename(t) for t in tags[:3]] - tag_str = ",".join(clean_tags) - else: - tag_str = "setu" - - random_suffix = os.urandom(6).hex() - filename = f"{tag_str}_{random_suffix}.docx" - output_path = temp_dir / filename - - doc.save(str(output_path)) - logger.info("Docx 文件已生成: %s", output_path) - return output_path - - except (OSError, ValueError, RuntimeError) as e: - logger.error("生成 Docx 文件失败: %s", e) - return None - - @staticmethod - def _sanitize_filename(text: str) -> str: - """清理文件名中的非法字符。""" - illegal_chars = r'[<>:"/\\|?*\x00-\x1f]' - sanitized = re.sub(illegal_chars, "", text) - - if len(sanitized) > 30: - sanitized = sanitized[:30] - - sanitized = sanitized.strip(" .") - - if not sanitized: - return "tag" - - return sanitized diff --git a/services/html.py b/services/html.py deleted file mode 100644 index 9e33422..0000000 --- a/services/html.py +++ /dev/null @@ -1,170 +0,0 @@ -"""HTML card renderer for safer image wrapping.""" - -from __future__ import annotations - -import asyncio -import base64 -import random -from pathlib import Path -from typing import Any - -from astrbot.api import logger - - -class HtmlCardRenderer: - """Wrap image(s) into compact HTML cards then render with AstrBot html_render.""" - - BG_COLORS = [ - "#f0f2f5", - "#f5f0f0", - "#f0f5f2", - "#f2f0f5", - "#faf8f5", - "#f5f8fa", - "#f8f5fa", - "#f5faf8", - "#fff5f0", - "#f0fff5", - "#f5f0ff", - "#fffff0", - ] - BORDER_COLORS = [ - "#e0e2e5", - "#e5e0e0", - "#e0e5e2", - "#e2e0e5", - "#d0d2d5", - "#d5d0d0", - "#d0d5d2", - "#d2d0d5", - ] - ROTATION_RANGE = (-2.0, 2.0) - - def __init__(self, template_path: Path | None = None): - self.template_path = ( - template_path or Path(__file__).parent.parent / "templates" / "main.html" - ) - self._template: str | None = None - - def _load_template(self) -> str: - if self._template is None: - self._template = self.template_path.read_text(encoding="utf-8") - return self._template - - def _generate_random_styles( - self, style_options: dict[str, Any] | None = None - ) -> dict[str, Any]: - style_options = style_options or {} - card_padding = int(style_options.get("card_padding", 6)) - card_gap = int(style_options.get("card_gap", 6)) - return { - "bg_color": random.choice(self.BG_COLORS), - "card_bg": "#ffffff", - "border_color": random.choice(self.BORDER_COLORS), - "rotation": random.uniform(*self.ROTATION_RANGE), - "noise_opacity": random.uniform(0.04, 0.10), - "grid_opacity": random.uniform(0.01, 0.04), - "border_radius": random.randint(8, 14), - "card_padding": max(2, card_padding), - "card_gap": max(2, card_gap), - "page_padding": 4, - } - - def _build_html( - self, images_b64: list[str], style_options: dict[str, Any] | None = None - ) -> str: - template = self._load_template() - styles = self._generate_random_styles(style_options) - - cards_html = [] - for index, img_b64 in enumerate(images_b64): - card_rotation = styles["rotation"] + random.uniform(-0.4, 0.4) - card_style = ( - f"background:{styles['card_bg']};" - f"border-radius:{styles['border_radius']}px;" - f"padding:{styles['card_padding']}px;" - f"margin-bottom:{styles['card_gap']}px;" - "box-shadow:0 1px 3px rgba(0,0,0,0.05);" - f"transform:rotate({card_rotation:.2f}deg);" - f"border:1px solid {styles['border_color']};" - "display:inline-block;" - ) - cards_html.append(f""" -
-
- img{index + 1} -
- """) - - body_style = ( - "margin:0;" - f"padding:{styles['page_padding']}px;" - f"background-color:{styles['bg_color']};" - "font-family:system-ui,-apple-system,sans-serif;" - "display:inline-block;" - ) - html = ( - template.replace("{cards}", "\n".join(cards_html)) - .replace("{body_style}", body_style) - .replace("{noise_opacity}", str(styles["noise_opacity"])) - .replace("{grid_opacity}", str(styles["grid_opacity"])) - ) - return html - - async def render_single_image( - self, context, image: bytes, style_options: dict[str, Any] | None = None - ) -> bytes | None: - """渲染单张图片,返回渲染后的图片字节数据。""" - if not image or not context: - return None - try: - img_b64 = base64.b64encode(image).decode("ascii") - html_content = self._build_html([img_b64], style_options=style_options) - render_options = {"full_page": True, "type": "png", "scale": "device"} - # 使用 return_url=False 获取文件路径,然后在线程池中读取字节数据,避免阻塞事件循环 - file_path = await context.html_render( - tmpl=html_content, data={}, return_url=False, options=render_options - ) - if file_path: - return await asyncio.to_thread(Path(file_path).read_bytes) - return None - except (ValueError, RuntimeError, OSError): - logger.exception("html single-card render failed") - return None - - async def render_images( - self, - context, - images: list[bytes], - options: dict[str, Any] | None = None, - style_options: dict[str, Any] | None = None, - ) -> bytes | None: - """渲染多张图片,返回渲染后的图片字节数据。""" - if not images or not context: - return None - try: - images_b64 = [base64.b64encode(img).decode("ascii") for img in images] - html_content = self._build_html(images_b64, style_options=style_options) - render_options = {"full_page": True, "type": "png", "scale": "device"} - if options: - render_options.update(options) - # 使用 return_url=False 获取文件路径,然后在线程池中读取字节数据,避免阻塞事件循环 - file_path = await context.html_render( - tmpl=html_content, data={}, return_url=False, options=render_options - ) - if file_path: - return await asyncio.to_thread(Path(file_path).read_bytes) - return None - except (ValueError, RuntimeError, OSError): - logger.exception("html card render failed") - return None - - def render_to_html_file(self, images: list[bytes], output_path: Path) -> bool: - """渲染为 HTML 文件(调试用)。""" - try: - images_b64 = [base64.b64encode(img).decode("ascii") for img in images] - output_path.write_text(self._build_html(images_b64), encoding="utf-8") - return True - except (OSError, ValueError): - logger.exception("html debug file save failed") - return False diff --git a/services/image.py b/services/image.py deleted file mode 100644 index 1b1fe72..0000000 --- a/services/image.py +++ /dev/null @@ -1,451 +0,0 @@ -"""图片下载/发送服务,支持 httpx 和 range 分段下载。""" - -from __future__ import annotations - -import asyncio -from typing import TYPE_CHECKING - -import httpx - -from astrbot.api import logger -from astrbot.api.event import AstrMessageEvent -from astrbot.api.message_components import Image, Node, Nodes, Plain - -from ..utils import obfuscate_image_bytes - -if TYPE_CHECKING: - from .cache import UrlImageDiskCache - -DEFAULT_HEADERS = { - "User-Agent": ( - "Mozilla/5.0 (Windows NT 10.0; Win64; x64) " - "AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36" - ), - "Accept": "image/avif,image/webp,image/apng,image/svg+xml,image/*,*/*;q=0.8", - "Accept-Language": "zh-CN,zh;q=0.9,en;q=0.8", - "Referer": "https://www.pixiv.net/", -} - -MAX_DOWNLOAD_SIZE_BYTES = 50 * 1024 * 1024 -CHUNK_SIZE = 65536 - - -class ImageService: - """图片下载/发送服务,支持 httpx 和 range 分段下载。""" - - def __init__( - self, - cache: UrlImageDiskCache | None = None, - concurrent_limit: int = 10, - timeout_seconds: int = 30, - tcp_connector_limit: int = 100, - tcp_connector_limit_per_host: int = 30, - enable_range_download: bool = False, - range_segments: int = 3, - range_threshold: int = 512, # KB - ): - self._cache = cache - self._download_semaphore = asyncio.Semaphore(max(1, concurrent_limit)) - self._timeout_seconds = max(10, timeout_seconds) - self._tcp_connector_limit = max(50, tcp_connector_limit) - self._tcp_connector_limit_per_host = max(20, tcp_connector_limit_per_host) - self._enable_range_download = enable_range_download - self._range_segments = max(2, min(8, range_segments)) - self._range_threshold = range_threshold * 1024 # Convert to bytes - - # httpx client (will be initialized lazily) - self._httpx_client: httpx.AsyncClient | None = None - - def _get_headers_for_url(self, url: str) -> dict[str, str]: - """根据 URL 获取合适的 headers。""" - headers = dict(DEFAULT_HEADERS) - if "i.pixiv.re" in url: - headers["Referer"] = "https://www.pixiv.net/" - return headers - - async def _get_httpx_client(self) -> httpx.AsyncClient: - """获取或创建 httpx 客户端。""" - if self._httpx_client is None: - limits = httpx.Limits( - max_connections=self._tcp_connector_limit, - max_keepalive_connections=self._tcp_connector_limit_per_host, - ) - timeout = httpx.Timeout( - connect=10.0, - read=self._timeout_seconds, - write=10.0, - pool=5.0, - ) - self._httpx_client = httpx.AsyncClient( - limits=limits, - timeout=timeout, - headers=DEFAULT_HEADERS, - http2=True, - follow_redirects=True, - ) - return self._httpx_client - - async def close(self): - """关闭所有客户端连接。""" - if self._httpx_client: - await self._httpx_client.aclose() - self._httpx_client = None - - async def _download_with_httpx(self, url: str, retry: int = 1) -> bytes | None: - """使用 httpx 下载单张图片,使用流式下载并强制大小限制。""" - client = await self._get_httpx_client() - headers = self._get_headers_for_url(url) - - for attempt in range(retry + 1): - try: - async with self._download_semaphore: - async with client.stream("GET", url, headers=headers) as response: - if response.status_code == 404: - logger.warning("image 404: %s", url) - return None - if response.status_code in (403, 401): - logger.warning( - "image access denied (%d): %s", - response.status_code, - url, - ) - return None - if response.status_code == 429: - logger.warning("rate limited: %s", url) - if attempt < retry: - await asyncio.sleep(1) - continue - return None - if response.status_code < 200 or response.status_code >= 300: - logger.warning( - "image download failed (%d): %s", - response.status_code, - url, - ) - if attempt < retry: - await asyncio.sleep(0.5) - continue - return None - - # Check Content-Length header if available - content_length = response.headers.get("content-length") - if content_length: - try: - total_size = int(content_length) - if total_size > MAX_DOWNLOAD_SIZE_BYTES: - logger.warning( - "image too large (%d > %d): %s", - total_size, - MAX_DOWNLOAD_SIZE_BYTES, - url, - ) - return None - except (ValueError, TypeError): - pass - - # Stream download with size enforcement - chunks: list[bytes] = [] - total_read = 0 - async for chunk in response.aiter_bytes(CHUNK_SIZE): - total_read += len(chunk) - if total_read > MAX_DOWNLOAD_SIZE_BYTES: - logger.warning( - "image download exceeds size limit: %s", url - ) - return None - chunks.append(chunk) - - data = b"".join(chunks) if chunks else b"" - if not data: - logger.warning("image download empty: %s", url) - return None - - return data - - except httpx.TimeoutException as exc: - logger.warning( - "httpx timeout (attempt %d/%d) url=%s: %s", - attempt + 1, - retry + 1, - url, - exc, - ) - if attempt < retry: - await asyncio.sleep(0.5) - continue - return None - except httpx.ConnectError as exc: - logger.warning( - "httpx connection error (attempt %d/%d) url=%s: %s", - attempt + 1, - retry + 1, - url, - exc, - ) - if attempt < retry: - await asyncio.sleep(0.5) - continue - return None - except Exception as exc: - logger.warning( - "httpx download error (attempt %d/%d) url=%s: %s", - attempt + 1, - retry + 1, - url, - exc, - exc_info=True, - ) - if attempt < retry: - await asyncio.sleep(0.5) - continue - return None - - return None - - async def _get_content_length(self, url: str) -> int | None: - """通过 HEAD 请求获取 Content-Length。""" - client = await self._get_httpx_client() - headers = self._get_headers_for_url(url) - - try: - async with self._download_semaphore: - response = await client.head( - url, headers=headers, follow_redirects=True - ) - if response.status_code >= 200 and response.status_code < 300: - content_length = response.headers.get("content-length") - if content_length: - return int(content_length) - except Exception as exc: - logger.debug("HEAD request failed for %s: %s", url, exc) - - return None - - async def _download_range_httpx( - self, url: str, start: int, end: int - ) -> bytes | None: - """使用 httpx 下载指定 range,并验证返回数据大小。""" - client = await self._get_httpx_client() - headers = self._get_headers_for_url(url) - headers["Range"] = f"bytes={start}-{end}" - expected_length = end - start + 1 - - try: - async with self._download_semaphore: - async with client.stream("GET", url, headers=headers) as response: - if response.status_code not in (200, 206): - logger.warning( - "range download failed (%d): %s", response.status_code, url - ) - return None - - # Stream the response and validate size - chunks: list[bytes] = [] - total_read = 0 - async for chunk in response.aiter_bytes(CHUNK_SIZE): - total_read += len(chunk) - # Each range segment should not exceed expected length by much - if total_read > expected_length + CHUNK_SIZE: - logger.warning( - "range download exceeded expected size for %s: expected %d, got >%d", - url, - expected_length, - total_read, - ) - return None - chunks.append(chunk) - - data = b"".join(chunks) if chunks else b"" - if not data: - logger.warning("range download returned empty body: %s", url) - return None - - actual_length = len(data) - # Allow some tolerance for servers that return slightly different sizes - if ( - actual_length != expected_length - and abs(actual_length - expected_length) > 1 - ): - logger.warning( - "range download size mismatch for %s: expected %d bytes, got %d", - url, - expected_length, - actual_length, - ) - return None - - return data - - except Exception as exc: - logger.warning("range download error for %s: %s", url, exc) - return None - - async def _download_single_with_range( - self, url: str, retry: int = 1 - ) -> bytes | None: - """使用 range 分段下载单张图片。""" - # 先尝试获取 Content-Length - total_size = await self._get_content_length(url) - - if total_size is None: - # 服务器不支持 HEAD 或没有 Content-Length,使用普通下载 - logger.debug("Content-Length not available, using normal download: %s", url) - return await self._download_with_httpx(url, retry) - - if total_size > MAX_DOWNLOAD_SIZE_BYTES: - logger.warning( - "image too large (%d > %d): %s", - total_size, - MAX_DOWNLOAD_SIZE_BYTES, - url, - ) - return None - - if total_size < self._range_threshold: - # 图片太小,不需要分段 - logger.debug( - "Image too small (%d bytes), using normal download: %s", total_size, url - ) - return await self._download_with_httpx(url, retry) - - # 计算每段大小 - segment_size = total_size // self._range_segments - ranges = [] - - for i in range(self._range_segments): - start = i * segment_size - if i == self._range_segments - 1: - end = total_size - 1 - else: - end = start + segment_size - 1 - ranges.append((start, end)) - - logger.debug("Downloading %s with %d ranges: %s", url, len(ranges), ranges) - - # 并行下载所有段 - tasks = [self._download_range_httpx(url, start, end) for start, end in ranges] - results = await asyncio.gather(*tasks) - - # 检查是否有失败的段 - if any(r is None for r in results): - logger.warning( - "Some range segments failed for %s, retrying normal download", url - ) - return await self._download_with_httpx(url, retry) - - # 合并所有段 - data = b"".join(results) - - if len(data) != total_size: - logger.warning( - "Range download size mismatch (%d vs %d) for %s", - len(data), - total_size, - url, - ) - return await self._download_with_httpx(url, retry) - - return data - - async def download_single(self, url: str, retry: int = 1) -> bytes | None: - """下载单张图片,带缓存和重试机制。""" - if not url: - return None - - # 检查缓存 - if self._cache: - try: - cached = await self._cache.get(url) - if cached: - return cached - except Exception as exc: - logger.exception("[setu.cache] read failed url=%s : %s", url, exc) - - # 选择下载方式 - data: bytes | None = None - - if self._enable_range_download: - # 使用 httpx + range 下载 - data = await self._download_single_with_range(url, retry) - else: - # 使用 httpx 普通下载 - data = await self._download_with_httpx(url, retry) - - # 写入缓存 - if data and self._cache: - try: - await self._cache.put(url, data) - except Exception as exc: - logger.exception("[setu.cache] write failed url=%s : %s", url, exc) - - return data - - async def download_parallel(self, urls: list[str]) -> list[bytes]: - """并发下载多张图片。""" - if not urls: - return [] - - tasks = [self.download_single(url, retry=2) for url in urls] - results = await asyncio.gather(*tasks, return_exceptions=True) - - downloaded: list[bytes] = [] - for i, result in enumerate(results): - if isinstance(result, bytes) and result: - downloaded.append(result) - elif isinstance(result, Exception): - logger.warning( - "download task failed for %s: %s", - urls[i] if i < len(urls) else "unknown", - result, - ) - - return downloaded - - @staticmethod - async def send_images( - event: AstrMessageEvent, - images: list[bytes], - found_message: str | None = None, - ): - """将所有图片一次性发送,失败时尝试混淆重发。""" - try: - message_chain = [Plain(found_message)] if found_message else [] - for img_data in images: - message_chain.append(Image.fromBytes(img_data)) - yield event.chain_result(message_chain) - return - except Exception as exc: - logger.warning("send_images direct failed, retry obfuscated: %s", exc) - - try: - retry_text = f"{found_message} (混淆重发)" if found_message else "" - message_chain = [Plain(retry_text)] if retry_text else [] - for img_data in images: - obf_data = obfuscate_image_bytes(img_data) - message_chain.append(Image.fromBytes(obf_data)) - yield event.chain_result(message_chain) - except Exception as exc: - logger.exception("send_images obfuscated retry failed: %s", exc) - yield event.plain_result("图片发送失败,可能被平台审核拦截。") - - @staticmethod - async def send_forward( - event: AstrMessageEvent, images: list[bytes], bot_name: str = "Bot" - ): - """以合并转发节点方式发送图片。""" - logger.info("[forward] building nodes total=%d", len(images)) - nodes = [] - for index, img_data in enumerate(images): - try: - node = Node( - uin=event.get_self_id(), - name=bot_name, - content=[Image.fromBytes(img_data)], - ) - nodes.append(node) - except Exception as exc: - logger.exception("[forward] build node failed index=%d: %s", index, exc) - if not nodes: - yield event.plain_result("合并转发构建失败,未发送任何图片。") - return - yield event.chain_result([Nodes(nodes)]) diff --git a/session_config.py b/session_config.py deleted file mode 100644 index d7b5fc9..0000000 --- a/session_config.py +++ /dev/null @@ -1,455 +0,0 @@ -"""会话级别配置管理器。 - -提供会话级别的配置覆盖功能,允许管理员在特定会话中设置不同的内容模式。 -配置存储在 AstrBotConfig 中,可在 WebUI 中管理。 -""" - -from __future__ import annotations - -import asyncio -import json -import warnings -from pathlib import Path -from typing import TYPE_CHECKING, Any - -from astrbot.api import logger - -from .session_config_base import SessionConfigBase - -if TYPE_CHECKING: - from astrbot.core import AstrBotConfig - - -class SessionConfigManager(SessionConfigBase): - """会话级别配置管理器。 - - 管理每个会话(群聊/私聊)的独立配置,支持内容模式覆盖。 - 配置存储在 AstrBotConfig 的 session_configs 字段中,可在 WebUI 中管理。 - - Attributes: - _config: AstrBotConfig 配置对象 - _data_dir: 插件数据目录(用于迁移) - _lock: 异步锁,防止并发读写不一致 - """ - - # 允许会话覆盖的配置项 - ALLOWED_KEYS = {"content_mode", "r18_docx_mode", "auto_revoke_r18", "send_mode"} - - # 有效的内容模式值 - VALID_CONTENT_MODES = {"sfw", "r18", "mix"} - - # 有效的发送模式值 - VALID_SEND_MODES = {"image", "forward", "auto"} - - def __init__(self, config: AstrBotConfig, data_dir: Path | None = None): - """初始化会话配置管理器。 - - 参数: - config: AstrBotConfig 配置对象 - data_dir: 插件数据目录(用于迁移旧配置) - """ - super().__init__(config, config_key="session_configs") - self._data_dir = data_dir - self._migrated = False - - async def initialize(self) -> None: - """初始化配置,迁移旧配置文件到新格式。""" - if self._data_dir and not self._migrated: - await self._migrate_old_config() - self._migrated = True - logger.info("[session_config] SessionConfigManager initialized") - - async def _migrate_old_config(self) -> None: - """迁移旧位置的配置文件到 AstrBotConfig。""" - old_config_file = self._data_dir / "setu" / "setu_session_config.json" - if not old_config_file.exists(): - # 检查旧的位置 - old_config_file = self._data_dir / "session_config.json" - - if not old_config_file.exists(): - return - - try: - content = await asyncio.to_thread( - old_config_file.read_text, encoding="utf-8" - ) - old_data = json.loads(content) - sessions = old_data.get("sessions", {}) - - if not sessions: - return - - # 转换格式到新的 template_list 格式 - new_configs = self._get_configs() - - for session_key, session_data in sessions.items(): - parsed = self._parse_session_key(session_key) - if not parsed: - continue - - session_type, session_id = parsed - is_group = session_type == "group" - - # 构建新的配置项 - new_item = { - "session_id": session_id, - "session_type": "group" if is_group else "private", - } - - # 迁移各配置项 - if "content_mode" in session_data: - new_item["content_mode"] = session_data["content_mode"] - if "r18_docx_mode" in session_data: - new_item["r18_docx_mode"] = ( - "enabled" if session_data["r18_docx_mode"] else "disabled" - ) - if "auto_revoke_r18" in session_data: - new_item["auto_revoke_r18"] = ( - "enabled" if session_data["auto_revoke_r18"] else "disabled" - ) - if "send_mode" in session_data: - new_item["send_mode"] = session_data["send_mode"] - - self.merge_session_item(new_configs, new_item) - - # 保存迁移后的配置 - self._save_configs(new_configs) - - # 删除旧配置文件 - await asyncio.to_thread(old_config_file.unlink) - logger.info( - "[session_config] Migrated old config to AstrBotConfig, %d sessions", - len(new_configs), - ) - except (OSError, json.JSONDecodeError) as exc: - logger.warning("[session_config] Failed to migrate old config: %s", exc) - - def _get_session_config(self, session_id: str, is_group: bool) -> dict | None: - """获取会话配置项。""" - return self._find_session_config(session_id, is_group) - - def _get_session_key(self, session_id: str, is_group: bool) -> str: - """生成会话键。""" - prefix = "group" if is_group else "private" - return f"{prefix}:{session_id}" - - async def set_config( - self, session_id: str, is_group: bool, key: str, value: Any - ) -> bool: - """设置会话配置项。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - key: 配置键 - value: 配置值 - - 返回: - 设置成功返回 True,否则返回 False - """ - if key not in self.ALLOWED_KEYS: - logger.warning("[session_config] Key %s is not allowed", key) - return False - - # 验证 content_mode 值 - if key == "content_mode" and value not in self.VALID_CONTENT_MODES: - logger.warning("[session_config] Invalid content_mode: %s", value) - return False - - # 验证 send_mode 值 - if key == "send_mode" and value not in self.VALID_SEND_MODES: - logger.warning("[session_config] Invalid send_mode: %s", value) - return False - - async with self._lock: - session_type = "group" if is_group else "private" - configs = self._get_configs() - - # 查找现有配置 - existing_idx = self._find_session_index(configs, session_id, is_group) - - # 根据配置项类型设置值 - config_value = value - if key == "r18_docx_mode": - config_value = "enabled" if value else "disabled" - elif key == "auto_revoke_r18": - config_value = "enabled" if value else "disabled" - - if existing_idx is not None: - # 更新现有配置 - configs[existing_idx][key] = config_value - else: - # 创建新配置 - new_item = { - "session_id": session_id, - "session_type": session_type, - key: config_value, - } - configs.append(new_item) - - self._save_configs(configs) - - session_key = self._get_session_key(session_id, is_group) - logger.info( - "[session_config] Set %s=%s for session %s", key, value, session_key - ) - return True - - async def get_config( - self, session_id: str, is_group: bool, key: str, default: Any = None - ) -> Any: - """获取会话配置项。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - key: 配置键 - default: 默认值 - - 返回: - 配置值,如果不存在则返回默认值 - """ - async with self._lock: - cfg = self._get_session_config(session_id, is_group) - if not cfg: - return default - - value = cfg.get(key) - if value is None or value == "": - return default - - # 转换布尔值类型 - if key == "r18_docx_mode": - return value == "enabled" - if key == "auto_revoke_r18": - return value == "enabled" - - return value - - async def clear_config(self, session_id: str, is_group: bool, key: str) -> bool: - """清除会话配置项。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - key: 配置键 - - 返回: - 清除成功返回 True,否则返回 False - """ - async with self._lock: - configs = self._get_configs() - - idx = self._find_session_index(configs, session_id, is_group) - if idx is None: - return False - - cfg = configs[idx] - if key not in cfg: - return False - - del cfg[key] - # 如果配置项都为空,删除整个会话配置 - if not any(k in cfg for k in self.ALLOWED_KEYS if k != key): - configs.pop(idx) - - self._save_configs(configs) - session_key = self._get_session_key(session_id, is_group) - logger.info("[session_config] Cleared %s for session %s", key, session_key) - return True - - async def get_session_content_mode( - self, session_id: str, is_group: bool - ) -> str | None: - """获取会话的内容模式。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - - 返回: - 内容模式 (sfw/r18/mix),如果未设置则返回 None - """ - return await self.get_config(session_id, is_group, "content_mode") - - async def set_session_content_mode( - self, session_id: str, is_group: bool, mode: str - ) -> bool: - """设置会话的内容模式。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - mode: 内容模式 (sfw/r18/mix) - - 返回: - 设置成功返回 True,否则返回 False - """ - return await self.set_config(session_id, is_group, "content_mode", mode) - - async def clear_session_content_mode(self, session_id: str, is_group: bool) -> bool: - """清除会话的内容模式设置。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - - 返回: - 清除成功返回 True,否则返回 False - """ - return await self.clear_config(session_id, is_group, "content_mode") - - async def get_session_r18_docx_mode( - self, session_id: str, is_group: bool - ) -> bool | None: - """获取会话的 R18 Docx 模式设置。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - - 返回: - R18 Docx 模式设置 (True/False),如果未设置则返回 None - """ - return await self.get_config(session_id, is_group, "r18_docx_mode") - - async def set_session_r18_docx_mode( - self, session_id: str, is_group: bool, enabled: bool - ) -> bool: - """设置会话的 R18 Docx 模式。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - enabled: 是否启用 R18 Docx 模式 - - 返回: - 设置成功返回 True,否则返回 False - """ - return await self.set_config( - session_id, is_group, "r18_docx_mode", bool(enabled) - ) - - async def clear_session_r18_docx_mode( - self, session_id: str, is_group: bool - ) -> bool: - """清除会话的 R18 Docx 模式设置。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - - 返回: - 清除成功返回 True,否则返回 False - """ - return await self.clear_config(session_id, is_group, "r18_docx_mode") - - async def get_session_auto_revoke_r18( - self, session_id: str, is_group: bool - ) -> bool | None: - """获取会话的自动撤回 R18 设置。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - - 返回: - 自动撤回设置 (True/False),如果未设置则返回 None - """ - return await self.get_config(session_id, is_group, "auto_revoke_r18") - - async def set_session_auto_revoke_r18( - self, session_id: str, is_group: bool, enabled: bool - ) -> bool: - """设置会话的自动撤回 R18。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - enabled: 是否启用自动撤回 - - 返回: - 设置成功返回 True,否则返回 False - """ - return await self.set_config( - session_id, is_group, "auto_revoke_r18", bool(enabled) - ) - - async def clear_session_auto_revoke_r18( - self, session_id: str, is_group: bool - ) -> bool: - """清除会话的自动撤回 R18 设置。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - - 返回: - 清除成功返回 True,否则返回 False - """ - return await self.clear_config(session_id, is_group, "auto_revoke_r18") - - async def get_session_send_mode( - self, session_id: str, is_group: bool - ) -> str | None: - """获取会话的发送模式设置。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - - 返回: - 发送模式 (image/forward/auto),如果未设置则返回 None - """ - return await self.get_config(session_id, is_group, "send_mode") - - async def set_session_send_mode( - self, session_id: str, is_group: bool, mode: str - ) -> bool: - """设置会话的发送模式。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - mode: 发送模式 (image/forward/auto) - - 返回: - 设置成功返回 True,否则返回 False - """ - return await self.set_config(session_id, is_group, "send_mode", mode) - - async def clear_session_send_mode(self, session_id: str, is_group: bool) -> bool: - """清除会话的发送模式设置。 - - 参数: - session_id: 会话ID - is_group: 是否为群聊 - - 返回: - 清除成功返回 True,否则返回 False - """ - return await self.clear_config(session_id, is_group, "send_mode") - - async def cleanup_expired_sessions(self, max_age_days: int = 30) -> int: - """[Deprecated] 兼容接口:清理过期会话配置(当前实现为 no-op)。 - - 历史版本曾基于时间戳删除会话配置;当前配置结构不再保存时间信息, - 因此该方法仅兼容旧调用方而保留。 - - 参数: - max_age_days: 兼容参数,当前版本中不会被使用 - - 返回: - 始终返回 0,不会执行任何删除操作 - """ - warnings.warn( - "SessionConfigManager.cleanup_expired_sessions() is deprecated and is now a no-op.", - DeprecationWarning, - stacklevel=2, - ) - logger.warning( - "[session_config] cleanup_expired_sessions is deprecated no-op; " - "max_age_days=%s is ignored", - max_age_days, - ) - return 0 diff --git a/session_config_base.py b/session_config_base.py deleted file mode 100644 index 1df376f..0000000 --- a/session_config_base.py +++ /dev/null @@ -1,92 +0,0 @@ -"""Shared helpers for AstrBotConfig-backed session-level configuration.""" - -from __future__ import annotations - -import asyncio -from typing import TYPE_CHECKING - -if TYPE_CHECKING: - from astrbot.core import AstrBotConfig - - -class SessionConfigBase: - """Base utilities for session config managers. - - Subclasses provide key-specific validation and value conversion logic. - """ - - def __init__( - self, - config: AstrBotConfig, - *, - config_key: str, - ): - self._config = config - self._config_key = config_key - self._lock = asyncio.Lock() - - @staticmethod - def _session_type(is_group: bool) -> str: - return "group" if is_group else "private" - - def _get_session_key(self, session_id: str, is_group: bool) -> str: - return f"{self._session_type(is_group)}:{session_id}" - - def _get_configs(self) -> list[dict]: - return list(self._config.get(self._config_key, [])) - - def _save_configs(self, configs: list[dict]) -> None: - self._config[self._config_key] = configs - self._config.save_config() - - def _find_session_index( - self, - configs: list[dict], - session_id: str, - is_group: bool, - ) -> int | None: - session_type = self._session_type(is_group) - for i, cfg in enumerate(configs): - if ( - cfg.get("session_id") == session_id - and cfg.get("session_type") == session_type - ): - return i - return None - - def _find_session_config( - self, - session_id: str, - is_group: bool, - ) -> dict | None: - configs = self._get_configs() - idx = self._find_session_index(configs, session_id, is_group) - if idx is None: - return None - return configs[idx] - - @staticmethod - def _parse_session_key(session_key: str) -> tuple[str, str] | None: - parts = session_key.split(":", 1) - if len(parts) != 2: - return None - session_type, session_id = parts - if session_type not in {"group", "private"}: - return None - return session_type, session_id - - @staticmethod - def merge_session_item(configs: list[dict], item: dict) -> None: - """Insert or replace a session config item by (session_id, session_type).""" - session_id = item.get("session_id") - session_type = item.get("session_type") - if not session_id or session_type not in {"group", "private"}: - return - for i, cfg in enumerate(configs): - if ( - cfg.get("session_id") == session_id - and cfg.get("session_type") == session_type - ): - configs[i] = item - return - configs.append(item) diff --git a/skills/get-setu/SKILL.md b/skills/get-setu/SKILL.md new file mode 100644 index 0000000..c726810 --- /dev/null +++ b/skills/get-setu/SKILL.md @@ -0,0 +1,92 @@ +--- +name: get-setu +description: Help users operate the current astrbot_plugin_setu LLM tools for fetching Setu images, reading or changing unified session settings, R18 delivery safeguards, and today's fortune. Use this skill whenever the user asks for 色图/瑟图/涩图/setu/random anime images, image tags, content mode, send mode, auto revoke, DOCX packaging, session config, or 今日运势, even if they do not mention the plugin or tool names explicitly. +--- + +# get-setu + +Use this skill to choose the correct `astrbot_plugin_setu` LLM tool and explain the result. Tools operate in the current AstrBot session. Session config tools only affect the current session; WebUI can manage other sessions, but these LLM tools cannot. + +## Current Tool Inventory + +Use only these registered tools: + +- `get_setu_image(count: integer, tags: string[])` +- `get_today_fortune()` +- `refresh_my_fortune()` +- `refresh_group_fortune()` +- `refresh_all_fortune()` +- `get_session_config(key?: string)` +- `set_session_config(key: string, value: string)` +- `clear_session_config(key?: string)` + +Do not call removed legacy tools such as `get_setu_content_mode`, `set_setu_content_mode`, `set_setu_r18_docx_mode`, `set_setu_auto_revoke`, `set_setu_send_mode`, `get_fortune_config`, or `set_fortune_config`. + +## Session Config Keys + +Use these keys with `get_session_config`, `set_session_config`, and `clear_session_config`: + +| Key | Values | Meaning | +|---|---|---| +| `setu.content_mode` | `sfw`, `r18`, `mix` | Content rating for Setu image requests. | +| `setu.r18_docx` | `true`, `false` | Whether R18 images are packaged as DOCX. | +| `setu.auto_revoke` | `true`, `false` | Whether R18 messages are auto-revoked. | +| `setu.send_mode` | `image`, `forward`, `auto` | Direct image, merged forward, or automatic send mode. | +| `fortune.tags` | string | Default tags for fortune images. | +| `fortune.content_mode` | `sfw`, `r18`, `mix` | Content rating for fortune images. | + +Boolean values can be passed as strings like `"true"` or `"false"`. + +Global delivery settings such as `napcat_stream_mode`, send cache TTL, and send cache size are not session preferences and are not exposed through the LLM session config tools. Tell users to change them in the AstrBot WebUI plugin config. + +## Choosing Tools + +For image requests, extract count and tags from the user. Default to `count=1` and `tags=[]` when unspecified. Preserve meaningful tags such as `白丝`, `猫耳`, `blue archive`, or `long_hair`. The plugin decides R18 behavior from the current effective `setu.content_mode`; do not pass a separate R18 flag because the current tool does not expose one. + +For settings questions, call `get_session_config`. If the user asks about one setting, pass the specific key. If they ask for all settings, omit `key`. The tool returns JSON; read `ok`, then summarize `data.effective`, `data.override`, and `data.global` as needed. + +For setting changes, call `set_session_config(key, value)` with the exact key. The tool returns JSON; if `ok` is false, report `message` directly. If `ok` is true, summarize the changed key and effective value. + +For clearing settings, call `clear_session_config(key)` for one override or `clear_session_config()` to clear all current-session overrides. Explain that the global config will apply afterward. + +For fortune requests, call `get_today_fortune`. For refresh requests, use `refresh_my_fortune`, `refresh_group_fortune`, or `refresh_all_fortune` based on scope. + +## Permission Boundaries + +Treat `set_session_config`, `clear_session_config`, and all refresh tools as privileged. If a tool reports insufficient permission, tell the user the operation requires an administrator or super administrator. Do not retry with a different tool to bypass permission checks. + +Use extra care with `r18`, `mix`, `setu.auto_revoke`, and `setu.r18_docx`. Change only what the user requested; do not silently enable unrelated safety or delivery settings. + +## Response Style + +Be concise. After a fetch or fortune tool returns, say what was requested and whether the plugin reported success or failure. Do not claim that an image was sent unless the tool result says it succeeded. + +For failures, surface the plugin's message directly and suggest a practical next step such as fewer images, different tags, or a different send mode. Do not tell the user to configure HTML fallback in WebUI after fallback failure; the plugin now reports fallback failure as a normal send failure. + +For slow image sending on NapCat/OneBot, explain that the plugin downloads images to local send cache first and defaults to `napcat_stream_mode=fallback`: it tries normal file-path sending, then uses NapCat stream upload if that fails. `always` can be enabled globally in WebUI when the platform supports stream upload and large original images are common. + +## Examples + +User: `来三张白丝猫耳` + +Action: `get_setu_image(count=3, tags=["白丝", "猫耳"])` + +User: `现在是 r18 吗` + +Action: `get_session_config(key="setu.content_mode")` + +User: `这个群以后用合并转发` + +Action: `set_session_config(key="setu.send_mode", value="forward")` + +User: `开启 r18 自动撤回` + +Action: `set_session_config(key="setu.auto_revoke", value="true")` + +User: `清掉这个群的色图模式` + +Action: `clear_session_config(key="setu.content_mode")` + +User: `今日运势` + +Action: `get_today_fortune()` diff --git a/src/__init__.py b/src/__init__.py new file mode 100644 index 0000000..77804d0 --- /dev/null +++ b/src/__init__.py @@ -0,0 +1,14 @@ +"""AstrBot Setu plugin internals.""" + +from __future__ import annotations + +from .infrastructure.astrbot import clear_config, get_config, init_config, set_config +from .shared import get_logger + +__all__ = [ + "clear_config", + "get_config", + "get_logger", + "init_config", + "set_config", +] diff --git a/src/application/__init__.py b/src/application/__init__.py new file mode 100644 index 0000000..cd39aa2 --- /dev/null +++ b/src/application/__init__.py @@ -0,0 +1,3 @@ +"""Application layer for use cases, ports, and settings facades.""" + +from __future__ import annotations diff --git a/src/application/ports/__init__.py b/src/application/ports/__init__.py new file mode 100644 index 0000000..822facf --- /dev/null +++ b/src/application/ports/__init__.py @@ -0,0 +1,16 @@ +"""Application ports implemented by infrastructure adapters.""" + +from __future__ import annotations + +from .access_control_repository import AccessControlRepository +from .fortune_repository import FortuneRepository +from .image_provider import ImageProvider, SetuImageProvider +from .session_config_repository import SessionConfigRepository + +__all__ = [ + "AccessControlRepository", + "FortuneRepository", + "ImageProvider", + "SessionConfigRepository", + "SetuImageProvider", +] diff --git a/src/application/ports/access_control_repository.py b/src/application/ports/access_control_repository.py new file mode 100644 index 0000000..139855d --- /dev/null +++ b/src/application/ports/access_control_repository.py @@ -0,0 +1,137 @@ +"""Repository interface for access control data. + +Defines the contract for accessing blacklist/whitelist data. +Implemented by infrastructure layer (e.g., FileBackedAccessControlRepo). +""" + +from __future__ import annotations + +from abc import ABC, abstractmethod + + +class AccessControlRepository(ABC): + """Repository interface for access control data.""" + + # Setu user access control + @abstractmethod + async def is_setu_user_blocked(self, user_id: str) -> bool: + """Check if user is in Setu blacklist.""" + ... + + @abstractmethod + async def is_setu_user_whitelisted(self, user_id: str) -> bool: + """Check if user is in Setu whitelist.""" + ... + + @abstractmethod + async def add_setu_blocked_user(self, user_id: str) -> bool: + """Add user to Setu blacklist.""" + ... + + @abstractmethod + async def remove_setu_blocked_user(self, user_id: str) -> bool: + """Remove user from Setu blacklist.""" + ... + + @abstractmethod + async def add_setu_whitelist_user(self, user_id: str) -> bool: + """Add user to Setu whitelist.""" + ... + + @abstractmethod + async def remove_setu_whitelist_user(self, user_id: str) -> bool: + """Remove user from Setu whitelist.""" + ... + + # Setu group access control + @abstractmethod + async def is_setu_group_blocked(self, group_id: str) -> bool: + """Check if group is in Setu blacklist.""" + ... + + @abstractmethod + async def is_setu_group_whitelisted(self, group_id: str) -> bool: + """Check if group is in Setu whitelist.""" + ... + + @abstractmethod + async def add_setu_blocked_group(self, group_id: str) -> bool: + """Add group to Setu blacklist.""" + ... + + @abstractmethod + async def remove_setu_blocked_group(self, group_id: str) -> bool: + """Remove group from Setu blacklist.""" + ... + + @abstractmethod + async def add_setu_whitelist_group(self, group_id: str) -> bool: + """Add group to Setu whitelist.""" + ... + + @abstractmethod + async def remove_setu_whitelist_group(self, group_id: str) -> bool: + """Remove group from Setu whitelist.""" + ... + + # Fortune user access control + @abstractmethod + async def is_fortune_user_blocked(self, user_id: str) -> bool: + """Check if user is in Fortune blacklist.""" + ... + + @abstractmethod + async def is_fortune_user_whitelisted(self, user_id: str) -> bool: + """Check if user is in Fortune whitelist.""" + ... + + @abstractmethod + async def add_fortune_blocked_user(self, user_id: str) -> bool: + """Add user to Fortune blacklist.""" + ... + + @abstractmethod + async def remove_fortune_blocked_user(self, user_id: str) -> bool: + """Remove user from Fortune blacklist.""" + ... + + @abstractmethod + async def add_fortune_whitelist_user(self, user_id: str) -> bool: + """Add user to Fortune whitelist.""" + ... + + @abstractmethod + async def remove_fortune_whitelist_user(self, user_id: str) -> bool: + """Remove user from Fortune whitelist.""" + ... + + # Fortune group access control + @abstractmethod + async def is_fortune_group_blocked(self, group_id: str) -> bool: + """Check if group is in Fortune blacklist.""" + ... + + @abstractmethod + async def is_fortune_group_whitelisted(self, group_id: str) -> bool: + """Check if group is in Fortune whitelist.""" + ... + + @abstractmethod + async def add_fortune_blocked_group(self, group_id: str) -> bool: + """Add group to Fortune blacklist.""" + ... + + @abstractmethod + async def remove_fortune_blocked_group(self, group_id: str) -> bool: + """Remove group from Fortune blacklist.""" + ... + + @abstractmethod + async def add_fortune_whitelist_group(self, group_id: str) -> bool: + """Add group to Fortune whitelist.""" + ... + + @abstractmethod + async def remove_fortune_whitelist_group(self, group_id: str) -> bool: + """Remove group from Fortune whitelist.""" + ... diff --git a/src/application/ports/fortune_repository.py b/src/application/ports/fortune_repository.py new file mode 100644 index 0000000..1f78918 --- /dev/null +++ b/src/application/ports/fortune_repository.py @@ -0,0 +1,66 @@ +"""Application port for fortune persistence.""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from typing import Any + +from ...domain.fortune.entities import FortuneGenerationRequest, FortuneRecord + + +class FortuneRepository(ABC): + """Repository interface for fortune data.""" + + @abstractmethod + async def get_today_fortune( + self, request: FortuneGenerationRequest + ) -> FortuneRecord | None: + """Get fortune for today.""" + ... + + @abstractmethod + async def save_fortune(self, record: FortuneRecord) -> bool: + """Save or update fortune record.""" + ... + + @abstractmethod + async def delete_fortune(self, user_id: str, date_str: str) -> bool: + """Delete a fortune record.""" + ... + + @abstractmethod + async def delete_group_fortunes(self, group_id: str, date_str: str) -> int: + """Delete all fortune records for a group on a date.""" + ... + + @abstractmethod + async def delete_all_fortunes(self, date_str: str) -> int: + """Delete all fortune records for a date.""" + ... + + @abstractmethod + async def get_active_users(self, days: int = 3) -> list[str]: + """Get active user IDs.""" + ... + + @abstractmethod + async def get_cached_image_path(self, user_id: str, date_str: str) -> Any | None: + """Get cached image path.""" + ... + + @abstractmethod + async def save_cached_image( + self, user_id: str, date_str: str, image_data: bytes, img_url: str | None + ) -> Any: + """Save cached image data.""" + ... + + @abstractmethod + async def delete_cached_image(self, user_id: str, date_str: str) -> bool: + """Delete cached image.""" + ... + + @abstractmethod + async def cleanup_expired_cache(self, date_str: str) -> int: + """Clean expired cache files.""" + ... diff --git a/src/application/ports/image_provider.py b/src/application/ports/image_provider.py new file mode 100644 index 0000000..12705b2 --- /dev/null +++ b/src/application/ports/image_provider.py @@ -0,0 +1,247 @@ +"""图片提供商基类。""" + +from __future__ import annotations + +import asyncio +from pathlib import Path +from urllib.parse import urlsplit, urlunsplit + +import httpx + +from ...domain.setu import SetuRequest +from ...shared import get_logger +from ...shared.send_cache import get_send_cache +from ..setu.dto import ImagePayload + +logger = get_logger() + + +class SetuImageProvider: + """色图图片提供商基类(策略模式)。""" + + async def fetch_image_urls( + self, + num: int, + tags: list[str], + r18: bool, + exclude_ai: bool = True, + ) -> list[str]: + """从 API 获取图片 URL 列表。 + + 参数: + num: 要获取的图片数量。 + tags: 搜索标签/关键词。 + r18: 是否请求 R18 内容。 + exclude_ai: 是否排除 AI 生成的作品。 + + 返回: + 图片 URL 列表。 + """ + raise NotImplementedError + + async def fetch_and_download(self, request: SetuRequest) -> ImagePayload: + """Fetch image URLs and download them to local files for delivery.""" + provider_name = self._provider_name() + logger.info( + "[provider] fetch start: provider=%s, count=%d, r18=%s, tags=%s, exclude_ai=%s", + provider_name, + request.count, + request.r18, + ",".join(request.tags) or "-", + request.exclude_ai, + ) + urls = await self.fetch_image_urls( + num=request.count, + tags=list(request.tags), + r18=request.r18, + exclude_ai=request.exclude_ai, + ) + if not urls: + logger.warning( + "[provider] no urls returned: provider=%s, count=%d, r18=%s, tags=%s", + provider_name, + request.count, + request.r18, + ",".join(request.tags) or "-", + ) + return ImagePayload( + urls=(), raw_bytes=(), file_paths=(), r18=request.r18, tags=request.tags + ) + + cache = get_send_cache() + cache_enabled = bool(cache and cache.enabled) + logger.info( + "[provider] fetch result: provider=%s, urls=%d, cache_enabled=%s", + provider_name, + len(urls), + cache_enabled, + ) + + async def download(client: httpx.AsyncClient, url: str) -> Path | bytes | None: + try: + if cache_enabled and cache: + cached = await cache.get(url) + if cached is not None: + logger.debug( + "[provider] cache hit: provider=%s, url=%s", + provider_name, + url, + ) + return cached + + if cache is not None: + write = None + try: + async with client.stream("GET", url) as response: + response.raise_for_status() + write = await cache.reserve( + url, response.headers.get("content-type") + ) + with write.temp_path.open("wb") as file: + async for chunk in response.aiter_bytes(): + if chunk: + await asyncio.to_thread(file.write, chunk) + final_path = await cache.commit(write) + logger.debug( + "[provider] download cached: provider=%s, url=%s, path=%s", + provider_name, + url, + final_path, + ) + return final_path + except Exception: + if write is not None: + await cache.discard(write) + raise + + response = await client.get(url) + response.raise_for_status() + logger.debug( + "[provider] download bytes: provider=%s, url=%s, bytes=%d", + provider_name, + url, + len(response.content), + ) + return response.content + except (httpx.HTTPError, asyncio.TimeoutError) as exc: + logger.warning( + "[provider] download failed: provider=%s, url=%s, error=%s", + provider_name, + url, + exc, + ) + return None + + async with httpx.AsyncClient(timeout=30.0, follow_redirects=True) as client: + results = await asyncio.gather(*(download(client, url) for url in urls)) + + items = tuple(item for item in results if item is not None) + raw_bytes = tuple(item for item in items if isinstance(item, bytes)) + file_paths = tuple(item for item in items if isinstance(item, Path)) + logger.info( + "[provider] download summary: provider=%s, requested=%d, succeeded=%d, failed=%d, bytes=%d, files=%d", + provider_name, + len(urls), + len(items), + len(urls) - len(items), + len(raw_bytes), + len(file_paths), + ) + if not items: + logger.error( + "[provider] all downloads failed: provider=%s, requested=%d, tags=%s", + provider_name, + len(urls), + ",".join(request.tags) or "-", + ) + return ImagePayload( + urls=tuple(urls), + raw_bytes=raw_bytes, + file_paths=file_paths, + items=items, + r18=request.r18, + tags=request.tags, + ) + + def _provider_name(self) -> str: + return self.__class__.__name__ + + def _apply_proxy_to_url(self, url: str, proxy: str | None) -> str: + """Rewrite Pixiv-style image URLs to the configured reverse proxy host.""" + proxy_host = (proxy or "").strip() + if not proxy_host: + return url + + try: + parsed = urlsplit(url) + except ValueError: + return url + + if parsed.scheme not in ("http", "https") or not parsed.netloc: + return url + + original_host = parsed.hostname or "" + proxy_targets = { + "i.pximg.net", + "i.pixiv.re", + "i.pixiv.cat", + "proxy.pixivel.moe", + } + if original_host not in proxy_targets: + return url + + if ":" in proxy_host: + netloc = proxy_host + else: + netloc = proxy_host + return urlunsplit( + (parsed.scheme, netloc, parsed.path, parsed.query, parsed.fragment) + ) + + def _apply_proxy_to_urls( + self, urls: list[str], proxy: str | None, provider_name: str + ) -> list[str]: + if not proxy: + return urls + + rewritten = [self._apply_proxy_to_url(url, proxy) for url in urls] + changed = sum( + 1 for old, new in zip(urls, rewritten, strict=False) if old != new + ) + if changed: + logger.info( + "[provider] proxy rewritten: provider=%s, proxy=%s, changed=%d", + provider_name, + proxy, + changed, + ) + else: + logger.debug( + "[provider] proxy rewrite skipped: provider=%s, proxy=%s, urls=%d", + provider_name, + proxy, + len(urls), + ) + return rewritten + + @staticmethod + def _normalize_bool(value, default: bool = True) -> bool: + """Normalize possibly-dirty boolean input from config/runtime sources.""" + if value is None: + return default + if isinstance(value, bool): + return value + if isinstance(value, (int, float)): + return bool(value) + if isinstance(value, str): + lowered = value.strip().lower() + if lowered in {"true", "1", "yes", "on"}: + return True + if lowered in {"false", "0", "no", "off"}: + return False + if lowered in {"none", "null", ""}: + return default + return default + + +ImageProvider = SetuImageProvider diff --git a/src/application/ports/session_config_repository.py b/src/application/ports/session_config_repository.py new file mode 100644 index 0000000..df4d534 --- /dev/null +++ b/src/application/ports/session_config_repository.py @@ -0,0 +1,31 @@ +"""Application port for session configuration persistence.""" + +from __future__ import annotations + +from abc import ABC, abstractmethod + +from ..session_config.dto import SessionConfigRecord + + +class SessionConfigRepository(ABC): + """Repository interface for session-scoped config overrides.""" + + @abstractmethod + async def list_sessions(self) -> list[SessionConfigRecord]: + """List all stored session records.""" + ... + + @abstractmethod + async def get_session(self, session_id: str) -> SessionConfigRecord | None: + """Get a session record by ID.""" + ... + + @abstractmethod + async def upsert_session(self, record: SessionConfigRecord) -> SessionConfigRecord: + """Create or replace a session record.""" + ... + + @abstractmethod + async def delete_session(self, session_id: str) -> bool: + """Delete a session record.""" + ... diff --git a/src/application/session_config/__init__.py b/src/application/session_config/__init__.py new file mode 100644 index 0000000..2764c54 --- /dev/null +++ b/src/application/session_config/__init__.py @@ -0,0 +1,29 @@ +"""Session-scoped configuration application services.""" + +from __future__ import annotations + +from .dto import JsonValue, SessionConfigRecord, SessionConfigSnapshot, SessionType +from .keys import ( + SESSION_CONFIG_KEYS, + SessionConfigKey, + SessionConfigValidationError, + get_key_definition, + normalize_config_value, + normalize_session_type, +) +from .service import SessionConfigService, get_global_session_config_values + +__all__ = [ + "JsonValue", + "SESSION_CONFIG_KEYS", + "SessionConfigKey", + "SessionConfigRecord", + "SessionConfigService", + "SessionConfigSnapshot", + "SessionConfigValidationError", + "SessionType", + "get_global_session_config_values", + "get_key_definition", + "normalize_config_value", + "normalize_session_type", +] diff --git a/src/application/session_config/dto.py b/src/application/session_config/dto.py new file mode 100644 index 0000000..5d14bfa --- /dev/null +++ b/src/application/session_config/dto.py @@ -0,0 +1,51 @@ +"""DTOs for per-session configuration overrides.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Literal + +JsonValue = str | int | float | bool | None | list[Any] | dict[str, Any] +SessionType = Literal["group", "private"] + + +@dataclass(frozen=True, slots=True) +class SessionConfigRecord: + """Stored override record for one chat session.""" + + session_id: str + session_type: SessionType = "private" + display_name: str = "" + overrides: dict[str, JsonValue] = field(default_factory=dict) + + def to_dict(self) -> dict[str, Any]: + """Return a JSON-serializable representation.""" + return { + "session_id": self.session_id, + "session_type": self.session_type, + "display_name": self.display_name, + "overrides": dict(self.overrides), + } + + +@dataclass(frozen=True, slots=True) +class SessionConfigSnapshot: + """Effective configuration for one session.""" + + session_id: str + session_type: SessionType + display_name: str + overrides: dict[str, JsonValue] + global_values: dict[str, JsonValue] + effective: dict[str, JsonValue] + + def to_dict(self) -> dict[str, Any]: + """Return a JSON-serializable representation.""" + return { + "session_id": self.session_id, + "session_type": self.session_type, + "display_name": self.display_name, + "overrides": dict(self.overrides), + "global": dict(self.global_values), + "effective": dict(self.effective), + } diff --git a/src/application/session_config/keys.py b/src/application/session_config/keys.py new file mode 100644 index 0000000..294708b --- /dev/null +++ b/src/application/session_config/keys.py @@ -0,0 +1,153 @@ +"""Known session-scoped configuration keys and validation.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Literal + +from .dto import JsonValue + +ValueType = Literal["enum", "bool", "string"] + + +@dataclass(frozen=True, slots=True) +class SessionConfigKey: + """Metadata for a supported session override key.""" + + key: str + label: str + value_type: ValueType + description: str + options: tuple[str, ...] = () + + def to_dict(self) -> dict[str, Any]: + """Return a JSON-serializable schema fragment.""" + return { + "key": self.key, + "label": self.label, + "type": self.value_type, + "description": self.description, + "options": list(self.options), + } + + +SESSION_CONFIG_KEYS: dict[str, SessionConfigKey] = { + "setu.content_mode": SessionConfigKey( + key="setu.content_mode", + label="色图内容分级", + value_type="enum", + options=("sfw", "r18", "mix"), + description="当前会话请求色图时使用的内容分级。", + ), + "setu.r18_docx": SessionConfigKey( + key="setu.r18_docx", + label="R18 Docx 打包", + value_type="bool", + description="当前会话发送 R18 图片时是否优先打包为 Docx。", + ), + "setu.auto_revoke": SessionConfigKey( + key="setu.auto_revoke", + label="R18 自动撤回", + value_type="bool", + description="当前会话发送 R18 内容后是否自动撤回。", + ), + "setu.send_mode": SessionConfigKey( + key="setu.send_mode", + label="图片发送模式", + value_type="enum", + options=("image", "forward", "auto"), + description="当前会话发送图片时使用直发、合并转发或自动模式。", + ), + "fortune.tags": SessionConfigKey( + key="fortune.tags", + label="运势图片标签", + value_type="string", + description="当前会话生成运势图片时使用的默认标签。", + ), + "fortune.content_mode": SessionConfigKey( + key="fortune.content_mode", + label="运势内容分级", + value_type="enum", + options=("sfw", "r18", "mix"), + description="当前会话生成运势图片时使用的内容分级。", + ), +} + +TRUE_VALUES = { + "1", + "true", + "yes", + "y", + "on", + "enable", + "enabled", + "开", + "开启", + "启用", + "是", +} +FALSE_VALUES = { + "0", + "false", + "no", + "n", + "off", + "disable", + "disabled", + "关", + "关闭", + "禁用", + "否", +} + + +class SessionConfigValidationError(ValueError): + """Raised when a session configuration key or value is invalid.""" + + +def get_key_definition(key: str) -> SessionConfigKey: + """Return metadata for a key or raise a validation error.""" + normalized = key.strip() + definition = SESSION_CONFIG_KEYS.get(normalized) + if definition is None: + raise SessionConfigValidationError( + f"未知配置项:{key}。可用配置项:{', '.join(SESSION_CONFIG_KEYS)}" + ) + return definition + + +def normalize_session_type(value: str | None) -> Literal["group", "private"]: + """Normalize a session type from user/API input.""" + normalized = (value or "private").strip().lower() + if normalized not in ("group", "private"): + raise SessionConfigValidationError("session_type 必须是 group 或 private") + return normalized # type: ignore[return-value] + + +def normalize_config_value(key: str, value: Any) -> JsonValue: + """Normalize and validate a value for a known session config key.""" + definition = get_key_definition(key) + if definition.value_type == "bool": + return _normalize_bool(value) + if definition.value_type == "enum": + text = str(value).strip().lower() + if text not in definition.options: + raise SessionConfigValidationError( + f"{definition.label} 的值必须是:{', '.join(definition.options)}" + ) + return text + return "" if value is None else str(value).strip() + + +def _normalize_bool(value: Any) -> bool: + """Normalize common boolean spellings.""" + if isinstance(value, bool): + return value + if isinstance(value, int) and value in (0, 1): + return bool(value) + text = str(value).strip().lower() + if text in TRUE_VALUES: + return True + if text in FALSE_VALUES: + return False + raise SessionConfigValidationError("布尔值必须是 true/false、on/off 或 1/0") diff --git a/src/application/session_config/service.py b/src/application/session_config/service.py new file mode 100644 index 0000000..59095f0 --- /dev/null +++ b/src/application/session_config/service.py @@ -0,0 +1,191 @@ +"""Application service for session-scoped configuration overrides.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from ..settings import get_delivery_settings, get_fortune_settings, get_setu_settings +from .dto import JsonValue, SessionConfigRecord, SessionConfigSnapshot +from .keys import ( + SESSION_CONFIG_KEYS, + get_key_definition, + normalize_config_value, + normalize_session_type, +) + +if TYPE_CHECKING: + from ..ports.session_config_repository import SessionConfigRepository + + +class SessionConfigService: + """Coordinate validation, persistence, and effective session config snapshots.""" + + def __init__(self, repository: SessionConfigRepository) -> None: + self._repo = repository + + async def list_snapshots(self) -> list[SessionConfigSnapshot]: + """Return effective snapshots for all stored sessions.""" + records = await self._repo.list_sessions() + return [self._snapshot_from_record(record) for record in records] + + async def get_snapshot( + self, + session_id: str, + session_type: str = "private", + display_name: str = "", + ) -> SessionConfigSnapshot: + """Return the effective config snapshot for a session.""" + record = await self._repo.get_session(_normalize_session_id(session_id)) + if record is None: + record = SessionConfigRecord( + session_id=_normalize_session_id(session_id), + session_type=normalize_session_type(session_type), + display_name=display_name.strip(), + ) + return self._snapshot_from_record(record) + + async def get_effective_value( + self, + session_id: str, + key: str, + session_type: str = "private", + display_name: str = "", + ) -> JsonValue: + """Return one effective value for a session.""" + snapshot = await self.get_snapshot(session_id, session_type, display_name) + return snapshot.effective[key] + + async def set_value( + self, + session_id: str, + session_type: str, + key: str, + value: Any, + display_name: str = "", + ) -> SessionConfigSnapshot: + """Set one override value and return the updated snapshot.""" + session_id = _normalize_session_id(session_id) + normalized_key = key.strip() + record = await self._repo.get_session(session_id) + if record is None: + record = SessionConfigRecord( + session_id=session_id, + session_type=normalize_session_type(session_type), + display_name=display_name.strip(), + ) + + overrides = dict(record.overrides) + overrides[normalized_key] = normalize_config_value(normalized_key, value) + updated = SessionConfigRecord( + session_id=session_id, + session_type=normalize_session_type(session_type or record.session_type), + display_name=(display_name or record.display_name).strip(), + overrides=overrides, + ) + saved = await self._repo.upsert_session(updated) + return self._snapshot_from_record(saved) + + async def clear( + self, + session_id: str, + session_type: str, + key: str | None = None, + display_name: str = "", + ) -> SessionConfigSnapshot: + """Clear one override or all overrides while keeping the session record.""" + session_id = _normalize_session_id(session_id) + record = await self._repo.get_session(session_id) + if record is None: + record = SessionConfigRecord( + session_id=session_id, + session_type=normalize_session_type(session_type), + display_name=display_name.strip(), + ) + + overrides = dict(record.overrides) + if key: + get_key_definition(key) + overrides.pop(key.strip(), None) + else: + overrides.clear() + + updated = SessionConfigRecord( + session_id=session_id, + session_type=normalize_session_type(session_type or record.session_type), + display_name=(display_name or record.display_name).strip(), + overrides=overrides, + ) + saved = await self._repo.upsert_session(updated) + return self._snapshot_from_record(saved) + + async def upsert_session( + self, + session_id: str, + session_type: str, + display_name: str = "", + overrides: dict[str, Any] | None = None, + ) -> SessionConfigSnapshot: + """Create or replace a session config record from WebUI input.""" + normalized_overrides = { + key.strip(): normalize_config_value(key, value) + for key, value in (overrides or {}).items() + } + record = SessionConfigRecord( + session_id=_normalize_session_id(session_id), + session_type=normalize_session_type(session_type), + display_name=display_name.strip(), + overrides=normalized_overrides, + ) + saved = await self._repo.upsert_session(record) + return self._snapshot_from_record(saved) + + async def delete_session(self, session_id: str) -> bool: + """Delete a whole session record.""" + return await self._repo.delete_session(_normalize_session_id(session_id)) + + def _snapshot_from_record( + self, record: SessionConfigRecord + ) -> SessionConfigSnapshot: + global_values = get_global_session_config_values() + overrides = _filter_known_overrides(record.overrides) + effective = dict(global_values) + effective.update(overrides) + return SessionConfigSnapshot( + session_id=record.session_id, + session_type=record.session_type, + display_name=record.display_name, + overrides=overrides, + global_values=global_values, + effective=effective, + ) + + +def get_global_session_config_values() -> dict[str, JsonValue]: + """Return current read-only global values relevant to session overrides.""" + setu = get_setu_settings() + delivery = get_delivery_settings() + fortune = get_fortune_settings() + return { + "setu.content_mode": setu.content_mode, + "setu.r18_docx": delivery.r18_docx_mode, + "setu.auto_revoke": delivery.auto_revoke_r18, + "setu.send_mode": delivery.send_mode, + "fortune.tags": fortune.tags, + "fortune.content_mode": fortune.content_mode, + } + + +def _filter_known_overrides(overrides: dict[str, Any]) -> dict[str, JsonValue]: + """Keep only known and valid override values.""" + result: dict[str, JsonValue] = {} + for key, value in overrides.items(): + if key in SESSION_CONFIG_KEYS: + result[key] = normalize_config_value(key, value) + return result + + +def _normalize_session_id(session_id: str) -> str: + normalized = str(session_id).strip() + if not normalized: + raise ValueError("session_id 不能为空") + return normalized diff --git a/src/application/settings.py b/src/application/settings.py new file mode 100644 index 0000000..a1a87f4 --- /dev/null +++ b/src/application/settings.py @@ -0,0 +1,120 @@ +"""Read-only settings facade for application code. + +Infrastructure owns the full AstrBot/Pydantic config object. Application code reads +small settings snapshots so it does not reach into the full runtime config tree. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +from ..shared.config import SetuPluginConfig + + +@dataclass(frozen=True) +class SetuSettings: + """Settings needed by Setu use cases.""" + + max_count: int + content_mode: str + exclude_ai: bool + fetch_timeout_seconds: float = 60.0 + + +@dataclass(frozen=True) +class DeliverySettings: + """Settings needed by adapter-level delivery.""" + + send_mode: str + html_card_strategy: str + auto_revoke_r18: bool + auto_revoke_delay: int + r18_docx_mode: bool + html_card_padding: int + html_card_gap: int + napcat_stream_mode: str + + +@dataclass(frozen=True) +class FortuneSettings: + """Settings needed by fortune use cases.""" + + enabled: bool + content_mode: str + tags: str + allow_user_refresh: bool + auto_refresh: bool + + +_config: SetuPluginConfig | None = None + + +def set_application_config(config: SetuPluginConfig) -> None: + """Set the config snapshot used by application settings getters.""" + global _config + _config = config + + +def clear_application_config() -> None: + """Clear application config for tests and plugin shutdown.""" + global _config + _config = None + + +def get_setu_settings() -> SetuSettings: + """Return Setu settings from the current config snapshot.""" + if _config is None: + return SetuSettings( + max_count=10, + content_mode="sfw", + exclude_ai=True, + ) + return SetuSettings( + max_count=_config.max_count, + content_mode=_config.content_mode, + exclude_ai=_config.exclude_ai, + ) + + +def get_delivery_settings() -> DeliverySettings: + """Return delivery settings from the current config snapshot.""" + if _config is None: + return DeliverySettings( + send_mode="image", + html_card_strategy="never", + auto_revoke_r18=False, + auto_revoke_delay=30, + r18_docx_mode=False, + html_card_padding=6, + html_card_gap=6, + napcat_stream_mode="fallback", + ) + return DeliverySettings( + send_mode=_config.send_mode, + html_card_strategy=_config.html_card_strategy, + auto_revoke_r18=_config.auto_revoke_r18, + auto_revoke_delay=_config.auto_revoke_delay, + r18_docx_mode=_config.r18_docx_mode, + html_card_padding=_config.html_card_padding, + html_card_gap=_config.html_card_gap, + napcat_stream_mode=_config.napcat_stream_mode, + ) + + +def get_fortune_settings() -> FortuneSettings: + """Return fortune settings from the current config snapshot.""" + if _config is None: + return FortuneSettings( + enabled=True, + content_mode="sfw", + tags="", + allow_user_refresh=False, + auto_refresh=True, + ) + return FortuneSettings( + enabled=_config.fortune.enabled, + content_mode=_config.fortune.content_mode.value, + tags=_config.fortune.tags, + allow_user_refresh=_config.fortune.allow_user_refresh, + auto_refresh=_config.fortune.auto_refresh, + ) diff --git a/src/application/setu/__init__.py b/src/application/setu/__init__.py new file mode 100644 index 0000000..d25f6af --- /dev/null +++ b/src/application/setu/__init__.py @@ -0,0 +1,7 @@ +"""Setu application use cases.""" + +from __future__ import annotations + +from .dto import ImagePayload, SetuImagesResult + +__all__ = ["ImagePayload", "SetuImagesResult"] diff --git a/src/application/setu/dto.py b/src/application/setu/dto.py new file mode 100644 index 0000000..2580154 --- /dev/null +++ b/src/application/setu/dto.py @@ -0,0 +1,43 @@ +"""Setu application DTOs.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + + +@dataclass(frozen=True) +class ImagePayload: + """Fetched image data ready for adapter-level delivery.""" + + urls: tuple[str, ...] + raw_bytes: tuple[bytes, ...] + r18: bool + tags: tuple[str, ...] + file_paths: tuple[Path, ...] = () + items: tuple[Path | bytes, ...] = () + + @property + def is_empty(self) -> bool: + """Check if payload contains no data.""" + return ( + not self.urls + and not self.raw_bytes + and not self.file_paths + and not self.items + ) + + @property + def count(self) -> int: + """Return number of images in payload.""" + if self.items: + return len(self.items) + return max(len(self.urls), len(self.raw_bytes), len(self.file_paths)) + + +@dataclass(frozen=True) +class SetuImagesResult: + """Application result for a Setu image request.""" + + payload: ImagePayload | None + notice: str | None = None diff --git a/src/application/setu/get_images.py b/src/application/setu/get_images.py new file mode 100644 index 0000000..0a22ad4 --- /dev/null +++ b/src/application/setu/get_images.py @@ -0,0 +1,39 @@ +"""Use case for fetching Setu images.""" + +from __future__ import annotations + +import asyncio + +from ...domain.setu import SetuRequest +from ..ports import ImageProvider +from ..settings import get_setu_settings +from .dto import SetuImagesResult + + +class GetSetuImagesUseCase: + """Fetch Setu image data without knowing how a chat platform will send it.""" + + def __init__(self, image_provider: ImageProvider) -> None: + self._provider = image_provider + + async def execute(self, count: int, tags: list[str], r18: bool) -> SetuImagesResult: + """Fetch image files according to current settings.""" + settings = get_setu_settings() + request = SetuRequest.from_user_input( + count=count, + tags=tags, + r18=r18, + exclude_ai=settings.exclude_ai, + ) + + payload = await asyncio.wait_for( + self._provider.fetch_and_download(request), + timeout=settings.fetch_timeout_seconds, + ) + + if payload.is_empty: + tags_info = f"标签: {', '.join(tags)}" if tags else "" + return SetuImagesResult( + payload=None, notice=f"未找到{tags_info}符合要求的图片~" + ) + return SetuImagesResult(payload=payload) diff --git a/src/domain/__init__.py b/src/domain/__init__.py new file mode 100644 index 0000000..21d5956 --- /dev/null +++ b/src/domain/__init__.py @@ -0,0 +1,64 @@ +"""Domain layer for AstrBot Setu plugin. + +Contains value objects, exceptions, enums, and domain services following +Domain-Driven Design principles. +""" + +from __future__ import annotations + +from .access_control import AccessPolicy +from .constants import COMMAND_PATTERN, FORTUNE_PATTERN, HTTP_TIMEOUT_SECONDS +from .enums import ( + AccessControlMode, + ApiType, + ContentMode, + HtmlCardStrategy, + MultiApiStrategy, + SendMode, +) +from .exceptions import ( + AccessDeniedError, + FortuneException, + FortuneNotFoundError, + ProviderError, + SendError, + SetuException, + ValidationError, +) +from .fortune import ( + FortuneGenerationRequest, + FortuneRecord, + FortuneResult, + FortuneSeed, +) +from .setu import SetuRequest, TagResolverService + +__all__ = [ + # Constants + "COMMAND_PATTERN", + "FORTUNE_PATTERN", + "HTTP_TIMEOUT_SECONDS", + # Enums + "ApiType", + "ContentMode", + "SendMode", + "HtmlCardStrategy", + "MultiApiStrategy", + "AccessControlMode", + # Exceptions + "SetuException", + "ProviderError", + "SendError", + "AccessDeniedError", + "ValidationError", + "FortuneException", + "FortuneNotFoundError", + # Value Objects + "SetuRequest", + "AccessPolicy", + "FortuneSeed", + "FortuneResult", + "FortuneGenerationRequest", + "FortuneRecord", + "TagResolverService", +] diff --git a/src/domain/access_control/__init__.py b/src/domain/access_control/__init__.py new file mode 100644 index 0000000..5d8e85c --- /dev/null +++ b/src/domain/access_control/__init__.py @@ -0,0 +1,7 @@ +"""Access-control domain value objects.""" + +from __future__ import annotations + +from .value_objects import AccessPolicy + +__all__ = ["AccessPolicy"] diff --git a/src/domain/access_control/service.py b/src/domain/access_control/service.py new file mode 100644 index 0000000..1805a17 --- /dev/null +++ b/src/domain/access_control/service.py @@ -0,0 +1,134 @@ +"""Domain service for access control logic. + +Centralizes access control decisions that were previously duplicated +across SetuCore, CommandHandler, and FortuneCommandHandler. +""" + +from __future__ import annotations + +from ...application.ports import AccessControlRepository +from .value_objects import AccessPolicy + + +class AccessControlService: + """Domain service for access control decisions. + + Provides a single place for access control logic, eliminating + duplication across SetuCore, CommandHandler, and FortuneCommandHandler. + """ + + def __init__(self, repository: AccessControlRepository) -> None: + """Initialize access control service. + + Args: + repository: Repository for accessing blacklist/whitelist data + """ + self._repo = repository + + async def check_setu_access(self, policy: AccessPolicy) -> tuple[bool, str | None]: + """Check if Setu access is allowed. + + Args: + policy: Access policy containing user/group IDs and modes + + Returns: + Tuple of (allowed, denial_reason). If allowed=True, reason is None. + """ + return await self._check_access( + policy, + user_blacklist_fn=self._repo.is_setu_user_blocked, + user_whitelist_fn=self._repo.is_setu_user_whitelisted, + group_blacklist_fn=self._repo.is_setu_group_blocked, + group_whitelist_fn=self._repo.is_setu_group_whitelisted, + feature_name="setu", + ) + + async def check_fortune_access( + self, policy: AccessPolicy + ) -> tuple[bool, str | None]: + """Check if Fortune access is allowed. + + Args: + policy: Access policy containing user/group IDs and modes + + Returns: + Tuple of (allowed, denial_reason). If allowed=True, reason is None. + """ + return await self._check_access( + policy, + user_blacklist_fn=self._repo.is_fortune_user_blocked, + user_whitelist_fn=self._repo.is_fortune_user_whitelisted, + group_blacklist_fn=self._repo.is_fortune_group_blocked, + group_whitelist_fn=self._repo.is_fortune_group_whitelisted, + feature_name="fortune", + ) + + async def _check_access( + self, + policy: AccessPolicy, + user_blacklist_fn, + user_whitelist_fn, + group_blacklist_fn, + group_whitelist_fn, + feature_name: str, + ) -> tuple[bool, str | None]: + """Internal access check implementation. + + Args: + policy: Access policy + user_blacklist_fn: Async function to check if user is blacklisted + user_whitelist_fn: Async function to check if user is whitelisted + group_blacklist_fn: Async function to check if group is blacklisted + group_whitelist_fn: Async function to check if group is whitelisted + feature_name: Feature name for logging + + Returns: + Tuple of (allowed, denial_reason) + """ + from astrbot.api import logger + + # Check user access control + if policy.user_id is not None: + uid = str(policy.user_id) + + if policy.user_mode == "blacklist": + if await user_blacklist_fn(uid): + logger.info( + "[%s] Access denied for user=%s: blacklist mode", + feature_name, + uid, + ) + return False, "用户被禁用" + + elif policy.user_mode == "whitelist": + if not await user_whitelist_fn(uid): + logger.info( + "[%s] Access denied for user=%s: not in whitelist", + feature_name, + uid, + ) + return False, "用户不在白名单中" + + # Check group access control + if policy.group_id is not None: + gid = str(policy.group_id) + + if policy.group_mode == "blacklist": + if await group_blacklist_fn(gid): + logger.info( + "[%s] Access denied for group=%s: blacklist mode", + feature_name, + gid, + ) + return False, "群组被禁用" + + elif policy.group_mode == "whitelist": + if not await group_whitelist_fn(gid): + logger.info( + "[%s] Access denied for group=%s: not in whitelist", + feature_name, + gid, + ) + return False, "群组不在白名单中" + + return True, None diff --git a/src/domain/access_control/value_objects.py b/src/domain/access_control/value_objects.py new file mode 100644 index 0000000..817e11a --- /dev/null +++ b/src/domain/access_control/value_objects.py @@ -0,0 +1,36 @@ +"""Access-control domain value objects.""" + +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class AccessPolicy: + """Value object for access-control policy.""" + + user_id: str | None + group_id: str | None + user_mode: str + group_mode: str + + @classmethod + def for_user(cls, user_id: str, user_mode: str = "none") -> AccessPolicy: + """Create policy for user-only access control.""" + return cls(user_id, None, user_mode, "none") + + @classmethod + def for_group(cls, group_id: str, group_mode: str = "none") -> AccessPolicy: + """Create policy for group-only access control.""" + return cls(None, group_id, "none", group_mode) + + @classmethod + def for_session( + cls, + user_id: str | None, + group_id: str | None, + user_mode: str = "none", + group_mode: str = "none", + ) -> AccessPolicy: + """Create policy for full session access control.""" + return cls(user_id, group_id, user_mode, group_mode) diff --git a/constants.py b/src/domain/constants.py similarity index 93% rename from constants.py rename to src/domain/constants.py index fbf674b..5d7188b 100644 --- a/constants.py +++ b/src/domain/constants.py @@ -1,4 +1,4 @@ -"""Setu 插件的常量和映射配置。""" +"""Setu 插件的领域常量。""" from __future__ import annotations diff --git a/src/domain/enums.py b/src/domain/enums.py new file mode 100644 index 0000000..b282128 --- /dev/null +++ b/src/domain/enums.py @@ -0,0 +1,55 @@ +"""Setu 插件的领域枚举。""" + +from __future__ import annotations + +from enum import Enum + + +class ContentMode(str, Enum): + """内容分级模式。""" + + SFW = "sfw" + R18 = "r18" + MIX = "mix" + + +class SendMode(str, Enum): + """图片发送模式。""" + + IMAGE = "image" + FORWARD = "forward" + AUTO = "auto" + + +class HtmlCardStrategy(str, Enum): + """HTML 卡片策略。""" + + NEVER = "never" + FALLBACK = "fallback" + ALWAYS = "always" + + +class ApiType(str, Enum): + """API 提供商类型。""" + + LOLICON = "lolicon" + ATRI = "atri" + SEXNYAN = "sexnyan" + CUSTOM = "custom" + ALL = "all" + + +class MultiApiStrategy(str, Enum): + """多 API 策略。""" + + ROUND_ROBIN = "round_robin" + RANDOM = "random" + FAILOVER = "failover" + + +class AccessControlMode(str, Enum): + """访问控制模式。""" + + NONE = "none" + BLACKLIST = "blacklist" + WHITELIST = "whitelist" diff --git a/src/domain/exceptions.py b/src/domain/exceptions.py new file mode 100644 index 0000000..9fb0962 --- /dev/null +++ b/src/domain/exceptions.py @@ -0,0 +1,47 @@ +"""Domain exceptions for Setu and Fortune bounded contexts.""" + +from __future__ import annotations + + +class SetuException(Exception): + """Base exception for Setu domain errors.""" + + pass + + +class ProviderError(SetuException): + """Raised when image provider fails.""" + + pass + + +class SendError(SetuException): + """Raised when image sending fails.""" + + pass + + +class AccessDeniedError(SetuException): + """Raised when access control denies a request.""" + + def __init__(self, reason: str = "") -> None: + self.reason = reason + super().__init__(f"Access denied: {reason}" if reason else "Access denied") + + +class ValidationError(SetuException): + """Raised when input validation fails.""" + + pass + + +class FortuneException(Exception): + """Base exception for Fortune domain errors.""" + + pass + + +class FortuneNotFoundError(FortuneException): + """Raised when fortune record is not found.""" + + pass diff --git a/src/domain/fortune/__init__.py b/src/domain/fortune/__init__.py new file mode 100644 index 0000000..7a774a7 --- /dev/null +++ b/src/domain/fortune/__init__.py @@ -0,0 +1,22 @@ +"""Fortune domain entities and value objects.""" + +from __future__ import annotations + +from .entities import ( + FortuneConfig, + FortuneGenerationRequest, + FortuneRecord, + FortuneTheme, + FortuneWeights, +) +from .value_objects import FortuneResult, FortuneSeed + +__all__ = [ + "FortuneConfig", + "FortuneGenerationRequest", + "FortuneRecord", + "FortuneTheme", + "FortuneWeights", + "FortuneResult", + "FortuneSeed", +] diff --git a/src/domain/fortune/entities.py b/src/domain/fortune/entities.py new file mode 100644 index 0000000..cf41ae0 --- /dev/null +++ b/src/domain/fortune/entities.py @@ -0,0 +1,333 @@ +"""Fortune domain entities and value objects. + +Extracts domain logic from FortuneCore into proper DDD entities. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import date + +# Default weights array (corresponding to 0-7 stars) +DEFAULT_WEIGHTS = [0.1, 0.15, 0.2, 0.25, 0.15, 0.12, 0.07, 0.005] + +# Default fortune titles (corresponding to stars 0-7) +DEFAULT_TITLES = ["凶", "末吉", "末小吉", "小吉", "中吉", "吉", "大吉", "超大吉"] + +# Default fortune descriptions (corresponding to stars 0-7) +DEFAULT_MESSAGES = [ + "长夜再暗,火种仍在,转机终会到来。", + "微光不灭,步步向前,黎明就在眼前。", + "心怀希冀,顺流而行,好事悄然靠近。", + "逆境翻篇,机遇迎面,惊喜不期而至。", + "小吉随身,难题化易,幸运与你并肩。", + "吉星高照,所行皆坦,所愿皆如愿。", + "福泽深厚,大吉加身,一路花开有声。", + "七星同耀,奇迹频现,今日万事皆成。", +] + + +@dataclass +class FortuneConfig: + """Fortune feature configuration. + + Attributes: + enabled: Whether fortune feature is enabled + api_type: API type for fortune images + default_tags: Default tags for fortune images + content_mode: Default content mode + allow_user_refresh: Whether users can refresh their own fortune + auto_refresh: Whether to auto-refresh daily + """ + + enabled: bool = True + api_type: str = "inherit" + default_tags: str = "" + content_mode: str = "sfw" + allow_user_refresh: bool = False + auto_refresh: bool = True + + +@dataclass(frozen=True, slots=True) +class FortuneRecord: + """Entity representing a user's fortune record. + + Attributes: + user_id: User identifier + username: User display name + date_str: Date string in ISO format + title: Fortune title (e.g., "大吉") + star_count: Star rating (0-7) + description: Fortune description + extra_message: Extra message (usually empty) + theme_color: Theme color class + image_cached: Whether image is cached + img_url: URL of background image + last_view_date: Last date this fortune was viewed + group_id: Group ID where fortune was generated + """ + + user_id: str + username: str + date_str: str + title: str + star_count: int + description: str + extra_message: str + theme_color: str + image_cached: bool + img_url: str | None + last_view_date: str + group_id: str | None = None + + @property + def max_stars(self) -> int: + """Maximum possible stars.""" + return 7 + + @property + def is_expired(self) -> bool: + """Check if fortune is from a previous day.""" + try: + record_date = date.fromisoformat(self.date_str) + today = date.today() + return record_date < today + except ValueError: + return False + + def with_last_view_date(self, last_view_date: str) -> FortuneRecord: + """Return new record with updated last_view_date. + + Args: + last_view_date: New last view date string + + Returns: + New FortuneRecord with updated last_view_date + """ + return FortuneRecord( + user_id=self.user_id, + username=self.username, + date_str=self.date_str, + title=self.title, + star_count=self.star_count, + description=self.description, + extra_message=self.extra_message, + theme_color=self.theme_color, + image_cached=self.image_cached, + img_url=self.img_url, + last_view_date=last_view_date, + group_id=self.group_id, + ) + + def with_refreshed_data( + self, title: str, star_count: int, description: str, theme_color: str + ) -> FortuneRecord: + """Return new record with refreshed fortune data. + + Args: + title: New title + star_count: New star count + description: New description + theme_color: New theme color + + Returns: + New FortuneRecord with updated data + """ + return FortuneRecord( + user_id=self.user_id, + username=self.username, + date_str=self.date_str, + title=title, + star_count=star_count, + description=description, + extra_message=self.extra_message, + theme_color=theme_color, + image_cached=False, + img_url=None, + last_view_date=self.last_view_date, + group_id=self.group_id, + ) + + @classmethod + def create_new( + cls, + *, + user_id: str, + username: str, + date_str: str, + title: str, + star_count: int, + description: str, + extra_message: str, + theme_color: str, + group_id: str | None = None, + ) -> FortuneRecord: + """Create a new fortune record with defaults for image fields.""" + return cls( + user_id=user_id, + username=username, + date_str=date_str, + title=title, + star_count=star_count, + description=description, + extra_message=extra_message, + theme_color=theme_color, + image_cached=False, + img_url=None, + last_view_date=date_str, + group_id=group_id, + ) + + def with_image_cache(self, img_url: str | None) -> FortuneRecord: + """Return new record with image cache marked as cached. + + Args: + img_url: Image URL + + Returns: + New FortuneRecord with image_cached=True + """ + return FortuneRecord( + user_id=self.user_id, + username=self.username, + date_str=self.date_str, + title=self.title, + star_count=self.star_count, + description=self.description, + extra_message=self.extra_message, + theme_color=self.theme_color, + image_cached=True, + img_url=img_url, + last_view_date=self.last_view_date, + group_id=self.group_id, + ) + + +@dataclass(frozen=True, slots=True) +class FortuneWeights: + """Value object for fortune calculation weights. + + Attributes: + weights: Weight array for each star level (0-7) + """ + + weights: tuple[float, ...] = tuple(DEFAULT_WEIGHTS) + + def calculate_star(self) -> int: + """Calculate fortune star based on weights. + + Returns: + Star rating (0-7) + """ + import random + + total_weight = sum(self.weights) + random_value = random.random() * total_weight + current_weight = 0.0 + + for i, w in enumerate(self.weights): + current_weight += w + if random_value <= current_weight: + return i + return 0 + + @classmethod + def default(cls) -> FortuneWeights: + """Get default weights.""" + return cls() + + +@dataclass(frozen=True, slots=True) +class FortuneTheme: + """Value object for fortune theme configuration. + + Attributes: + titles: Title array for each star level + messages: Message array for each star level + extra_message: Default extra message + """ + + titles: tuple[str, ...] = tuple(DEFAULT_TITLES) + messages: tuple[str, ...] = tuple(DEFAULT_MESSAGES) + extra_message: str = "" + + def get_title(self, star_count: int) -> str: + """Get title for star count. + + Args: + star_count: Star rating + + Returns: + Title string + """ + idx = min(star_count, len(self.titles) - 1) + return self.titles[idx] + + def get_message(self, star_count: int) -> str: + """Get message for star count. + + Args: + star_count: Star rating + + Returns: + Message string + """ + idx = min(star_count, len(self.messages) - 1) + return self.messages[idx] + + def get_theme_color(self, star_count: int) -> str: + """Get theme color class for star count. + + Args: + star_count: Star rating + + Returns: + Theme color class name + """ + if star_count in (7, 6): + return "theme-red" + elif star_count in (5, 4): + return "theme-gold" + elif star_count in (1, 0): + return "theme-gray" + else: + return "theme-blue" + + @classmethod + def default(cls) -> FortuneTheme: + """Get default theme.""" + return cls() + + +@dataclass +class FortuneGenerationRequest: + """Value object representing a fortune generation request. + + Attributes: + user_id: User identifier + username: User display name + date_str: Date string (ISO format) + group_id: Optional group ID + """ + + user_id: str + username: str + date_str: str + group_id: str | None = None + + @classmethod + def for_today( + cls, user_id: str, username: str, group_id: str | None = None + ) -> FortuneGenerationRequest: + """Create request for today's fortune. + + Args: + user_id: User identifier + username: User display name + group_id: Optional group ID + + Returns: + New FortuneGenerationRequest + """ + today_str = date.today().isoformat() + return cls(user_id, username, today_str, group_id) diff --git a/src/domain/fortune/service.py b/src/domain/fortune/service.py new file mode 100644 index 0000000..d79ff9e --- /dev/null +++ b/src/domain/fortune/service.py @@ -0,0 +1,230 @@ +"""Fortune domain service for generating fortunes. + +Extracts business logic from FortuneCore into a pure domain service. +""" + +from __future__ import annotations + +import logging +from datetime import date + +from ...application.ports import FortuneRepository +from .entities import ( + FortuneGenerationRequest, + FortuneRecord, + FortuneTheme, + FortuneWeights, +) + + +class FortuneService: + """Domain service for fortune generation and retrieval. + + Provides pure business logic for fortune operations, + separate from persistence and infrastructure concerns. + """ + + def __init__( + self, + repository: FortuneRepository, + weights: FortuneWeights | None = None, + theme: FortuneTheme | None = None, + ) -> None: + """Initialize fortune service. + + Args: + repository: Fortune repository for persistence + weights: Fortune weights (uses default if None) + theme: Fortune theme (uses default if None) + """ + self._repo = repository + self._weights = weights or FortuneWeights.default() + self._theme = theme or FortuneTheme.default() + + async def get_or_create_fortune( + self, request: FortuneGenerationRequest, force_refresh: bool = False + ) -> FortuneRecord: + """Get or create fortune for user. + + Args: + request: Fortune generation request + force_refresh: If True, generate new fortune even if one exists + + Returns: + FortuneRecord + + Raises: + FortuneException: If fortune generation fails + """ + # Try to get existing fortune + existing = await self._repo.get_today_fortune(request) + if existing and not force_refresh: + updated = existing.with_last_view_date(request.date_str) + await self._repo.save_fortune(updated) + return updated + + # Generate new fortune + star_count = self._weights.calculate_star() + title = self._theme.get_title(star_count) + description = self._theme.get_message(star_count) + theme_color = self._theme.get_theme_color(star_count) + extra_message = self._theme.extra_message + + record = FortuneRecord.create_new( + user_id=request.user_id, + username=request.username, + date_str=request.date_str, + title=title, + star_count=star_count, + description=description, + extra_message=extra_message, + theme_color=theme_color, + group_id=request.group_id, + ) + + await self._repo.save_fortune(record) + return record + + async def refresh_fortune(self, request: FortuneGenerationRequest) -> FortuneRecord: + """Refresh user's fortune (generate new). + + Args: + request: Fortune generation request + + Returns: + New FortuneRecord + """ + # Delete existing fortune first + await self._repo.delete_fortune(request.user_id, request.date_str) + + # Delete cached image + await self._repo.delete_cached_image(request.user_id, request.date_str) + + # Generate new fortune + return await self.get_or_create_fortune(request, force_refresh=True) + + async def refresh_group_fortunes( + self, group_id: str, date_str: str | None = None + ) -> int: + """Refresh all fortunes for a group. + + Args: + group_id: Group identifier + date_str: Date string (uses today if None) + + Returns: + Number of fortunes refreshed + """ + if date_str is None: + date_str = date.today().isoformat() + + count = await self._repo.delete_group_fortunes(group_id, date_str) + + # Note: This deletes records but doesn't regenerate them. + # Regeneration happens when users request their fortune again. + + return count + + async def refresh_all_fortunes(self, date_str: str | None = None) -> int: + """Refresh all fortunes for a given date. + + Args: + date_str: Date string (uses today if None) + + Returns: + Number of fortunes deleted + """ + if date_str is None: + date_str = date.today().isoformat() + + return await self._repo.delete_all_fortunes(date_str) + + async def pregenerate_active_users(self, days: int = 3) -> int: + """Pregenerate fortunes for active users. + + Args: + days: Number of days to look back for activity + + Returns: + Number of fortunes pregenerated + """ + today_str = date.today().isoformat() + active_users = await self._repo.get_active_users(days) + + pregenerated = 0 + for user_id in active_users: + # Check if already has fortune today + request = FortuneGenerationRequest(user_id, "指挥官", today_str) + existing = await self._repo.get_today_fortune(request) + if not existing: + # Generate new fortune + await self.get_or_create_fortune(request) + pregenerated += 1 + + return pregenerated + + async def update_image_cache( + self, record: FortuneRecord, image_data: bytes, img_url: str | None + ) -> FortuneRecord: + """Update image cache for fortune record. + + Args: + record: Fortune record + image_data: Image bytes + img_url: Image URL + + Returns: + Updated FortuneRecord with image_cached=True + """ + await self._repo.save_cached_image( + record.user_id, record.date_str, image_data, img_url + ) + return record.with_image_cache(img_url) + + async def get_cached_image(self, user_id: str, date_str: str) -> bytes | None: + """Get cached image data. + + Args: + user_id: User identifier + date_str: Date string + + Returns: + Image bytes if cache exists, None otherwise + """ + cache_path = await self._repo.get_cached_image_path(user_id, date_str) + if cache_path: + try: + import asyncio + + return await asyncio.to_thread(cache_path.read_bytes) + except OSError as e: + logging.debug("Failed to read cached image: %s", e) + return None + + async def cleanup_cache(self, date_str: str | None = None) -> int: + """Clean up expired cache files. + + Args: + date_str: Date string (uses today if None) + + Returns: + Number of files deleted + """ + if date_str is None: + date_str = date.today().isoformat() + + return await self._repo.cleanup_expired_cache(date_str) + + def format_stars(self, star_count: int, max_count: int = 7) -> str: + """Format star display. + + Args: + star_count: Number of stars + max_count: Maximum stars + + Returns: + Formatted star string (e.g., "★★★★☆☆☆☆") + """ + filled = "★" * star_count + empty = "☆" * (max_count - star_count) + return filled + empty diff --git a/src/domain/fortune/value_objects.py b/src/domain/fortune/value_objects.py new file mode 100644 index 0000000..45ed53c --- /dev/null +++ b/src/domain/fortune/value_objects.py @@ -0,0 +1,41 @@ +"""Fortune domain value objects.""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import date + + +@dataclass(frozen=True) +class FortuneSeed: + """Value object identifying a fortune request.""" + + user_id: str + date_str: str + + @classmethod + def for_today(cls, user_id: str) -> FortuneSeed: + """Create seed for today's fortune.""" + return cls(user_id, date.today().isoformat()) + + @property + def cache_key(self) -> str: + """Return cache key for this fortune.""" + return f"{self.user_id}_{self.date_str}" + + +@dataclass(frozen=True) +class FortuneResult: + """Value object containing fortune generation result.""" + + seed: FortuneSeed + title: str + star_count: int + description: str + extra_message: str + theme_color: str + + @property + def max_stars(self) -> int: + """Maximum possible stars.""" + return 7 diff --git a/src/domain/setu/__init__.py b/src/domain/setu/__init__.py new file mode 100644 index 0000000..743acee --- /dev/null +++ b/src/domain/setu/__init__.py @@ -0,0 +1,8 @@ +"""Setu domain rules and value objects.""" + +from __future__ import annotations + +from .tag_resolver import TagResolverService +from .value_objects import SetuRequest + +__all__ = ["SetuRequest", "TagResolverService"] diff --git a/src/domain/setu/tag_resolver.py b/src/domain/setu/tag_resolver.py new file mode 100644 index 0000000..ed49dc1 --- /dev/null +++ b/src/domain/setu/tag_resolver.py @@ -0,0 +1,150 @@ +"""Domain service for tag alias resolution. + +Extracted from SafetyConfigMixin.resolve_tags() to separate domain logic +from infrastructure/config layer. +""" + +from __future__ import annotations + + +class TagResolverService: + """Domain service for resolving tag aliases. + + Provides pure functions for tag normalization and alias lookup. + """ + + DEFAULT_TAG_ALIAS: dict[str, list[str]] = { + "萝莉": ["loli", "roricon"], + "少女": ["girl", "girls"], + "猫耳": ["cat_ears", "nekomimi"], + "狗耳": ["dog_ears", "inumimi"], + "长发": ["long_hair"], + "短发": ["short_hair"], + "双马尾": ["twintails", "twin_tails"], + "丝袜": ["pantyhose", "stockings"], + "白丝": ["white_stockings", "white_pantyhose"], + "黑丝": ["black_stockings", "black_pantyhose"], + "泳装": ["swimsuit", "mizugi"], + "校服": ["school_uniform", "seifuku"], + } + + def __init__(self, alias_map: dict[str, list[str]] | None = None) -> None: + """Initialize tag resolver. + + Args: + alias_map: Custom tag alias mapping (canonical -> [aliases]). + If None, uses DEFAULT_TAG_ALIAS. + """ + self._alias_map = alias_map if alias_map is not None else self.DEFAULT_TAG_ALIAS + + def resolve_tags(self, raw_tags: str) -> list[str]: + """Resolve and normalize tag string to canonical tag names. + + Args: + raw_tags: Comma or space-separated tag string + + Returns: + List of canonical tag names + """ + if not raw_tags: + return [] + + # Normalize separators + normalized = raw_tags.replace(",", ",").replace(" ", ",") + raw_list = [t.strip() for t in normalized.split(",") if t.strip()] + + return [self._resolve_single_tag(tag) for tag in raw_list] + + def _resolve_single_tag(self, tag: str) -> str: + """Resolve a single tag to its canonical name. + + Args: + tag: Tag name (may be alias or canonical) + + Returns: + Canonical tag name, or original tag if not found in alias map + """ + canonical = self._find_canonical_tag(tag) + return canonical if canonical else tag + + def _find_canonical_tag(self, tag: str) -> str | None: + """Find canonical tag name from alias map. + + Args: + tag: Tag name to look up + + Returns: + Canonical tag name, or None if not found + """ + normalized = tag.lower() + + for canonical, aliases in self._alias_map.items(): + if not isinstance(canonical, str): + continue + + if normalized == canonical.lower(): + return canonical + + if isinstance(aliases, list): + for alias in aliases: + if isinstance(alias, str) and normalized == alias.lower(): + return canonical + + return None + + @classmethod + def parse_alias_map_from_string(cls, alias_str: str) -> dict[str, list[str]]: + """Parse alias map from config string format. + + Format: "canonical=alias1,alias2" (one per line, # or ; for comments) + + Args: + alias_str: Raw alias string from config + + Returns: + Parsed alias map + """ + if not alias_str or not isinstance(alias_str, str): + return {} + + result: dict[str, list[str]] = {} + lines = alias_str.strip().replace("\r\n", "\n").split("\n") + + for line in lines: + line = line.strip() + + # Skip comments and empty lines + if not line or line.startswith(("#", ";")): + continue + + if "=" not in line: + continue + + key, value = line.split("=", 1) + key = key.strip() + value = value.strip() + + if not key or not value: + continue + + aliases = [a.strip() for a in value.split(",") if a.strip()] + if aliases: + result[key] = aliases + + return result + + def get_alias_map(self) -> dict[str, list[str]]: + """Get the current alias map. + + Returns: + Copy of the alias map + """ + return dict(self._alias_map) + + def update_alias_map(self, new_map: dict[str, list[str]]) -> None: + """Update the alias map. + + Args: + new_map: New alias map to use + """ + self._alias_map = dict(new_map) if new_map else {} diff --git a/src/domain/setu/value_objects.py b/src/domain/setu/value_objects.py new file mode 100644 index 0000000..638026f --- /dev/null +++ b/src/domain/setu/value_objects.py @@ -0,0 +1,26 @@ +"""Setu domain value objects.""" + +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class SetuRequest: + """Value object representing a Setu image request.""" + + count: int + tags: tuple[str, ...] + r18: bool + exclude_ai: bool + + @classmethod + def from_user_input( + cls, count: int, tags: list[str], r18: bool, exclude_ai: bool = True + ) -> SetuRequest: + """Create from user input, normalizing tags to tuple.""" + return cls(count, tuple(tags), r18, exclude_ai) + + def with_tags(self, new_tags: list[str]) -> SetuRequest: + """Return a new SetuRequest with different tags.""" + return SetuRequest(self.count, tuple(new_tags), self.r18, self.exclude_ai) diff --git a/src/infrastructure/__init__.py b/src/infrastructure/__init__.py new file mode 100644 index 0000000..5445792 --- /dev/null +++ b/src/infrastructure/__init__.py @@ -0,0 +1,69 @@ +"""Infrastructure layer — external concerns and technical details. + +Contains implementations for: +- Sending strategies (direct, forward, HTML card) +- Persistence (repositories, file I/O) +- External API clients +- Caching +- Permission checking +""" + +from __future__ import annotations + +from ..shared import get_logger +from .permission_service import PermissionService +from .persistence import ( + FileBackedAccessControlRepo, + JsonSessionConfigRepository, + SQLiteFortuneRepo, + clear_fortune_repo, + clear_repo, + clear_session_config_repo, + get_access_control_repo, + get_fortune_repo, + get_session_config_repo, + init_access_control_repo, + init_fortune_repo, + init_session_config_repo, +) +from .providers import ( + clear_provider, + get_provider, + init_provider, + init_provider_from_config, +) +from .sending import ( + DirectSendStrategy, + ForwardSendStrategy, + HtmlCardFallbackStrategy, + ImageSender, + resolve_send_mode, +) + +__all__ = [ + "ImageSender", + "DirectSendStrategy", + "ForwardSendStrategy", + "HtmlCardFallbackStrategy", + "resolve_send_mode", + "FileBackedAccessControlRepo", + "JsonSessionConfigRepository", + "SQLiteFortuneRepo", + "PermissionService", + # Singleton getters + "get_provider", + "init_provider", + "init_provider_from_config", + "clear_provider", + "get_access_control_repo", + "init_access_control_repo", + "clear_repo", + "get_fortune_repo", + "init_fortune_repo", + "clear_fortune_repo", + "get_session_config_repo", + "init_session_config_repo", + "clear_session_config_repo", + # Logger + "get_logger", +] diff --git a/src/infrastructure/astrbot/__init__.py b/src/infrastructure/astrbot/__init__.py new file mode 100644 index 0000000..d309ba0 --- /dev/null +++ b/src/infrastructure/astrbot/__init__.py @@ -0,0 +1,21 @@ +"""AstrBot infrastructure adapters.""" + +from __future__ import annotations + +from .config import ( + clear_config, + get_config, + get_plugin_context, + init_config, + set_config, + set_plugin_context, +) + +__all__ = [ + "clear_config", + "get_config", + "get_plugin_context", + "init_config", + "set_config", + "set_plugin_context", +] diff --git a/src/infrastructure/astrbot/commands/__init__.py b/src/infrastructure/astrbot/commands/__init__.py new file mode 100644 index 0000000..5a77a3e --- /dev/null +++ b/src/infrastructure/astrbot/commands/__init__.py @@ -0,0 +1,43 @@ +"""AstrBot command adapters.""" + +from __future__ import annotations + +from .fortune import ( + FortuneCommandHandler, +) +from .fortune import ( + register_llm_tools as register_fortune_llm_tools, +) +from .fortune import ( + unregister_llm_tools as unregister_fortune_llm_tools, +) +from .session_config import ( + SessionConfigCommandHandler, +) +from .session_config import ( + register_llm_tools as register_session_config_llm_tools, +) +from .session_config import ( + unregister_llm_tools as unregister_session_config_llm_tools, +) +from .setu import ( + SetuCommandHandler, +) +from .setu import ( + register_llm_tools as register_setu_llm_tools, +) +from .setu import ( + unregister_llm_tools as unregister_setu_llm_tools, +) + +__all__ = [ + "FortuneCommandHandler", + "SessionConfigCommandHandler", + "SetuCommandHandler", + "register_fortune_llm_tools", + "register_session_config_llm_tools", + "register_setu_llm_tools", + "unregister_fortune_llm_tools", + "unregister_session_config_llm_tools", + "unregister_setu_llm_tools", +] diff --git a/src/infrastructure/astrbot/commands/fortune.py b/src/infrastructure/astrbot/commands/fortune.py new file mode 100644 index 0000000..81dea93 --- /dev/null +++ b/src/infrastructure/astrbot/commands/fortune.py @@ -0,0 +1,434 @@ +"""Fortune command handler - all Fortune-related commands.""" + +from __future__ import annotations + +from collections.abc import AsyncGenerator +from typing import Any + +from astrbot.api.event import AstrMessageEvent +from astrbot.core.provider.register import llm_tools + +from ....domain.access_control import AccessPolicy +from ....domain.fortune import ( + FortuneGenerationRequest, + FortuneRecord, +) +from ....domain.fortune.service import FortuneService +from ....shared import get_logger +from ... import get_access_control_repo +from ...permission_service import PermissionService +from ...persistence import get_fortune_repo +from ..config import get_config + +logger = get_logger() + +# Regex pattern directly in decorator (not in constants) +FORTUNE_REGEX_PATTERN = r"^(?!/)(今日运势|jrys)$" + + +class FortuneCommandHandler: + """Handles all Fortune-related commands. + + Uses singleton pattern for config, fortune repo, and access control repo. + Commands are auto-registered by AstrBot decorators. + """ + + def __init__(self) -> None: + pass + + # ==================== Command Handlers ==================== + + async def fortune_command( + self, event: AstrMessageEvent + ) -> AsyncGenerator[Any, None]: + """Handle /今日运势 command (/今日运势, /jrys).""" + config = get_config() + if not config: + yield event.result("配置未加载") + return + + has_perm, msg = await self._check_access(event, config) + if not has_perm: + yield event.result(msg) + return + + request = self._build_fortune_request(event) + + try: + repo = get_fortune_repo() + service = FortuneService(repo=repo) + result = await service.get_or_create_fortune(request) + yield event.result(self._format_fortune(result)) + except Exception as e: + yield event.result(f"获取运势失败: {e}") + + async def refresh_fortune_command( + self, event: AstrMessageEvent + ) -> AsyncGenerator[Any, None]: + """Handle /刷新今日运势 command.""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + yield event.result(msg) + return + + config = get_config() + if not config: + yield event.result("配置未加载") + return + + request = self._build_fortune_request(event) + + try: + repo = get_fortune_repo() + service = FortuneService(repo=repo) + result = await service.refresh_fortune(request) + yield event.result(self._format_fortune(result)) + except Exception as e: + yield event.result(f"刷新运势失败: {e}") + + async def refresh_group_fortune_command( + self, event: AstrMessageEvent + ) -> AsyncGenerator[Any, None]: + """Handle /刷新本群今日运势 command.""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + yield event.result(msg) + return + + group_id = event.get_group_id() + if not group_id: + yield event.result("此命令仅支持群聊") + return + + config = get_config() + if not config: + yield event.result("配置未加载") + return + + try: + repo = get_fortune_repo() + service = FortuneService(repo=repo) + refreshed_count = await service.pregenerate_active_users() + yield event.result(f"已刷新本群 {refreshed_count} 位用户的今日运势") + except Exception as e: + yield event.result(f"刷新群运势失败: {e}") + + async def refresh_all_fortune_command( + self, event: AstrMessageEvent + ) -> AsyncGenerator[Any, None]: + """Handle /刷新全局今日运势 command.""" + has_perm, msg = PermissionService.require_super_user(event) + if not has_perm: + yield event.result(msg) + return + + config = get_config() + if not config: + yield event.result("配置未加载") + return + + try: + repo = get_fortune_repo() + service = FortuneService(repo=repo) + refreshed_count = await service.pregenerate_active_users() + yield event.result(f"已刷新全局 {refreshed_count} 位用户的今日运势") + except Exception as e: + yield event.result(f"刷新全局运势失败: {e}") + + async def enable_fortune_group_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """Handle /开启运势 command (enable Fortune for current group).""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + yield event.result(msg) + return + + group_id = event.get_group_id() + if not group_id: + yield event.result("此命令仅支持群聊") + return + + repo = get_access_control_repo() + await repo.remove_fortune_blocked_group(str(group_id)) + yield event.result("运势功能已开启") + + async def disable_fortune_group_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """Handle /关闭运势 command (disable Fortune for current group).""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + yield event.result(msg) + return + + group_id = event.get_group_id() + if not group_id: + yield event.result("此命令仅支持群聊") + return + + repo = get_access_control_repo() + await repo.add_fortune_blocked_group(str(group_id)) + yield event.result("运势功能已关闭") + + async def block_fortune_user_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """Handle /拉黑运势用户 command (add user to Fortune blacklist).""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + yield event.result(msg) + return + + target_id = args.strip() or event.get_sender_id() + if not target_id: + yield event.result("请指定用户ID") + return + + repo = get_access_control_repo() + await repo.add_fortune_blocked_user(str(target_id)) + yield event.result(f"用户 {target_id} 已添加到运势黑名单") + + async def unblock_fortune_user_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """Handle /解除运势拉黑 command (remove user from Fortune blacklist).""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + yield event.result(msg) + return + + target_id = args.strip() + if not target_id: + yield event.result("请指定用户ID") + return + + repo = get_access_control_repo() + await repo.remove_fortune_blocked_user(str(target_id)) + yield event.result(f"用户 {target_id} 已从运势黑名单移除") + + async def trust_fortune_user_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """Handle /信任运势用户 command (add user to Fortune whitelist).""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + yield event.result(msg) + return + + target_id = args.strip() or event.get_sender_id() + if not target_id: + yield event.result("请指定用户ID") + return + + repo = get_access_control_repo() + await repo.add_fortune_whitelist_user(str(target_id)) + yield event.result(f"用户 {target_id} 已添加到运势白名单") + + async def untrust_fortune_user_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """Handle /取消运势信任 command (remove user from Fortune whitelist).""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + yield event.result(msg) + return + + target_id = args.strip() + if not target_id: + yield event.result("请指定用户ID") + return + + repo = get_access_control_repo() + await repo.remove_fortune_whitelist_user(str(target_id)) + yield event.result(f"用户 {target_id} 已从运势白名单移除") + + # ==================== LLM Tool Handlers ==================== + + async def _llm_get_fortune(self, event: AstrMessageEvent) -> str: + """LLM tool handler for getting today's fortune.""" + config = get_config() + if not config: + return "配置未加载" + + has_perm, msg = await self._check_access(event, config) + if not has_perm: + return msg + + request = self._build_fortune_request(event) + + try: + repo = get_fortune_repo() + service = FortuneService(repo=repo) + result = await service.get_or_create_fortune(request) + return f"今日运势: {result.title}, 星级: {result.star_count}/{result.max_stars}" + except Exception as e: + return f"获取运势失败: {e}" + + async def _llm_refresh_fortune(self, event: AstrMessageEvent) -> str: + """LLM tool handler for refreshing today's fortune.""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + return msg + + config = get_config() + if not config: + return "配置未加载" + + request = self._build_fortune_request(event) + + try: + repo = get_fortune_repo() + service = FortuneService(repo=repo) + result = await service.refresh_fortune(request) + return f"今日运势已刷新: {result.title}" + except Exception as e: + return f"刷新运势失败: {e}" + + async def _llm_refresh_group_fortune(self, event: AstrMessageEvent) -> str: + """LLM tool handler for refreshing group fortunes.""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + return msg + + group_id = event.get_group_id() + if not group_id: + return "此命令仅支持群聊" + + config = get_config() + if not config: + return "配置未加载" + + try: + repo = get_fortune_repo() + service = FortuneService(repo=repo) + refreshed_count = await service.pregenerate_active_users() + return f"已刷新本群 {refreshed_count} 位用户的今日运势" + except Exception as e: + return f"刷新群运势失败: {e}" + + async def _llm_refresh_all_fortune(self, event: AstrMessageEvent) -> str: + """LLM tool handler for refreshing all fortunes.""" + has_perm, msg = PermissionService.require_super_user(event) + if not has_perm: + return msg + + config = get_config() + if not config: + return "配置未加载" + + try: + repo = get_fortune_repo() + service = FortuneService(repo=repo) + refreshed_count = await service.pregenerate_active_users() + return f"已刷新全局 {refreshed_count} 位用户的今日运势" + except Exception as e: + return f"刷新全局运势失败: {e}" + + # ==================== Helper Methods ==================== + + async def _check_access(self, event: AstrMessageEvent, config) -> tuple[bool, str]: + """Check if user/group has access to Fortune feature.""" + from ....domain.access_control.service import AccessControlService + + user_id = event.get_sender_id() + group_id = event.get_group_id() + + policy = AccessPolicy.for_session( + user_id=user_id, + group_id=group_id, + user_mode=config.fortune_user_access_control_mode if config else "none", + group_mode=config.fortune_group_access_control_mode if config else "none", + ) + + repo = get_access_control_repo() + service = AccessControlService(repo) + return await service.check_fortune_access(policy) + + def _build_fortune_request( + self, event: AstrMessageEvent + ) -> FortuneGenerationRequest: + """Build a framework-free fortune generation request from an AstrBot event.""" + user_id = str(event.get_sender_id()) + username = user_id + message_obj = getattr(event, "message_obj", None) + sender = getattr(message_obj, "sender", None) + for attr in ("nickname", "name", "user_name"): + value = getattr(sender, attr, None) + if value: + username = str(value) + break + group_id = event.get_group_id() + return FortuneGenerationRequest.for_today( + user_id=user_id, + username=username, + group_id=str(group_id) if group_id else None, + ) + + def _format_fortune(self, result: FortuneRecord) -> str: + """Format fortune result for message sending.""" + return ( + f"🔮 {result.date_str} 运势\n" + f"📊 {result.title}\n" + f"⭐ {'★' * result.star_count}{'☆' * (result.max_stars - result.star_count)}\n" + f"💬 {result.description}" + ) + + +# ==================== LLM Tools Registration ==================== + + +def register_llm_tools() -> None: + """Register Fortune LLM tools.""" + _handler = FortuneCommandHandler() + tools = [ + ( + "get_today_fortune", + _handler._llm_get_fortune, + [], + "Get today's fortune for the user.", + ), + ( + "refresh_my_fortune", + _handler._llm_refresh_fortune, + [], + "Refresh my today's fortune (admin only).", + ), + ( + "refresh_group_fortune", + _handler._llm_refresh_group_fortune, + [], + "Refresh today's fortune for the current group (admin only).", + ), + ( + "refresh_all_fortune", + _handler._llm_refresh_all_fortune, + [], + "Refresh today's fortune for all users (super admin only).", + ), + ] + + for name, handler, args, desc in tools: + try: + llm_tools.add_func(name=name, func_args=args, desc=desc, handler=handler) + tool = llm_tools.get_func(name) + if tool: + tool.handler_module_path = __name__ + except (AttributeError, RuntimeError): + pass + + +def unregister_llm_tools() -> None: + """Unregister Fortune LLM tools.""" + tool_names = [ + "get_today_fortune", + "refresh_my_fortune", + "refresh_group_fortune", + "refresh_all_fortune", + ] + + for name in tool_names: + try: + llm_tools.remove_func(name) + except (AttributeError, RuntimeError): + pass diff --git a/src/infrastructure/astrbot/commands/session_config.py b/src/infrastructure/astrbot/commands/session_config.py new file mode 100644 index 0000000..9d0434a --- /dev/null +++ b/src/infrastructure/astrbot/commands/session_config.py @@ -0,0 +1,395 @@ +"""Unified session configuration command and LLM tools.""" + +from __future__ import annotations + +import json +from collections.abc import AsyncGenerator +from typing import Any + +from astrbot.api.event import AstrMessageEvent +from astrbot.core.provider.register import llm_tools + +from ....application.session_config import ( + SESSION_CONFIG_KEYS, + SessionConfigService, + SessionConfigSnapshot, + get_key_definition, +) +from ....application.session_config.keys import SessionConfigValidationError +from ...permission_service import PermissionService +from ...persistence import get_session_config_repo +from ..session_identity import get_event_session_identity + +LEGACY_CONFIG_TOOL_NAMES = ( + "get_setu_content_mode", + "set_setu_content_mode", + "set_setu_r18_docx_mode", + "set_setu_auto_revoke", + "set_setu_send_mode", + "get_fortune_config", + "set_fortune_config", +) + + +class SessionConfigCommandHandler: + """Handles unified per-session configuration operations.""" + + async def session_config_command( + self, event: AstrMessageEvent, args: str = "" + ) -> AsyncGenerator[Any, None]: + """Handle `/session_config`.""" + args = (args or "").strip() + if not args: + async for result in self._handle_get(event, []): + yield result + return + + action = args.split(maxsplit=1)[0].lower() + if action in ("get", "show", "status"): + tail = args.split(maxsplit=1)[1] if " " in args else "" + async for result in self._handle_get(event, tail.split()): + yield result + return + if action == "set": + async for result in self._handle_set(event, args): + yield result + return + if action == "clear": + async for result in self._handle_clear(event, args): + yield result + return + + yield event.plain_result(self._usage(f"未知子命令:{action}")) + + async def _handle_get( + self, event: AstrMessageEvent, tokens: list[str] + ) -> AsyncGenerator[Any, None]: + try: + key, as_json = _parse_get_tokens(tokens) + snapshot = await _current_snapshot(event) + if as_json: + yield event.plain_result(_json(_snapshot_payload(snapshot, key))) + elif key: + yield event.plain_result(_format_one(snapshot, key)) + else: + yield event.plain_result(_format_all(snapshot)) + except (RuntimeError, ValueError, SessionConfigValidationError) as exc: + yield event.plain_result(str(exc)) + + async def _handle_set( + self, event: AstrMessageEvent, args: str + ) -> AsyncGenerator[Any, None]: + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + yield event.plain_result(msg) + return + + parts = args.split(maxsplit=2) + if len(parts) < 3: + yield event.plain_result( + self._usage("用法:/session_config set ") + ) + return + + key = parts[1].strip() + value = _strip_quotes(parts[2]) + try: + service = SessionConfigService(get_session_config_repo()) + identity = get_event_session_identity(event) + snapshot = await service.set_value( + identity.session_id, + identity.session_type, + key, + value, + identity.display_name, + ) + definition = get_key_definition(key) + effective = snapshot.effective[key] + yield event.plain_result( + f"已设置 {definition.label}({key})为:{_format_value(effective)}" + ) + except (RuntimeError, ValueError, SessionConfigValidationError) as exc: + yield event.plain_result(str(exc)) + + async def _handle_clear( + self, event: AstrMessageEvent, args: str + ) -> AsyncGenerator[Any, None]: + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + yield event.plain_result(msg) + return + + parts = args.split(maxsplit=1) + key = parts[1].strip() if len(parts) > 1 else None + try: + service = SessionConfigService(get_session_config_repo()) + identity = get_event_session_identity(event) + await service.clear( + identity.session_id, + identity.session_type, + key, + identity.display_name, + ) + if key: + definition = get_key_definition(key) + yield event.plain_result( + f"已清除 {definition.label}({key})的会话覆盖" + ) + else: + yield event.plain_result("已清除当前会话的全部配置覆盖") + except (RuntimeError, ValueError, SessionConfigValidationError) as exc: + yield event.plain_result(str(exc)) + + async def _llm_get_session_config( + self, event: AstrMessageEvent, key: str = "" + ) -> str: + """LLM tool handler for reading current session config.""" + try: + key = key.strip() + if key: + get_key_definition(key) + snapshot = await _current_snapshot(event) + return _json({"ok": True, "data": _snapshot_payload(snapshot, key or None)}) + except Exception as exc: + return _json({"ok": False, "message": str(exc)}) + + async def _llm_set_session_config( + self, event: AstrMessageEvent, key: str, value: str + ) -> str: + """LLM tool handler for setting current session config.""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + return _json({"ok": False, "message": msg}) + try: + normalized_key = key.strip() + service = SessionConfigService(get_session_config_repo()) + identity = get_event_session_identity(event) + snapshot = await service.set_value( + identity.session_id, + identity.session_type, + normalized_key, + value, + identity.display_name, + ) + return _json( + { + "ok": True, + "message": "updated", + "key": normalized_key, + "value": snapshot.overrides[normalized_key], + "effective": snapshot.effective[normalized_key], + } + ) + except Exception as exc: + return _json({"ok": False, "message": str(exc)}) + + async def _llm_clear_session_config( + self, event: AstrMessageEvent, key: str = "" + ) -> str: + """LLM tool handler for clearing current session config.""" + has_perm, msg = PermissionService.require_admin(event) + if not has_perm: + return _json({"ok": False, "message": msg}) + try: + service = SessionConfigService(get_session_config_repo()) + identity = get_event_session_identity(event) + normalized_key = key.strip() or None + snapshot = await service.clear( + identity.session_id, + identity.session_type, + normalized_key, + identity.display_name, + ) + return _json( + { + "ok": True, + "message": "cleared", + "key": normalized_key, + "data": snapshot.to_dict(), + } + ) + except Exception as exc: + return _json({"ok": False, "message": str(exc)}) + + def _usage(self, prefix: str = "") -> str: + lines = [ + prefix, + "用法:", + "/session_config get [key] [json]", + "/session_config set ", + "/session_config clear [key]", + "可用 key:", + *[ + f"- {key}: {definition.label}" + for key, definition in SESSION_CONFIG_KEYS.items() + ], + ] + return "\n".join(line for line in lines if line) + + +async def _current_snapshot(event: AstrMessageEvent) -> SessionConfigSnapshot: + service = SessionConfigService(get_session_config_repo()) + identity = get_event_session_identity(event) + return await service.get_snapshot( + identity.session_id, + identity.session_type, + identity.display_name, + ) + + +def _parse_get_tokens(tokens: list[str]) -> tuple[str | None, bool]: + as_json = False + cleaned: list[str] = [] + for token in tokens: + if token.lower() == "json": + as_json = True + else: + cleaned.append(token) + if len(cleaned) > 1: + raise ValueError("用法:/session_config get [key] [json]") + key = cleaned[0] if cleaned else None + if key: + get_key_definition(key) + return key, as_json + + +def _snapshot_payload( + snapshot: SessionConfigSnapshot, key: str | None +) -> dict[str, Any]: + if not key: + return snapshot.to_dict() + definition = get_key_definition(key) + return { + "session_id": snapshot.session_id, + "session_type": snapshot.session_type, + "display_name": snapshot.display_name, + "key": key, + "label": definition.label, + "override": snapshot.overrides.get(key), + "global": snapshot.global_values[key], + "effective": snapshot.effective[key], + } + + +def _format_all(snapshot: SessionConfigSnapshot) -> str: + lines = [ + "当前会话配置", + f"会话:{snapshot.display_name or snapshot.session_id} ({snapshot.session_type})", + f"ID:{snapshot.session_id}", + "", + ] + for key, definition in SESSION_CONFIG_KEYS.items(): + lines.append(f"{definition.label} ({key})") + lines.append(f" 会话覆盖:{_format_override(snapshot.overrides, key)}") + lines.append(f" 全局配置:{_format_value(snapshot.global_values[key])}") + lines.append(f" 生效值:{_format_value(snapshot.effective[key])}") + return "\n".join(lines) + + +def _format_one(snapshot: SessionConfigSnapshot, key: str) -> str: + definition = get_key_definition(key) + return "\n".join( + [ + f"{definition.label} ({key})", + f"会话覆盖:{_format_override(snapshot.overrides, key)}", + f"全局配置:{_format_value(snapshot.global_values[key])}", + f"生效值:{_format_value(snapshot.effective[key])}", + ] + ) + + +def _format_override(overrides: dict[str, Any], key: str) -> str: + if key not in overrides: + return "未设置" + return _format_value(overrides[key]) + + +def _format_value(value: Any) -> str: + if isinstance(value, bool): + return "启用" if value else "禁用" + if value == "": + return "空" + return str(value) + + +def _strip_quotes(value: str) -> str: + value = value.strip() + if len(value) >= 2 and value[0] == value[-1] and value[0] in ("'", '"'): + return value[1:-1] + return value + + +def _json(payload: dict[str, Any]) -> str: + return json.dumps(payload, ensure_ascii=False) + + +def register_llm_tools() -> None: + """Register unified session configuration LLM tools.""" + _remove_tools(LEGACY_CONFIG_TOOL_NAMES) + handler = SessionConfigCommandHandler() + tools = [ + ( + "get_session_config", + handler._llm_get_session_config, + [ + { + "name": "key", + "type": "string", + "description": "Optional config key, e.g. setu.content_mode.", + } + ], + "Get current session config as JSON.", + ), + ( + "set_session_config", + handler._llm_set_session_config, + [ + {"name": "key", "type": "string", "description": "Config key."}, + {"name": "value", "type": "string", "description": "Config value."}, + ], + "Set one current-session config override as JSON result.", + ), + ( + "clear_session_config", + handler._llm_clear_session_config, + [ + { + "name": "key", + "type": "string", + "description": "Optional config key. Empty clears all overrides.", + } + ], + "Clear current-session config override(s) as JSON result.", + ), + ] + + for name, handler_func, args, desc in tools: + try: + llm_tools.add_func( + name=name, func_args=args, desc=desc, handler=handler_func + ) + tool = llm_tools.get_func(name) + if tool: + tool.handler_module_path = __name__ + except (AttributeError, RuntimeError): + pass + + +def unregister_llm_tools() -> None: + """Unregister unified session configuration LLM tools.""" + _remove_tools( + ( + "get_session_config", + "set_session_config", + "clear_session_config", + *LEGACY_CONFIG_TOOL_NAMES, + ) + ) + + +def _remove_tools(names: tuple[str, ...]) -> None: + for name in names: + try: + llm_tools.remove_func(name) + except (AttributeError, RuntimeError): + pass diff --git a/src/infrastructure/astrbot/commands/setu.py b/src/infrastructure/astrbot/commands/setu.py new file mode 100644 index 0000000..50e63c9 --- /dev/null +++ b/src/infrastructure/astrbot/commands/setu.py @@ -0,0 +1,403 @@ +"""Setu command handler - all Setu-related commands.""" + +from __future__ import annotations + +import asyncio +import random +import re +import time +from collections.abc import AsyncGenerator +from typing import Any + +from astrbot.api.event import AstrMessageEvent +from astrbot.core.provider.register import llm_tools + +from ....application.session_config import SessionConfigService +from ....application.setu.get_images import GetSetuImagesUseCase +from ....domain.access_control import AccessPolicy +from ....domain.access_control.service import AccessControlService +from ....domain.setu import SetuRequest +from ....shared import get_logger +from ... import get_access_control_repo, get_provider +from ...persistence import get_session_config_repo +from ..config import get_config +from ..session_identity import get_event_session_identity +from ...providers import init_provider_from_config + +logger = get_logger() + +# Regex pattern directly in decorator (not in constants) +SETU_REGEX_PATTERN = r"^/?(来\s*(.*?)(份|个|张|点))(.*?)(?:福利|色|瑟|涩|塞)?图$" + + +class RateLimiter: + """Simple rate limiter to prevent concurrent requests from same user.""" + + MAX_LOCKS = 1000 + LOCK_TTL = 120 # Auto-release locks after 120s + + def __init__(self) -> None: + self._locks: dict[str, asyncio.Lock] = {} + self._lock_times: dict[str, float] = {} + + async def acquire(self, event: AstrMessageEvent) -> bool: + """Acquire lock for event. Returns True if acquired, False if already processing.""" + key = f"user:{event.get_sender_id()}" + lock = self._locks.setdefault(key, asyncio.Lock()) + + # Auto-release stale locks (safety net for leaked locks) + if lock.locked(): + acquire_time = self._lock_times.get(key, 0) + if time.monotonic() - acquire_time > self.LOCK_TTL: + try: + lock.release() + except RuntimeError: + pass + + if lock.locked(): + return False + await lock.acquire() + self._lock_times[key] = time.monotonic() + return True + + async def release(self, event: AstrMessageEvent) -> None: + """Release lock for event and evict stale entries.""" + key = f"user:{event.get_sender_id()}" + if key in self._locks: + self._locks[key].release() + self._lock_times.pop(key, None) + if len(self._locks) > self.MAX_LOCKS: + stale = [k for k, v in self._locks.items() if not v.locked()] + for k in stale[: len(stale) // 2]: + del self._locks[k] + self._lock_times.pop(k, None) + + +# Module-level rate limiter singleton +_rate_limiter = RateLimiter() + + +class SetuCommandHandler: + """Handles all Setu-related commands. + + Uses singleton pattern for config, provider, and access control repo. + Commands are auto-registered by AstrBot decorators. + """ + + # ==================== Command Handlers ==================== + + async def get_random_picture( + self, event: AstrMessageEvent + ) -> AsyncGenerator[Any, None]: + """Handle natural language setu requests (regex trigger).""" + if not await _rate_limiter.acquire(event): + yield event.plain_result("你有一个请求正在处理中,请稍后再试~") + return + + try: + async for result in self._handle_random_picture_internal(event): + yield result + finally: + await _rate_limiter.release(event) + + async def _handle_random_picture_internal( + self, event: AstrMessageEvent + ) -> AsyncGenerator[Any, None]: + """Internal handler for regex-triggered setu requests.""" + config = get_config() + if not config: + yield event.plain_result("配置未加载") + return + + match = re.match(SETU_REGEX_PATTERN, event.message_str.strip()) + if not match: + return + + has_perm, msg = await self._check_access(event, config) + if not has_perm: + yield event.plain_result(msg) + return + + num_str = match.group(2) + num = self._parse_count(num_str) + max_count = config.max_count or 10 + + if num < 1 or num > max_count: + if num == -1: + yield event.plain_result( + f"数量解析失败,图片数量必须在1-{max_count}之间" + ) + elif num > max_count: + yield event.plain_result(f"一次最多只能获取{max_count}张哦~") + else: + yield event.plain_result(f"图片数量必须在1-{max_count}之间哦~") + return + + tag_str = match.group(4).strip() + tags = tag_str.replace(",", " ").split() if tag_str else [] + + effective_mode = await self._get_effective_content_mode(event) + is_r18 = self._mode_requires_r18(effective_mode) + + try: + async for result in self._fetch_and_send_images( + event, num, tags, is_r18, config + ): + yield result + except asyncio.TimeoutError: + logger.warning("get_random_picture timeout (>60s)") + yield event.plain_result("获取图片超时,网络可能不稳定,请稍后再试。") + except Exception: + logger.exception("get_random_picture failed") + yield event.plain_result("获取图片失败,请稍后再试") + + async def setu_command( + self, event: AstrMessageEvent, count: str = "1", *, tags: str = "" + ) -> AsyncGenerator[Any, None]: + """Handle /setu command. + + Usage: /setu [count] [tags...] + Example: /setu 3 girl cute + """ + if not await _rate_limiter.acquire(event): + yield event.plain_result("你有一个请求正在处理中,请稍后再试~") + return + + try: + async for result in self._handle_setu_command_internal(event, count, tags): + yield result + finally: + await _rate_limiter.release(event) + + async def _handle_setu_command_internal( + self, event: AstrMessageEvent, count: str, tags: str + ) -> AsyncGenerator[Any, None]: + """Internal handler for /setu command.""" + config = get_config() + if not config: + yield event.plain_result("配置未加载") + return + + has_perm, msg = await self._check_access(event, config) + if not has_perm: + yield event.plain_result(msg) + return + + max_count = config.max_count or 10 + num = self._parse_count(count) + extra_tag = "" + + if num == -1: + num = 1 + extra_tag = count + + all_tags = tags + if extra_tag: + all_tags = f"{extra_tag} {all_tags}".strip() + + if num > max_count: + yield event.plain_result(f"一次最多只能获取{max_count}张哦~") + return + + parsed_tags = [t.strip() for t in all_tags.split() if t.strip()] + + effective_mode = await self._get_effective_content_mode(event) + is_r18 = self._mode_requires_r18(effective_mode) + + try: + async for result in self._fetch_and_send_images( + event, num, parsed_tags, is_r18, config + ): + yield result + except asyncio.TimeoutError: + logger.warning("setu command timeout (>60s)") + yield event.plain_result("获取图片超时,网络可能不稳定,请稍后再试。") + except Exception: + logger.exception("setu command failed") + yield event.plain_result("获取图片失败,网络或服务异常,请稍后再试。") + + # ==================== LLM Tool Handlers ==================== + + async def _llm_get_setu_handler( + self, event: AstrMessageEvent, count: int = 1, tags: list[str] | None = None + ) -> str: + """LLM tool handler for getting Setu images.""" + config = get_config() + if not config: + return "配置未加载" + + has_perm, msg = await self._check_access(event, config) + if not has_perm: + return msg + + try: + init_provider_from_config(config) + provider = get_provider() + effective_mode = await self._get_effective_content_mode(event) + request = SetuRequest.from_user_input( + count=count, + tags=tags or [], + r18=self._mode_requires_r18(effective_mode), + exclude_ai=config.exclude_ai, + ) + payload = await provider.fetch_and_download(request) + from ...sending import ImageSender + + sender = ImageSender(config, logger) + async for _ in sender.send_images(payload, event): + pass + return f"Successfully fetched {payload.count} images" + except Exception: + return "获取图片失败,请稍后再试" + + # ==================== Helper Methods ==================== + + async def _check_access(self, event: AstrMessageEvent, config) -> tuple[bool, str]: + """Check if user/group has access to Setu feature.""" + repo = get_access_control_repo() + user_id = event.get_sender_id() + group_id = event.get_group_id() + policy = AccessPolicy.for_session( + user_id=user_id, + group_id=group_id, + user_mode=config.setu_user_access_control_mode, + group_mode=config.setu_group_access_control_mode, + ) + return await AccessControlService(repo).check_setu_access(policy) + + async def _fetch_and_send_images( + self, event: AstrMessageEvent, num: int, tags: list[str], is_r18: bool, config + ) -> AsyncGenerator[Any, None]: + """Fetch images and send to user.""" + init_provider_from_config(config) + provider = get_provider() + use_case = GetSetuImagesUseCase(provider) + + try: + result = await use_case.execute(num, tags, is_r18) + except asyncio.TimeoutError: + logger.warning("image fetch timeout (>60s)") + yield event.plain_result("获取图片超时,网络可能不稳定,请稍后再试。") + return + + if result.notice: + yield event.plain_result(result.notice) + return + + payload = result.payload + if payload is None: + yield event.plain_result("未找到符合要求的图片~") + return + + from ...sending import ImageSender + + sender = ImageSender(config, logger) + async for send_result in sender.send_images(payload, event): + yield send_result + + async def _get_effective_content_mode(self, event: AstrMessageEvent) -> str: + """Get effective content mode for session.""" + config = get_config() + global_mode = (config.content_mode if config else None) or "sfw" + try: + identity = get_event_session_identity(event) + service = SessionConfigService(get_session_config_repo()) + value = await service.get_effective_value( + identity.session_id, + "setu.content_mode", + identity.session_type, + identity.display_name, + ) + return str(value) + except Exception as exc: + logger.debug("Failed to read session content mode: %s", exc) + return global_mode + + @staticmethod + def _mode_requires_r18(mode: str) -> bool: + """Resolve content mode to the provider R18 flag.""" + if mode == "r18": + return True + if mode == "mix": + return random.random() > 0.5 + return False + + def _parse_count(self, count_str: str) -> int: + """Parse count from string, handling Chinese numbers.""" + if not count_str: + return 1 + + try: + return int(count_str) + except ValueError: + pass + + chinese_nums = { + "一": 1, + "二": 2, + "三": 3, + "四": 4, + "五": 5, + "六": 6, + "七": 7, + "八": 8, + "九": 9, + "十": 10, + } + if count_str in chinese_nums: + return chinese_nums[count_str] + + if count_str.startswith("十"): + if len(count_str) == 1: + return 10 + try: + return 10 + int(count_str[1]) + except ValueError: + return 10 + + return -1 + + +# ==================== LLM Tools Registration ==================== + + +def register_llm_tools() -> None: + """Register Setu LLM tools.""" + _handler = SetuCommandHandler() + tools = [ + ( + "get_setu_image", + _handler._llm_get_setu_handler, + [ + { + "name": "count", + "type": "integer", + "description": "Number of images.", + }, + {"name": "tags", "type": "array", "items": {"type": "string"}}, + ], + "Fetch random anime images.", + ), + ] + + for name, handler, args, desc in tools: + try: + llm_tools.add_func(name=name, func_args=args, desc=desc, handler=handler) + tool = llm_tools.get_func(name) + if tool: + tool.handler_module_path = __name__ + except (AttributeError, RuntimeError): + pass + + +def unregister_llm_tools() -> None: + """Unregister Setu LLM tools.""" + tool_names = [ + "get_setu_image", + ] + + for name in tool_names: + try: + llm_tools.remove_func(name) + except (AttributeError, RuntimeError): + pass diff --git a/src/infrastructure/astrbot/config.py b/src/infrastructure/astrbot/config.py new file mode 100644 index 0000000..17d7d21 --- /dev/null +++ b/src/infrastructure/astrbot/config.py @@ -0,0 +1,55 @@ +"""Configuration singleton manager. + +Provides module-level singleton access to the plugin configuration, +following the rsshub plugin pattern. +""" + +from __future__ import annotations + +from ...application.settings import clear_application_config, set_application_config +from ...shared.config import SetuPluginConfig + +_config: SetuPluginConfig | None = None +_plugin_context: object | None = None + + +def get_config() -> SetuPluginConfig | None: + """Get the plugin configuration singleton. + + Returns: + The current SetuPluginConfig instance or None if not initialized. + """ + return _config + + +def set_config(config: SetuPluginConfig) -> None: + """Set the plugin configuration singleton.""" + global _config + _config = config + + +def get_plugin_context() -> object | None: + """Get the AstrBot plugin context singleton.""" + return _plugin_context + + +def set_plugin_context(context: object) -> None: + """Set the AstrBot plugin context singleton.""" + global _plugin_context + _plugin_context = context + + +def init_config(config_dict: dict) -> SetuPluginConfig: + """Initialize config from dict. Called once at plugin init.""" + global _config + _config = SetuPluginConfig(**config_dict) + set_application_config(_config) + return _config + + +def clear_config() -> None: + """Clear config singleton (for testing).""" + global _config, _plugin_context + _config = None + _plugin_context = None + clear_application_config() diff --git a/src/infrastructure/astrbot/session_config_api.py b/src/infrastructure/astrbot/session_config_api.py new file mode 100644 index 0000000..580feb1 --- /dev/null +++ b/src/infrastructure/astrbot/session_config_api.py @@ -0,0 +1,115 @@ +"""WebUI API handlers for session configuration management.""" + +from __future__ import annotations + +from typing import Any + +from quart import jsonify, request + +from ...application.session_config import ( + SESSION_CONFIG_KEYS, + SessionConfigService, + get_global_session_config_values, +) +from ..persistence import get_session_config_repo +from .config import get_config + +PLUGIN_NAME = "astrbot_plugin_setu" + + +class SessionConfigApi: + """Quart handlers registered through AstrBot's plugin Web API bridge.""" + + async def list_sessions(self): + """Return session config schema, globals, and all records.""" + try: + service = SessionConfigService(get_session_config_repo()) + snapshots = await service.list_snapshots() + return jsonify( + { + "success": True, + "keys": [item.to_dict() for item in SESSION_CONFIG_KEYS.values()], + "global": get_global_session_config_values(), + "sessions": [snapshot.to_dict() for snapshot in snapshots], + "config_loaded": get_config() is not None, + } + ) + except Exception as exc: + return jsonify({"success": False, "error": str(exc)}), 500 + + async def upsert_session(self): + """Create or replace one session record.""" + try: + payload = await request.get_json() + payload = payload or {} + service = SessionConfigService(get_session_config_repo()) + snapshot = await service.upsert_session( + session_id=str(payload.get("session_id", "")), + session_type=str(payload.get("session_type", "private")), + display_name=str(payload.get("display_name", "")), + overrides=_dict(payload.get("overrides")), + ) + return jsonify({"success": True, "data": snapshot.to_dict()}) + except Exception as exc: + return jsonify({"success": False, "error": str(exc)}), 400 + + async def delete_session(self): + """Delete a session record.""" + try: + payload = await request.get_json() + payload = payload or {} + service = SessionConfigService(get_session_config_repo()) + deleted = await service.delete_session(str(payload.get("session_id", ""))) + return jsonify({"success": True, "deleted": deleted}) + except Exception as exc: + return jsonify({"success": False, "error": str(exc)}), 400 + + async def clear_session(self): + """Clear all overrides or one key for a session while keeping the record.""" + try: + payload = await request.get_json() + payload = payload or {} + service = SessionConfigService(get_session_config_repo()) + key = str(payload.get("key", "")).strip() or None + snapshot = await service.clear( + session_id=str(payload.get("session_id", "")), + session_type=str(payload.get("session_type", "private")), + key=key, + display_name=str(payload.get("display_name", "")), + ) + return jsonify({"success": True, "data": snapshot.to_dict()}) + except Exception as exc: + return jsonify({"success": False, "error": str(exc)}), 400 + + +def register_session_config_web_apis(context: Any) -> None: + """Register WebUI APIs for the sessionConfig page.""" + api = SessionConfigApi() + context.register_web_api( + f"/{PLUGIN_NAME}/session-config", + api.list_sessions, + ["GET"], + "List session config records", + ) + context.register_web_api( + f"/{PLUGIN_NAME}/session-config/upsert", + api.upsert_session, + ["POST"], + "Create or update session config", + ) + context.register_web_api( + f"/{PLUGIN_NAME}/session-config/delete", + api.delete_session, + ["POST"], + "Delete session config", + ) + context.register_web_api( + f"/{PLUGIN_NAME}/session-config/clear", + api.clear_session, + ["POST"], + "Clear session config overrides", + ) + + +def _dict(value: Any) -> dict[str, Any]: + return value if isinstance(value, dict) else {} diff --git a/src/infrastructure/astrbot/session_identity.py b/src/infrastructure/astrbot/session_identity.py new file mode 100644 index 0000000..ece64b5 --- /dev/null +++ b/src/infrastructure/astrbot/session_identity.py @@ -0,0 +1,41 @@ +"""Helpers for mapping AstrBot events to session config identities.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + + +@dataclass(frozen=True, slots=True) +class SessionIdentity: + """Normalized AstrBot session identity.""" + + session_id: str + session_type: str + display_name: str + + +def get_event_session_identity(event: Any) -> SessionIdentity: + """Return the session identity used by per-session config.""" + session_id = _call_or_attr(event, "unified_msg_origin") + if not session_id: + session_id = _call_or_attr(event, "get_session_id") + if not session_id: + raise ValueError("无法获取当前会话 ID") + + group_id = _call_or_attr(event, "get_group_id") + sender_id = _call_or_attr(event, "get_sender_id") + session_type = "group" if group_id else "private" + display_name = str(group_id or sender_id or session_id) + return SessionIdentity( + session_id=str(session_id), + session_type=session_type, + display_name=display_name, + ) + + +def _call_or_attr(obj: Any, name: str) -> Any: + value = getattr(obj, name, None) + if callable(value): + return value() + return value diff --git a/src/infrastructure/permission_service.py b/src/infrastructure/permission_service.py new file mode 100644 index 0000000..ca0dc71 --- /dev/null +++ b/src/infrastructure/permission_service.py @@ -0,0 +1,112 @@ +"""Permission service for checking user roles. + +Centralizes admin/super-user checks that were duplicated across +CommandHandler, LlmHandlers, and Fortune handlers. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +from ..shared import get_logger + +if TYPE_CHECKING: + from astrbot.api.event import AstrMessageEvent + +logger = get_logger() + + +class PermissionService: + """Service for checking user permissions. + + Provides a single place for admin/super-user checks, + eliminating duplication across handlers. + """ + + @staticmethod + def is_admin(event: AstrMessageEvent) -> bool: + """Check if user is admin or super user. + + Args: + event: Message event + + Returns: + True if user is admin or super user + """ + try: + # Check is_admin method + if hasattr(event, "is_admin") and callable(getattr(event, "is_admin")): + if event.is_admin(): + return True + + # Check is_super_user method + if hasattr(event, "is_super_user") and callable( + getattr(event, "is_super_user") + ): + if event.is_super_user(): + return True + + # Check message_obj.sender.role + if hasattr(event, "message_obj"): + msg_obj = event.message_obj + if hasattr(msg_obj, "sender") and hasattr(msg_obj.sender, "role"): + role = msg_obj.sender.role + if role in ("admin", "owner"): + return True + + except AttributeError as e: + logger.debug("Permission check attr error: %s", e) + pass + + return False + + @staticmethod + def is_super_user(event: AstrMessageEvent) -> bool: + """Check if user is super user. + + Args: + event: Message event + + Returns: + True if user is super user + """ + try: + if hasattr(event, "is_super_user") and callable( + getattr(event, "is_super_user") + ): + if event.is_super_user(): + return True + + except AttributeError as e: + logger.debug("Super user check attr error: %s", e) + pass + + return False + + @staticmethod + def require_admin(event: AstrMessageEvent) -> tuple[bool, str]: + """Check admin permission and return result with message. + + Args: + event: Message event + + Returns: + Tuple of (has_permission, error_message) + """ + if PermissionService.is_admin(event): + return True, "" + return False, "❌ 权限不足:此命令仅限管理员或超级管理员使用。" + + @staticmethod + def require_super_user(event: AstrMessageEvent) -> tuple[bool, str]: + """Check super user permission and return result with message. + + Args: + event: Message event + + Returns: + Tuple of (has_permission, error_message) + """ + if PermissionService.is_super_user(event): + return True, "" + return False, "❌ 权限不足:此命令仅限超级管理员使用。" diff --git a/src/infrastructure/persistence/__init__.py b/src/infrastructure/persistence/__init__.py new file mode 100644 index 0000000..eec6c52 --- /dev/null +++ b/src/infrastructure/persistence/__init__.py @@ -0,0 +1,127 @@ +"""Persistence layer — data access and storage. + +Contains repository implementations for: +- Access control (blacklist/whitelist) +- Session configuration +- Fortune data +- Cache management +""" + +from __future__ import annotations + +from pathlib import Path +from typing import TYPE_CHECKING + +from .access_control_repo import FileBackedAccessControlRepo +from .session_config_json_repository import JsonSessionConfigRepository +from .sqlite_fortune_repository import SQLiteFortuneRepo + +if TYPE_CHECKING: + from astrbot.core import AstrBotConfig + +__all__ = [ + "FileBackedAccessControlRepo", + "JsonSessionConfigRepository", + "SQLiteFortuneRepo", + "get_access_control_repo", + "init_access_control_repo", + "clear_repo", + "get_fortune_repo", + "init_fortune_repo", + "clear_fortune_repo", + "get_session_config_repo", + "init_session_config_repo", + "clear_session_config_repo", +] + +# ==================== Singleton Pattern ==================== + +_repo: FileBackedAccessControlRepo | None = None +_fortune_repo: SQLiteFortuneRepo | None = None +_session_config_repo: JsonSessionConfigRepository | None = None + + +def get_access_control_repo() -> FileBackedAccessControlRepo: + """Get the access control repository singleton. + + Returns: + The current FileBackedAccessControlRepo instance. + + Raises: + RuntimeError: If repo not initialized. + """ + if _repo is None: + raise RuntimeError( + "Access control repo not initialized. Call init_access_control_repo() first." + ) + return _repo + + +async def init_access_control_repo( + data_dir: Path, astrbot_config: AstrBotConfig | None = None +) -> FileBackedAccessControlRepo: + """Initialize access control repository singleton. + + Args: + data_dir: Plugin data directory. + astrbot_config: AstrBot config dict for WebUI sync. + + Returns: + The initialized repository instance. + """ + global _repo + _repo = FileBackedAccessControlRepo(data_dir, astrbot_config) + await _repo.initialize() + return _repo + + +def clear_repo() -> None: + """Clear repo singleton (for testing).""" + global _repo + _repo = None + + +def get_fortune_repo() -> SQLiteFortuneRepo: + """Get the fortune repository singleton.""" + if _fortune_repo is None: + raise RuntimeError( + "Fortune repo not initialized. Call init_fortune_repo() first." + ) + return _fortune_repo + + +async def init_fortune_repo(data_dir: Path) -> SQLiteFortuneRepo: + """Initialize fortune repository singleton.""" + global _fortune_repo + _fortune_repo = SQLiteFortuneRepo(data_dir / "fortune") + await _fortune_repo.initialize() + return _fortune_repo + + +def clear_fortune_repo() -> None: + """Clear fortune repo singleton.""" + global _fortune_repo + _fortune_repo = None + + +def get_session_config_repo() -> JsonSessionConfigRepository: + """Get the session configuration repository singleton.""" + if _session_config_repo is None: + raise RuntimeError( + "Session config repo not initialized. Call init_session_config_repo() first." + ) + return _session_config_repo + + +async def init_session_config_repo(data_dir: Path) -> JsonSessionConfigRepository: + """Initialize session configuration repository singleton.""" + global _session_config_repo + _session_config_repo = JsonSessionConfigRepository(data_dir) + await _session_config_repo.initialize() + return _session_config_repo + + +def clear_session_config_repo() -> None: + """Clear session configuration repo singleton.""" + global _session_config_repo + _session_config_repo = None diff --git a/src/infrastructure/persistence/access_control_repo.py b/src/infrastructure/persistence/access_control_repo.py new file mode 100644 index 0000000..42b1750 --- /dev/null +++ b/src/infrastructure/persistence/access_control_repo.py @@ -0,0 +1,311 @@ +"""File-backed repository for access control data. + +Implements AccessControlRepository interface using JSON file persistence. +Extracted from ConfigManager to separate persistence from domain logic. +""" + +from __future__ import annotations + +import asyncio +import json +from pathlib import Path +from typing import TYPE_CHECKING, Any + +from astrbot.api import logger + +from ...application.ports import AccessControlRepository + +if TYPE_CHECKING: + from astrbot.core import AstrBotConfig + + +class FileBackedAccessControlRepo(AccessControlRepository): + """File-backed repository for access control lists. + + Stores blacklist/whitelist data in JSON file with async lock protection. + Syncs with AstrBotConfig for WebUI compatibility. + """ + + SAFETY_LIST_KEYS = ( + "setu_blocked_users", + "setu_whitelist_users", + "setu_blocked_groups", + "setu_whitelist_groups", + "fortune_blocked_users", + "fortune_whitelist_users", + "fortune_blocked_groups", + "fortune_whitelist_groups", + ) + + def __init__( + self, data_dir: Path, astrbot_config: AstrBotConfig | None = None + ) -> None: + """Initialize repository. + + Args: + data_dir: Plugin data directory + astrbot_config: AstrBot config for WebUI sync + """ + self._data_dir = data_dir + self._config_file = data_dir / "config.json" + self._cache: dict[str, Any] = {} + self._astrbot_config = astrbot_config + self._main_config_cache: dict[str, Any] | None = None + self._main_config_cache_mtime: float | None = None + self._main_config_cache_path: Path | None = None + + async def initialize(self) -> None: + """Initialize repository, load existing config.""" + self._load_config() + imported = self._sync_from_astrbot_config() + if not imported: + self._sync_to_astrbot_config() + if not self._config_file.exists(): + await self._save_config() + + def _load_config(self) -> None: + """Load config from file.""" + if not self._config_file.exists(): + self._cache = {} + return + + try: + with open(self._config_file, encoding="utf-8") as f: + self._cache = json.load(f) + except (json.JSONDecodeError, OSError) as e: + logger.warning("Failed to load config file: %s", e) + self._cache = {} + + async def _save_config(self) -> bool: + """Save config to file via executor. + + Returns: + True if save succeeded + """ + try: + self._data_dir.mkdir(parents=True, exist_ok=True) + data = json.dumps(self._cache, ensure_ascii=False, indent=2) + loop = asyncio.get_running_loop() + await loop.run_in_executor(None, self._write_config_file, data) + self._sync_to_astrbot_config() + return True + except (OSError, TypeError) as e: + logger.error("Failed to save config file: %s", e) + return False + + def _write_config_file(self, data: str) -> None: + """Write config data to file (called in thread pool).""" + with open(self._config_file, "w", encoding="utf-8") as f: + f.write(data) + + # Setu user access control + async def add_setu_blocked_user(self, user_id: str) -> bool: + """Add user to Setu blacklist.""" + user_id = str(user_id).strip() + if not user_id: + return False + await self.remove_setu_whitelist_user(user_id) + return await self._add_to_list("setu_blocked_users", user_id) + + async def remove_setu_blocked_user(self, user_id: str) -> bool: + """Remove user from Setu blacklist.""" + return await self._remove_from_list("setu_blocked_users", user_id) + + async def is_setu_user_blocked(self, user_id: str) -> bool: + """Check if user is in Setu blacklist.""" + return self._is_in_list("setu_blocked_users", user_id) + + async def add_setu_whitelist_user(self, user_id: str) -> bool: + """Add user to Setu whitelist.""" + user_id = str(user_id).strip() + if not user_id: + return False + await self.remove_setu_blocked_user(user_id) + return await self._add_to_list("setu_whitelist_users", user_id) + + async def remove_setu_whitelist_user(self, user_id: str) -> bool: + """Remove user from Setu whitelist.""" + return await self._remove_from_list("setu_whitelist_users", user_id) + + async def is_setu_user_whitelisted(self, user_id: str) -> bool: + """Check if user is in Setu whitelist.""" + return self._is_in_list("setu_whitelist_users", user_id) + + # Setu group access control + async def add_setu_blocked_group(self, group_id: str) -> bool: + """Add group to Setu blacklist.""" + return await self._add_to_list("setu_blocked_groups", group_id) + + async def remove_setu_blocked_group(self, group_id: str) -> bool: + """Remove group from Setu blacklist.""" + return await self._remove_from_list("setu_blocked_groups", group_id) + + async def is_setu_group_blocked(self, group_id: str) -> bool: + """Check if group is in Setu blacklist.""" + return self._is_in_list("setu_blocked_groups", group_id) + + async def add_setu_whitelist_group(self, group_id: str) -> bool: + """Add group to Setu whitelist.""" + return await self._add_to_list("setu_whitelist_groups", group_id) + + async def remove_setu_whitelist_group(self, group_id: str) -> bool: + """Remove group from Setu whitelist.""" + return await self._remove_from_list("setu_whitelist_groups", group_id) + + async def is_setu_group_whitelisted(self, group_id: str) -> bool: + """Check if group is in Setu whitelist.""" + return self._is_in_list("setu_whitelist_groups", group_id) + + # Fortune user access control + async def add_fortune_blocked_user(self, user_id: str) -> bool: + """Add user to Fortune blacklist.""" + user_id = str(user_id).strip() + if not user_id: + return False + await self.remove_fortune_whitelist_user(user_id) + return await self._add_to_list("fortune_blocked_users", user_id) + + async def remove_fortune_blocked_user(self, user_id: str) -> bool: + """Remove user from Fortune blacklist.""" + return await self._remove_from_list("fortune_blocked_users", user_id) + + async def is_fortune_user_blocked(self, user_id: str) -> bool: + """Check if user is in Fortune blacklist.""" + return self._is_in_list("fortune_blocked_users", user_id) + + async def add_fortune_whitelist_user(self, user_id: str) -> bool: + """Add user to Fortune whitelist.""" + user_id = str(user_id).strip() + if not user_id: + return False + await self.remove_fortune_blocked_user(user_id) + return await self._add_to_list("fortune_whitelist_users", user_id) + + async def remove_fortune_whitelist_user(self, user_id: str) -> bool: + """Remove user from Fortune whitelist.""" + return await self._remove_from_list("fortune_whitelist_users", user_id) + + async def is_fortune_user_whitelisted(self, user_id: str) -> bool: + """Check if user is in Fortune whitelist.""" + return self._is_in_list("fortune_whitelist_users", user_id) + + # Fortune group access control + async def add_fortune_blocked_group(self, group_id: str) -> bool: + """Add group to Fortune blacklist.""" + return await self._add_to_list("fortune_blocked_groups", group_id) + + async def remove_fortune_blocked_group(self, group_id: str) -> bool: + """Remove group from Fortune blacklist.""" + return await self._remove_from_list("fortune_blocked_groups", group_id) + + async def is_fortune_group_blocked(self, group_id: str) -> bool: + """Check if group is in Fortune blacklist.""" + return self._is_in_list("fortune_blocked_groups", group_id) + + async def add_fortune_whitelist_group(self, group_id: str) -> bool: + """Add group to Fortune whitelist.""" + return await self._add_to_list("fortune_whitelist_groups", group_id) + + async def remove_fortune_whitelist_group(self, group_id: str) -> bool: + """Remove group from Fortune whitelist.""" + return await self._remove_from_list("fortune_whitelist_groups", group_id) + + async def is_fortune_group_whitelisted(self, group_id: str) -> bool: + """Check if group is in Fortune whitelist.""" + return self._is_in_list("fortune_whitelist_groups", group_id) + + # Helper methods + async def _add_to_list(self, key: str, item: str) -> bool: + """Add item to list (normalized on write).""" + current = self._cache.setdefault(key, []) + item_str = str(item).strip() + if not item_str or item_str in current: + return True + current.append(item_str) + return await self._save_config() + + async def _remove_from_list(self, key: str, item: str) -> bool: + """Remove item from list.""" + current = self._cache.get(key, []) + item_str = str(item).strip() + if item_str not in current: + return True + current.remove(item_str) + return await self._save_config() + + def _is_in_list(self, key: str, item: str) -> bool: + """Check if item is in list.""" + current = self._cache.get(key, []) + return str(item).strip() in current + + # WebUI sync methods + def _sync_to_astrbot_config(self) -> None: + """Sync local config to AstrBotConfig for WebUI.""" + if self._astrbot_config is None: + return + + try: + updated = False + + if "safety" not in self._astrbot_config: + self._astrbot_config["safety"] = {} + updated = True + + safety_config = self._astrbot_config["safety"] + if not isinstance(safety_config, dict): + safety_config = {} + self._astrbot_config["safety"] = safety_config + updated = True + + for key in self.SAFETY_LIST_KEYS: + if key not in self._cache: + continue + value = self._cache.get(key) + if not isinstance(value, list): + value = [] + if safety_config.get(key) != value: + safety_config[key] = value + updated = True + + if updated and hasattr(self._astrbot_config, "save_config"): + if callable(getattr(self._astrbot_config, "save_config")): + self._astrbot_config.save_config() + + except Exception as e: + logger.debug("Failed to sync to AstrBot config: %s", e) + + def _sync_from_astrbot_config(self) -> bool: + """Sync from AstrBotConfig to local cache.""" + if self._astrbot_config is None: + return False + + try: + imported = False + updated = False + + safety_config = self._astrbot_config.get("safety", {}) + if not isinstance(safety_config, dict): + return False + + for key in self.SAFETY_LIST_KEYS: + if key not in safety_config: + continue + imported = True + value = safety_config.get(key) + if not isinstance(value, list): + value = [] + value = [str(v).strip() for v in value if str(v).strip()] + if self._cache.get(key) != value: + self._cache[key] = value + updated = True + + if updated: + self._data_dir.mkdir(parents=True, exist_ok=True) + with open(self._config_file, "w", encoding="utf-8") as f: + json.dump(self._cache, f, ensure_ascii=False, indent=2) + + return imported + + except Exception as e: + logger.debug("Failed to sync from AstrBot config: %s", e) + return False diff --git a/src/infrastructure/persistence/session_config_json_repository.py b/src/infrastructure/persistence/session_config_json_repository.py new file mode 100644 index 0000000..b36ca25 --- /dev/null +++ b/src/infrastructure/persistence/session_config_json_repository.py @@ -0,0 +1,133 @@ +"""JSON-backed repository for session configuration overrides.""" + +from __future__ import annotations + +import asyncio +import json +from pathlib import Path +from typing import Any + +from astrbot.api import logger + +from ...application.ports import SessionConfigRepository +from ...application.session_config import ( + SessionConfigRecord, + normalize_config_value, + normalize_session_type, +) + +SESSION_CONFIG_VERSION = 1 + + +class JsonSessionConfigRepository(SessionConfigRepository): + """Persist session overrides in `/session_overrides.json`.""" + + def __init__(self, data_dir: Path) -> None: + self._path = data_dir / "session_overrides.json" + self._lock = asyncio.Lock() + self._sessions: dict[str, SessionConfigRecord] = {} + + async def initialize(self) -> None: + """Load existing records or create an empty store.""" + async with self._lock: + self._path.parent.mkdir(parents=True, exist_ok=True) + if not self._path.exists(): + self._sessions = {} + await self._save_unlocked() + return + self._sessions = await asyncio.to_thread(self._load_from_disk) + + async def list_sessions(self) -> list[SessionConfigRecord]: + """List all stored session records.""" + async with self._lock: + return sorted( + self._sessions.values(), + key=lambda record: ( + record.session_type, + record.display_name, + record.session_id, + ), + ) + + async def get_session(self, session_id: str) -> SessionConfigRecord | None: + """Get a session record by ID.""" + async with self._lock: + return self._sessions.get(str(session_id).strip()) + + async def upsert_session(self, record: SessionConfigRecord) -> SessionConfigRecord: + """Create or replace a session record.""" + normalized = _normalize_record(record) + async with self._lock: + self._sessions[normalized.session_id] = normalized + await self._save_unlocked() + return normalized + + async def delete_session(self, session_id: str) -> bool: + """Delete a session record.""" + async with self._lock: + existed = self._sessions.pop(str(session_id).strip(), None) is not None + if existed: + await self._save_unlocked() + return existed + + def _load_from_disk(self) -> dict[str, SessionConfigRecord]: + try: + data = json.loads(self._path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + logger.warning("Failed to load session overrides: %s", exc) + return {} + + sessions: dict[str, SessionConfigRecord] = {} + for item in data.get("sessions", []): + if not isinstance(item, dict): + continue + try: + record = _record_from_dict(item) + except (TypeError, ValueError) as exc: + logger.warning("Ignoring invalid session override record: %s", exc) + continue + sessions[record.session_id] = record + return sessions + + async def _save_unlocked(self) -> None: + data = { + "version": SESSION_CONFIG_VERSION, + "sessions": [record.to_dict() for record in self._sessions.values()], + } + await asyncio.to_thread(self._write_json, data) + + def _write_json(self, data: dict[str, Any]) -> None: + tmp_path = self._path.with_suffix(".tmp") + tmp_path.write_text( + json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8" + ) + tmp_path.replace(self._path) + + +def _record_from_dict(data: dict[str, Any]) -> SessionConfigRecord: + raw_overrides = ( + data.get("overrides") if isinstance(data.get("overrides"), dict) else {} + ) + overrides = { + str(key).strip(): normalize_config_value(str(key).strip(), value) + for key, value in raw_overrides.items() + } + return SessionConfigRecord( + session_id=str(data.get("session_id", "")).strip(), + session_type=normalize_session_type(str(data.get("session_type", "private"))), + display_name=str(data.get("display_name", "")).strip(), + overrides=overrides, + ) + + +def _normalize_record(record: SessionConfigRecord) -> SessionConfigRecord: + overrides = { + key.strip(): normalize_config_value(key, value) + for key, value in record.overrides.items() + } + return SessionConfigRecord( + session_id=str(record.session_id).strip(), + session_type=normalize_session_type(record.session_type), + display_name=record.display_name.strip(), + overrides=overrides, + ) diff --git a/src/infrastructure/persistence/sqlite_fortune_repository.py b/src/infrastructure/persistence/sqlite_fortune_repository.py new file mode 100644 index 0000000..356eccb --- /dev/null +++ b/src/infrastructure/persistence/sqlite_fortune_repository.py @@ -0,0 +1,320 @@ +"""SQLite implementation for fortune data persistence.""" + +from __future__ import annotations + +import asyncio +import datetime +from pathlib import Path +from typing import Any + +from astrbot.api import logger + +from ...application.ports import FortuneRepository +from ...domain.fortune.entities import FortuneGenerationRequest, FortuneRecord + + +class SQLiteFortuneRepo(FortuneRepository): + """SQLite-backed fortune repository implementation.""" + + def __init__(self, data_dir: Path) -> None: + self._data_dir = data_dir + self._db_path = data_dir / "fortune.db" + self._cache_dir = data_dir / "cache" + self._db_lock = asyncio.Lock() + self._cache_lock = asyncio.Lock() + self._db_inited = False + self._last_cleanup_date: str | None = None + + async def initialize(self) -> None: + """Initialize database and cache directory.""" + self._data_dir.mkdir(parents=True, exist_ok=True) + self._cache_dir.mkdir(parents=True, exist_ok=True) + await self._init_db() + today = datetime.date.today().isoformat() + await self.cleanup_expired_cache(today) + self._db_inited = True + logger.info("[fortune] Database initialized at %s", self._db_path) + + async def _init_db(self) -> None: + """Initialize database tables.""" + import aiosqlite + + async with self._db_lock: + async with aiosqlite.connect(str(self._db_path)) as db: + await db.execute( + """ + CREATE TABLE IF NOT EXISTS fortune_data ( + user_id TEXT NOT NULL, + date_str TEXT NOT NULL, + title TEXT NOT NULL, + stars INTEGER NOT NULL, + desc_text TEXT NOT NULL, + extra TEXT NOT NULL, + theme TEXT NOT NULL, + image_cached INTEGER DEFAULT 0, + img_url TEXT, + last_view_date TEXT, + group_id TEXT, + PRIMARY KEY (user_id, date_str) + ) + """ + ) + await db.commit() + await self._migrate_db(db) + + async def _migrate_db(self, db: Any) -> None: + """Database migration: add missing columns.""" + cursor = await db.execute("PRAGMA table_info(fortune_data)") + columns = await cursor.fetchall() + column_names = [col[1] for col in columns] + + if "img_url" not in column_names: + await db.execute("ALTER TABLE fortune_data ADD COLUMN img_url TEXT") + await db.commit() + + if "last_view_date" not in column_names: + await db.execute("ALTER TABLE fortune_data ADD COLUMN last_view_date TEXT") + await db.commit() + + if "group_id" not in column_names: + await db.execute("ALTER TABLE fortune_data ADD COLUMN group_id TEXT") + await db.commit() + await db.execute( + "CREATE INDEX IF NOT EXISTS idx_fortune_group ON fortune_data(group_id, date_str)" + ) + await db.commit() + + def _row_to_record( + self, row: tuple, user_id: str, username: str, date_str: str + ) -> FortuneRecord: + """Convert a database row to FortuneRecord.""" + return FortuneRecord( + user_id=user_id, + username=username, + date_str=date_str, + title=row[0], + star_count=row[1], + description=row[2], + extra_message=row[3], + theme_color=row[4], + image_cached=bool(row[5]), + img_url=row[6], + last_view_date=row[7] if len(row) > 7 else date_str, + group_id=row[8] if len(row) > 8 else None, + ) + + async def get_today_fortune( + self, request: FortuneGenerationRequest + ) -> FortuneRecord | None: + """Get fortune for today.""" + import aiosqlite + + async with self._db_lock: + async with aiosqlite.connect(str(self._db_path)) as db: + cursor = await db.execute( + "SELECT title, stars, desc_text, extra, theme, image_cached, img_url, last_view_date, group_id " + "FROM fortune_data WHERE user_id = ? AND date_str = ?", + (request.user_id, request.date_str), + ) + row = await cursor.fetchone() + if row: + return self._row_to_record( + row, request.user_id, request.username, request.date_str + ) + return None + + async def save_fortune(self, record: FortuneRecord) -> bool: + """Save or update fortune record.""" + import aiosqlite + + try: + async with self._db_lock: + async with aiosqlite.connect(str(self._db_path)) as db: + await db.execute( + "INSERT OR REPLACE INTO fortune_data " + "(user_id, date_str, title, stars, desc_text, extra, theme, image_cached, img_url, last_view_date, group_id) " + "VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", + ( + record.user_id, + record.date_str, + record.title, + record.star_count, + record.description, + record.extra_message, + record.theme_color, + int(record.image_cached), + record.img_url, + record.last_view_date, + record.group_id, + ), + ) + await db.commit() + return True + except Exception as exc: + logger.warning("[fortune] Failed to save fortune: %s", exc) + return False + + async def delete_fortune(self, user_id: str, date_str: str) -> bool: + """Delete fortune record.""" + import aiosqlite + + try: + async with self._db_lock: + async with aiosqlite.connect(str(self._db_path)) as db: + await db.execute( + "DELETE FROM fortune_data WHERE user_id = ? AND date_str = ?", + (user_id, date_str), + ) + await db.commit() + return True + except Exception as exc: + logger.warning("[fortune] Failed to delete fortune: %s", exc) + return False + + async def delete_group_fortunes(self, group_id: str, date_str: str) -> int: + """Delete all fortune records for a group on a given date.""" + import aiosqlite + + async with self._db_lock: + async with aiosqlite.connect(str(self._db_path)) as db: + cursor = await db.execute( + "SELECT COUNT(*) FROM fortune_data WHERE date_str = ? AND (group_id = ? OR group_id IS NULL)", + (date_str, group_id), + ) + row = await cursor.fetchone() + count = row[0] if row else 0 + + await db.execute( + "DELETE FROM fortune_data WHERE date_str = ? AND (group_id = ? OR group_id IS NULL)", + (date_str, group_id), + ) + await db.commit() + + return count + + async def delete_all_fortunes(self, date_str: str) -> int: + """Delete all fortune records for a given date.""" + import aiosqlite + + async with self._db_lock: + async with aiosqlite.connect(str(self._db_path)) as db: + cursor = await db.execute( + "SELECT COUNT(*) FROM fortune_data WHERE date_str = ?", + (date_str,), + ) + row = await cursor.fetchone() + count = row[0] if row else 0 + + await db.execute( + "DELETE FROM fortune_data WHERE date_str = ?", + (date_str,), + ) + await db.commit() + + return count + + async def get_active_users(self, days: int = 3) -> list[str]: + """Get list of active users (viewed fortune within N days).""" + import aiosqlite + + cutoff = (datetime.date.today() - datetime.timedelta(days=days)).isoformat() + async with self._db_lock: + async with aiosqlite.connect(str(self._db_path)) as db: + cursor = await db.execute( + "SELECT DISTINCT user_id FROM fortune_data WHERE last_view_date >= ?", + (cutoff,), + ) + rows = await cursor.fetchall() + return [row[0] for row in rows] + + async def get_cached_image_path(self, user_id: str, date_str: str) -> Any | None: + """Get cached image path for fortune.""" + cache_path = self._cache_dir / f"{user_id}_{date_str}.jpg" + if cache_path.exists(): + return cache_path + return None + + async def save_cached_image( + self, user_id: str, date_str: str, image_data: bytes, img_url: str | None + ) -> Any: + """Save image to cache.""" + cache_path = self._cache_dir / f"{user_id}_{date_str}.jpg" + + async with self._cache_lock: + await asyncio.to_thread(cache_path.write_bytes, image_data) + + import aiosqlite + + async with self._db_lock: + async with aiosqlite.connect(str(self._db_path)) as db: + await db.execute( + "UPDATE fortune_data SET image_cached = 1, img_url = ? " + "WHERE user_id = ? AND date_str = ?", + (img_url, user_id, date_str), + ) + await db.commit() + + return cache_path + + async def delete_cached_image(self, user_id: str, date_str: str) -> bool: + """Delete cached image.""" + cache_path = self._cache_dir / f"{user_id}_{date_str}.jpg" + try: + async with self._cache_lock: + if cache_path.exists(): + cache_path.unlink() + return True + except OSError: + pass + return False + + async def cleanup_expired_cache(self, date_str: str) -> int: + """Clean up cache files from before given date.""" + if self._last_cleanup_date == date_str: + return 0 + + removed = 0 + removed_size = 0 + + try: + async with self._cache_lock: + for file_path in self._cache_dir.iterdir(): + if not file_path.is_file(): + continue + try: + file_name = file_path.stem + parts = file_name.rsplit("_", 1) + if len(parts) >= 2: + file_date = parts[-1] + if file_date != date_str: + file_size = file_path.stat().st_size + file_path.unlink() + removed += 1 + removed_size += file_size + else: + file_size = file_path.stat().st_size + file_path.unlink() + removed += 1 + removed_size += file_size + except (OSError, ValueError): + try: + file_size = file_path.stat().st_size + file_path.unlink() + removed += 1 + removed_size += file_size + except OSError: + pass + + if removed > 0: + logger.info( + "[fortune] Cleaned up %d expired cache files (%.2f MB)", + removed, + removed_size / 1024 / 1024, + ) + + self._last_cleanup_date = date_str + + except OSError as exc: + logger.warning("[fortune] Failed to cleanup cache: %s", exc) + + return removed diff --git a/providers/__init__.py b/src/infrastructure/providers/__init__.py similarity index 56% rename from providers/__init__.py rename to src/infrastructure/providers/__init__.py index bd14c84..f6b50bd 100644 --- a/providers/__init__.py +++ b/src/infrastructure/providers/__init__.py @@ -4,15 +4,16 @@ from typing import Any -from astrbot.api import logger - +from ...application.ports import SetuImageProvider +from ...shared import get_logger from .atri import AtriProvider -from .base import SetuImageProvider from .custom import CustomApiProvider from .lolicon import LoliconProvider from .multi import MultiApiProvider from .sexnyan import SexNyanRunProvider +logger = get_logger() + # 提供商注册表 PROVIDERS: dict[str, type[SetuImageProvider]] = { "lolicon": LoliconProvider, @@ -23,8 +24,98 @@ # 多API策略下的内置提供商列表 BUILTIN_PROVIDERS = ["lolicon", "atri", "sexnyan"] +# ==================== Singleton Pattern ==================== + +_provider: SetuImageProvider | None = None + + +def get_provider() -> SetuImageProvider: + """Get the provider singleton. + + Returns: + The current SetuImageProvider instance. + + Raises: + RuntimeError: If provider not initialized. + """ + if _provider is None: + raise RuntimeError("Provider not initialized. Call init_provider() first.") + return _provider + -def get_provider( +def init_provider( + api_type: str, + custom_config: dict[str, Any] | None = None, + parser_config: dict[str, Any] | None = None, + custom_api_configs: list[dict[str, Any]] | None = None, + multi_api_strategy: str = "round_robin", + lolicon_config: dict[str, Any] | None = None, + atri_config: dict[str, Any] | None = None, +) -> SetuImageProvider: + """Initialize provider singleton. + + Args: + api_type: Provider type ('lolicon', 'atri', 'sexnyan', 'custom', 'all'). + custom_config: Custom API config for 'custom' type. + parser_config: Parser config for 'custom' type. + custom_api_configs: Custom API configs list (template_list format). + multi_api_strategy: Multi-API strategy ('round_robin', 'random', 'failover'). + lolicon_config: Lolicon provider config. + atri_config: Atri provider config. + + Returns: + The initialized provider instance. + """ + global _provider + _provider = _create_provider( + api_type=api_type, + custom_config=custom_config, + parser_config=parser_config, + custom_api_configs=custom_api_configs, + multi_api_strategy=multi_api_strategy, + lolicon_config=lolicon_config, + atri_config=atri_config, + ) + return _provider + + +def init_provider_from_config(config: Any) -> SetuImageProvider: + """Initialize provider singleton from the current plugin config snapshot.""" + provider = init_provider( + api_type=(getattr(config, "api_type", None) or "lolicon"), + custom_api_configs=getattr(config, "custom_api_configs", None) or None, + multi_api_strategy=getattr(config, "multi_api_strategy", "round_robin"), + lolicon_config={ + "image_size": getattr(config, "image_size", None) or "original", + "proxy": getattr(config, "proxy", None) or "", + "aspect_ratio": getattr(config, "aspect_ratio", None) or "", + "uid": getattr(config, "uid", None) or [], + "keyword": getattr(config, "keyword", None) or "", + }, + atri_config={ + "image_size": getattr(config, "atri_image_size", None) or "original", + "proxy": getattr(config, "atri_proxy", None) or "", + "aspect_ratio": getattr(config, "atri_aspect_ratio", None) or "", + "uid": getattr(config, "atri_uid", None) or [], + "keyword": getattr(config, "atri_keyword", None) or "", + }, + ) + logger.info( + "[provider] initialized from config: api_type=%s, lolicon_proxy=%s, atri_proxy=%s", + getattr(config, "api_type", None) or "lolicon", + getattr(config, "proxy", None) or "-", + getattr(config, "atri_proxy", None) or "-", + ) + return provider + + +def clear_provider() -> None: + """Clear provider singleton (for testing).""" + global _provider + _provider = None + + +def _create_provider( api_type: str, custom_config: dict[str, Any] | None = None, parser_config: dict[str, Any] | None = None, @@ -126,4 +217,7 @@ def get_provider( "SexNyanRunProvider", "CustomApiProvider", "get_provider", + "init_provider", + "init_provider_from_config", + "clear_provider", ] diff --git a/providers/atri.py b/src/infrastructure/providers/atri.py similarity index 82% rename from providers/atri.py rename to src/infrastructure/providers/atri.py index 907b7fd..f88c4cd 100644 --- a/providers/atri.py +++ b/src/infrastructure/providers/atri.py @@ -9,10 +9,11 @@ import httpx -from astrbot.api import logger +from ...application.ports import SetuImageProvider +from ...domain import HTTP_TIMEOUT_SECONDS +from ...shared import get_logger -from ..constants import HTTP_TIMEOUT_SECONDS -from .base import SetuImageProvider +logger = get_logger() class AtriProvider(SetuImageProvider): @@ -53,6 +54,15 @@ async def fetch_image_urls( r18: bool, exclude_ai: bool = True, ) -> list[str]: + logger.info( + "[provider] Atri request: count=%d, r18=%s, tags=%s, size=%s, proxy=%s, exclude_ai=%s", + num, + r18, + ",".join(tags) or "-", + self.image_size, + self.proxy or "-", + exclude_ai, + ) exclude_ai_flag = self._normalize_bool(exclude_ai, default=True) query_params: list[tuple[str, str | int]] = [ ("r18", 1 if r18 else 0), @@ -103,4 +113,10 @@ async def fetch_image_urls( img_url = urls_obj.get(self.image_size) if img_url: urls.append(img_url) + urls = self._apply_proxy_to_urls(urls, self.proxy, "AtriProvider") + logger.info( + "[provider] Atri response: requested=%d, returned=%d", + num, + len(urls), + ) return urls diff --git a/providers/custom.py b/src/infrastructure/providers/custom.py similarity index 66% rename from providers/custom.py rename to src/infrastructure/providers/custom.py index ae365c9..149bcf0 100644 --- a/providers/custom.py +++ b/src/infrastructure/providers/custom.py @@ -5,14 +5,120 @@ from __future__ import annotations +import asyncio +import ipaddress from typing import Any +from urllib.parse import quote, urlparse import httpx -from astrbot.api import logger +from ...application.ports import SetuImageProvider +from ...domain import HTTP_TIMEOUT_SECONDS +from ...shared import get_logger + +logger = get_logger() + +# Blocked private/network IP ranges for SSRF protection +_BLOCKED_NETWORKS = [ + ipaddress.ip_network("0.0.0.0/8"), + ipaddress.ip_network("10.0.0.0/8"), + ipaddress.ip_network("100.64.0.0/10"), + ipaddress.ip_network("127.0.0.0/8"), + ipaddress.ip_network("169.254.0.0/16"), + ipaddress.ip_network("172.16.0.0/12"), + ipaddress.ip_network("192.0.0.0/29"), + ipaddress.ip_network("192.168.0.0/16"), + ipaddress.ip_network("198.18.0.0/15"), + ipaddress.ip_network("224.0.0.0/4"), + ipaddress.ip_network("240.0.0.0/4"), + ipaddress.ip_network("::1/128"), + ipaddress.ip_network("fc00::/7"), + ipaddress.ip_network("fe80::/10"), +] + +_BLOCKED_HEADERS = frozenset( + { + "host", + "authorization", + "cookie", + "proxy-authorization", + "x-forwarded-for", + "x-forwarded-host", + "x-forwarded-proto", + "forwarded", + "proxy-connection", + } +) + + +def _sanitize_headers(headers: dict[str, str]) -> dict[str, str]: + """Remove sensitive headers that could enable SSRF or credential leakage.""" + return {k: v for k, v in headers.items() if k.lower() not in _BLOCKED_HEADERS} + + +async def _validate_url(url: str) -> tuple[str, str | None]: + """Validate URL scheme and block private IPs (SSRF protection). + + Resolves DNS, validates IPs, and pins the URL to the resolved IP + to prevent DNS rebinding attacks. + + Args: + url: URL to validate. + + Returns: + Tuple of (ip_pinned_url, original_hostname). + + Raises: + ValueError: If URL scheme is invalid or points to private IP. + """ + parsed = urlparse(url) + if parsed.scheme not in ("http", "https"): + raise ValueError( + f"Invalid URL scheme: {parsed.scheme}. Only http/https allowed." + ) + if not parsed.hostname: + raise ValueError("URL has no hostname.") + + import socket + + try: + loop = asyncio.get_running_loop() + resolved = await loop.run_in_executor( + None, + lambda: socket.getaddrinfo( + parsed.hostname, None, socket.AF_UNSPEC, socket.SOCK_STREAM + ), + ) + except socket.gaierror as e: + raise ValueError(f"Cannot resolve hostname: {parsed.hostname}") from e + + first_ip: str | None = None + for family, _type, _proto, _canon, addr in resolved: + ip = ipaddress.ip_address(addr[0]) + for network in _BLOCKED_NETWORKS: + if ip in network: + raise ValueError(f"URL resolves to blocked private IP: {ip}") + if first_ip is None: + first_ip = str(addr[0]) + + # Pin hostname to resolved IP to prevent DNS rebinding + if first_ip and parsed.hostname != first_ip: + if parsed.port: + pinned_netloc = ( + f"[{first_ip}]:{parsed.port}" + if ":" in first_ip + else f"{first_ip}:{parsed.port}" + ) + else: + pinned_netloc = f"[{first_ip}]" if ":" in first_ip else first_ip + pinned_url = url.replace( + f"{parsed.scheme}://{parsed.netloc}", + f"{parsed.scheme}://{pinned_netloc}", + 1, + ) + return pinned_url, parsed.hostname -from ..constants import HTTP_TIMEOUT_SECONDS -from .base import SetuImageProvider + return url, None class CustomApiProvider(SetuImageProvider): @@ -35,31 +141,31 @@ def __init__( self.api_config = api_config or {} self.parser_config = parser_config or {} - def _build_url( + async def _build_url( self, num: int, tags: list[str], r18: bool, exclude_ai: bool - ) -> tuple[str, dict[str, Any] | None]: - """构建请求 URL 和请求体。""" + ) -> tuple[str, dict[str, Any] | None, str | None]: + """构建请求 URL、请求体和原始主机名。""" url_template = self.api_config.get("url", "") method = self.api_config.get("method", "GET").upper() # 替换 URL 中的占位符 - tags_str = ",".join(tags) if tags else "" + tags_str = quote(",".join(tags)) if tags else "" url = url_template.replace("{num}", str(num)) url = url.replace("{r18}", "1" if r18 else "0") url = url.replace("{tags}", tags_str) + pinned_url, original_host = await _validate_url(url) + if method == "POST": - # POST 请求,参数放 body body = { "num": num, "r18": r18, "tags": tags, "exclude_ai": exclude_ai, } - return url, body + return pinned_url, body, original_host else: - # GET 请求 - return url, None + return pinned_url, None, original_host def _parse_response(self, data: Any) -> list[str]: """解析 API 响应,提取图片 URL 列表。""" @@ -223,14 +329,16 @@ async def fetch_image_urls( exclude_ai: bool = True, ) -> list[str]: """从自定义 API 获取图片 URL 列表。""" - url, body = self._build_url(num, tags, r18, exclude_ai) + url, body, original_host = await self._build_url(num, tags, r18, exclude_ai) if not url: logger.error("自定义 API URL 未配置") return [] timeout = self.api_config.get("timeout", HTTP_TIMEOUT_SECONDS) - headers = self.api_config.get("headers", {}) + headers = _sanitize_headers(self.api_config.get("headers", {})) + if original_host: + headers["Host"] = original_host try: async with httpx.AsyncClient(timeout=timeout) as client: diff --git a/providers/lolicon.py b/src/infrastructure/providers/lolicon.py similarity index 65% rename from providers/lolicon.py rename to src/infrastructure/providers/lolicon.py index c495845..2dae0f9 100644 --- a/providers/lolicon.py +++ b/src/infrastructure/providers/lolicon.py @@ -10,10 +10,11 @@ import httpx -from astrbot.api import logger +from ...application.ports import SetuImageProvider +from ...domain import HTTP_TIMEOUT_SECONDS +from ...shared import get_logger -from ..constants import HTTP_TIMEOUT_SECONDS -from .base import SetuImageProvider +logger = get_logger() class LoliconProvider(SetuImageProvider): @@ -54,29 +55,39 @@ async def fetch_image_urls( r18: bool, exclude_ai: bool = True, ) -> list[str]: + logger.info( + "[provider] Lolicon request: count=%d, r18=%s, tags=%s, size=%s, proxy=%s, exclude_ai=%s", + num, + r18, + ",".join(tags) or "-", + self.image_size, + self.proxy or "-", + exclude_ai, + ) exclude_ai_flag = self._normalize_bool(exclude_ai, default=True) - params: dict[str, str | int | list] = { - "r18": 1 if r18 else 0, - "num": num, - "excludeAI": str(exclude_ai_flag).lower(), - "size": self.image_size, - } + query_params: list[tuple[str, str | int]] = [ + ("r18", 1 if r18 else 0), + ("num", num), + ("excludeAI", str(exclude_ai_flag).lower()), + ("size", self.image_size), + ] # 添加可选参数 if self.proxy: - params["proxy"] = self.proxy + query_params.append(("proxy", self.proxy)) if self.aspect_ratio: - params["aspectRatio"] = self.aspect_ratio - if self.uid: - params["uid"] = self.uid + query_params.append(("aspectRatio", self.aspect_ratio)) if self.keyword: - params["keyword"] = self.keyword + query_params.append(("keyword", self.keyword)) + + # Handle uid as repeated parameters + for uid in self.uid: + if uid is not None: + query_params.append(("uid", uid)) - # 构建带多个 'tag' 参数的 URL(使用 URL 编码防止特殊字符问题) + # 构建 URL — handle repeated tag params separately tag_params = "&".join(f"tag={quote(t, safe='')}" for t in tags) if tags else "" - base_params = "&".join( - f"{k}={quote(str(v), safe='')}" for k, v in params.items() - ) + base_params = "&".join(f"{k}={quote(str(v), safe='')}" for k, v in query_params) url = f"{self.API_URL}?{base_params}" if tag_params: url += f"&{tag_params}" @@ -106,4 +117,10 @@ async def fetch_image_urls( img_url = urls_obj.get(self.image_size) if img_url: urls.append(img_url) + urls = self._apply_proxy_to_urls(urls, self.proxy, "LoliconProvider") + logger.info( + "[provider] Lolicon response: requested=%d, returned=%d", + num, + len(urls), + ) return urls diff --git a/providers/multi.py b/src/infrastructure/providers/multi.py similarity index 95% rename from providers/multi.py rename to src/infrastructure/providers/multi.py index 9aff444..9675511 100644 --- a/providers/multi.py +++ b/src/infrastructure/providers/multi.py @@ -3,15 +3,14 @@ from __future__ import annotations import random -from typing import TYPE_CHECKING -from astrbot.api import logger +from ...application.ports import SetuImageProvider +from ...shared import get_logger -if TYPE_CHECKING: - from .base import SetuImageProvider +logger = get_logger() -class MultiApiProvider: +class MultiApiProvider(SetuImageProvider): """多 API 提供商,支持轮询、随机和故障转移策略。""" def __init__( diff --git a/providers/sexnyan.py b/src/infrastructure/providers/sexnyan.py similarity index 93% rename from providers/sexnyan.py rename to src/infrastructure/providers/sexnyan.py index 57b4227..9260ade 100644 --- a/providers/sexnyan.py +++ b/src/infrastructure/providers/sexnyan.py @@ -10,10 +10,11 @@ import httpx -from astrbot.api import logger +from ...application.ports import SetuImageProvider +from ...domain import HTTP_TIMEOUT_SECONDS +from ...shared import get_logger -from ..constants import HTTP_TIMEOUT_SECONDS -from .base import SetuImageProvider +logger = get_logger() class SexNyanRunProvider(SetuImageProvider): diff --git a/src/infrastructure/sending/__init__.py b/src/infrastructure/sending/__init__.py new file mode 100644 index 0000000..bf2f5c0 --- /dev/null +++ b/src/infrastructure/sending/__init__.py @@ -0,0 +1,36 @@ +"""Sending layer — image sending strategies and implementations.""" + +from __future__ import annotations + +from .dto import SendOptions +from .image_sender import ImageSender +from .send_filters import ( + SendFilter, + SendResult, + direct_send_filter, + forward_send_filter, + html_card_filter, + send_with_filter_chain, +) +from .send_strategies import ( + DirectSendStrategy, + ForwardSendStrategy, + HtmlCardFallbackStrategy, + resolve_send_mode, +) + +__all__ = [ + "ImageSender", + "DirectSendStrategy", + "ForwardSendStrategy", + "HtmlCardFallbackStrategy", + "SendOptions", + "resolve_send_mode", + # Filter chain (new) + "send_with_filter_chain", + "SendResult", + "SendFilter", + "direct_send_filter", + "forward_send_filter", + "html_card_filter", +] diff --git a/src/infrastructure/sending/dto.py b/src/infrastructure/sending/dto.py new file mode 100644 index 0000000..22e9935 --- /dev/null +++ b/src/infrastructure/sending/dto.py @@ -0,0 +1,20 @@ +"""Infrastructure DTOs for sending strategies.""" + +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class SendOptions: + """Value object for send strategy options.""" + + send_mode: str + use_html_card: bool + auto_revoke: bool + revoke_delay: int + r18_docx_mode: bool + html_padding: int = 6 + html_gap: int = 6 + html_card_strategy: str = "fallback" + napcat_stream_mode: str = "fallback" diff --git a/src/infrastructure/sending/image_sender.py b/src/infrastructure/sending/image_sender.py new file mode 100644 index 0000000..014bdc7 --- /dev/null +++ b/src/infrastructure/sending/image_sender.py @@ -0,0 +1,617 @@ +"""Image sender service for adapter-level image delivery.""" + +from __future__ import annotations + +import asyncio +from collections.abc import AsyncGenerator +from pathlib import Path +from typing import Any + +import astrbot.api.message_components as Comp +from astrbot.api.event import AstrMessageEvent + +from ...application.session_config import SessionConfigService +from ...application.setu.dto import ImagePayload +from ...shared import get_logger +from ...shared.send_cache import schedule_send_cache_cleanup +from ..astrbot.config import get_config, get_plugin_context +from ..astrbot.session_identity import get_event_session_identity +from ..persistence import get_session_config_repo +from .dto import SendOptions +from .napcat_stream import upload_file_stream +from .send_strategies import ( + DirectSendStrategy, + ForwardSendStrategy, + HtmlCardFallbackStrategy, + resolve_send_mode, +) + +ImageItem = Path | bytes | Comp.Image +logger = get_logger() + + +class ImageSender: + """Send images with send-mode, stream-upload, and fallback strategies.""" + + def __init__(self, config: Any = None, log: Any = None) -> None: + self._user_config = config + self._log = log or logger + self._html_renderer: Any = None + self._forward_supported_cache: dict[str, bool] = {} + + @property + def _config(self): + return self._user_config or get_config() + + @property + def _context(self): + ctx = get_plugin_context() + if ctx is None: + raise RuntimeError("Plugin context not initialized") + return ctx + + async def _build_options( + self, event: AstrMessageEvent, is_r18: bool = False + ) -> SendOptions: + config = self._config + if not config: + return SendOptions( + send_mode="image", + use_html_card=False, + auto_revoke=False, + revoke_delay=30, + r18_docx_mode=False, + napcat_stream_mode="fallback", + ) + + send_mode = config.send_mode + auto_revoke = config.auto_revoke_r18 if is_r18 else False + r18_docx_mode = config.r18_docx_mode if is_r18 else False + + try: + identity = get_event_session_identity(event) + service = SessionConfigService(get_session_config_repo()) + snapshot = await service.get_snapshot( + identity.session_id, + identity.session_type, + identity.display_name, + ) + send_mode = str(snapshot.effective["setu.send_mode"]) + if is_r18: + auto_revoke = bool(snapshot.effective["setu.auto_revoke"]) + r18_docx_mode = bool(snapshot.effective["setu.r18_docx"]) + except Exception as exc: + self._log.debug( + "[send] failed to apply session overrides: session=%s, error=%s", + getattr(identity, "session_id", "unknown"), + exc, + ) + + html_card_strategy = config.html_card_strategy + return SendOptions( + send_mode=send_mode, + use_html_card=html_card_strategy != "never", + auto_revoke=auto_revoke, + revoke_delay=config.auto_revoke_delay, + r18_docx_mode=r18_docx_mode, + html_padding=config.html_card_padding, + html_gap=config.html_card_gap, + html_card_strategy=html_card_strategy, + napcat_stream_mode=config.napcat_stream_mode, + ) + + def set_html_renderer(self, renderer: Any) -> None: + """Set HTML card renderer.""" + self._html_renderer = renderer + + async def send_images( + self, + payload: ImagePayload, + event: AstrMessageEvent, + ) -> AsyncGenerator[Any, None]: + """Send a fetched image payload to the current AstrBot event.""" + options = await self._build_options(event, payload.r18) + items = self._payload_items(payload) + self._log_send_summary(event, payload, items, options) + if not items: + self._log.warning( + "[send] empty payload: session=%s, tags=%s, urls=%d", + self._session_label(event), + ",".join(payload.tags) or "-", + len(payload.urls), + ) + yield event.plain_result("运气不好,一张图都没拿到...") + return + + if payload.r18 and options.r18_docx_mode: + docx_images = await self._read_image_bytes(items) + docx_yielded = False + if docx_images: + async for result in self._send_r18_docx( + event, docx_images, payload.tags, options + ): + docx_yielded = True + yield result + if docx_yielded: + schedule_send_cache_cleanup() + return + + found_message = self._format_found_message( + payload.count, + options.revoke_delay if options.auto_revoke else None, + ) + chain = self._build_image_chain(items) + + if options.html_card_strategy == "always": + self._log.info( + "[send] html-card-only mode: session=%s, count=%d", + self._session_label(event), + payload.count, + ) + if found_message: + await self._send_plain_text(event, found_message) + success = await self._try_html_card_fallback( + event, chain, options, payload.r18 + ) + if not success: + yield event.plain_result(self._send_failed_message()) + else: + yield {"send_success": True, "image_count": payload.count} + schedule_send_cache_cleanup() + return + + if found_message: + await self._send_plain_text(event, found_message) + + supports_forward = self._is_forward_supported(event) + effective_mode = resolve_send_mode( + options.send_mode, payload.count, supports_forward + ) + self._log.info( + "[send] dispatch start: session=%s, platform=%s, mode=%s, supports_forward=%s, napcat_stream=%s, html_fallback=%s", + self._session_label(event), + self._get_platform_name(event) or "unknown", + effective_mode, + supports_forward, + options.napcat_stream_mode, + options.use_html_card, + ) + send_success = await self._send_chain(event, chain, effective_mode, options) + + if ( + not send_success + and options.napcat_stream_mode == "fallback" + and self._has_local_image_paths(chain) + ): + self._log.warning( + "[send] primary send failed, trying NapCat stream fallback: session=%s, mode=%s", + self._session_label(event), + effective_mode, + ) + stream_chain, changed = await self._stream_upload_chain(event, chain) + if changed: + self._log.info( + "[send] NapCat stream upload rebuilt chain: session=%s, count=%d", + self._session_label(event), + len(stream_chain), + ) + send_success = await self._send_chain( + event, + stream_chain, + effective_mode, + self._without_stream_upload(options), + ) + else: + self._log.warning( + "[send] NapCat stream fallback skipped: session=%s, no files uploaded", + self._session_label(event), + ) + + if not send_success and options.use_html_card: + self._log.warning( + "[send] image send failed, attempting HTML card fallback: session=%s, mode=%s", + self._session_label(event), + effective_mode, + ) + send_success = await self._try_html_card_fallback( + event, chain, options, payload.r18 + ) + + if not send_success: + self._log.error( + "[send] all send strategies failed: session=%s, count=%d, mode=%s, html_strategy=%s, napcat_stream=%s", + self._session_label(event), + payload.count, + effective_mode, + options.html_card_strategy, + options.napcat_stream_mode, + ) + yield event.plain_result(self._send_failed_message()) + else: + self._log.info( + "[send] completed: session=%s, count=%d, mode=%s", + self._session_label(event), + payload.count, + effective_mode, + ) + yield {"send_success": True, "image_count": payload.count} + + schedule_send_cache_cleanup() + + async def _send_chain( + self, + event: AstrMessageEvent, + chain: list[Comp.Image], + effective_mode: str, + options: SendOptions, + ) -> bool: + """Send an already-built image chain.""" + if options.napcat_stream_mode == "always": + self._log.info( + "[send] pre-upload via NapCat stream: session=%s, count=%d", + self._session_label(event), + len(chain), + ) + streamed_chain, changed = await self._stream_upload_chain(event, chain) + if changed: + chain = streamed_chain + + if effective_mode == "forward": + return await ForwardSendStrategy(self._context).send( + event, chain, options.auto_revoke + ) + chain = await self._materialize_local_chain(chain) + return await DirectSendStrategy(self._context).send( + event, chain, options.auto_revoke + ) + + async def _materialize_local_chain( + self, chain: list[Comp.Image] + ) -> list[Comp.Image]: + """Convert readable local-file images to in-memory payloads before send.""" + materialized: list[Comp.Image] = [] + for comp in chain: + file_path = self._local_file_path(comp) + if file_path is None: + materialized.append(comp) + continue + try: + data = await asyncio.to_thread(file_path.read_bytes) + except OSError as exc: + self._log.warning( + "[send] failed to read image before send: path=%s, error=%s", + file_path, + exc, + ) + materialized.append(comp) + continue + self._log.debug("[send] materialized local image: path=%s", file_path) + materialized.append(Comp.Image.fromBytes(data)) + return materialized + + def _without_stream_upload(self, options: SendOptions) -> SendOptions: + """Return options for a retry after stream upload has already run.""" + return SendOptions( + send_mode=options.send_mode, + use_html_card=options.use_html_card, + auto_revoke=options.auto_revoke, + revoke_delay=options.revoke_delay, + r18_docx_mode=options.r18_docx_mode, + html_padding=options.html_padding, + html_gap=options.html_gap, + html_card_strategy=options.html_card_strategy, + napcat_stream_mode="disabled", + ) + + async def _send_r18_docx( + self, + event: AstrMessageEvent, + images: tuple[bytes, ...], + tags: tuple[str, ...], + options: SendOptions, + ) -> AsyncGenerator[Any, None]: + """Send R18 images packaged as DOCX when a docx service is available.""" + docx_service = getattr(self, "_docx_service", None) + if not docx_service: + self._log.debug("[send] docx service unavailable, fallback to image send") + return + + docx_path = docx_service.create_docx_with_images(list(images), tags=list(tags)) + if docx_path: + if options.auto_revoke: + message_id = await self._send_file_with_revoke( + event, str(docx_path), docx_path.name + ) + if message_id: + await self._schedule_revoke(event, message_id, options.revoke_delay) + found_msg = self._format_found_message( + len(images), options.revoke_delay + ) + if found_msg: + yield event.plain_result(found_msg) + return + + found_msg = self._format_found_message(len(images)) + if found_msg: + yield event.plain_result(found_msg) + yield event.chain_result( + [Comp.File(file=str(docx_path), name=docx_path.name)] + ) + else: + yield event.plain_result("R18 Docx 封装失败,请稍后再试或联系管理员。") + + async def _try_html_card_fallback( + self, + event: AstrMessageEvent, + chain: list[Comp.Image], + options: SendOptions, + is_r18: bool, + ) -> bool: + """Try HTML card fallback.""" + if not self._html_renderer: + return False + + strategy = HtmlCardFallbackStrategy( + self._context, + self._html_renderer, + { + "card_padding": options.html_padding, + "card_gap": options.html_gap, + }, + ) + return await strategy.send(event, chain, is_r18 and options.auto_revoke) + + def _payload_items(self, payload: ImagePayload) -> tuple[ImageItem, ...]: + if payload.items: + return tuple(payload.items) + items: list[ImageItem] = [] + items.extend(payload.file_paths) + items.extend(payload.raw_bytes) + return tuple(items) + + def _build_image_chain(self, images: tuple[ImageItem, ...]) -> list[Comp.Image]: + """Build image components from local paths or in-memory bytes.""" + chain: list[Comp.Image] = [] + for item in images: + if isinstance(item, Comp.Image): + chain.append(item) + elif isinstance(item, bytes): + chain.append(Comp.Image.fromBytes(item)) + elif isinstance(item, Path): + chain.append(Comp.Image.fromFileSystem(str(item))) + return chain + + async def _read_image_bytes( + self, images: tuple[ImageItem, ...] + ) -> tuple[bytes, ...]: + """Materialize image items as bytes for DOCX/HTML-only paths.""" + result: list[bytes] = [] + for item in images: + if isinstance(item, bytes): + result.append(item) + elif isinstance(item, Path): + try: + result.append(await asyncio.to_thread(item.read_bytes)) + except OSError as exc: + self._log.warning( + "[send] failed to read cached image: path=%s, error=%s", + item, + exc, + ) + elif isinstance(item, Comp.Image): + try: + file_path = await item.convert_to_file_path() + result.append(await asyncio.to_thread(Path(file_path).read_bytes)) + except Exception as exc: + self._log.warning( + "[send] failed to read image component: file=%s, error=%s", + getattr(item, "file", None), + exc, + ) + return tuple(result) + + async def _stream_upload_chain( + self, event: AstrMessageEvent, chain: list[Comp.Image] + ) -> tuple[list[Comp.Image], bool]: + """Upload local image files through NapCat Stream API and rebuild the chain.""" + changed = False + streamed: list[Comp.Image] = [] + for comp in chain: + file_path = self._local_file_path(comp) + if file_path is None: + streamed.append(comp) + continue + + uploaded_path = await upload_file_stream(event, file_path) + if uploaded_path: + self._log.debug( + "[send] stream upload success: local=%s, remote=%s", + file_path, + uploaded_path, + ) + streamed.append(self._image_from_ref(uploaded_path)) + changed = True + else: + self._log.warning( + "[send] stream upload failed, keep original image: local=%s", + file_path, + ) + streamed.append(comp) + return streamed, changed + + def _has_local_image_paths(self, chain: list[Comp.Image]) -> bool: + return any(self._local_file_path(comp) is not None for comp in chain) + + def _local_file_path(self, comp: Comp.Image) -> Path | None: + path_value = getattr(comp, "path", None) + if path_value: + path = Path(str(path_value)) + if path.exists(): + return path + + file_value = getattr(comp, "file", None) + if not isinstance(file_value, str) or not file_value: + return None + if file_value.startswith("file:///"): + path = Path(file_value[8:]) + else: + path = Path(file_value) + return path if path.exists() else None + + def _image_from_ref(self, ref: str | Path) -> Comp.Image: + text = str(ref) + if text.startswith("file:///"): + return Comp.Image(file=text) + path = Path(text) + if path.exists(): + return Comp.Image.fromFileSystem(str(path)) + return Comp.Image(file=text) + + async def _send_plain_text(self, event: AstrMessageEvent, text: str) -> bool: + """Send plain text message.""" + try: + result = event.plain_result(text) + await self._context.send_message(event.unified_msg_origin, result) + return True + except Exception as exc: + self._log.warning("[send] failed to send plain text: error=%s", exc) + return False + + def _is_forward_supported(self, event: AstrMessageEvent) -> bool: + """Check if platform supports forward messages (cached per platform).""" + platform_name = self._get_platform_name(event) + if platform_name in self._forward_supported_cache: + return self._forward_supported_cache[platform_name] + + supported = self._check_forward_support(platform_name, event) + self._forward_supported_cache[platform_name or ""] = supported + return supported + + def _get_platform_name(self, event: AstrMessageEvent) -> str | None: + """Extract platform name from event.""" + if hasattr(event, "platform") and event.platform: + if hasattr(event.platform, "name"): + return event.platform.name + + if hasattr(event, "get_platform_name"): + try: + return event.get_platform_name() + except Exception: + pass + + return None + + def _check_forward_support( + self, platform_name: str | None, event: AstrMessageEvent + ) -> bool: + """Check forward support from platform info.""" + if not platform_name and hasattr(event, "bot") and event.bot: + if hasattr(event.bot, "call_action"): + return True + + if platform_name: + supported = ( + "aiocqhttp", + "onebot11", + "onebot", + "go-cqhttp", + "napcat", + "llonebot", + ) + return any(p in platform_name.lower() for p in supported) + + return False + + async def _send_with_revoke_support( + self, + event: AstrMessageEvent, + chain: list[Any], + is_group: bool, + target_id: str, + ) -> str | None: + """Send message with revoke support.""" + try: + result = event.chain_result(chain) + send_result = await self._context.send_message( + event.unified_msg_origin, result + ) + if isinstance(send_result, dict): + return send_result.get("message_id") + if isinstance(send_result, str): + return send_result + return None + except Exception as exc: + self._log.warning("[send] revoke-capable send failed: error=%s", exc) + return None + + async def _send_file_with_revoke( + self, event: AstrMessageEvent, file_path: str, file_name: str + ) -> str | None: + """Send a file and return its message id when available.""" + return await self._send_with_revoke_support( + event, + [Comp.File(file=file_path, name=file_name)], + bool(event.get_group_id()), + event.get_group_id() or event.get_sender_id(), + ) + + async def _schedule_revoke( + self, event: AstrMessageEvent, message_id: str, delay: int + ) -> None: + """Schedule message revocation.""" + try: + scheduler = getattr(self, "_revoke_scheduler", None) + if scheduler: + await scheduler.schedule_revoke(event, message_id, delay) + except AttributeError: + self._log.debug("[send] revoke scheduler missing schedule_revoke") + + def _format_found_message( + self, count: int, revoke_delay: int | None = None + ) -> str | None: + """Format found message with optional revoke delay.""" + config = self._config + if config and hasattr(config, "format_found_message"): + return config.format_found_message(count, revoke_delay) + if revoke_delay and revoke_delay > 0: + return f"找到 {count} 张图,将在 {revoke_delay} 秒后撤回" + return f"找到 {count} 张图" + + def _send_failed_message(self) -> str: + config = self._config + if config and getattr(config, "msg_send_failed_text", None): + return str(config.msg_send_failed_text) + return "图片发送失败,请稍后再试。" + + def _log_send_summary( + self, + event: AstrMessageEvent, + payload: ImagePayload, + items: tuple[ImageItem, ...], + options: SendOptions, + ) -> None: + local_paths = sum(isinstance(item, Path) for item in items) + raw_bytes = sum(isinstance(item, bytes) for item in items) + image_components = sum(isinstance(item, Comp.Image) for item in items) + self._log.info( + "[send] payload ready: session=%s, platform=%s, count=%d, r18=%s, tags=%s, items=%d(paths=%d,bytes=%d,components=%d), config_mode=%s, html_strategy=%s, napcat_stream=%s", + self._session_label(event), + self._get_platform_name(event) or "unknown", + payload.count, + payload.r18, + ",".join(payload.tags) or "-", + len(items), + local_paths, + raw_bytes, + image_components, + options.send_mode, + options.html_card_strategy, + options.napcat_stream_mode, + ) + + def _session_label(self, event: AstrMessageEvent) -> str: + group_id = event.get_group_id() + sender_id = event.get_sender_id() + if group_id: + return f"group:{group_id}/user:{sender_id}" + return f"user:{sender_id}" diff --git a/src/infrastructure/sending/napcat_stream.py b/src/infrastructure/sending/napcat_stream.py new file mode 100644 index 0000000..009353c --- /dev/null +++ b/src/infrastructure/sending/napcat_stream.py @@ -0,0 +1,163 @@ +"""NapCat stream upload helpers for local image files.""" + +from __future__ import annotations + +import asyncio +import base64 +import hashlib +import math +import uuid +from pathlib import Path +from typing import Any + +from ...shared import get_logger + +logger = get_logger() + +DEFAULT_STREAM_CHUNK_SIZE = 64 * 1024 +DEFAULT_FILE_RETENTION_MS = 30 * 1000 + + +def _get_bot_client(event: Any) -> Any | None: + return getattr(event, "bot", None) or getattr(event, "_bot", None) + + +def _supports_call_action(bot_client: Any) -> bool: + return ( + hasattr(bot_client, "api") and hasattr(bot_client.api, "call_action") + ) or hasattr(bot_client, "call_action") + + +async def _call_action(bot_client: Any, action: str, params: dict[str, Any]) -> Any: + if hasattr(bot_client, "api") and hasattr(bot_client.api, "call_action"): + return await bot_client.api.call_action(action, **params) + if hasattr(bot_client, "call_action"): + return await bot_client.call_action(action, **params) + return None + + +def _extract_response_data(response: Any) -> dict[str, Any]: + if response is None: + raise RuntimeError("NapCat Stream API 未返回响应") + if not isinstance(response, dict): + raise RuntimeError(f"NapCat Stream API 返回格式异常: {type(response).__name__}") + + status = response.get("status") + if status == "failed": + message = response.get("message") or response.get("wording") or response + raise RuntimeError(f"NapCat Stream API 返回失败: {message}") + retcode = response.get("retcode") + if retcode not in (None, 0): + message = response.get("message") or response.get("wording") or response + raise RuntimeError(f"NapCat Stream API 返回错误: retcode={retcode}, {message}") + + data = response.get("data") + if isinstance(data, dict): + return data + return response + + +def _extract_uploaded_path(response: Any) -> str | None: + data = _extract_response_data(response) + for key in ("file_path", "file", "path"): + value = data.get(key) + if isinstance(value, str) and value.strip(): + return value.strip() + return None + + +def _calculate_sha256(file_path: Path) -> str: + hasher = hashlib.sha256() + with file_path.open("rb") as file: + while True: + chunk = file.read(DEFAULT_STREAM_CHUNK_SIZE) + if not chunk: + break + hasher.update(chunk) + return hasher.hexdigest() + + +async def upload_file_stream( + event: Any, + file_path: str | Path, + *, + chunk_size: int = DEFAULT_STREAM_CHUNK_SIZE, + file_retention_ms: int = DEFAULT_FILE_RETENTION_MS, +) -> str | None: + """Upload one local file through NapCat's stream action.""" + bot_client = _get_bot_client(event) + if not bot_client or not _supports_call_action(bot_client): + logger.debug( + "[send] NapCat stream unavailable: missing bot call_action support" + ) + return None + + path = Path(file_path) + if not path.exists() or not path.is_file(): + logger.warning("[send] NapCat stream skipped: invalid file path=%s", path) + return None + + file_size = path.stat().st_size + if file_size <= 0: + logger.warning("[send] NapCat stream skipped: empty file path=%s", path) + return None + + chunk_size = max(1, int(chunk_size or DEFAULT_STREAM_CHUNK_SIZE)) + total_chunks = max(1, math.ceil(file_size / chunk_size)) + stream_id = str(uuid.uuid4()) + + try: + expected_sha256 = await asyncio.to_thread(_calculate_sha256, path) + current_size = path.stat().st_size + if current_size != file_size: + raise RuntimeError("文件在上传前大小发生变化") + + logger.info( + "[send] NapCat stream upload start: file=%s, size=%d, chunks=%d", + path, + file_size, + total_chunks, + ) + + with path.open("rb") as file: + for chunk_index in range(total_chunks): + chunk = file.read(chunk_size) + if not chunk: + raise RuntimeError("文件在上传过程中提前结束") + response = await _call_action( + bot_client, + "upload_file_stream", + { + "stream_id": stream_id, + "chunk_data": base64.b64encode(chunk).decode("utf-8"), + "chunk_index": chunk_index, + "total_chunks": total_chunks, + "file_size": file_size, + "expected_sha256": expected_sha256, + "filename": path.name, + "file_retention": file_retention_ms, + }, + ) + _extract_response_data(response) + + complete_response = await _call_action( + bot_client, + "upload_file_stream", + {"stream_id": stream_id, "is_complete": True}, + ) + uploaded = _extract_uploaded_path(complete_response) + logger.info( + "[send] NapCat stream upload completed: file=%s, remote=%s", + path, + uploaded, + ) + return uploaded + except asyncio.CancelledError: + raise + except Exception as exc: + logger.exception( + "[send] NapCat stream upload failed: file=%s, error=%s", + path, + exc, + ) + return None diff --git a/src/infrastructure/sending/send_filters.py b/src/infrastructure/sending/send_filters.py new file mode 100644 index 0000000..774dda9 --- /dev/null +++ b/src/infrastructure/sending/send_filters.py @@ -0,0 +1,194 @@ +"""Send filter chain for image sending. + +A simpler alternative to the strategy pattern - each filter tries to send +images and returns a result. The chain is configured based on send_mode. +""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import astrbot.api.message_components as Comp +from astrbot.api.event import AstrMessageEvent + +from ...shared import get_logger + +logger = get_logger() + +SendFilter = Callable[[list[Path], AstrMessageEvent, Any], Awaitable["SendResult"]] + + +@dataclass +class SendResult: + """Result of a send filter attempt.""" + + success: bool + message: str | None = None + images_sent: int = 0 + error: str | None = None + + +async def direct_send_filter( + images: list[Path], event: AstrMessageEvent, config: Any +) -> SendResult: + """Direct send filter - send images directly in message chain. + + Args: + images: List of image file paths + event: Message event + config: Plugin config + + Returns: + SendResult indicating success or failure + """ + try: + for img in images: + chain = [Comp.Image.fromFileSystem(str(img))] + result = event.chain_result(chain) + await event.ctx.send_message(event.unified_msg_origin, result) + + logger.debug("[direct] Sent %d images", len(images)) + return SendResult(success=True, images_sent=len(images)) + except Exception as e: + logger.warning("[direct] Direct send failed: %s", e) + return SendResult(success=False, error=str(e)) + + +async def forward_send_filter( + images: list[Path], event: AstrMessageEvent, config: Any +) -> SendResult: + """Forward send filter - merge images into a single forward message. + + Args: + images: List of image file paths + event: Message event + config: Plugin config + + Returns: + SendResult indicating success or failure + """ + try: + nodes = [] + for img in images: + node = Comp.Node( + uin=event.get_self_id(), + name="色图", + content=[Comp.Image.fromFileSystem(str(img))], + ) + nodes.append(node) + + forward_chain = [Comp.Forward(node) for node in nodes] + result = event.chain_result(forward_chain) + await event.ctx.send_message(event.unified_msg_origin, result) + + logger.debug("[forward] Sent %d images as forward", len(images)) + return SendResult(success=True, images_sent=len(images)) + except Exception as e: + logger.warning("[forward] Forward send failed: %s", e) + return SendResult(success=False, error=str(e)) + + +async def html_card_filter( + images: list[Path], event: AstrMessageEvent, config: Any +) -> SendResult: + """HTML card filter - wrap images in HTML card to bypass censorship. + + Args: + images: List of image file paths + event: Message event + config: Plugin config + + Returns: + SendResult indicating success or failure + """ + try: + html_content = _build_html_card(images, config) + await event.send(html_content) + logger.debug("[html] Sent %d images as HTML card", len(images)) + return SendResult( + success=True, images_sent=len(images), message="HTML卡片发送成功" + ) + except Exception as e: + logger.warning("[html] HTML card send failed: %s", e) + return SendResult(success=False, error=str(e)) + + +def _build_html_card(images: list[Path], config: Any) -> str: + """Build HTML card with images. + + Args: + images: List of image file paths + config: Plugin config + + Returns: + HTML content as string + """ + opts = config.html_card if config else None + padding = opts.card_padding if opts else 10 + + img_tags = [] + for img in images: + uri = img.resolve().as_uri() + img_tags.append( + f'' + ) + + gap_img = '' + + content = gap_img.join(img_tags) + return f'{content}' + + +async def send_with_filter_chain( + images: list[Path], + event: AstrMessageEvent, + config: Any, +) -> str: + """Send images using filter chain based on config. + + Args: + images: List of image file paths + event: Message event + config: Plugin config + + Returns: + Result message string + """ + send_mode = config.delivery.send_mode + if send_mode: + send_mode_str = ( + send_mode.value if hasattr(send_mode, "value") else str(send_mode) + ) + else: + send_mode_str = "auto" + + # Build filter chain based on send mode + if send_mode_str == "forward": + chain = [forward_send_filter, direct_send_filter, html_card_filter] + elif send_mode_str == "image": + chain = [direct_send_filter, html_card_filter] + else: # auto or unknown + chain = [direct_send_filter, html_card_filter] + + # Try filters in order + for filter_func in chain: + result = await filter_func(images, event, config) + if result.success: + return result.message or "发送成功" + + # All filters failed + fail_msg = config.msg_send_failed_text if config else None + return fail_msg or "图片发送失败" + + +__all__ = [ + "SendResult", + "SendFilter", + "direct_send_filter", + "forward_send_filter", + "html_card_filter", + "send_with_filter_chain", +] diff --git a/src/infrastructure/sending/send_strategies.py b/src/infrastructure/sending/send_strategies.py new file mode 100644 index 0000000..7c69f11 --- /dev/null +++ b/src/infrastructure/sending/send_strategies.py @@ -0,0 +1,363 @@ +"""Send strategy pattern for image delivery. + +Defines the strategy interface and implementations for different send modes: +- Direct send: send images directly in message chain +- Forward send: use merge forward (OneBot v11 feature) +- HTML card fallback: wrap images in HTML cards to bypass censorship +""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from pathlib import Path +from typing import Any + +import astrbot.api.message_components as Comp +from astrbot.api.event import AstrMessageEvent + +from ...shared import get_logger + +logger = get_logger() + + +class SendStrategy(ABC): + """Abstract base class for send strategies.""" + + @abstractmethod + async def send( + self, + event: AstrMessageEvent, + chain: list[Any], + auto_revoke: bool = False, + ) -> bool: + """Send message chain using this strategy. + + Args: + event: Message event + chain: Message chain (list of components) + auto_revoke: Whether to schedule auto-revoke after send + + Returns: + True if send succeeded, False otherwise + """ + ... + + +class DirectSendStrategy(SendStrategy): + """Direct send strategy — sends images in a single message chain.""" + + def __init__(self, plugin_context: Any) -> None: + """Initialize direct send strategy. + + Args: + plugin_context: Plugin context for sending messages + """ + self._context = plugin_context + + async def send( + self, + event: AstrMessageEvent, + chain: list[Any], + auto_revoke: bool = False, + ) -> bool: + """Send images directly in message chain. + + Args: + event: Message event + chain: Message chain + auto_revoke: Not supported for direct send (ignored) + + Returns: + True if send succeeded + """ + try: + send_result = await self._send_message(event, chain) + + platform_name = getattr(event.platform, "name", "unknown") + + # Only strict-check return value for OneBot platforms + if send_result is None and platform_name == "aiocqhttp": + logger.warning( + "[send] direct send returned None: platform=%s, chain=%d", + platform_name, + len(chain), + ) + return False + + logger.info( + "[send] direct send completed: platform=%s, chain=%d", + platform_name, + len(chain), + ) + return True + except TimeoutError: + logger.warning( + "[send] direct send timed out: platform=%s, chain=%d", + getattr(event.platform, "name", "unknown"), + len(chain), + ) + return False + except Exception as exc: + logger.exception( + "[send] direct send failed: platform=%s, chain=%d, error=%s", + getattr(event.platform, "name", "unknown"), + len(chain), + exc, + ) + return False + + async def _send_message(self, event: AstrMessageEvent, chain: list[Any]) -> Any: + if self._requires_onebot_passthrough(event, chain): + send_result = await self._send_onebot_image_chain(event, chain) + if send_result is not None: + return send_result + result = event.chain_result(chain) + return await self._context.send_message(event.unified_msg_origin, result) + + def _requires_onebot_passthrough( + self, event: AstrMessageEvent, chain: list[Any] + ) -> bool: + platform_name = getattr(getattr(event, "platform", None), "name", "") or "" + return "aiocqhttp" in platform_name.lower() and any( + self._is_onebot_image_ref(comp) + for comp in chain + if isinstance(comp, Comp.Image) + ) + + def _is_onebot_image_ref(self, comp: Comp.Image) -> bool: + file_value = getattr(comp, "file", None) + return ( + isinstance(file_value, str) + and "://" in file_value + and not ( + file_value.startswith("file:///") + or file_value.startswith("http://") + or file_value.startswith("https://") + or file_value.startswith("base64://") + ) + ) + + async def _send_onebot_image_chain( + self, event: AstrMessageEvent, chain: list[Any] + ) -> Any | None: + bot = getattr(event, "bot", None) + if bot is None: + return None + + message: list[dict[str, Any]] = [] + for comp in chain: + if isinstance(comp, Comp.Image): + message.append( + { + "type": "image", + "data": {"file": str(getattr(comp, "file", ""))}, + } + ) + else: + message.append(comp.toDict()) + + is_group = bool(event.get_group_id()) + session_id = event.get_group_id() if is_group else event.get_sender_id() + if not session_id or not str(session_id).isdigit(): + logger.debug( + "[send] skip OneBot passthrough: invalid session_id=%s", + session_id, + ) + return None + + if is_group: + return await bot.send_group_msg(group_id=int(session_id), message=message) + return await bot.send_private_msg(user_id=int(session_id), message=message) + + +class ForwardSendStrategy(SendStrategy): + """Forward send strategy — uses merge forward (OneBot v11).""" + + def __init__(self, plugin_context: Any) -> None: + """Initialize forward send strategy. + + Args: + plugin_context: Plugin context for sending messages + """ + self._context = plugin_context + + async def send( + self, + event: AstrMessageEvent, + chain: list[Any], + auto_revoke: bool = False, + ) -> bool: + """Send images using merge forward. + + Args: + event: Message event + chain: Message chain (list of image components) + auto_revoke: Not supported (handled by caller) + + Returns: + True if send succeeded + """ + import time + + build_start = time.monotonic() + nodes = [] + for comp in chain: + if isinstance(comp, Comp.Image): + node = Comp.Node( + uin=event.get_self_id(), + name="色图", + content=[comp], + ) + nodes.append(node) + + build_end = time.monotonic() + logger.debug( + "[forward] built nodes: count=%d, elapsed=%.3fs", + len(nodes), + build_end - build_start, + ) + + return await self._send_nodes_direct(event, nodes) + + async def _send_nodes_direct( + self, event: AstrMessageEvent, nodes: list[Comp.Node] + ) -> bool: + """Send forward nodes directly. + + Args: + event: Message event + nodes: Forward nodes + + Returns: + True if send succeeded + """ + try: + forward_chain = [Comp.Forward(node) for node in nodes] + result = event.chain_result(forward_chain) + await self._context.send_message(event.unified_msg_origin, result) + + logger.info("[forward] send completed: nodes=%d", len(nodes)) + return True + except Exception as exc: + logger.exception( + "[forward] send failed: nodes=%d, error=%s", + len(nodes), + exc, + ) + return False + + +class HtmlCardFallbackStrategy(SendStrategy): + """HTML card fallback strategy — wraps images in HTML cards.""" + + def __init__( + self, + plugin_context: Any, + html_renderer: Any, + style_options: dict[str, int] | None = None, + ) -> None: + """Initialize HTML card fallback strategy. + + Args: + plugin_context: Plugin context + html_renderer: HtmlCardRenderer instance + style_options: Style options (card_padding, card_gap) + """ + self._context = plugin_context + self._renderer = html_renderer + self._style_options = style_options or {} + + async def send( + self, + event: AstrMessageEvent, + chain: list[Any], + auto_revoke: bool = False, + ) -> bool: + """Send images wrapped in HTML cards. + + Args: + event: Message event + chain: Message chain (list of image components) + auto_revoke: Not supported (ignored) + + Returns: + True if send succeeded + """ + if not self._renderer: + logger.warning("[html_fallback] renderer unavailable") + return False + + images: list[bytes] = [] + for comp in chain: + if not isinstance(comp, Comp.Image): + continue + if isinstance(comp.file, bytes): + images.append(comp.file) + continue + + path_value = getattr(comp, "path", None) or getattr(comp, "file", None) + if isinstance(path_value, str): + candidate = ( + Path(path_value[8:]) + if path_value.startswith("file:///") + else Path(path_value) + ) + if candidate.exists(): + try: + import asyncio + + data = await asyncio.to_thread(candidate.read_bytes) + images.append(data) + except OSError: + logger.warning( + "[html_fallback] failed to read image path=%s", + candidate, + ) + + if not images: + logger.warning("[html_fallback] no images available after materialization") + return False + + rendered_images = [] + for i, img_data in enumerate(images): + logger.debug("[html_fallback] Rendering image %d/%d", i + 1, len(images)) + rendered = await self._renderer.render_single_image( + context=self._context, + image=img_data, + style_options=self._style_options, + ) + if rendered: + rendered_images.append(rendered) + else: + logger.warning("[html_fallback] failed to render image index=%d", i + 1) + + if not rendered_images: + logger.warning("[html_fallback] renderer produced no images") + return False + + # Send rendered images + chain = [Comp.Image.fromBytes(img) for img in rendered_images] + logger.info("[html_fallback] rendered images: count=%d", len(rendered_images)) + return await DirectSendStrategy(self._context).send(event, chain, False) + + +def resolve_send_mode( + send_mode: str, + image_count: int, + supports_forward: bool = True, +) -> str: + """Resolve effective send mode. + + Args: + send_mode: Configured send mode (image/forward/auto) + image_count: Number of images to send + supports_forward: Whether platform supports forward + + Returns: + Effective send mode (image or forward) + """ + if send_mode == "auto": + return "forward" if image_count > 1 else "image" + if send_mode == "forward" and not supports_forward: + return "image" + return send_mode diff --git a/src/shared/__init__.py b/src/shared/__init__.py new file mode 100644 index 0000000..f4036e3 --- /dev/null +++ b/src/shared/__init__.py @@ -0,0 +1,7 @@ +"""Shared low-level helpers and configuration models.""" + +from __future__ import annotations + +from .logging import PrefixedLogger, clear_logger, get_logger + +__all__ = ["PrefixedLogger", "clear_logger", "get_logger"] diff --git a/src/shared/config/__init__.py b/src/shared/config/__init__.py new file mode 100644 index 0000000..a13b1b6 --- /dev/null +++ b/src/shared/config/__init__.py @@ -0,0 +1,49 @@ +"""Shared configuration models.""" + +from __future__ import annotations + +from .models import ( + AccessControlModeStr, + ApiSectionConfig, + ApiTypeStr, + AtriConfig, + CacheConfig, + ContentModeStr, + CustomApiConfig, + DeliveryConfig, + FortuneConfig, + HtmlCardConfig, + HtmlCardStrategyStr, + LoliconConfig, + MessagesConfig, + MultiApiStrategyStr, + NapcatStreamModeStr, + PerformanceConfig, + ProviderConfig, + SendModeStr, + SetuGeneralConfig, + SetuPluginConfig, +) + +__all__ = [ + "AccessControlModeStr", + "ApiSectionConfig", + "ApiTypeStr", + "AtriConfig", + "CacheConfig", + "ContentModeStr", + "CustomApiConfig", + "DeliveryConfig", + "FortuneConfig", + "HtmlCardConfig", + "HtmlCardStrategyStr", + "LoliconConfig", + "MessagesConfig", + "MultiApiStrategyStr", + "NapcatStreamModeStr", + "PerformanceConfig", + "ProviderConfig", + "SendModeStr", + "SetuGeneralConfig", + "SetuPluginConfig", +] diff --git a/src/shared/config/models.py b/src/shared/config/models.py new file mode 100644 index 0000000..66b48b0 --- /dev/null +++ b/src/shared/config/models.py @@ -0,0 +1,595 @@ +"""Pydantic v2 configuration models for Setu plugin. + +Replaces the ConfigBase + mixin chain with type-safe, validated models. +Each section mirrors the _conf_schema.json structure for WebUI compatibility. +""" + +from __future__ import annotations + +from enum import Enum +from typing import Any + +from pydantic import BaseModel, Field, field_validator + + +class ImageSize(str, Enum): + """Valid image sizes for Lolicon/Atri APIs.""" + + ORIGINAL = "original" + REGULAR = "regular" + SMALL = "small" + THUMB = "thumb" + MINI = "mini" + + +class AspectRatio(str, Enum): + """Valid aspect ratios for filtering.""" + + HORIZONTAL = "horizontal" + VERTICAL = "vertical" + SQUARE = "square" + + +class ContentModeStr(str, Enum): + """Content rating modes.""" + + SFW = "sfw" + R18 = "r18" + MIX = "mix" + + +class ApiTypeStr(str, Enum): + """API provider types.""" + + LOLICON = "lolicon" + ATRI = "atri" + SEXNYAN = "sexnyan" + CUSTOM = "custom" + ALL = "all" + + +class MultiApiStrategyStr(str, Enum): + """Multi-API strategy types.""" + + ROUND_ROBIN = "round_robin" + RANDOM = "random" + FAILOVER = "failover" + + +class SendModeStr(str, Enum): + """Image send modes.""" + + IMAGE = "image" + FORWARD = "forward" + AUTO = "auto" + + +class NapcatStreamModeStr(str, Enum): + """NapCat stream upload modes.""" + + DISABLED = "disabled" + FALLBACK = "fallback" + ALWAYS = "always" + + +class HtmlCardStrategyStr(str, Enum): + """HTML card strategies.""" + + NEVER = "never" + FALLBACK = "fallback" + ALWAYS = "always" + + +class AccessControlModeStr(str, Enum): + """Access control modes.""" + + NONE = "none" + BLACKLIST = "blacklist" + WHITELIST = "whitelist" + + +class ProviderConfig(BaseModel): + """Configuration for a single API provider (Lolicon or Atri).""" + + image_size: ImageSize = ImageSize.ORIGINAL + proxy: str = "i.pixiv.re" + aspect_ratio: AspectRatio | None = None + uid: list[int] = Field(default_factory=list) + keyword: str = "" + exclude_ai: bool = True + + @field_validator("aspect_ratio", mode="before") + @classmethod + def normalize_empty_aspect_ratio(cls, value: Any) -> Any: + """Treat empty-string aspect ratio from schema/config files as None.""" + if value == "": + return None + return value + + +class LoliconConfig(ProviderConfig): + """Lolicon API specific configuration.""" + + pass + + +class AtriConfig(ProviderConfig): + """Atri API specific configuration.""" + + pass + + +class CustomApiConfig(BaseModel): + """Custom API configuration.""" + + name: str = "My custom API" + url: str = "" + method: str = "GET" + timeout: int = 30 + parser_type: str = "auto" + json_path: str = "" + + +class SetuGeneralConfig(BaseModel): + """Setu general configuration.""" + + api_type: ApiTypeStr = ApiTypeStr.LOLICON + multi_api_strategy: MultiApiStrategyStr = MultiApiStrategyStr.ROUND_ROBIN + content_mode: ContentModeStr = ContentModeStr.SFW + max_count: int = Field(default=10, ge=1, le=10) + max_replenish_rounds: int = Field(default=3, ge=1, le=3) + tag_alias: str = "" + + +class DeliveryConfig(BaseModel): + """Image delivery configuration.""" + + send_mode: SendModeStr = SendModeStr.IMAGE + r18_docx_mode: bool = True + auto_handle_send_failure: bool = True + auto_revoke_r18: bool = False + auto_revoke_delay: int = Field(default=30, ge=5, le=300) + napcat_stream_mode: NapcatStreamModeStr = NapcatStreamModeStr.FALLBACK + + +class HtmlCardConfig(BaseModel): + """HTML card wrapping configuration.""" + + strategy: HtmlCardStrategyStr = HtmlCardStrategyStr.FALLBACK + mode: str = "single" + card_padding: int = Field(default=6, ge=0, le=30) + card_gap: int = Field(default=6, ge=0, le=30) + + +class FortuneConfig(BaseModel): + """Fortune (Today's Luck) configuration.""" + + enabled: bool = True + api_type: str = "inherit" + tags: str = "" + content_mode: ContentModeStr = ContentModeStr.SFW + allow_user_refresh: bool = False + auto_refresh: bool = True + + +class CacheConfig(BaseModel): + """Disk cache used before adapter-level image sending.""" + + enabled: bool = True + ttl_hours: int = Field(default=2, ge=1, le=168) + max_items: int = Field(default=200, ge=1, le=1000) + cleanup_on_start: bool = True + + +class PerformanceConfig(BaseModel): + """Performance tuning configuration.""" + + enable_range_download: bool = False + range_segments: int = Field(default=3, ge=2, le=6) + range_download_threshold: int = Field(default=512, ge=256, le=2048) + download_concurrent_limit: int = Field(default=10, ge=1, le=50) + download_timeout_seconds: int = Field(default=30, ge=5, le=120) + + +class MessagesFetchingConfig(BaseModel): + """Fetching message configuration.""" + + enabled: bool = True + text: str = "正在获取图片,请稍候..." + + +class MessagesFoundConfig(BaseModel): + """Found message configuration.""" + + enabled: bool = True + text: str = "找到 {count} 张符合要求的图片~" + + +class MessagesSendFailedConfig(BaseModel): + """Send failed message configuration.""" + + text: str = "图片发送失败,请稍后再试。" + + +class MessagesConfig(BaseModel): + """User-facing message configuration.""" + + fetching: MessagesFetchingConfig = Field(default_factory=MessagesFetchingConfig) + found: MessagesFoundConfig = Field(default_factory=MessagesFoundConfig) + send_failed: MessagesSendFailedConfig = Field( + default_factory=MessagesSendFailedConfig + ) + + +class SafetyConfig(BaseModel): + """Safety and access control configuration.""" + + setu_user_access_control_mode: AccessControlModeStr = AccessControlModeStr.NONE + setu_group_access_control_mode: AccessControlModeStr = AccessControlModeStr.NONE + setu_blocked_users: list[str] = Field(default_factory=list) + setu_whitelist_users: list[str] = Field(default_factory=list) + setu_blocked_groups: list[str] = Field(default_factory=list) + setu_whitelist_groups: list[str] = Field(default_factory=list) + fortune_user_access_control_mode: AccessControlModeStr = AccessControlModeStr.NONE + fortune_group_access_control_mode: AccessControlModeStr = AccessControlModeStr.NONE + fortune_blocked_users: list[str] = Field(default_factory=list) + fortune_whitelist_users: list[str] = Field(default_factory=list) + fortune_blocked_groups: list[str] = Field(default_factory=list) + fortune_whitelist_groups: list[str] = Field(default_factory=list) + + +class SessionTemplateItem(BaseModel): + """Session configuration template item.""" + + session_id: str = "" + session_type: str = "group" + content_mode: str = "" + r18_docx_mode: str = "" + auto_revoke_r18: str = "" + send_mode: str = "" + + +class FortuneSessionTemplateItem(BaseModel): + """Fortune session configuration template item.""" + + session_id: str = "" + session_type: str = "group" + tags: str = "" + content_mode: str = "" + + +class ApiSectionConfig(BaseModel): + """API section configuration.""" + + lolicon: LoliconConfig = Field(default_factory=LoliconConfig) + atri: AtriConfig = Field(default_factory=AtriConfig) + custom_api_configs: list[CustomApiConfig] = Field(default_factory=list) + + +class SetuPluginConfig(BaseModel): + """Root Pydantic model for Setu plugin configuration. + + This replaces the ConfigBase + mixin chain with a single validated model. + The structure mirrors _conf_schema.json for WebUI compatibility. + + The model can be instantiated directly from AstrBotConfig dict: + config = SetuPluginConfig(**astrbot_config) + """ + + setu_general: SetuGeneralConfig = Field(default_factory=SetuGeneralConfig) + delivery: DeliveryConfig = Field(default_factory=DeliveryConfig) + html_card: HtmlCardConfig = Field(default_factory=HtmlCardConfig) + fortune: FortuneConfig = Field(default_factory=FortuneConfig) + cache: CacheConfig = Field(default_factory=CacheConfig) + api: ApiSectionConfig = Field(default_factory=ApiSectionConfig) + messages: MessagesConfig = Field(default_factory=MessagesConfig) + safety: SafetyConfig = Field(default_factory=SafetyConfig) + performance: PerformanceConfig = Field(default_factory=PerformanceConfig) + session_configs: list[SessionTemplateItem] = Field(default_factory=list) + fortune_session_configs: list[FortuneSessionTemplateItem] = Field( + default_factory=list + ) + + @field_validator("session_configs") + @classmethod + def validate_session_configs( + cls, v: list[SessionTemplateItem] + ) -> list[SessionTemplateItem]: + """Validate session configurations.""" + for item in v: + if item.session_type not in ("group", "private"): + raise ValueError(f"Invalid session_type: {item.session_type}") + return v + + @field_validator("fortune_session_configs") + @classmethod + def validate_fortune_session_configs( + cls, v: list[FortuneSessionTemplateItem] + ) -> list[FortuneSessionTemplateItem]: + """Validate fortune session configurations.""" + for item in v: + if item.session_type not in ("group", "private"): + raise ValueError(f"Invalid session_type: {item.session_type}") + return v + + # Convenience properties for backward compatibility + @property + def api_type(self) -> str: + """Get API type.""" + return self.setu_general.api_type.value + + @property + def multi_api_strategy(self) -> str: + """Get multi-API strategy.""" + return self.setu_general.multi_api_strategy.value + + @property + def content_mode(self) -> str: + """Get content mode.""" + return self.setu_general.content_mode.value + + @property + def max_count(self) -> int: + """Get max count.""" + return self.setu_general.max_count + + @property + def max_replenish_rounds(self) -> int: + """Get max replenish rounds.""" + return self.setu_general.max_replenish_rounds + + @property + def tag_alias(self) -> str: + """Get tag alias string.""" + return self.setu_general.tag_alias + + @property + def send_mode(self) -> str: + """Get send mode.""" + return self.delivery.send_mode.value + + @property + def r18_docx_mode(self) -> bool: + """Get R18 DOCX mode.""" + return self.delivery.r18_docx_mode + + @property + def auto_revoke_r18(self) -> bool: + """Get auto-revoke R18.""" + return self.delivery.auto_revoke_r18 + + @property + def auto_revoke_delay(self) -> int: + """Get auto-revoke delay.""" + return self.delivery.auto_revoke_delay + + @property + def napcat_stream_mode(self) -> str: + """Get NapCat stream upload mode.""" + return self.delivery.napcat_stream_mode.value + + @property + def html_card_strategy(self) -> str: + """Get HTML card strategy.""" + return self.html_card.strategy.value + + @property + def html_card_padding(self) -> int: + """Get HTML card padding.""" + return self.html_card.card_padding + + @property + def html_card_gap(self) -> int: + """Get HTML card gap.""" + return self.html_card.card_gap + + @property + def cache_enabled(self) -> bool: + """Get cache enabled.""" + return self.cache.enabled + + @property + def cache_ttl_hours(self) -> int: + """Get cache TTL hours.""" + return self.cache.ttl_hours + + @property + def cache_max_items(self) -> int: + """Get cache max items.""" + return self.cache.max_items + + @property + def cache_cleanup_on_start(self) -> bool: + """Get cache cleanup on start.""" + return self.cache.cleanup_on_start + + @property + def download_concurrent_limit(self) -> int: + """Get download concurrent limit.""" + return self.performance.download_concurrent_limit + + @property + def download_timeout_seconds(self) -> int: + """Get download timeout seconds.""" + return self.performance.download_timeout_seconds + + @property + def enable_range_download(self) -> bool: + """Get enable range download.""" + return self.performance.enable_range_download + + @property + def range_segments(self) -> int: + """Get range segments.""" + return self.performance.range_segments + + @property + def range_threshold(self) -> int: + """Get range threshold.""" + return self.performance.range_download_threshold + + @property + def exclude_ai(self) -> bool: + """Get exclude AI.""" + return self.api.lolicon.exclude_ai + + @property + def image_size(self) -> str: + """Get image size.""" + return self.api.lolicon.image_size.value + + @property + def proxy(self) -> str: + """Get proxy.""" + return self.api.lolicon.proxy + + @property + def aspect_ratio(self) -> str: + """Get aspect ratio.""" + return ( + self.api.lolicon.aspect_ratio.value if self.api.lolicon.aspect_ratio else "" + ) + + @property + def uid(self) -> list[int]: + """Get UID list.""" + return self.api.lolicon.uid + + @property + def keyword(self) -> str: + """Get keyword.""" + return self.api.lolicon.keyword + + @property + def atri_image_size(self) -> str: + """Get Atri image size.""" + return self.api.atri.image_size.value + + @property + def atri_proxy(self) -> str: + """Get Atri proxy.""" + return self.api.atri.proxy + + @property + def atri_aspect_ratio(self) -> str: + """Get Atri aspect ratio.""" + return self.api.atri.aspect_ratio.value if self.api.atri.aspect_ratio else "" + + @property + def atri_uid(self) -> list[int]: + """Get Atri UID list.""" + return self.api.atri.uid + + @property + def atri_keyword(self) -> str: + """Get Atri keyword.""" + return self.api.atri.keyword + + @property + def atri_exclude_ai(self) -> bool: + """Get Atri exclude AI.""" + return self.api.atri.exclude_ai + + @property + def fortune_api_type(self) -> str: + """Get fortune API type.""" + return self.fortune.api_type + + @property + def custom_api(self) -> dict[str, Any]: + """Get custom API config.""" + if self.api.custom_api_configs: + cfg = self.api.custom_api_configs[0] + return { + "url": cfg.url, + "method": cfg.method, + "timeout": cfg.timeout, + } + return {"url": "", "method": "GET", "timeout": 30} + + @property + def api_response_parser(self) -> dict[str, Any]: + """Get API response parser config.""" + if self.api.custom_api_configs: + cfg = self.api.custom_api_configs[0] + return { + "type": cfg.parser_type, + "json_path": cfg.json_path, + } + return {"type": "auto", "json_path": ""} + + @property + def custom_api_configs(self) -> list[dict[str, Any]]: + """Get custom API configs.""" + return [cfg.model_dump() for cfg in self.api.custom_api_configs] + + def get_custom_api_config(self, name: str | None = None) -> dict[str, Any] | None: + """Get custom API config by name.""" + configs = self.api.custom_api_configs + if not configs: + return None + if name: + for cfg in configs: + if cfg.name == name: + return cfg.model_dump() + return None + return configs[0].model_dump() + + @property + def msg_fetching_enabled(self) -> bool: + """Get fetching message enabled.""" + return self.messages.fetching.enabled + + @property + def msg_fetching_text(self) -> str: + """Get fetching message text.""" + return self.messages.fetching.text + + @property + def msg_found_enabled(self) -> bool: + """Get found message enabled.""" + return self.messages.found.enabled + + @property + def msg_found_text(self) -> str: + """Get found message text.""" + return self.messages.found.text + + @property + def msg_send_failed_text(self) -> str: + """Get send failed message text.""" + return self.messages.send_failed.text + + @property + def setu_user_access_control_mode(self) -> str: + """Get Setu user access control mode.""" + return self.safety.setu_user_access_control_mode.value + + @property + def setu_group_access_control_mode(self) -> str: + """Get Setu group access control mode.""" + return self.safety.setu_group_access_control_mode.value + + @property + def fortune_user_access_control_mode(self) -> str: + """Get Fortune user access control mode.""" + return self.safety.fortune_user_access_control_mode.value + + @property + def fortune_group_access_control_mode(self) -> str: + """Get Fortune group access control mode.""" + return self.safety.fortune_group_access_control_mode.value + + def format_found_message(self, count: int, revoke_delay: int | None = None) -> str: + """Format found message with placeholders.""" + result = self.msg_found_text.replace("{count}", str(count)) + if revoke_delay is not None: + result = result.replace("{revoke_delay}", str(revoke_delay)) + return result + + def get_effective_fortune_api_type(self) -> str: + """Get effective fortune API type.""" + fortune_api = self.fortune.api_type + if fortune_api == "inherit": + return self.api_type + return fortune_api diff --git a/src/shared/logging.py b/src/shared/logging.py new file mode 100644 index 0000000..43cf582 --- /dev/null +++ b/src/shared/logging.py @@ -0,0 +1,86 @@ +"""AstrBot Setu Plugin 日志包装器 + +提供带插件前缀的日志记录器,避免直接导出单例实例。 + +使用示例: + from .logger import get_logger + + logger = get_logger() + logger.info("消息") +""" + +from __future__ import annotations + +from typing import Any + +from astrbot.api import logger as _astrbot_logger + + +class PrefixedLogger: + """AstrBot 日志包装器,添加插件前缀并正确显示调用位置。""" + + PREFIX = "[setu] " + CALLER_STACKLEVEL = 2 + + def _add_prefix(self, msg: object) -> str: + """为消息添加前缀。""" + return self.PREFIX + str(msg) + + def _with_stacklevel(self, kwargs: dict[str, Any]) -> dict[str, Any]: + """确保 stacklevel 参数正确传递以显示真实调用位置。""" + copied = dict(kwargs) + if "stacklevel" not in copied: + copied["stacklevel"] = self.CALLER_STACKLEVEL + return copied + + def debug(self, msg: object, *args: Any, **kwargs: Any) -> None: + _astrbot_logger.debug( + self._add_prefix(msg), *args, **self._with_stacklevel(kwargs) + ) + + def info(self, msg: object, *args: Any, **kwargs: Any) -> None: + _astrbot_logger.info( + self._add_prefix(msg), *args, **self._with_stacklevel(kwargs) + ) + + def warning(self, msg: object, *args: Any, **kwargs: Any) -> None: + _astrbot_logger.warning( + self._add_prefix(msg), *args, **self._with_stacklevel(kwargs) + ) + + def error(self, msg: object, *args: Any, **kwargs: Any) -> None: + _astrbot_logger.error( + self._add_prefix(msg), *args, **self._with_stacklevel(kwargs) + ) + + def exception(self, msg: object, *args: Any, **kwargs: Any) -> None: + _astrbot_logger.exception( + self._add_prefix(msg), *args, **self._with_stacklevel(kwargs) + ) + + def critical(self, msg: object, *args: Any, **kwargs: Any) -> None: + _astrbot_logger.critical( + self._add_prefix(msg), *args, **self._with_stacklevel(kwargs) + ) + + +# 内部缓存,禁止直接导出 +_logger_instance: PrefixedLogger | None = None + + +def get_logger() -> PrefixedLogger: + """获取插件日志记录器单例。 + + Returns: + PrefixedLogger 实例 + """ + global _logger_instance + if _logger_instance is None: + _logger_instance = PrefixedLogger() + return _logger_instance + + +def clear_logger() -> None: + """清除日志记录器单例(用于测试)。""" + global _logger_instance + _logger_instance = None diff --git a/src/shared/send_cache.py b/src/shared/send_cache.py new file mode 100644 index 0000000..64da45f --- /dev/null +++ b/src/shared/send_cache.py @@ -0,0 +1,210 @@ +"""Disk-backed send cache for image delivery.""" + +from __future__ import annotations + +import asyncio +import hashlib +import mimetypes +import time +from dataclasses import dataclass +from pathlib import Path +from urllib.parse import urlparse + +from .logging import get_logger + +logger = get_logger() + +_IMAGE_SUFFIXES = { + "image/jpeg": ".jpg", + "image/jpg": ".jpg", + "image/png": ".png", + "image/gif": ".gif", + "image/webp": ".webp", + "image/bmp": ".bmp", +} + + +def guess_file_suffix(source: str, content_type: str | None = None) -> str: + """Guess a file suffix for an image source.""" + normalized_type = (content_type or "").split(";", 1)[0].strip().lower() + if normalized_type in _IMAGE_SUFFIXES: + return _IMAGE_SUFFIXES[normalized_type] + + parsed = urlparse(source) + guessed = Path(parsed.path).suffix.lower() + if guessed in {".jpg", ".jpeg", ".png", ".gif", ".webp", ".bmp"}: + return ".jpg" if guessed == ".jpeg" else guessed + + fallback = mimetypes.guess_extension(normalized_type) if normalized_type else "" + return fallback or ".img" + + +@dataclass(frozen=True) +class SendCacheWrite: + """Reserved file paths for one cache write.""" + + temp_path: Path + final_path: Path + + +class DiskSendCache: + """Small URL-addressed disk cache used before platform image sending.""" + + def __init__( + self, + root: Path, + *, + enabled: bool = True, + ttl_hours: int = 2, + max_items: int = 200, + ) -> None: + self.root = root + self.enabled = enabled + self.ttl_seconds = max(1, int(ttl_hours)) * 60 * 60 + self.max_items = max(1, int(max_items)) + self._lock = asyncio.Lock() + + def _key(self, source: str) -> str: + return hashlib.sha256(source.encode("utf-8")).hexdigest() + + async def get(self, source: str) -> Path | None: + """Return a fresh cached file for the source URL, if present.""" + if not self.enabled: + return None + + key = self._key(source) + now = time.time() + + def find() -> Path | None: + if not self.root.exists(): + return None + for path in self.root.glob(f"{key}.*"): + if path.suffix == ".part" or not path.is_file(): + continue + if now - path.stat().st_mtime > self.ttl_seconds: + continue + return path + return None + + return await asyncio.to_thread(find) + + async def reserve( + self, source: str, content_type: str | None = None + ) -> SendCacheWrite: + """Reserve a temp path and final path for a source URL.""" + suffix = guess_file_suffix(source, content_type) + key = self._key(source) + final_path = self.root / f"{key}{suffix}" + temp_path = self.root / f"{key}.{time.monotonic_ns()}{suffix}.part" + + await asyncio.to_thread(self.root.mkdir, parents=True, exist_ok=True) + return SendCacheWrite(temp_path=temp_path, final_path=final_path) + + async def commit(self, write: SendCacheWrite) -> Path: + """Atomically promote a temp cache file to its final path.""" + + def replace() -> None: + write.temp_path.replace(write.final_path) + write.final_path.touch() + + async with self._lock: + await asyncio.to_thread(replace) + return write.final_path + + async def discard(self, write: SendCacheWrite) -> None: + """Delete a failed temp cache file.""" + + def unlink() -> None: + if write.temp_path.exists(): + write.temp_path.unlink() + + await asyncio.to_thread(unlink) + + async def cleanup(self) -> int: + """Delete expired and overflow cache files.""" + now = time.time() + + def clean() -> int: + if not self.root.exists(): + return 0 + + removed = 0 + files: list[Path] = [] + for path in self.root.iterdir(): + if not path.is_file(): + continue + try: + age = now - path.stat().st_mtime + except OSError: + continue + if path.suffix == ".part" and age > 3600: + path.unlink(missing_ok=True) + removed += 1 + continue + if age > self.ttl_seconds: + path.unlink(missing_ok=True) + removed += 1 + continue + files.append(path) + + files.sort(key=lambda p: p.stat().st_mtime, reverse=True) + for path in files[self.max_items :]: + path.unlink(missing_ok=True) + removed += 1 + return removed + + removed = await asyncio.to_thread(clean) + if removed: + logger.debug("[send_cache] cleaned %d cached files", removed) + return removed + + +_send_cache: DiskSendCache | None = None + + +async def init_send_cache( + data_dir: Path, + *, + enabled: bool, + ttl_hours: int, + max_items: int, + cleanup_on_start: bool, +) -> DiskSendCache: + """Initialize the process-wide send cache.""" + global _send_cache + _send_cache = DiskSendCache( + data_dir / "send_cache", + enabled=enabled, + ttl_hours=ttl_hours, + max_items=max_items, + ) + if cleanup_on_start: + await _send_cache.cleanup() + return _send_cache + + +def get_send_cache() -> DiskSendCache | None: + """Return the configured send cache, if initialized.""" + return _send_cache + + +def clear_send_cache() -> None: + """Clear the process-wide send cache reference.""" + global _send_cache + _send_cache = None + + +def schedule_send_cache_cleanup(delay_seconds: float = 300.0) -> None: + """Schedule a best-effort cache cleanup after current sends settle.""" + cache = _send_cache + if cache is None: + return + + async def cleanup_later() -> None: + await asyncio.sleep(max(0.0, delay_seconds)) + await cache.cleanup() + + try: + asyncio.create_task(cleanup_later()) + except RuntimeError: + logger.debug("[send_cache] no running loop for delayed cleanup") diff --git a/templates/fortune.html b/templates/fortune.html deleted file mode 100644 index 1023ee3..0000000 --- a/templates/fortune.html +++ /dev/null @@ -1,219 +0,0 @@ - - - - - - Fortune Card - - - -
- -
-
{{ username }} 的运势
-
{{ date_str }}
-
- - -
- {% if image_base64 %} - Fortune Image - {% else %} -
今日运势
- {% endif %} -
- - -
-
{{ title }}
-
{{ stars_display }}
-
-
{{ description }}
- {% if extra_message %}
{{ extra_message }}
{% endif %} -
-
- - - -
- - diff --git a/templates/res/fonts/NotoSansSC-Bold.woff2 b/templates/res/fonts/NotoSansSC-Bold.woff2 deleted file mode 100644 index 29dceb2..0000000 Binary files a/templates/res/fonts/NotoSansSC-Bold.woff2 and /dev/null differ diff --git a/templates/res/fonts/NotoSansSC-Regular.woff2 b/templates/res/fonts/NotoSansSC-Regular.woff2 deleted file mode 100644 index a658a4d..0000000 Binary files a/templates/res/fonts/NotoSansSC-Regular.woff2 and /dev/null differ diff --git a/templates/res/fonts/SSFangTangTi.woff2 b/templates/res/fonts/SSFangTangTi.woff2 deleted file mode 100644 index 0d8f73d..0000000 Binary files a/templates/res/fonts/SSFangTangTi.woff2 and /dev/null differ diff --git a/templates/setu.html b/templates/setu.html deleted file mode 100644 index 6b82065..0000000 --- a/templates/setu.html +++ /dev/null @@ -1,94 +0,0 @@ - - - - - - Safe Render Template - - - -
- -
- -
- {cards} -
-
- - diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..040cdd1 --- /dev/null +++ b/tests/__init__.py @@ -0,0 +1,18 @@ +"""Test configuration for AstrBot Setu plugin tests.""" + +from __future__ import annotations + +import pytest + + +@pytest.fixture +def sample_access_policy(): + """Create sample access policy for testing.""" + from astrbot_plugin_setu.src.domain.access_control import AccessPolicy + + return AccessPolicy.for_session( + user_id="test_user", + group_id="test_group", + user_mode="none", + group_mode="none", + ) diff --git a/tests/application/test_image_provider_logging.py b/tests/application/test_image_provider_logging.py new file mode 100644 index 0000000..3343a15 --- /dev/null +++ b/tests/application/test_image_provider_logging.py @@ -0,0 +1,48 @@ +from __future__ import annotations + +from unittest.mock import MagicMock + +import httpx +import pytest + +from astrbot_plugin_setu.src.application.ports import ( + image_provider as image_provider_module, +) +from astrbot_plugin_setu.src.application.ports.image_provider import SetuImageProvider +from astrbot_plugin_setu.src.domain.setu import SetuRequest + + +class DummyProvider(SetuImageProvider): + async def fetch_image_urls( + self, + num: int, + tags: list[str], + r18: bool, + exclude_ai: bool = True, + ) -> list[str]: + return ["https://example.com/a.jpg"] + + +@pytest.mark.asyncio +async def test_fetch_and_download_logs_when_all_downloads_fail(monkeypatch) -> None: + provider = DummyProvider() + request = SetuRequest.from_user_input( + count=1, tags=["cat"], r18=False, exclude_ai=True + ) + fake_logger = MagicMock() + monkeypatch.setattr(image_provider_module, "logger", fake_logger) + + async def fail_get(_self, url: str): + raise httpx.ConnectError("boom", request=httpx.Request("GET", url)) + + monkeypatch.setattr(httpx.AsyncClient, "get", fail_get) + + payload = await provider.fetch_and_download(request) + + assert payload.items == () + warning_messages = [call.args[0] for call in fake_logger.warning.call_args_list] + error_messages = [call.args[0] for call in fake_logger.error.call_args_list] + assert any("[provider] download failed:" in message for message in warning_messages) + assert any( + "[provider] all downloads failed:" in message for message in error_messages + ) diff --git a/tests/application/test_provider_proxy_rewrite.py b/tests/application/test_provider_proxy_rewrite.py new file mode 100644 index 0000000..a74084a --- /dev/null +++ b/tests/application/test_provider_proxy_rewrite.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +from astrbot_plugin_setu.src.infrastructure.providers.atri import AtriProvider +from astrbot_plugin_setu.src.infrastructure.providers.lolicon import LoliconProvider + + +def test_lolicon_provider_rewrites_pixiv_proxy_host() -> None: + provider = LoliconProvider(proxy="my.proxy.local") + + rewritten = provider._apply_proxy_to_urls( + ["https://i.pximg.net/img-original/img/2024/01/01/00/00/00/1_p0.jpg"], + provider.proxy, + "LoliconProvider", + ) + + assert rewritten == [ + "https://my.proxy.local/img-original/img/2024/01/01/00/00/00/1_p0.jpg" + ] + + +def test_atri_provider_keeps_non_pixiv_host_unchanged() -> None: + provider = AtriProvider(proxy="my.proxy.local") + + rewritten = provider._apply_proxy_to_urls( + ["https://cdn.example.com/image.jpg"], + provider.proxy, + "AtriProvider", + ) + + assert rewritten == ["https://cdn.example.com/image.jpg"] diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..46adba0 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,223 @@ +"""Shared fixtures for AstrBot Setu plugin tests.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any +from unittest.mock import MagicMock + +import pytest + +from astrbot.api.event import AstrMessageEvent +from astrbot.core import AstrBotConfig + + +@pytest.fixture +def temp_data_dir(tmp_path: Path) -> Path: + """Create temporary data directory for tests. + + Args: + tmp_path: Pytest temporary path fixture + + Returns: + Temporary data directory + """ + data_dir = tmp_path / "plugin_data" + data_dir.mkdir(parents=True, exist_ok=True) + return data_dir + + +@pytest.fixture +def mock_astrbot_config() -> AstrBotConfig: + """Create mock AstrBotConfig. + + Returns: + Mock AstrBotConfig with basic structure + """ + config = MagicMock(spec=AstrBotConfig) + config.__setitem__ = MagicMock() + config.__getitem__ = MagicMock(return_value={}) + config.get = MagicMock(return_value=None) + return config + + +@pytest.fixture +def mock_event() -> AstrMessageEvent: + """Create mock AstrMessageEvent. + + Returns: + Mock event with common attributes + """ + event = MagicMock(spec=AstrMessageEvent) + event.get_sender_id = MagicMock(return_value="test_user_123") + event.get_group_id = MagicMock(return_value="test_group_456") + event.get_session_id = MagicMock(return_value="test_group_456_test_user_123") + event.get_self_id = MagicMock(return_value="test_bot_789") + event.get_messages = MagicMock(return_value=[]) + event.get_message_str = MagicMock(return_value="/setu 3") + event.message_str = "/setu 3" + event.unified_msg_origin = "test_platform:12345" + event.platform = MagicMock() + event.platform.name = "test_platform" + event.is_at_or_wake_command = False + + # Message chain results + def plain_result_fn(text: str) -> Any: + result = MagicMock() + result.result_chain = [MagicMock(text=text)] + return result + + event.plain_result = MagicMock(side_effect=plain_result_fn) + + def chain_result_fn(chain: list[Any]) -> Any: + result = MagicMock() + result.result_chain = chain + return result + + event.chain_result = MagicMock(side_effect=chain_result_fn) + + return event + + +@pytest.fixture +def mock_plugin_context() -> Any: + """Create mock plugin context. + + Returns: + Mock context with send_message method + """ + context = MagicMock() + context.send_message = MagicMock() + + async def mock_send(origin: str, result: Any) -> Any: + return {"message_id": "test_msg_123"} + + context.send_message.side_effect = mock_send + return context + + +@pytest.fixture +def sample_image_bytes() -> bytes: + """Create sample image bytes for testing. + + Returns: + Small PNG image bytes + """ + # Minimal PNG (1x1 transparent pixel) + return ( + b"\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01" + b"\x08\x06\x00\x00\x00\x1f\x15\xc4\x89\x00\x00\x00\nIDATx\x9cc\x00\x01" + b"\x00\x00\x05\x00\x01\r\n-\xb4\x00\x00\x00\x00IEND\xaeB`\x82" + ) + + +@pytest.fixture +def mock_provider() -> Any: + """Create mock image provider. + + Returns: + Mock provider with fetch_image_urls method + """ + provider = MagicMock() + + async def mock_fetch( + num: int, tags: list[str], r18: bool, exclude_ai: bool + ) -> list[str]: + return [f"https://example.com/image{i}.jpg" for i in range(num)] + + provider.fetch_image_urls = MagicMock(side_effect=mock_fetch) + return provider + + +@pytest.fixture +def sample_config_dict() -> dict[str, Any]: + """Create sample plugin configuration. + + Returns: + Configuration dict matching _conf_schema.json structure + """ + return { + "setu_general": { + "api_type": "lolicon", + "multi_api_strategy": "round_robin", + "content_mode": "sfw", + "max_count": 10, + "max_replenish_rounds": 3, + "tag_alias": "", + }, + "delivery": { + "send_mode": "auto", + "r18_docx_mode": True, + "auto_handle_send_failure": True, + "auto_revoke_r18": False, + "auto_revoke_delay": 30, + "napcat_stream_mode": "fallback", + }, + "html_card": { + "strategy": "fallback", + "mode": "single", + "card_padding": 6, + "card_gap": 6, + }, + "fortune": { + "enabled": True, + "api_type": "inherit", + "tags": "", + "content_mode": "sfw", + "allow_user_refresh": False, + "auto_refresh": True, + }, + "cache": { + "enabled": True, + "ttl_hours": 2, + "max_items": 200, + "cleanup_on_start": True, + }, + "api": { + "lolicon": { + "image_size": "original", + "proxy": "i.pixiv.re", + "aspect_ratio": None, + "uid": [], + "keyword": "", + "exclude_ai": True, + }, + "atri": { + "image_size": "original", + "proxy": "i.pixiv.re", + "aspect_ratio": None, + "uid": [], + "keyword": "", + "exclude_ai": True, + }, + "custom_api_configs": [], + }, + "messages": { + "fetching": {"enabled": True, "text": "正在获取图片,请稍候..."}, + "found": {"enabled": True, "text": "找到 {count} 张符合要求的图片~"}, + "send_failed": {"text": "图片发送失败,请稍后再试。"}, + }, + "safety": { + "setu_user_access_control_mode": "none", + "setu_group_access_control_mode": "none", + "setu_blocked_users": [], + "setu_whitelist_users": [], + "setu_blocked_groups": [], + "setu_whitelist_groups": [], + "fortune_user_access_control_mode": "none", + "fortune_group_access_control_mode": "none", + "fortune_blocked_users": [], + "fortune_whitelist_users": [], + "fortune_blocked_groups": [], + "fortune_whitelist_groups": [], + }, + "performance": { + "enable_range_download": False, + "range_segments": 3, + "range_download_threshold": 512, + "download_concurrent_limit": 10, + "download_timeout_seconds": 30, + }, + "session_configs": [], + "fortune_session_configs": [], + } diff --git a/tests/domain/test_exceptions.py b/tests/domain/test_exceptions.py new file mode 100644 index 0000000..8fa431c --- /dev/null +++ b/tests/domain/test_exceptions.py @@ -0,0 +1,85 @@ +"""Tests for domain exceptions.""" + +from __future__ import annotations + +import pytest + +from astrbot_plugin_setu.src.domain.exceptions import ( + AccessDeniedError, + FortuneException, + FortuneNotFoundError, + ProviderError, + SendError, + SetuException, + ValidationError, +) + + +class TestSetuException: + """Test SetuException base class.""" + + def test_setu_exception(self) -> None: + """Test basic SetuException.""" + with pytest.raises(SetuException): + raise SetuException("Test error") + + +class TestProviderError: + """Test ProviderError.""" + + def test_provider_error(self) -> None: + """Test ProviderError is SetuException.""" + with pytest.raises(SetuException): + raise ProviderError("Provider failed") + + +class TestSendError: + """Test SendError.""" + + def test_send_error(self) -> None: + """Test SendError is SetuException.""" + with pytest.raises(SetuException): + raise SendError("Send failed") + + +class TestAccessDeniedError: + """Test AccessDeniedError.""" + + def test_access_denied_with_reason(self) -> None: + """Test AccessDeniedError with reason.""" + error = AccessDeniedError("User is blocked") + assert str(error) == "Access denied: User is blocked" + assert error.reason == "User is blocked" + + def test_access_denied_no_reason(self) -> None: + """Test AccessDeniedError without reason.""" + error = AccessDeniedError() + assert str(error) == "Access denied" + assert error.reason == "" + + +class TestValidationError: + """Test ValidationError.""" + + def test_validation_error(self) -> None: + """Test ValidationError is SetuException.""" + with pytest.raises(SetuException): + raise ValidationError("Invalid input") + + +class TestFortuneException: + """Test FortuneException base class.""" + + def test_fortune_exception(self) -> None: + """Test basic FortuneException.""" + with pytest.raises(FortuneException): + raise FortuneException("Fortune error") + + +class TestFortuneNotFoundError: + """Test FortuneNotFoundError.""" + + def test_fortune_not_found(self) -> None: + """Test FortuneNotFoundError is FortuneException.""" + with pytest.raises(FortuneException): + raise FortuneNotFoundError("Fortune not found") diff --git a/tests/domain/test_tag_resolver.py b/tests/domain/test_tag_resolver.py new file mode 100644 index 0000000..1020f67 --- /dev/null +++ b/tests/domain/test_tag_resolver.py @@ -0,0 +1,110 @@ +"""Tests for TagResolverService domain service.""" + +from __future__ import annotations + + +from astrbot_plugin_setu.src.domain.setu import TagResolverService + + +class TestTagResolverService: + """Test TagResolverService.""" + + def test_init_default(self) -> None: + """Test initialization with default alias map.""" + resolver = TagResolverService() + alias_map = resolver.get_alias_map() + assert "萝莉" in alias_map + assert "loli" in alias_map["萝莉"] + + def test_init_custom(self) -> None: + """Test initialization with custom alias map.""" + custom_map = {"test": ["alias1", "alias2"]} + resolver = TagResolverService(custom_map) + assert resolver.get_alias_map() == custom_map + + def test_resolve_tags_empty(self) -> None: + """Test resolving empty tag string.""" + resolver = TagResolverService() + assert resolver.resolve_tags("") == [] + + def test_resolve_tags_single(self) -> None: + """Test resolving single tag.""" + resolver = TagResolverService() + tags = resolver.resolve_tags("girl") + assert tags == ["少女"] + + def test_resolve_tags_multiple(self) -> None: + """Test resolving multiple tags.""" + resolver = TagResolverService() + tags = resolver.resolve_tags("girl,cute,long_hair") + assert tags == ["少女", "cute", "长发"] + + def test_resolve_tags_with_alias(self) -> None: + """Test resolving tag with alias.""" + resolver = TagResolverService() + tags = resolver.resolve_tags("loli") + assert tags == ["萝莉"] + + def test_resolve_tags_chinese_comma(self) -> None: + """Test resolving tags with Chinese comma.""" + resolver = TagResolverService() + tags = resolver.resolve_tags("girl,cute") + assert tags == ["少女", "cute"] + + def test_resolve_tags_with_spaces(self) -> None: + """Test resolving tags with spaces.""" + resolver = TagResolverService() + tags = resolver.resolve_tags("girl cute long_hair") + assert tags == ["少女", "cute", "长发"] + + def test_resolve_tags_mixed_separators(self) -> None: + """Test resolving tags with mixed separators.""" + resolver = TagResolverService() + tags = resolver.resolve_tags("girl, cute,long_hair cat") + assert tags == ["少女", "cute", "长发", "cat"] + + def test_parse_alias_map_from_string(self) -> None: + """Test parsing alias map from config string.""" + alias_str = """# Comment line +canonical1=alias1,alias2 +canonical2=alias3 + +# Another comment +canonical3=alias4,alias5 +""" + result = TagResolverService.parse_alias_map_from_string(alias_str) + assert result == { + "canonical1": ["alias1", "alias2"], + "canonical2": ["alias3"], + "canonical3": ["alias4", "alias5"], + } + + def test_parse_alias_map_empty(self) -> None: + """Test parsing empty alias string.""" + result = TagResolverService.parse_alias_map_from_string("") + assert result == {} + + def test_parse_alias_map_invalid_lines(self) -> None: + """Test parsing alias string with invalid lines.""" + alias_str = """ +valid=test1,alias1 +invalid line +=alias2 +test2= +""" + result = TagResolverService.parse_alias_map_from_string(alias_str) + assert result == {"valid": ["test1", "alias1"]} + + def test_update_alias_map(self) -> None: + """Test updating alias map.""" + resolver = TagResolverService() + new_map = {"custom": ["alias1", "alias2"]} + resolver.update_alias_map(new_map) + assert resolver.get_alias_map() == new_map + + def test_find_canonical_tag_case_insensitive(self) -> None: + """Test canonical tag lookup is case-insensitive.""" + resolver = TagResolverService() + assert resolver._find_canonical_tag("萝莉") == "萝莉" + assert resolver._find_canonical_tag("LOLI") == "萝莉" + assert resolver._find_canonical_tag("loli") == "萝莉" diff --git a/tests/domain/test_value_objects.py b/tests/domain/test_value_objects.py new file mode 100644 index 0000000..c95fac5 --- /dev/null +++ b/tests/domain/test_value_objects.py @@ -0,0 +1,143 @@ +"""Tests for domain value objects.""" + +from __future__ import annotations + + +from astrbot_plugin_setu.src.application.setu import ImagePayload +from astrbot_plugin_setu.src.domain.access_control import AccessPolicy +from astrbot_plugin_setu.src.domain.fortune import FortuneResult, FortuneSeed +from astrbot_plugin_setu.src.domain.setu import SetuRequest +from astrbot_plugin_setu.src.infrastructure.sending import SendOptions + + +class TestSetuRequest: + """Test SetuRequest value object.""" + + def test_from_user_input(self) -> None: + """Test creating SetuRequest from user input.""" + request = SetuRequest.from_user_input(3, ["girl", "cute"], False) + assert request.count == 3 + assert request.tags == ("girl", "cute") + assert request.r18 is False + assert request.exclude_ai is True + + def test_with_tags(self) -> None: + """Test creating new request with different tags.""" + request = SetuRequest.from_user_input(1, ["girl"], False) + new_request = request.with_tags(["cat", "cute"]) + assert new_request.count == 1 + assert new_request.tags == ("cat", "cute") + assert new_request.r18 is False + + +class TestImagePayload: + """Test ImagePayload value object.""" + + def test_empty_payload(self) -> None: + """Test empty payload detection.""" + payload = ImagePayload((), (), False, ()) + assert payload.is_empty + assert payload.count == 0 + + def test_payload_with_urls(self) -> None: + """Test payload with URLs.""" + payload = ImagePayload(("url1", "url2"), (), False, ("tag1",)) + assert not payload.is_empty + assert payload.count == 2 + + def test_payload_with_bytes(self) -> None: + """Test payload with bytes.""" + payload = ImagePayload((), (b"data1", b"data2"), False, ()) + assert not payload.is_empty + assert payload.count == 2 + + def test_payload_count_uses_max(self) -> None: + """Test count uses max of URLs and bytes.""" + payload = ImagePayload(("url1", "url2"), (b"data1",), False, ()) + assert payload.count == 2 + + +class TestAccessPolicy: + """Test AccessPolicy value object.""" + + def test_for_user(self) -> None: + """Test creating policy for user.""" + policy = AccessPolicy.for_user("user123", "blacklist") + assert policy.user_id == "user123" + assert policy.group_id is None + assert policy.user_mode == "blacklist" + assert policy.group_mode == "none" + + def test_for_group(self) -> None: + """Test creating policy for group.""" + policy = AccessPolicy.for_group("group456", "whitelist") + assert policy.user_id is None + assert policy.group_id == "group456" + assert policy.user_mode == "none" + assert policy.group_mode == "whitelist" + + def test_for_session(self) -> None: + """Test creating policy for full session.""" + policy = AccessPolicy.for_session( + "user123", "group456", "blacklist", "whitelist" + ) + assert policy.user_id == "user123" + assert policy.group_id == "group456" + assert policy.user_mode == "blacklist" + assert policy.group_mode == "whitelist" + + +class TestFortuneSeed: + """Test FortuneSeed value object.""" + + def test_for_today(self) -> None: + """Test creating seed for today.""" + from datetime import date + + seed = FortuneSeed.for_today("user123") + assert seed.user_id == "user123" + assert seed.date_str == date.today().isoformat() + + def test_cache_key(self) -> None: + """Test cache key generation.""" + seed = FortuneSeed("user123", "2026-05-09") + assert seed.cache_key == "user123_2026-05-09" + + +class TestFortuneResult: + """Test FortuneResult value object.""" + + def test_max_stars(self) -> None: + """Test max stars property.""" + seed = FortuneSeed("user123", "2026-05-09") + result = FortuneResult( + seed=seed, + title="大吉", + star_count=6, + description="Lucky day!", + extra_message="", + theme_color="theme-red", + ) + assert result.max_stars == 7 + assert result.star_count == 6 + + +class TestSendOptions: + """Test SendOptions value object.""" + + def test_defaults(self) -> None: + """Test default values.""" + options = SendOptions( + send_mode="auto", + use_html_card=False, + auto_revoke=False, + revoke_delay=30, + r18_docx_mode=True, + ) + assert options.send_mode == "auto" + assert options.use_html_card is False + assert options.auto_revoke is False + assert options.revoke_delay == 30 + assert options.r18_docx_mode is True + assert options.html_padding == 6 + assert options.html_gap == 6 diff --git a/tests/infrastructure/test_access_control_repo.py b/tests/infrastructure/test_access_control_repo.py new file mode 100644 index 0000000..09ca0e8 --- /dev/null +++ b/tests/infrastructure/test_access_control_repo.py @@ -0,0 +1,137 @@ +"""Tests for FileBackedAccessControlRepo.""" + +from __future__ import annotations + +import pytest + +from astrbot_plugin_setu.src.infrastructure.persistence import ( + FileBackedAccessControlRepo, +) + + +class TestFileBackedAccessControlRepo: + """Test FileBackedAccessControlRepo.""" + + @pytest.mark.asyncio + async def test_initialize(self, temp_data_dir, mock_astrbot_config) -> None: + """Test repository initialization.""" + repo = FileBackedAccessControlRepo(temp_data_dir, mock_astrbot_config) + await repo.initialize() + assert repo._config_file.exists() + + @pytest.mark.asyncio + async def test_setu_user_blacklist( + self, temp_data_dir, mock_astrbot_config + ) -> None: + """Test Setu user blacklist operations.""" + repo = FileBackedAccessControlRepo(temp_data_dir, mock_astrbot_config) + await repo.initialize() + + # Add to blacklist + assert await repo.add_setu_blocked_user("user123") is True + assert await repo.is_setu_user_blocked("user123") is True + + # Remove from blacklist + assert await repo.remove_setu_blocked_user("user123") is True + assert await repo.is_setu_user_blocked("user123") is False + + @pytest.mark.asyncio + async def test_setu_user_whitelist( + self, temp_data_dir, mock_astrbot_config + ) -> None: + """Test Setu user whitelist operations.""" + repo = FileBackedAccessControlRepo(temp_data_dir, mock_astrbot_config) + await repo.initialize() + + # Add to whitelist + assert await repo.add_setu_whitelist_user("user456") is True + assert await repo.is_setu_user_whitelisted("user456") is True + + # Remove from whitelist + assert await repo.remove_setu_whitelist_user("user456") is True + assert await repo.is_setu_user_whitelisted("user456") is False + + @pytest.mark.asyncio + async def test_setu_group_blacklist( + self, temp_data_dir, mock_astrbot_config + ) -> None: + """Test Setu group blacklist operations.""" + repo = FileBackedAccessControlRepo(temp_data_dir, mock_astrbot_config) + await repo.initialize() + + # Add to blacklist + assert await repo.add_setu_blocked_group("group789") is True + assert await repo.is_setu_group_blocked("group789") is True + + # Remove from blacklist + assert await repo.remove_setu_blocked_group("group789") is True + assert await repo.is_setu_group_blocked("group789") is False + + @pytest.mark.asyncio + async def test_setu_group_whitelist( + self, temp_data_dir, mock_astrbot_config + ) -> None: + """Test Setu group whitelist operations.""" + repo = FileBackedAccessControlRepo(temp_data_dir, mock_astrbot_config) + await repo.initialize() + + # Add to whitelist + assert await repo.add_setu_whitelist_group("group101") is True + assert await repo.is_setu_group_whitelisted("group101") is True + + # Remove from whitelist + assert await repo.remove_setu_whitelist_group("group101") is True + assert await repo.is_setu_group_whitelisted("group101") is False + + @pytest.mark.asyncio + async def test_fortune_user_operations( + self, temp_data_dir, mock_astrbot_config + ) -> None: + """Test Fortune user operations.""" + repo = FileBackedAccessControlRepo(temp_data_dir, mock_astrbot_config) + await repo.initialize() + + # Blacklist + assert await repo.add_fortune_blocked_user("user222") is True + assert await repo.is_fortune_user_blocked("user222") is True + + # Whitelist + assert await repo.add_fortune_whitelist_user("user333") is True + assert await repo.is_fortune_user_whitelisted("user333") is True + + @pytest.mark.asyncio + async def test_fortune_group_operations( + self, temp_data_dir, mock_astrbot_config + ) -> None: + """Test Fortune group operations.""" + repo = FileBackedAccessControlRepo(temp_data_dir, mock_astrbot_config) + await repo.initialize() + + # Blacklist + assert await repo.add_fortune_blocked_group("group444") is True + assert await repo.is_fortune_group_blocked("group444") is True + + # Whitelist + assert await repo.add_fortune_whitelist_group("group555") is True + assert await repo.is_fortune_group_whitelisted("group555") is True + + @pytest.mark.asyncio + async def test_mutual_exclusion_blacklist_whitelist( + self, temp_data_dir, mock_astrbot_config + ) -> None: + """Test that adding to blacklist removes from whitelist and vice versa.""" + repo = FileBackedAccessControlRepo(temp_data_dir, mock_astrbot_config) + await repo.initialize() + + # Add to blacklist, then whitelist + await repo.add_setu_blocked_user("user666") + assert await repo.is_setu_user_blocked("user666") is True + + await repo.add_setu_whitelist_user("user666") + assert await repo.is_setu_user_whitelisted("user666") is True + assert await repo.is_setu_user_blocked("user666") is False + + # Add to whitelist, then blacklist + await repo.add_setu_blocked_user("user666") + assert await repo.is_setu_user_blocked("user666") is True + assert await repo.is_setu_user_whitelisted("user666") is False diff --git a/tests/infrastructure/test_image_sender.py b/tests/infrastructure/test_image_sender.py new file mode 100644 index 0000000..b732e7e --- /dev/null +++ b/tests/infrastructure/test_image_sender.py @@ -0,0 +1,150 @@ +"""Tests for image sender transport behavior.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest +import astrbot.api.message_components as Comp +from astrbot_plugin_setu.src.application.setu import ImagePayload +from astrbot_plugin_setu.src.infrastructure.astrbot.config import ( + clear_config, + set_plugin_context, +) +from astrbot_plugin_setu.src.infrastructure.sending.image_sender import ImageSender +from astrbot_plugin_setu.src.infrastructure.sending.send_strategies import ( + DirectSendStrategy, +) +from astrbot_plugin_setu.src.shared.config import SetuPluginConfig + + +@pytest.fixture(autouse=True) +def reset_singletons() -> None: + """Keep config/context singletons isolated.""" + clear_config() + yield + clear_config() + + +@pytest.mark.asyncio +async def test_send_images_streams_on_fallback( + tmp_path: Path, mock_event, sample_config_dict +) -> None: + """A failed normal path send uses NapCat stream upload and retries once.""" + image_path = tmp_path / "image.jpg" + image_path.write_bytes(b"image-data") + + context = MagicMock() + sent_results: list[Any] = [] + + async def send_message(_origin: str, result: Any) -> Any: + sent_results.append(result) + if len(sent_results) == 1: + return {"message_id": "found"} + if len(sent_results) == 2: + return None + return {"message_id": "streamed"} + + context.send_message.side_effect = send_message + set_plugin_context(context) + + mock_event.platform.name = "aiocqhttp" + mock_event.bot = MagicMock() + mock_event.bot.api = None + + async def call_action(action: str, **params: Any) -> dict[str, Any]: + assert action == "upload_file_stream" + if params.get("is_complete"): + return { + "status": "ok", + "retcode": 0, + "data": {"file_path": "stream://image"}, + } + return {"status": "ok", "retcode": 0, "data": {}} + + mock_event.bot.call_action = AsyncMock(side_effect=call_action) + + config = SetuPluginConfig(**sample_config_dict) + payload = ImagePayload( + urls=("https://example.com/image.jpg",), + raw_bytes=(), + file_paths=(image_path,), + items=(image_path,), + r18=False, + tags=(), + ) + + results = [ + item async for item in ImageSender(config).send_images(payload, mock_event) + ] + + assert results == [{"send_success": True, "image_count": 1}] + assert len(sent_results) == 3 + retried_chain = sent_results[-1].result_chain + assert retried_chain[0].file == "stream://image" + assert mock_event.bot.call_action.called + + +@pytest.mark.asyncio +async def test_send_images_materializes_local_files_before_direct_send( + tmp_path: Path, mock_event, sample_config_dict +) -> None: + image_path = tmp_path / "image.jpg" + image_path.write_bytes(b"image-data") + + context = MagicMock() + sent_results: list[Any] = [] + + async def send_message(_origin: str, result: Any) -> dict[str, str]: + sent_results.append(result) + return {"message_id": "ok"} + + context.send_message = AsyncMock(side_effect=send_message) + set_plugin_context(context) + + config_dict = sample_config_dict.copy() + config_dict["delivery"] = { + **sample_config_dict["delivery"], + "napcat_stream_mode": "disabled", + } + config = SetuPluginConfig(**config_dict) + payload = ImagePayload( + urls=(), + raw_bytes=(), + file_paths=(image_path,), + items=(image_path,), + r18=False, + tags=(), + ) + + results = [ + item async for item in ImageSender(config).send_images(payload, mock_event) + ] + + assert results == [{"send_success": True, "image_count": 1}] + image_comp = sent_results[-1].result_chain[0] + assert isinstance(image_comp, Comp.Image) + assert isinstance(image_comp.file, str) + assert image_comp.file.startswith("base64://") + + +@pytest.mark.asyncio +async def test_direct_send_strategy_passthroughs_onebot_stream_refs(mock_event) -> None: + context = MagicMock() + strategy = DirectSendStrategy(context) + + mock_event.platform.name = "aiocqhttp" + mock_event.get_group_id.return_value = "123456" + mock_event.get_sender_id.return_value = "654321" + mock_event.bot = MagicMock() + mock_event.bot.send_group_msg = AsyncMock(return_value={"message_id": "ok"}) + + success = await strategy.send(mock_event, [Comp.Image(file="stream://image")]) + + assert success is True + mock_event.bot.send_group_msg.assert_awaited_once_with( + group_id=123456, + message=[{"type": "image", "data": {"file": "stream://image"}}], + ) diff --git a/tests/infrastructure/test_permission_service.py b/tests/infrastructure/test_permission_service.py new file mode 100644 index 0000000..264e471 --- /dev/null +++ b/tests/infrastructure/test_permission_service.py @@ -0,0 +1,75 @@ +"""Tests for PermissionService.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + + +from astrbot_plugin_setu.src.infrastructure.permission_service import PermissionService + + +class TestPermissionService: + """Test PermissionService.""" + + def test_is_admin_with_role(self, mock_event) -> None: + """Test is_admin with admin role.""" + mock_event.is_admin = MagicMock(return_value=True) + assert PermissionService.is_admin(mock_event) is True + + def test_is_admin_with_super_user(self, mock_event) -> None: + """Test is_admin with super user.""" + mock_event.is_admin = MagicMock(return_value=False) + mock_event.is_super_user = MagicMock(return_value=True) + assert PermissionService.is_admin(mock_event) is True + + def test_is_admin_with_sender_role(self, mock_event) -> None: + """Test is_admin with sender role.""" + mock_event.is_admin = MagicMock(return_value=False) + mock_event.is_super_user = MagicMock(return_value=False) + + sender = MagicMock() + sender.role = "admin" + mock_event.message_obj = MagicMock() + mock_event.message_obj.sender = sender + + assert PermissionService.is_admin(mock_event) is True + + def test_is_admin_false(self, mock_event) -> None: + """Test is_admin returns False for regular user.""" + mock_event.is_admin = MagicMock(return_value=False) + mock_event.is_super_user = MagicMock(return_value=False) + assert PermissionService.is_admin(mock_event) is False + + def test_is_super_user(self, mock_event) -> None: + """Test is_super_user.""" + mock_event.is_super_user = MagicMock(return_value=True) + assert PermissionService.is_super_user(mock_event) is True + + def test_require_admin_granted(self, mock_event) -> None: + """Test require_admin when permission granted.""" + mock_event.is_admin = MagicMock(return_value=True) + has_perm, msg = PermissionService.require_admin(mock_event) + assert has_perm is True + assert msg == "" + + def test_require_admin_denied(self, mock_event) -> None: + """Test require_admin when permission denied.""" + mock_event.is_admin = MagicMock(return_value=False) + mock_event.is_super_user = MagicMock(return_value=False) + has_perm, msg = PermissionService.require_admin(mock_event) + assert has_perm is False + assert "权限不足" in msg + + def test_require_super_user_granted(self, mock_event) -> None: + """Test require_super_user when permission granted.""" + mock_event.is_super_user = MagicMock(return_value=True) + has_perm, msg = PermissionService.require_super_user(mock_event) + assert has_perm is True + assert msg == "" + + def test_require_super_user_denied(self, mock_event) -> None: + """Test require_super_user when permission denied.""" + mock_event.is_super_user = MagicMock(return_value=False) + has_perm, msg = PermissionService.require_super_user(mock_event) + assert has_perm is False + assert "权限不足" in msg diff --git a/tests/infrastructure/test_provider_init_from_config.py b/tests/infrastructure/test_provider_init_from_config.py new file mode 100644 index 0000000..23df637 --- /dev/null +++ b/tests/infrastructure/test_provider_init_from_config.py @@ -0,0 +1,36 @@ +from __future__ import annotations + +from astrbot_plugin_setu.src.infrastructure.providers import ( + clear_provider, + get_provider, + init_provider_from_config, +) +from astrbot_plugin_setu.src.shared.config import SetuPluginConfig + + +def test_init_provider_from_config_uses_current_proxy_values( + sample_config_dict, +) -> None: + clear_provider() + config_dict = sample_config_dict.copy() + config_dict["setu_general"] = { + **sample_config_dict["setu_general"], + "api_type": "lolicon", + } + config_dict["api"] = { + **sample_config_dict["api"], + "lolicon": { + **sample_config_dict["api"]["lolicon"], + "proxy": "proxy.example.com", + }, + "atri": { + **sample_config_dict["api"]["atri"], + "proxy": "atri-proxy.example.com", + }, + } + config = SetuPluginConfig(**config_dict) + + init_provider_from_config(config) + provider = get_provider() + + assert getattr(provider, "proxy", None) == "proxy.example.com" diff --git a/tests/infrastructure/test_session_config_repo.py b/tests/infrastructure/test_session_config_repo.py new file mode 100644 index 0000000..0f72aa6 --- /dev/null +++ b/tests/infrastructure/test_session_config_repo.py @@ -0,0 +1,70 @@ +"""Tests for JSON-backed session configuration.""" + +from __future__ import annotations + +import json + +import pytest +from astrbot_plugin_setu.src.application.session_config import ( + SessionConfigService, + SessionConfigValidationError, +) +from astrbot_plugin_setu.src.application.settings import clear_application_config +from astrbot_plugin_setu.src.infrastructure.persistence.session_config_json_repository import ( + JsonSessionConfigRepository, +) + +pytestmark = pytest.mark.asyncio + + +@pytest.fixture(autouse=True) +def clear_config() -> None: + """Keep global config defaults deterministic.""" + clear_application_config() + + +async def test_session_config_repo_creates_initial_file(temp_data_dir) -> None: + """Repository initializes the independent session override store.""" + repo = JsonSessionConfigRepository(temp_data_dir) + await repo.initialize() + + data = json.loads((temp_data_dir / "session_overrides.json").read_text()) + assert data == {"version": 1, "sessions": []} + + +async def test_session_config_service_sets_and_clears_values(temp_data_dir) -> None: + """Service validates values and keeps a record when clearing overrides.""" + repo = JsonSessionConfigRepository(temp_data_dir) + await repo.initialize() + service = SessionConfigService(repo) + + snapshot = await service.set_value( + "platform:group:1000", "group", "setu.content_mode", "r18", "测试群" + ) + assert snapshot.overrides["setu.content_mode"] == "r18" + assert snapshot.effective["setu.content_mode"] == "r18" + + snapshot = await service.set_value( + "platform:group:1000", "group", "setu.auto_revoke", "on", "测试群" + ) + assert snapshot.overrides["setu.auto_revoke"] is True + + snapshot = await service.clear( + "platform:group:1000", "group", "setu.content_mode", "测试群" + ) + assert "setu.content_mode" not in snapshot.overrides + assert snapshot.effective["setu.content_mode"] == "sfw" + + sessions = await service.list_snapshots() + assert len(sessions) == 1 + assert sessions[0].session_id == "platform:group:1000" + + +async def test_session_config_service_rejects_unknown_key(temp_data_dir) -> None: + """Unknown config keys are rejected before persistence.""" + repo = JsonSessionConfigRepository(temp_data_dir) + await repo.initialize() + service = SessionConfigService(repo) + + with pytest.raises(SessionConfigValidationError): + await service.set_value("session", "private", "unknown.key", "value") diff --git a/tests/modules/__init__.py b/tests/modules/__init__.py new file mode 100644 index 0000000..d6cce43 --- /dev/null +++ b/tests/modules/__init__.py @@ -0,0 +1 @@ +"""Module tests for AstrBot Setu plugin.""" diff --git a/tests/modules/test_fortune_module.py b/tests/modules/test_fortune_module.py new file mode 100644 index 0000000..5ebc4ca --- /dev/null +++ b/tests/modules/test_fortune_module.py @@ -0,0 +1,58 @@ +"""Tests for FortuneModule class.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +from astrbot_plugin_setu.src.infrastructure.astrbot.commands import ( + FortuneCommandHandler, +) + + +@pytest.fixture +def mock_plugin_context() -> MagicMock: + """Create mock AstrBot plugin context.""" + context = MagicMock() + context.get_config = MagicMock(return_value={}) + context.logger = MagicMock() + return context + + +@pytest.fixture +def mock_config() -> MagicMock: + """Create mock AstrBot config.""" + config = MagicMock() + config.__getitem__ = MagicMock(return_value={}) + config.get = MagicMock(return_value=None) + return config + + +class TestFortuneCommandHandler: + """Test Fortune command adapter initialization.""" + + def test_module_init(self, mock_plugin_context, mock_config) -> None: + """Test handler initialization.""" + handler = FortuneCommandHandler() + assert handler is not None + + def test_module_has_required_methods( + self, mock_plugin_context, mock_config + ) -> None: + """Test module has required methods.""" + handler = FortuneCommandHandler() + + assert hasattr(handler, "fortune_command") + assert hasattr(handler, "refresh_fortune_command") + assert callable(handler.fortune_command) + assert callable(handler.refresh_fortune_command) + + def test_module_has_decorators(self, mock_plugin_context, mock_config) -> None: + """Test command adapter methods exist.""" + handler = FortuneCommandHandler() + + assert hasattr(handler, "refresh_group_fortune_command") + assert hasattr(handler, "refresh_all_fortune_command") + assert hasattr(handler, "enable_fortune_group_command") + assert hasattr(handler, "disable_fortune_group_command") diff --git a/tests/modules/test_setu_module.py b/tests/modules/test_setu_module.py new file mode 100644 index 0000000..5ef0f3d --- /dev/null +++ b/tests/modules/test_setu_module.py @@ -0,0 +1,71 @@ +"""Tests for SetuModule class.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +from astrbot_plugin_setu.src.infrastructure.astrbot.commands import ( + SessionConfigCommandHandler, + SetuCommandHandler, +) + + +@pytest.fixture +def mock_plugin_context() -> MagicMock: + """Create mock AstrBot plugin context.""" + context = MagicMock() + context.get_config = MagicMock(return_value={}) + context.logger = MagicMock() + return context + + +@pytest.fixture +def mock_config() -> MagicMock: + """Create mock AstrBot config.""" + config = MagicMock() + config.__getitem__ = MagicMock(return_value={}) + config.get = MagicMock(return_value=None) + return config + + +class TestSetuCommandHandler: + """Test Setu command adapter initialization.""" + + def test_module_init(self, mock_plugin_context, mock_config) -> None: + """Test handler initialization.""" + handler = SetuCommandHandler() + assert handler is not None + + def test_module_has_required_methods( + self, mock_plugin_context, mock_config + ) -> None: + """Test module has required methods.""" + handler = SetuCommandHandler() + + assert hasattr(handler, "get_random_picture") + assert hasattr(handler, "setu_command") + assert hasattr(handler, "_get_effective_content_mode") + assert callable(handler.get_random_picture) + assert callable(handler.setu_command) + + def test_module_has_decorators(self, mock_plugin_context, mock_config) -> None: + """Test command adapter methods exist.""" + handler = SetuCommandHandler() + + assert hasattr(handler, "setu_command") + assert hasattr(handler, "_fetch_and_send_images") + + +class TestSessionConfigCommandHandler: + """Test unified session config adapter initialization.""" + + def test_module_has_required_methods(self) -> None: + """Test command adapter methods exist.""" + handler = SessionConfigCommandHandler() + + assert hasattr(handler, "session_config_command") + assert hasattr(handler, "_llm_get_session_config") + assert hasattr(handler, "_llm_set_session_config") + assert hasattr(handler, "_llm_clear_session_config") diff --git a/tests/shared/test_config_models.py b/tests/shared/test_config_models.py new file mode 100644 index 0000000..4844424 --- /dev/null +++ b/tests/shared/test_config_models.py @@ -0,0 +1,10 @@ +from __future__ import annotations + +from astrbot_plugin_setu.src.shared.config import SetuPluginConfig + + +def test_provider_config_accepts_empty_aspect_ratio(sample_config_dict) -> None: + config = SetuPluginConfig(**sample_config_dict) + + assert config.api.lolicon.aspect_ratio is None + assert config.api.atri.aspect_ratio is None diff --git a/tests/test_main_command_routing.py b/tests/test_main_command_routing.py new file mode 100644 index 0000000..30f0a39 --- /dev/null +++ b/tests/test_main_command_routing.py @@ -0,0 +1,37 @@ +from __future__ import annotations + +from astrbot_plugin_setu.main import ( + _resolve_fortune_refresh_target, + _resolve_fortune_toggle_action, + _resolve_fortune_user_action, +) + + +def test_resolve_fortune_refresh_target_from_new_command(mock_event) -> None: + mock_event.message_str = "/运势刷新 本群" + assert _resolve_fortune_refresh_target(mock_event, "本群") == "group" + + +def test_resolve_fortune_refresh_target_from_legacy_alias(mock_event) -> None: + mock_event.message_str = "/刷新全局今日运势" + assert _resolve_fortune_refresh_target(mock_event, "") == "all" + + +def test_resolve_fortune_toggle_action_from_new_command(mock_event) -> None: + mock_event.message_str = "/运势开关 关" + assert _resolve_fortune_toggle_action(mock_event, "关") == "disable" + + +def test_resolve_fortune_toggle_action_from_legacy_alias(mock_event) -> None: + mock_event.message_str = "/开启运势" + assert _resolve_fortune_toggle_action(mock_event, "") == "enable" + + +def test_resolve_fortune_user_action_from_new_command(mock_event) -> None: + mock_event.message_str = "/运势用户 拉黑 12345" + assert _resolve_fortune_user_action(mock_event, "拉黑 12345") == ("block", "12345") + + +def test_resolve_fortune_user_action_from_legacy_alias(mock_event) -> None: + mock_event.message_str = "/取消运势信任 12345" + assert _resolve_fortune_user_action(mock_event, "12345") == ("untrust", "12345") diff --git a/tests/test_main_config_source.py b/tests/test_main_config_source.py new file mode 100644 index 0000000..1436d77 --- /dev/null +++ b/tests/test_main_config_source.py @@ -0,0 +1,81 @@ +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from astrbot_plugin_setu.main import SetuPlugin + + +@pytest.mark.asyncio +async def test_initialize_uses_plugin_config_not_context_config( + monkeypatch, sample_config_dict +) -> None: + context = MagicMock() + context.get_config.return_value = { + "api": {"lolicon": {"proxy": "wrong.example.com"}} + } + config = MagicMock() + config.items.return_value = sample_config_dict.items() + + plugin = SetuPlugin(context, config) + plugin.name = "astrbot_plugin_setu" + + captured: dict[str, object] = {} + + def fake_init_config(raw_config): + captured["raw_config"] = raw_config + return MagicMock( + api_type="lolicon", + custom_api_configs=[], + multi_api_strategy="round_robin", + proxy=raw_config["api"]["lolicon"]["proxy"], + image_size="original", + aspect_ratio="", + uid=[], + keyword="", + atri_proxy=raw_config["api"]["atri"]["proxy"], + atri_image_size="original", + atri_aspect_ratio="", + atri_uid=[], + atri_keyword="", + cache_enabled=True, + cache_ttl_hours=2, + cache_max_items=200, + cache_cleanup_on_start=True, + ) + + monkeypatch.setattr("astrbot_plugin_setu.main.init_config", fake_init_config) + monkeypatch.setattr( + "astrbot_plugin_setu.main.set_plugin_context", lambda _ctx: None + ) + monkeypatch.setattr( + "astrbot_plugin_setu.main.StarTools.get_data_dir", lambda _name: "/tmp" + ) + monkeypatch.setattr( + "astrbot_plugin_setu.main.init_provider_from_config", lambda _cfg: None + ) + monkeypatch.setattr( + "astrbot_plugin_setu.main.init_access_control_repo", AsyncMock() + ) + monkeypatch.setattr("astrbot_plugin_setu.main.init_fortune_repo", AsyncMock()) + monkeypatch.setattr( + "astrbot_plugin_setu.main.init_session_config_repo", AsyncMock() + ) + monkeypatch.setattr("astrbot_plugin_setu.main.init_send_cache", AsyncMock()) + monkeypatch.setattr( + "astrbot_plugin_setu.main.register_setu_llm_tools", lambda: None + ) + monkeypatch.setattr( + "astrbot_plugin_setu.main.register_fortune_llm_tools", lambda: None + ) + monkeypatch.setattr( + "astrbot_plugin_setu.main.register_session_config_llm_tools", lambda: None + ) + + await plugin.initialize() + + assert ( + captured["raw_config"]["api"]["lolicon"]["proxy"] + == sample_config_dict["api"]["lolicon"]["proxy"] + ) diff --git a/utils.py b/utils.py deleted file mode 100644 index e70a022..0000000 --- a/utils.py +++ /dev/null @@ -1,115 +0,0 @@ -"""Setu 插件的工具函数。""" - -from __future__ import annotations - -import base64 -import random - -from astrbot.api.event import MessageChain -from astrbot.api.message_components import Image - - -def obfuscate_image_bytes(data: bytes) -> bytes: - """对图片数据进行最小程度的混淆以绕过图片哈希检测。 - - 在图片数据末尾添加随机字节来改变其哈希值,同时保持图片视觉内容不变。 - """ - noise = bytes(random.randint(0, 255) for _ in range(8)) - return data + noise - - -def create_image_chain(images: list[bytes], text: str | None = None) -> MessageChain: - """创建包含文本和图片的消息链。 - - 参数: - images: 图片字节数据列表。 - text: 可选的文本消息。 - - 返回: - 可发送的消息链。 - """ - chain = MessageChain() - if text: - chain.message(text) - for img_data in images: - b64 = base64.b64encode(img_data).decode("ascii") - chain.chain.append(Image.fromBase64(b64)) - return chain - - -def create_obfuscated_image_chain( - images: list[bytes], text: str | None = None -) -> MessageChain: - """创建包含混淆图片的消息链。 - - 参数: - images: 图片字节数据列表。 - text: 可选的文本消息。 - - 返回: - 包含混淆图片的消息链。 - """ - chain = MessageChain() - if text: - chain.message(text) - for img_data in images: - obf = obfuscate_image_bytes(img_data) - b64 = base64.b64encode(obf).decode("ascii") - chain.chain.append(Image.fromBase64(b64)) - return chain - - -def cn_to_an(cn_num: str) -> int: - """中文数字转阿拉伯数字 - - 参数: - cn_num: 中文数字。 - - 返回: - 阿拉伯数字。 - """ - - # 定义对应关系 - num_dict = { - "零": 0, - "一": 1, - "二": 2, - "两": 2, - "三": 3, - "四": 4, - "五": 5, - "六": 6, - "七": 7, - "八": 8, - "九": 9, - } - unit_dict = {"十": 10, "百": 100, "千": 1000, "万": 10000, "亿": 100000000} - - res = 0 - unit = 1 # 当前单位 - temp = 0 # 累加临时值 - - # 倒序处理 - for i in range(len(cn_num) - 1, -1, -1): - char = cn_num[i] - if char in unit_dict: - val = unit_dict[char] - if val >= 10000: # 处理万、亿大单位 - if val > unit: - unit = val - temp = 0 # 重置临时值 - else: - unit *= val - else: - if val >= temp: - temp = val - else: - temp *= val - elif char in num_dict: - res += num_dict[char] * (temp if temp > 0 else 1) * unit - temp = 0 - - # 特殊处理开头是“十”的情况,如“十四” - if cn_num.startswith("十"): - res += 10 - return res