216 lines
7.7 KiB
Python
216 lines
7.7 KiB
Python
"""Reduce task-specific format measurements into a conservative Skill policy."""
|
||
|
||
from __future__ import annotations
|
||
|
||
from dataclasses import asdict, dataclass
|
||
import re
|
||
from typing import Any
|
||
|
||
from .document import CANONICAL_HEADINGS, HEADING_RE, parse_document, section_key
|
||
from .models import Operation
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class FormatStyle:
|
||
id: str
|
||
source_format: str | None
|
||
strict_accuracy: float | None
|
||
prior_rank: int
|
||
|
||
def to_dict(self) -> dict[str, Any]:
|
||
return asdict(self)
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class FormatPolicy:
|
||
enabled: bool
|
||
classification: str
|
||
strict_accuracy_spread: float | None
|
||
styles: tuple[FormatStyle, ...]
|
||
avoid_patterns: tuple[str, ...]
|
||
cautions: tuple[str, ...]
|
||
|
||
def to_dict(self) -> dict[str, Any]:
|
||
return {
|
||
"enabled": self.enabled,
|
||
"classification": self.classification,
|
||
"strict_accuracy_spread": self.strict_accuracy_spread,
|
||
"styles": [style.to_dict() for style in self.styles],
|
||
"avoid_patterns": list(self.avoid_patterns),
|
||
"cautions": list(self.cautions),
|
||
}
|
||
|
||
|
||
def _number(value: Any) -> float | None:
|
||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||
return None
|
||
result = float(value)
|
||
return result if 0.0 <= result <= 1.0 else None
|
||
|
||
|
||
def _infer_style(prompt_format: str) -> str | None:
|
||
"""Map a short-template result to a heading-label hypothesis.
|
||
|
||
The format benchmark does not test complete Skills. The mapping therefore
|
||
deliberately captures only casing and the label delimiter; Markdown heading
|
||
structure is retained by the renderer.
|
||
"""
|
||
|
||
labels = re.findall(r"([A-Za-z]+)\s*([:-])\s*\{\}", prompt_format)
|
||
if not labels:
|
||
labels = re.findall(r"([A-Za-z]+)\s+([-])\s+\{\}", prompt_format)
|
||
if not labels:
|
||
return None
|
||
words = [word for word, _ in labels]
|
||
delimiters = {delimiter for _, delimiter in labels}
|
||
if len(delimiters) != 1:
|
||
return None
|
||
delimiter = next(iter(delimiters))
|
||
casing = "uppercase" if all(word.isupper() for word in words) else "title"
|
||
suffix = "hyphen" if delimiter == "-" else "colon"
|
||
return f"{casing}-{suffix}-labels"
|
||
|
||
|
||
def _avoid_patterns(formats: Any) -> tuple[str, ...]:
|
||
if not isinstance(formats, list):
|
||
return ()
|
||
patterns: list[str] = []
|
||
for item in formats:
|
||
if not isinstance(item, dict) or not isinstance(item.get("prompt_format"), str):
|
||
continue
|
||
value = item["prompt_format"]
|
||
if "<sep>" in value and "synthetic-separator-token" not in patterns:
|
||
patterns.append("synthetic-separator-token")
|
||
if re.search(r"\n[ \t]+[A-Za-z]", value) and "indented-label" not in patterns:
|
||
patterns.append("indented-label")
|
||
labels = re.findall(r"\b([A-Za-z]+)\s*[:-]", value)
|
||
if labels and len({word.isupper() for word in labels}) > 1:
|
||
if "mixed-label-casing" not in patterns:
|
||
patterns.append("mixed-label-casing")
|
||
if " " in value and "inconsistent-spacing" not in patterns:
|
||
patterns.append("inconsistent-spacing")
|
||
return tuple(patterns)
|
||
|
||
|
||
def reduce_format_policy(profile: dict[str, Any]) -> FormatPolicy:
|
||
raw = profile.get("format_preference")
|
||
if not isinstance(raw, dict):
|
||
return FormatPolicy(
|
||
enabled=False,
|
||
classification="unavailable",
|
||
strict_accuracy_spread=None,
|
||
styles=(),
|
||
avoid_patterns=(),
|
||
cautions=("No format_preference object is present in the profile.",),
|
||
)
|
||
|
||
classification = str(raw.get("classification", "unknown"))
|
||
spread = _number(raw.get("strict_accuracy_spread"))
|
||
enabled = classification == "format_sensitive" and spread is not None and spread >= 0.10
|
||
styles: list[FormatStyle] = []
|
||
best = raw.get("best_formats")
|
||
if isinstance(best, list):
|
||
for index, item in enumerate(best):
|
||
if not isinstance(item, dict) or not isinstance(item.get("prompt_format"), str):
|
||
continue
|
||
style_id = _infer_style(item["prompt_format"])
|
||
if style_id is None:
|
||
continue
|
||
styles.append(
|
||
FormatStyle(
|
||
id=style_id,
|
||
source_format=item["prompt_format"],
|
||
strict_accuracy=_number(item.get("strict_accuracy")),
|
||
prior_rank=index + 1,
|
||
)
|
||
)
|
||
|
||
return FormatPolicy(
|
||
enabled=enabled and bool(styles),
|
||
classification=classification,
|
||
strict_accuracy_spread=spread,
|
||
styles=tuple(styles) if enabled else (),
|
||
avoid_patterns=_avoid_patterns(raw.get("worst_formats")),
|
||
cautions=(
|
||
"Format scores are priors from short templates, not proof of whole-Skill quality.",
|
||
"The compiler applies all ranked safe surface styles in order; the final style wins.",
|
||
),
|
||
)
|
||
|
||
|
||
def _style_heading(title: str, language: str, style_id: str) -> str:
|
||
key = section_key(title)
|
||
canonical = (
|
||
CANONICAL_HEADINGS[language][key]
|
||
if key is not None
|
||
else title.strip()
|
||
)
|
||
delimiter = "-" if style_id.endswith("-hyphen-labels") else ":"
|
||
if canonical.endswith(delimiter):
|
||
return (
|
||
canonical.upper()
|
||
if language == "en" and style_id.startswith("uppercase-")
|
||
else canonical
|
||
)
|
||
canonical = canonical.rstrip("::-–—")
|
||
labelled = re.match(r"^([^::]+)[::]\s*(.+)$", canonical)
|
||
if labelled is None and delimiter == "-":
|
||
existing_hyphen = re.match(r"^([^-]+)-\s+(.+)$", canonical)
|
||
if existing_hyphen and existing_hyphen.group(1).isupper():
|
||
labelled = existing_hyphen
|
||
if labelled:
|
||
label, payload = labelled.groups()
|
||
if language == "en" and style_id.startswith("uppercase-"):
|
||
label, payload = label.upper(), payload.upper()
|
||
return f"{label}{delimiter} {payload}"
|
||
if language == "en" and style_id.startswith("uppercase-"):
|
||
canonical = canonical.upper()
|
||
return canonical + delimiter
|
||
|
||
|
||
def apply_format_style(
|
||
content: str, style: FormatStyle
|
||
) -> tuple[str, list[Operation]]:
|
||
"""Apply one profile-selected style to safe H2 section-label surfaces."""
|
||
|
||
document = parse_document(content)
|
||
patches: list[tuple[int, int, str]] = []
|
||
operations: list[Operation] = []
|
||
for block in document.blocks:
|
||
# Keep the Skill title and step-level prose unchanged. Inline literals
|
||
# in headings are protected because they may be paths or identifiers.
|
||
if (
|
||
block.kind != "heading"
|
||
or block.heading_level != 2
|
||
or block.protected_spans
|
||
):
|
||
continue
|
||
match = HEADING_RE.match(block.text)
|
||
if match is None:
|
||
continue
|
||
replacement = (
|
||
f"{match.group(1)} "
|
||
f"{_style_heading(match.group(2), document.language, style.id)}"
|
||
)
|
||
if replacement == block.text:
|
||
continue
|
||
patches.append(
|
||
(block.start_offset, block.start_offset + len(block.text), replacement)
|
||
)
|
||
operations.append(
|
||
Operation(
|
||
type="FORMAT_HEADING_LABEL",
|
||
signal="format_preference",
|
||
block_id=block.id,
|
||
quote=block.text,
|
||
replacement=replacement,
|
||
target_section=section_key(match.group(2)),
|
||
)
|
||
)
|
||
if not patches:
|
||
return content, []
|
||
body = document.body
|
||
for start, end, replacement in sorted(patches, reverse=True):
|
||
body = body[:start] + replacement + body[end:]
|
||
return document.frontmatter + body, operations
|