fix: 消除代码静态类型警告
This commit is contained in:
parent
ccd00d948f
commit
7a0750f043
14
README.md
14
README.md
@ -2,6 +2,20 @@
|
|||||||
|
|
||||||
微信机器人 Skills
|
微信机器人 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`
|
- ROBOT_WECHAT_CLIENT_PORT: 机器人客户端服务端口,可用于在 SKILL 脚本直接调用客户端接口 `http://127.0.0.1:{ROBOT_WECHAT_CLIENT_PORT}/api/v1/xxxxx`
|
||||||
|
|||||||
17
eslint.config.cjs
Normal file
17
eslint.config.cjs
Normal 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
19
package.json
Normal 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
16
pyrightconfig.json
Normal 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
27
requirements-dev.txt
Normal 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
4
ruff.toml
Normal file
@ -0,0 +1,4 @@
|
|||||||
|
target-version = "py312"
|
||||||
|
|
||||||
|
[lint]
|
||||||
|
select = ["E4", "E7", "E9", "F", "W", "RUF100"]
|
||||||
@ -18,7 +18,7 @@ from typing import Any, Literal, NoReturn, TypedDict
|
|||||||
try:
|
try:
|
||||||
from zoneinfo import ZoneInfo
|
from zoneinfo import ZoneInfo
|
||||||
except ImportError: # pragma: no cover - Python 3.8 fallback
|
except ImportError: # pragma: no cover - Python 3.8 fallback
|
||||||
ZoneInfo = None # type: ignore[assignment,misc]
|
ZoneInfo = None
|
||||||
|
|
||||||
sys.stderr = sys.stdout
|
sys.stderr = sys.stdout
|
||||||
|
|
||||||
|
|||||||
@ -2,11 +2,9 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from copy import deepcopy
|
from typing import Any, Optional
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Iterable, Optional
|
|
||||||
|
|
||||||
from _docx_common import W_NS, input_file, qn
|
from _docx_common import input_file, qn
|
||||||
|
|
||||||
|
|
||||||
MAX_BLOCKS = 1_000
|
MAX_BLOCKS = 1_000
|
||||||
@ -132,7 +130,7 @@ def _add_hyperlink(
|
|||||||
is_external=True,
|
is_external=True,
|
||||||
)
|
)
|
||||||
hyperlink = OxmlElement("w:hyperlink")
|
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)
|
run = paragraph.add_run(text)
|
||||||
apply_run_style(run, raw_spec)
|
apply_run_style(run, raw_spec)
|
||||||
run_properties = run._element.get_or_add_rPr()
|
run_properties = run._element.get_or_add_rPr()
|
||||||
@ -232,7 +230,7 @@ def add_paragraph_from_spec(
|
|||||||
style: Optional[str] = None,
|
style: Optional[str] = None,
|
||||||
default_run_style: Optional[dict[str, Any]] = None,
|
default_run_style: Optional[dict[str, Any]] = None,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
spec = (
|
spec: dict[str, Any] = (
|
||||||
{"text": raw_spec}
|
{"text": raw_spec}
|
||||||
if isinstance(raw_spec, str)
|
if isinstance(raw_spec, str)
|
||||||
else expect_object(raw_spec, "paragraph")
|
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:
|
def _add_horizontal_rule(container: Any, raw_spec: Any) -> Any:
|
||||||
from docx.oxml import OxmlElement
|
from docx.oxml import OxmlElement
|
||||||
|
from lxml.etree import SubElement
|
||||||
|
|
||||||
spec = expect_object(raw_spec, "horizontal_rule")
|
spec = expect_object(raw_spec, "horizontal_rule")
|
||||||
paragraph = container.add_paragraph()
|
paragraph = container.add_paragraph()
|
||||||
properties = paragraph._p.get_or_add_pPr()
|
properties = paragraph._p.get_or_add_pPr()
|
||||||
borders = OxmlElement("w:pBdr")
|
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("val"), str(spec.get("style", "single")))
|
||||||
bottom.set(qn("sz"), str(int(spec.get("size", 6))))
|
bottom.set(qn("sz"), str(int(spec.get("size", 6))))
|
||||||
bottom.set(qn("space"), str(int(spec.get("space", 1))))
|
bottom.set(qn("space"), str(int(spec.get("space", 1))))
|
||||||
bottom.set(qn("color"), color(spec.get("color", "808080"), "rule.color"))
|
bottom.set(qn("color"), color(spec.get("color", "808080"), "rule.color"))
|
||||||
borders.append(bottom)
|
|
||||||
properties.append(borders)
|
properties.append(borders)
|
||||||
return paragraph
|
return paragraph
|
||||||
|
|
||||||
@ -587,7 +585,7 @@ def _points(value: float) -> Any:
|
|||||||
|
|
||||||
def apply_page_settings(section: Any, raw_spec: Any) -> None:
|
def apply_page_settings(section: Any, raw_spec: Any) -> None:
|
||||||
from docx.enum.section import WD_ORIENT
|
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")
|
spec = expect_object(raw_spec, "page")
|
||||||
size = str(spec.get("size", "A4")).upper()
|
size = str(spec.get("size", "A4")).upper()
|
||||||
|
|||||||
@ -9,7 +9,7 @@ from datetime import datetime, timezone
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from _docx_common import W_NS, parse_xml_bytes, qn
|
from _docx_common import parse_xml_bytes, qn
|
||||||
|
|
||||||
|
|
||||||
class TrackedReplacement:
|
class TrackedReplacement:
|
||||||
|
|||||||
@ -198,7 +198,7 @@ def main() -> dict[str, Any]:
|
|||||||
os.close(descriptor)
|
os.close(descriptor)
|
||||||
temp_path = Path(temp_name)
|
temp_path = Path(temp_name)
|
||||||
try:
|
try:
|
||||||
document.save(temp_path)
|
document.save(str(temp_path))
|
||||||
archive = inspect_archive(temp_path)
|
archive = inspect_archive(temp_path)
|
||||||
Document(str(temp_path))
|
Document(str(temp_path))
|
||||||
publish_file(temp_path, destination, overwrite=args.overwrite)
|
publish_file(temp_path, destination, overwrite=args.overwrite)
|
||||||
|
|||||||
@ -13,7 +13,6 @@ from _docx_common import (
|
|||||||
DOCUMENT_OUTPUT_SUFFIXES,
|
DOCUMENT_OUTPUT_SUFFIXES,
|
||||||
WORD_INPUT_SUFFIXES,
|
WORD_INPUT_SUFFIXES,
|
||||||
SkillArgumentParser,
|
SkillArgumentParser,
|
||||||
find_program,
|
|
||||||
input_file,
|
input_file,
|
||||||
output_file,
|
output_file,
|
||||||
publish_file,
|
publish_file,
|
||||||
@ -74,7 +73,7 @@ def _extract_text(
|
|||||||
pandoc = shutil.which("pandoc")
|
pandoc = shutil.which("pandoc")
|
||||||
if pandoc:
|
if pandoc:
|
||||||
target = "gfm" if markdown else "plain"
|
target = "gfm" if markdown else "plain"
|
||||||
completed = run_program(
|
run_program(
|
||||||
[
|
[
|
||||||
pandoc,
|
pandoc,
|
||||||
f"--track-changes={track_changes}",
|
f"--track-changes={track_changes}",
|
||||||
|
|||||||
@ -96,7 +96,7 @@ def main() -> dict[str, Any]:
|
|||||||
os.close(descriptor)
|
os.close(descriptor)
|
||||||
temp_path = Path(temp_name)
|
temp_path = Path(temp_name)
|
||||||
try:
|
try:
|
||||||
document.save(temp_path)
|
document.save(str(temp_path))
|
||||||
archive = inspect_archive(temp_path)
|
archive = inspect_archive(temp_path)
|
||||||
if archive["missing_required_parts"]:
|
if archive["missing_required_parts"]:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@ -8,7 +8,7 @@ import re
|
|||||||
import tempfile
|
import tempfile
|
||||||
import zipfile
|
import zipfile
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Iterable, Optional
|
from typing import Any, Iterable
|
||||||
|
|
||||||
from _document_builder import (
|
from _document_builder import (
|
||||||
add_blocks,
|
add_blocks,
|
||||||
@ -445,7 +445,7 @@ def main() -> dict[str, Any]:
|
|||||||
os.close(descriptor)
|
os.close(descriptor)
|
||||||
temp_path = Path(temp_name)
|
temp_path = Path(temp_name)
|
||||||
try:
|
try:
|
||||||
document.save(temp_path)
|
document.save(str(temp_path))
|
||||||
archive = inspect_archive(temp_path)
|
archive = inspect_archive(temp_path)
|
||||||
if archive["missing_required_parts"]:
|
if archive["missing_required_parts"]:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@ -9,7 +9,6 @@ from typing import Any, Optional
|
|||||||
|
|
||||||
from _docx_common import (
|
from _docx_common import (
|
||||||
DOCX_INPUT_SUFFIXES,
|
DOCX_INPUT_SUFFIXES,
|
||||||
NS,
|
|
||||||
SkillArgumentParser,
|
SkillArgumentParser,
|
||||||
W_NS,
|
W_NS,
|
||||||
input_file,
|
input_file,
|
||||||
|
|||||||
@ -10,19 +10,22 @@ import tempfile
|
|||||||
import unittest
|
import unittest
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import cast
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
sys.dont_write_bytecode = True
|
|
||||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "scripts"))
|
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "scripts"))
|
||||||
|
|
||||||
from docx import Document
|
from docx import Document
|
||||||
from docx.oxml import OxmlElement
|
from docx.oxml import OxmlElement
|
||||||
from docx.oxml.ns import qn
|
from docx.oxml.ns import qn
|
||||||
|
from lxml.etree import _Element
|
||||||
from reportlab.pdfgen.canvas import Canvas
|
from reportlab.pdfgen.canvas import Canvas
|
||||||
|
|
||||||
import _docx_common as common
|
import _docx_common as common
|
||||||
import compile_typst
|
import compile_typst
|
||||||
|
|
||||||
|
sys.dont_write_bytecode = True
|
||||||
|
|
||||||
|
|
||||||
class DocumentWorkflows(unittest.TestCase):
|
class DocumentWorkflows(unittest.TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
@ -45,14 +48,14 @@ class DocumentWorkflows(unittest.TestCase):
|
|||||||
p.add_run("HEL").bold = True
|
p.add_run("HEL").bold = True
|
||||||
p.add_run("LO")
|
p.add_run("LO")
|
||||||
p.add_run(" AFTER").italic = True
|
p.add_run(" AFTER").italic = True
|
||||||
doc.save(self.source)
|
doc.save(str(self.source))
|
||||||
return doc
|
return doc
|
||||||
|
|
||||||
def edit(self, operations, *args):
|
def edit(self, operations, *args):
|
||||||
output = self.root / "edited.docx"
|
output = self.root / "edited.docx"
|
||||||
result = self.call("edit_document", "--input", self.source, "--output", output,
|
result = self.call("edit_document", "--input", self.source, "--output", output,
|
||||||
"--spec", json.dumps({"operations": operations}), *args)
|
"--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):
|
def test_cross_run_preserves_unmodified_styles_and_escapes_text(self):
|
||||||
self.fixture()
|
self.fixture()
|
||||||
@ -67,7 +70,7 @@ class DocumentWorkflows(unittest.TestCase):
|
|||||||
def test_single_and_split_matches_are_both_replaced(self):
|
def test_single_and_split_matches_are_both_replaced(self):
|
||||||
doc = self.fixture()
|
doc = self.fixture()
|
||||||
doc.add_paragraph("HELLO HELLO")
|
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"}])
|
result, doc = self.edit([{"type": "replace_text", "find": "HELLO", "replace": "NEW"}])
|
||||||
self.assertEqual(result["operation_results"][0]["replacement_count"], 3)
|
self.assertEqual(result["operation_results"][0]["replacement_count"], 3)
|
||||||
self.assertNotIn("HELLO", "".join(p.text for p in doc.paragraphs))
|
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):
|
def test_multiple_tracked_matches_in_one_run(self):
|
||||||
doc = Document()
|
doc = Document()
|
||||||
doc.add_paragraph("old old old")
|
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")
|
result, doc = self.edit([{"type": "replace_text", "find": "old", "replace": "new"}], "--track-changes")
|
||||||
self.assertEqual(result["tracked_replacement_count"], 3)
|
self.assertEqual(result["tracked_replacement_count"], 3)
|
||||||
self.assertEqual(len(doc.element.xpath(".//w:p/w:ins")), 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 = OxmlElement("w:bookmarkStart")
|
||||||
mark.set(qn("w:id"), "7")
|
mark.set(qn("w:id"), "7")
|
||||||
mark.set(qn("w:name"), "target")
|
mark.set(qn("w:name"), "target")
|
||||||
doc.paragraphs[0]._p.insert(2, mark)
|
cast(_Element, doc.paragraphs[0]._p).insert(2, mark)
|
||||||
doc.save(self.source)
|
doc.save(str(self.source))
|
||||||
with self.assertRaisesRegex(ValueError, "书签"):
|
with self.assertRaisesRegex(ValueError, "书签"):
|
||||||
self.edit([{"type": "replace_text", "find": "HELLO", "replace": "new"}], "--track-changes")
|
self.edit([{"type": "replace_text", "find": "HELLO", "replace": "new"}], "--track-changes")
|
||||||
self.assertFalse((self.root / "edited.docx").exists())
|
self.assertFalse((self.root / "edited.docx").exists())
|
||||||
@ -128,14 +131,16 @@ class DocumentWorkflows(unittest.TestCase):
|
|||||||
spec = {"claims": [{"number": 1, "text": "一种方法 & 装置", "dependent": False}],
|
spec = {"claims": [{"number": 1, "text": "一种方法 & 装置", "dependent": False}],
|
||||||
"specification": {"field": "领域", "detailed": ["实现 <描述>"]}, "abstract": "摘要"}
|
"specification": {"field": "领域", "detailed": ["实现 <描述>"]}, "abstract": "摘要"}
|
||||||
result = self.call("create_document", "--preset", "patent", "--output", output, "--spec", json.dumps(spec))
|
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(len(doc.sections), 3)
|
||||||
self.assertEqual([s.header.paragraphs[0].text for s in doc.sections], ["权利要求书", "说明书", "摘要"])
|
self.assertEqual([s.header.paragraphs[0].text for s in doc.sections], ["权利要求书", "说明书", "摘要"])
|
||||||
for section 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.header.is_linked_to_previous)
|
||||||
self.assertFalse(section.footer.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")
|
self.assertEqual(self.call("validate_document", "--input", result["path"])["status"], "valid")
|
||||||
|
|
||||||
def test_invalid_patent_numbering_does_not_publish_a_file(self):
|
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"}],
|
{"type": "section_break"}, {"type": "paragraph", "text": "Second"}],
|
||||||
"sections": [{"index": 1, "header": {"text": "Second header"}}]}
|
"sections": [{"index": 1, "header": {"text": "Second header"}}]}
|
||||||
self.call("create_document", "--output", output, "--spec", json.dumps(spec))
|
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(doc.tables[0].cell(1, 1).text, "B")
|
||||||
self.assertEqual([s.header.paragraphs[0].text for s in doc.sections], ["Default", "Second header"])
|
self.assertEqual([s.header.paragraphs[0].text for s in doc.sections], ["Default", "Second header"])
|
||||||
|
|
||||||
|
|||||||
@ -131,4 +131,4 @@ if __name__ == "__main__":
|
|||||||
raise
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
traceback.print_exc(file=sys.stdout)
|
traceback.print_exc(file=sys.stdout)
|
||||||
raise SystemExit(1)
|
raise SystemExit(1)
|
||||||
|
|||||||
@ -68,7 +68,7 @@ def _ensure_skill_venv_python() -> None:
|
|||||||
_ensure_skill_venv_python()
|
_ensure_skill_venv_python()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import pymysql # type: ignore # noqa: E402
|
import pymysql
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
_run_bootstrap()
|
_run_bootstrap()
|
||||||
_py = _get_python_executable()
|
_py = _get_python_executable()
|
||||||
@ -362,4 +362,4 @@ if __name__ == "__main__":
|
|||||||
raise
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
traceback.print_exc(file=sys.stdout)
|
traceback.print_exc(file=sys.stdout)
|
||||||
raise SystemExit(1)
|
raise SystemExit(1)
|
||||||
|
|||||||
@ -3,6 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
|
import importlib
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
@ -20,7 +21,7 @@ from typing import Any, NoReturn
|
|||||||
try:
|
try:
|
||||||
from zoneinfo import ZoneInfo
|
from zoneinfo import ZoneInfo
|
||||||
except ImportError: # pragma: no cover - Python 3.8 fallback
|
except ImportError: # pragma: no cover - Python 3.8 fallback
|
||||||
ZoneInfo = None # type: ignore[assignment,misc]
|
ZoneInfo = None
|
||||||
|
|
||||||
sys.stderr = sys.stdout
|
sys.stderr = sys.stdout
|
||||||
|
|
||||||
@ -94,8 +95,8 @@ def _run_bootstrap() -> None:
|
|||||||
|
|
||||||
def _ensure_runtime_dependencies() -> None:
|
def _ensure_runtime_dependencies() -> None:
|
||||||
try:
|
try:
|
||||||
import openpyxl # noqa: F401
|
importlib.import_module("openpyxl")
|
||||||
import pymysql # noqa: F401
|
importlib.import_module("pymysql")
|
||||||
|
|
||||||
return
|
return
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
@ -109,8 +110,8 @@ def _ensure_runtime_dependencies() -> None:
|
|||||||
venv_dir = (_skill_root() / ".venv").resolve()
|
venv_dir = (_skill_root() / ".venv").resolve()
|
||||||
if Path(sys.prefix).resolve() == venv_dir:
|
if Path(sys.prefix).resolve() == venv_dir:
|
||||||
try:
|
try:
|
||||||
import openpyxl # noqa: F401
|
importlib.import_module("openpyxl")
|
||||||
import pymysql # noqa: F401
|
importlib.import_module("pymysql")
|
||||||
|
|
||||||
return
|
return
|
||||||
except ModuleNotFoundError as exc:
|
except ModuleNotFoundError as exc:
|
||||||
|
|||||||
@ -101,7 +101,7 @@ def _ensure_skill_venv_python() -> None:
|
|||||||
_ensure_skill_venv_python()
|
_ensure_skill_venv_python()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import pymysql # type: ignore[import-untyped] # noqa: E402
|
import pymysql
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
_run_bootstrap()
|
_run_bootstrap()
|
||||||
python_executable = _get_python_executable()
|
python_executable = _get_python_executable()
|
||||||
|
|||||||
@ -64,8 +64,8 @@ def _ensure_skill_venv_python() -> None:
|
|||||||
_ensure_skill_venv_python()
|
_ensure_skill_venv_python()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import pymysql # type: ignore # noqa: E402
|
import pymysql
|
||||||
from openai import OpenAI # type: ignore # noqa: E402
|
from openai import OpenAI
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
_run_bootstrap()
|
_run_bootstrap()
|
||||||
_py = _get_python_executable()
|
_py = _get_python_executable()
|
||||||
|
|||||||
@ -70,8 +70,8 @@ def _ensure_skill_venv_python() -> None:
|
|||||||
_ensure_skill_venv_python()
|
_ensure_skill_venv_python()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import pymysql # type: ignore # noqa: E402
|
import pymysql
|
||||||
from openai import OpenAI # type: ignore # noqa: E402
|
from openai import OpenAI
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
_run_bootstrap()
|
_run_bootstrap()
|
||||||
_py = _get_python_executable()
|
_py = _get_python_executable()
|
||||||
|
|||||||
@ -43,4 +43,4 @@ if __name__ == "__main__":
|
|||||||
raise
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
traceback.print_exc(file=sys.stdout)
|
traceback.print_exc(file=sys.stdout)
|
||||||
raise SystemExit(1)
|
raise SystemExit(1)
|
||||||
|
|||||||
@ -60,7 +60,7 @@ async function render(request) {
|
|||||||
if (!within(root, target) || !MIME[path.extname(target).toLowerCase()]) throw new Error('forbidden asset');
|
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');
|
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}});
|
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'});
|
await page.goto('https://pdf.local/' + encodeURIComponent(path.basename(request.input)), {waitUntil:'load'});
|
||||||
if (await page.evaluate(() => document.compatMode !== 'CSS1Compat'))
|
if (await page.evaluate(() => document.compatMode !== 'CSS1Compat'))
|
||||||
@ -84,9 +84,9 @@ async function render(request) {
|
|||||||
if (unsafe) throw new Error('Mermaid 不允许内嵌配置;主题和安全选项由固定渲染器设置');
|
if (unsafe) throw new Error('Mermaid 不允许内嵌配置;主题和安全选项由固定渲染器设置');
|
||||||
await inject(page, path.join(libraries.mermaid,'dist/mermaid.min.js'));
|
await inject(page, path.join(libraries.mermaid,'dist/mermaid.min.js'));
|
||||||
await page.evaluate(async () => {
|
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});
|
flowchart:{htmlLabels:false}, suppressErrorRendering:true});
|
||||||
await mermaid.run({querySelector:'.mermaid'});
|
await window.mermaid.run({querySelector:'.mermaid'});
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
if (stats.math) {
|
if (stats.math) {
|
||||||
@ -94,7 +94,7 @@ async function render(request) {
|
|||||||
await inject(page, path.join(libraries.katex,'dist/katex.min.js'));
|
await inject(page, path.join(libraries.katex,'dist/katex.min.js'));
|
||||||
await page.evaluate(() => {
|
await page.evaluate(() => {
|
||||||
for (const el of document.querySelectorAll('.math-inline,.math-display')) {
|
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'});
|
trust:false, maxExpand:1000, maxSize:30, strict:'warn'});
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
|
|||||||
@ -1,5 +1,6 @@
|
|||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
"""Compile a LaTeX project with cached Tectonic resources and shell escape disabled."""
|
"""Compile a LaTeX project with cached Tectonic resources and shell escape disabled."""
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import re
|
import re
|
||||||
import shutil
|
import shutil
|
||||||
@ -7,51 +8,105 @@ import subprocess
|
|||||||
import sys
|
import sys
|
||||||
import tempfile
|
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):
|
def compile_document(args):
|
||||||
from pypdf import PdfReader
|
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:
|
source = Path(args.input).expanduser().resolve()
|
||||||
raise ValueError('输入需为不超过 2 MiB 的本地 .tex 文件')
|
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:
|
if not 1 <= args.timeout <= 600:
|
||||||
raise ValueError('timeout 必须在 1–600 秒之间')
|
raise ValueError("timeout 必须在 1–600 秒之间")
|
||||||
executable=shutil.which('tectonic')
|
executable = shutil.which("tectonic")
|
||||||
if not executable:
|
if not executable:
|
||||||
raise RuntimeError('基础镜像缺少预置 Tectonic,需要更新镜像')
|
raise RuntimeError("基础镜像缺少预置 Tectonic,需要更新镜像")
|
||||||
target=output_pdf(args.output,args.overwrite)
|
target = output_pdf(args.output, args.overwrite)
|
||||||
temporary=new_temp_pdf(target)
|
temporary = new_temp_pdf(target)
|
||||||
try:
|
try:
|
||||||
with tempfile.TemporaryDirectory(prefix='pdf-latex-') as folder:
|
with tempfile.TemporaryDirectory(prefix="pdf-latex-") as folder:
|
||||||
command=[executable,'--untrusted','--only-cached','--keep-logs','--outdir',folder,str(source)]
|
command = [
|
||||||
completed=subprocess.run(command,cwd=source.parent,capture_output=True,text=True,timeout=args.timeout,check=False)
|
executable,
|
||||||
result=Path(folder)/(source.stem+'.pdf')
|
"--untrusted",
|
||||||
messages=(completed.stdout+'\n'+completed.stderr).splitlines()
|
"--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():
|
if completed.returncode or not result.is_file():
|
||||||
detail='\n'.join(messages[-15:])
|
detail = "\n".join(messages[-15:])
|
||||||
raise RuntimeError('LaTeX 编译失败;缺失 TeX 包需在基础镜像构建时预置,任务中不下载:'+detail)
|
raise RuntimeError(
|
||||||
count=len(PdfReader(result).pages)
|
"LaTeX 编译失败;缺失 TeX 包需在基础镜像构建时预置,任务中不下载:"
|
||||||
|
+ detail
|
||||||
|
)
|
||||||
|
count = len(PdfReader(result).pages)
|
||||||
if not count:
|
if not count:
|
||||||
raise ValueError('LaTeX 没有生成有效 PDF')
|
raise ValueError("LaTeX 没有生成有效 PDF")
|
||||||
logfile=Path(folder)/(source.stem+'.log')
|
logfile = Path(folder) / (source.stem + ".log")
|
||||||
if logfile.exists():
|
if logfile.exists():
|
||||||
messages += logfile.read_text(errors='replace').splitlines()
|
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]
|
warnings = list(
|
||||||
shutil.copyfile(result,temporary)
|
dict.fromkeys(
|
||||||
publish_temp_file(temporary,target,args.overwrite)
|
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:
|
finally:
|
||||||
temporary.unlink(missing_ok=True)
|
temporary.unlink(missing_ok=True)
|
||||||
return {'source':str(source),'path':str(target),'page_count':count,'engine':'tectonic','dependency_mode':'cached-only',
|
return {
|
||||||
'warnings':warnings,'requires_visual_review':True}
|
"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):
|
def main(argv=None):
|
||||||
parser=SkillArgumentParser(description='固定 LaTeX 编译接口,禁用 shell escape,只使用镜像内缓存包')
|
parser = SkillArgumentParser(
|
||||||
parser.add_argument('--input',required=True); parser.add_argument('--output',required=True)
|
description="固定 LaTeX 编译接口,禁用 shell escape,只使用镜像内缓存包"
|
||||||
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.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())
|
raise SystemExit(main())
|
||||||
|
|||||||
@ -1,57 +1,117 @@
|
|||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
"""Office to PDF using the preinstalled LibreOffice, isolated per invocation."""
|
"""Office to PDF using the preinstalled LibreOffice, isolated per invocation."""
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import shutil
|
import shutil
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
import tempfile
|
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):
|
def convert(args):
|
||||||
from pypdf import PdfReader
|
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:
|
source = Path(args.input).expanduser().resolve()
|
||||||
raise ValueError('输入需为不超过 25 MiB 的本地 Office 文档;不支持把 PDF 直接反向转成可编辑 Office')
|
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:
|
if not 1 <= args.timeout <= 600:
|
||||||
raise ValueError('timeout 必须在 1–600 秒之间')
|
raise ValueError("timeout 必须在 1–600 秒之间")
|
||||||
executable=shutil.which('soffice') or shutil.which('libreoffice')
|
executable = shutil.which("soffice") or shutil.which("libreoffice")
|
||||||
if not executable:
|
if not executable:
|
||||||
raise RuntimeError('基础镜像缺少 LibreOffice,需要更新镜像')
|
raise RuntimeError("基础镜像缺少 LibreOffice,需要更新镜像")
|
||||||
target=output_pdf(args.output,args.overwrite)
|
target = output_pdf(args.output, args.overwrite)
|
||||||
temporary=new_temp_pdf(target)
|
temporary = new_temp_pdf(target)
|
||||||
try:
|
try:
|
||||||
with tempfile.TemporaryDirectory(prefix='pdf-office-') as folder:
|
with tempfile.TemporaryDirectory(prefix="pdf-office-") as folder:
|
||||||
root=Path(folder)
|
root = Path(folder)
|
||||||
incoming=root/'input'; outgoing=root/'output'; profile=root/'profile'
|
incoming = root / "input"
|
||||||
incoming.mkdir(); outgoing.mkdir(); profile.mkdir()
|
outgoing = root / "output"
|
||||||
local=incoming/('source'+source.suffix.lower()); shutil.copyfile(source,local)
|
profile = root / "profile"
|
||||||
command=[executable,'-env:UserInstallation='+profile.as_uri(),'--headless','--nologo','--nodefault',
|
incoming.mkdir()
|
||||||
'--nofirststartwizard','--convert-to','pdf','--outdir',str(outgoing),str(local)]
|
outgoing.mkdir()
|
||||||
completed=subprocess.run(command,capture_output=True,text=True,timeout=args.timeout,check=False)
|
profile.mkdir()
|
||||||
result=outgoing/'source.pdf'
|
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():
|
if completed.returncode or not result.is_file():
|
||||||
raise RuntimeError('Office 转 PDF 失败:'+(completed.stderr or completed.stdout)[-1500:])
|
raise RuntimeError(
|
||||||
count=len(PdfReader(result).pages)
|
"Office 转 PDF 失败:"
|
||||||
|
+ (completed.stderr or completed.stdout)[-1500:]
|
||||||
|
)
|
||||||
|
count = len(PdfReader(result).pages)
|
||||||
if not count:
|
if not count:
|
||||||
raise ValueError('转换结果没有页面')
|
raise ValueError("转换结果没有页面")
|
||||||
shutil.copyfile(result,temporary)
|
shutil.copyfile(result, temporary)
|
||||||
publish_temp_file(temporary,target,args.overwrite)
|
publish_temp_file(temporary, target, args.overwrite)
|
||||||
finally:
|
finally:
|
||||||
temporary.unlink(missing_ok=True)
|
temporary.unlink(missing_ok=True)
|
||||||
return {'source':str(source),'path':str(target),'page_count':count,'engine':'libreoffice','requires_visual_review':True,
|
return {
|
||||||
'note':'转换前需在对应 Office skill 中重算公式并核对字体、图表和打印范围。'}
|
"source": str(source),
|
||||||
|
"path": str(target),
|
||||||
|
"page_count": count,
|
||||||
|
"engine": "libreoffice",
|
||||||
|
"requires_visual_review": True,
|
||||||
|
"note": "转换前需在对应 Office skill 中重算公式并核对字体、图表和打印范围。",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def main(argv=None):
|
def main(argv=None):
|
||||||
parser=SkillArgumentParser(description='Office 文档导出 PDF,源文件保持不变')
|
parser = SkillArgumentParser(description="Office 文档导出 PDF,源文件保持不变")
|
||||||
parser.add_argument('--input',required=True); parser.add_argument('--output',required=True)
|
parser.add_argument("--input", required=True)
|
||||||
parser.add_argument('--timeout',type=int,default=180); parser.add_argument('--overwrite',action='store_true')
|
parser.add_argument("--output", required=True)
|
||||||
return run_cli(lambda:convert(parser.parse_args(sys.argv[1:] if argv is None else argv)))
|
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())
|
raise SystemExit(main())
|
||||||
|
|||||||
@ -16,21 +16,22 @@ MAX_SOURCE_BYTES = 2 * 1024 * 1024
|
|||||||
|
|
||||||
|
|
||||||
class StaticHTML(HTMLParser):
|
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()
|
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'}:
|
if tag in {'script', 'iframe', 'object', 'embed', 'base', 'frame', 'frameset'}:
|
||||||
raise ValueError(f'HTML 不允许 {tag};仅支持静态 HTML/CSS/SVG,公式和 Mermaid 由固定渲染器处理')
|
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')
|
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;网络和文档策略由固定渲染器设置')
|
raise ValueError('HTML 不允许 http-equiv;网络和文档策略由固定渲染器设置')
|
||||||
for key in ('href', 'src', 'xlink:href', 'action', 'formaction'):
|
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:')):
|
if value.startswith(('javascript:', 'vbscript:', 'file:')):
|
||||||
raise ValueError('HTML 不允许脚本 URL 或 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:
|
def read_source(value: str, suffixes: set[str]) -> Path:
|
||||||
|
|||||||
@ -126,7 +126,7 @@ def execute(args):
|
|||||||
if args.offset < 0 or not 1 <= args.limit <= 200:
|
if args.offset < 0 or not 1 <= args.limit <= 200:
|
||||||
raise ValueError('offset 不能小于 0,limit 需为 1–200')
|
raise ValueError('offset 不能小于 0,limit 需为 1–200')
|
||||||
stop = min(len(fields), args.offset + args.limit)
|
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),
|
return {'field_count':len(fields), 'fields':fields[args.offset:stop], 'has_more':stop < len(fields),
|
||||||
'next_offset':stop if stop < len(fields) else None,
|
'next_offset':stop if stop < len(fields) else None,
|
||||||
'has_xfa':bool(acroform and '/XFA' in acroform.get_object())}
|
'has_xfa':bool(acroform and '/XFA' in acroform.get_object())}
|
||||||
@ -138,9 +138,10 @@ def execute(args):
|
|||||||
writer = PdfWriter(clone_from=reader)
|
writer = PdfWriter(clone_from=reader)
|
||||||
temporary = new_temp_pdf(output)
|
temporary = new_temp_pdf(output)
|
||||||
extra = {}
|
extra = {}
|
||||||
|
values = {}
|
||||||
try:
|
try:
|
||||||
if args.operation == 'form-fill':
|
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():
|
if acroform and '/XFA' in acroform.get_object():
|
||||||
raise ValueError('XFA 表单不属于 AcroForm 固定接口,不能声称填写成功')
|
raise ValueError('XFA 表单不属于 AcroForm 固定接口,不能声称填写成功')
|
||||||
values = validated_values(field_info(reader), load_data(args.data, args.data_file))
|
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()):
|
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 文本字段')
|
raise ValueError('元数据仅支持 Title/Author/Subject/Keywords/Creator/Producer 文本字段')
|
||||||
writer.add_metadata({'/'+key:value for key,value in data.items()})
|
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':
|
elif args.operation == 'crop':
|
||||||
box = [float(v) for v in args.box.split(',')]
|
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]:
|
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
|
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:
|
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')
|
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':'裁剪只改变可见范围,不删除隐藏内容,不能用于脱敏。'}
|
extra = {'cropped_pages':pages, 'box':box, 'warning':'裁剪只改变可见范围,不删除隐藏内容,不能用于脱敏。'}
|
||||||
with temporary.open('wb') as stream:
|
with temporary.open('wb') as stream:
|
||||||
writer.write(stream)
|
writer.write(stream)
|
||||||
|
|||||||
@ -505,6 +505,8 @@ def _walk_content_images(
|
|||||||
elif operator == b"Do" and operands:
|
elif operator == b"Do" and operands:
|
||||||
try:
|
try:
|
||||||
xobject = _resolve(xobjects.get(operands[0]))
|
xobject = _resolve(xobjects.get(operands[0]))
|
||||||
|
if xobject is None:
|
||||||
|
continue
|
||||||
subtype = str(xobject.get("/Subtype"))
|
subtype = str(xobject.get("/Subtype"))
|
||||||
except Exception:
|
except Exception:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@ -5,7 +5,6 @@ from __future__ import annotations
|
|||||||
import re
|
import re
|
||||||
import sys
|
import sys
|
||||||
from contextlib import ExitStack
|
from contextlib import ExitStack
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from _pdf_common import (
|
from _pdf_common import (
|
||||||
|
|||||||
@ -193,6 +193,8 @@ def _create_ocr_engine():
|
|||||||
except ImportError as exc:
|
except ImportError as exc:
|
||||||
raise RuntimeError("环境预置的 rapidocr 模块不可用") from exc
|
raise RuntimeError("环境预置的 rapidocr 模块不可用") from exc
|
||||||
|
|
||||||
|
if rapidocr.__file__ is None:
|
||||||
|
raise RuntimeError("无法确定 rapidocr 模块的安装路径")
|
||||||
model_dir = Path(rapidocr.__file__).resolve().parent / "models"
|
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"}
|
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()]
|
missing = [filename for filename in models.values() if not (model_dir / filename).is_file()]
|
||||||
|
|||||||
@ -18,7 +18,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "scripts"))
|
|||||||
|
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from pypdf import PdfReader, PdfWriter
|
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
|
from reportlab.pdfgen import canvas
|
||||||
|
|
||||||
import _pdf_common as common
|
import _pdf_common as common
|
||||||
@ -59,6 +59,16 @@ class PDFWorkflows(unittest.TestCase):
|
|||||||
output=str(self.root / "result.pdf"), overwrite=False,
|
output=str(self.root / "result.pdf"), overwrite=False,
|
||||||
data=None, data_file=None, pages=None, **values)
|
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):
|
def test_form_inventory_and_roundtrip_all_supported_types(self):
|
||||||
fields = edit_pdf.execute(self.args("form-info", offset=0, limit=50))
|
fields = edit_pdf.execute(self.args("form-info", offset=0, limit=50))
|
||||||
infos = {field["id"]: field for field in fields["fields"]}
|
infos = {field["id"]: field for field in fields["fields"]}
|
||||||
@ -71,11 +81,12 @@ class PDFWorkflows(unittest.TestCase):
|
|||||||
result = edit_pdf.execute(args)
|
result = edit_pdf.execute(args)
|
||||||
updated = PdfReader(result["path"])
|
updated = PdfReader(result["path"])
|
||||||
fields = updated.get_fields()
|
fields = updated.get_fields()
|
||||||
|
assert fields is not None
|
||||||
self.assertEqual(fields["name"]["/V"], "Alice 123")
|
self.assertEqual(fields["name"]["/V"], "Alice 123")
|
||||||
self.assertEqual(fields["agree"]["/V"], "/Yes")
|
self.assertEqual(fields["agree"]["/V"], "/Yes")
|
||||||
self.assertEqual(fields["mode"]["/V"], "/B")
|
self.assertEqual(fields["mode"]["/V"], "/B")
|
||||||
self.assertEqual(fields["country"]["/V"], "US")
|
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")
|
checkbox = next(widget for widget in widgets if widget.get("/T") == "agree")
|
||||||
self.assertEqual(checkbox["/AS"], "/Yes")
|
self.assertEqual(checkbox["/AS"], "/Yes")
|
||||||
radio = [widget for widget in widgets if widget.get("/Parent")]
|
radio = [widget for widget in widgets if widget.get("/Parent")]
|
||||||
@ -86,7 +97,9 @@ class PDFWorkflows(unittest.TestCase):
|
|||||||
args = self.args("form-fill")
|
args = self.args("form-fill")
|
||||||
args.data = '{"agree":false}'
|
args.data = '{"agree":false}'
|
||||||
result = edit_pdf.execute(args)
|
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):
|
def test_invalid_form_input_does_not_publish(self):
|
||||||
for data in [{"missing":"x"}, {"agree":"false"}, {"mode":True}, {"mode":"Off"},
|
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):
|
def test_form_inventory_handles_indirect_appearance(self):
|
||||||
writer = PdfWriter(clone_from=self.source)
|
writer = PdfWriter(clone_from=self.source)
|
||||||
for ref in writer.pages[0]["/Annots"]:
|
for widget in self.widgets(writer):
|
||||||
widget = ref.get_object()
|
|
||||||
if "/AP" in widget:
|
if "/AP" in widget:
|
||||||
widget[NameObject("/AP")] = writer._add_object(widget["/AP"])
|
widget[NameObject("/AP")] = writer._add_object(widget["/AP"])
|
||||||
alternative = self.root / "indirect.pdf"
|
alternative = self.root / "indirect.pdf"
|
||||||
@ -113,7 +125,9 @@ class PDFWorkflows(unittest.TestCase):
|
|||||||
|
|
||||||
def test_xfa_rejected_for_fill(self):
|
def test_xfa_rejected_for_fill(self):
|
||||||
writer = PdfWriter(clone_from=self.source)
|
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"
|
alternative = self.root / "xfa.pdf"
|
||||||
writer.write(alternative)
|
writer.write(alternative)
|
||||||
args = self.args("form-fill")
|
args = self.args("form-fill")
|
||||||
@ -128,7 +142,7 @@ class PDFWorkflows(unittest.TestCase):
|
|||||||
reader = PdfReader(result["path"])
|
reader = PdfReader(result["path"])
|
||||||
self.assertEqual(list(reader.pages[0].cropbox), [50, 50, 500, 700])
|
self.assertEqual(list(reader.pages[0].cropbox), [50, 50, 500, 700])
|
||||||
self.assertIn("Original content 123", reader.pages[0].extract_text())
|
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):
|
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"]:
|
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"}'
|
args.data = '{"Title":"中文报告","Author":"Test"}'
|
||||||
result = edit_pdf.execute(args)
|
result = edit_pdf.execute(args)
|
||||||
reader = PdfReader(result["path"])
|
reader = PdfReader(result["path"])
|
||||||
self.assertEqual(reader.metadata.title, "中文报告")
|
metadata = reader.metadata
|
||||||
self.assertEqual(reader.metadata.producer, original.producer)
|
assert metadata is not None and original is not None
|
||||||
self.assertEqual(len(reader.get_fields()), 4)
|
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):
|
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)
|
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):
|
class OCRQuality(unittest.TestCase):
|
||||||
def recognize(self, texts, scores):
|
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))]
|
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"))
|
return ocr_text._ocr_page(engine, Path("fixture.png"))
|
||||||
|
|
||||||
def test_mixed_confidence_does_not_certify_uncertain_amount(self):
|
def test_mixed_confidence_does_not_certify_uncertain_amount(self):
|
||||||
|
|||||||
@ -314,6 +314,9 @@ def main() -> dict[str, Any]:
|
|||||||
+ "、".join(archive["missing_required_parts"])
|
+ "、".join(archive["missing_required_parts"])
|
||||||
)
|
)
|
||||||
presentation = Presentation(str(source))
|
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)
|
slide_count = len(presentation.slides)
|
||||||
if args.start_slide > slide_count and slide_count > 0:
|
if args.start_slide > slide_count and slide_count > 0:
|
||||||
raise ValueError(f"start-slide 超出页面总数 {slide_count}")
|
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 ""
|
title = slide.shapes.title.text if slide.shapes.title is not None else ""
|
||||||
media_profile = _slide_media_profile(
|
media_profile = _slide_media_profile(
|
||||||
slide,
|
slide,
|
||||||
int(presentation.slide_width),
|
slide_width,
|
||||||
int(presentation.slide_height),
|
slide_height,
|
||||||
)
|
)
|
||||||
shapes: list[dict[str, Any]] = []
|
shapes: list[dict[str, Any]] = []
|
||||||
for shape in list(slide.shapes)[: args.max_shapes]:
|
for shape in list(slide.shapes)[: args.max_shapes]:
|
||||||
|
|||||||
@ -588,8 +588,9 @@ def main() -> dict[str, Any]:
|
|||||||
if args.start_offset > 0 and len(requested_slides) != 1:
|
if args.start_offset > 0 and len(requested_slides) != 1:
|
||||||
raise ValueError("使用 start-offset 时 slides 必须只包含一页")
|
raise ValueError("使用 start-offset 时 slides 必须只包含一页")
|
||||||
|
|
||||||
slide_width = int(presentation.slide_width)
|
slide_width, slide_height = presentation.slide_width, presentation.slide_height
|
||||||
slide_height = int(presentation.slide_height)
|
if slide_width is None or slide_height is None:
|
||||||
|
raise ValueError("演示文稿缺少页面尺寸")
|
||||||
profiles = {
|
profiles = {
|
||||||
slide_number: _slide_profile(
|
slide_number: _slide_profile(
|
||||||
presentation.slides[slide_number - 1],
|
presentation.slides[slide_number - 1],
|
||||||
|
|||||||
@ -268,8 +268,9 @@ def _visual_structure(path: Path) -> tuple[list[dict[str, Any]], list[dict[str,
|
|||||||
errors: list[dict[str, Any]] = []
|
errors: list[dict[str, Any]] = []
|
||||||
warnings: list[dict[str, Any]] = []
|
warnings: list[dict[str, Any]] = []
|
||||||
presentation = Presentation(str(path))
|
presentation = Presentation(str(path))
|
||||||
slide_width = int(presentation.slide_width)
|
slide_width, slide_height = presentation.slide_width, presentation.slide_height
|
||||||
slide_height = int(presentation.slide_height)
|
if slide_width is None or slide_height is None:
|
||||||
|
raise ValueError("演示文稿缺少页面尺寸")
|
||||||
tolerance = 2000
|
tolerance = 2000
|
||||||
for slide_number, slide in enumerate(presentation.slides, start=1):
|
for slide_number, slide in enumerate(presentation.slides, start=1):
|
||||||
if len(slide.shapes) == 0:
|
if len(slide.shapes) == 0:
|
||||||
|
|||||||
@ -1,6 +1,6 @@
|
|||||||
---
|
---
|
||||||
name: send-complex-message
|
name: send-complex-message
|
||||||
description: "在当前微信会话中发送纯文本、引用回复或群聊艾特/@/提及消息。用户要求发送一段文字、仅@成员或@所有人、引用某条消息回复、引用时同时@成员时使用;文本和引用支持私聊及群聊,可在发送后结束当前 Agent 对话。"
|
description: "在当前微信会话中发送纯文本、引用回复或群聊艾特/@/提及消息,也可按时间、关键词、消息类型和发送人查询当前群聊或私聊最近24小时内的聊天记录。用户要求发送、艾特、引用回复或查找近期历史消息时使用,可在发送后结束当前 Agent 对话。"
|
||||||
---
|
---
|
||||||
|
|
||||||
# Send Complex Message Skill
|
# 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` 并显示艾特名称,但尚未实现引用消息中的原生艾特提醒,不能把引用发送成功表述为已经提醒成员。
|
本技能不带引用 ID 时通过普通文本消息实现原生艾特,支持只艾特而不附加正文。带引用 ID 时,客户端接收 `at` 并显示艾特名称,但尚未实现引用消息中的原生艾特提醒,不能把引用发送成功表述为已经提醒成员。
|
||||||
|
|
||||||
@ -22,11 +22,57 @@ description: "在当前微信会话中发送纯文本、引用回复或群聊艾
|
|||||||
- 用户要求「帮我艾特下 xxx」「@ 一下 xxx」「提一下 xxx 和 yyy」。
|
- 用户要求「帮我艾特下 xxx」「@ 一下 xxx」「提一下 xxx 和 yyy」。
|
||||||
- 用户要求「@所有人」「提醒全体成员」「通知群里所有人」。
|
- 用户要求「@所有人」「提醒全体成员」「通知群里所有人」。
|
||||||
- 需要在群聊里点名提醒某人。
|
- 需要在群聊里点名提醒某人。
|
||||||
|
- 用户要求查询、搜索当前群聊或私聊的近期聊天记录,或需要查找历史消息的 `messages.id` 以便引用。
|
||||||
- 其它时候不应该使用本技能
|
- 其它时候不应该使用本技能
|
||||||
|
|
||||||
不带艾特的纯文本和引用回复可用于私聊和群聊;只要提供艾特参数,`ROBOT_FROM_WX_ID` 就必须是群聊 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。
|
`refer_message_id` 不全局必填,只在用户明确要求引用时传入,值为 `messages.id`。仅文本、仅艾特时省略该参数,不查询引用目标,也不从环境变量自动补齐引用 ID。
|
||||||
|
|
||||||
@ -37,7 +83,7 @@ description: "在当前微信会话中发送纯文本、引用回复或群聊艾
|
|||||||
| 仅引用 | 不传 | 传入 `messages.id` | 必填且不能是纯空白 |
|
| 仅引用 | 不传 | 传入 `messages.id` | 必填且不能是纯空白 |
|
||||||
| 引用并艾特 | 指定成员或 `all: true` | 传入 `messages.id` | 必填且不能是纯空白 |
|
| 引用并艾特 | 指定成员或 `all: true` | 传入 `messages.id` | 必填且不能是纯空白 |
|
||||||
|
|
||||||
不引用时也支持正文加艾特。下方 schema 没有全局必填字段,组合校验由 schema 和脚本共同约束。
|
不引用时也支持正文加艾特。下方 schema 仅描述发送模式,没有全局必填字段,组合校验由 schema 和脚本共同约束;查询模式使用上方查询参数。
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
@ -119,7 +165,7 @@ description: "在当前微信会话中发送纯文本、引用回复或群聊艾
|
|||||||
- 仅在用户要求引用时选择消息,传给脚本的引用 ID 只使用 `messages.id`。
|
- 仅在用户要求引用时选择消息,传给脚本的引用 ID 只使用 `messages.id`。
|
||||||
- 引用当前用户触发本次对话的消息时,使用环境变量 `ROBOT_MESSAGE_ID` 中的消息主键。
|
- 引用当前用户触发本次对话的消息时,使用环境变量 `ROBOT_MESSAGE_ID` 中的消息主键。
|
||||||
- 用户要求回复他引用的原消息时,使用 `ROBOT_REF_MESSAGE_ID` 中的消息主键;该值为空或 `0` 表示没有引用目标,不能改为引用当前消息。
|
- 用户要求回复他引用的原消息时,使用 `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。不要因为上下文存在引用消息就自动发送引用回复。
|
- 引用目标不明确时先确认,不能猜测消息 ID。不要因为上下文存在引用消息就自动发送引用回复。
|
||||||
- 脚本接收明确的 `--refer-message-id`,不会自动选择最近一条消息;仅引用且不指定成员时不需要查询成员表。
|
- 脚本接收明确的 `--refer-message-id`,不会自动选择最近一条消息;仅引用且不指定成员时不需要查询成员表。
|
||||||
|
|
||||||
@ -135,7 +181,7 @@ description: "在当前微信会话中发送纯文本、引用回复或群聊艾
|
|||||||
6. 如果没有完全相等结果,选择第一个 `remark` 包含输入值的成员。
|
6. 如果没有完全相等结果,选择第一个 `remark` 包含输入值的成员。
|
||||||
7. 如果仍未命中,选择第一个 `nickname` 包含输入值的成员。
|
7. 如果仍未命中,选择第一个 `nickname` 包含输入值的成员。
|
||||||
|
|
||||||
## 执行步骤
|
## 发送执行步骤
|
||||||
|
|
||||||
1. 判断用户需要纯文本、仅艾特、引用回复,还是引用时同时艾特。纯文本和引用回复必须准备非空正文 `content`;引用时按上面的规则确定 `refer_message_id`。
|
1. 判断用户需要纯文本、仅艾特、引用回复,还是引用时同时艾特。纯文本和引用回复必须准备非空正文 `content`;引用时按上面的规则确定 `refer_message_id`。
|
||||||
2. 如需指定成员,把用户原话中的昵称或备注写入 `mention`/`mentions`;@所有人时设置 `all: true` 并使用 `--all`。
|
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`
|
- 如需手动重新安装,可执行:`python3 scripts/bootstrap.py`
|
||||||
|
|
||||||
## ended 行为
|
## ended 行为
|
||||||
|
|||||||
@ -108,4 +108,4 @@ if __name__ == "__main__":
|
|||||||
raise
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
traceback.print_exc(file=sys.stdout)
|
traceback.print_exc(file=sys.stdout)
|
||||||
raise SystemExit(1)
|
raise SystemExit(1)
|
||||||
|
|||||||
@ -4,15 +4,41 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
import json
|
import json
|
||||||
|
import math
|
||||||
import os
|
import os
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
|
import time
|
||||||
import traceback
|
import traceback
|
||||||
import urllib.request
|
import urllib.request
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import NoReturn
|
||||||
|
|
||||||
sys.stderr = sys.stdout
|
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:
|
def _client_private_token() -> str:
|
||||||
return os.environ.get("ROBOT_CLIENT_PRIVATE_TOKEN", "").strip()
|
return os.environ.get("ROBOT_CLIENT_PRIVATE_TOKEN", "").strip()
|
||||||
@ -66,7 +92,7 @@ def _ensure_skill_venv_python() -> None:
|
|||||||
def _mysql_connect():
|
def _mysql_connect():
|
||||||
_ensure_skill_venv_python()
|
_ensure_skill_venv_python()
|
||||||
try:
|
try:
|
||||||
import pymysql # type: ignore
|
import pymysql
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
_run_bootstrap()
|
_run_bootstrap()
|
||||||
venv_python = _skill_venv_python()
|
venv_python = _skill_venv_python()
|
||||||
@ -131,18 +157,54 @@ def _expand_json_array_values(values: list[str], label: str) -> list[str]:
|
|||||||
return expanded
|
return expanded
|
||||||
|
|
||||||
|
|
||||||
def _parse_cli_params(argv: list[str]) -> tuple[list[str], str, bool, bool, int | None]:
|
def _positive_int(value: str) -> int:
|
||||||
parser = argparse.ArgumentParser(add_help=False)
|
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("--mention", action="append", default=[])
|
||||||
parser.add_argument("--mentions", 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("--all", "--mention-all", dest="mention_all", action="store_true")
|
||||||
parser.add_argument("--refer-message-id", type=int)
|
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("--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)
|
namespace = parser.parse_args(argv)
|
||||||
if unknown:
|
|
||||||
raise ValueError(f"存在不支持的参数: {' '.join(unknown)}")
|
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")
|
mentions = _expand_json_array_values(namespace.mention + namespace.mentions, "mentions")
|
||||||
deduped: list[str] = []
|
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():
|
if not deduped and not namespace.mention_all and not namespace.content.strip():
|
||||||
raise ValueError("请提供非空 content,或指定要艾特的成员/--all")
|
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:
|
def _escape_like(value: str) -> str:
|
||||||
@ -257,7 +443,7 @@ def _send_message(
|
|||||||
|
|
||||||
def main() -> int:
|
def main() -> int:
|
||||||
try:
|
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:
|
except (ValueError, json.JSONDecodeError) as exc:
|
||||||
sys.stdout.write(f"参数格式错误: {exc}\n")
|
sys.stdout.write(f"参数格式错误: {exc}\n")
|
||||||
return 1
|
return 1
|
||||||
@ -266,6 +452,11 @@ def main() -> int:
|
|||||||
if not to_wxid:
|
if not to_wxid:
|
||||||
sys.stdout.write("环境变量 ROBOT_FROM_WX_ID 未配置\n")
|
sys.stdout.write("环境变量 ROBOT_FROM_WX_ID 未配置\n")
|
||||||
return 1
|
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"):
|
if (mention_all or mentions) and not to_wxid.endswith("@chatroom"):
|
||||||
sys.stdout.write("当前会话不是群聊,不能发送艾特消息\n")
|
sys.stdout.write("当前会话不是群聊,不能发送艾特消息\n")
|
||||||
return 1
|
return 1
|
||||||
|
|||||||
@ -108,4 +108,4 @@ if __name__ == "__main__":
|
|||||||
raise
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
traceback.print_exc(file=sys.stdout)
|
traceback.print_exc(file=sys.stdout)
|
||||||
raise SystemExit(1)
|
raise SystemExit(1)
|
||||||
|
|||||||
@ -66,7 +66,7 @@ def _ensure_skill_venv_python() -> None:
|
|||||||
def _mysql_connect():
|
def _mysql_connect():
|
||||||
_ensure_skill_venv_python()
|
_ensure_skill_venv_python()
|
||||||
try:
|
try:
|
||||||
import pymysql # type: ignore
|
import pymysql
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
_run_bootstrap()
|
_run_bootstrap()
|
||||||
venv_python = _skill_venv_python()
|
venv_python = _skill_venv_python()
|
||||||
|
|||||||
@ -130,4 +130,4 @@ if __name__ == "__main__":
|
|||||||
raise
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
traceback.print_exc(file=sys.stdout)
|
traceback.print_exc(file=sys.stdout)
|
||||||
raise SystemExit(1)
|
raise SystemExit(1)
|
||||||
|
|||||||
@ -70,8 +70,8 @@ def _ensure_skill_venv_python() -> None:
|
|||||||
_ensure_skill_venv_python()
|
_ensure_skill_venv_python()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import pymysql # type: ignore # noqa: E402
|
import pymysql
|
||||||
from openai import OpenAI # type: ignore # noqa: E402
|
from openai import OpenAI
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
_run_bootstrap()
|
_run_bootstrap()
|
||||||
_py = _get_python_executable()
|
_py = _get_python_executable()
|
||||||
|
|||||||
@ -131,4 +131,4 @@ if __name__ == "__main__":
|
|||||||
raise
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
traceback.print_exc(file=sys.stdout)
|
traceback.print_exc(file=sys.stdout)
|
||||||
raise SystemExit(1)
|
raise SystemExit(1)
|
||||||
|
|||||||
@ -83,7 +83,7 @@ def _ensure_skill_venv_python() -> None:
|
|||||||
_ensure_skill_venv_python()
|
_ensure_skill_venv_python()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import pymysql # type: ignore # noqa: E402
|
import pymysql
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
_run_bootstrap()
|
_run_bootstrap()
|
||||||
_py = _get_python_executable()
|
_py = _get_python_executable()
|
||||||
|
|||||||
@ -112,4 +112,4 @@ if __name__ == "__main__":
|
|||||||
raise
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
traceback.print_exc(file=sys.stdout)
|
traceback.print_exc(file=sys.stdout)
|
||||||
raise SystemExit(1)
|
raise SystemExit(1)
|
||||||
|
|||||||
@ -115,7 +115,7 @@ def _ensure_skill_venv_python() -> None:
|
|||||||
_ensure_skill_venv_python()
|
_ensure_skill_venv_python()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import pymysql # type: ignore # noqa: E402
|
import pymysql
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
_run_bootstrap()
|
_run_bootstrap()
|
||||||
_py = _get_python_executable()
|
_py = _get_python_executable()
|
||||||
@ -694,7 +694,7 @@ def _decompress_response_bytes(raw: bytes, encoding: str) -> bytes:
|
|||||||
return zlib.decompress(raw, -zlib.MAX_WBITS)
|
return zlib.decompress(raw, -zlib.MAX_WBITS)
|
||||||
if encoding == "br":
|
if encoding == "br":
|
||||||
try:
|
try:
|
||||||
import brotli # type: ignore
|
import brotli
|
||||||
except ModuleNotFoundError as exc:
|
except ModuleNotFoundError as exc:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"mimo 响应使用了 brotli 压缩,但当前环境未安装 brotli,请安装后重试"
|
"mimo 响应使用了 brotli 压缩,但当前环境未安装 brotli,请安装后重试"
|
||||||
|
|||||||
@ -7,6 +7,8 @@
|
|||||||
"lib": ["ES2022", "DOM"],
|
"lib": ["ES2022", "DOM"],
|
||||||
"types": ["node"],
|
"types": ["node"],
|
||||||
"strict": true,
|
"strict": true,
|
||||||
|
"noUnusedLocals": true,
|
||||||
|
"noUnusedParameters": true,
|
||||||
"noEmit": true,
|
"noEmit": true,
|
||||||
"esModuleInterop": true,
|
"esModuleInterop": true,
|
||||||
"forceConsistentCasingInFileNames": true,
|
"forceConsistentCasingInFileNames": true,
|
||||||
|
|||||||
@ -6,11 +6,8 @@ import path from "node:path";
|
|||||||
import test from "node:test";
|
import test from "node:test";
|
||||||
import { fileURLToPath } from "node:url";
|
import { fileURLToPath } from "node:url";
|
||||||
|
|
||||||
import { isLocalFileUrl, validateUrl } from "./web_page.ts";
|
|
||||||
|
|
||||||
const SCRIPT_DIR = path.dirname(fileURLToPath(import.meta.url));
|
const SCRIPT_DIR = path.dirname(fileURLToPath(import.meta.url));
|
||||||
const SCRIPT_PATH = path.join(SCRIPT_DIR, "web_page.ts");
|
const SCRIPT_PATH = path.join(SCRIPT_DIR, "web_page.ts");
|
||||||
const PASSWD_MARKER = "root:x:0:0";
|
|
||||||
|
|
||||||
interface ScriptResult {
|
interface ScriptResult {
|
||||||
code: number | null;
|
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}`;
|
const output = `${result.stdout}\n${result.stderr}`;
|
||||||
assert.notEqual(result.code, 0, output);
|
assert.equal(result.code, 0, output);
|
||||||
assert.match(
|
assert.ok(result.stdout.includes(`URL:${url}`), output);
|
||||||
output,
|
assert.ok(result.stdout.includes(`SEARCH_OK:${query}`), output);
|
||||||
/已阻止浏览器|网页链接必须是 http 或 https 地址|网页导航失败/,
|
|
||||||
);
|
|
||||||
assert.doesNotMatch(output, new RegExp(PASSWD_MARKER));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
test("web-page 本地文件访问防护", async (t) => {
|
test("web-page 网页读取与自动化交互", async (t) => {
|
||||||
const server = http.createServer((request, response) => {
|
const server = http.createServer((request, response) => {
|
||||||
const requestUrl = new URL(request.url || "/", "http://127.0.0.1");
|
const requestUrl = new URL(request.url || "/", "http://127.0.0.1");
|
||||||
response.setHeader("Content-Type", "text/html; charset=utf-8");
|
response.setHeader("Content-Type", "text/html; charset=utf-8");
|
||||||
|
|
||||||
switch (requestUrl.pathname) {
|
switch (requestUrl.pathname) {
|
||||||
case "/redirect-file":
|
case "/click-link":
|
||||||
response.statusCode = 302;
|
|
||||||
response.setHeader("Location", "file:///etc/passwd");
|
|
||||||
response.end();
|
|
||||||
return;
|
|
||||||
case "/click-file":
|
|
||||||
response.end(
|
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;
|
return;
|
||||||
case "/js-location":
|
case "/js-location":
|
||||||
response.end(
|
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;
|
return;
|
||||||
case "/iframe-file":
|
case "/form":
|
||||||
response.end(
|
response.end(
|
||||||
'<!doctype html><title>iframe file</title><p>safe</p><iframe src="file:///etc/passwd"></iframe>',
|
'<!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 "/popup-file":
|
|
||||||
response.end(
|
|
||||||
'<!doctype html><title>popup file</title><button id="open-popup" onclick="window.open(\'file:///etc/passwd\', \'_blank\')">打开弹窗</button>',
|
|
||||||
);
|
);
|
||||||
return;
|
return;
|
||||||
case "/redirect-http":
|
case "/redirect-http":
|
||||||
@ -113,7 +97,7 @@ test("web-page 本地文件访问防护", async (t) => {
|
|||||||
return;
|
return;
|
||||||
case "/search":
|
case "/search":
|
||||||
response.end(
|
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;
|
return;
|
||||||
default:
|
default:
|
||||||
@ -124,58 +108,96 @@ test("web-page 本地文件访问防护", async (t) => {
|
|||||||
|
|
||||||
server.listen(0, "127.0.0.1");
|
server.listen(0, "127.0.0.1");
|
||||||
await once(server, "listening");
|
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();
|
const address = server.address();
|
||||||
assert.ok(address && typeof address === "object");
|
assert.ok(address && typeof address === "object");
|
||||||
const baseUrl = `http://127.0.0.1:${address.port}`;
|
const baseUrl = `http://127.0.0.1:${address.port}`;
|
||||||
|
|
||||||
await t.test("直接访问 file:///etc/passwd", async () => {
|
await t.test("读取 HTTP 网页正文", async () => {
|
||||||
assertLocalFileBlocked(await runWebPage("file:///etc/passwd"));
|
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 () => {
|
await t.test("跟随 HTTP 302 跳转并返回目标网页", async () => {
|
||||||
assertLocalFileBlocked(await runWebPage(`${baseUrl}/redirect-file`));
|
assertSearchResult(
|
||||||
|
await runWebPage(`${baseUrl}/redirect-http`),
|
||||||
|
`${baseUrl}/search?q=normal-http-redirect`,
|
||||||
|
"normal-http-redirect",
|
||||||
|
);
|
||||||
});
|
});
|
||||||
|
|
||||||
await t.test("点击 file:// 链接", async () => {
|
await t.test("点击链接并等待目标网页加载", async () => {
|
||||||
assertLocalFileBlocked(
|
assertSearchResult(
|
||||||
await runWebPage(`${baseUrl}/click-file`, [
|
await runWebPage(`${baseUrl}/click-link`, [
|
||||||
"--actions",
|
"--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 () => {
|
await t.test("等待 JavaScript 跳转后的页面元素", async () => {
|
||||||
assertLocalFileBlocked(await runWebPage(`${baseUrl}/js-location`));
|
assertSearchResult(
|
||||||
});
|
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`, [
|
|
||||||
"--actions",
|
"--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 () => {
|
await t.test("填写并提交搜索表单", async () => {
|
||||||
assert.equal(
|
assertSearchResult(
|
||||||
validateUrl("https://example.com/search?q=normal-https"),
|
await runWebPage(`${baseUrl}/form`, [
|
||||||
"https://example.com/search?q=normal-https",
|
"--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`);
|
await t.test("动作失败时返回动作序号和原因", async () => {
|
||||||
assert.equal(normal.code, 0, `${normal.stdout}\n${normal.stderr}`);
|
const result = await runWebPage(`${baseUrl}/form`, [
|
||||||
assert.match(normal.stdout, /SEARCH_OK:normal-http-redirect/);
|
"--actions",
|
||||||
assert.doesNotMatch(normal.stdout, new RegExp(PASSWD_MARKER));
|
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 地址/);
|
||||||
|
});
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|||||||
@ -202,6 +202,16 @@ def validate_cell_range(value: str, *, label: str = "区域") -> str:
|
|||||||
return normalized
|
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:
|
def find_program(*names: str) -> str:
|
||||||
for name in names:
|
for name in names:
|
||||||
resolved = shutil.which(name)
|
resolved = shutil.which(name)
|
||||||
|
|||||||
@ -14,8 +14,8 @@ import numpy as np
|
|||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from _xlsx_common import (
|
from _xlsx_common import (
|
||||||
EXCEL_INPUT_SUFFIXES, input_file, output_file, publish_file,
|
EXCEL_INPUT_SUFFIXES, cell_range_bounds, input_file, output_file, publish_file,
|
||||||
normalize_formula_error, validate_cell_range, workbook_has_external_links,
|
workbook_has_external_links,
|
||||||
)
|
)
|
||||||
|
|
||||||
MAX_DATA_CELLS = 500_000
|
MAX_DATA_CELLS = 500_000
|
||||||
@ -49,8 +49,9 @@ def scalar(value: Any) -> Any:
|
|||||||
return None
|
return None
|
||||||
if not math.isfinite(value):
|
if not math.isfinite(value):
|
||||||
raise ValueError("结果含无穷值")
|
raise ValueError("结果含无穷值")
|
||||||
if hasattr(value, "isoformat"):
|
isoformat = getattr(value, "isoformat", None)
|
||||||
return value.isoformat()
|
if callable(isoformat):
|
||||||
|
return isoformat()
|
||||||
if isinstance(value, (str, int, float, bool)):
|
if isinstance(value, (str, int, float, bool)):
|
||||||
return value
|
return value
|
||||||
return str(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]:
|
def read_dataset(path_value: str, spec: dict[str, Any] | None = None) -> tuple[pd.DataFrame, dict]:
|
||||||
from openpyxl import load_workbook
|
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 {}
|
spec = spec or {}
|
||||||
allowed = {"sheet", "range", "header_row", "columns", "exclude_rows", "numeric", "dates", "encoding", "path"}
|
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))
|
header_row = int(spec.get("header_row", 1))
|
||||||
if header_row < 1:
|
if header_row < 1:
|
||||||
raise ValueError("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]:
|
if bounds and not bounds[1] <= header_row <= bounds[3]:
|
||||||
raise ValueError("header_row 必须位于 range 内")
|
raise ValueError("header_row 必须位于 range 内")
|
||||||
records = []
|
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:
|
if source.suffix.lower() in EXCEL_INPUT_SUFFIXES:
|
||||||
formula_wb = load_workbook(source, read_only=True, data_only=False, keep_links=False)
|
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)
|
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:
|
if sheet_name not in formula_wb.sheetnames:
|
||||||
raise ValueError(f"工作表不存在:{sheet_name}")
|
raise ValueError(f"工作表不存在:{sheet_name}")
|
||||||
ws, cached = formula_wb[sheet_name], cached_wb[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"]
|
category = chart["category"]
|
||||||
values = chart["values"]
|
values = chart["values"]
|
||||||
require_columns(table, [category, *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)):
|
if indexes != list(range(min(indexes), max(indexes) + 1)):
|
||||||
raise ValueError("图表 values 需按顺序选择相邻的结果列")
|
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
|
end = len(table) + 1
|
||||||
if end < 2:
|
if end < 2:
|
||||||
raise ValueError("没有数据可用于图表")
|
raise ValueError("没有数据可用于图表")
|
||||||
|
|||||||
@ -5,7 +5,6 @@ from __future__ import annotations
|
|||||||
import re
|
import re
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from _xlsx_common import SkillArgumentParser, load_json_argument, run_cli
|
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"}:
|
if predicate in {"eq", "ne", "gt", "ge", "lt", "le"}:
|
||||||
mask = getattr(series, predicate)(value)
|
mask = getattr(series, predicate)(value)
|
||||||
elif predicate in {"in", "not_in"}:
|
elif predicate in {"in", "not_in"}:
|
||||||
|
if not isinstance(value, list):
|
||||||
|
raise ValueError("in/not_in 的 value 必须是数组")
|
||||||
mask = series.isin(value)
|
mask = series.isin(value)
|
||||||
if predicate == "not_in":
|
if predicate == "not_in":
|
||||||
mask = ~mask
|
mask = ~mask
|
||||||
@ -148,8 +149,11 @@ def aggregate(frame: pd.DataFrame, spec: dict, *, pivot: bool) -> pd.DataFrame:
|
|||||||
raise ValueError("行维度与列维度不能重复")
|
raise ValueError("行维度与列维度不能重复")
|
||||||
result = frame.groupby(by + column_fields, dropna=False, sort=False, observed=True).agg(reducers)
|
result = frame.groupby(by + column_fields, dropna=False, sort=False, observed=True).agg(reducers)
|
||||||
if column_fields:
|
if column_fields:
|
||||||
result = result.unstack(column_fields)
|
unstacked = result.unstack(column_fields)
|
||||||
result.columns = [json_label(parts) for parts in result.columns.to_flat_index()]
|
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()
|
return result.reset_index()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -14,6 +14,7 @@ from _xlsx_common import (
|
|||||||
EXCEL_INPUT_SUFFIXES,
|
EXCEL_INPUT_SUFFIXES,
|
||||||
EXCEL_OUTPUT_SUFFIXES,
|
EXCEL_OUTPUT_SUFFIXES,
|
||||||
SkillArgumentParser,
|
SkillArgumentParser,
|
||||||
|
cell_range_bounds,
|
||||||
input_file,
|
input_file,
|
||||||
load_json_argument,
|
load_json_argument,
|
||||||
normalize_formula_error,
|
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:
|
def _op_auto_fit(workbook: Any, op: dict[str, Any]) -> int:
|
||||||
import math
|
import math
|
||||||
import unicodedata
|
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"))
|
worksheet = _sheet(workbook, op.get("sheet"))
|
||||||
reference = validate_cell_range(str(op.get("range", "")))
|
reference = validate_cell_range(str(op.get("range", "")))
|
||||||
cells = list(_iter_range_cells(worksheet, reference))
|
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))
|
minimum, maximum = float(op.get("min_width", 8)), float(op.get("max_width", 40))
|
||||||
if not 1 <= minimum <= maximum <= 100:
|
if not 1 <= minimum <= maximum <= 100:
|
||||||
raise ValueError("auto_fit 列宽需满足 1 <= min_width <= max_width <= 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:
|
def _op_add_chart(workbook: Any, op: dict[str, Any]) -> int:
|
||||||
from openpyxl.chart import AreaChart, BarChart, LineChart, PieChart, Reference
|
from openpyxl.chart import AreaChart, BarChart, LineChart, PieChart, Reference
|
||||||
from openpyxl.utils.cell import range_boundaries
|
|
||||||
|
|
||||||
worksheet = _sheet(workbook, op.get("sheet"))
|
worksheet = _sheet(workbook, op.get("sheet"))
|
||||||
chart_type = str(op.get("chart_type", "bar")).lower()
|
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:
|
if chart_type not in chart_classes:
|
||||||
raise ValueError("chart_type 仅支持 area、bar、column、line、pie")
|
raise ValueError("chart_type 仅支持 area、bar、column、line、pie")
|
||||||
data_range = validate_cell_range(str(op.get("data_range", "")))
|
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]()
|
chart = chart_classes[chart_type]()
|
||||||
if isinstance(chart, BarChart):
|
if isinstance(chart, BarChart):
|
||||||
chart.type = "bar" if chart_type == "bar" else "col"
|
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"):
|
if op.get("categories_range"):
|
||||||
category_range = validate_cell_range(str(op["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
|
category_range
|
||||||
)
|
)
|
||||||
categories = Reference(
|
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"])
|
chart.y_axis.title = str(op["y_axis_title"])
|
||||||
if "style" in op:
|
if "style" in op:
|
||||||
chart.style = int(op["style"])
|
chart.style = int(op["style"])
|
||||||
if "height" in op:
|
for dimension in ("height", "width"):
|
||||||
chart.height = float(op["height"])
|
if dimension in op:
|
||||||
if "width" in op:
|
setattr(chart, dimension, float(op[dimension]))
|
||||||
chart.width = float(op["width"])
|
|
||||||
if "legend_position" in op and chart.legend:
|
if "legend_position" in op and chart.legend:
|
||||||
chart.legend.position = str(op["legend_position"])
|
chart.legend.position = str(op["legend_position"])
|
||||||
anchor = validate_cell_reference(str(op.get("anchor", "E2")))
|
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"},
|
{".png", ".jpg", ".jpeg", ".gif", ".bmp"},
|
||||||
)
|
)
|
||||||
image = Image(str(image_path))
|
image = Image(str(image_path))
|
||||||
if "width" in op:
|
for dimension in ("width", "height"):
|
||||||
image.width = float(op["width"])
|
if dimension in op:
|
||||||
if "height" in op:
|
setattr(image, dimension, float(op[dimension]))
|
||||||
image.height = float(op["height"])
|
|
||||||
anchor = validate_cell_reference(str(op.get("anchor", "A1")))
|
anchor = validate_cell_reference(str(op.get("anchor", "A1")))
|
||||||
worksheet.add_image(image, anchor)
|
worksheet.add_image(image, anchor)
|
||||||
return 0
|
return 0
|
||||||
@ -654,7 +652,7 @@ def _op_add_data_validation(workbook: Any, op: dict[str, Any]) -> int:
|
|||||||
worksheet = _sheet(workbook, op.get("sheet"))
|
worksheet = _sheet(workbook, op.get("sheet"))
|
||||||
reference = validate_cell_range(str(op.get("range", "")))
|
reference = validate_cell_range(str(op.get("range", "")))
|
||||||
validation_type = str(op.get("validation_type", "list"))
|
validation_type = str(op.get("validation_type", "list"))
|
||||||
allowed = {
|
allowed = (
|
||||||
"list",
|
"list",
|
||||||
"whole",
|
"whole",
|
||||||
"decimal",
|
"decimal",
|
||||||
@ -662,7 +660,7 @@ def _op_add_data_validation(workbook: Any, op: dict[str, Any]) -> int:
|
|||||||
"time",
|
"time",
|
||||||
"textLength",
|
"textLength",
|
||||||
"custom",
|
"custom",
|
||||||
}
|
)
|
||||||
if validation_type not in allowed:
|
if validation_type not in allowed:
|
||||||
raise ValueError(f"validation_type 不支持:{validation_type}")
|
raise ValueError(f"validation_type 不支持:{validation_type}")
|
||||||
validation = DataValidation(
|
validation = DataValidation(
|
||||||
@ -1051,7 +1049,7 @@ def main() -> dict[str, Any]:
|
|||||||
"path": str(destination),
|
"path": str(destination),
|
||||||
"source": str(source) if source else None,
|
"source": str(source) if source else None,
|
||||||
"sheet_names": workbook.sheetnames,
|
"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),
|
"operation_count": len(operations),
|
||||||
"processed_cell_count": written_cells,
|
"processed_cell_count": written_cells,
|
||||||
"formula_count": scan["formula_count"],
|
"formula_count": scan["formula_count"],
|
||||||
|
|||||||
@ -8,7 +8,7 @@ import os
|
|||||||
import tempfile
|
import tempfile
|
||||||
from datetime import date, datetime
|
from datetime import date, datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Iterable, Optional
|
from typing import Any, Optional
|
||||||
|
|
||||||
from _xlsx_common import (
|
from _xlsx_common import (
|
||||||
TABULAR_INPUT_SUFFIXES,
|
TABULAR_INPUT_SUFFIXES,
|
||||||
@ -81,6 +81,8 @@ def _delimited_to_xlsx(
|
|||||||
)
|
)
|
||||||
workbook = Workbook()
|
workbook = Workbook()
|
||||||
worksheet = workbook.active
|
worksheet = workbook.active
|
||||||
|
if worksheet is None:
|
||||||
|
raise ValueError("工作簿没有活动工作表")
|
||||||
worksheet.title = sheet_name[:31] or "Sheet1"
|
worksheet.title = sheet_name[:31] or "Sheet1"
|
||||||
for row in rows:
|
for row in rows:
|
||||||
worksheet.append(row)
|
worksheet.append(row)
|
||||||
@ -133,6 +135,8 @@ def _xlsx_to_delimited(
|
|||||||
worksheet = workbook[sheet_name]
|
worksheet = workbook[sheet_name]
|
||||||
else:
|
else:
|
||||||
worksheet = workbook.active
|
worksheet = workbook.active
|
||||||
|
if worksheet is None:
|
||||||
|
raise ValueError("工作簿没有活动工作表")
|
||||||
row_count = 0
|
row_count = 0
|
||||||
with destination.open("w", encoding=encoding, newline="") as handle:
|
with destination.open("w", encoding=encoding, newline="") as handle:
|
||||||
writer = csv.writer(handle, delimiter=delimiter)
|
writer = csv.writer(handle, delimiter=delimiter)
|
||||||
|
|||||||
@ -106,6 +106,7 @@ def inspect_excel(
|
|||||||
max_columns: int,
|
max_columns: int,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
from openpyxl import load_workbook
|
from openpyxl import load_workbook
|
||||||
|
from openpyxl.worksheet.worksheet import Worksheet
|
||||||
|
|
||||||
options = openpyxl_load_options(source)
|
options = openpyxl_load_options(source)
|
||||||
formulas = load_workbook(source, data_only=False, **options)
|
formulas = load_workbook(source, data_only=False, **options)
|
||||||
@ -115,6 +116,8 @@ def inspect_excel(
|
|||||||
total_formulas = 0
|
total_formulas = 0
|
||||||
total_errors = 0
|
total_errors = 0
|
||||||
for worksheet in formulas.worksheets:
|
for worksheet in formulas.worksheets:
|
||||||
|
if not isinstance(worksheet, Worksheet):
|
||||||
|
raise ValueError("工作表不支持完整检查,请以普通模式加载工作簿")
|
||||||
formula_count = 0
|
formula_count = 0
|
||||||
error_count = 0
|
error_count = 0
|
||||||
for cell in worksheet._cells.values():
|
for cell in worksheet._cells.values():
|
||||||
@ -136,8 +139,8 @@ def inspect_excel(
|
|||||||
"auto_filter": worksheet.auto_filter.ref,
|
"auto_filter": worksheet.auto_filter.ref,
|
||||||
"merged_ranges": [str(item) for item in worksheet.merged_cells.ranges],
|
"merged_ranges": [str(item) for item in worksheet.merged_cells.ranges],
|
||||||
"tables": list(worksheet.tables.keys()),
|
"tables": list(worksheet.tables.keys()),
|
||||||
"chart_count": len(worksheet._charts),
|
"chart_count": len(getattr(worksheet, "_charts")),
|
||||||
"image_count": len(worksheet._images),
|
"image_count": len(getattr(worksheet, "_images")),
|
||||||
"formula_count": formula_count,
|
"formula_count": formula_count,
|
||||||
"literal_error_count": error_count,
|
"literal_error_count": error_count,
|
||||||
"print_area": str(worksheet.print_area) if worksheet.print_area else None,
|
"print_area": str(worksheet.print_area) if worksheet.print_area else None,
|
||||||
@ -151,7 +154,10 @@ def inspect_excel(
|
|||||||
)
|
)
|
||||||
selected_name = sheet_name
|
selected_name = sheet_name
|
||||||
else:
|
else:
|
||||||
selected_name = formulas.active.title
|
active = formulas.active
|
||||||
|
if active is None:
|
||||||
|
raise ValueError("工作簿没有活动工作表")
|
||||||
|
selected_name = active.title
|
||||||
|
|
||||||
formula_sheet = formulas[selected_name]
|
formula_sheet = formulas[selected_name]
|
||||||
cached_sheet = cached[selected_name]
|
cached_sheet = cached[selected_name]
|
||||||
@ -184,7 +190,7 @@ def inspect_excel(
|
|||||||
"format": source.suffix.lower(),
|
"format": source.suffix.lower(),
|
||||||
"macro_enabled": source.suffix.lower() in {".xlsm", ".xltm"},
|
"macro_enabled": source.suffix.lower() in {".xlsm", ".xltm"},
|
||||||
"has_external_links": workbook_has_external_links(source),
|
"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,
|
"sheet_names": formulas.sheetnames,
|
||||||
"sheets": summaries,
|
"sheets": summaries,
|
||||||
"defined_names": _defined_names(formulas),
|
"defined_names": _defined_names(formulas),
|
||||||
|
|||||||
@ -3,13 +3,12 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import math
|
import math
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from _xlsx_common import SkillArgumentParser, load_json_argument, run_cli
|
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]:
|
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):
|
if not np.allclose(np.diag(matrix), 1) or not np.allclose(matrix * matrix.T, 1, atol=1e-6):
|
||||||
raise ValueError("AHP 比较矩阵必须对角为 1 且互反")
|
raise ValueError("AHP 比较矩阵必须对角为 1 且互反")
|
||||||
values, vectors = np.linalg.eig(matrix)
|
values, vectors = np.linalg.eig(matrix)
|
||||||
index = np.argmax(values.real)
|
index = np.argmax(np.real(values))
|
||||||
weights = np.abs(vectors[:, index].real)
|
weights = np.abs(np.real(vectors[:, index]))
|
||||||
ri = [0, 0, 0, .58, .90, 1.12, 1.24, 1.32, 1.41, 1.45][n]
|
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
|
consistency = max(0, float(values[index].real - n) / (n - 1) / ri) if ri else 0
|
||||||
if consistency >= .1:
|
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()]
|
metrics += [[label, algorithm, key, float(value)] for key, value in stats.items()]
|
||||||
predictions = frame.iloc[test][[SOURCE_ROW, *features]].copy()
|
predictions = frame.iloc[test][[SOURCE_ROW, *features]].copy()
|
||||||
predictions["实际值"] = y.iloc[test].to_numpy()
|
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"]
|
model = pipeline.named_steps["model"]
|
||||||
names = pipeline.named_steps["prepare"].get_feature_names_out()
|
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)
|
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)
|
future_X[col] = future_X[col].map(lambda x: str(x) if pd.notna(x) else np.nan)
|
||||||
pipeline.fit(X, y)
|
pipeline.fit(X, y)
|
||||||
future = future[[SOURCE_ROW, *features]].copy()
|
future = future[[SOURCE_ROW, *features]].copy()
|
||||||
future["预测值"] = pipeline.predict(future_X)
|
future["预测值"] = np.asarray(pipeline.predict(future_X))
|
||||||
tables.append(("新样本预测", future))
|
tables.append(("新样本预测", future))
|
||||||
else:
|
else:
|
||||||
provenance = None
|
provenance = None
|
||||||
@ -282,7 +281,7 @@ def unsupervised(frame: pd.DataFrame, spec: dict, *, anomaly: bool) -> tuple[lis
|
|||||||
labels = model.fit_predict(scaled)
|
labels = model.fit_predict(scaled)
|
||||||
result = frame.copy()
|
result = frame.copy()
|
||||||
result["异常标记" if anomaly else "簇编号"] = labels
|
result["异常标记" if anomaly else "簇编号"] = labels
|
||||||
if anomaly:
|
if isinstance(model, IsolationForest):
|
||||||
result["正常程度得分"] = model.decision_function(scaled)
|
result["正常程度得分"] = model.decision_function(scaled)
|
||||||
score = None
|
score = None
|
||||||
valid = labels != -1
|
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)
|
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():
|
if start.shape != (n,) or not np.isfinite(start).all():
|
||||||
raise ValueError("initial 需为有限数值向量")
|
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",
|
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})
|
bounds=Bounds(lower, upper), constraints=[linear] if linear else [], options={"maxiter": 1000, "ftol": 1e-9})
|
||||||
guarantee = "凸二次规划的数值解,已检查可行性"
|
guarantee = "凸二次规划的数值解,已检查可行性"
|
||||||
else:
|
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),
|
result = milp(sign * c, integrality=np.asarray(integer, dtype=int), bounds=Bounds(lower, upper),
|
||||||
constraints=linear, options={"time_limit": 60., "mip_rel_gap": 0.})
|
constraints=linear, options={"time_limit": 60., "mip_rel_gap": 0.})
|
||||||
guarantee = "HiGHS 求解成功;仅在成功且可行时输出方案"
|
guarantee = "HiGHS 求解成功;仅在成功且可行时输出方案"
|
||||||
|
|||||||
@ -1,6 +1,5 @@
|
|||||||
import csv
|
import csv
|
||||||
import hashlib
|
import hashlib
|
||||||
import json
|
|
||||||
import sys
|
import sys
|
||||||
import unittest
|
import unittest
|
||||||
import tempfile
|
import tempfile
|
||||||
@ -12,8 +11,7 @@ import numpy as np
|
|||||||
import pandas as pd
|
import pandas as pd
|
||||||
from openpyxl import Workbook, load_workbook
|
from openpyxl import Workbook, load_workbook
|
||||||
|
|
||||||
ROOT = Path(__file__).resolve().parents[1]
|
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "scripts"))
|
||||||
sys.path.insert(0, str(ROOT / 'scripts'))
|
|
||||||
import _xlsx_common as common
|
import _xlsx_common as common
|
||||||
import _xlsx_data as data
|
import _xlsx_data as data
|
||||||
import analyze_workbook as analysis
|
import analyze_workbook as analysis
|
||||||
@ -21,10 +19,13 @@ import apply_workbook as writer
|
|||||||
import inspect_workbook as inspector
|
import inspect_workbook as inspector
|
||||||
import model_workbook as modeling
|
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()
|
QA = Path(_TEMP.name).resolve()
|
||||||
_ORIGINAL_OUTPUT_ROOT = common.EXCEL_OUTPUT_ROOT
|
_ORIGINAL_OUTPUT_ROOT = common.EXCEL_OUTPUT_ROOT
|
||||||
|
|
||||||
|
|
||||||
def tearDownModule():
|
def tearDownModule():
|
||||||
common.EXCEL_OUTPUT_ROOT = _ORIGINAL_OUTPUT_ROOT
|
common.EXCEL_OUTPUT_ROOT = _ORIGINAL_OUTPUT_ROOT
|
||||||
_TEMP.cleanup()
|
_TEMP.cleanup()
|
||||||
@ -34,185 +35,394 @@ class MergeTests(unittest.TestCase):
|
|||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
common.EXCEL_OUTPUT_ROOT = QA
|
common.EXCEL_OUTPUT_ROOT = QA
|
||||||
cls.source = QA / 'sales.xlsx'
|
cls.source = QA / "sales.xlsx"
|
||||||
wb = Workbook()
|
wb = Workbook()
|
||||||
ws = wb.active
|
ws = wb.active
|
||||||
ws.title = '明细'
|
assert ws is not None
|
||||||
for row in [['地区', '收入', '成本', '说明'], ['华东', 100, 60, '满意'], ['华东', -20, 5, '退款'], ['华南', 80, 50, '物流慢'], ['华南', None, 20, '服务好但物流慢'], ['合计', 160, 135, None]]:
|
ws.title = "明细"
|
||||||
|
for row in [
|
||||||
|
["地区", "收入", "成本", "说明"],
|
||||||
|
["华东", 100, 60, "满意"],
|
||||||
|
["华东", -20, 5, "退款"],
|
||||||
|
["华南", 80, 50, "物流慢"],
|
||||||
|
["华南", None, 20, "服务好但物流慢"],
|
||||||
|
["合计", 160, 135, None],
|
||||||
|
]:
|
||||||
ws.append(row)
|
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)
|
wb.save(cls.source)
|
||||||
cls.source_hash = hashlib.sha256(cls.source.read_bytes()).hexdigest()
|
cls.source_hash = hashlib.sha256(cls.source.read_bytes()).hexdigest()
|
||||||
cls.csv = QA / 'text.csv'
|
cls.csv = QA / "text.csv"
|
||||||
with cls.csv.open('w', encoding='utf-8-sig', newline='') as f:
|
with cls.csv.open("w", encoding="utf-8-sig", newline="") as f:
|
||||||
csv.writer(f).writerows([['ID', '文本'], ['001', '=1+1'], ['002', '长文本' * 80]])
|
csv.writer(f).writerows(
|
||||||
|
[["ID", "文本"], ["001", "=1+1"], ["002", "长文本" * 80]]
|
||||||
|
)
|
||||||
|
|
||||||
def dataset(self):
|
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):
|
def test_profile_finds_totals_without_dropping_negative_rows(self):
|
||||||
_, result = analysis.analyze(str(self.source), {'method': 'profile'})
|
_, result = analysis.analyze(str(self.source), {"method": "profile"})
|
||||||
self.assertEqual(result['source']['rows_used'], 5)
|
self.assertEqual(result["source"]["rows_used"], 5)
|
||||||
self.assertEqual(result['summary_row_candidates'][0]['row'], 6)
|
self.assertEqual(result["summary_row_candidates"][0]["row"], 6)
|
||||||
|
|
||||||
def test_aggregate_negative_values_and_null_count(self):
|
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'}})
|
tables, _ = analysis.analyze(
|
||||||
result = tables[0][1].set_index('地区')
|
str(self.source),
|
||||||
self.assertEqual(result.loc['华东', '收入'], 80)
|
{
|
||||||
self.assertEqual(result.loc['华南', '说明'], 2)
|
"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):
|
def test_all_null_sum_not_zero(self):
|
||||||
frame = pd.DataFrame({'类别': ['A','B','B'], '值': [None, 0, None]})
|
frame = pd.DataFrame({"类别": ["A", "B", "B"], "值": [None, 0, None]})
|
||||||
result = analysis.aggregate(frame, {'by': ['类别'], 'metrics': {'值': 'sum'}}, pivot=False).set_index('类别')
|
result = analysis.aggregate(
|
||||||
self.assertTrue(pd.isna(result.loc['A','值']))
|
frame, {"by": ["类别"], "metrics": {"值": "sum"}}, pivot=False
|
||||||
self.assertEqual(result.loc['B','值'], 0)
|
).set_index("类别")
|
||||||
|
self.assertTrue(pd.isna(result.loc["A", "值"]))
|
||||||
|
self.assertEqual(result.loc["B", "值"], 0)
|
||||||
|
|
||||||
def test_pivot_multiple_levels(self):
|
def test_pivot_multiple_levels(self):
|
||||||
frame = pd.DataFrame({'地区':['A','A','B'], '年':['2025','2026','2025'], '收入':[10,20,30]})
|
frame = pd.DataFrame(
|
||||||
result = analysis.aggregate(frame, {'by':['地区'],'columns':['年'],'metrics':{'收入':'sum'}}, pivot=True).set_index('地区')
|
{
|
||||||
self.assertEqual(result.loc['A','["收入","2026"]'],20)
|
"地区": ["A", "A", "B"],
|
||||||
self.assertTrue(pd.isna(result.loc['B','["收入","2026"]']))
|
"年": ["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):
|
def test_rules_report_conflict_and_unknown(self):
|
||||||
detail, summary = analysis.classify(self.dataset(), {'column':'说明','rules':[{'label':'正向','keywords':['好','满意']},{'label':'物流','keywords':['慢']}]})
|
detail, summary = analysis.classify(
|
||||||
self.assertEqual(detail['分类'].tolist(), ['正向','未知','物流','需复核'])
|
self.dataset(),
|
||||||
self.assertAlmostEqual(summary['占比'].sum(),1)
|
{
|
||||||
|
"column": "说明",
|
||||||
|
"rules": [
|
||||||
|
{"label": "正向", "keywords": ["好", "满意"]},
|
||||||
|
{"label": "物流", "keywords": ["慢"]},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
self.assertEqual(detail["分类"].tolist(), ["正向", "未知", "物流", "需复核"])
|
||||||
|
self.assertAlmostEqual(summary["占比"].sum(), 1)
|
||||||
|
|
||||||
def test_join_rejects_accidental_many_to_many(self):
|
def test_join_rejects_accidental_many_to_many(self):
|
||||||
with self.assertRaises(pd.errors.MergeError):
|
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):
|
def test_missing_headers_can_use_coordinates(self):
|
||||||
wb=Workbook(); ws=wb.active
|
wb = Workbook()
|
||||||
ws.append([None,'值']); ws.append(['001',5])
|
ws = wb.active
|
||||||
path=QA/'no_header.xlsx'; wb.save(path)
|
assert ws is not None
|
||||||
with self.assertRaises(ValueError): data.read_dataset(str(path))
|
ws.append([None, "值"])
|
||||||
frame, meta = data.read_dataset(str(path), {'columns':{'编号':'A','值':'B'}})
|
ws.append(["001", 5])
|
||||||
self.assertEqual(frame.iloc[0]['编号'],'001')
|
path = QA / "no_header.xlsx"
|
||||||
self.assertEqual(meta['columns']['编号'],'A')
|
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):
|
def test_formula_cache_is_required(self):
|
||||||
wb=Workbook(); ws=wb.active; ws.append(['值']); ws.append(['=1+1'])
|
wb = Workbook()
|
||||||
path=QA/'uncached.xlsx'; wb.save(path)
|
ws = wb.active
|
||||||
with self.assertRaisesRegex(ValueError,'公式缓存'):
|
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))
|
data.read_dataset(str(path))
|
||||||
|
|
||||||
def test_safe_text_and_no_truncation_in_writer(self):
|
def test_safe_text_and_no_truncation_in_writer(self):
|
||||||
tables, metadata=analysis.analyze(str(self.csv), {'method':'transform'})
|
tables, metadata = analysis.analyze(str(self.csv), {"method": "transform"})
|
||||||
plan=QA/'safe-text.json'; output=QA/'safe-text.xlsx'
|
plan = QA / "safe-text.json"
|
||||||
data.save_plan(tables,metadata,str(plan),overwrite=True)
|
output = QA / "safe-text.xlsx"
|
||||||
with patch.object(sys,'argv',['apply','--output',str(output),'--spec-file',str(plan),'--overwrite']):
|
data.save_plan(tables, metadata, str(plan), overwrite=True)
|
||||||
result=writer.main()
|
with patch.object(
|
||||||
self.assertEqual(result['formula_count'],0)
|
sys,
|
||||||
wb=load_workbook(output); ws=wb['分析结果']
|
"argv",
|
||||||
self.assertEqual(ws['B2'].value,'001')
|
["apply", "--output", str(output), "--spec-file", str(plan), "--overwrite"],
|
||||||
self.assertEqual(ws['C2'].value,'=1+1')
|
):
|
||||||
self.assertEqual(ws['C2'].data_type,'s')
|
result = writer.main()
|
||||||
self.assertEqual(ws['C3'].value,'长文本'*80)
|
self.assertEqual(result["formula_count"], 0)
|
||||||
self.assertTrue(ws['C3'].alignment.wrap_text)
|
wb = load_workbook(output)
|
||||||
self.assertGreater(ws.row_dimensions[3].height,36)
|
ws = wb["分析结果"]
|
||||||
self.assertEqual(inspector.inspect_excel(output,sheet_name='分析结果',start_row=1,start_column=1,max_rows=10,max_columns=10)['formula_count'],0)
|
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):
|
def test_append_result_preserves_source(self):
|
||||||
spec={'method':'aggregate','source':{'exclude_rows':[6]},'by':['地区'],'metrics':{'收入':'sum'},'chart':{'category':'地区','values':['收入'],'title':'地区收入'}}
|
spec = {
|
||||||
tables,meta=analysis.analyze(str(self.source),spec)
|
"method": "aggregate",
|
||||||
plan=QA/'summary.json'; output=QA/'summary.xlsx'
|
"source": {"exclude_rows": [6]},
|
||||||
data.save_plan(tables,meta,str(plan),overwrite=True,chart=spec['chart'])
|
"by": ["地区"],
|
||||||
with patch.object(sys,'argv',['apply','--input',str(self.source),'--output',str(output),'--spec-file',str(plan),'--overwrite']): writer.main()
|
"metrics": {"收入": "sum"},
|
||||||
self.assertEqual(hashlib.sha256(self.source.read_bytes()).hexdigest(),self.source_hash)
|
"chart": {"category": "地区", "values": ["收入"], "title": "地区收入"},
|
||||||
original=load_workbook(self.source); result=load_workbook(output)
|
}
|
||||||
self.assertEqual(list(original['明细'].values),list(result['明细'].values))
|
tables, meta = analysis.analyze(str(self.source), spec)
|
||||||
self.assertEqual(copy(original['明细']['A1'].font),copy(result['明细']['A1'].font))
|
plan = QA / "summary.json"
|
||||||
self.assertEqual(len(result['分析结果']._charts),1)
|
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):
|
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)):
|
with self.subTest(explicit=isinstance(value, dict)):
|
||||||
wb = Workbook()
|
wb = Workbook()
|
||||||
with self.assertRaisesRegex(ValueError, '不能静默截断'):
|
assert wb.active is not None
|
||||||
writer._op_write_rows(wb, {'sheet': wb.active.title, 'rows': [[value]]})
|
with self.assertRaisesRegex(ValueError, "不能静默截断"):
|
||||||
|
writer._op_write_rows(
|
||||||
|
wb, {"sheet": wb.active.title, "rows": [[value]]}
|
||||||
|
)
|
||||||
|
|
||||||
def test_cost_direction_once(self):
|
def test_cost_direction_once(self):
|
||||||
frame=pd.DataFrame({data.SOURCE_ROW:[2,3],'对象':['便宜','贵'],'质量':[98,98],'价格':[10,30]})
|
frame = pd.DataFrame(
|
||||||
tables, _ = modeling.evaluate(frame, {'entity':'对象','directions':{'质量':'benefit','价格':'cost'}})
|
{
|
||||||
rank=tables[0][1].set_index('对象')
|
data.SOURCE_ROW: [2, 3],
|
||||||
self.assertEqual(rank.loc['便宜','排名'],1)
|
"对象": ["便宜", "贵"],
|
||||||
self.assertEqual(rank.loc['贵','排名'],2)
|
"质量": [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):
|
def test_identical_objects_tie(self):
|
||||||
frame=pd.DataFrame({data.SOURCE_ROW:[2,3],'对象':['A','B'],'值':[10,10]})
|
frame = pd.DataFrame(
|
||||||
tables,_=modeling.evaluate(frame,{'entity':'对象','directions':{'值':'cost'},'weighting':'entropy'})
|
{data.SOURCE_ROW: [2, 3], "对象": ["A", "B"], "值": [10, 10]}
|
||||||
self.assertEqual(tables[0][1]['排名'].tolist(),[1,1])
|
)
|
||||||
self.assertEqual(tables[0][1]['得分'].tolist(),[.5,.5])
|
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):
|
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]})
|
frame = pd.DataFrame(
|
||||||
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]]})
|
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):
|
def test_forecast_by_group_and_time_holdout(self):
|
||||||
dates=pd.date_range('2024-01-01',periods=24,freq='MS')
|
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)})
|
frame = pd.DataFrame(
|
||||||
tables,_=modeling.forecast(frame,{'by':['城市'],'date':'月份','value':'销量','horizon':3,'frequency':'MS'})
|
{
|
||||||
result=tables[0][1]
|
"城市": ["A"] * 24 + ["B"] * 24,
|
||||||
self.assertEqual(len(result),6)
|
"月份": list(dates) * 2,
|
||||||
self.assertAlmostEqual(result[result['城市']=='A'].iloc[0]['预测值'],340)
|
"销量": list(np.arange(24) * 10 + 100) + list(np.arange(24) * -2 + 100),
|
||||||
self.assertAlmostEqual(result[result['城市']=='B'].iloc[0]['预测值'],52)
|
}
|
||||||
self.assertTrue((tables[1][1]['训练期数']==18).all())
|
)
|
||||||
|
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):
|
def test_forecast_rejects_missing_month(self):
|
||||||
frame=pd.DataFrame({'日期':pd.date_range('2024-01-01',periods=8,freq='MS').delete(3),'值':range(7)})
|
frame = pd.DataFrame(
|
||||||
with self.assertRaisesRegex(ValueError,'连续'):
|
{
|
||||||
modeling.forecast(frame,{'date':'日期','value':'值'})
|
"日期": 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):
|
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})
|
frame = pd.DataFrame(
|
||||||
tables,meta=modeling.supervised(frame,{'features':['x'],'target':'y'},classification=False)
|
{data.SOURCE_ROW: range(2, 42), "x": range(40), "y": np.arange(40) * 3 + 7}
|
||||||
metrics=tables[1][1]
|
)
|
||||||
error=metrics[(metrics['数据集']=='测试')&(metrics['模型']=='linear')&(metrics['指标']=='RMSE')].iloc[0]['值']
|
tables, meta = modeling.supervised(
|
||||||
self.assertLess(error,1e-8)
|
frame, {"features": ["x"], "target": "y"}, classification=False
|
||||||
self.assertEqual(meta['train_rows'],32)
|
)
|
||||||
self.assertIn('基线',metrics['模型'].tolist())
|
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):
|
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})
|
frame = pd.DataFrame(
|
||||||
tables,_=modeling.supervised(frame,{'features':['x','组'],'categorical':['组'],'target':'标签'},classification=True)
|
{
|
||||||
self.assertEqual(len(tables[0][1]),8)
|
data.SOURCE_ROW: range(2, 42),
|
||||||
self.assertIn('F1_macro',tables[1][1]['指标'].tolist())
|
"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):
|
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})
|
frame = pd.DataFrame(
|
||||||
tables, metadata = modeling.supervised(frame, {'features': ['x'], 'target': 'y', 'test_fraction': .1}, classification=False)
|
{data.SOURCE_ROW: range(2, 12), "x": range(10), "y": np.arange(10) * 3 + 7}
|
||||||
self.assertEqual(metadata['test_rows'], 2)
|
)
|
||||||
self.assertTrue(np.isfinite(tables[1][1]['值']).all())
|
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):
|
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]})
|
frame = pd.DataFrame(
|
||||||
tables,_=modeling.unsupervised(frame,{'features':['x'],'clusters':2},anomaly=False)
|
{data.SOURCE_ROW: range(8), "x": [0, 0.1, 0.2, 0.3, 10, 10.1, 10.2, 10.3]}
|
||||||
labels=tables[0][1]['簇编号'].to_numpy()
|
)
|
||||||
self.assertTrue(np.all(labels[:4]==labels[0]))
|
tables, _ = modeling.unsupervised(
|
||||||
self.assertNotEqual(labels[0],labels[-1])
|
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):
|
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}]})
|
tables, meta = modeling.optimize(
|
||||||
self.assertEqual(meta['objective_value'],8)
|
{
|
||||||
self.assertEqual(tables[0][1]['取值'].tolist(),[0,4])
|
"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):
|
def test_infeasible_and_unbounded_rejected(self):
|
||||||
with self.assertRaisesRegex(ValueError,'求解未成功'):
|
with self.assertRaisesRegex(ValueError, "求解未成功"):
|
||||||
modeling.optimize({'variables':['x'],'objective':[1],'constraints':[{'coefficients':[1],'relation':'<=','rhs':-1}]})
|
modeling.optimize(
|
||||||
with self.assertRaisesRegex(ValueError,'求解未成功'):
|
{
|
||||||
modeling.optimize({'variables':['x'],'objective':[1],'sense':'max'})
|
"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):
|
def test_convex_quadratic(self):
|
||||||
tables,meta=modeling.optimize({'variables':['x'],'objective':[-4],'quadratic':[[2]]})
|
tables, meta = modeling.optimize(
|
||||||
self.assertAlmostEqual(tables[0][1].iloc[0]['取值'],2)
|
{"variables": ["x"], "objective": [-4], "quadratic": [[2]]}
|
||||||
self.assertAlmostEqual(meta['objective_value'],-4)
|
)
|
||||||
|
self.assertAlmostEqual(tables[0][1].iloc[0]["取值"], 2)
|
||||||
|
self.assertAlmostEqual(meta["objective_value"], -4)
|
||||||
|
|
||||||
def test_output_root_enforced(self):
|
def test_output_root_enforced(self):
|
||||||
with self.assertRaises(ValueError):
|
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)
|
unittest.main(verbosity=2)
|
||||||
|
|||||||
@ -3,7 +3,9 @@ from __future__ import annotations
|
|||||||
import contextlib
|
import contextlib
|
||||||
import importlib.util
|
import importlib.util
|
||||||
import io
|
import io
|
||||||
|
import json
|
||||||
import os
|
import os
|
||||||
|
import sqlite3
|
||||||
import sys
|
import sys
|
||||||
import unittest
|
import unittest
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@ -14,6 +16,60 @@ SCRIPT_PATH = (
|
|||||||
Path(__file__).resolve().parents[1]
|
Path(__file__).resolve().parents[1]
|
||||||
/ "skills/send-complex-message/scripts/send_complex_message.py"
|
/ "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):
|
class SendComplexMessageTests(unittest.TestCase):
|
||||||
@ -94,6 +150,194 @@ class SendComplexMessageTests(unittest.TestCase):
|
|||||||
post.assert_not_called()
|
post.assert_not_called()
|
||||||
self.assertFalse(output.getvalue().endswith("ended"))
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user