Skip to content

Commit 78051ec

Browse files
committed
feat: support combining results from multiple runs
Closes #63.
1 parent 19cb2dc commit 78051ec

1 file changed

Lines changed: 24 additions & 30 deletions

File tree

adbc_drivers_validation/generate_documentation.py

Lines changed: 24 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -282,11 +282,10 @@ def render_part(template, kwargs) -> str:
282282

283283

284284
def load_testcases(
285-
quirks: model.DriverQuirks, results_path: Path, query_set: model.QuerySet
285+
get_quirks: typing.Callable[[str], model.DriverQuirks], results_path: Path
286286
) -> None:
287287
"""Load test case data into DuckDB."""
288288
report = xml.etree.ElementTree.parse(results_path).getroot()
289-
driver_name = f"{quirks.name}:{quirks.short_version}"
290289

291290
testcases = []
292291
for testcase in report.findall(".//testsuite[@name='validation']/testcase"):
@@ -304,9 +303,9 @@ def load_testcases(
304303
for prop in testcase.findall(".//properties/property"):
305304
properties[prop.get("name")] = prop.get("value")
306305

307-
driver = properties["driver"]
308-
if driver != driver_name:
309-
continue
306+
driver, _, version = properties["driver"].partition(":")
307+
quirks = get_quirks(version)
308+
query_set = quirks.query_set
310309
driver = quirks.name
311310
version = quirks.short_version
312311

@@ -511,18 +510,18 @@ def render(
511510

512511

513512
def generate_includes(
514-
all_quirks: list[model.DriverQuirks], query_sets: dict[str, model.QuerySet]
513+
get_quirks: typing.Callable[[str], model.DriverQuirks],
515514
) -> ValidationReport:
516515
# Handle different versions of one vendor
517-
report = ValidationReport(
518-
driver=all_quirks[0].name,
519-
versions={
520-
quirks.short_version: DriverTypeTable(
521-
quirks=quirks, features=quirks.features
522-
)
523-
for quirks in all_quirks
524-
},
525-
)
516+
driver = duckdb.sql("FROM testcases SELECT DISTINCT driver").fetchall()
517+
assert len(driver) == 1, f"Expected exactly one driver, got {driver}"
518+
driver = driver[0][0]
519+
versions = {}
520+
for row in duckdb.sql("FROM testcases SELECT DISTINCT vendor_version").fetchall():
521+
quirks = get_quirks(row[0])
522+
versions[row[0]] = DriverTypeTable(quirks=quirks, features=quirks.features)
523+
524+
report = ValidationReport(driver=driver, versions=versions)
526525

527526
# Version
528527
version = (
@@ -583,7 +582,9 @@ def generate_includes(
583582
for test_case in type_tests:
584583
arrow_type_names = set()
585584
for query_name in test_case["query_names"]:
586-
query = query_sets[test_case["vendor_version"]].queries[query_name]
585+
query = get_quirks(test_case["vendor_version"]).query_set.queries[
586+
query_name
587+
]
587588
show_type_parameters = query.metadata().tags.show_arrow_type_parameters
588589

589590
# Take the first field; some queries may select additional things
@@ -640,7 +641,7 @@ def generate_includes(
640641
.to_pylist()
641642
)
642643
for test_case in type_tests:
643-
query_set = query_sets[test_case["vendor_version"]]
644+
query_set = get_quirks(test_case["vendor_version"]).query_set
644645
arrow_type_names = set()
645646
for query_name in test_case["query_names"]:
646647
arrow_type_names.add(query_set.queries[query_name].arrow_type_name)
@@ -676,7 +677,7 @@ def generate_includes(
676677
.to_pylist()
677678
)
678679
for test_case in type_tests:
679-
query_set = query_sets[test_case["vendor_version"]]
680+
query_set = get_quirks(test_case["vendor_version"]).query_set
680681
query_name = test_case["query_name"]
681682
arrow_type = html.escape(test_case["arrow_type_name"])
682683
sql_type = html.escape(test_case["sql_type"])
@@ -798,20 +799,13 @@ def generate_includes(
798799

799800

800801
def generate(
801-
all_quirks: list[model.DriverQuirks],
802-
test_results: Path,
802+
get_quirks: typing.Callable[[str], model.DriverQuirks],
803+
test_results: list[Path],
803804
driver_template: Path,
804805
output: Path,
805806
) -> None:
806-
if len({quirks.name for quirks in all_quirks}) != 1:
807-
raise ValueError("All quirks must be for the same driver")
808-
if len({quirks.short_version for quirks in all_quirks}) != len(all_quirks):
809-
raise ValueError("All quirks must be for the different versions")
810-
811-
query_sets = {}
812-
for quirks in all_quirks:
813-
load_testcases(quirks, test_results, quirks.query_set)
814-
query_sets[quirks.short_version] = quirks.query_set
815-
report = generate_includes(all_quirks, query_sets)
807+
for results in test_results:
808+
load_testcases(get_quirks, results)
809+
report = generate_includes(get_quirks)
816810
print(report.pprint())
817811
render(report, driver_template, output)

0 commit comments

Comments
 (0)