From 7a0750f04339b797332b24bda1005a7556a0c4c0 Mon Sep 17 00:00:00 2001 From: hp0912 <809211365@qq.com> Date: Sun, 13 Sep 2026 12:56:56 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E6=B6=88=E9=99=A4=E4=BB=A3=E7=A0=81?= =?UTF-8?q?=E9=9D=99=E6=80=81=E7=B1=BB=E5=9E=8B=E8=AD=A6=E5=91=8A?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 14 + eslint.config.cjs | 17 + package.json | 19 + pyrightconfig.json | 16 + requirements-dev.txt | 27 ++ ruff.toml | 4 + .../scripts/create_scheduled_task.py | 2 +- skills/docx/scripts/_document_builder.py | 16 +- skills/docx/scripts/_tracked_replace.py | 2 +- skills/docx/scripts/add_comment.py | 2 +- skills/docx/scripts/convert_document.py | 3 +- skills/docx/scripts/create_document.py | 2 +- skills/docx/scripts/edit_document.py | 4 +- skills/docx/scripts/inspect_document.py | 1 - skills/docx/tests/test_document_workflows.py | 27 +- .../scripts/bootstrap.py | 2 +- .../scripts/video_understanding.py | 4 +- .../scripts/export_chat_history.py | 11 +- .../scripts/find_recent_chat_media.py | 2 +- .../scripts/image_recognition.py | 4 +- .../image-to-image/scripts/image_to_image.py | 4 +- skills/kfc/scripts/kfc.py | 2 +- skills/pdf/scripts/_render_html.cjs | 8 +- skills/pdf/scripts/compile_latex.py | 115 +++-- skills/pdf/scripts/convert_to_pdf.py | 122 +++-- skills/pdf/scripts/create_design_pdf.py | 13 +- skills/pdf/scripts/edit_pdf.py | 9 +- skills/pdf/scripts/extract_text.py | 2 + skills/pdf/scripts/manage_pdf.py | 1 - skills/pdf/scripts/ocr_text.py | 2 + skills/pdf/tests/test_pdf_workflows.py | 39 +- skills/pptx/scripts/inspect_presentation.py | 7 +- skills/pptx/scripts/ocr_presentation.py | 5 +- skills/pptx/scripts/validate_presentation.py | 5 +- skills/send-complex-message/SKILL.md | 60 ++- .../send-complex-message/scripts/bootstrap.py | 2 +- .../scripts/send_complex_message.py | 209 +++++++- .../send-mention-message/scripts/bootstrap.py | 2 +- .../scripts/send_mention_message.py | 2 +- skills/text-to-image/scripts/bootstrap.py | 2 +- skills/text-to-image/scripts/text_to_image.py | 4 +- skills/video-generation/scripts/bootstrap.py | 2 +- .../scripts/video_generation.py | 2 +- skills/voice-message/scripts/bootstrap.py | 2 +- skills/voice-message/scripts/voice_message.py | 4 +- skills/web-page/scripts/tsconfig.json | 2 + skills/web-page/scripts/web_page.test.ts | 140 +++--- skills/xlsx/scripts/_xlsx_common.py | 10 + skills/xlsx/scripts/_xlsx_data.py | 23 +- skills/xlsx/scripts/analyze_workbook.py | 10 +- skills/xlsx/scripts/apply_workbook.py | 30 +- skills/xlsx/scripts/convert_workbook.py | 6 +- skills/xlsx/scripts/inspect_workbook.py | 14 +- skills/xlsx/scripts/model_workbook.py | 19 +- skills/xlsx/tests/test_data_workflows.py | 456 +++++++++++++----- tests/test_send_complex_message.py | 244 ++++++++++ 56 files changed, 1371 insertions(+), 387 deletions(-) create mode 100644 eslint.config.cjs create mode 100644 package.json create mode 100644 pyrightconfig.json create mode 100644 requirements-dev.txt create mode 100644 ruff.toml diff --git a/README.md b/README.md index 2cb3978..9ac83e4 100644 --- a/README.md +++ b/README.md @@ -2,6 +2,20 @@ 微信机器人 Skills +**开发检查** + +使用 Python 3.12 和 Node.js 24+,在仓库根目录运行: + +```sh +python3.12 -m venv .venv +source .venv/bin/activate +python -m pip install -r requirements-dev.txt +npm install +npm run check +``` + +检查覆盖所有 Python、JavaScript、TypeScript 脚本及测试,包含 Ruff、Pyright、TypeScript 和 ESLint;类型错误、未使用代码及检查警告会导致命令失败。VS Code / Pylance 请选择 `.venv/bin/python` 作为解释器,以使用相同的依赖和类型信息。 + **系统自动注入的环境变量** - ROBOT_WECHAT_CLIENT_PORT: 机器人客户端服务端口,可用于在 SKILL 脚本直接调用客户端接口 `http://127.0.0.1:{ROBOT_WECHAT_CLIENT_PORT}/api/v1/xxxxx` diff --git a/eslint.config.cjs b/eslint.config.cjs new file mode 100644 index 0000000..940acf3 --- /dev/null +++ b/eslint.config.cjs @@ -0,0 +1,17 @@ +const js = require('@eslint/js'); +const globals = require('globals'); + +module.exports = [ + { + ...js.configs.recommended, + files: ['skills/**/*.js', 'skills/**/*.cjs', 'eslint.config.cjs'], + languageOptions: { + sourceType: 'commonjs', + globals: globals.node, + }, + }, + { + files: ['skills/pdf/scripts/_render_html.cjs'], + languageOptions: { globals: globals.browser }, + }, +]; diff --git a/package.json b/package.json new file mode 100644 index 0000000..c3022b6 --- /dev/null +++ b/package.json @@ -0,0 +1,19 @@ +{ + "name": "wechat-robot-skills-checks", + "private": true, + "scripts": { + "check": "npm run lint:python && npm run typecheck:python && npm run typecheck:ts && npm run lint:js", + "lint:python": "ruff check .", + "typecheck:python": "pyright --warnings", + "typecheck:ts": "tsc --noEmit --project skills/web-page/scripts/tsconfig.json", + "lint:js": "eslint eslint.config.cjs \"skills/**/*.js\" \"skills/**/*.cjs\" --max-warnings 0" + }, + "devDependencies": { + "@eslint/js": "9.39.5", + "@types/node": "25.9.1", + "eslint": "9.39.5", + "globals": "17.12.0", + "pyright": "1.1.414", + "typescript": "6.0.3" + } +} diff --git a/pyrightconfig.json b/pyrightconfig.json new file mode 100644 index 0000000..2503993 --- /dev/null +++ b/pyrightconfig.json @@ -0,0 +1,16 @@ +{ + "include": ["skills", "tests"], + "exclude": ["**/.venv", "**/node_modules", "**/__pycache__"], + "typeCheckingMode": "standard", + "pythonVersion": "3.12", + "executionEnvironments": [ + {"root": "skills/docx", "extraPaths": ["skills/docx/scripts"]}, + {"root": "skills/pdf", "extraPaths": ["skills/pdf/scripts"]}, + {"root": "skills/pptx", "extraPaths": ["skills/pptx/scripts"]}, + {"root": "skills/xlsx", "extraPaths": ["skills/xlsx/scripts"]}, + {"root": "."} + ], + "reportUnusedImport": "warning", + "reportUnusedVariable": "warning", + "reportUnnecessaryTypeIgnoreComment": "warning" +} diff --git a/requirements-dev.txt b/requirements-dev.txt new file mode 100644 index 0000000..a329dde --- /dev/null +++ b/requirements-dev.txt @@ -0,0 +1,27 @@ +# 本地静态检查、类型信息及 Python 回归测试依赖。 +ruff==0.16.7 +brotli==1.2.0 +defusedxml==0.7.1 +lxml==6.1.1 +markitdown==0.1.7 +numpy==2.5.3 +openai==3.13.0 +openpyxl==3.1.5 +pandas==2.2.3 +pdfplumber==0.11.9 +pillow==12.3.0 +pymysql==1.2.0 +pypdf==6.10.0 +python-docx==1.2.0 +python-pptx==1.0.2 +rapidocr==3.9.2 +reportlab==4.4.9 +scikit-learn==1.9.1 +scipy==1.18.1 + +# 补齐第三方库的静态类型信息。 +pandas-stubs==3.0.5.260730 +scikit-learn-stubs==0.0.3 +scipy-stubs==1.18.1.0 +types-lxml==2026.2.16 +types-openpyxl==3.1.5.20260827 diff --git a/ruff.toml b/ruff.toml new file mode 100644 index 0000000..f706d7a --- /dev/null +++ b/ruff.toml @@ -0,0 +1,4 @@ +target-version = "py312" + +[lint] +select = ["E4", "E7", "E9", "F", "W", "RUF100"] diff --git a/skills/create-scheduled-task/scripts/create_scheduled_task.py b/skills/create-scheduled-task/scripts/create_scheduled_task.py index 3377d6e..4141949 100644 --- a/skills/create-scheduled-task/scripts/create_scheduled_task.py +++ b/skills/create-scheduled-task/scripts/create_scheduled_task.py @@ -18,7 +18,7 @@ from typing import Any, Literal, NoReturn, TypedDict try: from zoneinfo import ZoneInfo except ImportError: # pragma: no cover - Python 3.8 fallback - ZoneInfo = None # type: ignore[assignment,misc] + ZoneInfo = None sys.stderr = sys.stdout diff --git a/skills/docx/scripts/_document_builder.py b/skills/docx/scripts/_document_builder.py index 6f04f8a..2aa36df 100644 --- a/skills/docx/scripts/_document_builder.py +++ b/skills/docx/scripts/_document_builder.py @@ -2,11 +2,9 @@ from __future__ import annotations -from copy import deepcopy -from pathlib import Path -from typing import Any, Iterable, Optional +from typing import Any, Optional -from _docx_common import W_NS, input_file, qn +from _docx_common import input_file, qn MAX_BLOCKS = 1_000 @@ -132,7 +130,7 @@ def _add_hyperlink( is_external=True, ) hyperlink = OxmlElement("w:hyperlink") - hyperlink.set(f"{{http://schemas.openxmlformats.org/officeDocument/2006/relationships}}id", relationship_id) + hyperlink.set("{http://schemas.openxmlformats.org/officeDocument/2006/relationships}id", relationship_id) run = paragraph.add_run(text) apply_run_style(run, raw_spec) run_properties = run._element.get_or_add_rPr() @@ -232,7 +230,7 @@ def add_paragraph_from_spec( style: Optional[str] = None, default_run_style: Optional[dict[str, Any]] = None, ) -> Any: - spec = ( + spec: dict[str, Any] = ( {"text": raw_spec} if isinstance(raw_spec, str) else expect_object(raw_spec, "paragraph") @@ -442,17 +440,17 @@ def _add_toc(container: Any, raw_spec: Any) -> Any: def _add_horizontal_rule(container: Any, raw_spec: Any) -> Any: from docx.oxml import OxmlElement + from lxml.etree import SubElement spec = expect_object(raw_spec, "horizontal_rule") paragraph = container.add_paragraph() properties = paragraph._p.get_or_add_pPr() borders = OxmlElement("w:pBdr") - bottom = OxmlElement("w:bottom") + bottom = SubElement(borders, qn("bottom")) bottom.set(qn("val"), str(spec.get("style", "single"))) bottom.set(qn("sz"), str(int(spec.get("size", 6)))) bottom.set(qn("space"), str(int(spec.get("space", 1)))) bottom.set(qn("color"), color(spec.get("color", "808080"), "rule.color")) - borders.append(bottom) properties.append(borders) return paragraph @@ -587,7 +585,7 @@ def _points(value: float) -> Any: def apply_page_settings(section: Any, raw_spec: Any) -> None: from docx.enum.section import WD_ORIENT - from docx.shared import Cm, Inches, Mm + from docx.shared import Inches, Mm spec = expect_object(raw_spec, "page") size = str(spec.get("size", "A4")).upper() diff --git a/skills/docx/scripts/_tracked_replace.py b/skills/docx/scripts/_tracked_replace.py index 03e0444..6507f39 100644 --- a/skills/docx/scripts/_tracked_replace.py +++ b/skills/docx/scripts/_tracked_replace.py @@ -9,7 +9,7 @@ from datetime import datetime, timezone from pathlib import Path from typing import Any -from _docx_common import W_NS, parse_xml_bytes, qn +from _docx_common import parse_xml_bytes, qn class TrackedReplacement: diff --git a/skills/docx/scripts/add_comment.py b/skills/docx/scripts/add_comment.py index c15eeda..158a1b5 100644 --- a/skills/docx/scripts/add_comment.py +++ b/skills/docx/scripts/add_comment.py @@ -198,7 +198,7 @@ def main() -> dict[str, Any]: os.close(descriptor) temp_path = Path(temp_name) try: - document.save(temp_path) + document.save(str(temp_path)) archive = inspect_archive(temp_path) Document(str(temp_path)) publish_file(temp_path, destination, overwrite=args.overwrite) diff --git a/skills/docx/scripts/convert_document.py b/skills/docx/scripts/convert_document.py index 5ca615f..5676af8 100644 --- a/skills/docx/scripts/convert_document.py +++ b/skills/docx/scripts/convert_document.py @@ -13,7 +13,6 @@ from _docx_common import ( DOCUMENT_OUTPUT_SUFFIXES, WORD_INPUT_SUFFIXES, SkillArgumentParser, - find_program, input_file, output_file, publish_file, @@ -74,7 +73,7 @@ def _extract_text( pandoc = shutil.which("pandoc") if pandoc: target = "gfm" if markdown else "plain" - completed = run_program( + run_program( [ pandoc, f"--track-changes={track_changes}", diff --git a/skills/docx/scripts/create_document.py b/skills/docx/scripts/create_document.py index b31aa89..1aa1b9f 100644 --- a/skills/docx/scripts/create_document.py +++ b/skills/docx/scripts/create_document.py @@ -96,7 +96,7 @@ def main() -> dict[str, Any]: os.close(descriptor) temp_path = Path(temp_name) try: - document.save(temp_path) + document.save(str(temp_path)) archive = inspect_archive(temp_path) if archive["missing_required_parts"]: raise ValueError( diff --git a/skills/docx/scripts/edit_document.py b/skills/docx/scripts/edit_document.py index e3dd6b6..6ec9180 100644 --- a/skills/docx/scripts/edit_document.py +++ b/skills/docx/scripts/edit_document.py @@ -8,7 +8,7 @@ import re import tempfile import zipfile from pathlib import Path -from typing import Any, Iterable, Optional +from typing import Any, Iterable from _document_builder import ( add_blocks, @@ -445,7 +445,7 @@ def main() -> dict[str, Any]: os.close(descriptor) temp_path = Path(temp_name) try: - document.save(temp_path) + document.save(str(temp_path)) archive = inspect_archive(temp_path) if archive["missing_required_parts"]: raise ValueError( diff --git a/skills/docx/scripts/inspect_document.py b/skills/docx/scripts/inspect_document.py index a28ffa0..b10a632 100644 --- a/skills/docx/scripts/inspect_document.py +++ b/skills/docx/scripts/inspect_document.py @@ -9,7 +9,6 @@ from typing import Any, Optional from _docx_common import ( DOCX_INPUT_SUFFIXES, - NS, SkillArgumentParser, W_NS, input_file, diff --git a/skills/docx/tests/test_document_workflows.py b/skills/docx/tests/test_document_workflows.py index d1f6976..97c66a2 100644 --- a/skills/docx/tests/test_document_workflows.py +++ b/skills/docx/tests/test_document_workflows.py @@ -10,19 +10,22 @@ import tempfile import unittest from copy import deepcopy from pathlib import Path +from typing import cast from unittest.mock import patch -sys.dont_write_bytecode = True sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "scripts")) from docx import Document from docx.oxml import OxmlElement from docx.oxml.ns import qn +from lxml.etree import _Element from reportlab.pdfgen.canvas import Canvas import _docx_common as common import compile_typst +sys.dont_write_bytecode = True + class DocumentWorkflows(unittest.TestCase): def setUp(self): @@ -45,14 +48,14 @@ class DocumentWorkflows(unittest.TestCase): p.add_run("HEL").bold = True p.add_run("LO") p.add_run(" AFTER").italic = True - doc.save(self.source) + doc.save(str(self.source)) return doc def edit(self, operations, *args): output = self.root / "edited.docx" result = self.call("edit_document", "--input", self.source, "--output", output, "--spec", json.dumps({"operations": operations}), *args) - return result, Document(output) + return result, Document(str(output)) def test_cross_run_preserves_unmodified_styles_and_escapes_text(self): self.fixture() @@ -67,7 +70,7 @@ class DocumentWorkflows(unittest.TestCase): def test_single_and_split_matches_are_both_replaced(self): doc = self.fixture() doc.add_paragraph("HELLO HELLO") - doc.save(self.source) + doc.save(str(self.source)) result, doc = self.edit([{"type": "replace_text", "find": "HELLO", "replace": "NEW"}]) self.assertEqual(result["operation_results"][0]["replacement_count"], 3) self.assertNotIn("HELLO", "".join(p.text for p in doc.paragraphs)) @@ -99,7 +102,7 @@ class DocumentWorkflows(unittest.TestCase): def test_multiple_tracked_matches_in_one_run(self): doc = Document() doc.add_paragraph("old old old") - doc.save(self.source) + doc.save(str(self.source)) result, doc = self.edit([{"type": "replace_text", "find": "old", "replace": "new"}], "--track-changes") self.assertEqual(result["tracked_replacement_count"], 3) self.assertEqual(len(doc.element.xpath(".//w:p/w:ins")), 3) @@ -117,8 +120,8 @@ class DocumentWorkflows(unittest.TestCase): mark = OxmlElement("w:bookmarkStart") mark.set(qn("w:id"), "7") mark.set(qn("w:name"), "target") - doc.paragraphs[0]._p.insert(2, mark) - doc.save(self.source) + cast(_Element, doc.paragraphs[0]._p).insert(2, mark) + doc.save(str(self.source)) with self.assertRaisesRegex(ValueError, "书签"): self.edit([{"type": "replace_text", "find": "HELLO", "replace": "new"}], "--track-changes") self.assertFalse((self.root / "edited.docx").exists()) @@ -128,14 +131,16 @@ class DocumentWorkflows(unittest.TestCase): spec = {"claims": [{"number": 1, "text": "一种方法 & 装置", "dependent": False}], "specification": {"field": "领域", "detailed": ["实现 <描述>"]}, "abstract": "摘要"} result = self.call("create_document", "--preset", "patent", "--output", output, "--spec", json.dumps(spec)) - doc = Document(output) + doc = Document(str(output)) self.assertEqual(len(doc.sections), 3) self.assertEqual([s.header.paragraphs[0].text for s in doc.sections], ["权利要求书", "说明书", "摘要"]) for section in doc.sections: - self.assertEqual(section._sectPr.find(qn("w:pgNumType")).get(qn("w:start")), "1") + page_numbers = section._sectPr.find(qn("w:pgNumType")) + assert page_numbers is not None + self.assertEqual(page_numbers.get(qn("w:start")), "1") self.assertFalse(section.header.is_linked_to_previous) self.assertFalse(section.footer.is_linked_to_previous) - self.assertTrue(any(p.style.name == "Heading 2" for p in doc.paragraphs)) + self.assertTrue(any(p.style is not None and p.style.name == "Heading 2" for p in doc.paragraphs)) self.assertEqual(self.call("validate_document", "--input", result["path"])["status"], "valid") def test_invalid_patent_numbering_does_not_publish_a_file(self): @@ -153,7 +158,7 @@ class DocumentWorkflows(unittest.TestCase): {"type": "section_break"}, {"type": "paragraph", "text": "Second"}], "sections": [{"index": 1, "header": {"text": "Second header"}}]} self.call("create_document", "--output", output, "--spec", json.dumps(spec)) - doc = Document(output) + doc = Document(str(output)) self.assertEqual(doc.tables[0].cell(1, 1).text, "B") self.assertEqual([s.header.paragraphs[0].text for s in doc.sections], ["Default", "Second header"]) diff --git a/skills/doubao-video-understanding/scripts/bootstrap.py b/skills/doubao-video-understanding/scripts/bootstrap.py index 39d4579..8959bb4 100644 --- a/skills/doubao-video-understanding/scripts/bootstrap.py +++ b/skills/doubao-video-understanding/scripts/bootstrap.py @@ -131,4 +131,4 @@ if __name__ == "__main__": raise except Exception: traceback.print_exc(file=sys.stdout) - raise SystemExit(1) \ No newline at end of file + raise SystemExit(1) diff --git a/skills/doubao-video-understanding/scripts/video_understanding.py b/skills/doubao-video-understanding/scripts/video_understanding.py index ec78402..e1f7c47 100644 --- a/skills/doubao-video-understanding/scripts/video_understanding.py +++ b/skills/doubao-video-understanding/scripts/video_understanding.py @@ -68,7 +68,7 @@ def _ensure_skill_venv_python() -> None: _ensure_skill_venv_python() try: - import pymysql # type: ignore # noqa: E402 + import pymysql except ModuleNotFoundError: _run_bootstrap() _py = _get_python_executable() @@ -362,4 +362,4 @@ if __name__ == "__main__": raise except Exception: traceback.print_exc(file=sys.stdout) - raise SystemExit(1) \ No newline at end of file + raise SystemExit(1) diff --git a/skills/export-chat-history/scripts/export_chat_history.py b/skills/export-chat-history/scripts/export_chat_history.py index 2e70ea1..c7ba3a4 100644 --- a/skills/export-chat-history/scripts/export_chat_history.py +++ b/skills/export-chat-history/scripts/export_chat_history.py @@ -3,6 +3,7 @@ from __future__ import annotations import argparse +import importlib import json import os import re @@ -20,7 +21,7 @@ from typing import Any, NoReturn try: from zoneinfo import ZoneInfo except ImportError: # pragma: no cover - Python 3.8 fallback - ZoneInfo = None # type: ignore[assignment,misc] + ZoneInfo = None sys.stderr = sys.stdout @@ -94,8 +95,8 @@ def _run_bootstrap() -> None: def _ensure_runtime_dependencies() -> None: try: - import openpyxl # noqa: F401 - import pymysql # noqa: F401 + importlib.import_module("openpyxl") + importlib.import_module("pymysql") return except ModuleNotFoundError: @@ -109,8 +110,8 @@ def _ensure_runtime_dependencies() -> None: venv_dir = (_skill_root() / ".venv").resolve() if Path(sys.prefix).resolve() == venv_dir: try: - import openpyxl # noqa: F401 - import pymysql # noqa: F401 + importlib.import_module("openpyxl") + importlib.import_module("pymysql") return except ModuleNotFoundError as exc: diff --git a/skills/find-recent-chat-media/scripts/find_recent_chat_media.py b/skills/find-recent-chat-media/scripts/find_recent_chat_media.py index 25eb9ba..74f9156 100644 --- a/skills/find-recent-chat-media/scripts/find_recent_chat_media.py +++ b/skills/find-recent-chat-media/scripts/find_recent_chat_media.py @@ -101,7 +101,7 @@ def _ensure_skill_venv_python() -> None: _ensure_skill_venv_python() try: - import pymysql # type: ignore[import-untyped] # noqa: E402 + import pymysql except ModuleNotFoundError: _run_bootstrap() python_executable = _get_python_executable() diff --git a/skills/image-recognition/scripts/image_recognition.py b/skills/image-recognition/scripts/image_recognition.py index 4ee4082..7ef6378 100644 --- a/skills/image-recognition/scripts/image_recognition.py +++ b/skills/image-recognition/scripts/image_recognition.py @@ -64,8 +64,8 @@ def _ensure_skill_venv_python() -> None: _ensure_skill_venv_python() try: - import pymysql # type: ignore # noqa: E402 - from openai import OpenAI # type: ignore # noqa: E402 + import pymysql + from openai import OpenAI except ModuleNotFoundError: _run_bootstrap() _py = _get_python_executable() diff --git a/skills/image-to-image/scripts/image_to_image.py b/skills/image-to-image/scripts/image_to_image.py index 12e0609..be71626 100644 --- a/skills/image-to-image/scripts/image_to_image.py +++ b/skills/image-to-image/scripts/image_to_image.py @@ -70,8 +70,8 @@ def _ensure_skill_venv_python() -> None: _ensure_skill_venv_python() try: - import pymysql # type: ignore # noqa: E402 - from openai import OpenAI # type: ignore # noqa: E402 + import pymysql + from openai import OpenAI except ModuleNotFoundError: _run_bootstrap() _py = _get_python_executable() diff --git a/skills/kfc/scripts/kfc.py b/skills/kfc/scripts/kfc.py index d99c971..b638ca9 100644 --- a/skills/kfc/scripts/kfc.py +++ b/skills/kfc/scripts/kfc.py @@ -43,4 +43,4 @@ if __name__ == "__main__": raise except Exception: traceback.print_exc(file=sys.stdout) - raise SystemExit(1) \ No newline at end of file + raise SystemExit(1) diff --git a/skills/pdf/scripts/_render_html.cjs b/skills/pdf/scripts/_render_html.cjs index d735475..e5bf64a 100644 --- a/skills/pdf/scripts/_render_html.cjs +++ b/skills/pdf/scripts/_render_html.cjs @@ -60,7 +60,7 @@ async function render(request) { if (!within(root, target) || !MIME[path.extname(target).toLowerCase()]) throw new Error('forbidden asset'); if (fs.statSync(target).size > 25 * 1024 * 1024) throw new Error('asset too large'); return route.fulfill({body:fs.readFileSync(target), contentType:MIME[path.extname(target).toLowerCase()], headers:{'Content-Security-Policy':policy}}); - } catch (_) { blocked.push('缺失或不允许的本地资源:' + relative.slice(0, 200)); return route.abort(); } + } catch { blocked.push('缺失或不允许的本地资源:' + relative.slice(0, 200)); return route.abort(); } }); await page.goto('https://pdf.local/' + encodeURIComponent(path.basename(request.input)), {waitUntil:'load'}); if (await page.evaluate(() => document.compatMode !== 'CSS1Compat')) @@ -84,9 +84,9 @@ async function render(request) { if (unsafe) throw new Error('Mermaid 不允许内嵌配置;主题和安全选项由固定渲染器设置'); await inject(page, path.join(libraries.mermaid,'dist/mermaid.min.js')); await page.evaluate(async () => { - mermaid.initialize({startOnLoad:false, securityLevel:'strict', theme:'neutral', maxTextSize:50000, + window.mermaid.initialize({startOnLoad:false, securityLevel:'strict', theme:'neutral', maxTextSize:50000, flowchart:{htmlLabels:false}, suppressErrorRendering:true}); - await mermaid.run({querySelector:'.mermaid'}); + await window.mermaid.run({querySelector:'.mermaid'}); }); } if (stats.math) { @@ -94,7 +94,7 @@ async function render(request) { await inject(page, path.join(libraries.katex,'dist/katex.min.js')); await page.evaluate(() => { for (const el of document.querySelectorAll('.math-inline,.math-display')) { - katex.render(el.textContent, el, {displayMode:el.classList.contains('math-display'), throwOnError:true, + window.katex.render(el.textContent, el, {displayMode:el.classList.contains('math-display'), throwOnError:true, trust:false, maxExpand:1000, maxSize:30, strict:'warn'}); } }); diff --git a/skills/pdf/scripts/compile_latex.py b/skills/pdf/scripts/compile_latex.py index 99e57c7..2975626 100644 --- a/skills/pdf/scripts/compile_latex.py +++ b/skills/pdf/scripts/compile_latex.py @@ -1,5 +1,6 @@ #!/usr/bin/env python3 """Compile a LaTeX project with cached Tectonic resources and shell escape disabled.""" + from pathlib import Path import re import shutil @@ -7,51 +8,105 @@ import subprocess import sys import tempfile -from _pdf_common import SkillArgumentParser, output_pdf, new_temp_pdf, publish_temp_file, run_cli +from _pdf_common import ( + SkillArgumentParser, + output_pdf, + new_temp_pdf, + publish_temp_file, + run_cli, +) def compile_document(args): from pypdf import PdfReader - source=Path(args.input).expanduser().resolve() - if not source.is_file() or source.suffix.lower()!='.tex' or not 0 < source.stat().st_size <= 2*1024*1024: - raise ValueError('输入需为不超过 2 MiB 的本地 .tex 文件') + + source = Path(args.input).expanduser().resolve() + if ( + not source.is_file() + or source.suffix.lower() != ".tex" + or not 0 < source.stat().st_size <= 2 * 1024 * 1024 + ): + raise ValueError("输入需为不超过 2 MiB 的本地 .tex 文件") if not 1 <= args.timeout <= 600: - raise ValueError('timeout 必须在 1–600 秒之间') - executable=shutil.which('tectonic') + raise ValueError("timeout 必须在 1–600 秒之间") + executable = shutil.which("tectonic") if not executable: - raise RuntimeError('基础镜像缺少预置 Tectonic,需要更新镜像') - target=output_pdf(args.output,args.overwrite) - temporary=new_temp_pdf(target) + raise RuntimeError("基础镜像缺少预置 Tectonic,需要更新镜像") + target = output_pdf(args.output, args.overwrite) + temporary = new_temp_pdf(target) try: - with tempfile.TemporaryDirectory(prefix='pdf-latex-') as folder: - command=[executable,'--untrusted','--only-cached','--keep-logs','--outdir',folder,str(source)] - completed=subprocess.run(command,cwd=source.parent,capture_output=True,text=True,timeout=args.timeout,check=False) - result=Path(folder)/(source.stem+'.pdf') - messages=(completed.stdout+'\n'+completed.stderr).splitlines() + with tempfile.TemporaryDirectory(prefix="pdf-latex-") as folder: + command = [ + executable, + "--untrusted", + "--only-cached", + "--keep-logs", + "--outdir", + folder, + str(source), + ] + completed = subprocess.run( + command, + cwd=source.parent, + capture_output=True, + text=True, + timeout=args.timeout, + check=False, + ) + result = Path(folder) / (source.stem + ".pdf") + messages = (completed.stdout + "\n" + completed.stderr).splitlines() if completed.returncode or not result.is_file(): - detail='\n'.join(messages[-15:]) - raise RuntimeError('LaTeX 编译失败;缺失 TeX 包需在基础镜像构建时预置,任务中不下载:'+detail) - count=len(PdfReader(result).pages) + detail = "\n".join(messages[-15:]) + raise RuntimeError( + "LaTeX 编译失败;缺失 TeX 包需在基础镜像构建时预置,任务中不下载:" + + detail + ) + count = len(PdfReader(result).pages) if not count: - raise ValueError('LaTeX 没有生成有效 PDF') - logfile=Path(folder)/(source.stem+'.log') + raise ValueError("LaTeX 没有生成有效 PDF") + logfile = Path(folder) / (source.stem + ".log") if logfile.exists(): - messages += logfile.read_text(errors='replace').splitlines() - warnings=list(dict.fromkeys(line.strip() for line in messages if re.search(r'warning:|Overfull|Missing character|undefined references',line,re.I)))[:30] - shutil.copyfile(result,temporary) - publish_temp_file(temporary,target,args.overwrite) + messages += logfile.read_text(errors="replace").splitlines() + warnings = list( + dict.fromkeys( + line.strip() + for line in messages + if re.search( + r"warning:|Overfull|Missing character|undefined references", + line, + re.I, + ) + ) + )[:30] + shutil.copyfile(result, temporary) + publish_temp_file(temporary, target, args.overwrite) finally: temporary.unlink(missing_ok=True) - return {'source':str(source),'path':str(target),'page_count':count,'engine':'tectonic','dependency_mode':'cached-only', - 'warnings':warnings,'requires_visual_review':True} + return { + "source": str(source), + "path": str(target), + "page_count": count, + "engine": "tectonic", + "dependency_mode": "cached-only", + "warnings": warnings, + "requires_visual_review": True, + } def main(argv=None): - parser=SkillArgumentParser(description='固定 LaTeX 编译接口,禁用 shell escape,只使用镜像内缓存包') - parser.add_argument('--input',required=True); parser.add_argument('--output',required=True) - parser.add_argument('--timeout',type=int,default=180); parser.add_argument('--overwrite',action='store_true') - return run_cli(lambda:compile_document(parser.parse_args(sys.argv[1:] if argv is None else argv))) + parser = SkillArgumentParser( + description="固定 LaTeX 编译接口,禁用 shell escape,只使用镜像内缓存包" + ) + parser.add_argument("--input", required=True) + parser.add_argument("--output", required=True) + parser.add_argument("--timeout", type=int, default=180) + parser.add_argument("--overwrite", action="store_true") + return run_cli( + lambda: compile_document( + parser.parse_args(sys.argv[1:] if argv is None else argv) + ) + ) -if __name__=='__main__': +if __name__ == "__main__": raise SystemExit(main()) diff --git a/skills/pdf/scripts/convert_to_pdf.py b/skills/pdf/scripts/convert_to_pdf.py index 4bbf223..e069c46 100644 --- a/skills/pdf/scripts/convert_to_pdf.py +++ b/skills/pdf/scripts/convert_to_pdf.py @@ -1,57 +1,117 @@ #!/usr/bin/env python3 """Office to PDF using the preinstalled LibreOffice, isolated per invocation.""" + from pathlib import Path import shutil import subprocess import sys import tempfile -from _pdf_common import SkillArgumentParser, output_pdf, new_temp_pdf, publish_temp_file, run_cli +from _pdf_common import ( + SkillArgumentParser, + output_pdf, + new_temp_pdf, + publish_temp_file, + run_cli, +) -FORMATS={'.docx','.doc','.odt','.rtf','.pptx','.ppt','.odp','.xlsx','.xls','.ods'} +FORMATS = { + ".docx", + ".doc", + ".odt", + ".rtf", + ".pptx", + ".ppt", + ".odp", + ".xlsx", + ".xls", + ".ods", +} def convert(args): from pypdf import PdfReader - source=Path(args.input).expanduser().resolve() - if not source.is_file() or source.suffix.lower() not in FORMATS or not 0 < source.stat().st_size <= 25*1024*1024: - raise ValueError('输入需为不超过 25 MiB 的本地 Office 文档;不支持把 PDF 直接反向转成可编辑 Office') + + source = Path(args.input).expanduser().resolve() + if ( + not source.is_file() + or source.suffix.lower() not in FORMATS + or not 0 < source.stat().st_size <= 25 * 1024 * 1024 + ): + raise ValueError( + "输入需为不超过 25 MiB 的本地 Office 文档;不支持把 PDF 直接反向转成可编辑 Office" + ) if not 1 <= args.timeout <= 600: - raise ValueError('timeout 必须在 1–600 秒之间') - executable=shutil.which('soffice') or shutil.which('libreoffice') + raise ValueError("timeout 必须在 1–600 秒之间") + executable = shutil.which("soffice") or shutil.which("libreoffice") if not executable: - raise RuntimeError('基础镜像缺少 LibreOffice,需要更新镜像') - target=output_pdf(args.output,args.overwrite) - temporary=new_temp_pdf(target) + raise RuntimeError("基础镜像缺少 LibreOffice,需要更新镜像") + target = output_pdf(args.output, args.overwrite) + temporary = new_temp_pdf(target) try: - with tempfile.TemporaryDirectory(prefix='pdf-office-') as folder: - root=Path(folder) - incoming=root/'input'; outgoing=root/'output'; profile=root/'profile' - incoming.mkdir(); outgoing.mkdir(); profile.mkdir() - local=incoming/('source'+source.suffix.lower()); shutil.copyfile(source,local) - command=[executable,'-env:UserInstallation='+profile.as_uri(),'--headless','--nologo','--nodefault', - '--nofirststartwizard','--convert-to','pdf','--outdir',str(outgoing),str(local)] - completed=subprocess.run(command,capture_output=True,text=True,timeout=args.timeout,check=False) - result=outgoing/'source.pdf' + with tempfile.TemporaryDirectory(prefix="pdf-office-") as folder: + root = Path(folder) + incoming = root / "input" + outgoing = root / "output" + profile = root / "profile" + incoming.mkdir() + outgoing.mkdir() + profile.mkdir() + local = incoming / ("source" + source.suffix.lower()) + shutil.copyfile(source, local) + command = [ + executable, + "-env:UserInstallation=" + profile.as_uri(), + "--headless", + "--nologo", + "--nodefault", + "--nofirststartwizard", + "--convert-to", + "pdf", + "--outdir", + str(outgoing), + str(local), + ] + completed = subprocess.run( + command, + capture_output=True, + text=True, + timeout=args.timeout, + check=False, + ) + result = outgoing / "source.pdf" if completed.returncode or not result.is_file(): - raise RuntimeError('Office 转 PDF 失败:'+(completed.stderr or completed.stdout)[-1500:]) - count=len(PdfReader(result).pages) + raise RuntimeError( + "Office 转 PDF 失败:" + + (completed.stderr or completed.stdout)[-1500:] + ) + count = len(PdfReader(result).pages) if not count: - raise ValueError('转换结果没有页面') - shutil.copyfile(result,temporary) - publish_temp_file(temporary,target,args.overwrite) + raise ValueError("转换结果没有页面") + shutil.copyfile(result, temporary) + publish_temp_file(temporary, target, args.overwrite) finally: temporary.unlink(missing_ok=True) - return {'source':str(source),'path':str(target),'page_count':count,'engine':'libreoffice','requires_visual_review':True, - 'note':'转换前需在对应 Office skill 中重算公式并核对字体、图表和打印范围。'} + return { + "source": str(source), + "path": str(target), + "page_count": count, + "engine": "libreoffice", + "requires_visual_review": True, + "note": "转换前需在对应 Office skill 中重算公式并核对字体、图表和打印范围。", + } def main(argv=None): - parser=SkillArgumentParser(description='Office 文档导出 PDF,源文件保持不变') - parser.add_argument('--input',required=True); parser.add_argument('--output',required=True) - parser.add_argument('--timeout',type=int,default=180); parser.add_argument('--overwrite',action='store_true') - return run_cli(lambda:convert(parser.parse_args(sys.argv[1:] if argv is None else argv))) + parser = SkillArgumentParser(description="Office 文档导出 PDF,源文件保持不变") + parser.add_argument("--input", required=True) + parser.add_argument("--output", required=True) + parser.add_argument("--timeout", type=int, default=180) + parser.add_argument("--overwrite", action="store_true") + return run_cli( + lambda: convert(parser.parse_args(sys.argv[1:] if argv is None else argv)) + ) -if __name__=='__main__': +if __name__ == "__main__": raise SystemExit(main()) diff --git a/skills/pdf/scripts/create_design_pdf.py b/skills/pdf/scripts/create_design_pdf.py index a621ee0..bf9a03a 100644 --- a/skills/pdf/scripts/create_design_pdf.py +++ b/skills/pdf/scripts/create_design_pdf.py @@ -16,21 +16,22 @@ MAX_SOURCE_BYTES = 2 * 1024 * 1024 class StaticHTML(HTMLParser): - def handle_starttag(self, tag, attributes): + def handle_starttag(self, tag: str, attrs: list[tuple[str, str | None]]) -> None: tag = tag.lower() - attrs = {key.lower(): value or '' for key, value in attributes} + attributes = {key.lower(): value or '' for key, value in attrs} if tag in {'script', 'iframe', 'object', 'embed', 'base', 'frame', 'frameset'}: raise ValueError(f'HTML 不允许 {tag};仅支持静态 HTML/CSS/SVG,公式和 Mermaid 由固定渲染器处理') - if any(key.startswith('on') for key in attrs) or 'srcdoc' in attrs: + if any(key.startswith('on') for key in attributes) or 'srcdoc' in attributes: raise ValueError('HTML 不允许事件处理程序或 srcdoc') - if tag == 'meta' and 'http-equiv' in attrs: + if tag == 'meta' and 'http-equiv' in attributes: raise ValueError('HTML 不允许 http-equiv;网络和文档策略由固定渲染器设置') for key in ('href', 'src', 'xlink:href', 'action', 'formaction'): - value = ''.join(attrs.get(key, '').split()).lower() + value = ''.join(attributes.get(key, '').split()).lower() if value.startswith(('javascript:', 'vbscript:', 'file:')): raise ValueError('HTML 不允许脚本 URL 或 file: 资源;使用任务目录内相对路径') - handle_startendtag = handle_starttag + def handle_startendtag(self, tag: str, attrs: list[tuple[str, str | None]]) -> None: + self.handle_starttag(tag, attrs) def read_source(value: str, suffixes: set[str]) -> Path: diff --git a/skills/pdf/scripts/edit_pdf.py b/skills/pdf/scripts/edit_pdf.py index 0adc9b1..a1355c8 100644 --- a/skills/pdf/scripts/edit_pdf.py +++ b/skills/pdf/scripts/edit_pdf.py @@ -126,7 +126,7 @@ def execute(args): if args.offset < 0 or not 1 <= args.limit <= 200: raise ValueError('offset 不能小于 0,limit 需为 1–200') stop = min(len(fields), args.offset + args.limit) - acroform = reader.trailer['/Root'].get('/AcroForm') + acroform = reader.root_object.get('/AcroForm') return {'field_count':len(fields), 'fields':fields[args.offset:stop], 'has_more':stop < len(fields), 'next_offset':stop if stop < len(fields) else None, 'has_xfa':bool(acroform and '/XFA' in acroform.get_object())} @@ -138,9 +138,10 @@ def execute(args): writer = PdfWriter(clone_from=reader) temporary = new_temp_pdf(output) extra = {} + values = {} try: if args.operation == 'form-fill': - acroform = reader.trailer['/Root'].get('/AcroForm') + acroform = reader.root_object.get('/AcroForm') if acroform and '/XFA' in acroform.get_object(): raise ValueError('XFA 表单不属于 AcroForm 固定接口,不能声称填写成功') values = validated_values(field_info(reader), load_data(args.data, args.data_file)) @@ -152,7 +153,7 @@ def execute(args): if set(data) - allowed or not all(isinstance(v,str) and len(v) <= 4096 for v in data.values()): raise ValueError('元数据仅支持 Title/Author/Subject/Keywords/Creator/Producer 文本字段') writer.add_metadata({'/'+key:value for key,value in data.items()}) - extra = {'updated_keys':list(data), 'xmp_preserved':'/Metadata' in reader.trailer['/Root']} + extra = {'updated_keys':list(data), 'xmp_preserved':'/Metadata' in reader.root_object} elif args.operation == 'crop': box = [float(v) for v in args.box.split(',')] if len(box) != 4 or not all(math.isfinite(v) for v in box) or box[0] >= box[2] or box[1] >= box[3]: @@ -162,7 +163,7 @@ def execute(args): media = writer.pages[number-1].mediabox if box[0] < media.left or box[1] < media.bottom or box[2] > media.right or box[3] > media.top: raise ValueError(f'裁剪框超出第 {number} 页 MediaBox') - writer.pages[number-1].cropbox = RectangleObject(box) + writer.pages[number-1].cropbox = RectangleObject((box[0], box[1], box[2], box[3])) extra = {'cropped_pages':pages, 'box':box, 'warning':'裁剪只改变可见范围,不删除隐藏内容,不能用于脱敏。'} with temporary.open('wb') as stream: writer.write(stream) diff --git a/skills/pdf/scripts/extract_text.py b/skills/pdf/scripts/extract_text.py index 7268f9c..96bc026 100644 --- a/skills/pdf/scripts/extract_text.py +++ b/skills/pdf/scripts/extract_text.py @@ -505,6 +505,8 @@ def _walk_content_images( elif operator == b"Do" and operands: try: xobject = _resolve(xobjects.get(operands[0])) + if xobject is None: + continue subtype = str(xobject.get("/Subtype")) except Exception: continue diff --git a/skills/pdf/scripts/manage_pdf.py b/skills/pdf/scripts/manage_pdf.py index 7474fab..abdc8df 100644 --- a/skills/pdf/scripts/manage_pdf.py +++ b/skills/pdf/scripts/manage_pdf.py @@ -5,7 +5,6 @@ from __future__ import annotations import re import sys from contextlib import ExitStack -from pathlib import Path from typing import Any from _pdf_common import ( diff --git a/skills/pdf/scripts/ocr_text.py b/skills/pdf/scripts/ocr_text.py index dcfa8e5..f02d293 100644 --- a/skills/pdf/scripts/ocr_text.py +++ b/skills/pdf/scripts/ocr_text.py @@ -193,6 +193,8 @@ def _create_ocr_engine(): except ImportError as exc: raise RuntimeError("环境预置的 rapidocr 模块不可用") from exc + if rapidocr.__file__ is None: + raise RuntimeError("无法确定 rapidocr 模块的安装路径") model_dir = Path(rapidocr.__file__).resolve().parent / "models" models = {"Det": "PP-OCRv6_det_small.onnx", "Cls": "ch_ppocr_mobile_v2.0_cls_mobile.onnx", "Rec": "PP-OCRv6_rec_small.onnx"} missing = [filename for filename in models.values() if not (model_dir / filename).is_file()] diff --git a/skills/pdf/tests/test_pdf_workflows.py b/skills/pdf/tests/test_pdf_workflows.py index 7937754..a128ff6 100644 --- a/skills/pdf/tests/test_pdf_workflows.py +++ b/skills/pdf/tests/test_pdf_workflows.py @@ -18,7 +18,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "scripts")) from PIL import Image from pypdf import PdfReader, PdfWriter -from pypdf.generic import DictionaryObject, NameObject, TextStringObject +from pypdf.generic import ArrayObject, DictionaryObject, NameObject, TextStringObject from reportlab.pdfgen import canvas import _pdf_common as common @@ -59,6 +59,16 @@ class PDFWorkflows(unittest.TestCase): output=str(self.root / "result.pdf"), overwrite=False, data=None, data_file=None, pages=None, **values) + def widgets(self, document: PdfReader | PdfWriter) -> list[DictionaryObject]: + annotations = document.pages[0]["/Annots"] + assert isinstance(annotations, ArrayObject) + widgets = [] + for reference in annotations: + widget = reference.get_object() + assert isinstance(widget, DictionaryObject) + widgets.append(widget) + return widgets + def test_form_inventory_and_roundtrip_all_supported_types(self): fields = edit_pdf.execute(self.args("form-info", offset=0, limit=50)) infos = {field["id"]: field for field in fields["fields"]} @@ -71,11 +81,12 @@ class PDFWorkflows(unittest.TestCase): result = edit_pdf.execute(args) updated = PdfReader(result["path"]) fields = updated.get_fields() + assert fields is not None self.assertEqual(fields["name"]["/V"], "Alice 123") self.assertEqual(fields["agree"]["/V"], "/Yes") self.assertEqual(fields["mode"]["/V"], "/B") self.assertEqual(fields["country"]["/V"], "US") - widgets = [ref.get_object() for ref in updated.pages[0]["/Annots"]] + widgets = self.widgets(updated) checkbox = next(widget for widget in widgets if widget.get("/T") == "agree") self.assertEqual(checkbox["/AS"], "/Yes") radio = [widget for widget in widgets if widget.get("/Parent")] @@ -86,7 +97,9 @@ class PDFWorkflows(unittest.TestCase): args = self.args("form-fill") args.data = '{"agree":false}' result = edit_pdf.execute(args) - self.assertEqual(PdfReader(result["path"]).get_fields()["agree"]["/V"], "/Off") + fields = PdfReader(result["path"]).get_fields() + assert fields is not None + self.assertEqual(fields["agree"]["/V"], "/Off") def test_invalid_form_input_does_not_publish(self): for data in [{"missing":"x"}, {"agree":"false"}, {"mode":True}, {"mode":"Off"}, @@ -100,8 +113,7 @@ class PDFWorkflows(unittest.TestCase): def test_form_inventory_handles_indirect_appearance(self): writer = PdfWriter(clone_from=self.source) - for ref in writer.pages[0]["/Annots"]: - widget = ref.get_object() + for widget in self.widgets(writer): if "/AP" in widget: widget[NameObject("/AP")] = writer._add_object(widget["/AP"]) alternative = self.root / "indirect.pdf" @@ -113,7 +125,9 @@ class PDFWorkflows(unittest.TestCase): def test_xfa_rejected_for_fill(self): writer = PdfWriter(clone_from=self.source) - writer.root_object["/AcroForm"][NameObject("/XFA")] = TextStringObject("unsupported") + acroform = writer.root_object["/AcroForm"] + assert isinstance(acroform, DictionaryObject) + acroform[NameObject("/XFA")] = TextStringObject("unsupported") alternative = self.root / "xfa.pdf" writer.write(alternative) args = self.args("form-fill") @@ -128,7 +142,7 @@ class PDFWorkflows(unittest.TestCase): reader = PdfReader(result["path"]) self.assertEqual(list(reader.pages[0].cropbox), [50, 50, 500, 700]) self.assertIn("Original content 123", reader.pages[0].extract_text()) - self.assertEqual(len(reader.get_fields()), 4) + self.assertEqual(len(reader.get_fields() or {}), 4) def test_crop_rejects_out_of_bounds_and_nonfinite_values(self): for box in ["-1,0,300,400", "0,0,601,800", "10,0,0,20", "0,0,nan,20"]: @@ -142,9 +156,11 @@ class PDFWorkflows(unittest.TestCase): args.data = '{"Title":"中文报告","Author":"Test"}' result = edit_pdf.execute(args) reader = PdfReader(result["path"]) - self.assertEqual(reader.metadata.title, "中文报告") - self.assertEqual(reader.metadata.producer, original.producer) - self.assertEqual(len(reader.get_fields()), 4) + metadata = reader.metadata + assert metadata is not None and original is not None + self.assertEqual(metadata.title, "中文报告") + self.assertEqual(metadata.producer, original.producer) + self.assertEqual(len(reader.get_fields() or {}), 4) def test_embedded_image_has_original_dimensions(self): args = self.args("extract-images", output_dir=str(self.root / "images"), start_image=0, max_images=20) @@ -189,7 +205,8 @@ class PDFWorkflows(unittest.TestCase): class OCRQuality(unittest.TestCase): def recognize(self, texts, scores): boxes = [[[0, i*30], [100, i*30], [100, i*30+20], [0, i*30+20]] for i in range(len(texts))] - engine = lambda _: SimpleNamespace(txts=texts, scores=scores, boxes=boxes) + def engine(_): + return SimpleNamespace(txts=texts, scores=scores, boxes=boxes) return ocr_text._ocr_page(engine, Path("fixture.png")) def test_mixed_confidence_does_not_certify_uncertain_amount(self): diff --git a/skills/pptx/scripts/inspect_presentation.py b/skills/pptx/scripts/inspect_presentation.py index 044b242..32f5630 100644 --- a/skills/pptx/scripts/inspect_presentation.py +++ b/skills/pptx/scripts/inspect_presentation.py @@ -314,6 +314,9 @@ def main() -> dict[str, Any]: + "、".join(archive["missing_required_parts"]) ) presentation = Presentation(str(source)) + slide_width, slide_height = presentation.slide_width, presentation.slide_height + if slide_width is None or slide_height is None: + raise ValueError("演示文稿缺少页面尺寸") slide_count = len(presentation.slides) if args.start_slide > slide_count and slide_count > 0: raise ValueError(f"start-slide 超出页面总数 {slide_count}") @@ -327,8 +330,8 @@ def main() -> dict[str, Any]: title = slide.shapes.title.text if slide.shapes.title is not None else "" media_profile = _slide_media_profile( slide, - int(presentation.slide_width), - int(presentation.slide_height), + slide_width, + slide_height, ) shapes: list[dict[str, Any]] = [] for shape in list(slide.shapes)[: args.max_shapes]: diff --git a/skills/pptx/scripts/ocr_presentation.py b/skills/pptx/scripts/ocr_presentation.py index d7bf74b..0caa295 100644 --- a/skills/pptx/scripts/ocr_presentation.py +++ b/skills/pptx/scripts/ocr_presentation.py @@ -588,8 +588,9 @@ def main() -> dict[str, Any]: if args.start_offset > 0 and len(requested_slides) != 1: raise ValueError("使用 start-offset 时 slides 必须只包含一页") - slide_width = int(presentation.slide_width) - slide_height = int(presentation.slide_height) + slide_width, slide_height = presentation.slide_width, presentation.slide_height + if slide_width is None or slide_height is None: + raise ValueError("演示文稿缺少页面尺寸") profiles = { slide_number: _slide_profile( presentation.slides[slide_number - 1], diff --git a/skills/pptx/scripts/validate_presentation.py b/skills/pptx/scripts/validate_presentation.py index f8c0c7f..8218eb0 100644 --- a/skills/pptx/scripts/validate_presentation.py +++ b/skills/pptx/scripts/validate_presentation.py @@ -268,8 +268,9 @@ def _visual_structure(path: Path) -> tuple[list[dict[str, Any]], list[dict[str, errors: list[dict[str, Any]] = [] warnings: list[dict[str, Any]] = [] presentation = Presentation(str(path)) - slide_width = int(presentation.slide_width) - slide_height = int(presentation.slide_height) + slide_width, slide_height = presentation.slide_width, presentation.slide_height + if slide_width is None or slide_height is None: + raise ValueError("演示文稿缺少页面尺寸") tolerance = 2000 for slide_number, slide in enumerate(presentation.slides, start=1): if len(slide.shapes) == 0: diff --git a/skills/send-complex-message/SKILL.md b/skills/send-complex-message/SKILL.md index 2b88c5d..47f562f 100644 --- a/skills/send-complex-message/SKILL.md +++ b/skills/send-complex-message/SKILL.md @@ -1,6 +1,6 @@ --- name: send-complex-message -description: "在当前微信会话中发送纯文本、引用回复或群聊艾特/@/提及消息。用户要求发送一段文字、仅@成员或@所有人、引用某条消息回复、引用时同时@成员时使用;文本和引用支持私聊及群聊,可在发送后结束当前 Agent 对话。" +description: "在当前微信会话中发送纯文本、引用回复或群聊艾特/@/提及消息,也可按时间、关键词、消息类型和发送人查询当前群聊或私聊最近24小时内的聊天记录。用户要求发送、艾特、引用回复或查找近期历史消息时使用,可在发送后结束当前 Agent 对话。" --- # Send Complex Message Skill @@ -9,7 +9,7 @@ description: "在当前微信会话中发送纯文本、引用回复或群聊艾 本技能是当前微信会话中纯文本、艾特和引用回复的统一发送入口。仅艾特时不需要正文;引用消息必须有正文,也可以附带成员艾特参数。艾特支持指定一个或多个成员,也支持微信原生的 `@所有人`。 -技能脚本位于 `scripts/send_complex_message.py`,统一调用客户端的 `/message/send/refermessage` 接口。`refer_message_id` 是可选参数:不传时由客户端调用普通文本消息方法,传入时发送引用消息。指定成员时,根据昵称或备注查询当前群内未退群成员;@所有人时使用客户端协议值 `notify@all`,不要把 `@昵称` 或 `@所有人` 当普通正文拼接。 +技能脚本位于 `scripts/send_complex_message.py`。加上 `--query-history` 时只查询当前会话历史消息;发送模式统一调用客户端的 `/message/send/refermessage` 接口。`refer_message_id` 是可选参数:不传时由客户端调用普通文本消息方法,传入时发送引用消息。指定成员时,根据昵称或备注查询当前群内未退群成员;@所有人时使用客户端协议值 `notify@all`,不要把 `@昵称` 或 `@所有人` 当普通正文拼接。 本技能不带引用 ID 时通过普通文本消息实现原生艾特,支持只艾特而不附加正文。带引用 ID 时,客户端接收 `at` 并显示艾特名称,但尚未实现引用消息中的原生艾特提醒,不能把引用发送成功表述为已经提醒成员。 @@ -22,11 +22,57 @@ description: "在当前微信会话中发送纯文本、引用回复或群聊艾 - 用户要求「帮我艾特下 xxx」「@ 一下 xxx」「提一下 xxx 和 yyy」。 - 用户要求「@所有人」「提醒全体成员」「通知群里所有人」。 - 需要在群聊里点名提醒某人。 +- 用户要求查询、搜索当前群聊或私聊的近期聊天记录,或需要查找历史消息的 `messages.id` 以便引用。 - 其它时候不应该使用本技能 不带艾特的纯文本和引用回复可用于私聊和群聊;只要提供艾特参数,`ROBOT_FROM_WX_ID` 就必须是群聊 ID。 -## 入参规范 +## 历史聊天记录查询 + +使用 `--query-history` 进入只读查询模式。会话固定取系统注入的 `ROBOT_FROM_WX_ID`:群聊只能查询该群,私聊只能查询与当前好友的私聊,包含该会话双方的收发消息。查询同时限定 `from_wxid` 和 `is_chat_room`,不能按发送人跨群或跨私聊搜索,也不能修改会话环境变量来切换查询对象。 + +所有时间范围必须落在**执行查询时的最近 24 小时内**。24 小时是最大回溯范围,不是固定查询时长;可以查询最近半小时、2 小时,或这 24 小时内任意更短的起止区间。超过上限、落在过去更早日期、包含未来或起止倒置的时间范围会报错,不会静默扩大或改写用户指定的范围。 + +| 查询参数 | 说明 | +| --- | --- | +| `--query-history` | 必须提供,不能与发送参数或 `--ended` 同用 | +| `--hours <小时数>` | 最近多少小时,支持小数,`0 < hours <= 24`,精度为整秒且至少 1 秒;如 `0.5` 表示最近 30 分钟 | +| `--start-time <时间>` | 起点,支持 Unix 秒或北京时间 `YYYY-MM-DD HH:mm[:ss]` | +| `--end-time <时间>` | 终点,格式同上;与起点均为包含边界 | +| `--keyword <关键词>` | 对 `content` 和 `display_full_content` 做包含匹配;可重复,多个关键词必须全部命中,`%`、`_`、反斜杠按字面量匹配 | +| `--message-type <类型>` | 消息类型编号或名称,可重复,多个类型匹配任一即可;省略时查询所有类型 | +| `--app-msg-type <子类型编号>` | 可重复,如 `57` 引用、`6` 文件、`5` 链接;自动限定消息类型为 `49`,不能与其他消息类型组合 | +| `--sender-wxid <微信ID>` | 只在当前会话内按发送人微信 ID 精确过滤 | +| `--limit <条数>` | 单页条数,默认 `50`,范围 `1..200` | +| `--offset <偏移量>` | 分页偏移量,默认 `0`,必须非负 | + +`--hours` 与 `--start-time`/`--end-time` 互斥。未指定时间参数时默认查询最近 24 小时;只传起点时终点为现在,只传终点时起点为现在减 24 小时。 + +类型名称支持 `text=1`、`image=3`、`voice=34`、`card=42`、`video=43`、`emoji=47`、`location=48`、`app=49`、`system=10000`、`recall=10002`,其他类型可直接传数字编号。不同类别的过滤条件同时生效。 + +查询最近 30 分钟包含“安排”的文本消息: + +```bash +python3 scripts/send_complex_message.py --query-history --hours 0.5 --keyword '安排' --message-type text +``` + +查询最近 6 小时某人发送的引用消息: + +```bash +python3 scripts/send_complex_message.py --query-history --hours 6 --sender-wxid 'wxid_zhangsan' --app-msg-type 57 --limit 20 +``` + +按明确起止时间查询(先根据用户指定的区间设置两个 Unix 秒时间戳,且必须在最近 24 小时内): + +```bash +python3 scripts/send_complex_message.py --query-history --start-time "$START_TIMESTAMP" --end-time "$END_TIMESTAMP" --message-type image --message-type video +``` + +查询成功时输出 JSON 对象,包含当前会话、实际 `start_time`/`end_time`(Unix 秒)、`count`、`has_more`、`next_offset` 和 `messages`。消息按 `created_at DESC, id DESC` 排列,每条包含数据库主键 `id`、会话及发送人微信 ID、消息类型/子类型、正文、显示内容、撤回标记和发送时间(Unix 秒)。`messages: []` 表示没有匹配记录。`has_more: true` 时可以保留过滤条件,使用返回的 `next_offset` 作为 `--offset` 继续查询,不能把单页结果表述为全部记录。 + +查询不发送微信消息、不输出 `ended`,也不要求配置客户端端口。需要引用查询结果时,先根据消息内容、发送人和时间确定目标,再单独以返回的 `id` 调用发送模式。不能绕过脚本直接查询其他会话或更早记录。 + +## 发送入参规范 `refer_message_id` 不全局必填,只在用户明确要求引用时传入,值为 `messages.id`。仅文本、仅艾特时省略该参数,不查询引用目标,也不从环境变量自动补齐引用 ID。 @@ -37,7 +83,7 @@ description: "在当前微信会话中发送纯文本、引用回复或群聊艾 | 仅引用 | 不传 | 传入 `messages.id` | 必填且不能是纯空白 | | 引用并艾特 | 指定成员或 `all: true` | 传入 `messages.id` | 必填且不能是纯空白 | -不引用时也支持正文加艾特。下方 schema 没有全局必填字段,组合校验由 schema 和脚本共同约束。 +不引用时也支持正文加艾特。下方 schema 仅描述发送模式,没有全局必填字段,组合校验由 schema 和脚本共同约束;查询模式使用上方查询参数。 ```json { @@ -119,7 +165,7 @@ description: "在当前微信会话中发送纯文本、引用回复或群聊艾 - 仅在用户要求引用时选择消息,传给脚本的引用 ID 只使用 `messages.id`。 - 引用当前用户触发本次对话的消息时,使用环境变量 `ROBOT_MESSAGE_ID` 中的消息主键。 - 用户要求回复他引用的原消息时,使用 `ROBOT_REF_MESSAGE_ID` 中的消息主键;该值为空或 `0` 表示没有引用目标,不能改为引用当前消息。 -- 引用其他历史消息时,从当前机器人数据库 `messages` 表查找,限定 `from_wxid = ROBOT_FROM_WX_ID`,按用户描述确认原消息后取 `id`。不要使用 `msg_id`、`client_msg_id` 或 XML 中的 `svrid`。 +- 引用其他历史消息时,使用本脚本 `--query-history`,在当前会话最近 24 小时内按时间、关键词、类型或发送人查找,按用户描述确认原消息后取 `id`。不要使用 `msg_id`、`client_msg_id` 或 XML 中的 `svrid`。 - 引用目标不明确时先确认,不能猜测消息 ID。不要因为上下文存在引用消息就自动发送引用回复。 - 脚本接收明确的 `--refer-message-id`,不会自动选择最近一条消息;仅引用且不指定成员时不需要查询成员表。 @@ -135,7 +181,7 @@ description: "在当前微信会话中发送纯文本、引用回复或群聊艾 6. 如果没有完全相等结果,选择第一个 `remark` 包含输入值的成员。 7. 如果仍未命中,选择第一个 `nickname` 包含输入值的成员。 -## 执行步骤 +## 发送执行步骤 1. 判断用户需要纯文本、仅艾特、引用回复,还是引用时同时艾特。纯文本和引用回复必须准备非空正文 `content`;引用时按上面的规则确定 `refer_message_id`。 2. 如需指定成员,把用户原话中的昵称或备注写入 `mention`/`mentions`;@所有人时设置 `all: true` 并使用 `--all`。 @@ -206,7 +252,7 @@ python3 scripts/send_complex_message.py --mention '张三' --content '看一下 ## 依赖安装 -- 指定成员、需要查询数据库时,脚本会自动创建虚拟环境并安装依赖;纯文本、仅引用或 `--all` 不需要安装数据库依赖。 +- 查询历史记录或指定成员、需要查询数据库时,脚本会自动创建虚拟环境并安装依赖,使用 `MYSQL_HOST`、`MYSQL_PORT`、`MYSQL_USER`、`MYSQL_PASSWORD` 连接 `ROBOT_CODE` 对应的机器人数据库;纯文本、仅引用或 `--all` 不需要安装数据库依赖。 - 如需手动重新安装,可执行:`python3 scripts/bootstrap.py` ## ended 行为 diff --git a/skills/send-complex-message/scripts/bootstrap.py b/skills/send-complex-message/scripts/bootstrap.py index 45005f7..0e96053 100644 --- a/skills/send-complex-message/scripts/bootstrap.py +++ b/skills/send-complex-message/scripts/bootstrap.py @@ -108,4 +108,4 @@ if __name__ == "__main__": raise except Exception: traceback.print_exc(file=sys.stdout) - raise SystemExit(1) \ No newline at end of file + raise SystemExit(1) diff --git a/skills/send-complex-message/scripts/send_complex_message.py b/skills/send-complex-message/scripts/send_complex_message.py index 7bb602a..b8be29a 100644 --- a/skills/send-complex-message/scripts/send_complex_message.py +++ b/skills/send-complex-message/scripts/send_complex_message.py @@ -4,15 +4,41 @@ from __future__ import annotations import argparse import json +import math import os import subprocess import sys +import time import traceback import urllib.request +from datetime import datetime, timedelta, timezone from pathlib import Path +from typing import NoReturn sys.stderr = sys.stdout +MAX_HISTORY_SECONDS = 24 * 60 * 60 +DEFAULT_HISTORY_LIMIT = 50 +MAX_HISTORY_LIMIT = 200 +SHANGHAI_TZ = timezone(timedelta(hours=8)) +MESSAGE_TYPES = { + "text": 1, + "image": 3, + "voice": 34, + "card": 42, + "video": 43, + "emoji": 47, + "location": 48, + "app": 49, + "system": 10000, + "recall": 10002, +} + + +class SkillArgumentParser(argparse.ArgumentParser): + def error(self, message: str) -> NoReturn: + raise ValueError(message) + def _client_private_token() -> str: return os.environ.get("ROBOT_CLIENT_PRIVATE_TOKEN", "").strip() @@ -66,7 +92,7 @@ def _ensure_skill_venv_python() -> None: def _mysql_connect(): _ensure_skill_venv_python() try: - import pymysql # type: ignore + import pymysql except ModuleNotFoundError: _run_bootstrap() venv_python = _skill_venv_python() @@ -131,18 +157,54 @@ def _expand_json_array_values(values: list[str], label: str) -> list[str]: return expanded -def _parse_cli_params(argv: list[str]) -> tuple[list[str], str, bool, bool, int | None]: - parser = argparse.ArgumentParser(add_help=False) +def _positive_int(value: str) -> int: + number = int(value) + if not 0 < number <= 2**63 - 1: + raise ValueError("必须是 int64 范围内的正整数") + return number + + +def _message_type(value: str) -> int: + normalized = value.strip().lower() + if normalized in MESSAGE_TYPES: + return MESSAGE_TYPES[normalized] + return _positive_int(normalized) + + +def _parse_cli_params(argv: list[str]) -> argparse.Namespace: + parser = SkillArgumentParser(description="发送消息或查询当前会话最近 24 小时内的聊天记录", allow_abbrev=False) parser.add_argument("--mention", action="append", default=[]) parser.add_argument("--mentions", action="append", default=[]) parser.add_argument("--all", "--mention-all", dest="mention_all", action="store_true") parser.add_argument("--refer-message-id", type=int) - parser.add_argument("--content", default="") + parser.add_argument("--content") parser.add_argument("--ended", action="store_true", default=False) + parser.add_argument("--query-history", action="store_true", help="只查询历史聊天记录,不发送消息") + parser.add_argument("--hours", type=float, help="查询最近多少小时,支持小数,最多 24 小时") + parser.add_argument("--start-time", help="开始时间:Unix 秒或北京时间 YYYY-MM-DD HH:mm[:ss]") + parser.add_argument("--end-time", help="结束时间:Unix 秒或北京时间 YYYY-MM-DD HH:mm[:ss]") + parser.add_argument("--keyword", action="append", help="正文或显示内容包含的关键词,可重复,多个词须全部匹配") + parser.add_argument("--message-type", action="append", type=_message_type, help="消息类型名称或编号,可重复") + parser.add_argument("--app-msg-type", action="append", type=_positive_int, help="APP 消息子类型编号,可重复") + parser.add_argument("--sender-wxid", help="按发送人微信 ID 精确过滤") + parser.add_argument("--limit", type=int, help="单页条数,默认 50,最多 200") + parser.add_argument("--offset", type=int, help="分页偏移量,默认 0") - namespace, unknown = parser.parse_known_args(argv) - if unknown: - raise ValueError(f"存在不支持的参数: {' '.join(unknown)}") + namespace = parser.parse_args(argv) + + if namespace.query_history: + if (namespace.mention or namespace.mentions or namespace.mention_all + or namespace.refer_message_id is not None or namespace.content is not None + or namespace.ended): + raise ValueError("--query-history 不能与发送参数或 --ended 同时使用") + _validate_history_params(namespace) + return namespace + + history_fields = ("hours", "start_time", "end_time", "keyword", "message_type", + "app_msg_type", "sender_wxid", "limit", "offset") + if any(getattr(namespace, field) is not None for field in history_fields): + raise ValueError("聊天记录过滤参数必须与 --query-history 一起使用") + namespace.content = namespace.content or "" mentions = _expand_json_array_values(namespace.mention + namespace.mentions, "mentions") deduped: list[str] = [] @@ -165,7 +227,131 @@ def _parse_cli_params(argv: list[str]) -> tuple[list[str], str, bool, bool, int if not deduped and not namespace.mention_all and not namespace.content.strip(): raise ValueError("请提供非空 content,或指定要艾特的成员/--all") - return deduped, namespace.content, namespace.ended, namespace.mention_all, namespace.refer_message_id + namespace.mentions = deduped + return namespace + + +def _parse_history_time(value: str, field_name: str) -> int: + text = value.strip() + if text.isascii() and text.isdigit(): + return int(text) + for pattern in ("%Y-%m-%d %H:%M:%S", "%Y-%m-%d %H:%M"): + try: + return int(datetime.strptime(text, pattern).replace(tzinfo=SHANGHAI_TZ).timestamp()) + except ValueError: + continue + raise ValueError(f"{field_name} 必须是 Unix 秒或北京时间 YYYY-MM-DD HH:mm[:ss]") + + +def _resolve_history_time_range(args: argparse.Namespace) -> tuple[int, int]: + now = int(time.time()) + earliest = now - MAX_HISTORY_SECONDS + if args.hours is not None: + if args.start_time is not None or args.end_time is not None: + raise ValueError("--hours 不能与 --start-time/--end-time 同时使用") + if not math.isfinite(args.hours) or not 0 < args.hours <= 24: + raise ValueError("hours 必须大于 0 且不超过 24") + seconds = int(args.hours * 3600) + if seconds < 1: + raise ValueError("hours 对应的时间范围不能小于 1 秒") + return now - seconds, now + + start = _parse_history_time(args.start_time, "start_time") if args.start_time is not None else earliest + end = _parse_history_time(args.end_time, "end_time") if args.end_time is not None else now + if start < earliest or end > now: + raise ValueError("只能查询最近 24 小时内的聊天记录,不能查询更早或未来的时间") + if start >= end: + raise ValueError("结束时间必须晚于开始时间") + return start, end + + +def _validate_history_params(args: argparse.Namespace) -> None: + _resolve_history_time_range(args) + args.limit = DEFAULT_HISTORY_LIMIT if args.limit is None else args.limit + args.offset = 0 if args.offset is None else args.offset + if not 1 <= args.limit <= MAX_HISTORY_LIMIT: + raise ValueError(f"limit 必须在 1 到 {MAX_HISTORY_LIMIT} 之间") + if not 0 <= args.offset <= 2**63 - 1: + raise ValueError("offset 必须是 int64 范围内的非负整数") + args.keyword = [keyword.strip() for keyword in (args.keyword or [])] + if any(not keyword for keyword in args.keyword): + raise ValueError("keyword 不能为空或纯空白") + if args.sender_wxid is not None: + args.sender_wxid = args.sender_wxid.strip() + if not args.sender_wxid: + raise ValueError("sender_wxid 不能为空或纯空白") + if args.app_msg_type and args.message_type and set(args.message_type) != {49}: + raise ValueError("app_msg_type 只能与 APP 消息类型 49(app)一起使用") + + +def _query_history(conn, conversation_id: str, args: argparse.Namespace) -> dict: + # 建立连接/安装依赖可能耗时;在真正查询前重新确定最近 24 小时的边界。 + start, end = _resolve_history_time_range(args) + is_chat_room = conversation_id.endswith("@chatroom") + conditions = [ + "from_wxid = %s", + "is_chat_room = %s", + "created_at >= %s", + "created_at <= %s", + ] + params: list[object] = [conversation_id, is_chat_room, start, end] + for keyword in args.keyword: + conditions.append("(content LIKE %s ESCAPE '\\\\' OR display_full_content LIKE %s ESCAPE '\\\\')") + pattern = f"%{_escape_like(keyword)}%" + params.extend((pattern, pattern)) + if args.sender_wxid: + conditions.append("sender_wxid = %s") + params.append(args.sender_wxid) + if args.message_type: + conditions.append(f"`type` IN ({', '.join(['%s'] * len(args.message_type))})") + params.extend(args.message_type) + if args.app_msg_type: + conditions.append("`type` = 49") + conditions.append(f"app_msg_type IN ({', '.join(['%s'] * len(args.app_msg_type))})") + params.extend(args.app_msg_type) + sql = f""" + SELECT id, from_wxid, sender_wxid, to_wxid, is_chat_room, + type, app_msg_type, content, display_full_content, is_recalled, created_at + FROM messages + WHERE {' AND '.join(conditions)} + ORDER BY created_at DESC, id DESC + LIMIT %s OFFSET %s + """ + params.extend((args.limit + 1, args.offset)) + with conn.cursor() as cursor: + cursor.execute(sql, tuple(params)) + rows = list(cursor.fetchall()) + has_more = len(rows) > args.limit + messages = rows[:args.limit] + return { + "conversation_id": conversation_id, + "is_chat_room": is_chat_room, + "start_time": start, + "end_time": end, + "limit": args.limit, + "offset": args.offset, + "count": len(messages), + "has_more": has_more, + "next_offset": args.offset + len(messages) if has_more else None, + "messages": messages, + } + + +def _run_history_query(conversation_id: str, args: argparse.Namespace) -> int: + try: + conn = _mysql_connect() + except Exception as exc: + sys.stdout.write(f"数据库连接失败: {exc}\n") + return 1 + try: + result = _query_history(conn, conversation_id, args) + sys.stdout.write(json.dumps(result, ensure_ascii=False) + "\n") + return 0 + except Exception as exc: + sys.stdout.write(f"查询聊天记录失败: {exc}\n") + return 1 + finally: + conn.close() def _escape_like(value: str) -> str: @@ -257,7 +443,7 @@ def _send_message( def main() -> int: try: - mentions, content, ended, mention_all, refer_message_id = _parse_cli_params(sys.argv[1:]) + args = _parse_cli_params(sys.argv[1:]) except (ValueError, json.JSONDecodeError) as exc: sys.stdout.write(f"参数格式错误: {exc}\n") return 1 @@ -266,6 +452,11 @@ def main() -> int: if not to_wxid: sys.stdout.write("环境变量 ROBOT_FROM_WX_ID 未配置\n") return 1 + if args.query_history: + return _run_history_query(to_wxid, args) + + mentions, content = args.mentions, args.content + ended, mention_all, refer_message_id = args.ended, args.mention_all, args.refer_message_id if (mention_all or mentions) and not to_wxid.endswith("@chatroom"): sys.stdout.write("当前会话不是群聊,不能发送艾特消息\n") return 1 diff --git a/skills/send-mention-message/scripts/bootstrap.py b/skills/send-mention-message/scripts/bootstrap.py index 45005f7..0e96053 100644 --- a/skills/send-mention-message/scripts/bootstrap.py +++ b/skills/send-mention-message/scripts/bootstrap.py @@ -108,4 +108,4 @@ if __name__ == "__main__": raise except Exception: traceback.print_exc(file=sys.stdout) - raise SystemExit(1) \ No newline at end of file + raise SystemExit(1) diff --git a/skills/send-mention-message/scripts/send_mention_message.py b/skills/send-mention-message/scripts/send_mention_message.py index c66305d..aa0aa61 100644 --- a/skills/send-mention-message/scripts/send_mention_message.py +++ b/skills/send-mention-message/scripts/send_mention_message.py @@ -66,7 +66,7 @@ def _ensure_skill_venv_python() -> None: def _mysql_connect(): _ensure_skill_venv_python() try: - import pymysql # type: ignore + import pymysql except ModuleNotFoundError: _run_bootstrap() venv_python = _skill_venv_python() diff --git a/skills/text-to-image/scripts/bootstrap.py b/skills/text-to-image/scripts/bootstrap.py index 0d2cb77..4ebdb30 100644 --- a/skills/text-to-image/scripts/bootstrap.py +++ b/skills/text-to-image/scripts/bootstrap.py @@ -130,4 +130,4 @@ if __name__ == "__main__": raise except Exception: traceback.print_exc(file=sys.stdout) - raise SystemExit(1) \ No newline at end of file + raise SystemExit(1) diff --git a/skills/text-to-image/scripts/text_to_image.py b/skills/text-to-image/scripts/text_to_image.py index 51c8ca8..820ba05 100644 --- a/skills/text-to-image/scripts/text_to_image.py +++ b/skills/text-to-image/scripts/text_to_image.py @@ -70,8 +70,8 @@ def _ensure_skill_venv_python() -> None: _ensure_skill_venv_python() try: - import pymysql # type: ignore # noqa: E402 - from openai import OpenAI # type: ignore # noqa: E402 + import pymysql + from openai import OpenAI except ModuleNotFoundError: _run_bootstrap() _py = _get_python_executable() diff --git a/skills/video-generation/scripts/bootstrap.py b/skills/video-generation/scripts/bootstrap.py index 39d4579..8959bb4 100644 --- a/skills/video-generation/scripts/bootstrap.py +++ b/skills/video-generation/scripts/bootstrap.py @@ -131,4 +131,4 @@ if __name__ == "__main__": raise except Exception: traceback.print_exc(file=sys.stdout) - raise SystemExit(1) \ No newline at end of file + raise SystemExit(1) diff --git a/skills/video-generation/scripts/video_generation.py b/skills/video-generation/scripts/video_generation.py index 7b11644..edc5865 100644 --- a/skills/video-generation/scripts/video_generation.py +++ b/skills/video-generation/scripts/video_generation.py @@ -83,7 +83,7 @@ def _ensure_skill_venv_python() -> None: _ensure_skill_venv_python() try: - import pymysql # type: ignore # noqa: E402 + import pymysql except ModuleNotFoundError: _run_bootstrap() _py = _get_python_executable() diff --git a/skills/voice-message/scripts/bootstrap.py b/skills/voice-message/scripts/bootstrap.py index caecf37..0a45d98 100644 --- a/skills/voice-message/scripts/bootstrap.py +++ b/skills/voice-message/scripts/bootstrap.py @@ -112,4 +112,4 @@ if __name__ == "__main__": raise except Exception: traceback.print_exc(file=sys.stdout) - raise SystemExit(1) \ No newline at end of file + raise SystemExit(1) diff --git a/skills/voice-message/scripts/voice_message.py b/skills/voice-message/scripts/voice_message.py index 7cf0e80..582a1d7 100644 --- a/skills/voice-message/scripts/voice_message.py +++ b/skills/voice-message/scripts/voice_message.py @@ -115,7 +115,7 @@ def _ensure_skill_venv_python() -> None: _ensure_skill_venv_python() try: - import pymysql # type: ignore # noqa: E402 + import pymysql except ModuleNotFoundError: _run_bootstrap() _py = _get_python_executable() @@ -694,7 +694,7 @@ def _decompress_response_bytes(raw: bytes, encoding: str) -> bytes: return zlib.decompress(raw, -zlib.MAX_WBITS) if encoding == "br": try: - import brotli # type: ignore + import brotli except ModuleNotFoundError as exc: raise RuntimeError( "mimo 响应使用了 brotli 压缩,但当前环境未安装 brotli,请安装后重试" diff --git a/skills/web-page/scripts/tsconfig.json b/skills/web-page/scripts/tsconfig.json index 22d9c10..19c652a 100644 --- a/skills/web-page/scripts/tsconfig.json +++ b/skills/web-page/scripts/tsconfig.json @@ -7,6 +7,8 @@ "lib": ["ES2022", "DOM"], "types": ["node"], "strict": true, + "noUnusedLocals": true, + "noUnusedParameters": true, "noEmit": true, "esModuleInterop": true, "forceConsistentCasingInFileNames": true, diff --git a/skills/web-page/scripts/web_page.test.ts b/skills/web-page/scripts/web_page.test.ts index c027787..e532000 100644 --- a/skills/web-page/scripts/web_page.test.ts +++ b/skills/web-page/scripts/web_page.test.ts @@ -6,11 +6,8 @@ import path from "node:path"; import test from "node:test"; import { fileURLToPath } from "node:url"; -import { isLocalFileUrl, validateUrl } from "./web_page.ts"; - const SCRIPT_DIR = path.dirname(fileURLToPath(import.meta.url)); const SCRIPT_PATH = path.join(SCRIPT_DIR, "web_page.ts"); -const PASSWD_MARKER = "root:x:0:0"; interface ScriptResult { code: number | null; @@ -65,45 +62,32 @@ function runWebPage(url: string, args: string[] = []): Promise { }); } -function assertLocalFileBlocked(result: ScriptResult): void { +function assertSearchResult(result: ScriptResult, url: string, query: string): void { const output = `${result.stdout}\n${result.stderr}`; - assert.notEqual(result.code, 0, output); - assert.match( - output, - /已阻止浏览器|网页链接必须是 http 或 https 地址|网页导航失败/, - ); - assert.doesNotMatch(output, new RegExp(PASSWD_MARKER)); + assert.equal(result.code, 0, output); + assert.ok(result.stdout.includes(`URL:${url}`), output); + assert.ok(result.stdout.includes(`SEARCH_OK:${query}`), output); } -test("web-page 本地文件访问防护", async (t) => { +test("web-page 网页读取与自动化交互", async (t) => { const server = http.createServer((request, response) => { const requestUrl = new URL(request.url || "/", "http://127.0.0.1"); response.setHeader("Content-Type", "text/html; charset=utf-8"); switch (requestUrl.pathname) { - case "/redirect-file": - response.statusCode = 302; - response.setHeader("Location", "file:///etc/passwd"); - response.end(); - return; - case "/click-file": + case "/click-link": response.end( - 'click file打开本地文件', + 'link查看结果', ); return; case "/js-location": response.end( - 'js location

