Skip to content

Commit 78f02f8

Browse files
authored
refactor(render_params): RenderParams base class for the 5 *RenderParams (#720)
1 parent 7a1b3a3 commit 78f02f8

2 files changed

Lines changed: 216 additions & 30 deletions

File tree

src/spatialdata_plot/pl/render_params.py

Lines changed: 26 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -231,13 +231,28 @@ class ScalebarParams:
231231
scalebar_kwargs: Mapping[str, Any] = field(default_factory=dict)
232232

233233

234-
@dataclass
235-
class ShapesRenderParams:
234+
@dataclass(kw_only=True)
235+
class RenderParams:
236+
"""Fields shared by every ``*RenderParams``.
237+
238+
These four are the only fields with identical type and default across all five renderers.
239+
``kw_only=True`` keeps field order irrelevant across inheritance (subclasses add required fields
240+
such as ``cmap_params``), and all call sites construct these dataclasses by keyword. Subclasses
241+
carry their renderer-specific fields (including ``cmap_params``, whose type varies per renderer).
242+
"""
243+
244+
element: str
245+
zorder: int = 0
246+
colorbar: bool | str | None = "auto"
247+
colorbar_params: dict[str, object] | None = None
248+
249+
250+
@dataclass(kw_only=True)
251+
class ShapesRenderParams(RenderParams):
236252
"""Shapes render parameters."""
237253

238254
cmap_params: CmapParams
239255
outline_params: OutlineParams
240-
element: str
241256
color: Color | None = None
242257
col_for_color: str | None = None
243258
col_for_outline_color: str | None = None
@@ -249,27 +264,23 @@ class ShapesRenderParams:
249264
scale: float = 1.0
250265
transfunc: Callable[[float], float] | None = None
251266
method: str | None = None
252-
zorder: int = 0
253267
table_name: str | None = None
254268
table_layer: str | None = None
255269
shape: Literal["circle", "hex", "visium_hex", "square"] | None = None
256270
# Fast mode: render each shape as a single dot at its centroid instead of its geometry.
257271
as_points: bool = False
258272
size: float = 1.0 # marker size for as_points (matplotlib scatter ``s``)
259273
ds_reduction: _DsReduction | None = None
260-
colorbar: bool | str | None = "auto"
261-
colorbar_params: dict[str, object] | None = None
262274
# Multi-panel color: when set, this render entry belongs to the panel identified by this
263275
# color key. ``None`` means the entry is shared across every panel (e.g. a background layer).
264276
panel_key: str | None = None
265277

266278

267-
@dataclass
268-
class PointsRenderParams:
279+
@dataclass(kw_only=True)
280+
class PointsRenderParams(RenderParams):
269281
"""Points render parameters."""
270282

271283
cmap_params: CmapParams
272-
element: str
273284
color: Color | None = None
274285
col_for_color: str | None = None
275286
groups: str | list[str] | None = None
@@ -278,42 +289,34 @@ class PointsRenderParams:
278289
size: float = 1.0
279290
transfunc: Callable[[float], float] | None = None
280291
method: str | None = None
281-
zorder: int = 0
282292
table_name: str | None = None
283293
table_layer: str | None = None
284294
ds_reduction: _DsReduction | None = None
285-
colorbar: bool | str | None = "auto"
286-
colorbar_params: dict[str, object] | None = None
287295
density: bool = False
288296
density_how: Literal["linear", "log", "cbrt", "eq_hist"] = "linear"
289297

290298

291-
@dataclass
292-
class ImageRenderParams:
299+
@dataclass(kw_only=True)
300+
class ImageRenderParams(RenderParams):
293301
"""Image render parameters."""
294302

295303
cmap_params: list[CmapParams] | CmapParams
296-
element: str
297304
channel: list[str] | list[int] | int | str | None = None
298305
palette: ListedColormap | list[str] | None = None
299306
alpha: float = 1.0
300307
scale: str | None = None
301-
zorder: int = 0
302-
colorbar: bool | str | None = "auto"
303-
colorbar_params: dict[str, object] | None = None
304308
transfunc: Callable[[np.ndarray], np.ndarray] | list[Callable[[np.ndarray], np.ndarray]] | None = None
305309
grayscale: bool = False
306310
channels_as_legend: bool = False
307311
method: Literal["matplotlib", "datashader"] | None = None
308312
ds_reduction: _ImageDsReduction | None = None
309313

310314

311-
@dataclass
312-
class LabelsRenderParams:
315+
@dataclass(kw_only=True)
316+
class LabelsRenderParams(RenderParams):
313317
"""Labels render parameters."""
314318

315319
cmap_params: CmapParams
316-
element: str
317320
color: Color | None = None
318321
col_for_color: str | None = None
319322
col_for_outline_color: str | None = None
@@ -328,9 +331,6 @@ class LabelsRenderParams:
328331
table_name: str | None = None
329332
table_layer: str | None = None
330333
transfunc: Callable[[float], float] | None = None
331-
zorder: int = 0
332-
colorbar: bool | str | None = "auto"
333-
colorbar_params: dict[str, object] | None = None
334334
# Fast mode: render each label as a single dot at its centroid instead of the mask.
335335
as_points: bool = False
336336
size: float = 1.0 # marker size for as_points (matplotlib scatter ``s``)
@@ -341,11 +341,10 @@ class LabelsRenderParams:
341341
panel_key: str | None = None
342342

343343

344-
@dataclass
345-
class GraphRenderParams:
344+
@dataclass(kw_only=True)
345+
class GraphRenderParams(RenderParams):
346346
"""Graph render parameters."""
347347

348-
element: str
349348
connectivity_obsp_key: str = "spatial_connectivities"
350349
table_name: str | None = None
351350
color: Color | None = None
@@ -363,6 +362,3 @@ class GraphRenderParams:
363362
linestyle: str | Sequence[str] = "solid"
364363
rasterize: bool = True
365364
include_self_loops: bool = False
366-
zorder: int = 0
367-
colorbar: bool | str | None = "auto"
368-
colorbar_params: dict[str, object] | None = None

tests/pl/test_render_params.py

Lines changed: 190 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,190 @@
1+
"""Structural tests for the ``*RenderParams`` dataclass hierarchy.
2+
3+
Locks the ``RenderParams`` base-class refactor: the five renderers share one base holding only the
4+
universal fields, and each subclass must keep exactly the field set (and the universal defaults) it
5+
had before the reparent. A dropped field or a flipped default is invisible to the image-baseline
6+
suite (CI-only), so it is asserted directly here.
7+
"""
8+
9+
import dataclasses
10+
11+
import pytest
12+
13+
from spatialdata_plot.pl.render_params import (
14+
GraphRenderParams,
15+
ImageRenderParams,
16+
LabelsRenderParams,
17+
PointsRenderParams,
18+
RenderParams,
19+
ShapesRenderParams,
20+
)
21+
22+
ALL_SUBCLASSES = [
23+
ShapesRenderParams,
24+
PointsRenderParams,
25+
ImageRenderParams,
26+
LabelsRenderParams,
27+
GraphRenderParams,
28+
]
29+
30+
# The only fields with identical type+default across all five renderers; everything else lives on
31+
# the subclasses. ``element`` is required; the other three carry their canonical defaults.
32+
UNIVERSAL_FIELDS = {"element", "zorder", "colorbar", "colorbar_params"}
33+
34+
# Frozen per-subclass field-name snapshot (captured from the pre-refactor dataclasses). The set must
35+
# stay identical through the reparent; adding/removing a renderer field is a deliberate change that
36+
# updates this map.
37+
EXPECTED_FIELD_NAMES = {
38+
ShapesRenderParams: {
39+
"cmap_params",
40+
"outline_params",
41+
"element",
42+
"color",
43+
"col_for_color",
44+
"col_for_outline_color",
45+
"outline_table_name",
46+
"groups",
47+
"palette",
48+
"outline_alpha",
49+
"fill_alpha",
50+
"scale",
51+
"transfunc",
52+
"method",
53+
"zorder",
54+
"table_name",
55+
"table_layer",
56+
"shape",
57+
"as_points",
58+
"size",
59+
"ds_reduction",
60+
"colorbar",
61+
"colorbar_params",
62+
"panel_key",
63+
},
64+
PointsRenderParams: {
65+
"cmap_params",
66+
"element",
67+
"color",
68+
"col_for_color",
69+
"groups",
70+
"palette",
71+
"alpha",
72+
"size",
73+
"transfunc",
74+
"method",
75+
"zorder",
76+
"table_name",
77+
"table_layer",
78+
"ds_reduction",
79+
"colorbar",
80+
"colorbar_params",
81+
"density",
82+
"density_how",
83+
},
84+
ImageRenderParams: {
85+
"cmap_params",
86+
"element",
87+
"channel",
88+
"palette",
89+
"alpha",
90+
"scale",
91+
"zorder",
92+
"colorbar",
93+
"colorbar_params",
94+
"transfunc",
95+
"grayscale",
96+
"channels_as_legend",
97+
"method",
98+
"ds_reduction",
99+
},
100+
LabelsRenderParams: {
101+
"cmap_params",
102+
"element",
103+
"color",
104+
"col_for_color",
105+
"col_for_outline_color",
106+
"outline_table_name",
107+
"groups",
108+
"contour_px",
109+
"palette",
110+
"outline_alpha",
111+
"outline_color",
112+
"fill_alpha",
113+
"scale",
114+
"table_name",
115+
"table_layer",
116+
"transfunc",
117+
"zorder",
118+
"colorbar",
119+
"colorbar_params",
120+
"as_points",
121+
"size",
122+
"method",
123+
"panel_key",
124+
},
125+
GraphRenderParams: {
126+
"element",
127+
"connectivity_obsp_key",
128+
"table_name",
129+
"color",
130+
"obs_col",
131+
"obsp_key",
132+
"cmap_params",
133+
"palette_map",
134+
"na_color",
135+
"color_source",
136+
"groups",
137+
"group_key",
138+
"edge_width",
139+
"edge_alpha",
140+
"weight_key",
141+
"linestyle",
142+
"rasterize",
143+
"include_self_loops",
144+
"zorder",
145+
"colorbar",
146+
"colorbar_params",
147+
},
148+
}
149+
150+
151+
@pytest.mark.parametrize("cls", ALL_SUBCLASSES)
152+
def test_is_renderparams_subclass(cls):
153+
assert issubclass(cls, RenderParams)
154+
155+
156+
def test_base_holds_only_universal_fields():
157+
assert {f.name for f in dataclasses.fields(RenderParams)} == UNIVERSAL_FIELDS
158+
159+
160+
def test_base_universal_defaults():
161+
defaults = {f.name: f.default for f in dataclasses.fields(RenderParams)}
162+
assert defaults["element"] is dataclasses.MISSING # required
163+
assert defaults["zorder"] == 0
164+
assert defaults["colorbar"] == "auto"
165+
assert defaults["colorbar_params"] is None
166+
167+
168+
@pytest.mark.parametrize("cls", ALL_SUBCLASSES)
169+
def test_field_names_preserved(cls):
170+
assert {f.name for f in dataclasses.fields(cls)} == EXPECTED_FIELD_NAMES[cls]
171+
172+
173+
@pytest.mark.parametrize("cls", ALL_SUBCLASSES)
174+
def test_universal_fields_inherited_with_defaults(cls):
175+
names = {f.name for f in dataclasses.fields(cls)}
176+
assert names >= UNIVERSAL_FIELDS
177+
defaults = {f.name: f.default for f in dataclasses.fields(cls)}
178+
assert defaults["zorder"] == 0
179+
assert defaults["colorbar"] == "auto"
180+
assert defaults["colorbar_params"] is None
181+
182+
183+
def test_keyword_construction_roundtrip():
184+
# All construction in basic.py is keyword-only; kw_only=True must not break it, and the inherited
185+
# universal fields must round-trip through the subclass constructor.
186+
p = GraphRenderParams(element="graph", zorder=3, colorbar=False)
187+
assert p.element == "graph"
188+
assert p.zorder == 3
189+
assert p.colorbar is False
190+
assert p.colorbar_params is None # inherited default

0 commit comments

Comments
 (0)