Files
playbook/tools/tsl-codegen/tests/test_value_domains.py
T

543 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.
import csv
import json
import re
import subprocess
import sys
import tempfile
import unittest
from collections import defaultdict
from pathlib import Path
ROOT = Path(__file__).resolve().parents[3]
SKILL_ROOT = ROOT / "skills" / "tsl-api-reference"
LOOKUP = SKILL_ROOT / "scripts" / "lookup.py"
DICTIONARY_LOOKUP = SKILL_ROOT / "scripts" / "dictionary_lookup.py"
class ValueDomainCliTest(unittest.TestCase):
def run_cli(self, script, *args):
return subprocess.run(
[sys.executable, str(script), *args],
cwd=ROOT,
capture_output=True,
text=True,
encoding="utf-8",
check=False,
)
def test_dictionary_query_resolves_classification_by_label_and_code(self):
for query in ("申万煤炭", "SWHY740000"):
with self.subTest(query=query):
result = self.run_cli(DICTIONARY_LOOKUP, "--query", query)
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("status: ok", result.stdout)
self.assertIn("类型:参数取值", result.stdout)
self.assertIn(
"参数域:classification_code(分类属性代码)",
result.stdout,
)
self.assertIn("取值:SWHY740000", result.stdout)
self.assertIn("名称:申万煤炭", result.stdout)
self.assertIn("父级:SWHY(申万行业)", result.stdout)
self.assertIn("层级:1", result.stdout)
self.assertIn("关联表:138 股票.股票行业分类信息", result.stdout)
self.assertIn("快照日期:2023-10-13", result.stdout)
self.assertIn("来源:faq:31494", result.stdout)
def test_dictionary_query_uses_strong_values_without_hijacking_table_names(self):
for query in ("SWHY740000 是什么", "申万煤炭分类代码"):
with self.subTest(query=query):
result = self.run_cli(DICTIONARY_LOOKUP, "--query", query)
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("类型:参数取值", result.stdout)
self.assertIn("取值:SWHY740000", result.stdout)
table = self.run_cli(DICTIONARY_LOOKUP, "--query", "申万行业配置")
self.assertEqual(0, table.returncode, table.stderr)
self.assertIn("类型:table", table.stdout)
self.assertIn("表 ID629", table.stdout)
self.assertNotIn("类型:参数取值", table.stdout)
def test_pattern_punctuation_is_significant(self):
pattern = self.run_cli(DICTIONARY_LOOKUP, "--query", ".N")
single_letter = self.run_cli(DICTIONARY_LOOKUP, "--query", "N")
self.assertEqual(0, pattern.returncode, pattern.stderr)
self.assertIn("类型:参数取值", pattern.stdout)
self.assertIn("取值:<分类属性代码>.N", pattern.stdout)
self.assertEqual(0, single_letter.returncode, single_letter.stderr)
self.assertNotIn("类型:参数取值", single_letter.stdout)
def test_shared_board_type_enum_maps_multiple_api_parameters(self):
label = self.run_cli(DICTIONARY_LOOKUP, "--query", "申万三级行业")
numeric = self.run_cli(DICTIONARY_LOOKUP, "--query", "板块类别 7")
bare_numeric = self.run_cli(DICTIONARY_LOOKUP, "--query", "7")
api = self.run_cli(LOOKUP, "--name", "getbktypelist")
bare_lookup = self.run_cli(LOOKUP, "--kw", "7")
self.assertEqual(0, label.returncode, label.stderr)
self.assertIn("status: ambiguous", label.stdout)
self.assertIn("取值:7", label.stdout)
self.assertIn("名称:申万三级行业", label.stdout)
self.assertIn("绑定 APIgetbktypelist.bktype; stocksbklist.bktype", label.stdout)
self.assertIn("字段:申万三级行业", label.stdout)
self.assertEqual(0, numeric.returncode, numeric.stderr)
self.assertIn("名称:申万三级行业", numeric.stdout)
self.assertEqual(0, bare_numeric.returncode, bare_numeric.stderr)
self.assertNotIn("类型:参数取值", bare_numeric.stdout)
self.assertEqual(0, api.returncode, api.stderr)
self.assertIn("market_board_type", api.stdout)
self.assertIn("完整", api.stdout)
self.assertIn("`7` — 申万三级行业", api.stdout)
self.assertEqual(0, bare_lookup.returncode, bare_lookup.stderr)
self.assertNotIn("参数取值域:getbktypelist.bktype", bare_lookup.stdout)
def test_historical_market_board_codes_use_the_shared_domain(self):
dictionary = self.run_cli(DICTIONARY_LOOKUP, "--query", "TSI000001")
lookup = self.run_cli(LOOKUP, "--kw", "TSI000001")
self.assertEqual(0, dictionary.returncode, dictionary.stderr)
self.assertIn("取值:TSI000001", dictionary.stdout)
self.assertIn("名称:A股板块", dictionary.stdout)
self.assertIn("绑定 APIgetBkByDate.index_id", dictionary.stdout)
self.assertEqual(0, lookup.returncode, lookup.stderr)
self.assertIn("历史市场板块代码:TSI000001A股板块)", lookup.stdout)
self.assertNotIn("参数取值域:getbktypelist.bktype", lookup.stdout)
def test_exact_code_value_ignores_incidental_dictionary_id_substrings(self):
result = self.run_cli(DICTIONARY_LOOKUP, "--query", "SWHY740000")
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("status: ok", result.stdout)
self.assertIn("取值:SWHY740000", result.stdout)
self.assertNotIn("字段:截止日", result.stdout)
def test_named_historical_board_domain_is_complete_and_api_specific(self):
dictionary = self.run_cli(DICTIONARY_LOOKUP, "--query", "中小企业板")
lookup = self.run_cli(LOOKUP, "--kw", "中小企业板")
exact = self.run_cli(LOOKUP, "--name", "getAbkbyDate")
self.assertEqual(0, dictionary.returncode, dictionary.stderr)
self.assertIn("绑定 APIgetAbkbyDate.bk_name", dictionary.stdout)
self.assertIn("2021-04-06 并入主板", dictionary.stdout)
self.assertEqual(0, lookup.returncode, lookup.stderr)
self.assertIn("历史市场板块名:中小企业板", lookup.stdout)
self.assertEqual(0, exact.returncode, exact.stderr)
self.assertIn("historical_a_share_board_name", exact.stdout)
self.assertIn("`中小企业板` — 中小企业板(历史兼容)", exact.stdout)
def test_scope_filters_multi_table_value_domains(self):
stock = self.run_cli(
DICTIONARY_LOOKUP, "--query", "属性代码", "--scope", "stock"
)
fund = self.run_cli(
DICTIONARY_LOOKUP, "--query", "属性代码", "--scope", "fund"
)
self.assertEqual(0, stock.returncode, stock.stderr)
self.assertIn("关联表:138 股票.股票行业分类信息", stock.stdout)
self.assertNotIn("355 基金.基金分类信息", stock.stdout)
self.assertEqual(0, fund.returncode, fund.stderr)
self.assertIn("关联表:355 基金.基金分类信息", fund.stdout)
self.assertNotIn("138 股票.股票行业分类信息", fund.stdout)
def test_keyword_query_maps_value_to_current_and_historical_apis(self):
result = self.run_cli(LOOKUP, "--kw", "申万", "煤炭")
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("getBk\tfunction\tgetBk(marketlist)", result.stdout)
self.assertIn(
"getBkByDate\tfunction\tgetBkByDate(index_id, end_t, extype)",
result.stdout,
)
self.assertIn("getBk.marketlist", result.stdout)
self.assertIn("当前板块候选:申万煤炭", result.stdout)
self.assertIn("运行时核验", result.stdout)
self.assertIn("getBkByDate.index_id", result.stdout)
self.assertIn("历史分类代码:SWHY740000", result.stdout)
def test_keyword_annotations_only_include_the_top_value_group(self):
exact_code = self.run_cli(LOOKUP, "--kw", "SWHY740000")
domain_query = self.run_cli(LOOKUP, "--kw", "属性", "代码")
self.assertEqual(0, exact_code.returncode, exact_code.stderr)
self.assertIn("历史分类代码:SWHY740000", exact_code.stdout)
self.assertNotIn("历史分类代码:SWHY(申万行业)", exact_code.stdout)
self.assertEqual(0, domain_query.returncode, domain_query.stderr)
self.assertIn("目录选择器:属性代码", domain_query.stdout)
self.assertIn("下级分类代码模式:<分类属性代码>.N", domain_query.stdout)
self.assertNotIn("历史分类代码:CAPCHY", domain_query.stdout)
def test_exact_getbk_query_explains_runtime_catalog_boundary(self):
result = self.run_cli(LOOKUP, "--name", "getBk")
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("参数取值域", result.stdout)
self.assertIn("marketlist", result.stdout)
self.assertIn("runtime_catalog", result.stdout)
self.assertIn("静态记录不完整", result.stdout)
self.assertIn("getBkList2", result.stdout)
self.assertIn("getUserBkList2", result.stdout)
self.assertIn("按值查询", result.stdout)
self.assertNotIn("`申万煤炭`", result.stdout)
def test_runtime_domain_query_returns_resolvers_not_unrelated_market_codes(self):
result = self.run_cli(DICTIONARY_LOOKUP, "--query", "market_board")
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("status: ok", result.stdout)
self.assertIn("类型:参数取值域", result.stdout)
self.assertIn("参数域:market_board", result.stdout)
self.assertIn("绑定 APIgetBk.marketlist", result.stdout)
self.assertIn("getBkList2(bktype)(系统目录)", result.stdout)
self.assertIn("getUserBkList2(bktype)(用户目录)", result.stdout)
self.assertIn("申万煤炭", result.stdout)
self.assertNotIn("TSI000001", result.stdout)
def test_domain_label_query_returns_one_domain_summary(self):
result = self.run_cli(DICTIONARY_LOOKUP, "--query", "分类属性代码")
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("status: ok", result.stdout)
self.assertIn("参数域:classification_code", result.stdout)
self.assertIn("绑定 APIgetBkByDate.index_id", result.stdout)
self.assertIn("SWHY740000(申万煤炭)", result.stdout)
self.assertNotIn("类型:field", result.stdout)
def test_exact_api_query_does_not_expand_versioned_catalog(self):
result = self.run_cli(LOOKUP, "--name", "getBkByDate")
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("分类属性代码", result.stdout)
self.assertIn("已记录值或候选:7 条", result.stdout)
self.assertIn("dictionary_lookup.py --query", result.stdout)
parameter_domain = result.stdout.split("### 参数取值域", maxsplit=1)[1]
self.assertNotIn("`SWHY740000`", parameter_domain)
def test_invalid_value_domain_data_is_a_deployment_error(self):
with tempfile.TemporaryDirectory() as temp_dir:
invalid = Path(temp_dir) / "value_domains.json"
invalid.write_text('{"version": 1, "domains": []}\n', encoding="utf-8")
for script, action in (
(LOOKUP, ("--kw", "申万")),
(DICTIONARY_LOOKUP, ("--query", "申万")),
):
with self.subTest(script=script.name):
result = self.run_cli(
script,
*action,
"--value-domains",
str(invalid),
)
self.assertEqual(1, result.returncode)
self.assertEqual("", result.stdout)
self.assertIn("domains must not be empty", result.stderr)
self.assertNotIn("Traceback", result.stderr)
def test_parameter_domain_binding_does_not_cross_scope_or_module(self):
with tempfile.TemporaryDirectory() as temp_dir:
data_dir = Path(temp_dir) / "data"
data_dir.mkdir()
tsv = data_dir / "function_index.tsv"
tsv.write_text(
"name\tscope\tmodule\tsignature\tpage\tanchor\ttags\tsummary\n"
"getBk\tdotnet\tdatawarehouse\tgetBk(marketlist)\t"
"dotnet/market.md\tgetbk\t\t系统板块\n"
"getBk\tdotnet\tdatawarehouse\tgetBk()\t"
"dotnet/market.md\tgetbk-empty\t\t无参数重载\n"
"getBk\tproject\tdemo\tgetBk(name)\tproject/demo.md\t"
"getbk\t\t项目函数\n",
encoding="utf-8",
)
domains = data_dir / "value_domains.json"
domains.write_text(
json.dumps(
{
"version": 1,
"domains": [
{
"id": "market_board",
"label": "市场板块",
"mode": "runtime_catalog",
"complete": False,
"bindings": [
{
"scope": "dotnet",
"module": "datawarehouse",
"api": "getBk",
"parameter": "marketlist",
"role": "current_components",
}
],
"resolvers": [
{
"scope": "dotnet",
"module": "datawarehouse",
"api": "getBkList2",
"parameter": "bktype",
"catalog": "system",
"source": "net_function:28964",
}
],
"related_tables": [],
"values": [
{
"value": "申万煤炭",
"label": "申万煤炭",
"sources": ["faq:31494"],
}
],
"sources": ["faq:31494"],
"as_of": "2026-08-20",
}
],
},
ensure_ascii=False,
),
encoding="utf-8",
)
result = self.run_cli(
LOOKUP,
"--kw",
"申万煤炭",
"--tsv",
str(tsv),
"--value-domains",
str(domains),
)
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("dotnet/market.md#getbk", result.stdout)
self.assertNotIn("dotnet/market.md#getbk-empty", result.stdout)
self.assertNotIn("project/demo.md#getbk", result.stdout)
def test_exact_query_attaches_domain_only_to_overload_with_bound_parameter(self):
with tempfile.TemporaryDirectory() as temp_dir:
root = Path(temp_dir)
data_dir = root / "data"
codegen = root / "references" / "codegen" / "dotnet"
data_dir.mkdir()
codegen.mkdir(parents=True)
tsv = data_dir / "function_index.tsv"
tsv.write_text(
"name\tscope\tmodule\tsignature\tpage\tanchor\ttags\tsummary\n"
"getBk\tdotnet\tdatawarehouse\tgetBk(marketlist)\t"
"dotnet/market.md\tgetbk\t\t系统板块\n"
"getBk\tdotnet\tdatawarehouse\tgetBk()\t"
"dotnet/market.md\tgetbk-1\t\t无参数重载\n",
encoding="utf-8",
)
(codegen / "market.md").write_text(
"# Dotnet\n\n"
"## `getBk(marketlist)`\n\n声明:function\n\n系统板块\n\n"
"## `getBk()`\n\n声明:function\n\n无参数重载\n",
encoding="utf-8",
)
domains = data_dir / "value_domains.json"
domains.write_text(
json.dumps(
{
"version": 1,
"domains": [
{
"id": "market_board",
"label": "市场板块",
"mode": "runtime_catalog",
"complete": False,
"bindings": [
{
"scope": "dotnet",
"module": "datawarehouse",
"api": "getBk",
"parameter": "marketlist",
"role": "current_components",
}
],
"resolvers": [
{
"scope": "dotnet",
"module": "datawarehouse",
"api": "getBkList2",
"parameter": "bktype",
"catalog": "system",
"source": "net_function:28964",
}
],
"related_tables": [],
"values": [],
"sources": ["net_function:28960"],
"as_of": "2026-08-20",
}
],
},
ensure_ascii=False,
),
encoding="utf-8",
)
result = self.run_cli(
LOOKUP,
"--name",
"getBk",
"--tsv",
str(tsv),
"--value-domains",
str(domains),
)
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("## `getBk(marketlist)`", result.stdout)
self.assertIn("## `getBk()`", result.stdout)
self.assertEqual(1, result.stdout.count("### 参数取值域"))
def test_exact_query_does_not_attach_domain_to_same_name_in_other_scope(self):
with tempfile.TemporaryDirectory() as temp_dir:
root = Path(temp_dir)
data_dir = root / "data"
codegen = root / "references" / "codegen"
(codegen / "dotnet").mkdir(parents=True)
(codegen / "project").mkdir(parents=True)
data_dir.mkdir()
tsv = data_dir / "function_index.tsv"
tsv.write_text(
"name\tscope\tmodule\tsignature\tpage\tanchor\ttags\tsummary\n"
"getBk\tdotnet\tdatawarehouse\tgetBk(marketlist)\t"
"dotnet/market.md\tgetbk\t\t系统板块\n"
"getBk\tproject\tdemo\tgetBk(name)\tproject/demo.md\t"
"getbk\t\t项目函数\n",
encoding="utf-8",
)
(codegen / "dotnet" / "market.md").write_text(
"# Dotnet\n\n## `getBk(marketlist)`\n\n声明:function\n\n系统板块\n",
encoding="utf-8",
)
(codegen / "project" / "demo.md").write_text(
"# Project\n\n## `getBk(name)`\n\n声明:function\n\n项目函数\n",
encoding="utf-8",
)
domains = data_dir / "value_domains.json"
domains.write_text(
json.dumps(
{
"version": 1,
"domains": [
{
"id": "market_board",
"label": "市场板块",
"mode": "runtime_catalog",
"complete": False,
"bindings": [
{
"scope": "dotnet",
"module": "datawarehouse",
"api": "getBk",
"parameter": "marketlist",
"role": "current_components",
}
],
"resolvers": [
{
"scope": "dotnet",
"module": "datawarehouse",
"api": "getBkList2",
"parameter": "bktype",
"catalog": "system",
"source": "net_function:28964",
}
],
"related_tables": [],
"values": [],
"sources": ["net_function:28960"],
"as_of": "2026-08-20",
}
],
},
ensure_ascii=False,
),
encoding="utf-8",
)
result = self.run_cli(
LOOKUP,
"--scope",
"project",
"--name",
"getBk",
"--tsv",
str(tsv),
"--value-domains",
str(domains),
)
self.assertEqual(0, result.returncode, result.stderr)
self.assertIn("项目函数", result.stdout)
self.assertNotIn("参数取值域", result.stdout)
def test_bundled_value_domains_reference_existing_apis_parameters_and_tables(self):
document = json.loads(
(SKILL_ROOT / "data" / "value_domains.json").read_text(encoding="utf-8")
)
with (SKILL_ROOT / "data" / "function_index.tsv").open(
encoding="utf-8", newline=""
) as handle:
rows = list(csv.DictReader(handle, delimiter="\t"))
with (SKILL_ROOT / "data" / "dictionary_index.tsv").open(
encoding="utf-8", newline=""
) as handle:
dictionary_rows = list(csv.DictReader(handle, delimiter="\t"))
apis = defaultdict(set)
for row in rows:
match = re.search(r"\((.*)\)", row["signature"])
if not match:
continue
key = (
row["scope"].casefold(),
row["module"].casefold(),
row["name"].casefold(),
)
for parameter in match.group(1).split(","):
parameter = parameter.strip().strip("[]")
if parameter:
apis[key].add(parameter.casefold())
table_pages = {
row["table_id"]: row["page"]
for row in dictionary_rows
if row["kind"] in {"table", "source"} and row["table_id"]
}
for domain in document["domains"]:
for binding in domain["bindings"]:
self.assert_api_parameter_exists(apis, binding)
for resolver in domain["resolvers"]:
self.assert_api_parameter_exists(apis, resolver)
for table in domain["related_tables"]:
self.assertEqual(table["page"], table_pages.get(table["id"]))
self.assertTrue((SKILL_ROOT / table["page"]).is_file())
for value in domain["values"]:
for relation in value.get("relations", []):
self.assert_api_parameter_exists(apis, relation["verification"])
def assert_api_parameter_exists(self, apis, reference):
key = (
reference["scope"].casefold(),
reference["module"].casefold(),
(reference.get("api") or reference.get("resolver")).casefold(),
)
parameter = reference.get("parameter")
self.assertIn(key, apis)
if parameter:
self.assertIn(parameter.casefold(), apis[key])
if __name__ == "__main__":
unittest.main()