safe

', + 'js location', ); return; - case "/iframe-file": + case "/form": response.end( - 'iframe file

safe

', - ); - return; - case "/popup-file": - response.end( - 'popup file', + 'search form
', ); return; case "/redirect-http": @@ -113,7 +97,7 @@ test("web-page 本地文件访问防护", async (t) => { return; case "/search": response.end( - `search ok
SEARCH_OK:${requestUrl.searchParams.get("q") || ""}
`, + `search ok
SEARCH_OK:${requestUrl.searchParams.get("q") || ""}
`, ); return; default: @@ -124,58 +108,96 @@ test("web-page 本地文件访问防护", async (t) => { server.listen(0, "127.0.0.1"); await once(server, "listening"); - t.after(() => server.close()); + t.after(() => new Promise((resolve, reject) => { + server.close((error) => error ? reject(error) : resolve()); + })); const address = server.address(); assert.ok(address && typeof address === "object"); const baseUrl = `http://127.0.0.1:${address.port}`; - await t.test("直接访问 file:///etc/passwd", async () => { - assertLocalFileBlocked(await runWebPage("file:///etc/passwd")); + await t.test("读取 HTTP 网页正文", async () => { + const url = `${baseUrl}/search?q=direct-http`; + const result = await runWebPage(url); + assertSearchResult(result, url, "direct-http"); + assert.match(result.stdout, /标题:search ok/); + assert.doesNotMatch(result.stdout, /console output is not page content/); }); - await t.test("HTTP 302 跳转到 file://", async () => { - assertLocalFileBlocked(await runWebPage(`${baseUrl}/redirect-file`)); + await t.test("跟随 HTTP 302 跳转并返回目标网页", async () => { + assertSearchResult( + await runWebPage(`${baseUrl}/redirect-http`), + `${baseUrl}/search?q=normal-http-redirect`, + "normal-http-redirect", + ); }); - await t.test("点击 file:// 链接", async () => { - assertLocalFileBlocked( - await runWebPage(`${baseUrl}/click-file`, [ + await t.test("点击链接并等待目标网页加载", async () => { + assertSearchResult( + await runWebPage(`${baseUrl}/click-link`, [ "--actions", - JSON.stringify([{ type: "click", selector: "#local-file" }]), + JSON.stringify([{ + type: "click", + selector: "#search-link", + wait_for_navigation: true, + }]), ]), + `${baseUrl}/search?q=clicked-link`, + "clicked-link", ); }); - await t.test("JavaScript 修改 location", async () => { - assertLocalFileBlocked(await runWebPage(`${baseUrl}/js-location`)); - }); - - await t.test("iframe 加载本地文件", async () => { - assertLocalFileBlocked(await runWebPage(`${baseUrl}/iframe-file`)); - }); - - await t.test("弹窗加载本地文件", async () => { - assertLocalFileBlocked( - await runWebPage(`${baseUrl}/popup-file`, [ + await t.test("等待 JavaScript 跳转后的页面元素", async () => { + assertSearchResult( + await runWebPage(`${baseUrl}/js-location`, [ "--actions", - JSON.stringify([{ type: "click", selector: "#open-popup" }]), + JSON.stringify([{ type: "wait_for_selector", selector: "#search-result" }]), ]), + `${baseUrl}/search?q=js-navigation`, + "js-navigation", ); }); - await t.test("正常 HTTP/HTTPS 搜索及浏览不受影响", async () => { - assert.equal( - validateUrl("https://example.com/search?q=normal-https"), - "https://example.com/search?q=normal-https", + await t.test("填写并提交搜索表单", async () => { + assertSearchResult( + await runWebPage(`${baseUrl}/form`, [ + "--actions", + JSON.stringify([ + { type: "fill", selector: "#query", value: "form-search" }, + { type: "click", selector: "#submit", wait_for_navigation: true }, + ]), + ]), + `${baseUrl}/search?q=form-search`, + "form-search", ); - assert.equal(isLocalFileUrl("https://example.com/file.txt"), false); - assert.equal(isLocalFileUrl("file:///etc/passwd"), true); - assert.equal(isLocalFileUrl("filesystem:https://example.com/temporary/a"), true); + }); - const normal = await runWebPage(`${baseUrl}/redirect-http`); - assert.equal(normal.code, 0, `${normal.stdout}\n${normal.stderr}`); - assert.match(normal.stdout, /SEARCH_OK:normal-http-redirect/); - assert.doesNotMatch(normal.stdout, new RegExp(PASSWD_MARKER)); + await t.test("动作失败时返回动作序号和原因", async () => { + const result = await runWebPage(`${baseUrl}/form`, [ + "--actions", + JSON.stringify([ + { type: "fill", selector: "#query", value: "unused" }, + { type: "click", selector: "#missing-button" }, + ]), + ]); + const output = `${result.stdout}\n${result.stderr}`; + assert.equal(result.code, 1, output); + assert.match(output, /第 2 个 action\(click\) 执行失败/); + assert.match(output, /未找到可操作元素: #missing-button/); }); }); + +test("web-page 命令行拒绝非 HTTP/HTTPS 协议", async (t) => { + for (const url of [ + "file:///tmp/web-page-test.html", + "filesystem:https://example.com/temporary/a", + "ftp://example.com/file.txt", + ]) { + await t.test(`拒绝 ${new URL(url).protocol} 地址`, async () => { + const result = await runWebPage(url); + const output = `${result.stdout}\n${result.stderr}`; + assert.equal(result.code, 1, output); + assert.match(output, /网页链接必须是 http 或 https 地址/); + }); + } +}); diff --git a/skills/xlsx/scripts/_xlsx_common.py b/skills/xlsx/scripts/_xlsx_common.py index 269d704..68bc478 100644 --- a/skills/xlsx/scripts/_xlsx_common.py +++ b/skills/xlsx/scripts/_xlsx_common.py @@ -202,6 +202,16 @@ def validate_cell_range(value: str, *, label: str = "区域") -> str: return normalized +def cell_range_bounds(value: str) -> tuple[int, int, int, int]: + from openpyxl.utils.cell import range_boundaries + + bounds = range_boundaries(validate_cell_range(value)) + min_col, min_row, max_col, max_row = bounds + if min_col is None or min_row is None or max_col is None or max_row is None: + raise ValueError(f"区域必须包含完整的行列边界:{value}") + return min_col, min_row, max_col, max_row + + def find_program(*names: str) -> str: for name in names: resolved = shutil.which(name) diff --git a/skills/xlsx/scripts/_xlsx_data.py b/skills/xlsx/scripts/_xlsx_data.py index 61ca709..eb09ccd 100644 --- a/skills/xlsx/scripts/_xlsx_data.py +++ b/skills/xlsx/scripts/_xlsx_data.py @@ -14,8 +14,8 @@ import numpy as np import pandas as pd from _xlsx_common import ( - EXCEL_INPUT_SUFFIXES, input_file, output_file, publish_file, - normalize_formula_error, validate_cell_range, workbook_has_external_links, + EXCEL_INPUT_SUFFIXES, cell_range_bounds, input_file, output_file, publish_file, + workbook_has_external_links, ) MAX_DATA_CELLS = 500_000 @@ -49,8 +49,9 @@ def scalar(value: Any) -> Any: return None if not math.isfinite(value): raise ValueError("结果含无穷值") - if hasattr(value, "isoformat"): - return value.isoformat() + isoformat = getattr(value, "isoformat", None) + if callable(isoformat): + return isoformat() if isinstance(value, (str, int, float, bool)): return value return str(value) @@ -58,7 +59,7 @@ def scalar(value: Any) -> Any: def read_dataset(path_value: str, spec: dict[str, Any] | None = None) -> tuple[pd.DataFrame, dict]: from openpyxl import load_workbook - from openpyxl.utils.cell import column_index_from_string, get_column_letter, range_boundaries + from openpyxl.utils.cell import column_index_from_string, get_column_letter spec = spec or {} allowed = {"sheet", "range", "header_row", "columns", "exclude_rows", "numeric", "dates", "encoding", "path"} @@ -72,7 +73,7 @@ def read_dataset(path_value: str, spec: dict[str, Any] | None = None) -> tuple[p header_row = int(spec.get("header_row", 1)) if header_row < 1: raise ValueError("header_row 从 1 开始") - bounds = range_boundaries(validate_cell_range(spec["range"])) if spec.get("range") else None + bounds = cell_range_bounds(spec["range"]) if spec.get("range") else None if bounds and not bounds[1] <= header_row <= bounds[3]: raise ValueError("header_row 必须位于 range 内") records = [] @@ -82,7 +83,8 @@ def read_dataset(path_value: str, spec: dict[str, Any] | None = None) -> tuple[p if source.suffix.lower() in EXCEL_INPUT_SUFFIXES: formula_wb = load_workbook(source, read_only=True, data_only=False, keep_links=False) cached_wb = load_workbook(source, read_only=True, data_only=True, keep_links=False) - sheet_name = spec.get("sheet") or formula_wb.active.title + active = formula_wb.active + sheet_name = spec.get("sheet") or (active.title if active is not None else None) if sheet_name not in formula_wb.sheetnames: raise ValueError(f"工作表不存在:{sheet_name}") ws, cached = formula_wb[sheet_name], cached_wb[sheet_name] @@ -230,10 +232,13 @@ def save_plan(tables: list[tuple[str, pd.DataFrame]], metadata: dict, destinatio category = chart["category"] values = chart["values"] require_columns(table, [category, *values]) - indexes = [table.columns.get_loc(value) + 1 for value in values] + if not table.columns.is_unique: + raise ValueError("图表来源包含重复字段名") + column_names = list(table.columns) + indexes = [column_names.index(value) + 1 for value in values] if indexes != list(range(min(indexes), max(indexes) + 1)): raise ValueError("图表 values 需按顺序选择相邻的结果列") - c = get_column_letter(table.columns.get_loc(category) + 1) + c = get_column_letter(column_names.index(category) + 1) end = len(table) + 1 if end < 2: raise ValueError("没有数据可用于图表") diff --git a/skills/xlsx/scripts/analyze_workbook.py b/skills/xlsx/scripts/analyze_workbook.py index a03cddd..92e37ca 100644 --- a/skills/xlsx/scripts/analyze_workbook.py +++ b/skills/xlsx/scripts/analyze_workbook.py @@ -5,7 +5,6 @@ from __future__ import annotations import re from typing import Any -import numpy as np import pandas as pd from _xlsx_common import SkillArgumentParser, load_json_argument, run_cli @@ -72,6 +71,8 @@ def transform(frame: pd.DataFrame, operations: list[dict], audit: list) -> pd.Da if predicate in {"eq", "ne", "gt", "ge", "lt", "le"}: mask = getattr(series, predicate)(value) elif predicate in {"in", "not_in"}: + if not isinstance(value, list): + raise ValueError("in/not_in 的 value 必须是数组") mask = series.isin(value) if predicate == "not_in": mask = ~mask @@ -148,8 +149,11 @@ def aggregate(frame: pd.DataFrame, spec: dict, *, pivot: bool) -> pd.DataFrame: raise ValueError("行维度与列维度不能重复") result = frame.groupby(by + column_fields, dropna=False, sort=False, observed=True).agg(reducers) if column_fields: - result = result.unstack(column_fields) - result.columns = [json_label(parts) for parts in result.columns.to_flat_index()] + unstacked = result.unstack(column_fields) + if not isinstance(unstacked, pd.DataFrame) or not isinstance(unstacked.columns, pd.MultiIndex): + raise ValueError("透视聚合未生成预期的多级字段表") + unstacked.columns = pd.Index([json_label(parts) for parts in unstacked.columns.to_flat_index()]) + result = unstacked return result.reset_index() diff --git a/skills/xlsx/scripts/apply_workbook.py b/skills/xlsx/scripts/apply_workbook.py index 9cf1cd9..f6f718b 100644 --- a/skills/xlsx/scripts/apply_workbook.py +++ b/skills/xlsx/scripts/apply_workbook.py @@ -14,6 +14,7 @@ from _xlsx_common import ( EXCEL_INPUT_SUFFIXES, EXCEL_OUTPUT_SUFFIXES, SkillArgumentParser, + cell_range_bounds, input_file, load_json_argument, normalize_formula_error, @@ -503,12 +504,12 @@ def _op_set_row_heights(workbook: Any, op: dict[str, Any]) -> int: def _op_auto_fit(workbook: Any, op: dict[str, Any]) -> int: import math import unicodedata - from openpyxl.utils.cell import get_column_letter, range_boundaries + from openpyxl.utils.cell import get_column_letter worksheet = _sheet(workbook, op.get("sheet")) reference = validate_cell_range(str(op.get("range", ""))) cells = list(_iter_range_cells(worksheet, reference)) - min_col, min_row, max_col, max_row = range_boundaries(reference) + min_col, min_row, max_col, max_row = cell_range_bounds(reference) minimum, maximum = float(op.get("min_width", 8)), float(op.get("max_width", 40)) if not 1 <= minimum <= maximum <= 100: raise ValueError("auto_fit 列宽需满足 1 <= min_width <= max_width <= 100") @@ -568,7 +569,6 @@ def _op_add_table(workbook: Any, op: dict[str, Any]) -> int: def _op_add_chart(workbook: Any, op: dict[str, Any]) -> int: from openpyxl.chart import AreaChart, BarChart, LineChart, PieChart, Reference - from openpyxl.utils.cell import range_boundaries worksheet = _sheet(workbook, op.get("sheet")) chart_type = str(op.get("chart_type", "bar")).lower() @@ -582,7 +582,7 @@ def _op_add_chart(workbook: Any, op: dict[str, Any]) -> int: if chart_type not in chart_classes: raise ValueError("chart_type 仅支持 area、bar、column、line、pie") data_range = validate_cell_range(str(op.get("data_range", ""))) - min_col, min_row, max_col, max_row = range_boundaries(data_range) + min_col, min_row, max_col, max_row = cell_range_bounds(data_range) chart = chart_classes[chart_type]() if isinstance(chart, BarChart): chart.type = "bar" if chart_type == "bar" else "col" @@ -600,7 +600,7 @@ def _op_add_chart(workbook: Any, op: dict[str, Any]) -> int: ) if op.get("categories_range"): category_range = validate_cell_range(str(op["categories_range"])) - c_min_col, c_min_row, c_max_col, c_max_row = range_boundaries( + c_min_col, c_min_row, c_max_col, c_max_row = cell_range_bounds( category_range ) categories = Reference( @@ -619,10 +619,9 @@ def _op_add_chart(workbook: Any, op: dict[str, Any]) -> int: chart.y_axis.title = str(op["y_axis_title"]) if "style" in op: chart.style = int(op["style"]) - if "height" in op: - chart.height = float(op["height"]) - if "width" in op: - chart.width = float(op["width"]) + for dimension in ("height", "width"): + if dimension in op: + setattr(chart, dimension, float(op[dimension])) if "legend_position" in op and chart.legend: chart.legend.position = str(op["legend_position"]) anchor = validate_cell_reference(str(op.get("anchor", "E2"))) @@ -639,10 +638,9 @@ def _op_add_image(workbook: Any, op: dict[str, Any]) -> int: {".png", ".jpg", ".jpeg", ".gif", ".bmp"}, ) image = Image(str(image_path)) - if "width" in op: - image.width = float(op["width"]) - if "height" in op: - image.height = float(op["height"]) + for dimension in ("width", "height"): + if dimension in op: + setattr(image, dimension, float(op[dimension])) anchor = validate_cell_reference(str(op.get("anchor", "A1"))) worksheet.add_image(image, anchor) return 0 @@ -654,7 +652,7 @@ def _op_add_data_validation(workbook: Any, op: dict[str, Any]) -> int: worksheet = _sheet(workbook, op.get("sheet")) reference = validate_cell_range(str(op.get("range", ""))) validation_type = str(op.get("validation_type", "list")) - allowed = { + allowed = ( "list", "whole", "decimal", @@ -662,7 +660,7 @@ def _op_add_data_validation(workbook: Any, op: dict[str, Any]) -> int: "time", "textLength", "custom", - } + ) if validation_type not in allowed: raise ValueError(f"validation_type 不支持:{validation_type}") validation = DataValidation( @@ -1051,7 +1049,7 @@ def main() -> dict[str, Any]: "path": str(destination), "source": str(source) if source else None, "sheet_names": workbook.sheetnames, - "active_sheet": workbook.active.title, + "active_sheet": workbook.active.title if workbook.active is not None else None, "operation_count": len(operations), "processed_cell_count": written_cells, "formula_count": scan["formula_count"], diff --git a/skills/xlsx/scripts/convert_workbook.py b/skills/xlsx/scripts/convert_workbook.py index 57c3e74..5ca5636 100644 --- a/skills/xlsx/scripts/convert_workbook.py +++ b/skills/xlsx/scripts/convert_workbook.py @@ -8,7 +8,7 @@ import os import tempfile from datetime import date, datetime from pathlib import Path -from typing import Any, Iterable, Optional +from typing import Any, Optional from _xlsx_common import ( TABULAR_INPUT_SUFFIXES, @@ -81,6 +81,8 @@ def _delimited_to_xlsx( ) workbook = Workbook() worksheet = workbook.active + if worksheet is None: + raise ValueError("工作簿没有活动工作表") worksheet.title = sheet_name[:31] or "Sheet1" for row in rows: worksheet.append(row) @@ -133,6 +135,8 @@ def _xlsx_to_delimited( worksheet = workbook[sheet_name] else: worksheet = workbook.active + if worksheet is None: + raise ValueError("工作簿没有活动工作表") row_count = 0 with destination.open("w", encoding=encoding, newline="") as handle: writer = csv.writer(handle, delimiter=delimiter) diff --git a/skills/xlsx/scripts/inspect_workbook.py b/skills/xlsx/scripts/inspect_workbook.py index ac76cef..8fea2ef 100644 --- a/skills/xlsx/scripts/inspect_workbook.py +++ b/skills/xlsx/scripts/inspect_workbook.py @@ -106,6 +106,7 @@ def inspect_excel( max_columns: int, ) -> dict[str, Any]: from openpyxl import load_workbook + from openpyxl.worksheet.worksheet import Worksheet options = openpyxl_load_options(source) formulas = load_workbook(source, data_only=False, **options) @@ -115,6 +116,8 @@ def inspect_excel( total_formulas = 0 total_errors = 0 for worksheet in formulas.worksheets: + if not isinstance(worksheet, Worksheet): + raise ValueError("工作表不支持完整检查,请以普通模式加载工作簿") formula_count = 0 error_count = 0 for cell in worksheet._cells.values(): @@ -136,8 +139,8 @@ def inspect_excel( "auto_filter": worksheet.auto_filter.ref, "merged_ranges": [str(item) for item in worksheet.merged_cells.ranges], "tables": list(worksheet.tables.keys()), - "chart_count": len(worksheet._charts), - "image_count": len(worksheet._images), + "chart_count": len(getattr(worksheet, "_charts")), + "image_count": len(getattr(worksheet, "_images")), "formula_count": formula_count, "literal_error_count": error_count, "print_area": str(worksheet.print_area) if worksheet.print_area else None, @@ -151,7 +154,10 @@ def inspect_excel( ) selected_name = sheet_name else: - selected_name = formulas.active.title + active = formulas.active + if active is None: + raise ValueError("工作簿没有活动工作表") + selected_name = active.title formula_sheet = formulas[selected_name] cached_sheet = cached[selected_name] @@ -184,7 +190,7 @@ def inspect_excel( "format": source.suffix.lower(), "macro_enabled": source.suffix.lower() in {".xlsm", ".xltm"}, "has_external_links": workbook_has_external_links(source), - "active_sheet": formulas.active.title, + "active_sheet": formulas.active.title if formulas.active is not None else None, "sheet_names": formulas.sheetnames, "sheets": summaries, "defined_names": _defined_names(formulas), diff --git a/skills/xlsx/scripts/model_workbook.py b/skills/xlsx/scripts/model_workbook.py index 79c617d..5dd91a9 100644 --- a/skills/xlsx/scripts/model_workbook.py +++ b/skills/xlsx/scripts/model_workbook.py @@ -3,13 +3,12 @@ from __future__ import annotations import math -from typing import Any import numpy as np import pandas as pd from _xlsx_common import SkillArgumentParser, load_json_argument, run_cli -from _xlsx_data import SOURCE_ROW, numeric, read_dataset, require_columns, save_plan, scalar +from _xlsx_data import SOURCE_ROW, numeric, read_dataset, require_columns, save_plan def evaluate(frame: pd.DataFrame, spec: dict) -> tuple[list, dict]: @@ -61,8 +60,8 @@ def evaluate(frame: pd.DataFrame, spec: dict) -> tuple[list, dict]: if not np.allclose(np.diag(matrix), 1) or not np.allclose(matrix * matrix.T, 1, atol=1e-6): raise ValueError("AHP 比较矩阵必须对角为 1 且互反") values, vectors = np.linalg.eig(matrix) - index = np.argmax(values.real) - weights = np.abs(vectors[:, index].real) + index = np.argmax(np.real(values)) + weights = np.abs(np.real(vectors[:, index])) ri = [0, 0, 0, .58, .90, 1.12, 1.24, 1.32, 1.41, 1.45][n] consistency = max(0, float(values[index].real - n) / (n - 1) / ri) if ri else 0 if consistency >= .1: @@ -229,7 +228,7 @@ def supervised(frame: pd.DataFrame, spec: dict, *, classification: bool) -> tupl metrics += [[label, algorithm, key, float(value)] for key, value in stats.items()] predictions = frame.iloc[test][[SOURCE_ROW, *features]].copy() predictions["实际值"] = y.iloc[test].to_numpy() - predictions["预测值"] = pipeline.predict(X.iloc[test]) + predictions["预测值"] = np.asarray(pipeline.predict(X.iloc[test])) model = pipeline.named_steps["model"] names = pipeline.named_steps["prepare"].get_feature_names_out() importance = model.feature_importances_ if hasattr(model, "feature_importances_") else np.mean(np.abs(np.atleast_2d(model.coef_)), axis=0) @@ -246,7 +245,7 @@ def supervised(frame: pd.DataFrame, spec: dict, *, classification: bool) -> tupl future_X[col] = future_X[col].map(lambda x: str(x) if pd.notna(x) else np.nan) pipeline.fit(X, y) future = future[[SOURCE_ROW, *features]].copy() - future["预测值"] = pipeline.predict(future_X) + future["预测值"] = np.asarray(pipeline.predict(future_X)) tables.append(("新样本预测", future)) else: provenance = None @@ -282,7 +281,7 @@ def unsupervised(frame: pd.DataFrame, spec: dict, *, anomaly: bool) -> tuple[lis labels = model.fit_predict(scaled) result = frame.copy() result["异常标记" if anomaly else "簇编号"] = labels - if anomaly: + if isinstance(model, IsolationForest): result["正常程度得分"] = model.decision_function(scaled) score = None valid = labels != -1 @@ -339,12 +338,14 @@ def optimize(spec: dict) -> tuple[list, dict]: start = np.asarray(spec.get("initial", np.clip(np.zeros(n), lower, upper)), dtype=float) if start.shape != (n,) or not np.isfinite(start).all(): raise ValueError("initial 需为有限数值向量") - objective = lambda x: float(c @ x + .5 * x @ Q @ x) + def objective(x): + return float(c @ x + .5 * x @ Q @ x) result = minimize(lambda x: sign * objective(x), start, jac=lambda x: sign * (c + Q @ x), method="SLSQP", bounds=Bounds(lower, upper), constraints=[linear] if linear else [], options={"maxiter": 1000, "ftol": 1e-9}) guarantee = "凸二次规划的数值解,已检查可行性" else: - objective = lambda x: float(c @ x) + def objective(x): + return float(c @ x) result = milp(sign * c, integrality=np.asarray(integer, dtype=int), bounds=Bounds(lower, upper), constraints=linear, options={"time_limit": 60., "mip_rel_gap": 0.}) guarantee = "HiGHS 求解成功;仅在成功且可行时输出方案" diff --git a/skills/xlsx/tests/test_data_workflows.py b/skills/xlsx/tests/test_data_workflows.py index ca3eb4a..23db77c 100644 --- a/skills/xlsx/tests/test_data_workflows.py +++ b/skills/xlsx/tests/test_data_workflows.py @@ -1,6 +1,5 @@ import csv import hashlib -import json import sys import unittest import tempfile @@ -12,8 +11,7 @@ import numpy as np import pandas as pd from openpyxl import Workbook, load_workbook -ROOT = Path(__file__).resolve().parents[1] -sys.path.insert(0, str(ROOT / 'scripts')) +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "scripts")) import _xlsx_common as common import _xlsx_data as data import analyze_workbook as analysis @@ -21,10 +19,13 @@ import apply_workbook as writer import inspect_workbook as inspector import model_workbook as modeling -_TEMP = tempfile.TemporaryDirectory(prefix='xlsx-tests-') +ROOT = Path(__file__).resolve().parents[1] + +_TEMP = tempfile.TemporaryDirectory(prefix="xlsx-tests-") QA = Path(_TEMP.name).resolve() _ORIGINAL_OUTPUT_ROOT = common.EXCEL_OUTPUT_ROOT + def tearDownModule(): common.EXCEL_OUTPUT_ROOT = _ORIGINAL_OUTPUT_ROOT _TEMP.cleanup() @@ -34,185 +35,394 @@ class MergeTests(unittest.TestCase): @classmethod def setUpClass(cls): common.EXCEL_OUTPUT_ROOT = QA - cls.source = QA / 'sales.xlsx' + cls.source = QA / "sales.xlsx" wb = Workbook() ws = wb.active - ws.title = '明细' - for row in [['地区', '收入', '成本', '说明'], ['华东', 100, 60, '满意'], ['华东', -20, 5, '退款'], ['华南', 80, 50, '物流慢'], ['华南', None, 20, '服务好但物流慢'], ['合计', 160, 135, None]]: + assert ws is not None + ws.title = "明细" + for row in [ + ["地区", "收入", "成本", "说明"], + ["华东", 100, 60, "满意"], + ["华东", -20, 5, "退款"], + ["华南", 80, 50, "物流慢"], + ["华南", None, 20, "服务好但物流慢"], + ["合计", 160, 135, None], + ]: ws.append(row) - font = copy(ws['A1'].font); font.bold = True; ws['A1'].font = font + font = copy(ws["A1"].font) + font.bold = True + ws["A1"].font = font wb.save(cls.source) cls.source_hash = hashlib.sha256(cls.source.read_bytes()).hexdigest() - cls.csv = QA / 'text.csv' - with cls.csv.open('w', encoding='utf-8-sig', newline='') as f: - csv.writer(f).writerows([['ID', '文本'], ['001', '=1+1'], ['002', '长文本' * 80]]) + cls.csv = QA / "text.csv" + with cls.csv.open("w", encoding="utf-8-sig", newline="") as f: + csv.writer(f).writerows( + [["ID", "文本"], ["001", "=1+1"], ["002", "长文本" * 80]] + ) def dataset(self): - return data.read_dataset(str(self.source), {'sheet': '明细', 'exclude_rows': [6]})[0] + return data.read_dataset( + str(self.source), {"sheet": "明细", "exclude_rows": [6]} + )[0] def test_profile_finds_totals_without_dropping_negative_rows(self): - _, result = analysis.analyze(str(self.source), {'method': 'profile'}) - self.assertEqual(result['source']['rows_used'], 5) - self.assertEqual(result['summary_row_candidates'][0]['row'], 6) + _, result = analysis.analyze(str(self.source), {"method": "profile"}) + self.assertEqual(result["source"]["rows_used"], 5) + self.assertEqual(result["summary_row_candidates"][0]["row"], 6) def test_aggregate_negative_values_and_null_count(self): - tables, _ = analysis.analyze(str(self.source), {'method': 'aggregate', 'source': {'exclude_rows': [6]}, 'by': ['地区'], 'metrics': {'收入': 'sum', '说明': 'count'}}) - result = tables[0][1].set_index('地区') - self.assertEqual(result.loc['华东', '收入'], 80) - self.assertEqual(result.loc['华南', '说明'], 2) + tables, _ = analysis.analyze( + str(self.source), + { + "method": "aggregate", + "source": {"exclude_rows": [6]}, + "by": ["地区"], + "metrics": {"收入": "sum", "说明": "count"}, + }, + ) + result = tables[0][1].set_index("地区") + self.assertEqual(result.loc["华东", "收入"], 80) + self.assertEqual(result.loc["华南", "说明"], 2) def test_all_null_sum_not_zero(self): - frame = pd.DataFrame({'类别': ['A','B','B'], '值': [None, 0, None]}) - result = analysis.aggregate(frame, {'by': ['类别'], 'metrics': {'值': 'sum'}}, pivot=False).set_index('类别') - self.assertTrue(pd.isna(result.loc['A','值'])) - self.assertEqual(result.loc['B','值'], 0) + frame = pd.DataFrame({"类别": ["A", "B", "B"], "值": [None, 0, None]}) + result = analysis.aggregate( + frame, {"by": ["类别"], "metrics": {"值": "sum"}}, pivot=False + ).set_index("类别") + self.assertTrue(pd.isna(result.loc["A", "值"])) + self.assertEqual(result.loc["B", "值"], 0) def test_pivot_multiple_levels(self): - frame = pd.DataFrame({'地区':['A','A','B'], '年':['2025','2026','2025'], '收入':[10,20,30]}) - result = analysis.aggregate(frame, {'by':['地区'],'columns':['年'],'metrics':{'收入':'sum'}}, pivot=True).set_index('地区') - self.assertEqual(result.loc['A','["收入","2026"]'],20) - self.assertTrue(pd.isna(result.loc['B','["收入","2026"]'])) + frame = pd.DataFrame( + { + "地区": ["A", "A", "B"], + "年": ["2025", "2026", "2025"], + "收入": [10, 20, 30], + } + ) + result = analysis.aggregate( + frame, + {"by": ["地区"], "columns": ["年"], "metrics": {"收入": "sum"}}, + pivot=True, + ).set_index("地区") + self.assertEqual(result.loc["A", '["收入","2026"]'], 20) + self.assertTrue(pd.isna(result.loc["B", '["收入","2026"]'])) def test_rules_report_conflict_and_unknown(self): - detail, summary = analysis.classify(self.dataset(), {'column':'说明','rules':[{'label':'正向','keywords':['好','满意']},{'label':'物流','keywords':['慢']}]}) - self.assertEqual(detail['分类'].tolist(), ['正向','未知','物流','需复核']) - self.assertAlmostEqual(summary['占比'].sum(),1) + detail, summary = analysis.classify( + self.dataset(), + { + "column": "说明", + "rules": [ + {"label": "正向", "keywords": ["好", "满意"]}, + {"label": "物流", "keywords": ["慢"]}, + ], + }, + ) + self.assertEqual(detail["分类"].tolist(), ["正向", "未知", "物流", "需复核"]) + self.assertAlmostEqual(summary["占比"].sum(), 1) def test_join_rejects_accidental_many_to_many(self): with self.assertRaises(pd.errors.MergeError): - analysis.transform(self.dataset(), [{'type':'merge','source':{'path':str(self.source),'exclude_rows':[6]},'on':['地区']}], []) + analysis.transform( + self.dataset(), + [ + { + "type": "merge", + "source": {"path": str(self.source), "exclude_rows": [6]}, + "on": ["地区"], + } + ], + [], + ) def test_missing_headers_can_use_coordinates(self): - wb=Workbook(); ws=wb.active - ws.append([None,'值']); ws.append(['001',5]) - path=QA/'no_header.xlsx'; wb.save(path) - with self.assertRaises(ValueError): data.read_dataset(str(path)) - frame, meta = data.read_dataset(str(path), {'columns':{'编号':'A','值':'B'}}) - self.assertEqual(frame.iloc[0]['编号'],'001') - self.assertEqual(meta['columns']['编号'],'A') + wb = Workbook() + ws = wb.active + assert ws is not None + ws.append([None, "值"]) + ws.append(["001", 5]) + path = QA / "no_header.xlsx" + wb.save(path) + with self.assertRaises(ValueError): + data.read_dataset(str(path)) + frame, meta = data.read_dataset( + str(path), {"columns": {"编号": "A", "值": "B"}} + ) + self.assertEqual(frame.iloc[0]["编号"], "001") + self.assertEqual(meta["columns"]["编号"], "A") def test_formula_cache_is_required(self): - wb=Workbook(); ws=wb.active; ws.append(['值']); ws.append(['=1+1']) - path=QA/'uncached.xlsx'; wb.save(path) - with self.assertRaisesRegex(ValueError,'公式缓存'): + wb = Workbook() + ws = wb.active + assert ws is not None + ws.append(["值"]) + ws.append(["=1+1"]) + path = QA / "uncached.xlsx" + wb.save(path) + with self.assertRaisesRegex(ValueError, "公式缓存"): data.read_dataset(str(path)) def test_safe_text_and_no_truncation_in_writer(self): - tables, metadata=analysis.analyze(str(self.csv), {'method':'transform'}) - plan=QA/'safe-text.json'; output=QA/'safe-text.xlsx' - data.save_plan(tables,metadata,str(plan),overwrite=True) - with patch.object(sys,'argv',['apply','--output',str(output),'--spec-file',str(plan),'--overwrite']): - result=writer.main() - self.assertEqual(result['formula_count'],0) - wb=load_workbook(output); ws=wb['分析结果'] - self.assertEqual(ws['B2'].value,'001') - self.assertEqual(ws['C2'].value,'=1+1') - self.assertEqual(ws['C2'].data_type,'s') - self.assertEqual(ws['C3'].value,'长文本'*80) - self.assertTrue(ws['C3'].alignment.wrap_text) - self.assertGreater(ws.row_dimensions[3].height,36) - self.assertEqual(inspector.inspect_excel(output,sheet_name='分析结果',start_row=1,start_column=1,max_rows=10,max_columns=10)['formula_count'],0) + tables, metadata = analysis.analyze(str(self.csv), {"method": "transform"}) + plan = QA / "safe-text.json" + output = QA / "safe-text.xlsx" + data.save_plan(tables, metadata, str(plan), overwrite=True) + with patch.object( + sys, + "argv", + ["apply", "--output", str(output), "--spec-file", str(plan), "--overwrite"], + ): + result = writer.main() + self.assertEqual(result["formula_count"], 0) + wb = load_workbook(output) + ws = wb["分析结果"] + self.assertEqual(ws["B2"].value, "001") + self.assertEqual(ws["C2"].value, "=1+1") + self.assertEqual(ws["C2"].data_type, "s") + self.assertEqual(ws["C3"].value, "长文本" * 80) + self.assertTrue(ws["C3"].alignment.wrap_text) + self.assertGreater(ws.row_dimensions[3].height, 36) + self.assertEqual( + inspector.inspect_excel( + output, + sheet_name="分析结果", + start_row=1, + start_column=1, + max_rows=10, + max_columns=10, + )["formula_count"], + 0, + ) def test_append_result_preserves_source(self): - spec={'method':'aggregate','source':{'exclude_rows':[6]},'by':['地区'],'metrics':{'收入':'sum'},'chart':{'category':'地区','values':['收入'],'title':'地区收入'}} - tables,meta=analysis.analyze(str(self.source),spec) - plan=QA/'summary.json'; output=QA/'summary.xlsx' - data.save_plan(tables,meta,str(plan),overwrite=True,chart=spec['chart']) - with patch.object(sys,'argv',['apply','--input',str(self.source),'--output',str(output),'--spec-file',str(plan),'--overwrite']): writer.main() - self.assertEqual(hashlib.sha256(self.source.read_bytes()).hexdigest(),self.source_hash) - original=load_workbook(self.source); result=load_workbook(output) - self.assertEqual(list(original['明细'].values),list(result['明细'].values)) - self.assertEqual(copy(original['明细']['A1'].font),copy(result['明细']['A1'].font)) - self.assertEqual(len(result['分析结果']._charts),1) + spec = { + "method": "aggregate", + "source": {"exclude_rows": [6]}, + "by": ["地区"], + "metrics": {"收入": "sum"}, + "chart": {"category": "地区", "values": ["收入"], "title": "地区收入"}, + } + tables, meta = analysis.analyze(str(self.source), spec) + plan = QA / "summary.json" + output = QA / "summary.xlsx" + data.save_plan(tables, meta, str(plan), overwrite=True, chart=spec["chart"]) + with patch.object( + sys, + "argv", + [ + "apply", + "--input", + str(self.source), + "--output", + str(output), + "--spec-file", + str(plan), + "--overwrite", + ], + ): + writer.main() + self.assertEqual( + hashlib.sha256(self.source.read_bytes()).hexdigest(), self.source_hash + ) + original = load_workbook(self.source) + result = load_workbook(output) + self.assertEqual(list(original["明细"].values), list(result["明细"].values)) + self.assertEqual( + copy(original["明细"]["A1"].font), copy(result["明细"]["A1"].font) + ) + self.assertEqual(len(getattr(result["分析结果"], "_charts")), 1) def test_writer_rejects_text_that_excel_would_truncate(self): - for value in ['长' * 32768, {'value': '长' * 32768}]: + for value in ["长" * 32768, {"value": "长" * 32768}]: with self.subTest(explicit=isinstance(value, dict)): wb = Workbook() - with self.assertRaisesRegex(ValueError, '不能静默截断'): - writer._op_write_rows(wb, {'sheet': wb.active.title, 'rows': [[value]]}) + assert wb.active is not None + with self.assertRaisesRegex(ValueError, "不能静默截断"): + writer._op_write_rows( + wb, {"sheet": wb.active.title, "rows": [[value]]} + ) def test_cost_direction_once(self): - frame=pd.DataFrame({data.SOURCE_ROW:[2,3],'对象':['便宜','贵'],'质量':[98,98],'价格':[10,30]}) - tables, _ = modeling.evaluate(frame, {'entity':'对象','directions':{'质量':'benefit','价格':'cost'}}) - rank=tables[0][1].set_index('对象') - self.assertEqual(rank.loc['便宜','排名'],1) - self.assertEqual(rank.loc['贵','排名'],2) + frame = pd.DataFrame( + { + data.SOURCE_ROW: [2, 3], + "对象": ["便宜", "贵"], + "质量": [98, 98], + "价格": [10, 30], + } + ) + tables, _ = modeling.evaluate( + frame, {"entity": "对象", "directions": {"质量": "benefit", "价格": "cost"}} + ) + rank = tables[0][1].set_index("对象") + self.assertEqual(rank.loc["便宜", "排名"], 1) + self.assertEqual(rank.loc["贵", "排名"], 2) def test_identical_objects_tie(self): - frame=pd.DataFrame({data.SOURCE_ROW:[2,3],'对象':['A','B'],'值':[10,10]}) - tables,_=modeling.evaluate(frame,{'entity':'对象','directions':{'值':'cost'},'weighting':'entropy'}) - self.assertEqual(tables[0][1]['排名'].tolist(),[1,1]) - self.assertEqual(tables[0][1]['得分'].tolist(),[.5,.5]) + frame = pd.DataFrame( + {data.SOURCE_ROW: [2, 3], "对象": ["A", "B"], "值": [10, 10]} + ) + tables, _ = modeling.evaluate( + frame, + {"entity": "对象", "directions": {"值": "cost"}, "weighting": "entropy"}, + ) + self.assertEqual(tables[0][1]["排名"].tolist(), [1, 1]) + self.assertEqual(tables[0][1]["得分"].tolist(), [0.5, 0.5]) def test_ahp_rejects_inconsistent_matrix(self): - frame=pd.DataFrame({data.SOURCE_ROW:[2,3],'对象':['A','B'],'a':[1,2],'b':[2,1],'c':[2,3]}) - with self.assertRaisesRegex(ValueError,'一致性'): - modeling.evaluate(frame,{'entity':'对象','directions':{'a':'benefit','b':'benefit','c':'benefit'},'weighting':'ahp','comparison_matrix':[[1,9,1/9],[1/9,1,9],[9,1/9,1]]}) + frame = pd.DataFrame( + { + data.SOURCE_ROW: [2, 3], + "对象": ["A", "B"], + "a": [1, 2], + "b": [2, 1], + "c": [2, 3], + } + ) + with self.assertRaisesRegex(ValueError, "一致性"): + modeling.evaluate( + frame, + { + "entity": "对象", + "directions": {"a": "benefit", "b": "benefit", "c": "benefit"}, + "weighting": "ahp", + "comparison_matrix": [[1, 9, 1 / 9], [1 / 9, 1, 9], [9, 1 / 9, 1]], + }, + ) def test_forecast_by_group_and_time_holdout(self): - dates=pd.date_range('2024-01-01',periods=24,freq='MS') - frame=pd.DataFrame({'城市':['A']*24+['B']*24,'月份':list(dates)*2,'销量':list(np.arange(24)*10+100)+list(np.arange(24)*-2+100)}) - tables,_=modeling.forecast(frame,{'by':['城市'],'date':'月份','value':'销量','horizon':3,'frequency':'MS'}) - result=tables[0][1] - self.assertEqual(len(result),6) - self.assertAlmostEqual(result[result['城市']=='A'].iloc[0]['预测值'],340) - self.assertAlmostEqual(result[result['城市']=='B'].iloc[0]['预测值'],52) - self.assertTrue((tables[1][1]['训练期数']==18).all()) + dates = pd.date_range("2024-01-01", periods=24, freq="MS") + frame = pd.DataFrame( + { + "城市": ["A"] * 24 + ["B"] * 24, + "月份": list(dates) * 2, + "销量": list(np.arange(24) * 10 + 100) + list(np.arange(24) * -2 + 100), + } + ) + tables, _ = modeling.forecast( + frame, + { + "by": ["城市"], + "date": "月份", + "value": "销量", + "horizon": 3, + "frequency": "MS", + }, + ) + result = tables[0][1] + self.assertEqual(len(result), 6) + self.assertAlmostEqual(result[result["城市"] == "A"].iloc[0]["预测值"], 340) + self.assertAlmostEqual(result[result["城市"] == "B"].iloc[0]["预测值"], 52) + self.assertTrue((tables[1][1]["训练期数"] == 18).all()) def test_forecast_rejects_missing_month(self): - frame=pd.DataFrame({'日期':pd.date_range('2024-01-01',periods=8,freq='MS').delete(3),'值':range(7)}) - with self.assertRaisesRegex(ValueError,'连续'): - modeling.forecast(frame,{'date':'日期','value':'值'}) + frame = pd.DataFrame( + { + "日期": pd.date_range("2024-01-01", periods=8, freq="MS").delete(3), + "值": range(7), + } + ) + with self.assertRaisesRegex(ValueError, "连续"): + modeling.forecast(frame, {"date": "日期", "value": "值"}) def test_regression_has_holdout_and_baseline(self): - frame=pd.DataFrame({data.SOURCE_ROW:range(2,42),'x':range(40),'y':np.arange(40)*3+7}) - tables,meta=modeling.supervised(frame,{'features':['x'],'target':'y'},classification=False) - metrics=tables[1][1] - error=metrics[(metrics['数据集']=='测试')&(metrics['模型']=='linear')&(metrics['指标']=='RMSE')].iloc[0]['值'] - self.assertLess(error,1e-8) - self.assertEqual(meta['train_rows'],32) - self.assertIn('基线',metrics['模型'].tolist()) + frame = pd.DataFrame( + {data.SOURCE_ROW: range(2, 42), "x": range(40), "y": np.arange(40) * 3 + 7} + ) + tables, meta = modeling.supervised( + frame, {"features": ["x"], "target": "y"}, classification=False + ) + metrics = tables[1][1] + error = metrics[ + (metrics["数据集"] == "测试") + & (metrics["模型"] == "linear") + & (metrics["指标"] == "RMSE") + ].iloc[0]["值"] + self.assertLess(error, 1e-8) + self.assertEqual(meta["train_rows"], 32) + self.assertIn("基线", metrics["模型"].tolist()) def test_classification_categorical_pipeline(self): - frame=pd.DataFrame({data.SOURCE_ROW:range(2,42),'x':list(range(20))*2,'组':['A']*20+['B']*20,'标签':['低']*20+['高']*20}) - tables,_=modeling.supervised(frame,{'features':['x','组'],'categorical':['组'],'target':'标签'},classification=True) - self.assertEqual(len(tables[0][1]),8) - self.assertIn('F1_macro',tables[1][1]['指标'].tolist()) + frame = pd.DataFrame( + { + data.SOURCE_ROW: range(2, 42), + "x": list(range(20)) * 2, + "组": ["A"] * 20 + ["B"] * 20, + "标签": ["低"] * 20 + ["高"] * 20, + } + ) + tables, _ = modeling.supervised( + frame, + {"features": ["x", "组"], "categorical": ["组"], "target": "标签"}, + classification=True, + ) + self.assertEqual(len(tables[0][1]), 8) + self.assertIn("F1_macro", tables[1][1]["指标"].tolist()) def test_small_regression_has_two_test_rows_for_r_squared(self): - frame = pd.DataFrame({data.SOURCE_ROW: range(2, 12), 'x': range(10), 'y': np.arange(10) * 3 + 7}) - tables, metadata = modeling.supervised(frame, {'features': ['x'], 'target': 'y', 'test_fraction': .1}, classification=False) - self.assertEqual(metadata['test_rows'], 2) - self.assertTrue(np.isfinite(tables[1][1]['值']).all()) + frame = pd.DataFrame( + {data.SOURCE_ROW: range(2, 12), "x": range(10), "y": np.arange(10) * 3 + 7} + ) + tables, metadata = modeling.supervised( + frame, + {"features": ["x"], "target": "y", "test_fraction": 0.1}, + classification=False, + ) + self.assertEqual(metadata["test_rows"], 2) + self.assertTrue(np.isfinite(tables[1][1]["值"]).all()) def test_clustering_separates_obvious_groups(self): - frame=pd.DataFrame({data.SOURCE_ROW:range(8),'x':[0,.1,.2,.3,10,10.1,10.2,10.3]}) - tables,_=modeling.unsupervised(frame,{'features':['x'],'clusters':2},anomaly=False) - labels=tables[0][1]['簇编号'].to_numpy() - self.assertTrue(np.all(labels[:4]==labels[0])) - self.assertNotEqual(labels[0],labels[-1]) + frame = pd.DataFrame( + {data.SOURCE_ROW: range(8), "x": [0, 0.1, 0.2, 0.3, 10, 10.1, 10.2, 10.3]} + ) + tables, _ = modeling.unsupervised( + frame, {"features": ["x"], "clusters": 2}, anomaly=False + ) + labels = tables[0][1]["簇编号"].to_numpy() + self.assertTrue(np.all(labels[:4] == labels[0])) + self.assertNotEqual(labels[0], labels[-1]) def test_integer_optimization(self): - tables,meta=modeling.optimize({'variables':['x','y'],'objective':[3,2],'sense':'max','integer':[True,True],'constraints':[{'coefficients':[2,1],'relation':'<=','rhs':4}]}) - self.assertEqual(meta['objective_value'],8) - self.assertEqual(tables[0][1]['取值'].tolist(),[0,4]) + tables, meta = modeling.optimize( + { + "variables": ["x", "y"], + "objective": [3, 2], + "sense": "max", + "integer": [True, True], + "constraints": [{"coefficients": [2, 1], "relation": "<=", "rhs": 4}], + } + ) + self.assertEqual(meta["objective_value"], 8) + self.assertEqual(tables[0][1]["取值"].tolist(), [0, 4]) def test_infeasible_and_unbounded_rejected(self): - with self.assertRaisesRegex(ValueError,'求解未成功'): - modeling.optimize({'variables':['x'],'objective':[1],'constraints':[{'coefficients':[1],'relation':'<=','rhs':-1}]}) - with self.assertRaisesRegex(ValueError,'求解未成功'): - modeling.optimize({'variables':['x'],'objective':[1],'sense':'max'}) + with self.assertRaisesRegex(ValueError, "求解未成功"): + modeling.optimize( + { + "variables": ["x"], + "objective": [1], + "constraints": [{"coefficients": [1], "relation": "<=", "rhs": -1}], + } + ) + with self.assertRaisesRegex(ValueError, "求解未成功"): + modeling.optimize({"variables": ["x"], "objective": [1], "sense": "max"}) def test_convex_quadratic(self): - tables,meta=modeling.optimize({'variables':['x'],'objective':[-4],'quadratic':[[2]]}) - self.assertAlmostEqual(tables[0][1].iloc[0]['取值'],2) - self.assertAlmostEqual(meta['objective_value'],-4) + tables, meta = modeling.optimize( + {"variables": ["x"], "objective": [-4], "quadratic": [[2]]} + ) + self.assertAlmostEqual(tables[0][1].iloc[0]["取值"], 2) + self.assertAlmostEqual(meta["objective_value"], -4) def test_output_root_enforced(self): with self.assertRaises(ValueError): - data.save_plan([('结果',pd.DataFrame({'x':[1]}))],{},'/private/tmp/outside-plan.json') + data.save_plan( + [("结果", pd.DataFrame({"x": [1]}))], + {}, + "/private/tmp/outside-plan.json", + ) -if __name__ == '__main__': +if __name__ == "__main__": unittest.main(verbosity=2) diff --git a/tests/test_send_complex_message.py b/tests/test_send_complex_message.py index cdd67cc..268681c 100644 --- a/tests/test_send_complex_message.py +++ b/tests/test_send_complex_message.py @@ -3,7 +3,9 @@ from __future__ import annotations import contextlib import importlib.util import io +import json import os +import sqlite3 import sys import unittest from pathlib import Path @@ -14,6 +16,60 @@ SCRIPT_PATH = ( Path(__file__).resolve().parents[1] / "skills/send-complex-message/scripts/send_complex_message.py" ) +NOW = 1_789_272_000 # 2026-09-13 12:00:00 +08:00 + + +class HistoryCursor: + """在内存数据库执行实际查询,仅转换 MySQL 的占位符和转义字符串语法。""" + + def __init__(self, database) -> None: + self.cursor = database.cursor() + + def __enter__(self): + return self + + def __exit__(self, *_): + self.cursor.close() + + def execute(self, sql, params) -> None: + sql = sql.replace("%s", "?").replace("ESCAPE '\\\\'", "ESCAPE '\\'") + self.cursor.execute(sql, params) + + def fetchall(self): + return [dict(row) for row in self.cursor.fetchall()] + + +class HistoryConnection: + def __init__(self, messages: list[dict]) -> None: + self.database = sqlite3.connect(":memory:") + self.database.row_factory = sqlite3.Row + self.closed = False + self.database.execute(""" + CREATE TABLE messages ( + id INTEGER PRIMARY KEY, from_wxid TEXT, sender_wxid TEXT, to_wxid TEXT, + is_chat_room INTEGER, type INTEGER, app_msg_type INTEGER, + content TEXT, display_full_content TEXT, is_recalled INTEGER, created_at INTEGER + ) + """) + for message in messages: + row = { + "from_wxid": "room@chatroom", "sender_wxid": "wxid_alice", + "to_wxid": "wxid_robot", "is_chat_room": 1, "type": 1, + "app_msg_type": 0, "content": "安排会议", "display_full_content": "", + "is_recalled": 0, "created_at": NOW - 60, + **message, + } + self.database.execute( + f"INSERT INTO messages ({', '.join(row)}) VALUES ({', '.join(['?'] * len(row))})", + tuple(row.values()), + ) + + def cursor(self): + return HistoryCursor(self.database) + + def close(self) -> None: + self.closed = True + self.database.close() class SendComplexMessageTests(unittest.TestCase): @@ -94,6 +150,194 @@ class SendComplexMessageTests(unittest.TestCase): post.assert_not_called() self.assertFalse(output.getvalue().endswith("ended")) + def run_history(self, rows, args=(), conversation="room@chatroom") -> dict: + connection = HistoryConnection(rows) + self.addCleanup(connection.close) + with contextlib.ExitStack() as stack: + # 查询无须客户端端口;也不应该调用任何发送接口。 + stack.enter_context(mock.patch.dict(os.environ, {"ROBOT_FROM_WX_ID": conversation}, clear=True)) + stack.enter_context(mock.patch.object(sys, "argv", [str(SCRIPT_PATH), "--query-history", *args])) + stack.enter_context(mock.patch.object(self.module.time, "time", return_value=NOW)) + stack.enter_context(mock.patch.object(self.module, "_mysql_connect", return_value=connection)) + send = stack.enter_context(mock.patch.object(self.module, "_http_post_json")) + output = stack.enter_context(contextlib.redirect_stdout(io.StringIO())) + + self.assertEqual(self.module.main(), 0, output.getvalue()) + self.assertTrue(connection.closed) + send.assert_not_called() + result = json.loads(output.getvalue()) + self.assertEqual(result["conversation_id"], conversation) + return result + + def test_history_group_isolation_and_24_hour_boundaries(self) -> None: + result = self.run_history([ + {"id": 1, "created_at": NOW - 86400}, + {"id": 2, "created_at": NOW}, + {"id": 3, "created_at": NOW - 86401}, + {"id": 4, "created_at": NOW + 1}, + {"id": 5, "from_wxid": "other@chatroom"}, + {"id": 6, "from_wxid": "wxid_alice", "is_chat_room": 0}, + {"id": 7, "is_chat_room": 0}, + ]) + self.assertEqual([row["id"] for row in result["messages"]], [2, 1]) + self.assertTrue(result["is_chat_room"]) + self.assertEqual((result["start_time"], result["end_time"]), (NOW - 86400, NOW)) + + def test_history_private_chat_includes_both_directions_only_for_current_friend(self) -> None: + result = self.run_history([ + {"id": 1, "from_wxid": "wxid_alice", "is_chat_room": 0}, + {"id": 2, "from_wxid": "wxid_alice", "is_chat_room": 0, "sender_wxid": "wxid_robot"}, + {"id": 3, "from_wxid": "wxid_bob", "is_chat_room": 0}, + {"id": 4}, # 同一个人的群消息不能混入私聊。 + {"id": 5, "from_wxid": "wxid_alice", "is_chat_room": 1}, + {"id": 6, "from_wxid": "wxid_alice", "is_chat_room": 0, "created_at": NOW - 86401}, + ], conversation="wxid_alice") + self.assertEqual([row["id"] for row in result["messages"]], [2, 1]) + self.assertFalse(result["is_chat_room"]) + + def test_history_accepts_shorter_relative_and_absolute_ranges(self) -> None: + rows = [ + {"id": 1, "created_at": NOW - 1801}, + {"id": 2, "created_at": NOW - 1800}, + {"id": 3, "created_at": NOW - 900}, + {"id": 4, "created_at": NOW - 899}, + {"id": 5, "created_at": NOW}, + ] + cases = [ + (["--hours", "0.5"], [5, 4, 3, 2], NOW - 1800, NOW), + (["--hours", "24"], [5, 4, 3, 2, 1], NOW - 86400, NOW), + (["--start-time", str(NOW - 1800), "--end-time", str(NOW - 900)], [3, 2], NOW - 1800, NOW - 900), + (["--start-time", "2026-09-13 11:30", "--end-time", "2026-09-13 11:45:00"], [3, 2], NOW - 1800, NOW - 900), + (["--start-time", str(NOW - 900)], [5, 4, 3], NOW - 900, NOW), + (["--end-time", str(NOW - 900)], [3, 2, 1], NOW - 86400, NOW - 900), + ] + for args, ids, start, end in cases: + with self.subTest(args=args): + result = self.run_history(rows, args) + self.assertEqual([row["id"] for row in result["messages"]], ids) + self.assertEqual((result["start_time"], result["end_time"]), (start, end)) + + def test_history_combines_keywords_types_and_sender(self) -> None: + result = self.run_history([ + {"id": 1, "content": "安排", "display_full_content": "会议通知", "type": 49, "app_msg_type": 57}, + {"id": 2, "content": "安排会议", "type": 49, "app_msg_type": 6}, + {"id": 3, "content": "安排会议", "type": 49, "app_msg_type": 5}, + {"id": 4, "content": "安排会议", "type": 1, "app_msg_type": 57}, + {"id": 5, "content": "安排会议", "type": 49, "app_msg_type": 57, "sender_wxid": "wxid_bob"}, + {"id": 6, "content": "安排", "type": 49, "app_msg_type": 57}, + {"id": 7, "content": "安排会议", "type": 49, "app_msg_type": 57, "from_wxid": "other@chatroom"}, + {"id": 8, "content": "安排会议", "type": 49, "app_msg_type": 57, "created_at": NOW - 3601}, + ], ["--hours", "1", "--keyword", "安排", "--keyword", "会议", "--message-type", "app", + "--app-msg-type", "57", "--app-msg-type", "6", "--sender-wxid", "wxid_alice"]) + self.assertEqual([row["id"] for row in result["messages"]], [2, 1]) + + def test_history_accepts_multiple_message_types_and_app_subtype_alone(self) -> None: + rows = [ + {"id": 1, "type": 3}, {"id": 2, "type": 43}, {"id": 3, "type": 34}, + {"id": 4, "type": 49, "app_msg_type": 57}, {"id": 5, "type": 1, "app_msg_type": 57}, + ] + result = self.run_history(rows, ["--message-type", "image", "--message-type", "43"]) + self.assertEqual([row["id"] for row in result["messages"]], [2, 1]) + result = self.run_history(rows, ["--app-msg-type", "57"]) + self.assertEqual([row["id"] for row in result["messages"]], [4]) + + def test_history_keywords_are_literal_and_cannot_bypass_scope(self) -> None: + for keyword in ["100%_\\", "' OR 1=1 --"]: + with self.subTest(keyword=keyword): + result = self.run_history([ + {"id": 1, "content": f"前缀{keyword}后缀"}, + {"id": 2, "content": "100AB\\"}, + {"id": 3, "content": keyword, "from_wxid": "other@chatroom"}, + ], ["--keyword", keyword]) + self.assertEqual([row["id"] for row in result["messages"]], [1]) + + def test_history_pagination_preserves_order_and_primary_ids(self) -> None: + rows = [{"id": 1}, {"id": 2}, {"id": 9007199254740993}] + first = self.run_history(rows, ["--limit", "2"]) + self.assertEqual([row["id"] for row in first["messages"]], [9007199254740993, 2]) + self.assertEqual(first["count"], 2) + self.assertTrue(first["has_more"]) + self.assertEqual(first["next_offset"], 2) + last = self.run_history(rows, ["--limit", "2", "--offset", str(first["next_offset"])]) + self.assertEqual([row["id"] for row in last["messages"]], [1]) + self.assertFalse(last["has_more"]) + self.assertIsNone(last["next_offset"]) + empty = self.run_history(rows, ["--offset", "3"]) + self.assertEqual(empty["messages"], []) + self.assertEqual(empty["count"], 0) + self.assertFalse(empty["has_more"]) + + def test_invalid_history_requests_neither_query_nor_send(self) -> None: + cases = [ + ["--hours", "0"], ["--hours", "-1"], ["--hours", "24.01"], + ["--hours", "nan"], ["--hours", "inf"], ["--hours", "0.00001"], ["--hours", "bad"], + ["--start-time", str(NOW - 86401)], + ["--start-time", str(NOW - 172800), "--end-time", str(NOW - 86400)], + ["--start-time", str(NOW - 172800), "--end-time", str(NOW - 169200)], + ["--end-time", str(NOW + 1)], ["--end-time", str(NOW - 86401)], + ["--start-time", str(NOW)], ["--start-time", str(NOW + 1)], + ["--start-time", "invalid"], ["--start-time", ""], ["--end-time", "2026-09-31 09:00"], + ["--start-time", str(NOW - 60), "--end-time", str(NOW - 120)], + ["--hours", "1", "--start-time", str(NOW - 60)], + ["--hours", "1", "--end-time", str(NOW - 60)], + ["--limit", "0"], ["--limit", "201"], ["--offset", "-1"], + ["--keyword", " "], ["--sender-wxid", " "], + ["--message-type", "bad"], ["--message-type", "0"], ["--app-msg-type", "-1"], + ["--message-type", "text", "--app-msg-type", "57"], + ["--content", "你好"], ["--content", ""], ["--mention", "张三"], ["--mentions", "[]"], + ["--all"], ["--refer-message-id", "1"], ["--ended"], + ["--conversation-id", "other@chatroom"], ["--from-wxid", "other@chatroom"], + ] + for args in cases: + with self.subTest(args=args), contextlib.ExitStack() as stack: + stack.enter_context(mock.patch.object(sys, "argv", [str(SCRIPT_PATH), "--query-history", *args])) + stack.enter_context(mock.patch.object(self.module.time, "time", return_value=NOW)) + connect = stack.enter_context(mock.patch.object(self.module, "_mysql_connect")) + send = stack.enter_context(mock.patch.object(self.module, "_http_post_json")) + output = stack.enter_context(contextlib.redirect_stdout(io.StringIO())) + self.assertEqual(self.module.main(), 1) + connect.assert_not_called() + send.assert_not_called() + self.assertFalse(output.getvalue().endswith("ended")) + + def test_history_filters_require_explicit_query_mode(self) -> None: + for flag, value in [("--hours", "1"), ("--limit", "50"), ("--keyword", "会议")]: + with self.subTest(flag=flag): + with self.assertRaises(ValueError): + self.module._parse_cli_params(["--content", "你好", flag, value]) + + def test_history_requires_current_conversation(self) -> None: + with mock.patch.dict(os.environ, {}, clear=True), \ + mock.patch.object(sys, "argv", [str(SCRIPT_PATH), "--query-history"]), \ + mock.patch.object(self.module, "_mysql_connect") as connect, \ + contextlib.redirect_stdout(io.StringIO()) as output: + self.assertEqual(self.module.main(), 1) + self.assertIn("ROBOT_FROM_WX_ID", output.getvalue()) + connect.assert_not_called() + + def test_history_recomputes_window_after_connecting(self) -> None: + connection = HistoryConnection([{"id": 1, "created_at": NOW - 86400}]) + with mock.patch.object(sys, "argv", [str(SCRIPT_PATH), "--query-history"]), \ + mock.patch.object(self.module.time, "time", side_effect=[NOW, NOW + 10]), \ + mock.patch.object(self.module, "_mysql_connect", return_value=connection), \ + contextlib.redirect_stdout(io.StringIO()) as output: + self.assertEqual(self.module.main(), 0) + self.assertEqual(json.loads(output.getvalue())["messages"], []) + self.assertTrue(connection.closed) + + def test_history_query_failure_closes_connection(self) -> None: + connection = mock.MagicMock() + connection.cursor.side_effect = RuntimeError("查询失败") + with mock.patch.object(sys, "argv", [str(SCRIPT_PATH), "--query-history"]), \ + mock.patch.object(self.module, "_mysql_connect", return_value=connection), \ + mock.patch.object(self.module, "_http_post_json") as send, \ + contextlib.redirect_stdout(io.StringIO()) as output: + self.assertEqual(self.module.main(), 1) + self.assertIn("查询失败", output.getvalue()) + self.assertFalse(output.getvalue().endswith("ended")) + connection.close.assert_called_once() + send.assert_not_called() + if __name__ == "__main__": unittest.main()