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/spec.py ADDED
@@ -0,0 +1,437 @@
1
+ """Column specifications and the table builder."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Iterable, Mapping
6
+ from dataclasses import dataclass
7
+ from typing import TYPE_CHECKING, Any, Literal
8
+
9
+ from coeftable.format import CIStyle, Format, Number
10
+ from coeftable.theme import DEFAULT, ColorRule, Direction, Theme
11
+
12
+ if TYPE_CHECKING:
13
+ from great_tables import GT
14
+
15
+ type Scale = Literal["table", "row_group", "split_column", "row"]
16
+
17
+ # Module-level singletons: frozen and shared, so they are safe as argument
18
+ # defaults where ruff B008 forbids a constructor call.
19
+ _DEFAULT_FMT = Number()
20
+ _DEFAULT_CI_STYLE = CIStyle()
21
+
22
+
23
+ class SpecError(ValueError):
24
+ """Raised when a table specification is internally inconsistent."""
25
+
26
+
27
+ class ColumnNotFoundError(KeyError):
28
+ """Raised when a specification names a column absent from the frame."""
29
+
30
+
31
+ @dataclass(frozen=True)
32
+ class Estimate:
33
+ """A column rendering a point estimate and its interval.
34
+
35
+ Parameters
36
+ ----------
37
+ label
38
+ Column header, and the name a `Forest` binds to.
39
+ value
40
+ Frame column holding the point estimate.
41
+ ci
42
+ Frame columns holding the lower and upper bounds, or None.
43
+ fmt
44
+ Callable applied to the estimate and both bounds.
45
+ ci_style
46
+ Assembly options for the rendered cell.
47
+ """
48
+
49
+ label: str
50
+ value: str
51
+ ci: tuple[str, str] | None = None
52
+ fmt: Format = _DEFAULT_FMT
53
+ ci_style: CIStyle = _DEFAULT_CI_STYLE
54
+
55
+
56
+ @dataclass(frozen=True)
57
+ class Forest:
58
+ """A column rendering an inline SVG interval bar.
59
+
60
+ Parameters
61
+ ----------
62
+ label
63
+ Column header.
64
+ of
65
+ Label of the `Estimate` this plot visualises.
66
+ ref
67
+ Reference value for the dashed line and for role resolution.
68
+ scale
69
+ Which set of bars share an x-domain.
70
+ domain
71
+ Explicit domain, overriding `scale`.
72
+ symmetric
73
+ When `domain` is not set, symmetrize the auto-computed domain
74
+ around `ref` instead of fitting tightly to the data.
75
+ width
76
+ Bar width in pixels.
77
+ height
78
+ Bar row height in pixels. `None` (the default) picks a height
79
+ that fills the row based on the bound estimate's `ci_style.layout`
80
+ so the reference line spans the full cell instead of a short
81
+ segment centred in a taller row.
82
+ show_axis
83
+ Emit an axis row for each distinct domain.
84
+ axis_fmt
85
+ Callable labelling axis ticks; defaults to the bound estimate's `fmt`.
86
+ """
87
+
88
+ label: str
89
+ of: str
90
+ ref: float = 0.0
91
+ scale: Scale = "table"
92
+ domain: tuple[float, float] | None = None
93
+ symmetric: bool = False
94
+ width: int = 220
95
+ height: int | None = None
96
+ show_axis: bool = True
97
+ axis_fmt: Format | None = None
98
+
99
+
100
+ @dataclass(frozen=True)
101
+ class Passthrough:
102
+ """A column rendering a frame column verbatim.
103
+
104
+ Parameters
105
+ ----------
106
+ label
107
+ Column header.
108
+ column
109
+ Frame column to display.
110
+ """
111
+
112
+ label: str
113
+ column: str
114
+
115
+
116
+ type Column = Estimate | Forest | Passthrough
117
+
118
+
119
+ def validate_columns(columns: tuple[Column, ...]) -> None:
120
+ """Check a column specification for internal consistency.
121
+
122
+ Parameters
123
+ ----------
124
+ columns
125
+ Declared columns, in display order.
126
+
127
+ Raises
128
+ ------
129
+ SpecError
130
+ When no columns are declared, labels collide, a `Forest` names an
131
+ undeclared estimate, or a `Forest` is bound to a CI-less estimate.
132
+ """
133
+ if not columns:
134
+ raise SpecError("Table has no columns; declare at least one.")
135
+
136
+ seen: set[str] = set()
137
+ for column in columns:
138
+ if column.label in seen:
139
+ raise SpecError(f"Table has duplicate column label {column.label!r}.")
140
+ seen.add(column.label)
141
+
142
+ estimates = {c.label: c for c in columns if isinstance(c, Estimate)}
143
+ for column in columns:
144
+ if not isinstance(column, Forest):
145
+ continue
146
+ target = estimates.get(column.of)
147
+ if target is None:
148
+ raise SpecError(
149
+ f"Forest column {column.label!r} references estimate {column.of!r}, "
150
+ f"which is not declared. Declared estimates: {sorted(estimates)}."
151
+ )
152
+ if target.ci is None:
153
+ raise SpecError(
154
+ f"Forest column {column.label!r} references estimate {column.of!r}, "
155
+ "which has no confidence interval to plot."
156
+ )
157
+
158
+
159
+ class CoefTable:
160
+ """A specification for a summary table over a frame of estimates.
161
+
162
+ Immutable by convention: every chain method returns a new instance.
163
+
164
+ Parameters
165
+ ----------
166
+ data
167
+ Any frame narwhals can read: pandas, polars or pyarrow. A plain dict is
168
+ not accepted; narwhals has no backend to build from.
169
+ rows
170
+ Frame column whose values become the leading row label.
171
+ nest
172
+ Frame column stacked beneath each row key.
173
+ groups
174
+ Frame column driving row-group section headers.
175
+ split_columns
176
+ Frame column whose values repeat the declared columns side by side.
177
+ columns
178
+ Declared columns, in display order.
179
+ estimate, ci
180
+ Sugar declaring a single `Estimate` labelled ``"Estimate"``, prepended
181
+ before any `columns` entries.
182
+ direction
183
+ Which side of a reference counts as favorable, table-wide or per row key.
184
+ color_rule
185
+ Callable overriding role resolution entirely.
186
+ theme
187
+ Colour and typography.
188
+ title, subtitle
189
+ Header text.
190
+ sort_rows
191
+ Sort row keys lexically instead of by first appearance.
192
+ """
193
+
194
+ def __init__(
195
+ self,
196
+ data: Any,
197
+ *,
198
+ rows: str | None = None,
199
+ nest: str | None = None,
200
+ groups: str | None = None,
201
+ split_columns: str | None = None,
202
+ columns: Iterable[Column] = (),
203
+ estimate: str | None = None,
204
+ ci: tuple[str, str] | None = None,
205
+ direction: Direction | Mapping[str, Direction] = "higher_is_better",
206
+ color_rule: ColorRule | None = None,
207
+ theme: Theme = DEFAULT,
208
+ title: str = "",
209
+ subtitle: str = "",
210
+ sort_rows: bool = False,
211
+ ) -> None:
212
+ declared = tuple(columns)
213
+ if estimate is not None:
214
+ declared = (Estimate("Estimate", estimate, ci=ci), *declared)
215
+ self.data = data
216
+ self.rows = rows
217
+ self.nest = nest
218
+ self.groups = groups
219
+ self.split_columns = split_columns
220
+ self.columns = declared
221
+ self.direction = direction
222
+ self.color_rule = color_rule
223
+ self.theme = theme
224
+ self.title = title
225
+ self.subtitle = subtitle
226
+ self.sort_rows = sort_rows
227
+ if declared:
228
+ validate_columns(declared)
229
+
230
+ def _with(self, **changes: Any) -> CoefTable:
231
+ settings: dict[str, Any] = {
232
+ "rows": self.rows,
233
+ "nest": self.nest,
234
+ "groups": self.groups,
235
+ "split_columns": self.split_columns,
236
+ "columns": self.columns,
237
+ "direction": self.direction,
238
+ "color_rule": self.color_rule,
239
+ "theme": self.theme,
240
+ "title": self.title,
241
+ "subtitle": self.subtitle,
242
+ "sort_rows": self.sort_rows,
243
+ }
244
+ settings.update(changes)
245
+ return CoefTable(self.data, **settings)
246
+
247
+ def _add(self, column: Column) -> CoefTable:
248
+ return self._with(columns=(*self.columns, column))
249
+
250
+ def estimate(
251
+ self,
252
+ label: str,
253
+ value: str,
254
+ *,
255
+ ci: tuple[str, str] | None = None,
256
+ fmt: Format = _DEFAULT_FMT,
257
+ ci_style: CIStyle = _DEFAULT_CI_STYLE,
258
+ ) -> CoefTable:
259
+ """Append an estimate column.
260
+
261
+ Parameters
262
+ ----------
263
+ label
264
+ Column header, and the name a `Forest` binds to.
265
+ value
266
+ Frame column holding the point estimate.
267
+ ci
268
+ Frame columns holding the lower and upper bounds.
269
+ fmt
270
+ Callable applied to the estimate and both bounds.
271
+ ci_style
272
+ Assembly options for the rendered cell.
273
+
274
+ Returns
275
+ -------
276
+ CoefTable
277
+ A new table with the column appended.
278
+ """
279
+ return self._add(Estimate(label, value, ci=ci, fmt=fmt, ci_style=ci_style))
280
+
281
+ def forest(
282
+ self,
283
+ label: str,
284
+ *,
285
+ of: str,
286
+ ref: float = 0.0,
287
+ scale: Scale = "table",
288
+ domain: tuple[float, float] | None = None,
289
+ symmetric: bool = False,
290
+ width: int = 220,
291
+ height: int | None = None,
292
+ show_axis: bool = True,
293
+ axis_fmt: Format | None = None,
294
+ ) -> CoefTable:
295
+ """Append a forest-plot column bound to an existing estimate.
296
+
297
+ Parameters
298
+ ----------
299
+ label
300
+ Column header.
301
+ of
302
+ Label of the `Estimate` to visualise.
303
+ ref
304
+ Reference value for the dashed line and role resolution.
305
+ scale
306
+ Which set of bars share an x-domain.
307
+ domain
308
+ Explicit domain, overriding `scale`.
309
+ symmetric
310
+ When `domain` is not set, symmetrize the auto-computed domain
311
+ around `ref` instead of fitting tightly to the data.
312
+ width
313
+ Bar width in pixels.
314
+ height
315
+ Bar row height in pixels. `None` (the default) picks a height
316
+ that fills the row based on the bound estimate's
317
+ `ci_style.layout`.
318
+ show_axis
319
+ Emit an axis row per distinct domain.
320
+ axis_fmt
321
+ Callable labelling axis ticks.
322
+
323
+ Returns
324
+ -------
325
+ CoefTable
326
+ A new table with the column appended.
327
+ """
328
+ return self._add(
329
+ Forest(
330
+ label,
331
+ of=of,
332
+ ref=ref,
333
+ scale=scale,
334
+ domain=domain,
335
+ symmetric=symmetric,
336
+ width=width,
337
+ height=height,
338
+ show_axis=show_axis,
339
+ axis_fmt=axis_fmt,
340
+ )
341
+ )
342
+
343
+ def passthrough(self, label: str, column: str) -> CoefTable:
344
+ """Append a column rendered verbatim from the frame.
345
+
346
+ Parameters
347
+ ----------
348
+ label
349
+ Column header.
350
+ column
351
+ Frame column to display.
352
+
353
+ Returns
354
+ -------
355
+ CoefTable
356
+ A new table with the column appended.
357
+ """
358
+ return self._add(Passthrough(label, column))
359
+
360
+ def header(self, title: str, subtitle: str = "") -> CoefTable:
361
+ """Set the header text.
362
+
363
+ Parameters
364
+ ----------
365
+ title
366
+ Title line.
367
+ subtitle
368
+ Subtitle line.
369
+
370
+ Returns
371
+ -------
372
+ CoefTable
373
+ A new table with the header set.
374
+ """
375
+ return self._with(title=title, subtitle=subtitle)
376
+
377
+ def with_theme(self, theme: Theme) -> CoefTable:
378
+ """Replace the theme.
379
+
380
+ Parameters
381
+ ----------
382
+ theme
383
+ Theme to use.
384
+
385
+ Returns
386
+ -------
387
+ CoefTable
388
+ A new table using `theme`.
389
+ """
390
+ return self._with(theme=theme)
391
+
392
+ def with_direction(self, direction: Direction | Mapping[str, Direction]) -> CoefTable:
393
+ """Replace the direction semantics.
394
+
395
+ Parameters
396
+ ----------
397
+ direction
398
+ Table-wide direction, or a mapping from row key to direction.
399
+
400
+ Returns
401
+ -------
402
+ CoefTable
403
+ A new table using `direction`.
404
+ """
405
+ return self._with(direction=direction)
406
+
407
+ def direction_for(self, row_key: str) -> Direction:
408
+ """Resolve the direction applying to a row key.
409
+
410
+ Parameters
411
+ ----------
412
+ row_key
413
+ Value of the `rows` column.
414
+
415
+ Returns
416
+ -------
417
+ Direction
418
+ The direction for that row, defaulting to ``"higher_is_better"``.
419
+ """
420
+ if isinstance(self.direction, Mapping):
421
+ return self.direction.get(row_key, "higher_is_better")
422
+ return self.direction
423
+
424
+ def gt(self) -> GT:
425
+ """Render to a `great_tables` object.
426
+
427
+ Returns
428
+ -------
429
+ GT
430
+ The rendered table.
431
+ """
432
+ from coeftable.render import to_gt
433
+
434
+ return to_gt(self)
435
+
436
+ def _repr_html_(self) -> str:
437
+ return self.gt()._repr_html_()
coeftable/svg.py ADDED
@@ -0,0 +1,210 @@
1
+ """Inline SVG emitters for forest bars and their shared axis."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+
7
+ from coeftable.format import Format, is_missing
8
+ from coeftable.theme import Theme
9
+
10
+ _TICK_STEPS = (1.0, 2.0, 2.5, 5.0, 10.0)
11
+
12
+
13
+ def nice_ticks(low: float, high: float, target: int = 4) -> list[float]:
14
+ """Return round tick positions spanning ``[low, high]``.
15
+
16
+ Parameters
17
+ ----------
18
+ low, high
19
+ Domain bounds.
20
+ target
21
+ Approximate number of ticks wanted.
22
+
23
+ Returns
24
+ -------
25
+ list of float
26
+ Tick positions, empty when the domain is invalid.
27
+ """
28
+ if not (math.isfinite(low) and math.isfinite(high)) or high < low:
29
+ return []
30
+ if high == low:
31
+ return [low]
32
+ raw = (high - low) / max(target, 1)
33
+ magnitude = 10.0 ** math.floor(math.log10(raw))
34
+ step = next((m * magnitude for m in _TICK_STEPS if raw <= m * magnitude), 10.0 * magnitude)
35
+ start = math.ceil(low / step) * step
36
+ count = math.floor((high - start) / step) + 1
37
+ return [round(start + i * step, 10) for i in range(max(count, 0))]
38
+
39
+
40
+ def _projector(domain: tuple[float, float], width: int, pad: int):
41
+ low, high = domain
42
+ span = high - low
43
+ if span <= 0:
44
+ span = 1.0
45
+ inner = width - 2 * pad
46
+
47
+ def project(value: float) -> float:
48
+ return pad + (value - low) / span * inner
49
+
50
+ return project
51
+
52
+
53
+ def _svg(width: int, height: int, body: str) -> str:
54
+ return (
55
+ f'<svg width="{width}" height="{height}" viewBox="0 0 {width} {height}" '
56
+ f'xmlns="http://www.w3.org/2000/svg" style="display:block;margin:0 auto">'
57
+ f"{body}</svg>"
58
+ )
59
+
60
+
61
+ def forest_bar(
62
+ estimate: float | None,
63
+ lower: float | None,
64
+ upper: float | None,
65
+ *,
66
+ domain: tuple[float, float],
67
+ ref: float,
68
+ color: str,
69
+ theme: Theme,
70
+ width: int = 220,
71
+ height: int = 18,
72
+ bar_height: int = 9,
73
+ pad: int = 3,
74
+ ) -> str:
75
+ """Render one interval as an inline SVG bar.
76
+
77
+ The bar spans the interval, a light tick marks the point estimate, and a
78
+ dashed line marks `ref` when it falls inside `domain`. A bound outside the
79
+ domain, including an unbounded one, draws to the edge with a triangular cap
80
+ so that clipping is visible rather than silently misleading.
81
+
82
+ Parameters
83
+ ----------
84
+ estimate
85
+ Point estimate; the tick is omitted when missing or outside `domain`.
86
+ lower, upper
87
+ Interval bounds. ``None`` means unbounded on that side.
88
+ domain
89
+ Shared x-domain the bar is drawn against.
90
+ ref
91
+ Reference value for the dashed line.
92
+ color
93
+ Bar colour, resolved from a semantic role by the caller.
94
+ theme
95
+ Supplies axis and surface colours.
96
+ width, height, bar_height, pad
97
+ Geometry in pixels.
98
+
99
+ Returns
100
+ -------
101
+ str
102
+ A complete ``<svg>`` element.
103
+ """
104
+ low, high = domain
105
+ project = _projector(domain, width, pad)
106
+ low_value = low if lower is None or is_missing(lower) else lower
107
+ high_value = high if upper is None or is_missing(upper) else upper
108
+ clipped_low = is_missing(lower) or low_value < low
109
+ clipped_high = is_missing(upper) or high_value > high
110
+
111
+ x0 = project(max(low_value, low))
112
+ x1 = project(min(high_value, high))
113
+ top = (height - bar_height) / 2
114
+ middle = height / 2
115
+ parts: list[str] = []
116
+
117
+ if low <= ref <= high:
118
+ ref_x = project(ref)
119
+ parts.append(
120
+ f'<line x1="{ref_x:.2f}" y1="0" x2="{ref_x:.2f}" y2="{height}" '
121
+ f'stroke="{theme.axis}" stroke-width="1" stroke-dasharray="2,2"/>'
122
+ )
123
+
124
+ parts.append(
125
+ f'<rect x="{x0:.2f}" y="{top:.2f}" width="{max(x1 - x0, 0.75):.2f}" '
126
+ f'height="{bar_height}" fill="{color}" fill-opacity="0.75" '
127
+ f'stroke="{color}" stroke-width="0.75"/>'
128
+ )
129
+
130
+ if estimate is not None and not is_missing(estimate) and low <= estimate <= high:
131
+ tick_x = project(estimate)
132
+ parts.append(
133
+ f'<line x1="{tick_x:.2f}" y1="{top:.2f}" x2="{tick_x:.2f}" '
134
+ f'y2="{top + bar_height:.2f}" stroke="{theme.surface}" stroke-width="1.5"/>'
135
+ )
136
+
137
+ cap = bar_height * 0.6
138
+ if clipped_high:
139
+ tip = width - pad / 2
140
+ parts.append(
141
+ f'<polygon points="{tip:.2f},{middle:.2f} {tip - cap:.2f},{middle - cap:.2f} '
142
+ f'{tip - cap:.2f},{middle + cap:.2f}" fill="{color}"/>'
143
+ )
144
+ if clipped_low:
145
+ tip = pad / 2
146
+ parts.append(
147
+ f'<polygon points="{tip:.2f},{middle:.2f} {tip + cap:.2f},{middle - cap:.2f} '
148
+ f'{tip + cap:.2f},{middle + cap:.2f}" fill="{color}"/>'
149
+ )
150
+
151
+ return _svg(width, height, "".join(parts))
152
+
153
+
154
+ def forest_axis(
155
+ *,
156
+ domain: tuple[float, float],
157
+ ref: float,
158
+ fmt: Format,
159
+ theme: Theme,
160
+ width: int = 220,
161
+ height: int = 22,
162
+ pad: int = 3,
163
+ target_ticks: int = 4,
164
+ ) -> str:
165
+ """Render the shared x-axis for a set of forest bars.
166
+
167
+ Parameters
168
+ ----------
169
+ domain
170
+ Shared x-domain.
171
+ ref
172
+ Reference value for the dashed line.
173
+ fmt
174
+ Callable used to label each tick.
175
+ theme
176
+ Supplies the axis colour and label size.
177
+ width, height, pad
178
+ Geometry in pixels.
179
+ target_ticks
180
+ Approximate number of ticks wanted.
181
+
182
+ Returns
183
+ -------
184
+ str
185
+ A complete ``<svg>`` element.
186
+ """
187
+ low, high = domain
188
+ project = _projector(domain, width, pad)
189
+ baseline = 4.0
190
+ parts = [
191
+ f'<line x1="{pad}" y1="{baseline:.2f}" x2="{width - pad}" y2="{baseline:.2f}" '
192
+ f'stroke="{theme.axis}" stroke-width="0.75"/>'
193
+ ]
194
+ if low <= ref <= high:
195
+ ref_x = project(ref)
196
+ parts.append(
197
+ f'<line x1="{ref_x:.2f}" y1="0" x2="{ref_x:.2f}" y2="{baseline:.2f}" '
198
+ f'stroke="{theme.axis}" stroke-width="1" stroke-dasharray="2,2"/>'
199
+ )
200
+ for tick in nice_ticks(low, high, target_ticks):
201
+ tick_x = project(tick)
202
+ parts.append(
203
+ f'<line x1="{tick_x:.2f}" y1="{baseline:.2f}" x2="{tick_x:.2f}" '
204
+ f'y2="{baseline + 3:.2f}" stroke="{theme.axis}" stroke-width="0.75"/>'
205
+ )
206
+ parts.append(
207
+ f'<text x="{tick_x:.2f}" y="{height - 2:.2f}" fill="{theme.axis}" '
208
+ f'font-size="9" text-anchor="middle">{fmt(tick)}</text>'
209
+ )
210
+ return _svg(width, height, "".join(parts))