coeftable 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.
coeftable/frame.py ADDED
@@ -0,0 +1,442 @@
1
+ """Resolve a table specification and a frame into rendered cells."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+ from dataclasses import dataclass, field
7
+ from typing import Any
8
+
9
+ import narwhals as nw
10
+
11
+ from coeftable.format import is_missing, render_interval
12
+ from coeftable.spec import (
13
+ CoefTable,
14
+ Column,
15
+ ColumnNotFoundError,
16
+ Estimate,
17
+ Forest,
18
+ Passthrough,
19
+ SpecError,
20
+ validate_columns,
21
+ )
22
+ from coeftable.svg import forest_axis, forest_bar
23
+ from coeftable.theme import role_for
24
+
25
+ SPLIT_JOINER = "\u2009|\u2009"
26
+
27
+
28
+ @dataclass(frozen=True)
29
+ class Resolved:
30
+ """A specification resolved against a frame, ready to render.
31
+
32
+ Parameters
33
+ ----------
34
+ frame
35
+ Native frame of rendered cell strings, in the caller's own backend.
36
+ display_columns
37
+ Output column names in display order, excluding layout columns.
38
+ labels
39
+ Mapping from output column name to the header text to show.
40
+ spanners
41
+ Mapping from split value to the output columns it spans.
42
+ group_column
43
+ Name of the row-group column, if any.
44
+ band_rows, divider_rows, axis_rows
45
+ Zero-based row indices for banding, dividers and axis rows.
46
+ markdown_columns
47
+ Output columns whose contents are HTML.
48
+ forest_columns
49
+ Output columns rendering `Forest` bars, so the renderer can trim
50
+ their cell padding to let the bar fill the row.
51
+ """
52
+
53
+ frame: Any
54
+ display_columns: list[str] = field(default_factory=list)
55
+ labels: dict[str, str] = field(default_factory=dict)
56
+ spanners: dict[str, list[str]] = field(default_factory=dict)
57
+ group_column: str | None = None
58
+ band_rows: list[int] = field(default_factory=list)
59
+ divider_rows: list[int] = field(default_factory=list)
60
+ axis_rows: list[int] = field(default_factory=list)
61
+ markdown_columns: list[str] = field(default_factory=list)
62
+ forest_columns: list[str] = field(default_factory=list)
63
+
64
+
65
+ def _required_columns(table: CoefTable) -> list[str]:
66
+ names: list[str] = []
67
+ for key in (table.rows, table.nest, table.groups, table.split_columns):
68
+ if key is not None:
69
+ names.append(key)
70
+ for column in table.columns:
71
+ if isinstance(column, Estimate):
72
+ names.append(column.value)
73
+ if column.ci is not None:
74
+ names.extend(column.ci)
75
+ elif isinstance(column, Passthrough):
76
+ names.append(column.column)
77
+ return names
78
+
79
+
80
+ def _check_columns(frame: nw.DataFrame, table: CoefTable) -> None:
81
+ available = list(frame.columns)
82
+ missing = [n for n in _required_columns(table) if n not in available]
83
+ if missing:
84
+ raise ColumnNotFoundError(
85
+ f"Columns {missing} are not in the frame. Available columns: {available}."
86
+ )
87
+
88
+
89
+ def _numeric(frame: nw.DataFrame, name: str) -> list[float | None]:
90
+ values = frame[name].to_list()
91
+ out: list[float | None] = []
92
+ for value in values:
93
+ if value is None:
94
+ out.append(None)
95
+ continue
96
+ try:
97
+ out.append(float(value))
98
+ except (TypeError, ValueError):
99
+ # Some missing-value sentinels (pd.NA, pd.NaT, masked arrays) are
100
+ # not None but should be treated as missing rather than rejected.
101
+ s = str(value)
102
+ if s in ("<NA>", "NaT"):
103
+ out.append(None)
104
+ else:
105
+ raise TypeError(
106
+ f"Column {name!r} must be numeric to be used as an estimate or "
107
+ f"bound; found {value!r}."
108
+ ) from None
109
+ return out
110
+
111
+
112
+ def _ordered_unique(values: list[Any], *, sort: bool) -> list[Any]:
113
+ seen: list[Any] = []
114
+ for value in values:
115
+ if value not in seen:
116
+ seen.append(value)
117
+ return sorted(seen, key=str) if sort else seen
118
+
119
+
120
+ def _first_source(
121
+ source_index: dict[tuple[tuple[Any, Any], Any], int],
122
+ identity: tuple[Any, Any],
123
+ splits: list[Any],
124
+ ) -> int:
125
+ """Return the input row backing `identity`, preferring the first split value.
126
+
127
+ Split-column data is often sparse, so the first split value may have no row
128
+ for a given identity. Falling back to any split keeps layout metadata such
129
+ as the row-group value resolvable.
130
+ """
131
+ for split in splits:
132
+ found = source_index.get((identity, split))
133
+ if found is not None:
134
+ return found
135
+ raise KeyError(f"No input row for {identity!r} under any split value.")
136
+
137
+
138
+ def _finite(values: list[float | None]) -> list[float]:
139
+ return [v for v in values if v is not None and math.isfinite(v)]
140
+
141
+
142
+ def _domain_key(column: Forest, row_key: Any, group: Any, split: Any) -> Any:
143
+ match column.scale:
144
+ case "table":
145
+ return ("table",)
146
+ case "row_group":
147
+ return ("group", group)
148
+ case "split_column":
149
+ return ("split", split)
150
+ case "row":
151
+ return ("row", row_key)
152
+
153
+
154
+ def _pad_domain(
155
+ values: list[float], ref: float, *, symmetric: bool = False
156
+ ) -> tuple[float, float]:
157
+ if not values:
158
+ return (ref - 1.0, ref + 1.0)
159
+ low, high = min(values), max(values)
160
+ low, high = min(low, ref), max(high, ref)
161
+ if low == high:
162
+ return (low - 1.0, high + 1.0)
163
+ margin = (high - low) * 0.08
164
+ low, high = low - margin, high + margin
165
+ if symmetric:
166
+ half = max(ref - low, high - ref)
167
+ return (ref - half, ref + half)
168
+ return (low, high)
169
+
170
+
171
+ # Content height (px) a forest bar needs to fill its row for each CI
172
+ # layout, measured against the theme's default font sizes. Approximate but
173
+ # close enough that the reference line spans the row instead of a short
174
+ # segment centred in a taller cell; `Forest.height` overrides this per column.
175
+ _LAYOUT_HEIGHTS = {"stacked": 48, "inline": 34, "value_only": 34}
176
+
177
+
178
+ def _forest_height(column: Forest, source: Estimate) -> int:
179
+ if column.height is not None:
180
+ return column.height
181
+ return _LAYOUT_HEIGHTS.get(source.ci_style.layout, 18)
182
+
183
+
184
+ def resolve(table: CoefTable) -> Resolved:
185
+ """Resolve `table` against its frame.
186
+
187
+ Parameters
188
+ ----------
189
+ table
190
+ The specification to resolve.
191
+
192
+ Returns
193
+ -------
194
+ Resolved
195
+ Rendered cells plus the layout metadata `render` needs.
196
+
197
+ Raises
198
+ ------
199
+ ColumnNotFoundError
200
+ When a named column is absent from the frame.
201
+ TypeError
202
+ When an estimate or bound column is not numeric.
203
+ SpecError
204
+ When the column specification is inconsistent.
205
+ """
206
+ validate_columns(table.columns)
207
+
208
+ # A column label colliding with a layout key silently overwrites the
209
+ # layout column in the output frame. Catch it here at spec-check time.
210
+ layout_keys = {n for n in (table.rows, table.nest, table.groups) if n is not None}
211
+ for column in table.columns:
212
+ if column.label in layout_keys:
213
+ raise SpecError(
214
+ f"Column label {column.label!r} collides with layout column "
215
+ f"(rows/nest/groups key {column.label!r}); choose a different label."
216
+ )
217
+
218
+ frame = nw.from_native(table.data, eager_only=True)
219
+ _check_columns(frame, table)
220
+
221
+ n = len(frame)
222
+ row_keys = frame[table.rows].to_list() if table.rows else [""] * n
223
+ nest_keys = frame[table.nest].to_list() if table.nest else [None] * n
224
+ group_keys = frame[table.groups].to_list() if table.groups else [None] * n
225
+ split_keys = frame[table.split_columns].to_list() if table.split_columns else [None] * n
226
+
227
+ numeric: dict[str, list[float | None]] = {}
228
+ verbatim: dict[str, list[Any]] = {}
229
+ for column in table.columns:
230
+ if isinstance(column, Estimate):
231
+ numeric[column.value] = _numeric(frame, column.value)
232
+ if column.ci is not None:
233
+ for name in column.ci:
234
+ numeric[name] = _numeric(frame, name)
235
+ elif isinstance(column, Passthrough):
236
+ verbatim[column.column] = frame[column.column].to_list()
237
+
238
+ # Forest domains, keyed by (forest label, domain key).
239
+ domains: dict[tuple[str, Any], tuple[float, float]] = {}
240
+ estimates = {c.label: c for c in table.columns if isinstance(c, Estimate)}
241
+ for column in table.columns:
242
+ if not isinstance(column, Forest):
243
+ continue
244
+ source = estimates[column.of]
245
+ assert source.ci is not None # noqa: S101 - guaranteed by validate_columns
246
+ low_name, high_name = source.ci
247
+ buckets: dict[Any, list[float]] = {}
248
+ for i in range(n):
249
+ key = _domain_key(column, row_keys[i], group_keys[i], split_keys[i])
250
+ bucket = buckets.setdefault(key, [])
251
+ bucket.extend(
252
+ _finite([numeric[source.value][i], numeric[low_name][i], numeric[high_name][i]])
253
+ )
254
+ for key, values in buckets.items():
255
+ domains[(column.label, key)] = column.domain or _pad_domain(
256
+ values, column.ref, symmetric=column.symmetric
257
+ )
258
+
259
+ # Output row identity: one output row per (row key, nest key).
260
+ identities = [(row_keys[i], nest_keys[i]) for i in range(n)]
261
+ unique_rows = _ordered_unique([r for r, _ in identities], sort=table.sort_rows)
262
+ ordered: list[tuple[Any, Any]] = []
263
+ for row_key in unique_rows:
264
+ for identity in identities:
265
+ if identity[0] == row_key and identity not in ordered:
266
+ ordered.append(identity)
267
+
268
+ splits = _ordered_unique(split_keys, sort=table.sort_rows) if table.split_columns else [None]
269
+ source_index: dict[tuple[tuple[Any, Any], Any], int] = {}
270
+ for i in range(n):
271
+ key = (identities[i], split_keys[i])
272
+ if key in source_index:
273
+ row_label, nest_label = identities[i]
274
+ extra = f", split={split_keys[i]!r}" if split_keys[i] is not None else ""
275
+ raise SpecError(
276
+ f"Duplicate input row for row={row_label!r}, nest={nest_label!r}{extra}"
277
+ f" — each (rows, nest, split_columns) combination "
278
+ f"must appear at most once."
279
+ )
280
+ source_index[key] = i
281
+
282
+ def output_name(column: Column, split: Any) -> str:
283
+ return column.label if split is None else f"{split}{SPLIT_JOINER}{column.label}"
284
+
285
+ display_columns: list[str] = []
286
+ labels: dict[str, str] = {}
287
+ spanners: dict[str, list[str]] = {}
288
+ forest_columns: list[str] = []
289
+ for split in splits:
290
+ for column in table.columns:
291
+ name = output_name(column, split)
292
+ display_columns.append(name)
293
+ labels[name] = column.label
294
+ if split is not None:
295
+ spanners.setdefault(str(split), []).append(name)
296
+ if isinstance(column, Forest):
297
+ forest_columns.append(name)
298
+
299
+ cells: dict[str, list[str]] = {name: [] for name in display_columns}
300
+ layout_rows: list[str] = []
301
+ layout_nest: list[str] = []
302
+ layout_group: list[Any] = []
303
+ band_rows: list[int] = []
304
+ divider_rows: list[int] = []
305
+ axis_rows: list[int] = []
306
+ emitted_axis: set[tuple[str, Any]] = set()
307
+
308
+ def blank_row() -> None:
309
+ for name in display_columns:
310
+ cells[name].append("")
311
+
312
+ previous_row_key: Any = None
313
+ for position, (row_key, nest_key) in enumerate(ordered):
314
+ first_of_key = row_key != previous_row_key
315
+ if first_of_key and previous_row_key is not None:
316
+ divider_rows.append(len(layout_rows))
317
+ if unique_rows.index(row_key) % 2 == 0:
318
+ band_rows.append(len(layout_rows))
319
+ layout_rows.append(f"<b>{row_key}</b>" if first_of_key else "")
320
+ layout_nest.append("" if nest_key is None else str(nest_key))
321
+ layout_group.append(group_keys[_first_source(source_index, (row_key, nest_key), splits)])
322
+ previous_row_key = row_key
323
+
324
+ direction = table.direction_for(str(row_key))
325
+ for split in splits:
326
+ index = source_index.get(((row_key, nest_key), split))
327
+ for column in table.columns:
328
+ name = output_name(column, split)
329
+ if index is None:
330
+ cells[name].append("")
331
+ elif isinstance(column, Passthrough):
332
+ cells[name].append(str(verbatim[column.column][index]))
333
+ elif isinstance(column, Estimate):
334
+ low, high = (None, None)
335
+ if column.ci is not None:
336
+ low = numeric[column.ci[0]][index]
337
+ high = numeric[column.ci[1]][index]
338
+ cells[name].append(
339
+ render_interval(
340
+ numeric[column.value][index],
341
+ low,
342
+ high,
343
+ fmt=column.fmt,
344
+ style=column.ci_style,
345
+ theme=table.theme,
346
+ )
347
+ )
348
+ else:
349
+ source = estimates[column.of]
350
+ assert source.ci is not None # noqa: S101
351
+ value = numeric[source.value][index]
352
+ low = numeric[source.ci[0]][index]
353
+ high = numeric[source.ci[1]][index]
354
+ if is_missing(value):
355
+ cells[name].append("")
356
+ continue
357
+ key = _domain_key(column, row_key, layout_group[-1], split)
358
+ domain = domains[(column.label, key)]
359
+ role = (
360
+ table.color_rule(value, low, high, column.ref)
361
+ if table.color_rule is not None
362
+ else role_for(low, high, column.ref, direction)
363
+ )
364
+ cells[name].append(
365
+ forest_bar(
366
+ value,
367
+ low,
368
+ high,
369
+ domain=domain,
370
+ ref=column.ref,
371
+ color=table.theme.color(role),
372
+ theme=table.theme,
373
+ width=column.width,
374
+ height=_forest_height(column, source),
375
+ )
376
+ )
377
+
378
+ # Emit axis rows after the last data row using each domain.
379
+ pending: list[Forest] = []
380
+ for column in table.columns:
381
+ if not isinstance(column, Forest) or not column.show_axis:
382
+ continue
383
+ keys = {_domain_key(column, row_key, layout_group[-1], split) for split in splits}
384
+ if any((column.label, k) in emitted_axis for k in keys):
385
+ continue
386
+ future = any(
387
+ _domain_key(
388
+ column,
389
+ later_row,
390
+ group_keys[_first_source(source_index, (later_row, later_nest), splits)],
391
+ split,
392
+ )
393
+ in keys
394
+ for later_row, later_nest in ordered[position + 1 :]
395
+ for split in splits
396
+ )
397
+ if not future:
398
+ pending.append(column)
399
+ if pending:
400
+ blank_row()
401
+ layout_rows.append("")
402
+ layout_nest.append("")
403
+ layout_group.append(layout_group[-1])
404
+ axis_rows.append(len(layout_rows) - 1)
405
+ for column in pending:
406
+ source = estimates[column.of]
407
+ for split in splits:
408
+ key = _domain_key(column, row_key, layout_group[-1], split)
409
+ emitted_axis.add((column.label, key))
410
+ cells[output_name(column, split)][-1] = forest_axis(
411
+ domain=domains[(column.label, key)],
412
+ ref=column.ref,
413
+ fmt=column.axis_fmt or source.fmt,
414
+ theme=table.theme,
415
+ width=column.width,
416
+ )
417
+
418
+ data: dict[str, list[Any]] = {}
419
+ if table.groups:
420
+ data[table.groups] = layout_group
421
+ if table.rows:
422
+ data[table.rows] = layout_rows
423
+ if table.nest:
424
+ data[table.nest] = layout_nest
425
+ for name in display_columns:
426
+ data[name] = cells[name]
427
+
428
+ leading = [c for c in (table.groups, table.rows, table.nest) if c]
429
+ markdown = [c for c in (table.rows, table.nest) if c] + display_columns
430
+
431
+ return Resolved(
432
+ frame=nw.from_dict(data, backend=nw.get_native_namespace(frame)).to_native(),
433
+ display_columns=[*(leading[1:] if table.groups else leading), *display_columns],
434
+ labels=labels,
435
+ spanners=spanners,
436
+ group_column=table.groups,
437
+ band_rows=band_rows,
438
+ divider_rows=divider_rows,
439
+ axis_rows=axis_rows,
440
+ markdown_columns=markdown,
441
+ forest_columns=forest_columns,
442
+ )
coeftable/render.py ADDED
@@ -0,0 +1,153 @@
1
+ """Turn a resolved specification into a great_tables object."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from great_tables import GT, loc, style
6
+
7
+ from coeftable.frame import resolve
8
+ from coeftable.spec import CoefTable
9
+
10
+
11
+ def to_gt(table: CoefTable) -> GT:
12
+ """Render `table` to a `great_tables` object.
13
+
14
+ Parameters
15
+ ----------
16
+ table
17
+ The specification to render.
18
+
19
+ Returns
20
+ -------
21
+ GT
22
+ A styled table; further native `great_tables` calls may be chained.
23
+ """
24
+ resolved = resolve(table)
25
+ theme = table.theme
26
+ border_color = theme.border_color or theme.header_bg
27
+
28
+ gt = GT(
29
+ resolved.frame,
30
+ groupname_col=resolved.group_column,
31
+ )
32
+
33
+ if table.title:
34
+ gt = gt.tab_header(title=table.title, subtitle=table.subtitle or None)
35
+ gt = gt.tab_style(
36
+ style=[
37
+ style.text(color=theme.header_fg, weight="bold", size="26px", align="left"),
38
+ style.fill(color=theme.header_bg),
39
+ ],
40
+ locations=loc.title(),
41
+ )
42
+ if table.subtitle:
43
+ gt = gt.tab_style(
44
+ style=[
45
+ style.text(color=theme.header_fg, size="16px", align="left"),
46
+ style.fill(color=theme.header_bg),
47
+ ],
48
+ locations=loc.subtitle(),
49
+ )
50
+
51
+ gt = gt.tab_style(
52
+ style=[
53
+ style.text(weight="bold", color=theme.header_fg, align="center", size="16px"),
54
+ style.fill(color=theme.column_label_bg),
55
+ ],
56
+ locations=loc.column_labels(),
57
+ )
58
+
59
+ for split_value, columns in resolved.spanners.items():
60
+ gt = gt.tab_spanner(label=split_value, columns=columns)
61
+
62
+ if resolved.labels:
63
+ gt = gt.cols_label(cases=dict(resolved.labels))
64
+
65
+ gt = gt.fmt_markdown(columns=resolved.markdown_columns).cols_align(align="center")
66
+
67
+ if resolved.band_rows:
68
+ gt = gt.tab_style(
69
+ style=style.fill(color=theme.band),
70
+ locations=loc.body(rows=resolved.band_rows),
71
+ )
72
+ if resolved.divider_rows:
73
+ # Deliberately lighter than border_color: this is a subtle divider
74
+ # between row-key blocks within the light body-row region, not a
75
+ # structural/chrome border. theme.rule is the intended field for
76
+ # "divider on a light background"; border_color is for the darker
77
+ # borders around table/column-label/row-group chrome.
78
+ gt = gt.tab_style(
79
+ style=style.borders(sides="top", color=theme.rule, weight="1px"),
80
+ locations=loc.body(rows=resolved.divider_rows),
81
+ )
82
+ if resolved.axis_rows:
83
+ gt = gt.tab_style(
84
+ style=[
85
+ style.fill(color=theme.surface),
86
+ style.borders(sides="top", color=theme.rule, weight="1px"),
87
+ style.borders(sides="bottom", color=theme.surface, weight="0px"),
88
+ ],
89
+ locations=loc.body(rows=resolved.axis_rows),
90
+ )
91
+ if resolved.group_column:
92
+ gt = gt.tab_style(
93
+ style=[
94
+ style.text(
95
+ weight="bold", color=theme.header_fg, size="16px", transform="uppercase"
96
+ ),
97
+ style.fill(color=theme.column_label_bg),
98
+ style.css("letter-spacing: 0.8px;"),
99
+ style.css(
100
+ f"border-top: 1px solid {border_color} !important;"
101
+ f"border-bottom: 1px solid {border_color};"
102
+ ),
103
+ ],
104
+ locations=loc.row_groups(),
105
+ )
106
+ if resolved.forest_columns:
107
+ # The bar SVG's own height fills the row's content box; the
108
+ # remaining gap around it is this cell's vertical padding, which
109
+ # is otherwise sized for the taller estimate/CI text next to it.
110
+ # Shrinking it here lets the bar (and its reference line) reach
111
+ # closer to the row's true top/bottom edge.
112
+ gt = gt.tab_style(
113
+ style=style.css("padding-top: 2px; padding-bottom: 2px;"),
114
+ locations=loc.body(columns=resolved.forest_columns),
115
+ )
116
+
117
+ side_border_style = "none" if theme.border_style == "minimal" else "solid"
118
+ return gt.tab_options(
119
+ table_font_size=theme.table_font_size,
120
+ column_labels_font_size="16px",
121
+ data_row_padding="10px",
122
+ column_labels_padding="12px",
123
+ data_row_padding_horizontal="16px",
124
+ column_labels_padding_horizontal="16px",
125
+ heading_border_bottom_color=border_color,
126
+ heading_border_bottom_style="solid",
127
+ heading_border_bottom_width="1px",
128
+ column_labels_border_top_color=border_color,
129
+ column_labels_border_top_style="solid",
130
+ column_labels_border_top_width="1px",
131
+ column_labels_border_bottom_color=border_color,
132
+ column_labels_border_bottom_style="solid",
133
+ column_labels_border_bottom_width="1px",
134
+ row_group_border_top_color=border_color,
135
+ row_group_border_top_style="solid",
136
+ row_group_border_top_width="1px",
137
+ row_group_border_bottom_color=border_color,
138
+ row_group_border_bottom_style="solid",
139
+ row_group_border_bottom_width="1px",
140
+ table_border_top_color=border_color,
141
+ table_border_top_style="solid",
142
+ table_border_bottom_color=border_color,
143
+ table_border_bottom_style="solid",
144
+ table_border_left_color=border_color,
145
+ table_border_left_style=side_border_style,
146
+ table_border_right_color=border_color,
147
+ table_border_right_style=side_border_style,
148
+ table_body_border_top_color=border_color,
149
+ table_body_border_top_style="solid",
150
+ table_body_border_top_width="1px",
151
+ table_body_border_bottom_color=border_color,
152
+ table_body_border_bottom_style="solid",
153
+ )