fix: 消除代码静态类型警告

This commit is contained in:
hp0912 2026-09-13 12:56:56 +08:00
parent ccd00d948f
commit 7a0750f043
56 changed files with 1371 additions and 387 deletions

View File

@ -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`

17
eslint.config.cjs Normal file
View File

@ -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 },
},
];

19
package.json Normal file
View File

@ -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"
}
}

16
pyrightconfig.json Normal file
View File

@ -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"
}

27
requirements-dev.txt Normal file
View File

@ -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

4
ruff.toml Normal file
View File

@ -0,0 +1,4 @@
target-version = "py312"
[lint]
select = ["E4", "E7", "E9", "F", "W", "RUF100"]

View File

@ -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

View File

@ -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()

View File

@ -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:

View File

@ -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)

View File

@ -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}",

View File

@ -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(

View File

@ -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(

View File

@ -9,7 +9,6 @@ from typing import Any, Optional
from _docx_common import (
DOCX_INPUT_SUFFIXES,
NS,
SkillArgumentParser,
W_NS,
input_file,

View File

@ -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"])

View File

@ -131,4 +131,4 @@ if __name__ == "__main__":
raise
except Exception:
traceback.print_exc(file=sys.stdout)
raise SystemExit(1)
raise SystemExit(1)

View File

@ -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)
raise SystemExit(1)

View File

@ -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:

View File

@ -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()

View File

@ -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()

View File

@ -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()

View File

@ -43,4 +43,4 @@ if __name__ == "__main__":
raise
except Exception:
traceback.print_exc(file=sys.stdout)
raise SystemExit(1)
raise SystemExit(1)

View File

@ -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'});
}
});

View File

@ -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())

View File

@ -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())

View File

@ -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:

View File

@ -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)

View File

@ -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

View File

@ -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 (

View File

@ -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()]

View File

@ -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):

View File

@ -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]:

View File

@ -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],

View File

@ -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:

View File

@ -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 行为

View File

@ -108,4 +108,4 @@ if __name__ == "__main__":
raise
except Exception:
traceback.print_exc(file=sys.stdout)
raise SystemExit(1)
raise SystemExit(1)

View File

@ -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

View File

@ -108,4 +108,4 @@ if __name__ == "__main__":
raise
except Exception:
traceback.print_exc(file=sys.stdout)
raise SystemExit(1)
raise SystemExit(1)

View File

@ -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()

View File

@ -130,4 +130,4 @@ if __name__ == "__main__":
raise
except Exception:
traceback.print_exc(file=sys.stdout)
raise SystemExit(1)
raise SystemExit(1)

View File

@ -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()

View File

@ -131,4 +131,4 @@ if __name__ == "__main__":
raise
except Exception:
traceback.print_exc(file=sys.stdout)
raise SystemExit(1)
raise SystemExit(1)

View File

@ -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()

View File

@ -112,4 +112,4 @@ if __name__ == "__main__":
raise
except Exception:
traceback.print_exc(file=sys.stdout)
raise SystemExit(1)
raise SystemExit(1)

View File

@ -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,请安装后重试"

View File

@ -7,6 +7,8 @@
"lib": ["ES2022", "DOM"],
"types": ["node"],
"strict": true,
"noUnusedLocals": true,
"noUnusedParameters": true,
"noEmit": true,
"esModuleInterop": true,
"forceConsistentCasingInFileNames": true,

View File

@ -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<ScriptResult> {
});
}
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(
'<!doctype html><title>click file</title><a id="local-file" href="file:///etc/passwd">打开本地文件</a>',
'<!doctype html><title>link</title><a id="search-link" href="/search?q=clicked-link">查看结果</a>',
);
return;
case "/js-location":
response.end(
'<!doctype html><title>js location</title><p>safe</p><script>setTimeout(() => { location.href = "file:///etc/passwd"; }, 50);</script>',
'<!doctype html><title>js location</title><script>setTimeout(() => { location.href = "/search?q=js-navigation"; }, 50);</script>',
);
return;
case "/iframe-file":
case "/form":
response.end(
'<!doctype html><title>iframe file</title><p>safe</p><iframe src="file:///etc/passwd"></iframe>',
);
return;
case "/popup-file":
response.end(
'<!doctype html><title>popup file</title><button id="open-popup" onclick="window.open(\'file:///etc/passwd\', \'_blank\')">打开弹窗</button>',
'<!doctype html><title>search form</title><form action="/search" method="get"><input id="query" name="q"><button id="submit" type="submit">搜索</button></form>',
);
return;
case "/redirect-http":
@ -113,7 +97,7 @@ test("web-page 本地文件访问防护", async (t) => {
return;
case "/search":
response.end(
`<!doctype html><title>search ok</title><main>SEARCH_OK:${requestUrl.searchParams.get("q") || ""}</main><script>console.log("file:///not-an-access-attempt")</script>`,
`<!doctype html><title>search ok</title><main id="search-result">SEARCH_OK:${requestUrl.searchParams.get("q") || ""}</main><script>console.log("console output is not page content")</script>`,
);
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<void>((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 地址/);
});
}
});

View File

@ -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)

View File

@ -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("没有数据可用于图表")

View File

@ -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()

View File

@ -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"],

View File

@ -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)

View File

@ -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),

View File

@ -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 求解成功;仅在成功且可行时输出方案"

View File

@ -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)

View File

@ -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()