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

567 lines
19 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 difflib
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]+")
INLINE_CODE_RE = re.compile(r"`([^`\n]+)`")
FENCED_CODE_RE = re.compile(r"```[^\n]*\n(.*?)```", re.DOTALL)
IDENTITY_PREFIX = "代码块身份:"
BLOCK_DESCRIPTION_PREFIX = "代码块说明:"
ALLOWED_IDENTITIES = {
"可直接照写示例",
"反例 / 不可照写",
"输出片段",
"配置片段 / 概念骨架",
"仅服务端可执行示例",
}
ROUTER_PHRASES = ("路由中心", "选择一个主专题", "候选页继续判断")
EXCLUDED_REFERENCE_FILES = {"index.md"}
HEADING_TOKEN_SCORE = 12
HEADING_EXACT_SCORE = 20
TERM_TOKEN_SCORE = 10
TERM_EXACT_SCORE = 16
PAGE_TITLE_TOKEN_SCORE = 5
PAGE_TITLE_EXACT_SCORE = 8
BODY_TOKEN_SCORE = 3
BODY_EXACT_SCORE = 4
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
priority: tuple[int, int, int, int, int] = (0, 0, 0, 0, 0)
@dataclass(frozen=True)
class ScoreBreakdown:
heading_path: int
exact_term: int
page_title: int
body: int
mode_boost: int
diagnose_priority: bool = False
@property
def lexical_total(self) -> int:
return self.heading_path + self.exact_term + self.page_title + self.body
@property
def total(self) -> int:
return self.lexical_total + self.mode_boost
@property
def priority(self) -> tuple[int, int, int, int, int]:
if self.diagnose_priority:
return (
self.mode_boost,
self.heading_path,
self.exact_term,
self.page_title,
self.body,
)
return (
self.heading_path,
self.exact_term,
self.page_title,
self.mode_boost,
self.body,
)
@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 _associated_identity(lines: list[str], opening_fence: int) -> str | None:
previous = opening_fence - 1
while previous >= 0 and not lines[previous].strip():
previous -= 1
while previous >= 0 and lines[previous].strip().startswith(BLOCK_DESCRIPTION_PREFIX):
previous -= 1
while previous >= 0 and not lines[previous].strip():
previous -= 1
if previous < 0:
return None
metadata = lines[previous].strip()
if not metadata.startswith(IDENTITY_PREFIX):
return None
identity = metadata[len(IDENTITY_PREFIX) :].strip()
earlier = previous - 1
while earlier >= 0 and not lines[earlier].strip():
earlier -= 1
if earlier >= 0 and lines[earlier].strip().startswith(IDENTITY_PREFIX):
return None
return identity if identity in ALLOWED_IDENTITIES else None
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:
identity = _associated_identity(lines, index)
if identity is not None:
identities.append(identity)
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
if _associated_identity(lines, index) is None:
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 _has_chinese_context(section: Section, query: str) -> bool:
runs = CHINESE_RUN_RE.findall(normalize(query))
for run in runs:
tokens = (
[run]
if len(run) < 2
else [run[index : index + 2] for index in range(len(run) - 1)]
)
matched = sum(token in section.searchable_text for token in tokens)
if matched >= (len(tokens) + 1) // 2:
return True
return not runs
def _code_text(body: str) -> str:
inline = INLINE_CODE_RE.findall(body)
fenced = FENCED_CODE_RE.findall(body)
return normalize("\n".join((*inline, *fenced)))
def _score_section(section: Section, query: str, mode: str) -> ScoreBreakdown:
normalized_query = normalize(query).strip()
tokens = query_tokens(query)
heading_text = normalize("\n".join(section.heading_path))
page_title_text = normalize(section.page_title)
body_text = normalize(section.body)
term_text = _code_text(section.body)
heading_score = 0
term_score = 0
page_title_score = 0
body_score = 0
for token in tokens:
if token in heading_text:
heading_score += HEADING_TOKEN_SCORE
if token in term_text:
term_score += TERM_TOKEN_SCORE
if token in page_title_text:
page_title_score += PAGE_TITLE_TOKEN_SCORE
if token in body_text:
body_score += BODY_TOKEN_SCORE
if normalized_query and normalized_query in heading_text:
heading_score += HEADING_EXACT_SCORE
if normalized_query and normalized_query in term_text:
term_score += TERM_EXACT_SCORE
if normalized_query and normalized_query in page_title_text:
page_title_score += PAGE_TITLE_EXACT_SCORE
if normalized_query and normalized_query in body_text:
body_score += BODY_EXACT_SCORE
mode_boost = 0
if mode == "write" and "可直接照写示例" in section.identities:
mode_boost += DIRECT_EXAMPLE_BOOST
if mode == "diagnose":
if section.page.name == "11_pitfalls.md":
mode_boost += PITFALL_PAGE_BOOST
if "反例 / 不可照写" in section.identities:
mode_boost += COUNTEREXAMPLE_BOOST
if normalized_query and normalized_query in section.searchable_text:
mode_boost += EXACT_ERROR_BOOST
return ScoreBreakdown(
heading_path=heading_score,
exact_term=term_score,
page_title=page_title_score,
body=body_score,
mode_boost=mode_boost,
diagnose_priority=mode == "diagnose",
)
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")
all_sections = load_sections(references_dir)
sections = all_sections
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)
]
has_chinese = bool(CHINESE_RUN_RE.search(normalize(query)))
minimum_score = MIXED_QUERY_MIN_SCORE if ascii_tokens and has_chinese else 1
ranked: list[QueryMatch] = []
for section in sections:
if ascii_tokens and has_chinese and not _has_chinese_context(section, query):
continue
score = _score_section(section, query, mode)
if score.lexical_total < minimum_score:
continue
ranked.append(QueryMatch(section, score.total, score.priority))
ranked.sort(
key=lambda match: (
*(-value for value in match.priority),
match.section.page.as_posix(),
match.section.id,
)
)
matches = ranked[:limit]
prelude = _write_prelude(all_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}`",
f"Source: `{section.page.as_posix()}`",
"",
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 _nearest_section_ids(
requested: str, sections: list[Section], limit: int = 5
) -> list[str]:
ranked = sorted(
sections,
key=lambda section: (
-difflib.SequenceMatcher(None, requested, section.id).ratio(),
section.id,
),
)
return [section.id for section in ranked[:limit]]
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:
sections = load_sections(args.references_dir)
section = next(
(item for item in sections if item.id == args.section),
None,
)
if section is None:
print(f"section not found: {args.section}", file=sys.stderr)
print("Nearest section IDs:", file=sys.stderr)
for candidate in _nearest_section_ids(args.section, sections):
print(f"- {candidate}", 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())