Commit dd919218 by ccran

feat: add optimize_review_prompt_flow

parent d9c25a8a
...@@ -6,16 +6,11 @@ ...@@ -6,16 +6,11 @@
# Keep Python source files # Keep Python source files
!**/*.py !**/*.py
!**/*.doc
!**/*.docx
!**/*.xlsx
!**/*.pdf
!**/*.xls
!workflow/** !workflow/**
!demo/** !demo/**
!skills/** !skills/**
!README.md !README.md
!data/*.xlsx
# Keep this file tracked # Keep this file tracked
!.gitignore !.gitignore
......
...@@ -16,31 +16,40 @@ from utils.http_util import upload_file, fastgpt_openai_chat, download_file ...@@ -16,31 +16,40 @@ from utils.http_util import upload_file, fastgpt_openai_chat, download_file
use_lufa = False use_lufa = False
batch_size = 5 batch_size = 5
if not use_lufa:
SUFFIX = "_麓发迁移" def get_params():
batch_input_dir_path = "jp-input" output_suffix = f"{time.strftime('%Y%m%d-%H%M%S')}-{time.time_ns() % 1_000_000:06d}"
batch_output_dir_path = f"/data/home/htsc/jp-contract/data/benchmark/results/jp-output-simple" if not use_lufa:
return {
"suffix": "_麓发迁移",
"batch_size": batch_size,
"batch_input_dir_path": "jp-input",
"batch_output_dir_path": f"/data/home/htsc/jp-contract/data/benchmark/results/jp-output-{output_suffix}",
# 金盘fastgpt接口 # 金盘fastgpt接口
url = "http://172.21.107.45:3002/api/v1/chat/completions" "url": "http://172.21.107.45:3002/api/v1/chat/completions",
# 金盘迁移麓发合同审查测试token # 金盘迁移麓发合同审查测试token
# token = "fastgpt-vykT6qs07g7hR4tL2MNJE6DdNCIxaQjEu3Cxw9nuTBFg8MAG3CkByvnXKxSNEyMK7" # "token": "fastgpt-vykT6qs07g7hR4tL2MNJE6DdNCIxaQjEu3Cxw9nuTBFg8MAG3CkByvnXKxSNEyMK7",
token = 'fastgpt-pYh0DgMVPDh9DptmznbSz7fRAqS41Z7gUWWIPM0112APpMlbb8mc1iztJi' "token": "fastgpt-pYh0DgMVPDh9DptmznbSz7fRAqS41Z7gUWWIPM0112APpMlbb8mc1iztJi",
# 人机交互测试(测试环境) # 人机交互测试(测试环境)
# token = 'fastgpt-p189K5zoTX5wjp0dBybFCwsbWm3juIwlJxt2wTGyiaOWOANI5Y10pKEZzyt' # "token": "fastgpt-p189K5zoTX5wjp0dBybFCwsbWm3juIwlJxt2wTGyiaOWOANI5Y10pKEZzyt",
# 人机交互测试(生产环境) # 人机交互测试(生产环境)
# token = "fastgpt-ry4jIjgNwmNgufMr5jR0ncvJVmSS4GZl4bx2ItsNPoncdQzW9Na3IP1Xrankr" # "token": "fastgpt-ry4jIjgNwmNgufMr5jR0ncvJVmSS4GZl4bx2ItsNPoncdQzW9Na3IP1Xrankr",
# 提取后审查测试 # 提取后审查测试
# token = 'fastgpt-n74gGX5ZqLT6o1ysMBSGUTjIciswYOWDRfQ75krMkE5gDVDkpzsbz8u' # "token": "fastgpt-n74gGX5ZqLT6o1ysMBSGUTjIciswYOWDRfQ75krMkE5gDVDkpzsbz8u",
else: }
SUFFIX = "_麓发"
batch_input_dir_path = "4.24测财务合同审核" return {
batch_output_dir_path = "4.24测财务合同审核-batch" "suffix": "_麓发",
"batch_size": batch_size,
"batch_input_dir_path": "4.24测财务合同审核",
"batch_output_dir_path": f"4.24测财务合同审核-batch-{output_suffix}",
# 麓发fastgpt接口 # 麓发fastgpt接口
url = "http://192.168.252.71:18089/api/v1/chat/completions" "url": "http://192.168.252.71:18089/api/v1/chat/completions",
# 麓发合同审查生产token # 麓发合同审查生产token
# token = "fastgpt-ek3Z6PxI6sXgYc0jxzZ5bVGqrxwM6aVyfSmA6JVErJYBMr2KmYxrHwEUOIMSYz" # "token": "fastgpt-ek3Z6PxI6sXgYc0jxzZ5bVGqrxwM6aVyfSmA6JVErJYBMr2KmYxrHwEUOIMSYz",
# 麓发合同审查生产token-标准化 # 麓发合同审查生产token-标准化
token = "fastgpt-mg5tQUgreJeF7peoOr5zqP0NR4EIrfS2bEVXge6FUL94Suu1TvEMR1sGNRSiV" "token": "fastgpt-mg5tQUgreJeF7peoOr5zqP0NR4EIrfS2bEVXge6FUL94Suu1TvEMR1sGNRSiV",
}
def extract_url(text): def extract_url(text):
...@@ -59,7 +68,7 @@ def extract_url(text): ...@@ -59,7 +68,7 @@ def extract_url(text):
def process_single_file( def process_single_file(
file, batch_input_dir_path, batch_output_dir_path, counter, start_file file, params, counter, start_file
): ):
""" """
单文件处理逻辑,可被线程池并发调用 单文件处理逻辑,可被线程池并发调用
...@@ -68,14 +77,20 @@ def process_single_file( ...@@ -68,14 +77,20 @@ def process_single_file(
if start_file > counter: if start_file > counter:
return return
batch_input_dir_path = params["batch_input_dir_path"]
batch_output_dir_path = params["batch_output_dir_path"]
suffix = params["suffix"]
url = params["url"]
token = params["token"]
# 提取文件前缀 # 提取文件前缀
file_name = file[: file.rfind(".")] file_name = file[: file.rfind(".")]
ext_name = file[file.rfind(".") :] ext_name = file[file.rfind(".") :]
# 源目标处理 # 源目标处理
original_file = f"{batch_input_dir_path}/{file}" original_file = f"{batch_input_dir_path}/{file}"
des_check_file = f"{batch_output_dir_path}/{file_name}.md" des_check_file = f"{batch_output_dir_path}/{file_name}.md"
des_excel_file = f"{batch_output_dir_path}/{file_name}{SUFFIX}.xlsx" des_excel_file = f"{batch_output_dir_path}/{file_name}{suffix}.xlsx"
des_doc_file = f"{batch_output_dir_path}/{file_name}{SUFFIX}{ext_name}" des_doc_file = f"{batch_output_dir_path}/{file_name}{suffix}{ext_name}"
try: try:
# 处理原文件 # 处理原文件
...@@ -117,7 +132,11 @@ def process_single_file( ...@@ -117,7 +132,11 @@ def process_single_file(
logger.error(traceback.print_exc()) logger.error(traceback.print_exc())
def execute_batch(max_workers: int = 4): def execute_batch(params=None):
params = params or get_params()
max_workers = params["batch_size"]
batch_input_dir_path = params["batch_input_dir_path"]
batch_output_dir_path = params["batch_output_dir_path"]
start_file = 1 start_file = 1
dirs = os.listdir(batch_input_dir_path) dirs = os.listdir(batch_input_dir_path)
os.makedirs(batch_output_dir_path, exist_ok=True) os.makedirs(batch_output_dir_path, exist_ok=True)
...@@ -127,8 +146,7 @@ def execute_batch(max_workers: int = 4): ...@@ -127,8 +146,7 @@ def execute_batch(max_workers: int = 4):
executor.submit( executor.submit(
process_single_file, process_single_file,
file, file,
batch_input_dir_path, params,
batch_output_dir_path,
counter, counter,
start_file, start_file,
) )
...@@ -142,6 +160,7 @@ def execute_batch(max_workers: int = 4): ...@@ -142,6 +160,7 @@ def execute_batch(max_workers: int = 4):
if __name__ == "__main__": if __name__ == "__main__":
import os import os
execute_batch(batch_size) params = get_params()
execute_batch(params)
print("all done!") print("all done!")
print("文件保存在: ", os.path.abspath(batch_output_dir_path)) print("文件保存在: ", os.path.abspath(params["batch_output_dir_path"]))
import argparse import argparse
from pathlib import Path from pathlib import Path
from typing import Iterable
import pandas as pd import pandas as pd
from rapidfuzz import fuzz from rapidfuzz import fuzz
from contextlib import redirect_stdout, redirect_stderr
import time import time
fuzz_score_threshold = 80 fuzz_score_threshold = 80
EXCEL_SHEET_NAME_LIMIT = 31
INVALID_SHEET_NAME_CHARS = set(r'[]:*?/\\')
def _normalize_cell(value: object) -> str: def _normalize_cell(value: object) -> str:
...@@ -39,9 +41,55 @@ def _load_rows(path: Path) -> list[tuple[str, str]]: ...@@ -39,9 +41,55 @@ def _load_rows(path: Path) -> list[tuple[str, str]]:
return rows return rows
def _compare_impl(val_dir: Path, answer_dir: Path) -> None: def _normalize_ignore_items(ignore_items: Iterable[str] | None) -> set[str]:
if not ignore_items:
return set()
return {item.strip() for item in ignore_items if item.strip()}
def _filter_ignored_rows(
rows: list[tuple[str, str]],
ignore_items: set[str],
) -> list[tuple[str, str]]:
if not ignore_items:
return rows
return [(item, text) for item, text in rows if item not in ignore_items]
def _safe_sheet_name(name: str, used_names: set[str]) -> str:
sheet_name = "".join(
"_" if char in INVALID_SHEET_NAME_CHARS else char for char in name.strip()
)
sheet_name = sheet_name or "空审查项"
sheet_name = sheet_name[:EXCEL_SHEET_NAME_LIMIT]
if sheet_name not in used_names:
used_names.add(sheet_name)
return sheet_name
counter = 2
while True:
suffix = f"_{counter}"
candidate = sheet_name[: EXCEL_SHEET_NAME_LIMIT - len(suffix)] + suffix
if candidate not in used_names:
used_names.add(candidate)
return candidate
counter += 1
def _compare_impl(
val_dir: Path,
answer_dir: Path,
print_result: bool = False,
ignore_items: Iterable[str] | None = None,
) -> Path | None:
val_dir = val_dir.resolve() val_dir = val_dir.resolve()
answer_dir = answer_dir.resolve() answer_dir = answer_dir.resolve()
ignore_item_set = _normalize_ignore_items(ignore_items)
def log(*args: object, **kwargs: object) -> None:
if print_result:
print(*args, **kwargs)
overall_val = overall_answer = overall_matched = 0 overall_val = overall_answer = overall_matched = 0
...@@ -50,15 +98,17 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None: ...@@ -50,15 +98,17 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None:
overall_item_matched: dict[str, int] = {} overall_item_matched: dict[str, int] = {}
overall_item_unmatched_answer: dict[str, int] = {} overall_item_unmatched_answer: dict[str, int] = {}
overall_item_unmatched_val: dict[str, int] = {} overall_item_unmatched_val: dict[str, int] = {}
unmatched_answer_details_by_item: dict[str, list[str]] = {}
unmatched_val_details_by_item: dict[str, list[str]] = {}
for val_file in sorted(val_dir.glob("*.xlsx")): for val_file in sorted(val_dir.glob("*.xlsx")):
answer_file = answer_dir / val_file.name answer_file = answer_dir / val_file.name
if not answer_file.exists(): if not answer_file.exists():
print(f"Skip {val_file.name}: missing in answer") log(f"Skip {val_file.name}: missing in answer")
continue continue
val_rows = _load_rows(val_file) val_rows = _filter_ignored_rows(_load_rows(val_file), ignore_item_set)
answer_rows = _load_rows(answer_file) answer_rows = _filter_ignored_rows(_load_rows(answer_file), ignore_item_set)
# Baseline: answer -> match val, consume val to keep 1-1, report leftover answers # Baseline: answer -> match val, consume val to keep 1-1, report leftover answers
answer_counts: dict[str, int] = {} answer_counts: dict[str, int] = {}
...@@ -136,23 +186,25 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None: ...@@ -136,23 +186,25 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None:
overall_item_unmatched_answer[it] = overall_item_unmatched_answer.get( overall_item_unmatched_answer[it] = overall_item_unmatched_answer.get(
it, 0 it, 0
) + len(lst) ) + len(lst)
unmatched_answer_details_by_item.setdefault(it, []).extend(lst)
for it, lst in unmatched_val_by_item.items(): for it, lst in unmatched_val_by_item.items():
overall_item_unmatched_val[it] = overall_item_unmatched_val.get( overall_item_unmatched_val[it] = overall_item_unmatched_val.get(
it, 0 it, 0
) + len(lst) ) + len(lst)
print("#" * 40) unmatched_val_details_by_item.setdefault(it, []).extend(lst)
print( log("#" * 40)
log(
f"{val_file.name}: matched {matched_total} | val {val_total} | answer {answer_total} " f"{val_file.name}: matched {matched_total} | val {val_total} | answer {answer_total} "
f"| unmatched val {unmatched_val_count} | unmatched answer {unmatched_answer_count} | precision {file_precision:.2%} | recall {file_recall:.2%} | f1 {file_f1:.2%} | false_positive_rate {file_false_positive_rate:.2%}" f"| unmatched val {unmatched_val_count} | unmatched answer {unmatched_answer_count} | precision {file_precision:.2%} | recall {file_recall:.2%} | f1 {file_f1:.2%} | false_positive_rate {file_false_positive_rate:.2%}"
) )
import json import json
print( log(
f"unmatched_val_by_item: {json.dumps(unmatched_val_by_item, ensure_ascii=False, indent=2)}" f"unmatched_val_by_item: {json.dumps(unmatched_val_by_item, ensure_ascii=False, indent=2)}"
) )
for item in sorted(answer_counts): for item in sorted(answer_counts):
item_matches = matched_by_item.get(item, []) item_matches = matched_by_item.get(item, [])
print( log(
f" 审查项 {item}: matched {len(item_matches)} / {answer_counts[item]}" f" 审查项 {item}: matched {len(item_matches)} / {answer_counts[item]}"
) )
# 匹配成功的结果 # 匹配成功的结果
...@@ -161,15 +213,15 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None: ...@@ -161,15 +213,15 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None:
ua = unmatched_answer_by_item.get(item, []) ua = unmatched_answer_by_item.get(item, [])
if ua: if ua:
print(f" 未匹配(answer 未被匹配){len(ua)} 条:") log(f" 未匹配(answer 未被匹配){len(ua)} 条:")
for t in ua: for t in ua:
print(f" answer: {t}") log(f" answer: {t}")
uv = unmatched_val_by_item.get(item, []) uv = unmatched_val_by_item.get(item, [])
if uv: if uv:
print(f" 未匹配(val 残留){len(uv)} 条:") log(f" 未匹配(val 残留){len(uv)} 条:")
for t in uv: for t in uv:
print(f" val: {t}") log(f" val: {t}")
# break # only first file for demo # break # only first file for demo
precision = overall_matched / overall_val if overall_val else 0 precision = overall_matched / overall_val if overall_val else 0
recall = overall_matched / overall_answer if overall_answer else 0 recall = overall_matched / overall_answer if overall_answer else 0
...@@ -177,14 +229,14 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None: ...@@ -177,14 +229,14 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None:
overall_false_positive_rate = ( overall_false_positive_rate = (
(overall_val - overall_matched) / overall_val if overall_val else 0 (overall_val - overall_matched) / overall_val if overall_val else 0
) )
print( log(
f"Overall: matched {overall_matched} | val {overall_val} | answer {overall_answer} | precision {precision:.2%} | recall {recall:.2%} | f1 {f1:.2%}" f"Overall: matched {overall_matched} | val {overall_val} | answer {overall_answer} | precision {precision:.2%} | recall {recall:.2%} | f1 {f1:.2%}"
) )
# 按“审查项”的 overall 结果 # 按“审查项”的 overall 结果
if overall_item_answer: if overall_item_answer:
print("#" * 40) log("#" * 40)
print("Overall by item:") log("Overall by item:")
all_items = sorted( all_items = sorted(
set( set(
list(overall_item_answer.keys()) list(overall_item_answer.keys())
...@@ -220,7 +272,7 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None: ...@@ -220,7 +272,7 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None:
"误报率(D/B+D)": item_false_positive_rate, "误报率(D/B+D)": item_false_positive_rate,
} }
) )
print( log(
f" 审查项 {it}: matched {mat} / answer {ans} | unmatched val {u_val} | unmatched answer {u_ans} | precision {item_precision:.2%} | recall {acc:.2%} | f1 {item_f1:.2%}" f" 审查项 {it}: matched {mat} / answer {ans} | unmatched val {u_val} | unmatched answer {u_ans} | precision {item_precision:.2%} | recall {acc:.2%} | f1 {item_f1:.2%}"
) )
...@@ -288,36 +340,45 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None: ...@@ -288,36 +340,45 @@ def _compare_impl(val_dir: Path, answer_dir: Path) -> None:
compare_dir_name = val_dir.name compare_dir_name = val_dir.name
results_dir = Path(__file__).parent / "results" results_dir = Path(__file__).parent / "results"
results_dir.mkdir(parents=True, exist_ok=True) results_dir.mkdir(parents=True, exist_ok=True)
output_excel = results_dir / f"合同审查结果_{compare_dir_name}.xlsx" output_excel = results_dir / f"{compare_dir_name}.xlsx"
used_sheet_names = {"对比结果"}
with pd.ExcelWriter(output_excel, engine="openpyxl") as writer: with pd.ExcelWriter(output_excel, engine="openpyxl") as writer:
combined_df.to_excel(writer, sheet_name="对比结果", index=False) combined_df.to_excel(writer, sheet_name="对比结果", index=False)
print( for item in all_items:
uncalled = unmatched_answer_details_by_item.get(item, [])
inaccurate = unmatched_val_details_by_item.get(item, [])
max_len = max(len(uncalled), len(inaccurate))
detail_df = pd.DataFrame(
{
"审查不合格但是没有查出来的": uncalled + [""] * (max_len - len(uncalled)),
"审查合格但是误判为不合格的": inaccurate + [""] * (max_len - len(inaccurate)),
}
)
detail_df.to_excel(
writer,
sheet_name=_safe_sheet_name(item, used_sheet_names),
index=False,
)
log(
f"Excel written to {output_excel}.\nEval Time: {time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}" f"Excel written to {output_excel}.\nEval Time: {time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}"
) )
return output_excel
def compare(val_dir: Path, answer_dir: Path) -> None: return None
_compare_impl(val_dir=val_dir, answer_dir=answer_dir)
def compare(
def compare_with_log( val_dir: Path,
val_dir: Path, answer_dir: Path, log_path: Path | None = None answer_dir: Path,
) -> Path: print_result: bool = False,
val_dir = val_dir.resolve() ignore_items: Iterable[str] | None = None,
if log_path is None: ) -> Path | None:
results_dir = Path(__file__).parent / "results" return _compare_impl(
results_dir.mkdir(parents=True, exist_ok=True) val_dir=val_dir,
log_path = results_dir / f"合同审查结果_{val_dir.name}.log" answer_dir=answer_dir,
else: print_result=print_result,
log_path = log_path.resolve() ignore_items=ignore_items,
log_path.parent.mkdir(parents=True, exist_ok=True) )
with open(log_path, "w", encoding="utf-8") as f, redirect_stdout(
f
), redirect_stderr(f):
_compare_impl(val_dir=val_dir, answer_dir=answer_dir)
return log_path
def _parse_args() -> argparse.Namespace: def _parse_args() -> argparse.Namespace:
...@@ -338,19 +399,26 @@ def _parse_args() -> argparse.Namespace: ...@@ -338,19 +399,26 @@ def _parse_args() -> argparse.Namespace:
help="Directory containing answer xlsx files.", help="Directory containing answer xlsx files.",
) )
parser.add_argument( parser.add_argument(
"--log-path", "--no-print",
type=Path, default=True,
default=None, help="Disable comparison progress and result printing.",
help="Optional explicit log path. Defaults to results/合同审查结果_<val_dir_name>.log", )
parser.add_argument(
"--ignore-items",
nargs="*",
default=['备料审查'],
help="Review item names to exclude from matching and statistics.",
) )
return parser.parse_args() return parser.parse_args()
if __name__ == "__main__": if __name__ == "__main__":
args = _parse_args() args = _parse_args()
final_log_path = compare_with_log( final_excel_path = compare(
val_dir=args.val_dir, val_dir=args.val_dir,
answer_dir=args.answer_dir, answer_dir=args.answer_dir,
log_path=args.log_path, print_result=not args.no_print,
ignore_items=args.ignore_items,
) )
print(f"Log written to {final_log_path}") if not args.no_print:
print(f"Excel written to {final_excel_path}")
...@@ -8,7 +8,7 @@ from typing import Iterable ...@@ -8,7 +8,7 @@ from typing import Iterable
import pandas as pd import pandas as pd
from spire.doc import Document from spire.doc import Document
from compare_annotation import compare_with_log from compare_annotation import compare
# Map raw comment authors to unified review item names. # Map raw comment authors to unified review item names.
COMMENT_AUTHOR_MAPPING: dict[str, str] = { COMMENT_AUTHOR_MAPPING: dict[str, str] = {
...@@ -86,8 +86,8 @@ def extract_annotaion( ...@@ -86,8 +86,8 @@ def extract_annotaion(
def compare_annotaion(val_dir: Path, answer_dir: Path) -> None: def compare_annotaion(val_dir: Path, answer_dir: Path) -> None:
"""Run benchmark comparison on extracted annotations.""" """Run benchmark comparison on extracted annotations."""
log_path = compare_with_log(val_dir=val_dir, answer_dir=answer_dir) output_excel = compare(val_dir=val_dir, answer_dir=answer_dir)
print(f"Compare log written to: {log_path}") print(f"Compare result written to: {output_excel}")
def _strip_suffix_once(stem: str, suffixes: Iterable[str]) -> str: def _strip_suffix_once(stem: str, suffixes: Iterable[str]) -> str:
...@@ -121,7 +121,7 @@ def _parse_args() -> argparse.Namespace: ...@@ -121,7 +121,7 @@ def _parse_args() -> argparse.Namespace:
parser.add_argument( parser.add_argument(
"--datasets-dir", "--datasets-dir",
type=Path, type=Path,
default=base / "results" / "jp-output-lufa-20260511-101828", default=base / "results" / "jp-output-lufa-20260514-182859",
help="Directory containing Word files with annotations.", help="Directory containing Word files with annotations.",
) )
parser.add_argument( parser.add_argument(
...@@ -143,7 +143,7 @@ if __name__ == "__main__": ...@@ -143,7 +143,7 @@ if __name__ == "__main__":
datasets_dir=args.datasets_dir, datasets_dir=args.datasets_dir,
answer_dir=base / "审查答案", answer_dir=base / "审查答案",
val_dir=args.datasets_dir.with_name( val_dir=args.datasets_dir.with_name(
f"{args.datasets_dir.name}-extract-comment" f"{args.datasets_dir.name}-测评结果"
), ),
strip_suffixes=args.strip_suffixes, strip_suffixes=args.strip_suffixes,
) )
from __future__ import annotations
import argparse
import asyncio
import json
import re
import shutil
import subprocess
import sys
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any
REPO_ROOT = Path(__file__).resolve().parents[1]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
BENCHMARK_DIR = REPO_ROOT / "data" / "benchmark"
if str(BENCHMARK_DIR) not in sys.path:
sys.path.insert(0, str(BENCHMARK_DIR))
from utils import excel_tool # noqa: E402
SUMMARY_SHEET = "对比结果"
FN_COL = "审查不合格但是没有查出来的"
FP_COL = "审查合格但是误判为不合格的"
F1_COL = "F1"
REVIEW_ITEM_COL = "审查项"
RULE_COL = "审查规则"
TRIGGER_COL = "触发词"
SUGGESTION_COL = "建议模板"
DEFAULT_RULES_FILE = REPO_ROOT / "data" / "rules.xlsx"
DEFAULT_REPORT_FILE = REPO_ROOT / "data" / "prompt_optimization_rounds.xlsx"
DEFAULT_MAX_WORKERS = 10
DEFAULT_APP_SESSION = "jp-contract-debug"
DEFAULT_APP_COMMAND = "python main.py"
@dataclass
class RuleLocation:
sheet: str
row_idx: int
row: dict[str, Any]
@dataclass
class EvalItem:
name: str
metrics: dict[str, Any]
false_negatives: list[str]
false_positives: list[str]
locations: list[RuleLocation]
@dataclass
class RuleTask:
item: EvalItem
location: RuleLocation
def clean_text(value: Any) -> str:
if value is None:
return ""
return str(value).strip()
def parse_float(value: Any, default: float = 0.0) -> float:
if value is None or value == "":
return default
if isinstance(value, (int, float)):
return float(value)
text = str(value).strip().rstrip("%")
try:
number = float(text)
except ValueError:
return default
return number / 100 if str(value).strip().endswith("%") else number
def compact_examples(values: list[str], limit: int) -> list[str]:
seen: set[str] = set()
examples: list[str] = []
for value in values:
text = clean_text(value)
if not text or text in seen:
continue
seen.add(text)
examples.append(text)
if len(examples) >= limit:
break
return examples
def safe_sheet_name(name: str, used: set[str]) -> str:
invalid_chars = set("[]:*?/\\")
base = "".join("_" if ch in invalid_chars else ch for ch in name)[:31] or "Sheet"
if base not in used:
used.add(base)
return base
idx = 2
while True:
suffix = f"_{idx}"
candidate = base[: 31 - len(suffix)] + suffix
if candidate not in used:
used.add(candidate)
return candidate
idx += 1
def json_dumps(value: Any) -> str:
return json.dumps(value, ensure_ascii=False, indent=2)
def parse_json_object(text: str) -> dict[str, Any]:
text = clean_text(text)
if not text:
return {}
try:
value = json.loads(text)
return value if isinstance(value, dict) else {}
except json.JSONDecodeError:
pass
fenced = re.search(r"```(?:json)?\s*(\{.*?\})\s*```", text, re.S)
if fenced:
try:
value = json.loads(fenced.group(1))
return value if isinstance(value, dict) else {}
except json.JSONDecodeError:
pass
blob = re.search(r"(\{.*\})", text, re.S)
if not blob:
return {}
try:
value = json.loads(blob.group(1))
except json.JSONDecodeError:
return {}
return value if isinstance(value, dict) else {}
def read_rule_locations(rules_file: Path) -> dict[str, list[RuleLocation]]:
locations: dict[str, list[RuleLocation]] = {}
for sheet in excel_tool.list_sheets(str(rules_file)):
rows = excel_tool.load_excel(str(rules_file), sheet=sheet, header=True)
for offset, row in enumerate(rows, start=2):
if not isinstance(row, dict):
continue
review_item = clean_text(row.get(REVIEW_ITEM_COL))
if not review_item:
continue
locations.setdefault(review_item, []).append(
RuleLocation(sheet=sheet, row_idx=offset, row=row)
)
return locations
def read_eval_items(
eval_excel: Path,
rules_file: Path,
f1_threshold: float,
example_limit: int,
) -> list[EvalItem]:
rule_locations = read_rule_locations(rules_file)
metric_rows = excel_tool.load_excel(str(eval_excel), sheet=SUMMARY_SHEET, header=True)
metrics_by_item = {
clean_text(row.get(REVIEW_ITEM_COL)): row
for row in metric_rows
if isinstance(row, dict) and clean_text(row.get(REVIEW_ITEM_COL))
}
items: list[EvalItem] = []
for sheet in excel_tool.list_sheets(str(eval_excel)):
if sheet == SUMMARY_SHEET:
continue
rows = excel_tool.load_excel(str(eval_excel), sheet=sheet, header=True)
false_negatives: list[str] = []
false_positives: list[str] = []
for row in rows:
if not isinstance(row, dict):
continue
fn = clean_text(row.get(FN_COL))
fp = clean_text(row.get(FP_COL))
if fn:
false_negatives.append(fn)
if fp:
false_positives.append(fp)
metrics = metrics_by_item.get(sheet, {})
f1 = parse_float(metrics.get(F1_COL), default=0.0) if metrics else 0.0
has_errors = bool(false_negatives or false_positives)
if not has_errors or f1 >= f1_threshold:
continue
locations = rule_locations.get(sheet, [])
if not locations:
continue
items.append(
EvalItem(
name=sheet,
metrics=dict(metrics),
false_negatives=compact_examples(false_negatives, example_limit),
false_positives=compact_examples(false_positives, example_limit),
locations=locations,
)
)
return items
def build_rule_tasks(items: list[EvalItem]) -> list[RuleTask]:
return [
RuleTask(item=item, location=location)
for item in items
for location in item.locations
]
def build_messages(task: RuleTask) -> list[dict[str, str]]:
item = task.item
loc = task.location
rule = {
"ID": loc.row.get("ID"),
"摘要项": loc.row.get("摘要项"),
REVIEW_ITEM_COL: loc.row.get(REVIEW_ITEM_COL),
RULE_COL: loc.row.get(RULE_COL),
"风险等级": loc.row.get("风险等级"),
SUGGESTION_COL: loc.row.get(SUGGESTION_COL),
TRIGGER_COL: loc.row.get(TRIGGER_COL),
}
payload = {
"审查项": item.name,
"本轮指标": item.metrics,
"当前规则": rule,
"漏检样例_应判不合格但未查出": item.false_negatives,
"误报样例_合格但误判不合格": item.false_positives,
"任务": (
"请优化该审查项的规则。漏检样例用于增强召回,误报样例用于补充合格边界或排除条件。"
"不要过拟合样例原文,要抽象为可泛化的合同审查规则。"
),
}
system_prompt = """
你是合同审查规则优化专家。你只输出 JSON,不输出解释性文字。
目标:
1. 基于漏检样例补充不合格判定条件,提高查全率。
2. 基于误报样例补充合格边界、排除条件和适用前提,提高查准率。
3. 保留原规则中仍然有效的核心约束,不要为了消除误报删除应检风险。
输出 JSON Schema:
{
"should_optimize": true,
"updated_rule": "优化后的审查规则完整文本",
"reason": "简要说明本次优化依据"
}
如果样例不足以支持修改,返回 should_optimize=false,并保持 updated_rule 和 reason 为空字符串。
""".strip()
return [
{"role": "system", "content": system_prompt},
{"role": "user", "content": json_dumps(payload)},
]
async def optimize_rule_tasks(tasks: list[RuleTask], max_workers: int) -> list[dict[str, Any]]:
if not tasks:
return []
from core.config import LLM
from utils.openai_util import OpenAITool
tool = OpenAITool(LLM["base_tool_llm"], max_workers=max_workers)
responses = await tool.mul_chat([build_messages(task) for task in tasks])
results: list[dict[str, Any]] = []
for task, response in zip(tasks, responses):
parsed = parse_json_object(response)
results.append(
{
"task": task,
"raw_response": response,
"parsed": parsed,
}
)
return results
def backup_rules(rules_file: Path, output_dir: Path, round_idx: int) -> Path:
output_dir.mkdir(parents=True, exist_ok=True)
backup = output_dir / f"{rules_file.stem}.round{round_idx}.{time.strftime('%Y%m%d-%H%M%S')}{rules_file.suffix}"
shutil.copy2(rules_file, backup)
return backup
def update_rules_file(rules_file: Path, optimizations: list[dict[str, Any]]) -> list[dict[str, Any]]:
from openpyxl import load_workbook
wb = load_workbook(rules_file)
applied: list[dict[str, Any]] = []
for optimization in optimizations:
task: RuleTask = optimization["task"]
item = task.item
loc = task.location
parsed = optimization["parsed"]
if not parsed.get("should_optimize"):
continue
updated_rule = clean_text(parsed.get("updated_rule"))
if not updated_rule:
continue
ws = wb[loc.sheet]
headers = [clean_text(cell.value) for cell in ws[1]]
header_map = {header: idx + 1 for idx, header in enumerate(headers) if header}
if RULE_COL not in header_map:
continue
old_rule = clean_text(ws.cell(loc.row_idx, header_map[RULE_COL]).value)
ws.cell(loc.row_idx, header_map[RULE_COL], updated_rule)
applied.append(
{
"审查项": item.name,
"规则sheet": loc.sheet,
"规则行": loc.row_idx,
"旧审查规则": old_rule,
"新审查规则": updated_rule,
"优化理由": clean_text(parsed.get("reason")),
}
)
if applied:
wb.save(rules_file)
wb.close()
return applied
def run_batch() -> Path:
from data.batch import batch as batch_module
params = batch_module.get_params()
batch_module.execute_batch(params)
return Path(params["batch_output_dir_path"]).resolve()
def restart_app(session: str, command: str) -> None:
result = subprocess.run(
["tmux", "has-session", "-t", session],
cwd=REPO_ROOT,
text=True,
capture_output=True,
check=False,
)
if result.returncode != 0:
raise RuntimeError(f"tmux session not found: {session}")
subprocess.run(["tmux", "send-keys", "-t", session, "C-c"], cwd=REPO_ROOT, check=True)
time.sleep(2)
subprocess.run(
["tmux", "send-keys", "-t", session, f"cd {shell_quote(str(REPO_ROOT))} && {command}", "C-m"],
cwd=REPO_ROOT,
check=True,
)
time.sleep(5)
def shell_quote(value: str) -> str:
return "'" + value.replace("'", "'\"'\"'") + "'"
def run_eval(batch_output_dir: Path) -> Path:
from data.benchmark import eval as eval_module
val_dir = batch_output_dir.with_name(f"{batch_output_dir.name}-测评结果")
eval_module.eval(
datasets_dir=batch_output_dir,
answer_dir=BENCHMARK_DIR / "审查答案",
val_dir=val_dir,
strip_suffixes=["_麓发改进", "_人机交互", "_麓发迁移"],
)
return BENCHMARK_DIR / "results" / f"{val_dir.name}.xlsx"
def append_round_report(
report_file: Path,
round_idx: int,
batch_output_dir: Path,
eval_excel: Path,
backup_file: Path | None,
eval_items: list[EvalItem],
optimizations: list[dict[str, Any]],
applied: list[dict[str, Any]],
f1_threshold: float,
) -> None:
from openpyxl import Workbook, load_workbook
from openpyxl.styles import Alignment, Font
report_file.parent.mkdir(parents=True, exist_ok=True)
wb = load_workbook(report_file) if report_file.exists() else Workbook()
if wb.active.title == "Sheet" and wb.active.max_row == 1 and wb.active["A1"].value is None:
wb.remove(wb.active)
used = set(wb.sheetnames)
sheet_name = safe_sheet_name(f"round_{round_idx}", used)
ws = wb.create_sheet(sheet_name)
headers = [
"轮次",
"时间",
"F1阈值",
"batch输出目录",
"eval Excel",
"规则备份",
"审查项",
"F1",
"查准率(B/B+D)",
"查全率(B/C)",
"误报率(D/B+D)",
"漏检样例数",
"误报样例数",
"是否优化",
"是否写回",
"规则位置",
"旧审查规则",
"新审查规则",
"优化理由",
"模型原始输出",
]
ws.append(headers)
for cell in ws[1]:
cell.font = Font(bold=True)
applied_by_item = {(row["审查项"], row["规则sheet"], row["规则行"]): row for row in applied}
optimization_by_location = {
(opt["task"].item.name, opt["task"].location.sheet, opt["task"].location.row_idx): opt
for opt in optimizations
}
if not eval_items:
ws.append(
[
round_idx,
time.strftime("%Y-%m-%d %H:%M:%S"),
f1_threshold,
str(batch_output_dir),
str(eval_excel),
str(backup_file or ""),
"",
"",
"",
"",
"",
0,
0,
False,
False,
"",
"",
"",
"所有审查项均达到阈值或无可优化错误",
"",
]
)
for item in eval_items:
for loc in item.locations:
opt = optimization_by_location.get((item.name, loc.sheet, loc.row_idx), {})
parsed = opt.get("parsed", {})
raw_response = clean_text(opt.get("raw_response"))
applied_row = applied_by_item.get((item.name, loc.sheet, loc.row_idx), {})
ws.append(
[
round_idx,
time.strftime("%Y-%m-%d %H:%M:%S"),
f1_threshold,
str(batch_output_dir),
str(eval_excel),
str(backup_file or ""),
item.name,
item.metrics.get(F1_COL),
item.metrics.get("查准率(B/B+D)"),
item.metrics.get("查全率(B/C)"),
item.metrics.get("误报率(D/B+D)"),
len(item.false_negatives),
len(item.false_positives),
bool(parsed.get("should_optimize")),
bool(applied_row),
f"{loc.sheet}!{loc.row_idx}",
applied_row.get("旧审查规则", loc.row.get(RULE_COL)),
applied_row.get("新审查规则", parsed.get("updated_rule", "")),
applied_row.get("优化理由", parsed.get("reason", "")),
raw_response,
]
)
for row in ws.iter_rows():
for cell in row:
cell.alignment = Alignment(vertical="top", wrap_text=True)
wb.save(report_file)
wb.close()
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Loop batch/eval and optimize data/rules.xlsx with LLM.")
parser.add_argument("--rules-file", type=Path, default=DEFAULT_RULES_FILE)
parser.add_argument("--report-file", type=Path, default=DEFAULT_REPORT_FILE)
parser.add_argument("--rounds", type=int, default=1)
parser.add_argument("--f1-threshold", type=float, default=0.9)
parser.add_argument("--example-limit", type=int, default=20)
parser.add_argument("--max-workers", type=int, default=DEFAULT_MAX_WORKERS)
parser.add_argument("--app-session", default=DEFAULT_APP_SESSION)
parser.add_argument("--app-command", default=DEFAULT_APP_COMMAND)
parser.add_argument(
"--skip-app-restart",
action="store_true",
help="Do not restart tmux app session before running batch.py.",
)
parser.add_argument(
"--batch-output",
type=Path,
help="Use an existing batch output dir for the first round instead of running batch.py.",
)
parser.add_argument(
"--dry-run",
action="store_true",
help="Run eval and LLM optimization but do not write updates to rules.xlsx.",
)
return parser.parse_args()
def main() -> int:
args = parse_args()
rules_file = args.rules_file.resolve()
report_file = args.report_file.resolve()
for round_idx in range(1, args.rounds + 1):
# 1. 重启程序 跑批量
if round_idx == 1 and args.batch_output:
batch_output_dir = args.batch_output.resolve()
else:
if not args.skip_app_restart:
restart_app(args.app_session, args.app_command)
batch_output_dir = run_batch()
# 2. 测评结果分析,生成优化建议,写回规则文件,记录优化报告
eval_excel = run_eval(batch_output_dir)
eval_items = read_eval_items(
eval_excel=eval_excel,
rules_file=rules_file,
f1_threshold=args.f1_threshold,
example_limit=args.example_limit,
)
rule_tasks = build_rule_tasks(eval_items)
optimizations = asyncio.run(optimize_rule_tasks(rule_tasks, max_workers=args.max_workers))
backup_file: Path | None = None
applied: list[dict[str, Any]] = []
if optimizations and not args.dry_run:
backup_file = backup_rules(rules_file, report_file.parent / "rule_backups", round_idx)
applied = update_rules_file(rules_file, optimizations)
append_round_report(
report_file=report_file,
round_idx=round_idx,
batch_output_dir=batch_output_dir,
eval_excel=eval_excel,
backup_file=backup_file,
eval_items=eval_items,
optimizations=optimizations,
applied=applied,
f1_threshold=args.f1_threshold,
)
print(
f"round={round_idx} batch_output={batch_output_dir} eval_excel={eval_excel} "
f"optimized_items={len(eval_items)} optimized_rules={len(rule_tasks)} "
f"applied_updates={len(applied)} report={report_file}",
flush=True,
)
if not eval_items:
break
return 0
if __name__ == "__main__":
raise SystemExit(main())
No preview for this file type
---
name: review-prompt-optimization-skill
description: 审查提示词/规则优化编排 Skill。通过 batch 批处理、benchmark eval、测评 Excel 归因,循环更新 data/rules.xlsx 中的审查规则。
---
# Review Prompt Optimization Skill
## 定位
`review-prompt-optimization-skill` 用于编排合同审查提示词和规则库优化闭环。
本 Skill 负责定义流程、判读方法、规则更新约束和循环停止条件。使用时,LLM 需要结合可用终端、Excel 工具和代码编辑能力执行流程;不要只给建议。
默认优化目标文件:
- 规则库:`data/rules.xlsx`
- 批处理脚本:`data/batch/batch.py`
- 测评脚本:`data/benchmark/eval.py`
- Excel 工具:`skills/doc-excel-skill/scripts/excel_tool.py`
- 编排辅助脚本:`data/optimize_review_prompt_flow.py`
## 标准闭环
1. 启动批处理
- 优先进入或使用 tmux session:`tmux attach -t batch`
- 在该 session 中运行:`python3 data/batch/batch.py`
- 也可直接使用 `data/optimize_review_prompt_flow.py` 串联执行。
2. 等待 batch 完成
- 每 5 分钟检查一次 tmux 输出。
- 完成标志是 `batch.py` 输出批处理结果路径,格式通常为:`文件保存在: /abs/path/to/output`
- 记录该路径为 `batch_output_dir`
3. 运行测评
- 执行:
```bash
python3 data/benchmark/eval.py --datasets-dir "$batch_output_dir"
```
- `eval.py` 会先抽取批注,再对比 `data/benchmark/审查答案`,并输出测评 Excel。
- 测评 Excel 通常位于:`data/benchmark/results/{batch输出目录名}-测评结果.xlsx`
4. 读取测评 Excel
- 先读取全部 sheet:
```bash
python3 skills/doc-excel-skill/scripts/excel_tool.py list-sheets "$eval_excel"
```
- 跳过汇总 sheet `对比结果`,遍历其余每个审查项 sheet。
- 每个审查项 sheet 有两类错误:
- 第一列 `审查不合格但是没有查出来的`:漏检,应增强召回。
- 第二列 `审查合格但是误判为不合格的`:误报,应收紧边界。
5. 更新规则库
- 更新目标是 `data/rules.xlsx` 中与 sheet 名或列 `审查项` 对应的规则行。
- 更新前必须备份,例如:
```bash
cp data/rules.xlsx "data/rules.xlsx.bak.$(date +%Y%m%d-%H%M%S).xlsx"
```
- 对漏检样例:
-`审查规则` 中补充该类风险的明确判定条件。
- 必要时补充 `触发词`,提高路由和规则命中。
- 如果样例反映的是规则缺项,而非单条规则表述不足,可新增或拆分审查项。
- 对误报样例:
-`审查规则` 中补充排除条件、合格边界、适用前提或反例。
- 必要时在 `建议模板` 中约束“满足某条件时无需修改”。
- 不要为了消除误报而删除应检风险的核心判定。
- 对同一审查项同时存在漏检和误报时:
- 先抽象共同原因,再修改规则。
- 避免简单堆砌样例,优先写可泛化的规则边界。
6. 循环验证
- 保存 `data/rules.xlsx` 后,再回到步骤 1 执行下一轮 batch/eval。
- 每轮记录:
- batch 输出目录
- eval Excel 路径
- 修改过的 sheet/审查项
- 漏检数量、误报数量、F1、查全率、查准率变化
- 当总体 F1 或目标审查项指标不再明显提升,或剩余错误已无法通过规则库优化稳定解决时停止。
## 辅助脚本
可使用辅助脚本完成“batch + eval + 规则优化 + 每轮结果落 Excel”:
```bash
python data/optimize_review_prompt_flow.py --rounds 1 --f1-threshold 0.9
```
默认行为:
- 调用 `data/batch/batch.py` 生成批处理输出目录。
- 调用 `data/benchmark/eval.py` 生成测评 Excel。
- 使用 `utils/excel_tool.py` 读取 `data/rules.xlsx` 和测评 Excel。
- 对 F1 低于阈值且存在漏检/误报的审查项,调用 `utils/openai_util.py` 批量优化规则。
- 写回 `data/rules.xlsx` 前自动备份。
- 每轮结果保存到 `data/benchmark/results/prompt_optimization_rounds.xlsx` 的新 sheet。
如果 batch 已经跑完,可跳过 tmux:
```bash
python data/optimize_review_prompt_flow.py \
--batch-output /abs/path/to/batch-output
```
常用参数:
- `--rounds 3`
- `--f1-threshold 0.9`
- `--example-limit 20`
- `--max-workers 10`
- `--dry-run`
## Excel 判读规则
### `对比结果`
用于确定优化优先级:
- 优先处理 `大模型未匹配上的不合格项(C-B)` 高的审查项,提高召回。
- 其次处理 `大模型其他不合格项``误报率(D/B+D)` 高的审查项,提高准确性。
- 同等情况下优先处理合同风险更高、出现频率更高的审查项。
### 审查项明细 sheet
每个明细 sheet 的 sheet 名即审查项名称。
- 第一列非空样例是“应该判不合格,但没有输出”的证据文本。
- 第二列非空样例是“实际合格,但被模型判不合格”的证据文本。
- 修改规则时,必须把样例上升为抽象规则,不要只复制样例原文。
## 更新 `data/rules.xlsx` 的约束
- 保留原有 sheet、列名、ID 和风险等级,除非确实需要新增规则。
- 优先更新与测评 sheet 同名或语义对应的 `审查项` 行。
- 只能改与本轮错误归因直接相关的规则,不做无关重构。
- 修改应同时兼顾召回和准确性,避免单向拉高导致另一类错误恶化。
- 每次修改后都要保存 Excel,并在最终回复中说明改了哪些审查项。
- 如果工具无法安全写入 Excel,应先导出修改建议 JSON/Markdown,不得破坏原文件。
## 推荐规则改写格式
`审查规则` 中优先使用以下结构:
```text
检查……。不合格情形包括:1)……;2)……。
合格/不判定为不合格的情形包括:1)……;2)……。
仅当合同原文明确体现……时输出不合格;不得因……直接推断为不合格。
```
`触发词` 中使用分号分隔:
```text
付款;支付;验收;尾款;发票
```
`建议模板` 中写可执行建议:
```text
如存在该风险,建议补充……;若合同已明确……则无需修改。
```
## 输出要求
每轮优化结束后,LLM 应输出:
- batch 输出目录
- eval Excel 路径
- 本轮修改的 `data/rules.xlsx` 备份路径
- 修改过的审查项列表
- 每个审查项的漏检/误报归因摘要
- 是否建议继续下一轮
## 注意事项
- 不要把 `对比结果` 当成规则明细 sheet。
- 不要把“审查合格但是误判为不合格的”样例加入不合格判定条件;它们用于写排除条件。
- 不要因单个样例过拟合规则,应提炼可泛化的合同审查边界。
- 批处理和测评耗时较长,轮询期间保持 tmux session 不被关闭。
#!/usr/bin/env python3
"""Standalone Excel CLI for table reads, row edits, and JSON sheet writes."""
from __future__ import annotations
import argparse
import csv
import json
import string
import zipfile
from pathlib import Path
from typing import Any
from xml.etree import ElementTree as ET
NS = {"a": "http://schemas.openxmlformats.org/spreadsheetml/2006/main", "r": "http://schemas.openxmlformats.org/officeDocument/2006/relationships", "rel": "http://schemas.openxmlformats.org/package/2006/relationships"}
class ExcelLoadError(Exception):
pass
def _json(v: Any) -> None:
print(json.dumps(v, ensure_ascii=False, indent=2))
def _load_json(v: str) -> Any:
return json.loads(Path(v[1:]).read_text(encoding="utf-8") if v.startswith("@") else v)
def _col_idx(ref: str) -> int:
n = 0
for ch in "".join(c for c in ref if c in string.ascii_letters).upper():
n = n * 26 + ord(ch) - 64
return max(n - 1, 0)
def _rows_to_result(rows: list, header: bool) -> list:
if not rows:
return []
if not header:
return [list(r) for r in rows]
heads = [str(h).strip() if h is not None else "" for h in rows[0]]
return [{heads[i] if i < len(heads) else f"col{i}": row[i] for i in range(len(row))} for row in rows[1:]]
def _sheet_map(zf: zipfile.ZipFile) -> list[tuple[str, str]]:
wb = ET.fromstring(zf.read("xl/workbook.xml"))
rels = ET.fromstring(zf.read("xl/_rels/workbook.xml.rels"))
rel_map = {r.attrib["Id"]: r.attrib["Target"] for r in rels.findall("rel:Relationship", NS)}
out = []
for s in wb.findall(".//a:sheets/a:sheet", NS):
target = rel_map.get(s.attrib.get(f"{{{NS['r']}}}id", ""), "")
out.append((s.attrib.get("name", ""), "xl/" + target.lstrip("/") if not target.startswith("xl/") else target))
return out
def _shared(zf: zipfile.ZipFile) -> list[str]:
try:
root = ET.fromstring(zf.read("xl/sharedStrings.xml"))
except KeyError:
return []
return ["".join(t.text or "" for t in item.findall(".//a:t", NS)) for item in root.findall(".//a:si", NS)]
def _load_std_xlsx(path: Path, sheet: str | None, header: bool) -> list:
with zipfile.ZipFile(path) as zf:
shared, sheets = _shared(zf), _sheet_map(zf)
if not sheets:
return []
sheet_path = next((p for n, p in sheets if n == sheet), sheets[0][1])
root = ET.fromstring(zf.read(sheet_path))
rows = []
for r in root.findall(".//a:sheetData/a:row", NS):
values = []
for c in r.findall("a:c", NS):
while len(values) < _col_idx(c.attrib.get("r", "")):
values.append(None)
raw = (c.find("a:v", NS).text if c.find("a:v", NS) is not None else None)
values.append(shared[int(raw)] if c.attrib.get("t") == "s" and raw is not None and int(raw) < len(shared) else raw)
rows.append(values)
return _rows_to_result(rows, header)
def load_excel(path: str, sheet: str | None = None, header: bool = True) -> list:
p = Path(path)
if p.suffix.lower() in {".csv", ".tsv"}:
with p.open(newline="", encoding="utf-8-sig", errors="replace") as f:
return _rows_to_result(list(csv.reader(f, delimiter="\t" if p.suffix.lower() == ".tsv" else ",")), header)
try:
import openpyxl # type: ignore
wb = openpyxl.load_workbook(p, data_only=True, read_only=True)
ws = wb[sheet] if sheet else wb.active
return _rows_to_result(list(ws.iter_rows(values_only=True)), header)
except ImportError:
if p.suffix.lower() != ".xlsx":
raise ExcelLoadError("openpyxl is required for non-xlsx files")
return _load_std_xlsx(p, sheet, header)
def list_sheets(path: str) -> list[str]:
try:
import openpyxl # type: ignore
return openpyxl.load_workbook(path, read_only=True).sheetnames
except ImportError:
with zipfile.ZipFile(path) as zf:
return [n for n, _ in _sheet_map(zf)]
def _load_workbook_for_write(path: str):
try:
import openpyxl # type: ignore
except ImportError as exc:
raise ExcelLoadError("openpyxl is required for write operations") from exc
return openpyxl.load_workbook(path)
def _get_sheet(wb: Any, sheet: str | None):
return wb[sheet] if sheet else wb.active
def _headers(ws: Any) -> list[str]:
return [str(cell.value).strip() if cell.value is not None else "" for cell in ws[1]]
def _header_map(ws: Any) -> dict[str, int]:
return {header: idx for idx, header in enumerate(_headers(ws), start=1) if header}
def _ensure_header_columns(ws: Any, keys: list[str]) -> dict[str, int]:
header_map = _header_map(ws)
next_col = ws.max_column + 1
for key in keys:
if key in header_map:
continue
ws.cell(row=1, column=next_col, value=key)
header_map[key] = next_col
next_col += 1
return header_map
def _row_dict_from_ws(ws: Any, row_idx: int, headers: list[str]) -> dict[str, Any]:
return {
header: ws.cell(row=row_idx, column=col_idx).value
for col_idx, header in enumerate(headers, start=1)
if header
}
def rows_as_dicts(path: str, sheet: str | None = None) -> list[dict[str, Any]]:
return search_rows(path, sheet, {})
def _require_dict(value: Any, name: str = "row") -> dict[str, Any]:
if not isinstance(value, dict):
raise ExcelLoadError(f"{name} must be a JSON object")
return value
def append_row(path: str, sheet: str | None, row_data: dict[str, Any]) -> dict[str, Any]:
wb = _load_workbook_for_write(path)
ws = _get_sheet(wb, sheet)
header_map = _ensure_header_columns(ws, [str(key) for key in row_data.keys()])
row_idx = ws.max_row + 1
for key, value in row_data.items():
ws.cell(row=row_idx, column=header_map[str(key)], value=value)
wb.save(path)
return {"file": path, "sheet": ws.title, "row": row_idx, "inserted": row_data}
def _matched_row_indices(ws: Any, criteria: dict[str, Any]) -> list[int]:
header_map = _header_map(ws)
missing_keys = [key for key in criteria if key not in header_map]
if missing_keys:
raise ExcelLoadError(f"criteria keys not found in header: {', '.join(missing_keys)}")
matched: list[int] = []
for row_idx in range(2, ws.max_row + 1):
if all(ws.cell(row=row_idx, column=header_map[key]).value == value for key, value in criteria.items()):
matched.append(row_idx)
return matched
def search_rows(path: str, sheet: str | None, criteria: dict[str, Any]) -> list[dict[str, Any]]:
rows = [row for row in load_excel(path, sheet=sheet, header=True) if isinstance(row, dict)]
if not criteria:
return rows
if rows:
missing_keys = [key for key in criteria if key not in rows[0]]
if missing_keys:
raise ExcelLoadError(f"criteria keys not found in header: {', '.join(missing_keys)}")
return [row for row in rows if all(row.get(key) == value for key, value in criteria.items())]
def update_rows(path: str, sheet: str | None, criteria: dict[str, Any], row_data: dict[str, Any]) -> dict[str, Any]:
if not criteria:
raise ExcelLoadError("criteria must not be empty for update_rows")
wb = _load_workbook_for_write(path)
ws = _get_sheet(wb, sheet)
row_indices = _matched_row_indices(ws, criteria)
if not row_indices:
return {"file": path, "sheet": ws.title, "updated": False, "rows": []}
update_keys = [str(key) for key in row_data.keys()]
header_map = _ensure_header_columns(ws, update_keys)
for row_idx in row_indices:
for key, value in row_data.items():
key = str(key)
ws.cell(row=row_idx, column=header_map[key], value=value)
wb.save(path)
return {"file": path, "sheet": ws.title, "updated": True, "rows": row_indices, "count": len(row_indices)}
def delete_rows(path: str, sheet: str | None, criteria: dict[str, Any]) -> dict[str, Any]:
if not criteria:
raise ExcelLoadError("criteria must not be empty for delete_rows")
wb = _load_workbook_for_write(path)
ws = _get_sheet(wb, sheet)
row_indices = _matched_row_indices(ws, criteria)
if not row_indices:
return {"file": path, "sheet": ws.title, "deleted": False, "rows": []}
for row_idx in sorted(row_indices, reverse=True):
ws.delete_rows(row_idx, 1)
wb.save(path)
return {"file": path, "sheet": ws.title, "deleted": True, "rows": row_indices, "count": len(row_indices)}
def _cell(v: Any) -> str:
return json.dumps(v, ensure_ascii=False, indent=2) if isinstance(v, (dict, list)) else ("" if v is None else str(v))
def _json_rows(data: Any) -> tuple[list[str], list[list[Any]]]:
if isinstance(data, dict):
headers = [str(key) for key in data.keys()]
return headers, [[data[key] for key in data.keys()]]
if not isinstance(data, list):
return ["value"], [[data]]
if not data:
return [], []
if all(isinstance(item, dict) for item in data):
headers: list[str] = []
for item in data:
for key in item.keys():
key = str(key)
if key not in headers:
headers.append(key)
rows = [[item.get(header) for header in headers] for item in data]
return headers, rows
if all(isinstance(item, (list, tuple)) for item in data):
max_len = max(len(item) for item in data)
headers = [f"col{idx + 1}" for idx in range(max_len)]
rows = [list(item) + [None] * (max_len - len(item)) for item in data]
return headers, rows
return ["value"], [[item] for item in data]
def json_to_sheet(data: Any, out: str, sheet: str = "Sheet1") -> str:
try:
from openpyxl import Workbook, load_workbook # type: ignore
from openpyxl.styles import Alignment, Font # type: ignore
except ImportError as exc:
raise ExcelLoadError("openpyxl is required for json-to-sheet") from exc
output_path = Path(out)
wb = load_workbook(output_path) if output_path.exists() else Workbook()
if sheet in wb.sheetnames:
old_sheet = wb[sheet]
old_index = wb.sheetnames.index(sheet)
wb.remove(old_sheet)
ws = wb.create_sheet(sheet, old_index)
else:
ws = wb.active if wb.active.title == "Sheet" and wb.active.max_row == 1 and wb.active.max_column == 1 and wb.active["A1"].value is None else wb.create_sheet(sheet)
ws.title = sheet
headers, rows = _json_rows(data)
for col_idx, header in enumerate(headers, start=1):
cell = ws.cell(row=1, column=col_idx, value=header)
cell.font = Font(bold=True)
for row_idx, row in enumerate(rows, start=2):
for col_idx, value in enumerate(row, start=1):
ws.cell(row=row_idx, column=col_idx, value=_cell(value))
if headers or rows:
for row in ws.iter_rows(
min_row=1,
max_row=max(1, len(rows) + 1),
max_col=max(1, len(headers)),
):
for cell in row:
cell.alignment = Alignment(vertical="top", wrap_text=True)
output_path.parent.mkdir(parents=True, exist_ok=True)
wb.save(output_path)
return str(output_path)
def main() -> int:
p = argparse.ArgumentParser(description="Standalone Excel CLI"); sub = p.add_subparsers(dest="cmd", required=True)
a = sub.add_parser("load-excel"); a.add_argument("file"); a.add_argument("--sheet-name"); a.add_argument("--no-header", action="store_true")
a = sub.add_parser("list-sheets"); a.add_argument("file")
a = sub.add_parser("search_rows"); a.add_argument("file"); a.add_argument("criteria", nargs="?", default="{}"); a.add_argument("--sheet-name")
a = sub.add_parser("append-row"); a.add_argument("file"); a.add_argument("row"); a.add_argument("--sheet-name")
a = sub.add_parser("update-rows"); a.add_argument("file"); a.add_argument("criteria"); a.add_argument("row"); a.add_argument("--sheet-name")
a = sub.add_parser("delete-rows"); a.add_argument("file"); a.add_argument("criteria"); a.add_argument("--sheet-name")
a = sub.add_parser("find-value"); a.add_argument("file"); a.add_argument("key_column"); a.add_argument("key_value"); a.add_argument("value_column"); a.add_argument("--sheet-name")
a = sub.add_parser("map-rows"); a.add_argument("file"); a.add_argument("column_map"); a.add_argument("--sheet-name")
a = sub.add_parser("json-to-sheet"); a.add_argument("json_data"); a.add_argument("output"); a.add_argument("--sheet-name", default="Sheet1")
x = p.parse_args()
if x.cmd == "load-excel": _json(load_excel(x.file, x.sheet_name, not x.no_header))
elif x.cmd == "list-sheets": _json(list_sheets(x.file))
elif x.cmd == "search_rows": _json(search_rows(x.file, x.sheet_name, _require_dict(_load_json(x.criteria), "criteria")))
elif x.cmd == "append-row": _json(append_row(x.file, x.sheet_name, _require_dict(_load_json(x.row))))
elif x.cmd == "update-rows": _json(update_rows(x.file, x.sheet_name, _require_dict(_load_json(x.criteria), "criteria"), _require_dict(_load_json(x.row))))
elif x.cmd == "delete-rows": _json(delete_rows(x.file, x.sheet_name, _require_dict(_load_json(x.criteria), "criteria")))
elif x.cmd == "find-value": _json(next((r.get(x.value_column) for r in load_excel(x.file, x.sheet_name) if isinstance(r, dict) and r.get(x.key_column) == x.key_value), None))
elif x.cmd == "map-rows": _json([{k: r.get(v) for k, v in json.loads(x.column_map).items()} for r in load_excel(x.file, x.sheet_name) if isinstance(r, dict)])
elif x.cmd == "json-to-sheet": print(json_to_sheet(_load_json(x.json_data), x.output, x.sheet_name))
return 0
if __name__ == "__main__":
raise SystemExit(main())
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or sign in to comment