Skip to content

Commit 9609a7d

Browse files
authored
fix: infer surviving dims in to_dataset for aggregations (closes #189) (#218)
1 parent d8ffa52 commit 9609a7d

5 files changed

Lines changed: 97 additions & 36 deletions

File tree

‎AGENTS.md‎

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,37 @@
1+
# AGENTS.md
2+
3+
Guidance for contributors (including AI assistants) working on `xarray-sql`. It
4+
summarizes recurring maintainer review feedback so changes land clean.
5+
6+
## Documentation and comments
7+
8+
- Keep docstrings and comments self-contained. Do **not** put GitHub issue or PR
9+
numbers in docstrings or code comments; a reader should not need the issue
10+
tracker to understand the code. Issue references belong in the commit message
11+
and PR description (e.g. `Closes #189`), not in the source.
12+
- Do not reference the review conversation, chat, or "the reporter" in comments.
13+
Describe the behavior, not how it came up.
14+
15+
## API surface
16+
17+
- Mark internal helpers private with a leading underscore when they are not part
18+
of the public API.
19+
20+
## Tests
21+
22+
- Test the public contract (values, dims, coords, attrs), not internal call
23+
counts or private classes, so the suite survives refactors.
24+
- Avoid redundant tests: if a public-path test already covers a behavior, do not
25+
add a second lower-level test for the same thing.
26+
- Make query results deterministic with `ORDER BY` so assertions do not have to
27+
re-sort the output.
28+
- Do not pass `dims=` to `to_dataset()` when inference already resolves them.
29+
Reserve explicit `dims=` / `template=` for genuinely ambiguous cases (multiple
30+
registered Datasets, or a test that is specifically exercising those
31+
arguments).
32+
33+
## Imports
34+
35+
- Keep imports at the top of the file. Assume transitive dependencies are safe
36+
to import non-locally, rather than deferring imports into functions to avoid
37+
a dependency.

‎README.md‎

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -47,9 +47,8 @@ clim = ctx.sql('''
4747
ORDER BY month
4848
''')
4949

50-
# Write the SQL result back to an Xarray Dataset. `month` is a derived
51-
# column, so name it as the dimension; the variable's units are recovered
52-
# from the registered table. The result is one value per month: air(month).
50+
# Round-trip the result back to Xarray. `month` is a derived column, so name
51+
# it as the dimension.
5352
clim_ds = clim.to_dataset(dims=["month"])
5453

5554
# Plot the annual cycle as a time series.
@@ -138,7 +137,9 @@ ctx.sql('''
138137
AND TIMESTAMP '2020-01-01 05:00:00'
139138
GROUP BY latitude, longitude
140139
ORDER BY latitude DESC, longitude
141-
''').to_dataset(dims=['latitude', 'longitude'], template=ds)
140+
# `latitude`/`longitude` are inferred from the registered table's surviving
141+
# dims; `template` is kept only to recover metadata (attrs, encoding).
142+
''').to_dataset(template=ds)
142143
# <xarray.Dataset> Size: 8MB
143144
# Dimensions: (latitude: 721, longitude: 1440)
144145
# Coordinates:

‎docs/examples.md‎

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -34,8 +34,10 @@ clim = ctx.sql('''
3434
clim.to_pandas().head()
3535

3636
# Option 2: round-trip back to an Xarray Dataset and plot the annual cycle as
37-
# a time series. `month` is a derived column, so name it as the dimension; the
38-
# variable's units are recovered from the registered table.
37+
# a time series. `to_dataset()` infers dimensions from the registered table's
38+
# surviving dims, so a GROUP BY on a real dimension needs no `dims=`. Here
39+
# `month` is a derived column, not a registered dim, so name it explicitly;
40+
# the variable's units are recovered from the registered table.
3941
clim_ds = clim.to_dataset(dims=["month"])
4042
clim_ds["air"].plot()
4143
```

‎tests/test_ds.py‎

Lines changed: 25 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
142159
def 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

291308
def 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

396411
def 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

‎xarray_sql/ds.py‎

Lines changed: 26 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -757,10 +757,12 @@ def to_dataset(
757757
758758
Args:
759759
dims: Result columns to use as Dataset dimensions. When
760-
``None``, defaults to the dims of the registered Dataset
761-
referenced by the SQL ``FROM`` clause (if exactly one
762-
matches), or any single registered Dataset whose dims are
763-
all present in the result columns.
760+
``None``, defaults to a registered Dataset's dimensions that
761+
survive into the result columns, so an aggregation that drops
762+
dims (e.g. ``GROUP BY time`` over a ``(time, lat, lon)`` grid)
763+
round-trips on the remaining dim. Raises when no dimension
764+
survives, or when several registered Datasets imply different
765+
dims (pass ``dims`` explicitly then).
764766
template: Source to recover metadata (attrs, encoding, non-dim
765767
coordinates, dim-coord dtype) from. Either an ``xr.Dataset``
766768
used directly, or the name of a registered table (e.g.
@@ -879,33 +881,37 @@ def _infer_dimension_columns(
879881
) -> list[str]:
880882
"""Pick a default ``dimension_columns`` from the registry, or raise.
881883
882-
Uses the data variable's dim order (via :func:`_ds_var_dims`) so
883-
the round-trip preserves the original axis order.
884+
A registered Dataset's dims that survive into the result columns
885+
become the dimensions, so aggregations that drop dims (e.g.
886+
``GROUP BY time`` over a ``(time, lat, lon)`` grid) round-trip on the
887+
surviving dim(s). Uses the data variable's dim order (via
888+
:func:`_ds_var_dims`) so the original axis order is preserved.
884889
"""
885890
result_cols = set(self._result_columns())
886-
if (
887-
preferred_template is not None
888-
and set(preferred_template.dims) <= result_cols
889-
):
890-
return _ds_var_dims(preferred_template)
891+
892+
def surviving(template: xr.Dataset) -> list[str]:
893+
# Template dims still present in the result, in var axis order.
894+
return [d for d in _ds_var_dims(template) if d in result_cols]
895+
896+
if preferred_template is not None:
897+
preferred = surviving(preferred_template)
898+
if preferred:
899+
return preferred
891900
if not self._templates:
892901
raise ValueError(
893902
"dims cannot be inferred (no registered "
894903
"Dataset on this result); pass dims=[...] "
895904
"explicitly."
896905
)
897-
candidates = [
898-
_ds_var_dims(t)
899-
for t in self._templates.values()
900-
if set(t.dims) <= result_cols
901-
]
906+
candidates = {tuple(surviving(t)) for t in self._templates.values()}
907+
candidates.discard(()) # templates with no surviving dim
902908
if len(candidates) == 1:
903-
return candidates[0]
909+
return list(next(iter(candidates)))
904910
if not candidates:
905911
raise ValueError(
906-
"dims cannot be inferred: no registered "
907-
"Dataset has all of its dims present in the result "
908-
"columns. Pass dims=[...] explicitly."
912+
"dims cannot be inferred: no registered Dataset "
913+
"dimension survives in the result columns. Pass "
914+
"dims=[...] explicitly."
909915
)
910916
raise ValueError(
911917
"dims cannot be inferred unambiguously: multiple "

0 commit comments

Comments
 (0)