Files
playbook/skills/tsl-api-reference/scripts/value_domains.py
T

738 lines
25 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
"""Load and query bundled TSL API parameter value domains."""
from __future__ import annotations
import json
import re
import unicodedata
from dataclasses import dataclass
from pathlib import Path
from typing import Iterable, Sequence
DOMAIN_MODES = {
"catalog",
"enum",
"pattern",
"runtime_catalog",
"versioned_catalog",
}
DOMAIN_ID_RE = re.compile(r"^[a-z][a-z0-9_]*$")
def _fold(value: str) -> str:
normalized = unicodedata.normalize("NFKC", value).casefold()
return "".join(
character
for character in normalized
if character in {"_", "."}
or unicodedata.category(character)[0] in {"L", "N"}
)
def _fold_value(value: str) -> str:
"""Normalize a domain value without conflating negative and positive numbers."""
normalized = unicodedata.normalize("NFKC", value).casefold()
return "".join(
character
for character in normalized
if character in {"-", "_", "."}
or unicodedata.category(character)[0] in {"L", "N"}
)
def _required_string(value: object, where: str) -> str:
if not isinstance(value, str) or not value.strip():
raise ValueError(f"{where} must be a non-empty string")
return value.strip()
def _optional_string(value: object, where: str) -> str:
if value is None:
return ""
if not isinstance(value, str):
raise ValueError(f"{where} must be a string")
return value.strip()
def _string_list(value: object, where: str, *, required: bool = False) -> tuple[str, ...]:
if not isinstance(value, list):
raise ValueError(f"{where} must be a list")
items = tuple(_required_string(item, f"{where}[]") for item in value)
if required and not items:
raise ValueError(f"{where} must not be empty")
if len({_fold(item) for item in items}) != len(items):
raise ValueError(f"{where} contains duplicate values")
return items
def _mapping(value: object, where: str) -> dict:
if not isinstance(value, dict):
raise ValueError(f"{where} must be an object")
return value
def _list(value: object, where: str) -> list:
if not isinstance(value, list):
raise ValueError(f"{where} must be a list")
return value
def _reject_unknown(value: dict, allowed: set[str], where: str) -> None:
unknown = sorted(set(value) - allowed)
if unknown:
raise ValueError(f"{where} contains unknown fields: {', '.join(unknown)}")
@dataclass(frozen=True)
class Binding:
scope: str
module: str
api: str
parameter: str
role: str
@dataclass(frozen=True)
class Resolver:
scope: str
module: str
api: str
parameter: str
catalog: str
source: str
@dataclass(frozen=True)
class RelatedTable:
table_id: str
name: str
scope: str
page: str
@dataclass(frozen=True)
class Verification:
status: str
scope: str
module: str
resolver: str
parameter: str
argument: str
@dataclass(frozen=True)
class Relation:
kind: str
domain_id: str
value: str
verification: Verification
sources: tuple[str, ...]
@dataclass(frozen=True)
class DomainValue:
value: str
label: str
aliases: tuple[str, ...]
parent: str
parent_label: str
level: int | None
valid_from: str
valid_to: str
note: str
related_table: str
as_of: str
sources: tuple[str, ...]
relations: tuple[Relation, ...]
@dataclass(frozen=True)
class ValueDomain:
domain_id: str
label: str
mode: str
complete: bool
bindings: tuple[Binding, ...]
resolvers: tuple[Resolver, ...]
related_tables: tuple[RelatedTable, ...]
values: tuple[DomainValue, ...]
sources: tuple[str, ...]
as_of: str
def related_table(self, table_id: str) -> RelatedTable | None:
return next(
(table for table in self.related_tables if table.table_id == table_id),
None,
)
@dataclass(frozen=True)
class ValueMatch:
domain: ValueDomain
value: DomainValue
score: int
exact: bool
@dataclass(frozen=True)
class ApiHint:
scope: str
module: str
api: str
parameter: str
role: str
domain: ValueDomain
value: str
label: str
verification: Verification | None
sources: tuple[str, ...]
@property
def target(self) -> str:
return f"{self.api}.{self.parameter}"
def summary(self) -> str:
if self.role == "historical_components":
text = f"历史分类代码:{self.value}"
if self.label and self.label != self.value:
text += f"{self.label}"
return text
if self.role == "historical_market_components":
text = f"历史市场板块代码:{self.value}"
if self.label and self.label != self.value:
text += f"{self.label}"
return text
if self.role == "historical_named_components":
return f"历史市场板块名:{self.value}"
if self.role == "current_components":
prefix = (
"当前板块候选"
if self.verification
and self.verification.status == "runtime_required"
else "当前板块名"
)
return f"{prefix}{self.value}"
if self.role == "catalog_selector":
return f"目录选择器:{self.value}"
if self.role == "child_codes":
return f"下级分类代码模式:{self.value}"
text = f"参数值:{self.value}"
if self.label and self.label != self.value:
text += f"{self.label}"
return text
def verification_text(self) -> str:
if not self.verification:
return ""
call = self.verification.resolver
if self.verification.argument:
call += f"('{self.verification.argument}')"
if self.verification.status == "runtime_required":
return f"运行时核验:{call}"
if self.verification.status == "documented":
return "文档已确认"
return self.verification.status
class ValueDomainCatalog:
def __init__(self, domains: Sequence[ValueDomain]):
self.domains = tuple(domains)
self._by_id = {domain.domain_id: domain for domain in self.domains}
@classmethod
def load(cls, path: Path) -> "ValueDomainCatalog":
return cls(_parse_document(json.loads(path.read_text(encoding="utf-8"))))
def search(self, terms: str | Iterable[str]) -> tuple[ValueMatch, ...]:
raw_terms = [terms] if isinstance(terms, str) else list(terms)
normalized_terms = tuple(
_fold_value(term) for term in raw_terms if _fold_value(term)
)
if not normalized_terms:
return ()
matches = []
for domain in self.domains:
for value in domain.values:
scored = _score_value(domain, value, normalized_terms)
if scored:
score, exact = scored
matches.append(ValueMatch(domain, value, score, exact))
matches.sort(
key=lambda match: (
-match.score,
_fold(match.domain.domain_id),
_fold(match.value.value),
)
)
return tuple(matches)
def exact_domains(self, query: str) -> tuple[ValueDomain, ...]:
query_key = _fold_value(query)
if not query_key:
return ()
return tuple(
domain
for domain in self.domains
if query_key
in {
_fold_value(domain.domain_id),
_fold_value(domain.label),
}
)
def recorded_candidates(
self, domain: ValueDomain
) -> tuple[tuple[str, str], ...]:
candidates = [(value.value, value.label) for value in domain.values]
candidates.extend(
(relation.value, value.label)
for source_domain in self.domains
for value in source_domain.values
for relation in value.relations
if relation.domain_id == domain.domain_id
)
unique = []
seen = set()
for value, label in candidates:
key = _fold_value(value)
if key in seen:
continue
seen.add(key)
unique.append((value, label))
return tuple(unique)
def api_hints(self, matches: Iterable[ValueMatch]) -> tuple[ApiHint, ...]:
hints = []
seen = set()
for match in matches:
for binding in match.domain.bindings:
hint = ApiHint(
scope=binding.scope,
module=binding.module,
api=binding.api,
parameter=binding.parameter,
role=binding.role,
domain=match.domain,
value=match.value.value,
label=match.value.label,
verification=None,
sources=match.value.sources,
)
key = _hint_key(hint)
if key not in seen:
seen.add(key)
hints.append(hint)
for relation in match.value.relations:
target = self._by_id[relation.domain_id]
for binding in target.bindings:
hint = ApiHint(
scope=binding.scope,
module=binding.module,
api=binding.api,
parameter=binding.parameter,
role=binding.role,
domain=target,
value=relation.value,
label=match.value.label,
verification=relation.verification,
sources=relation.sources,
)
key = _hint_key(hint)
if key not in seen:
seen.add(key)
hints.append(hint)
hints.sort(
key=lambda hint: (
_fold(hint.api),
_fold(hint.parameter),
hint.summary(),
)
)
return tuple(hints)
def describe_api(
self,
api: str,
*,
scope: str = "",
module: str = "",
parameters: Iterable[str] | None = None,
) -> str:
api_key = _fold(api)
scope_key = _fold(scope)
module_key = _fold(module)
parameter_keys = (
{_fold(parameter) for parameter in parameters}
if parameters is not None
else None
)
bound = [
(domain, binding)
for domain in self.domains
for binding in domain.bindings
if _fold(binding.api) == api_key
and (not scope_key or _fold(binding.scope) == scope_key)
and (not module_key or _fold(binding.module) == module_key)
and (
parameter_keys is None
or _fold(binding.parameter) in parameter_keys
)
]
if not bound:
return ""
lines = ["### 参数取值域", ""]
for domain, binding in bound:
completeness = "完整" if domain.complete else "静态记录不完整"
lines.append(
f"- `{binding.parameter}` -> `{domain.domain_id}`{domain.label}"
f"`{domain.mode}`{completeness}"
)
lines.append(f" - 快照日期:`{domain.as_of}`")
sources = ", ".join(f"`{source}`" for source in domain.sources)
lines.append(f" - 来源:{sources}")
for resolver in domain.resolvers:
catalog = {"system": "系统", "user": "用户"}.get(
resolver.catalog, resolver.catalog
)
lines.append(
f" - 运行时解析器:`{resolver.api}`{catalog}目录;"
f"来源 `{resolver.source}`"
)
recorded_count = len(self.recorded_candidates(domain))
if domain.complete and domain.mode in {"enum", "pattern"}:
for value in domain.values:
label = f" — {value.label}" if value.label != value.value else ""
note = f"{value.note}" if value.note else ""
lines.append(f" - `{value.value}`{label}{note}")
else:
lines.append(f" - 已记录值或候选:{recorded_count} 条(此处不展开)")
lines.append(
" - 按值查询:`dictionary_lookup.py --query "
"\"<域 ID、名称、代码或模式>\"`"
)
return "\n".join(lines) + "\n"
def domains_path_for_index(
index_path: Path, *, default_index: Path, default_domains: Path
) -> Path | None:
if index_path.parent.name == "data":
sibling = index_path.parent / default_domains.name
if sibling.is_file():
return sibling
try:
is_default = index_path.resolve() == default_index.resolve()
except OSError:
is_default = index_path == default_index
return default_domains if is_default else None
def _parse_document(document: object) -> tuple[ValueDomain, ...]:
root = _mapping(document, "value domain document")
_reject_unknown(root, {"version", "domains"}, "value domain document")
if root.get("version") != 1:
raise ValueError("value domain document version must be 1")
domain_items = _list(root.get("domains"), "domains")
if not domain_items:
raise ValueError("domains must not be empty")
domains = tuple(
_parse_domain(item, f"domains[{index}]")
for index, item in enumerate(domain_items)
)
ids = [domain.domain_id for domain in domains]
if len(set(ids)) != len(ids):
raise ValueError("domains contains duplicate ids")
known = set(ids)
by_id = {domain.domain_id: domain for domain in domains}
for domain in domains:
table_ids = {table.table_id for table in domain.related_tables}
for value in domain.values:
if value.related_table and value.related_table not in table_ids:
raise ValueError(
f"domain {domain.domain_id!r} value {value.value!r} references "
f"unknown related table {value.related_table!r}"
)
for relation in value.relations:
if relation.domain_id not in known:
raise ValueError(
f"domain {domain.domain_id!r} value {value.value!r} references "
f"unknown domain {relation.domain_id!r}"
)
target = by_id[relation.domain_id]
verification_key = (
_fold(relation.verification.scope),
_fold(relation.verification.module),
_fold(relation.verification.resolver),
_fold(relation.verification.parameter),
)
resolver_keys = {
(
_fold(resolver.scope),
_fold(resolver.module),
_fold(resolver.api),
_fold(resolver.parameter),
)
for resolver in target.resolvers
}
if verification_key not in resolver_keys:
raise ValueError(
f"domain {domain.domain_id!r} value {value.value!r} "
f"references a resolver not owned by {relation.domain_id!r}"
)
return domains
def _parse_domain(value: object, where: str) -> ValueDomain:
item = _mapping(value, where)
_reject_unknown(
item,
{
"id",
"label",
"mode",
"complete",
"bindings",
"resolvers",
"related_tables",
"values",
"sources",
"as_of",
},
where,
)
domain_id = _required_string(item.get("id"), f"{where}.id")
if not DOMAIN_ID_RE.fullmatch(domain_id):
raise ValueError(f"{where}.id must use lowercase snake_case")
mode = _required_string(item.get("mode"), f"{where}.mode")
if mode not in DOMAIN_MODES:
raise ValueError(f"{where}.mode must be one of {', '.join(sorted(DOMAIN_MODES))}")
complete = item.get("complete")
if not isinstance(complete, bool):
raise ValueError(f"{where}.complete must be a boolean")
if mode == "enum" and not complete:
raise ValueError(f"{where}.complete must be true for mode enum")
if mode == "runtime_catalog" and complete:
raise ValueError(f"{where}.complete must be false for mode runtime_catalog")
bindings = tuple(
_parse_binding(binding, f"{where}.bindings[{index}]")
for index, binding in enumerate(_list(item.get("bindings"), f"{where}.bindings"))
)
if not bindings:
raise ValueError(f"{where}.bindings must not be empty")
resolvers = tuple(
_parse_resolver(resolver, f"{where}.resolvers[{index}]")
for index, resolver in enumerate(_list(item.get("resolvers"), f"{where}.resolvers"))
)
if mode == "runtime_catalog" and not resolvers:
raise ValueError(f"{where}.resolvers must not be empty for mode runtime_catalog")
tables = tuple(
_parse_table(table, f"{where}.related_tables[{index}]")
for index, table in enumerate(
_list(item.get("related_tables"), f"{where}.related_tables")
)
)
values = tuple(
_parse_value(domain_value, f"{where}.values[{index}]")
for index, domain_value in enumerate(_list(item.get("values"), f"{where}.values"))
)
value_keys = [_fold_value(domain_value.value) for domain_value in values]
if len(set(value_keys)) != len(value_keys):
raise ValueError(f"{where}.values contains duplicate values")
if mode != "runtime_catalog" and not values:
raise ValueError(f"{where}.values must not be empty for mode {mode}")
return ValueDomain(
domain_id=domain_id,
label=_required_string(item.get("label"), f"{where}.label"),
mode=mode,
complete=complete,
bindings=bindings,
resolvers=resolvers,
related_tables=tables,
values=values,
sources=_string_list(item.get("sources"), f"{where}.sources", required=True),
as_of=_required_string(item.get("as_of"), f"{where}.as_of"),
)
def _parse_binding(value: object, where: str) -> Binding:
item = _mapping(value, where)
_reject_unknown(item, {"scope", "module", "api", "parameter", "role"}, where)
return Binding(
scope=_required_string(item.get("scope"), f"{where}.scope"),
module=_required_string(item.get("module"), f"{where}.module"),
api=_required_string(item.get("api"), f"{where}.api"),
parameter=_required_string(item.get("parameter"), f"{where}.parameter"),
role=_required_string(item.get("role"), f"{where}.role"),
)
def _parse_resolver(value: object, where: str) -> Resolver:
item = _mapping(value, where)
_reject_unknown(
item,
{"scope", "module", "api", "parameter", "catalog", "source"},
where,
)
return Resolver(
scope=_required_string(item.get("scope"), f"{where}.scope"),
module=_required_string(item.get("module"), f"{where}.module"),
api=_required_string(item.get("api"), f"{where}.api"),
parameter=_required_string(item.get("parameter"), f"{where}.parameter"),
catalog=_required_string(item.get("catalog"), f"{where}.catalog"),
source=_required_string(item.get("source"), f"{where}.source"),
)
def _parse_table(value: object, where: str) -> RelatedTable:
item = _mapping(value, where)
_reject_unknown(item, {"id", "name", "scope", "page"}, where)
return RelatedTable(
table_id=_required_string(item.get("id"), f"{where}.id"),
name=_required_string(item.get("name"), f"{where}.name"),
scope=_required_string(item.get("scope"), f"{where}.scope"),
page=_required_string(item.get("page"), f"{where}.page"),
)
def _parse_value(value: object, where: str) -> DomainValue:
item = _mapping(value, where)
_reject_unknown(
item,
{
"value",
"label",
"aliases",
"parent",
"parent_label",
"level",
"valid_from",
"valid_to",
"note",
"related_table",
"as_of",
"sources",
"relations",
},
where,
)
level = item.get("level")
if level is not None and (not isinstance(level, int) or level < 0):
raise ValueError(f"{where}.level must be a non-negative integer")
return DomainValue(
value=_required_string(item.get("value"), f"{where}.value"),
label=_required_string(item.get("label"), f"{where}.label"),
aliases=_string_list(item.get("aliases", []), f"{where}.aliases"),
parent=_optional_string(item.get("parent"), f"{where}.parent"),
parent_label=_optional_string(item.get("parent_label"), f"{where}.parent_label"),
level=level,
valid_from=_optional_string(item.get("valid_from"), f"{where}.valid_from"),
valid_to=_optional_string(item.get("valid_to"), f"{where}.valid_to"),
note=_optional_string(item.get("note"), f"{where}.note"),
related_table=_optional_string(
item.get("related_table"), f"{where}.related_table"
),
as_of=_optional_string(item.get("as_of"), f"{where}.as_of"),
sources=_string_list(item.get("sources"), f"{where}.sources", required=True),
relations=tuple(
_parse_relation(relation, f"{where}.relations[{index}]")
for index, relation in enumerate(
_list(item.get("relations", []), f"{where}.relations")
)
),
)
def _parse_relation(value: object, where: str) -> Relation:
item = _mapping(value, where)
_reject_unknown(item, {"kind", "domain", "value", "verification", "sources"}, where)
verification_item = _mapping(item.get("verification"), f"{where}.verification")
_reject_unknown(
verification_item,
{"status", "scope", "module", "resolver", "parameter", "argument"},
f"{where}.verification",
)
status = _required_string(
verification_item.get("status"), f"{where}.verification.status"
)
if status not in {"documented", "runtime_required"}:
raise ValueError(
f"{where}.verification.status must be documented or runtime_required"
)
return Relation(
kind=_required_string(item.get("kind"), f"{where}.kind"),
domain_id=_required_string(item.get("domain"), f"{where}.domain"),
value=_required_string(item.get("value"), f"{where}.value"),
verification=Verification(
status=status,
scope=_required_string(
verification_item.get("scope"), f"{where}.verification.scope"
),
module=_required_string(
verification_item.get("module"), f"{where}.verification.module"
),
resolver=_required_string(
verification_item.get("resolver"), f"{where}.verification.resolver"
),
parameter=_required_string(
verification_item.get("parameter"),
f"{where}.verification.parameter",
),
argument=_optional_string(
verification_item.get("argument"), f"{where}.verification.argument"
),
),
sources=_string_list(item.get("sources"), f"{where}.sources", required=True),
)
def _score_value(
domain: ValueDomain, value: DomainValue, terms: Sequence[str]
) -> tuple[int, bool] | None:
fields = (
(value.value, 600),
(value.label, 560),
*((alias, 520) for alias in value.aliases),
(value.parent, 260),
(value.parent_label, 300),
(domain.label, 180),
(domain.domain_id, 160),
)
folded_fields = tuple(
(_fold_value(text), weight) for text, weight in fields if _fold_value(text)
)
scores = []
exact = False
for term in terms:
best = 0
for field, weight in folded_fields:
if term == field:
best = max(best, weight + 500)
exact = exact or len(terms) == 1
elif len(term) < 2:
continue
elif term in field:
best = max(best, weight)
elif len(field) >= 2 and field in term:
best = max(best, weight - 80)
if not best:
return None
scores.append(best)
same_field_bonus = 300 if any(all(term in field for term in terms) for field, _ in folded_fields) else 0
return sum(scores) + same_field_bonus, exact
def _hint_key(hint: ApiHint) -> tuple[str, ...]:
return (
_fold(hint.scope),
_fold(hint.module),
_fold(hint.api),
_fold(hint.parameter),
_fold(hint.value),
hint.role,
hint.verification.status if hint.verification else "",
)