Skip to content

Commit 051fecf

Browse files
authored
Merge pull request #2045 from dbcli/RW/save-plots-like-parquets
Ability to save plots with the `.>` operator
2 parents 7a43084 + cbfcde9 commit 051fecf

6 files changed

Lines changed: 202 additions & 59 deletions

File tree

‎changelog.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ Features
77
* Allow file target of `$>` redirection to be quoted.
88
* Display of inline plots returned from `.|` operations.
99
* Don't attempt inline plots in the Windows console.
10+
* Save plots, like parquets, with the `.>` operator.
1011

1112

1213
Bug Fixes

‎doc/transforms.md‎

Lines changed: 13 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -22,9 +22,8 @@ functionality may still change.
2222
Here are some known limitations:
2323

2424
* transforms can't be composed with `$|` shell redirection
25-
* multiple transform steps are not permitted
25+
* multiple transform steps are not supported
2626
* results from `UNION`s may be unable to be transformed
27-
* images cannot yet be saved
2827
* PNG images are static and do not support all Altair features
2928

3029
And there are inherent limitations to the post-processing model: the entire
@@ -84,30 +83,31 @@ the `[dataframe]` section of `~/.myclirc`.
8483
## Saving
8584

8685
A query result, transformed `DataFrame`, or transformed `Series` can be
87-
written directly to a Parquet file with the `.>` operator.
86+
written directly to a Parquet file with the `.>` operator. An Altair plot can
87+
also be written to a PNG file with the same operator.
8888

8989
Save example:
9090

9191
```sql
9292
SELECT * FROM orders .> orders.parquet;
9393
```
9494

95-
The `.>` operator must be last, requires a `.parquet` destination, and
95+
The `.>` operator must be last, requires a `.parquet` or `.png` destination, and
9696
overwrites any existing file. Spaces may be required around the operator.
9797
Destination paths containing whitespace must be quoted. A successful write
98-
reports its destination and row count.
98+
reports its destination and row count if appropriate.
9999

100100
When `post_redirect_command` is set in `~/.myclirc`, the given command runs
101-
after a successful Parquet save.
101+
after a successful file save.
102102

103103
`.>` cannot be combined with the `\x` or `\G` special display terminators.
104104

105105
## Combining
106106

107-
Parquet saves may be combined with dataframe transforms. Again, `.>` must
108-
be the last operator.
107+
Parquet saves may be combined with dataframe transforms. Again, the save
108+
operator `.>` must be the last operator.
109109

110-
Combined transform and save example:
110+
Combined transform and save examples:
111111

