@@ -282,11 +282,10 @@ def render_part(template, kwargs) -> str:
282282
283283
284284def 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
513512def 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
800801def 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