Files
playbook/skills/tsl-syntax-reference/scripts/lookup.py
T

435 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
import argparse
import re
import sys
import unicodedata
from dataclasses import dataclass
from pathlib import Path
from urllib.parse import unquote
SKILL_ROOT = Path(__file__).resolve().parents[1]
DEFAULT_REFERENCES_DIR = SKILL_ROOT / "references"
HEADING_RE = re.compile(r"^(#{1,6})(?!#)\s+(.+?)\s*$")
FENCE_RE = re.compile(r"^```([^`]*)$")
MARKDOWN_LINK_RE = re.compile(r"!?\[[^\]]*\]\(([^)]+)\)")
ASCII_TOKEN_RE = re.compile(r"[a-z_][a-z0-9_.$:+-]*", re.IGNORECASE)
CHINESE_RUN_RE = re.compile(r"[\u3400-\u9fff]+")
IDENTITY_PREFIX = "代码块身份:"
ALLOWED_IDENTITIES = {
"可直接照写示例",
"反例 / 不可照写",
"输出片段",
"配置片段 / 概念骨架",
"仅服务端可执行示例",
}
ROUTER_PHRASES = ("路由中心", "选择一个主专题", "候选页继续判断")
EXCLUDED_REFERENCE_FILES = {"index.md"}
TOKEN_IN_HEADING_SCORE = 12
TOKEN_IN_BODY_SCORE = 3
EXACT_QUERY_SCORE = 20
PAGE_NAME_SCORE = 5
DIRECT_EXAMPLE_BOOST = 8
PITFALL_PAGE_BOOST = 18
COUNTEREXAMPLE_BOOST = 8
EXACT_ERROR_BOOST = 14
MIXED_QUERY_MIN_SCORE = 10
@dataclass(frozen=True)
class Section:
id: str
page: Path
page_title: str
heading_path: tuple[str, ...]
body: str
identities: tuple[str, ...]
searchable_text: str
@dataclass(frozen=True)
class ValidationProblem:
page: Path
line: int
message: str
@dataclass(frozen=True)
class QueryMatch:
section: Section
score: int
@dataclass
class QueryResult:
query: str
mode: str
matches: list[QueryMatch]
prelude: list[Section]
def normalize(value: str) -> str:
return unicodedata.normalize("NFKC", value).casefold()
def _slug(value: str) -> str:
normalized = normalize(value)
slug = re.sub(r"[^\w]+", "-", normalized, flags=re.UNICODE).strip("-_")
return slug or "section"
def section_id(relative_page: Path | str, heading_path: tuple[str, ...]) -> str:
page = Path(relative_page)
parts = [_slug(page.with_suffix("").as_posix()), *(_slug(item) for item in heading_path)]
return "--".join(parts)
def _reference_pages(references_dir: Path) -> list[Path]:
return [
page
for page in sorted(references_dir.glob("*.md"), key=lambda item: item.name)
if page.name not in EXCLUDED_REFERENCE_FILES
]
def _heading_records(lines: list[str]) -> tuple[str, list[tuple[int, int, str]]]:
page_title = ""
records: list[tuple[int, int, str]] = []
in_fence = False
for index, line in enumerate(lines):
if FENCE_RE.match(line.rstrip("\r\n")):
in_fence = not in_fence
continue
if in_fence:
continue
match = HEADING_RE.match(line.rstrip("\r\n"))
if not match:
continue
level = len(match.group(1))
title = match.group(2)
if level == 1 and not page_title:
page_title = title
elif level in (2, 3):
records.append((index, level, title))
return page_title, records
def _identities(body: str) -> tuple[str, ...]:
identities: list[str] = []
lines = body.splitlines()
in_fence = False
for index, line in enumerate(lines):
if not FENCE_RE.match(line):
continue
if not in_fence:
previous = index - 1
while previous >= 0 and not lines[previous].strip():
previous -= 1
if previous >= 0 and lines[previous].strip().startswith(IDENTITY_PREFIX):
identities.append(lines[previous].strip()[len(IDENTITY_PREFIX) :].strip())
in_fence = not in_fence
return tuple(identities)
def load_sections(references_dir: Path = DEFAULT_REFERENCES_DIR) -> list[Section]:
sections: list[Section] = []
id_counts: dict[str, int] = {}
for page in _reference_pages(Path(references_dir)):
text = page.read_text(encoding="utf-8")
lines = text.splitlines(keepends=True)
page_title, headings = _heading_records(lines)
h2_title = ""
for position, (start, level, title) in enumerate(headings):
if level == 2:
h2_title = title
heading_path = (title,)
else:
heading_path = (h2_title, title) if h2_title else (title,)
end = len(lines)
for next_start, next_level, _ in headings[position + 1 :]:
if next_level <= level:
end = next_start
break
body = "".join(lines[start:end])
base_id = section_id(page.relative_to(references_dir), heading_path)
id_counts[base_id] = id_counts.get(base_id, 0) + 1
unique_id = base_id
if id_counts[base_id] > 1:
unique_id = f"{base_id}-{id_counts[base_id]}"
searchable_text = normalize(
"\n".join((page.stem, page_title, *heading_path, body))
)
sections.append(
Section(
id=unique_id,
page=page,
page_title=page_title,
heading_path=heading_path,
body=body,
identities=_identities(body),
searchable_text=searchable_text,
)
)
return sections
def _identity_problems(page: Path, lines: list[str]) -> list[ValidationProblem]:
problems: list[ValidationProblem] = []
for index, line in enumerate(lines, start=1):
stripped = line.strip()
if stripped.startswith(IDENTITY_PREFIX):
identity = stripped[len(IDENTITY_PREFIX) :].strip()
if identity not in ALLOWED_IDENTITIES:
problems.append(ValidationProblem(page, index, f"未知身份:{identity}"))
in_fence = False
for index, line in enumerate(lines):
if not FENCE_RE.match(line):
continue
if in_fence:
in_fence = False
continue
previous = index - 1
identities: list[tuple[int, str]] = []
while previous >= 0:
stripped = lines[previous].strip()
if FENCE_RE.match(stripped) or HEADING_RE.match(stripped):
break
if stripped.startswith(IDENTITY_PREFIX):
identity = stripped[len(IDENTITY_PREFIX) :].strip()
identities.append((previous, identity))
previous -= 1
recognized = [identity for _, identity in identities if identity in ALLOWED_IDENTITIES]
if len(recognized) != 1 or len(identities) != 1:
problems.append(
ValidationProblem(page, index + 1, "每个代码围栏必须关联恰好一个代码块身份")
)
in_fence = True
if in_fence:
problems.append(ValidationProblem(page, len(lines), "代码围栏未闭合"))
return problems
def _local_link_problems(
page: Path, text: str, references_dir: Path
) -> list[ValidationProblem]:
problems: list[ValidationProblem] = []
without_fences = re.sub(r"```.*?```", "", text, flags=re.DOTALL)
searchable_markdown = re.sub(r"`[^`\n]*`", "", without_fences)
for match in MARKDOWN_LINK_RE.finditer(searchable_markdown):
target = match.group(1).strip().split(maxsplit=1)[0].strip("<>")
if target.startswith(("#", "http://", "https://", "mailto:")):
continue
target_path = unquote(target.split("#", 1)[0].replace("\\", "/"))
resolved = (page.parent / target_path).resolve()
try:
resolved.relative_to(references_dir.resolve())
except ValueError:
exists = False
else:
exists = resolved.is_file()
if not exists:
line = searchable_markdown.count("\n", 0, match.start()) + 1
problems.append(ValidationProblem(page, line, f"本地链接不存在:{target}"))
return problems
def validate_references(
references_dir: Path = DEFAULT_REFERENCES_DIR,
) -> list[ValidationProblem]:
references_dir = Path(references_dir)
problems: list[ValidationProblem] = []
index_page = references_dir / "index.md"
if index_page.exists():
problems.append(ValidationProblem(index_page, 1, "references 中不得保留 index.md"))
for page in sorted(references_dir.glob("*.md"), key=lambda item: item.name):
text = page.read_text(encoding="utf-8")
lines = text.splitlines()
problems.extend(_identity_problems(page, lines))
problems.extend(_local_link_problems(page, text, references_dir))
for phrase in ROUTER_PHRASES:
for index, line in enumerate(lines, start=1):
if phrase in line:
problems.append(
ValidationProblem(page, index, f"包含人工路由协议:{phrase}")
)
sections = load_sections(references_dir)
ids: dict[str, Section] = {}
for section in sections:
if section.id in ids:
problems.append(ValidationProblem(section.page, 1, f"重复 section ID{section.id}"))
ids[section.id] = section
return problems
def query_tokens(text: str) -> set[str]:
normalized = normalize(text)
tokens = set(ASCII_TOKEN_RE.findall(normalized))
for run in CHINESE_RUN_RE.findall(normalized):
tokens.add(run)
tokens.update(run[index : index + 2] for index in range(len(run) - 1))
return {token for token in tokens if token.strip()}
def _score_section(section: Section, query: str, mode: str) -> int:
normalized_query = normalize(query).strip()
tokens = query_tokens(query)
heading_text = normalize("\n".join((section.page_title, *section.heading_path)))
score = 0
for token in tokens:
if token in heading_text:
score += TOKEN_IN_HEADING_SCORE
elif token in section.searchable_text:
score += TOKEN_IN_BODY_SCORE
if normalized_query and normalized_query in section.searchable_text:
score += EXACT_QUERY_SCORE
if any(token in normalize(section.page.stem) for token in tokens):
score += PAGE_NAME_SCORE
if mode == "write" and "可直接照写示例" in section.identities:
score += DIRECT_EXAMPLE_BOOST
if mode == "diagnose":
if section.page.name == "11_pitfalls.md":
score += PITFALL_PAGE_BOOST
if "反例 / 不可照写" in section.identities:
score += COUNTEREXAMPLE_BOOST
if normalized_query and normalized_query in section.searchable_text:
score += EXACT_ERROR_BOOST
return score
def _write_prelude(sections: list[Section]) -> list[Section]:
preferred = [
("02_core_model.md", "文件模型核心规则"),
("01_quickstart.md", "语言核心事实速查"),
]
prelude: list[Section] = []
for page_name, heading in preferred:
match = next(
(
section
for section in sections
if section.page.name == page_name and heading in section.heading_path
),
None,
)
if match is not None:
prelude.append(match)
return prelude
def query_sections(
query: str,
mode: str,
limit: int = 5,
references_dir: Path = DEFAULT_REFERENCES_DIR,
) -> QueryResult:
if mode not in {"write", "diagnose", "explain"}:
raise ValueError(f"unsupported mode: {mode}")
if not 1 <= limit <= 10:
raise ValueError("limit must be between 1 and 10")
sections = load_sections(references_dir)
ascii_tokens = {
token for token in query_tokens(query) if ASCII_TOKEN_RE.fullmatch(token)
}
if ascii_tokens:
sections = [
section
for section in sections
if any(token in section.searchable_text for token in ascii_tokens)
]
ranked = [
QueryMatch(section, _score_section(section, query, mode)) for section in sections
]
has_chinese = bool(CHINESE_RUN_RE.search(normalize(query)))
minimum_score = MIXED_QUERY_MIN_SCORE if ascii_tokens and has_chinese else 1
ranked = [match for match in ranked if match.score >= minimum_score]
ranked.sort(key=lambda match: (-match.score, match.section.page.as_posix(), match.section.id))
matches = ranked[:limit]
prelude = _write_prelude(sections) if mode == "write" else []
match_ids = {match.section.id for match in matches}
prelude = [section for section in prelude if section.id not in match_ids]
return QueryResult(query=query, mode=mode, matches=matches, prelude=prelude)
def render_result(result: QueryResult) -> str:
lines = [
"# TSL Syntax Lookup",
"",
f"Mode: `{result.mode}`",
f"Query: {result.query}",
]
if result.prelude:
lines.extend(["", "## Required context"])
for section in result.prelude:
lines.extend(["", f"Section ID: `{section.id}`", "", section.body.rstrip()])
for index, match in enumerate(result.matches, start=1):
lines.extend(
[
"",
f"## Match {index}",
"",
f"Score: {match.score}",
f"Section ID: `{match.section.id}`",
f"Source: `{match.section.page.as_posix()}`",
"",
match.section.body.rstrip(),
]
)
lines.extend(["", "## Narrower section IDs"])
lines.extend(f"- `{match.section.id}`" for match in result.matches)
return "\n".join(lines).rstrip() + "\n"
def _configure_utf8() -> None:
for stream in (sys.stdout, sys.stderr):
reconfigure = getattr(stream, "reconfigure", None)
if reconfigure is not None:
reconfigure(encoding="utf-8")
def _parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description="Search TSL syntax reference sections")
action = parser.add_mutually_exclusive_group(required=True)
action.add_argument("--query")
action.add_argument("--section")
action.add_argument("--check", action="store_true")
parser.add_argument("--mode", choices=("write", "diagnose", "explain"))
parser.add_argument("--limit", type=int, default=5)
parser.add_argument("--references-dir", type=Path, default=DEFAULT_REFERENCES_DIR)
return parser
def main(argv: list[str] | None = None) -> int:
_configure_utf8()
parser = _parser()
args = parser.parse_args(argv)
if not 1 <= args.limit <= 10:
parser.error("--limit must be between 1 and 10")
if args.query is not None and args.mode is None:
parser.error("--mode is required with --query")
if args.check:
problems = validate_references(args.references_dir)
for problem in problems:
print(f"{problem.page}:{problem.line}: {problem.message}", file=sys.stderr)
return 1 if problems else 0
if args.section is not None:
section = next(
(item for item in load_sections(args.references_dir) if item.id == args.section),
None,
)
if section is None:
print(f"section not found: {args.section}", file=sys.stderr)
return 2
result = QueryResult("", "section", [QueryMatch(section, 0)], [])
else:
result = query_sections(args.query, args.mode, args.limit, args.references_dir)
if not result.matches:
print("no matching sections", file=sys.stderr)
return 2
print(render_result(result), end="")
return 0
if __name__ == "__main__":
raise SystemExit(main())