112112
```sql
113113
SELECT * FROM orders .| df.group_by('customer_id').len() .> customer_counts.parquet;
@@ -116,3 +116,7 @@ SELECT * FROM orders .| df.group_by('customer_id').len() .> customer_counts.parq
116116
```sql
117117
SELECT * FROM orders .| df['order_id'] .> order_ids.parquet;
118118
```
119+
120+
```sql
121+
SELECT * FROM orders .| df['total'].plot.hist() .> total_histogram.png;
122+
```

‎mycli/main_modes/repl.py‎

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -750,7 +750,7 @@ def _one_iteration(
750750
mycli.explorer_formatter.query = text
751751
if polars_transform is not None:
752752
assert polars_pipeline is not None
753-
if polars_pipeline.parquet_path is None:
753+
if polars_pipeline.output_path is None:
754754
polars_result = run_polars_transform(
755755
polars_transform,
756756
results,
@@ -763,22 +763,22 @@ def _one_iteration(
763763
polars_result = run_polars_transform(
764764
polars_transform,
765765
results,
766-
polars_pipeline.parquet_path,
766+
polars_pipeline.output_path,
767767
image_protocol=mycli.image_protocol,
768768
plot_scale_factor=mycli.plot_scale_factor,
769769
plot_ppi=mycli.plot_ppi,
770770
plot_theme=mycli.plot_theme,
771771
)
772-
if polars_pipeline.parquet_path is None:
772+
if polars_pipeline.output_path is None:
773773
if polars_pipeline.output_mode == 'explorer':
774774
special.set_explorer_output(True)
775775
elif polars_pipeline.output_mode == 'expanded':
776776
special.set_expanded_output(True)
777777
_output_results(mycli, state, iter([polars_result]), start)
778-
if polars_pipeline.parquet_path is not None:
778+
if polars_pipeline.output_path is not None:
779779
special.run_post_redirect_hook(
780780
mycli.post_redirect_command,
781-
polars_pipeline.parquet_path,
781+
polars_pipeline.output_path,
782782
)
783783
else:
784784
_output_results(mycli, state, results, start)

‎mycli/packages/polars_transform.py‎

Lines changed: 53 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@ class PolarsTransformError(RuntimeError):
2323
class PolarsPipeline:
2424
sql: str
2525
expression: str | None
26-
parquet_path: str | None
26+
output_path: str | None
2727
output_mode: OutputMode
2828

2929

@@ -37,7 +37,7 @@ class PolarsTransform:
3737

3838

3939
def parse_polars_transform(command: str) -> PolarsPipeline | None:
40-
"""Parse a SQL statement with optional Polars transform and Parquet output."""
40+
"""Parse a SQL statement with optional Polars transform and file output."""
4141
try:
4242
tokens = sqlglot.tokenize(command)
4343
except sqlglot.errors.TokenError as exc:
@@ -56,7 +56,7 @@ def parse_polars_transform(command: str) -> PolarsPipeline | None:
5656
if following.end + 1 >= len(command):
5757
if following.token_type == sqlglot.TokenType.PIPE:
5858
raise PolarsTransformError('Polars transforms require a Python expression.')
59-
raise PolarsTransformError('Parquet saves require a destination path.')
59+
raise PolarsTransformError('File saves require a destination path.')
6060
if not command[following.end + 1].isspace():
6161
continue
6262
if following.token_type == sqlglot.TokenType.PIPE:
@@ -65,7 +65,7 @@ def parse_polars_transform(command: str) -> PolarsPipeline | None:
6565
pipe_index = index
6666
else:
6767
if parquet_index is not None:
68-
raise PolarsTransformError('Parquet saves support only one ".>" operator.')
68+
raise PolarsTransformError('File saves support only one ".>" operator.')
6969
parquet_index = index
7070

7171
if pipe_index is None and parquet_index is None:
@@ -77,20 +77,20 @@ def parse_polars_transform(command: str) -> PolarsPipeline | None:
7777
assert first_index is not None
7878
sql = command[: tokens[first_index].start].strip()
7979
expression: str | None = None
80-
parquet_path: str | None = None
80+
output_path: str | None = None
8181
if pipe_index is not None:
8282
pipe_operator = tokens[pipe_index + 1]
8383
expression_end = tokens[parquet_index].start if parquet_index is not None else len(command)
8484
expression = command[pipe_operator.end + 1 : expression_end].strip()
8585
if parquet_index is not None:
8686
parquet_operator = tokens[parquet_index + 1]
87-
parquet_path = command[parquet_operator.end + 1 :].strip()
88-
parquet_path = parquet_path.removesuffix(delimiter_command.current).rstrip()
89-
parquet_path = parquet_path.removesuffix(r'\g').rstrip()
87+
output_path = command[parquet_operator.end + 1 :].strip()
88+
output_path = output_path.removesuffix(delimiter_command.current).rstrip()
89+
output_path = output_path.removesuffix(r'\g').rstrip()
9090

91-
has_display_terminator = any(value is not None and value.endswith((r'\x', r'\G')) for value in (sql, expression, parquet_path))
92-
if parquet_path is not None and has_display_terminator:
93-
raise PolarsTransformError('Parquet saves cannot use special display terminators.')
91+
has_display_terminator = any(value is not None and value.endswith((r'\x', r'\G')) for value in (sql, expression, output_path))
92+
if output_path is not None and has_display_terminator:
93+
raise PolarsTransformError('File saves cannot use special display terminators.')
9494
if sql.endswith(r'\x') or expression is not None and expression.endswith(r'\x'):
9595
output_mode: OutputMode = 'explorer'
9696
elif sql.endswith(r'\G') or expression is not None and expression.endswith(r'\G'):
@@ -107,23 +107,23 @@ def parse_polars_transform(command: str) -> PolarsPipeline | None:
107107
raise PolarsTransformError('Polars transforms require a SQL statement.')
108108
if expression is not None and not expression:
109109
raise PolarsTransformError('Polars transforms require a Python expression.')
110-
if parquet_path is not None:
111-
if not parquet_path:
112-
raise PolarsTransformError('Parquet saves require a destination path.')
113-
parquet_path = _parse_parquet_path(parquet_path)
110+
if output_path is not None:
111+
if not output_path:
112+
raise PolarsTransformError('File saves require a destination path.')
113+
output_path = _parse_output_path(output_path)
114114
_validate_sql(sql)
115-
return PolarsPipeline(sql=sql, expression=expression, parquet_path=parquet_path, output_mode=output_mode)
115+
return PolarsPipeline(sql=sql, expression=expression, output_path=output_path, output_mode=output_mode)
116116

117117

118-
def _parse_parquet_path(path: str) -> str:
118+
def _parse_output_path(path: str) -> str:
119119
if path[0] in ('\'', '"'):
120120
if len(path) < 2 or path[-1] != path[0]:
121-
raise PolarsTransformError('Parquet save paths must use matching quotes.')
121+
raise PolarsTransformError('File save paths must use matching quotes.')
122122
path = path[1:-1]
123123
elif any(character.isspace() for character in path):
124-
raise PolarsTransformError('Parquet save paths containing spaces must be quoted.')
125-
if not path.lower().endswith('.parquet'):
126-
raise PolarsTransformError('Parquet save paths must end in ".parquet".')
124+
raise PolarsTransformError('File save paths containing spaces must be quoted.')
125+
if not path.lower().endswith(('.parquet', '.png')):
126+
raise PolarsTransformError('File save paths must end in ".parquet" or ".png".')
127127
return path
128128

129129

@@ -181,7 +181,7 @@ def _load_vl_convert() -> None:
181181
def run_polars_transform(
182182
transform: PolarsTransform,
183183
results: Iterable[SQLResult],
184-
parquet_path: str | None = None,
184+
output_path: str | None = None,
185185
*,
186186
image_protocol: ImageProtocol = 'none',
187187
plot_scale_factor: float = 1.0,
@@ -212,34 +212,49 @@ def run_polars_transform(
212212
except Exception as exc:
213213
raise PolarsTransformError(f'Polars expression failed: {type(exc).__name__}: {exc}') from exc
214214
if isinstance(value, transform.polars.DataFrame):
215-
if parquet_path is not None:
215+
if output_path is not None:
216+
if not output_path.lower().endswith('.parquet'):
217+
raise PolarsTransformError('Polars DataFrame results can only be written to ".parquet" files.')
216218
try:
217-
value.write_parquet(parquet_path)
219+
value.write_parquet(output_path)
218220
except Exception as exc:
219-
raise PolarsTransformError(f'Unable to write Parquet file "{parquet_path}": {type(exc).__name__}: {exc}') from exc
220-
return SQLResult(status=f'Wrote {len(value)} rows to {parquet_path}.')
221+
raise PolarsTransformError(f'Unable to write Parquet file "{output_path}": {type(exc).__name__}: {exc}') from exc
222+
return SQLResult(status=f'Wrote {len(value)} rows to {output_path}.')
221223
return SQLResult(header=list(value.columns), rows=list(value.iter_rows()))
222224
if isinstance(value, transform.polars.Series):
223225
column_name = value.name or 'value'
224-
if parquet_path is not None:
226+
if output_path is not None:
227+
if not output_path.lower().endswith('.parquet'):
228+
raise PolarsTransformError('Polars Series results can only be written to ".parquet" files.')
225229
try:
226230
series_dataframe = value.rename(column_name).to_frame()
227-
series_dataframe.write_parquet(parquet_path)
231+
series_dataframe.write_parquet(output_path)
228232
except Exception as exc:
229-
raise PolarsTransformError(f'Unable to write Parquet file "{parquet_path}": {type(exc).__name__}: {exc}') from exc
230-
return SQLResult(status=f'Wrote {len(series_dataframe)} rows to {parquet_path}.')
233+
raise PolarsTransformError(f'Unable to write Parquet file "{output_path}": {type(exc).__name__}: {exc}') from exc
234+
return SQLResult(status=f'Wrote {len(series_dataframe)} rows to {output_path}.')
231235
return SQLResult(header=[column_name], rows=[(item,) for item in value])
232236
if transform.altair is not None and isinstance(value, transform.altair.TopLevelMixin):
233-
if parquet_path is not None:
234-
raise PolarsTransformError('Polars transforms must return a DataFrame or Series before writing Parquet output.')
235-
if image_protocol == 'none':
237+
if output_path is not None and not output_path.lower().endswith('.png'):
238+
raise PolarsTransformError('Altair charts can only be written to ".png" files.')
239+
if output_path is None and image_protocol == 'none':
236240
return SQLResult(status='image_protocol is unset in ~/.myclirc. Inline plotting is disabled.')
237241
_load_vl_convert()
238-
png = BytesIO()
239242
try:
240243
transform.altair.theme.enable(plot_theme)
241244
except Exception as exc:
242245
raise PolarsTransformError(f'Unable to enable Altair plot theme "{plot_theme}": {type(exc).__name__}: {exc}') from exc
246+
if output_path is not None:
247+
try:
248+
value.save(
249+
output_path,
250+
format='png',
251+
scale_factor=plot_scale_factor,
252+
ppi=plot_ppi,
253+
)
254+
except Exception as exc:
255+
raise PolarsTransformError(f'Unable to write PNG file "{output_path}": {type(exc).__name__}: {exc}') from exc
256+
return SQLResult(status=f'Wrote PNG image to {output_path}.')
257+
png = BytesIO()
243258
try:
244259
value.save(
245260
png,
@@ -250,6 +265,8 @@ def run_polars_transform(
250265
except Exception as exc:
251266
raise PolarsTransformError(f'Unable to render Altair chart: {type(exc).__name__}: {exc}') from exc
252267
return SQLResult(image=png.getvalue(), image_protocol=image_protocol)
253-
if parquet_path is not None:
254-
raise PolarsTransformError('Polars transforms must return a DataFrame or Series before writing Parquet output.')
268+
if output_path is not None:
269+
raise PolarsTransformError(
270+
'Polars transforms must return a DataFrame or Series for Parquet output, or an Altair chart for PNG output.'
271+
)
255272
return SQLResult(status=f'Nothing could be displayed for return type: {type(value)}')

‎test/pytests/test_main_modes_repl.py‎

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1374,6 +1374,61 @@ def run(
13741374
assert hook_calls == [('post {}', 'orders.parquet')]
13751375

13761376

1377+
def test_one_iteration_writes_polars_png_and_runs_post_redirect_hook(monkeypatch: pytest.MonkeyPatch) -> None:
1378+
patch_repl_runtime_defaults(monkeypatch)
1379+
1380+
class FakeSQLExecute:
1381+
dbname = 'db'
1382+
connection_id = 0
1383+
1384+
def run(self, text: str) -> Iterator[SQLResult]:
1385+
assert text == 'SELECT * FROM orders'
1386+
return iter([SQLResult(header=['id'], rows=[(1,)])])
1387+
1388+
cli = make_repl_cli(FakeSQLExecute())
1389+
cli.post_redirect_command = 'post {}'
1390+
transform = object()
1391+
run_calls: list[str] = []
1392+
hook_calls: list[tuple[str, str]] = []
1393+
1394+
monkeypatch.setattr(repl_mode, 'prepare_polars_transform', lambda sql, expression: transform)
1395+
1396+
def run(
1397+
received_transform: object,
1398+
results: Iterator[SQLResult],
1399+
path: str,
1400+
*,
1401+
image_protocol: str,
1402+
plot_scale_factor: float,
1403+
plot_ppi: int,
1404+
plot_theme: str,
1405+
) -> SQLResult:
1406+
assert received_transform is transform
1407+
assert list(results) == [SQLResult(header=['id'], rows=[(1,)])]
1408+
assert image_protocol == 'none'
1409+
assert plot_scale_factor == 1.0
1410+
assert plot_ppi == 200
1411+
assert plot_theme == 'carbong90'
1412+
run_calls.append(path)
1413+
return SQLResult(status=f'Wrote PNG image to {path}.')
1414+
1415+
monkeypatch.setattr(repl_mode, 'run_polars_transform', run)
1416+
monkeypatch.setattr(
1417+
repl_mode.special,
1418+
'run_post_redirect_hook',
1419+
lambda command, filename: hook_calls.append((command, filename)),
1420+
)
1421+
1422+
repl_mode._one_iteration(
1423+
cli,
1424+
repl_mode.ReplState(),
1425+
'SELECT * FROM orders .| alt.Chart(df) .> orders.png',
1426+
)
1427+
1428+
assert run_calls == ['orders.png']
1429+
assert hook_calls == [('post {}', 'orders.png')]
1430+
1431+
13771432
def test_one_iteration_reports_polars_post_redirect_hook_error(monkeypatch: pytest.MonkeyPatch) -> None:
13781433
patch_repl_runtime_defaults(monkeypatch)
13791434

0 commit comments

Comments
 (0)