streamlit-pivot 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,942 @@
1
+ # Copyright 2025 Snowflake Inc.
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+
16
+ """Streamlit Pivot Table -- BI-focused pivot table component for Streamlit."""
17
+
18
+ from __future__ import annotations
19
+
20
+ import json
21
+ import warnings
22
+ from typing import TYPE_CHECKING, Any, TypedDict, cast
23
+
24
+ if TYPE_CHECKING:
25
+ from collections.abc import Callable
26
+
27
+ try:
28
+ import streamlit as st
29
+ from streamlit.dataframe_util import convert_anything_to_pandas_df
30
+ except ImportError as e:
31
+ raise ImportError(
32
+ "streamlit_pivot requires Streamlit >= 1.51. "
33
+ "Install it with: pip install 'streamlit>=1.51'"
34
+ ) from e
35
+
36
+
37
+ # ---------------------------------------------------------------------------
38
+ # Config schema v1
39
+ # ---------------------------------------------------------------------------
40
+
41
+ CONFIG_SCHEMA_VERSION = 1
42
+
43
+
44
+ class SortConfig(TypedDict, total=False):
45
+ """Sort configuration for a pivot axis (rows or columns).
46
+
47
+ When ``by="key"`` (default), rows/columns are sorted alphabetically by
48
+ their dimension labels. When ``by="value"``, they are sorted by an
49
+ aggregated measure (optionally restricted to a specific column via
50
+ ``col_key`` for row sorting).
51
+ """
52
+
53
+ by: str # "key" | "value"
54
+ direction: str # "asc" | "desc"
55
+ value_field: str # required when by="value"
56
+ col_key: list[str] # for row sort: sort within this specific column
57
+ dimension: str # optional: scope sort to this dimension level and below
58
+
59
+
60
+ class PivotConfig(TypedDict, total=False):
61
+ """Versioned configuration schema for the pivot table (v1).
62
+
63
+ Persisted via setStateValue("config", ...) and round-tripped through
64
+ Python session state. Every field has an explicit default so partial
65
+ configs are safe.
66
+ """
67
+
68
+ version: int # always 1 for this schema
69
+ rows: list[str]
70
+ columns: list[str]
71
+ values: list[str]
72
+ synthetic_measures: list[dict[str, Any]]
73
+ aggregation: dict[str, str]
74
+ show_totals: bool
75
+ show_row_totals: bool | list[str]
76
+ show_column_totals: bool | list[str]
77
+ empty_cell_value: str
78
+ interactive: bool
79
+ row_sort: SortConfig
80
+ col_sort: SortConfig
81
+ show_subtotals: bool | list[str]
82
+ repeat_row_labels: bool
83
+ collapsed_groups: list[str]
84
+ collapsed_col_groups: list[str]
85
+ sticky_headers: bool
86
+ show_values_as: dict[str, str]
87
+ conditional_formatting: list[dict[str, Any]]
88
+ number_format: dict[str, str]
89
+ column_alignment: dict[str, str]
90
+
91
+
92
+ VALID_AGGREGATIONS = frozenset(
93
+ (
94
+ "sum",
95
+ "avg",
96
+ "count",
97
+ "min",
98
+ "max",
99
+ "count_distinct",
100
+ "median",
101
+ "percentile_90",
102
+ "first",
103
+ "last",
104
+ )
105
+ )
106
+ VALID_SHOW_VALUES_AS = frozenset(("raw", "pct_of_total", "pct_of_row", "pct_of_col"))
107
+ VALID_ALIGNMENTS = frozenset(("left", "center", "right"))
108
+ VALID_COND_FMT_TYPES = frozenset(("color_scale", "data_bars", "threshold"))
109
+ VALID_NULL_MODES = frozenset(("exclude", "zero", "separate"))
110
+
111
+
112
+ _warned_keys: set[str] = set()
113
+ _PYTHON_CONFIG_STATE_PREFIX = "__streamlit_pivot_python_config__:"
114
+
115
+
116
+ def _normalize_aggregation_config(
117
+ aggregation: str | dict[str, str] | None,
118
+ values: list[str],
119
+ ) -> dict[str, str]:
120
+ """Normalize aggregation input to a canonical per-value map."""
121
+ if not values:
122
+ return {}
123
+ if aggregation is None:
124
+ return {value: "sum" for value in values}
125
+ if isinstance(aggregation, str):
126
+ if aggregation not in VALID_AGGREGATIONS:
127
+ raise ValueError(
128
+ f"aggregation must be one of {sorted(VALID_AGGREGATIONS)}, got {aggregation!r}"
129
+ )
130
+ return {value: aggregation for value in values}
131
+ if not isinstance(aggregation, dict):
132
+ raise TypeError(
133
+ f"aggregation must be a str, dict[str, str], or None, got {type(aggregation).__name__}"
134
+ )
135
+ normalized: dict[str, str] = {}
136
+ for value in values:
137
+ agg = aggregation.get(value, "sum")
138
+ if not isinstance(agg, str):
139
+ raise TypeError(
140
+ f"aggregation[{value!r}] must be a string, got {type(agg).__name__}"
141
+ )
142
+ if agg not in VALID_AGGREGATIONS:
143
+ raise ValueError(
144
+ f"aggregation[{value!r}] must be one of {sorted(VALID_AGGREGATIONS)}, got {agg!r}"
145
+ )
146
+ normalized[value] = agg
147
+ return normalized
148
+
149
+
150
+ def _normalize_config_aggregation(config: Any) -> Any:
151
+ """Normalize a config object's aggregation field in-place-compatible form."""
152
+ if not isinstance(config, dict):
153
+ return config
154
+ values = config.get("values")
155
+ value_list = (
156
+ [v for v in values if isinstance(v, str)] if isinstance(values, list) else []
157
+ )
158
+ normalized = dict(config)
159
+ normalized["aggregation"] = _normalize_aggregation_config(
160
+ normalized.get("aggregation"),
161
+ value_list,
162
+ )
163
+ return normalized
164
+
165
+
166
+ def _stable_config_json(config: Any) -> str:
167
+ """Serialize a config deterministically after aggregation normalization."""
168
+ return json.dumps(
169
+ _normalize_config_aggregation(config),
170
+ sort_keys=True,
171
+ separators=(",", ":"),
172
+ )
173
+
174
+
175
+ def _resolve_config_to_send(
176
+ session_state: Any,
177
+ key: str,
178
+ initial_config: PivotConfig,
179
+ ) -> PivotConfig:
180
+ """Resolve Python config vs persisted frontend config precedence.
181
+
182
+ Preserve persisted user config across normal reruns, but when Python sends a
183
+ new config for the same component key, prefer that new Python config.
184
+ """
185
+ tracker_key = f"{_PYTHON_CONFIG_STATE_PREFIX}{key}"
186
+ initial_json = _stable_config_json(initial_config)
187
+
188
+ previous_python_json = None
189
+ try:
190
+ previous_python_json = session_state.get(tracker_key)
191
+ except (AttributeError, TypeError):
192
+ previous_python_json = None
193
+
194
+ persisted_config = None
195
+ try:
196
+ persisted_config = session_state.get(key, {}).get("config")
197
+ except (AttributeError, TypeError):
198
+ persisted_config = None
199
+
200
+ normalized_persisted = (
201
+ _normalize_config_aggregation(persisted_config)
202
+ if persisted_config is not None
203
+ else None
204
+ )
205
+ python_config_changed = (
206
+ previous_python_json is not None and previous_python_json != initial_json
207
+ )
208
+
209
+ try:
210
+ session_state[tracker_key] = initial_json
211
+ except Exception:
212
+ pass
213
+
214
+ if normalized_persisted is None or python_config_changed:
215
+ return initial_config
216
+ return normalized_persisted
217
+
218
+
219
+ def _validate_list_field(
220
+ items: list[str],
221
+ valid: list[str],
222
+ param_name: str,
223
+ valid_label: str,
224
+ ) -> bool | list[str]:
225
+ """Filter list to valid members, warn once per unknown entry, normalize."""
226
+ filtered: list[str] = []
227
+ for item in items:
228
+ if item in valid:
229
+ filtered.append(item)
230
+ else:
231
+ warn_key = f"{param_name}:{item}"
232
+ if warn_key not in _warned_keys:
233
+ _warned_keys.add(warn_key)
234
+ warnings.warn(
235
+ f"{param_name}: ignoring unknown entry {item!r} "
236
+ f"— valid {valid_label} are {valid}",
237
+ stacklevel=4,
238
+ )
239
+ if len(filtered) == 0:
240
+ return False
241
+ if len(filtered) == len(valid):
242
+ return True
243
+ return filtered
244
+
245
+
246
+ def _default_config(
247
+ rows: list[str] | None = None,
248
+ columns: list[str] | None = None,
249
+ values: list[str] | None = None,
250
+ synthetic_measures: list[dict[str, Any]] | None = None,
251
+ aggregation: str | dict[str, str] = "sum",
252
+ show_totals: bool = True,
253
+ show_row_totals: bool | list[str] | None = None,
254
+ show_column_totals: bool | list[str] | None = None,
255
+ empty_cell_value: str = "-",
256
+ interactive: bool = True,
257
+ row_sort: SortConfig | None = None,
258
+ col_sort: SortConfig | None = None,
259
+ sticky_headers: bool = True,
260
+ show_subtotals: bool | list[str] = False,
261
+ repeat_row_labels: bool = False,
262
+ show_values_as: dict[str, str] | None = None,
263
+ conditional_formatting: list[dict[str, Any]] | None = None,
264
+ number_format: str | dict[str, str] | None = None,
265
+ column_alignment: dict[str, str] | None = None,
266
+ ) -> PivotConfig:
267
+ _rows = rows or []
268
+ _values = values or []
269
+ cfg = PivotConfig(
270
+ version=CONFIG_SCHEMA_VERSION,
271
+ rows=_rows,
272
+ columns=columns or [],
273
+ values=_values,
274
+ synthetic_measures=synthetic_measures or [],
275
+ aggregation=_normalize_aggregation_config(aggregation, _values),
276
+ show_totals=show_totals,
277
+ show_row_totals=show_row_totals if show_row_totals is not None else show_totals,
278
+ show_column_totals=show_column_totals
279
+ if show_column_totals is not None
280
+ else show_totals,
281
+ empty_cell_value=empty_cell_value,
282
+ interactive=interactive,
283
+ )
284
+ if isinstance(cfg["show_row_totals"], list):
285
+ cfg["show_row_totals"] = _validate_list_field(
286
+ cfg["show_row_totals"],
287
+ _values,
288
+ "show_row_totals",
289
+ "values",
290
+ )
291
+ if isinstance(cfg["show_column_totals"], list):
292
+ cfg["show_column_totals"] = _validate_list_field(
293
+ cfg["show_column_totals"],
294
+ _values,
295
+ "show_column_totals",
296
+ "values",
297
+ )
298
+ if row_sort is not None:
299
+ cfg["row_sort"] = row_sort
300
+ if col_sort is not None:
301
+ cfg["col_sort"] = col_sort
302
+ if not sticky_headers:
303
+ cfg["sticky_headers"] = False
304
+ if show_subtotals:
305
+ validated: bool | list[str] = show_subtotals
306
+ if isinstance(validated, list):
307
+ validated = _validate_list_field(
308
+ validated, _rows[:-1], "show_subtotals", "rows"
309
+ )
310
+ cfg["show_subtotals"] = validated
311
+ if repeat_row_labels:
312
+ cfg["repeat_row_labels"] = True
313
+ if show_values_as is not None:
314
+ cfg["show_values_as"] = show_values_as
315
+ if conditional_formatting is not None:
316
+ cfg["conditional_formatting"] = conditional_formatting
317
+ if number_format is not None:
318
+ nf = (
319
+ {"__all__": number_format}
320
+ if isinstance(number_format, str)
321
+ else number_format
322
+ )
323
+ cfg["number_format"] = nf
324
+ if column_alignment is not None:
325
+ cfg["column_alignment"] = column_alignment
326
+ return cfg
327
+
328
+
329
+ # ---------------------------------------------------------------------------
330
+ # Event payload schemas
331
+ # ---------------------------------------------------------------------------
332
+
333
+
334
+ class CellClickPayload(TypedDict):
335
+ """Payload fired by setTriggerValue("cell_click", ...).
336
+
337
+ Canonical schema -- both frontend and Python must agree on this shape.
338
+ """
339
+
340
+ rowKey: list[str]
341
+ colKey: list[str]
342
+ value: float | None
343
+ filters: dict[str, str]
344
+ valueField: str
345
+
346
+
347
+ # ---------------------------------------------------------------------------
348
+ # Return type
349
+ # ---------------------------------------------------------------------------
350
+
351
+
352
+ class PivotTableResult(TypedDict, total=False):
353
+ """Value returned by st_pivot_table() to the caller."""
354
+
355
+ config: PivotConfig
356
+
357
+
358
+ # ---------------------------------------------------------------------------
359
+ # CCv2 component registration
360
+ # ---------------------------------------------------------------------------
361
+
362
+ # Registration key follows the CCv2 packaged-component convention:
363
+ # "<project.name>.<component.name>"
364
+ # where both segments come from the in-package manifest at
365
+ # streamlit_pivot/pyproject.toml:
366
+ # [project] name = "streamlit-pivot" -> project.name
367
+ # [[tool.streamlit.component.components]] name = ... -> component.name
368
+ # See component_manifest_handler.py line 65 for the join logic.
369
+ _component = st.components.v2.component(
370
+ "streamlit-pivot.streamlit_pivot",
371
+ js="index-*.js",
372
+ css="index-*.css",
373
+ html='<div class="react-root"></div>',
374
+ )
375
+
376
+
377
+ def _noop_callback(*_args: Any, **_kwargs: Any) -> None:
378
+ """No-op callback supplied at mount when user omits on_config_change."""
379
+
380
+
381
+ # ---------------------------------------------------------------------------
382
+ # Public API
383
+ # ---------------------------------------------------------------------------
384
+
385
+
386
+ def st_pivot_table(
387
+ data: Any,
388
+ *,
389
+ key: str,
390
+ rows: list[str] | None = None,
391
+ columns: list[str] | None = None,
392
+ values: list[str] | None = None,
393
+ synthetic_measures: list[dict[str, Any]] | None = None,
394
+ aggregation: str | dict[str, str] = "sum",
395
+ show_totals: bool = True,
396
+ show_row_totals: bool | list[str] | None = None,
397
+ show_column_totals: bool | list[str] | None = None,
398
+ empty_cell_value: str = "-",
399
+ interactive: bool = True,
400
+ height: int | None = None,
401
+ max_height: int = 500,
402
+ on_cell_click: Callable[[], None] | None = None,
403
+ on_config_change: Callable[[], None] | None = None,
404
+ # Phase 2 parameters
405
+ null_handling: str | dict[str, str] | None = None,
406
+ hidden_attributes: list[str] | None = None,
407
+ hidden_from_aggregators: list[str] | None = None,
408
+ frozen_columns: list[str] | None = None,
409
+ hidden_from_drag_drop: list[str]
410
+ | None = None, # deprecated alias for frozen_columns
411
+ sorters: dict[str, list[str]] | None = None,
412
+ locked: bool = False,
413
+ menu_limit: int | None = None,
414
+ row_sort: SortConfig | None = None,
415
+ col_sort: SortConfig | None = None,
416
+ # Phase 3 parameters
417
+ sticky_headers: bool = True,
418
+ show_subtotals: bool | list[str] = False,
419
+ repeat_row_labels: bool = False,
420
+ show_values_as: dict[str, str] | None = None,
421
+ conditional_formatting: list[dict[str, Any]] | None = None,
422
+ number_format: str | dict[str, str] | None = None,
423
+ column_alignment: dict[str, str] | None = None,
424
+ # Phase 4 parameters
425
+ enable_drilldown: bool = True,
426
+ export_filename: str | None = None,
427
+ ) -> PivotTableResult:
428
+ """Create a pivot table component.
429
+
430
+ Parameters
431
+ ----------
432
+ data : DataFrame-like
433
+ Source data for the pivot table. Accepts the same data types as
434
+ ``st.dataframe``: Pandas DataFrame/Series, Polars DataFrame/Series,
435
+ NumPy arrays, dicts, lists of dicts, pyarrow Tables, and any object
436
+ supporting the DataFrame Interchange Protocol or ``to_pandas()``.
437
+ key : str
438
+ **Required.** A unique string that identifies this component
439
+ instance. Used for state persistence: user config changes made
440
+ via the frontend are hydrated from ``st.session_state[key]`` on
441
+ rerun. Each pivot table on a page must have a distinct key.
442
+ rows : list[str] or None
443
+ Column names from *data* to use as row dimensions.
444
+ columns : list[str] or None
445
+ Column names from *data* to use as column dimensions.
446
+ values : list[str] or None
447
+ Column names from *data* to aggregate as measures.
448
+ aggregation : str or dict[str, str]
449
+ Aggregation setting for raw value fields. A single string applies to all
450
+ measures, while a dict maps each value field to its aggregation.
451
+ show_totals : bool
452
+ Whether to display grand total rows and columns. Acts as default
453
+ for ``show_row_totals`` and ``show_column_totals`` when unset.
454
+ show_row_totals : bool, list[str], or None
455
+ Show row totals column. ``True`` = all measures, ``False`` = none,
456
+ ``["Revenue"]`` = only listed measures (others show ``–``).
457
+ Defaults to ``show_totals`` when None.
458
+ show_column_totals : bool, list[str], or None
459
+ Show column totals row. Same semantics as ``show_row_totals``.
460
+ Defaults to ``show_totals`` when None.
461
+ empty_cell_value : str
462
+ Display string for cells with no data.
463
+ interactive : bool
464
+ If True, the user can reconfigure the pivot via toolbar controls and
465
+ header-menu actions. If False, the toolbar is hidden and header-menu
466
+ sort/filter/show-values-as actions are disabled.
467
+ height : int or None
468
+ Fixed height in pixels. None means auto-size (capped by ``max_height``).
469
+ max_height : int
470
+ Maximum height in pixels when ``height`` is None. The table becomes
471
+ scrollable with sticky headers once content exceeds this value.
472
+ Ignored when ``height`` is explicitly set. Default 500.
473
+ on_cell_click : callable or None
474
+ Called (with no arguments) when a user clicks a data cell. Read the
475
+ payload from ``st.session_state[key]`` after the callback fires.
476
+ Mapped internally to ``on_cell_click_change`` at mount time.
477
+ on_config_change : callable or None
478
+ Called (with no arguments) when the user changes the pivot config
479
+ via the toolbar. Read the updated config from
480
+ ``st.session_state[key]`` after the callback fires.
481
+ If None, a no-op is supplied at mount to satisfy the CCv2 contract
482
+ (every ``default={}`` key needs a matching ``on_<key>_change``).
483
+ null_handling : str or dict[str, str] or None
484
+ How to treat null/NaN values. Global mode ("exclude", "zero",
485
+ "separate") or per-field dict mapping column names to modes.
486
+ Defaults to None ("exclude").
487
+ hidden_attributes : list[str] or None
488
+ Column names to hide entirely from the UI.
489
+ hidden_from_aggregators : list[str] or None
490
+ Column names hidden from the values/aggregators dropdown only.
491
+ frozen_columns : list[str] or None
492
+ Column names that cannot be removed from their toolbar zone.
493
+ hidden_from_drag_drop : list[str] or None
494
+ Deprecated alias for ``frozen_columns``.
495
+ sorters : dict[str, list[str]] or None
496
+ Custom sort orderings per dimension. Maps column name to a list
497
+ of values in the desired order.
498
+ locked : bool
499
+ If True, toolbar config controls are disabled. The settings gear stays
500
+ visible so users can inspect current view status and expand/collapse
501
+ groups. Data export plus header-menu sorting, filtering, and show-values-as
502
+ remain available. Defaults to False.
503
+ menu_limit : int or None
504
+ Max items to show in the header-menu filter checklist. Defaults
505
+ to 50 when None.
506
+ row_sort : dict or None
507
+ Initial sort configuration for rows. Dict with keys ``by``
508
+ ("key" or "value"), ``direction`` ("asc" or "desc"), and
509
+ optionally ``value_field`` (str), ``col_key`` (list[str]), and
510
+ ``dimension`` (str). When ``dimension`` is set and subtotals
511
+ are enabled, only the targeted level and below sort — parent
512
+ groups keep their existing order (scoped sorting).
513
+ col_sort : dict or None
514
+ Initial sort configuration for columns. Same shape as row_sort
515
+ (without ``col_key``).
516
+ sticky_headers : bool
517
+ If True (default), column headers stick to the top of the
518
+ scrolling container so they remain visible when scrolling
519
+ down through large tables.
520
+ show_subtotals : bool or list[str]
521
+ ``True`` = subtotal rows at every parent level, ``False`` = none,
522
+ ``["Region"]`` = only Region-level subtotals. Requires 2+ row
523
+ dimensions. Defaults to False.
524
+ repeat_row_labels : bool
525
+ If True, row dimension labels are repeated on every row
526
+ instead of being merged/spanned. Defaults to False.
527
+ show_values_as : dict[str, str] or None
528
+ Per-field display mode. Maps value field names to one of
529
+ ``"raw"``, ``"pct_of_total"``, ``"pct_of_row"``, or
530
+ ``"pct_of_col"``. Omitted fields default to ``"raw"``.
531
+ conditional_formatting : list[dict] or None
532
+ List of conditional formatting rules applied to value cells.
533
+ Each rule is a dict with ``"type"`` (``"color_scale"``,
534
+ ``"data_bars"``, or ``"threshold"``), ``"apply_to"`` (list of
535
+ field names, empty = all), and type-specific keys.
536
+ number_format : str or dict[str, str] or None
537
+ Number format pattern(s). A single string applies to all
538
+ value fields; a dict maps field names to patterns. Use
539
+ ``"__all__"`` as a key for a default pattern. Patterns
540
+ follow a lightweight d3-style syntax, e.g. ``"$,.0f"``
541
+ for currency integers, ``",.2f"`` for grouped 2-decimal.
542
+ column_alignment : dict[str, str] or None
543
+ Per-field text alignment override. Maps value field names
544
+ to ``"left"``, ``"center"``, or ``"right"``.
545
+ enable_drilldown : bool
546
+ If True (default), clicking a data cell opens an inline
547
+ drill-down panel below the pivot table showing the
548
+ contributing source records. Set to False to disable.
549
+ export_filename : str or None
550
+ Base filename (without extension) used when exporting data.
551
+ The date and file extension are appended automatically.
552
+ Defaults to ``"pivot-table"`` when not set.
553
+
554
+ Returns
555
+ -------
556
+ PivotTableResult
557
+ A dict containing the current ``config`` state.
558
+ """
559
+ if not isinstance(key, str) or not key:
560
+ raise TypeError(
561
+ "key is required: pass a unique string that identifies this "
562
+ "pivot table instance (e.g. key='my_pivot')"
563
+ )
564
+
565
+ # --- Convert data to pandas DataFrame ---
566
+ try:
567
+ data = convert_anything_to_pandas_df(data)
568
+ except (ValueError, TypeError) as exc:
569
+ raise TypeError(
570
+ f"data must be a DataFrame-like object (Pandas/Polars DataFrame, "
571
+ f"dict, list, NumPy array, etc.), got {type(data).__name__}: {exc}"
572
+ ) from exc
573
+ if not isinstance(show_totals, bool):
574
+ raise TypeError(f"show_totals must be a bool, got {type(show_totals).__name__}")
575
+ if show_row_totals is not None and not isinstance(show_row_totals, (bool, list)):
576
+ raise TypeError(
577
+ f"show_row_totals must be bool, list[str], or None, got {type(show_row_totals).__name__}"
578
+ )
579
+ if isinstance(show_row_totals, list) and not all(
580
+ isinstance(s, str) for s in show_row_totals
581
+ ):
582
+ raise TypeError("show_row_totals list items must be strings")
583
+ if show_column_totals is not None and not isinstance(
584
+ show_column_totals, (bool, list)
585
+ ):
586
+ raise TypeError(
587
+ f"show_column_totals must be bool, list[str], or None, got {type(show_column_totals).__name__}"
588
+ )
589
+ if isinstance(show_column_totals, list) and not all(
590
+ isinstance(s, str) for s in show_column_totals
591
+ ):
592
+ raise TypeError("show_column_totals list items must be strings")
593
+ if not isinstance(empty_cell_value, str):
594
+ raise TypeError(
595
+ f"empty_cell_value must be a str, got {type(empty_cell_value).__name__}"
596
+ )
597
+ if not isinstance(interactive, bool):
598
+ raise TypeError(f"interactive must be a bool, got {type(interactive).__name__}")
599
+
600
+ # --- Phase 2 type validation ---
601
+ if null_handling is not None:
602
+ if isinstance(null_handling, str):
603
+ if null_handling not in VALID_NULL_MODES:
604
+ raise ValueError(
605
+ f"null_handling must be one of {sorted(VALID_NULL_MODES)}, got {null_handling!r}"
606
+ )
607
+ elif isinstance(null_handling, dict):
608
+ for k, v in null_handling.items():
609
+ if not isinstance(k, str) or not isinstance(v, str):
610
+ raise TypeError(
611
+ "null_handling dict keys and values must be strings"
612
+ )
613
+ if v not in VALID_NULL_MODES:
614
+ raise ValueError(
615
+ f"null_handling[{k!r}] must be one of {sorted(VALID_NULL_MODES)}, got {v!r}"
616
+ )
617
+ else:
618
+ raise TypeError(
619
+ f"null_handling must be str, dict, or None, got {type(null_handling).__name__}"
620
+ )
621
+
622
+ for param_name, hidden_list in [
623
+ ("hidden_attributes", hidden_attributes),
624
+ ("hidden_from_aggregators", hidden_from_aggregators),
625
+ ("frozen_columns", frozen_columns or hidden_from_drag_drop),
626
+ ]:
627
+ if hidden_list is not None:
628
+ if not isinstance(hidden_list, list) or not all(
629
+ isinstance(c, str) for c in hidden_list
630
+ ):
631
+ raise TypeError(f"{param_name} must be a list of strings")
632
+
633
+ if sorters is not None:
634
+ if not isinstance(sorters, dict):
635
+ raise TypeError(f"sorters must be a dict, got {type(sorters).__name__}")
636
+ for sorter_key, sorter_values in sorters.items():
637
+ if not isinstance(sorter_key, str):
638
+ raise TypeError(
639
+ f"sorters keys must be strings, got {type(sorter_key).__name__}"
640
+ )
641
+ if not isinstance(sorter_values, list) or not all(
642
+ isinstance(s, str) for s in sorter_values
643
+ ):
644
+ raise TypeError(f"sorters[{sorter_key!r}] must be a list of strings")
645
+
646
+ if not isinstance(locked, bool):
647
+ raise TypeError(f"locked must be a bool, got {type(locked).__name__}")
648
+
649
+ # --- Sort config validation ---
650
+ _VALID_SORT_BY = frozenset(("key", "value"))
651
+ _VALID_SORT_DIR = frozenset(("asc", "desc"))
652
+ for param_name, sort_cfg in [("row_sort", row_sort), ("col_sort", col_sort)]:
653
+ if sort_cfg is not None:
654
+ if not isinstance(sort_cfg, dict):
655
+ raise TypeError(
656
+ f"{param_name} must be a dict or None, got {type(sort_cfg).__name__}"
657
+ )
658
+ if sort_cfg.get("by") not in _VALID_SORT_BY:
659
+ raise ValueError(f"{param_name}['by'] must be 'key' or 'value'")
660
+ if sort_cfg.get("direction") not in _VALID_SORT_DIR:
661
+ raise ValueError(f"{param_name}['direction'] must be 'asc' or 'desc'")
662
+ if sort_cfg.get("by") == "value":
663
+ vf = sort_cfg.get("value_field")
664
+ if vf is not None and not isinstance(vf, str):
665
+ raise TypeError(f"{param_name}['value_field'] must be a string")
666
+ ck = sort_cfg.get("col_key")
667
+ if ck is not None:
668
+ if not isinstance(ck, list) or not all(isinstance(s, str) for s in ck):
669
+ raise TypeError(
670
+ f"{param_name}['col_key'] must be a list of strings"
671
+ )
672
+
673
+ # --- Column list type + membership validation ---
674
+ df_cols = set(data.columns)
675
+ for param_name, col_list in [
676
+ ("rows", rows),
677
+ ("columns", columns),
678
+ ("values", values),
679
+ ]:
680
+ if col_list is not None:
681
+ if not isinstance(col_list, list) or not all(
682
+ isinstance(c, str) for c in col_list
683
+ ):
684
+ raise TypeError(
685
+ f"{param_name} must be a list of strings, got {type(col_list).__name__}"
686
+ )
687
+ missing = [c for c in col_list if c not in df_cols]
688
+ if missing:
689
+ raise ValueError(
690
+ f"{param_name} contains columns not in DataFrame: {missing}. "
691
+ f"Available columns: {sorted(df_cols)}"
692
+ )
693
+
694
+ # --- Auto-detect dimensions/measures when not specified ---
695
+ resolved_rows = rows
696
+ resolved_columns = columns
697
+ resolved_values = values
698
+
699
+ if resolved_rows is None and resolved_columns is None and resolved_values is None:
700
+ numeric_cols = data.select_dtypes(include="number").columns.tolist()
701
+ categorical_cols = [c for c in data.columns if c not in numeric_cols]
702
+ # Heuristic: numeric columns with few unique values (<=20) likely
703
+ # represent dimensions (e.g. Year) rather than measures.
704
+ likely_measures = [c for c in numeric_cols if data[c].nunique() > 20]
705
+ likely_numeric_dims = [c for c in numeric_cols if data[c].nunique() <= 20]
706
+ # Treat low-cardinality numerics as dimensions alongside categoricals
707
+ all_dims = categorical_cols + likely_numeric_dims
708
+ resolved_rows = all_dims[:1] if all_dims else []
709
+ resolved_columns = all_dims[1:2] if len(all_dims) > 1 else []
710
+ resolved_values = likely_measures[:2] if likely_measures else numeric_cols[:1]
711
+
712
+ normalized_synthetic_measures: list[dict[str, Any]] = []
713
+ if synthetic_measures is not None:
714
+ if not isinstance(synthetic_measures, list):
715
+ raise TypeError("synthetic_measures must be a list of dicts")
716
+ seen_ids: set[str] = set()
717
+ seen_labels: set[str] = set()
718
+ valid_ops = {"sum_over_sum", "difference"}
719
+ for i, item in enumerate(synthetic_measures):
720
+ if not isinstance(item, dict):
721
+ raise TypeError(f"synthetic_measures[{i}] must be a dict")
722
+ sid = item.get("id")
723
+ label = item.get("label")
724
+ op = item.get("operation")
725
+ numerator = item.get("numerator")
726
+ denominator = item.get("denominator")
727
+ if not isinstance(sid, str) or sid == "":
728
+ raise ValueError(
729
+ f"synthetic_measures[{i}]['id'] must be a non-empty string"
730
+ )
731
+ if not isinstance(label, str) or label == "":
732
+ raise ValueError(
733
+ f"synthetic_measures[{i}]['label'] must be a non-empty string"
734
+ )
735
+ if sid in seen_ids:
736
+ raise ValueError(f"duplicate synthetic_measures id: {sid!r}")
737
+ if label in seen_labels:
738
+ raise ValueError(f"duplicate synthetic_measures label: {label!r}")
739
+ seen_ids.add(sid)
740
+ seen_labels.add(label)
741
+ if op not in valid_ops:
742
+ raise ValueError(
743
+ f"synthetic_measures[{i}]['operation'] must be one of {sorted(valid_ops)}"
744
+ )
745
+ if not isinstance(numerator, str) or numerator not in df_cols:
746
+ raise ValueError(
747
+ f"synthetic_measures[{i}]['numerator'] must be a DataFrame column name"
748
+ )
749
+ if not isinstance(denominator, str) or denominator not in df_cols:
750
+ raise ValueError(
751
+ f"synthetic_measures[{i}]['denominator'] must be a DataFrame column name"
752
+ )
753
+ normalized_synthetic_measures.append(
754
+ {
755
+ "id": sid,
756
+ "label": label,
757
+ "operation": op,
758
+ "numerator": numerator,
759
+ "denominator": denominator,
760
+ "format": item.get("format"),
761
+ }
762
+ )
763
+
764
+ normalized_aggregation = _normalize_aggregation_config(
765
+ aggregation, resolved_values or []
766
+ )
767
+
768
+ # --- Phase 3 validation ---
769
+ if not isinstance(show_subtotals, (bool, list)):
770
+ raise TypeError(
771
+ f"show_subtotals must be bool or list[str], got {type(show_subtotals).__name__}"
772
+ )
773
+ if isinstance(show_subtotals, list) and not all(
774
+ isinstance(s, str) for s in show_subtotals
775
+ ):
776
+ raise TypeError("show_subtotals list items must be strings")
777
+ if not isinstance(repeat_row_labels, bool):
778
+ raise TypeError(
779
+ f"repeat_row_labels must be a bool, got {type(repeat_row_labels).__name__}"
780
+ )
781
+
782
+ if show_values_as is not None:
783
+ if not isinstance(show_values_as, dict):
784
+ raise TypeError(
785
+ f"show_values_as must be a dict or None, got {type(show_values_as).__name__}"
786
+ )
787
+ for k, v in show_values_as.items():
788
+ if not isinstance(k, str) or not isinstance(v, str):
789
+ raise TypeError("show_values_as keys and values must be strings")
790
+ if v not in VALID_SHOW_VALUES_AS:
791
+ raise ValueError(
792
+ f"show_values_as[{k!r}] must be one of {sorted(VALID_SHOW_VALUES_AS)}, got {v!r}"
793
+ )
794
+
795
+ if conditional_formatting is not None:
796
+ if not isinstance(conditional_formatting, list):
797
+ raise TypeError(
798
+ f"conditional_formatting must be a list or None, got {type(conditional_formatting).__name__}"
799
+ )
800
+ for i, rule in enumerate(conditional_formatting):
801
+ if not isinstance(rule, dict):
802
+ raise TypeError(f"conditional_formatting[{i}] must be a dict")
803
+ rtype = rule.get("type")
804
+ if rtype not in VALID_COND_FMT_TYPES:
805
+ raise ValueError(
806
+ f"conditional_formatting[{i}]['type'] must be one of "
807
+ f"{sorted(VALID_COND_FMT_TYPES)}, got {rtype!r}"
808
+ )
809
+ apply_to = rule.get("apply_to", [])
810
+ if not isinstance(apply_to, list) or not all(
811
+ isinstance(a, str) for a in apply_to
812
+ ):
813
+ raise TypeError(
814
+ f"conditional_formatting[{i}]['apply_to'] must be a list of strings"
815
+ )
816
+ if rtype == "color_scale":
817
+ for color_key in ("min_color", "max_color"):
818
+ if not isinstance(rule.get(color_key, ""), str):
819
+ raise TypeError(
820
+ f"conditional_formatting[{i}][{color_key!r}] must be a string"
821
+ )
822
+ if not rule.get("min_color") or not rule.get("max_color"):
823
+ raise ValueError(
824
+ f"conditional_formatting[{i}]: color_scale requires 'min_color' and 'max_color'"
825
+ )
826
+ elif rtype == "threshold":
827
+ conditions = rule.get("conditions")
828
+ if not isinstance(conditions, list) or len(conditions) == 0:
829
+ raise ValueError(
830
+ f"conditional_formatting[{i}]: threshold requires non-empty 'conditions' list"
831
+ )
832
+ valid_ops = {"gt", "gte", "lt", "lte", "eq", "between"}
833
+ for j, cond in enumerate(conditions):
834
+ if not isinstance(cond, dict):
835
+ raise TypeError(
836
+ f"conditional_formatting[{i}]['conditions'][{j}] must be a dict"
837
+ )
838
+ op = cond.get("operator")
839
+ if op not in valid_ops:
840
+ raise ValueError(
841
+ f"conditional_formatting[{i}]['conditions'][{j}]['operator'] "
842
+ f"must be one of {sorted(valid_ops)}, got {op!r}"
843
+ )
844
+ if "value" not in cond:
845
+ raise ValueError(
846
+ f"conditional_formatting[{i}]['conditions'][{j}] requires 'value'"
847
+ )
848
+
849
+ if number_format is not None:
850
+ if isinstance(number_format, str):
851
+ pass # global format string
852
+ elif isinstance(number_format, dict):
853
+ for k, v in number_format.items():
854
+ if not isinstance(k, str) or not isinstance(v, str):
855
+ raise TypeError("number_format keys and values must be strings")
856
+ else:
857
+ raise TypeError(
858
+ f"number_format must be str, dict, or None, got {type(number_format).__name__}"
859
+ )
860
+
861
+ if column_alignment is not None:
862
+ if not isinstance(column_alignment, dict):
863
+ raise TypeError(
864
+ f"column_alignment must be a dict or None, got {type(column_alignment).__name__}"
865
+ )
866
+ for k, v in column_alignment.items():
867
+ if not isinstance(k, str) or v not in VALID_ALIGNMENTS:
868
+ raise ValueError(
869
+ f"column_alignment[{k!r}] must be one of {sorted(VALID_ALIGNMENTS)}, got {v!r}"
870
+ )
871
+
872
+ initial_config = _default_config(
873
+ rows=resolved_rows,
874
+ columns=resolved_columns,
875
+ values=resolved_values,
876
+ synthetic_measures=normalized_synthetic_measures,
877
+ aggregation=normalized_aggregation,
878
+ show_totals=show_totals,
879
+ show_row_totals=show_row_totals,
880
+ show_column_totals=show_column_totals,
881
+ empty_cell_value=empty_cell_value,
882
+ interactive=interactive,
883
+ row_sort=row_sort,
884
+ col_sort=col_sort,
885
+ sticky_headers=sticky_headers,
886
+ show_subtotals=show_subtotals,
887
+ repeat_row_labels=repeat_row_labels,
888
+ show_values_as=show_values_as,
889
+ conditional_formatting=conditional_formatting,
890
+ number_format=number_format,
891
+ column_alignment=column_alignment,
892
+ )
893
+
894
+ # Controlled-state hydration: preserve persisted user config across normal
895
+ # reruns, but let explicit Python config changes take precedence.
896
+ config_to_send = _resolve_config_to_send(st.session_state, key, initial_config)
897
+
898
+ data_payload: dict[str, Any] = {
899
+ "dataframe": data,
900
+ "height": height,
901
+ "max_height": max_height,
902
+ "config": config_to_send,
903
+ }
904
+ if null_handling is not None:
905
+ data_payload["null_handling"] = null_handling
906
+ if hidden_attributes is not None:
907
+ data_payload["hidden_attributes"] = hidden_attributes
908
+ if hidden_from_aggregators is not None:
909
+ data_payload["hidden_from_aggregators"] = hidden_from_aggregators
910
+ _frozen = frozen_columns or hidden_from_drag_drop
911
+ if _frozen is not None:
912
+ data_payload["hidden_from_drag_drop"] = _frozen
913
+ if sorters is not None:
914
+ data_payload["sorters"] = sorters
915
+ if locked:
916
+ data_payload["locked"] = True
917
+ if menu_limit is not None:
918
+ if (
919
+ isinstance(menu_limit, bool)
920
+ or not isinstance(menu_limit, int)
921
+ or menu_limit < 1
922
+ ):
923
+ raise ValueError(
924
+ f"menu_limit must be a positive integer, got {menu_limit!r}"
925
+ )
926
+ data_payload["menu_limit"] = menu_limit
927
+ if not enable_drilldown:
928
+ data_payload["enable_drilldown"] = False
929
+ if export_filename is not None:
930
+ data_payload["export_filename"] = export_filename
931
+
932
+ mount_kwargs: dict[str, Any] = {
933
+ "key": key,
934
+ "default": {"config": config_to_send},
935
+ "data": data_payload,
936
+ "on_config_change": on_config_change or _noop_callback,
937
+ }
938
+
939
+ if on_cell_click is not None:
940
+ mount_kwargs["on_cell_click_change"] = on_cell_click
941
+
942
+ return cast(PivotTableResult, _component(**mount_kwargs))