diff --git a/dayplot/calendar.py b/dayplot/calendar.py index 3aef50b..1d4614d 100644 --- a/dayplot/calendar.py +++ b/dayplot/calendar.py @@ -50,6 +50,7 @@ def __repr__(self): _DEFAULT_VMIN = _DefaultArg(None) _DEFAULT_VMAX = _DefaultArg(None) _DEFAULT_VCENTER = _DefaultArg(None) +_DEFAULT_LEGEND = _DefaultArg(False) _DEFAULT_LEGEND_BINS = _DefaultArg(4) _DEFAULT_LESS_LABEL = _DefaultArg("Less") _DEFAULT_MORE_LABEL = _DefaultArg("More") @@ -262,7 +263,7 @@ def calendar( vmax: Any = _DEFAULT_VMAX, vcenter: Any = _DEFAULT_VCENTER, boxstyle: Union[str, patches.BoxStyle] = "square", - legend: bool = False, + legend: Any = _DEFAULT_LEGEND, legend_bins: Any = _DEFAULT_LEGEND_BINS, legend_labels: Optional[Union[List, Literal["auto"]]] = None, legend_labels_precision: Optional[int] = None, @@ -334,7 +335,8 @@ def calendar( positive values, `vcenter` will default to 0. Providing vcenter overrides this automatic setting. boxstyle: The style of each box. This will be passed to `matplotlib.patches.FancyBboxPatch`. Available values are: "square", "circle", "ellipse", "larrow" - legend: Whether to display a legend for the color scale. + legend: Whether to display a legend for the color scale. When omitted, + providing another legend argument automatically displays the legend. legend_bins: Number of boxes/steps to display in the numeric legend. legend_labels: Labels for the legend boxes. Can be a list of strings or "auto" to generate labels from the data values. For categorical legends, None and @@ -366,6 +368,20 @@ def calendar( their values. For categorical data, the last entry for a date is used. """ _validate_inputs(boxstyle, dates, values) + + if legend is _DEFAULT_LEGEND: + legend = any( + ( + legend_bins is not _DEFAULT_LEGEND_BINS, + legend_labels is not None, + legend_labels_precision is not None, + legend_labels_kws is not None, + legend_kws is not None, + less_label is not _DEFAULT_LESS_LABEL, + more_label is not _DEFAULT_MORE_LABEL, + ) + ) + is_categorical = not _is_numeric_values(values) if is_categorical: @@ -683,7 +699,7 @@ def calendar( if legend_labels is not None: if legend_labels == "auto": - legend_label = round(val, ndigits=legend_labels_precision) + legend_label = str(round(val, ndigits=legend_labels_precision)) else: legend_label = str(cast(List, legend_labels)[i]) diff --git a/docs/tuto/legend.md b/docs/tuto/legend.md index 07f1c9a..1fa194c 100644 --- a/docs/tuto/legend.md +++ b/docs/tuto/legend.md @@ -2,6 +2,9 @@ You can add a very simple legend by using `legend=True`: +The `legend` argument can be omitted when another legend option is provided. +For example, `legend_bins=8` both enables the legend and sets its number of bins. + ```py hl_lines="12" # mkdocs: render import matplotlib.pyplot as plt diff --git a/tests/test_main.py b/tests/test_main.py index 0b5e887..c9c9eb6 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -451,6 +451,53 @@ def test_legend_works( plt.close("all") +@pytest.mark.parametrize( + "legend_kwargs", + [ + {"legend_bins": 2}, + {"legend_labels": "auto"}, + {"legend_labels": "auto", "legend_labels_precision": 1}, + {"legend_labels": "auto", "legend_labels_kws": {"color": "red"}}, + {"less_label": "Low"}, + {"more_label": "High"}, + ], +) +def test_numeric_legend_arguments_enable_legend_when_legend_is_omitted( + legend_kwargs, +): + """Test numeric legend options enable the legend without legend=True.""" + dates = [datetime(2024, 1, 1) + timedelta(days=i) for i in range(7)] + values = [1, 2, 3, 4, 5, 6, 7] + fig, ax = plt.subplots() + + rects = calendar(dates, values, ax=ax, **legend_kwargs) + + assert len(ax.patches) > len(rects) + plt.close("all") + + +@pytest.mark.parametrize( + "legend_kwargs", + [ + {"legend_labels": ["Working", "Resting"]}, + {"legend_labels_kws": {"color": "red"}}, + {"legend_kws": {"ncol": 1}}, + ], +) +def test_categorical_legend_arguments_enable_legend_when_legend_is_omitted( + legend_kwargs, +): + """Test categorical legend options enable the legend without legend=True.""" + dates = [datetime(2024, 1, 1), datetime(2024, 1, 2)] + values = ["work", "rest"] + fig, ax = plt.subplots() + + calendar(dates, values, ax=ax, **legend_kwargs) + + assert ax.get_legend() is not None + plt.close("all") + + @pytest.mark.parametrize("month_grid", [True, False]) @pytest.mark.parametrize( "month_grid_kws", [{}, {"edgecolor": (0, 0, 1, 1), "linestyle": "--"}]