@@ -125,7 +125,7 @@ def test_aggregation_drops_dim(air_dataset_small):
125125 ctx .from_dataset ("air" , air_dataset_small )
126126 out = ctx .sql (
127127 "SELECT lat, lon, AVG(air) AS air_avg FROM air GROUP BY lat, lon"
128- ).to_dataset (dims = [ "lat" , "lon" ] )
128+ ).to_dataset ()
129129 assert set (out .dims ) == {"lat" , "lon" }
130130 assert "air_avg" in out .data_vars
131131 assert "air" not in out .data_vars
@@ -139,6 +139,23 @@ def test_aggregation_drops_dim(air_dataset_small):
139139 np .testing .assert_allclose (actual , expected )
140140
141141
142+ def test_aggregation_infers_dims (air_dataset_small ):
143+ """to_dataset() infers the surviving GROUP BY dim when dims is omitted."""
144+ ctx = XarrayContext ()
145+ ctx .from_dataset ("air" , air_dataset_small )
146+
147+ # Grouping by the time coordinate keeps time as the sole dimension; the
148+ # ORDER BY makes the result order deterministic so no sort is needed below.
149+ out = ctx .sql (
150+ 'SELECT "time", AVG("air") AS air FROM "air" '
151+ 'GROUP BY "time" ORDER BY "time"'
152+ ).to_dataset ()
153+ assert set (out .dims ) == {"time" }
154+ assert "air" in out .data_vars
155+ expected = air_dataset_small .compute ().mean (dim = ["lat" , "lon" ])["air" ]
156+ np .testing .assert_allclose (out ["air" ].values , expected .values )
157+
158+
142159def test_barrier_query_scans_source_once (air_dataset_small ):
143160 """A barrier plan (aggregation) executes the source exactly once.
144161
@@ -166,7 +183,7 @@ def test_barrier_query_scans_source_once(air_dataset_small):
166183
167184 out = ctx .sql (
168185 "SELECT lat, lon, AVG(air) AS air_avg FROM air GROUP BY lat, lon"
169- ).to_dataset (dims = [ "lat" , "lon" ] )
186+ ).to_dataset ()
170187 reads_after_construct = len (reads )
171188 out .compute ()
172189 reads_after_compute = len (reads )
@@ -188,7 +205,7 @@ def test_order_by_direction_sets_dim_order(air_dataset_small):
188205 ctx .from_dataset ("air" , air_dataset_small )
189206 out = ctx .sql (
190207 "SELECT lat, AVG(air) AS air_avg FROM air GROUP BY lat ORDER BY lat DESC"
191- ).to_dataset (dims = [ "lat" ] )
208+ ).to_dataset ()
192209
193210 lat = out ["lat" ].values
194211 assert (np .diff (lat ) < 0 ).all (), f"expected descending lat, got { lat } "
@@ -289,7 +306,7 @@ def test_fast_path_uses_scanned_tables_coords_not_user_template(
289306
290307
291308def test_round_trip_preserves_descending_lat_on_lazy_path (air_dataset_small ):
292- """Lazy round-trip preserves source dim order (xarray-sql#171) .
309+ """Lazy round-trip preserves source dim order.
293310
294311 NCEP ``air_temperature`` ships descending lat (75.0 -> 15.0). The
295312 discovery path's ``.distinct().sort()`` previously flipped lat to
@@ -383,14 +400,12 @@ def test_to_dataset_multi_registered_requires_explicit_template(
383400 assert set (out .dims ) == {"time" , "lat" , "lon" }
384401
385402
386- def test_to_dataset_infer_fails_when_no_template_fits (air_dataset_small ):
387- """If no registered Dataset's dims fit the result -> clear error."""
403+ def test_to_dataset_infer_fails_when_no_dim_survives (air_dataset_small ):
404+ """A global aggregation leaves no registered dim in the result -> clear error."""
388405 ctx = XarrayContext ()
389406 ctx .from_dataset ("air" , air_dataset_small )
390407 with pytest .raises (ValueError , match = "dims cannot be inferred" ):
391- ctx .sql (
392- "SELECT lat, lon, AVG(air) AS air_avg FROM air GROUP BY lat, lon"
393- ).to_dataset ()
408+ ctx .sql ("SELECT AVG(air) AS air_avg FROM air" ).to_dataset ()
394409
395410
396411def test_template_accepts_name_or_dataset (air_dataset_small ):
@@ -447,7 +462,7 @@ def test_template_aggregation_alias_no_attrs(air_dataset_small):
447462 ctx .from_dataset ("air" , ds )
448463 out = ctx .sql (
449464 "SELECT lat, lon, AVG(air) AS air_avg FROM air GROUP BY lat, lon"
450- ).to_dataset (dims = [ "lat" , "lon" ] )
465+ ).to_dataset ()
451466 assert "air_avg" in out .data_vars
452467 assert out ["air_avg" ].attrs == {}
453468
0 commit comments