diff --git a/CHANGELOG.md b/CHANGELOG.md index 0583fec45c..40e51de329 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,7 +4,11 @@ This project adheres to [Semantic Versioning](http://semver.org/). ## Unreleased +### Added +- Add support for converting matplotlib reference lines (`axhline`, `axvline`, and `axline`) as well as axes-coordinate lines into Plotly layout shapes in `mpl_to_plotly`, including support for date axes, with thanks to @robertoffmoura for the contribution! + ### Fixed +- Fix `convert_dash` mapping dotted lines (`linestyle=':'`) to invalid `'circle'` in Plotly, correctly converting them to `'dot'` and preserving custom dash patterns as pixel lists, with thanks to @robertoffmoura for the contribution! - Fix concurrent first access to lazily initialized graph object properties, which could raise `ValueError("Invalid value")` [[#3441](https://github.com/plotly/plotly.py/issues/3441)], with thanks to @hb1915 for the contribution! - Fix `px.sunburst`, `px.treemap` and `px.icicle` listing sectors in a different order on every run when `path` is used with a Polars DataFrame; sectors now follow their order of first appearance for all dataframe backends [[#5765](https://github.com/plotly/plotly.py/issues/5765)], with thanks to @Irahan2 for the contribution! diff --git a/plotly/matplotlylib/mplexporter/exporter.py b/plotly/matplotlylib/mplexporter/exporter.py index bbd17568e9..32cae14af6 100644 --- a/plotly/matplotlylib/mplexporter/exporter.py +++ b/plotly/matplotlylib/mplexporter/exporter.py @@ -84,12 +84,6 @@ def process_transform( Data transformed to match the given coordinate code. Returned only if data is specified """ - if isinstance(transform, transforms.BlendedGenericTransform): - warnings.warn( - "Blended transforms not yet supported. " - "Zoom behavior may not work as expected." - ) - if force_trans is not None: if data is not None: data = (transform - force_trans).transform(data) diff --git a/plotly/matplotlylib/mpltools.py b/plotly/matplotlylib/mpltools.py index ad01b37520..ecc35b8a3d 100644 --- a/plotly/matplotlylib/mpltools.py +++ b/plotly/matplotlylib/mpltools.py @@ -55,34 +55,40 @@ def check_corners(inner_obj, outer_obj): return True +VALID_DASH = {"solid", "dot", "dash", "longdash", "dashdot", "longdashdot"} + + def convert_dash(mpl_dash): """Convert mpl line symbol to plotly line symbol and return symbol.""" + if not mpl_dash: + return "solid" if mpl_dash in DASH_MAP: return DASH_MAP[mpl_dash] - else: - dash_array = mpl_dash.split(",") + if mpl_dash in VALID_DASH: + return mpl_dash + if mpl_dash in ("dashed", "--"): + return "dash" + if mpl_dash in ("dotted", ":"): + return "dot" + + dash_array = mpl_dash.replace("px", "").replace(" ", "").split(",") + cleaned = ",".join(dash_array) + if cleaned in DASH_MAP: + return DASH_MAP[cleaned] - if len(dash_array) < 2: - return "solid" + if len(dash_array) < 2: + return "solid" - # Catch the exception where the off length is zero, in case - # matplotlib 'solid' changes from '10,0' to 'N,0' - if math.isclose(float(dash_array[1]), 0.0): - return "solid" + try: + off = float(dash_array[1]) + except ValueError: + return "solid" - # If we can't find the dash pattern in the map, convert it - # into custom values in px, e.g. '7,5' -> '7px,5px' - dashpx = ",".join([x + "px" for x in dash_array]) + if math.isclose(off, 0.0): + return "solid" - # TODO: rewrite the convert_dash code - # only strings 'solid', 'dashed', etc allowed - if dashpx == "7.4px,3.2px": - dashpx = "dashed" - elif dashpx == "12.8px,3.2px,2.0px,3.2px": - dashpx = "dashdot" - elif dashpx == "2.0px,3.3px": - dashpx = "dotted" - return dashpx + # Convert custom dash pattern into px list (e.g. '7,5' -> '7px,5px') + return ",".join([x + "px" for x in dash_array]) def convert_path(path): @@ -570,10 +576,12 @@ def mpl_dates_to_datestrings(dates, mpl_formatter): DASH_MAP = { "10,0": "solid", "6,6": "dash", - "2,2": "circle", + "2,2": "dot", "4,4,2,4": "dashdot", "none": "solid", "7.4,3.2": "dash", + "12.8,3.2,2.0,3.2": "dashdot", + "2.0,3.3": "dot", } PATH_MAP = { diff --git a/plotly/matplotlylib/renderer.py b/plotly/matplotlylib/renderer.py index 65bbcfabb1..105497a26f 100644 --- a/plotly/matplotlylib/renderer.py +++ b/plotly/matplotlylib/renderer.py @@ -7,12 +7,21 @@ """ +import datetime +import math import warnings +from matplotlib import dates as mdates +from matplotlib import lines as mlines +from matplotlib import transforms import plotly.graph_objs as go from plotly.matplotlylib.mplexporter import Renderer from plotly.matplotlylib import mpltools +# Artist class created by ``Axes.axline``: ``AxLine`` in matplotlib >= 3.8, +# ``_AxLine`` in earlier versions. +_AXLINE_CLASS = getattr(mlines, "AxLine", None) or getattr(mlines, "_AxLine", ()) + def _export_color(color): """Export a matplotlib color for use as a plotly color. @@ -458,9 +467,11 @@ def draw_marked_line(self, **props): marked_line["x"] = self._convert_x_dates(marked_line["x"]) self.plotly_fig.add_trace(marked_line) self.msg += " Heck yeah, I drew that line\n" - elif props["coordinates"] == "axes": + elif props["coordinates"] == "axes" and self._processing_legend: # dealing with legend graphical elements self.msg += " Using native legend\n" + elif self._is_axes_reference_line(props): + self._draw_axes_line(props) else: self.msg += " Line didn't have 'data' coordinates, not drawing\n" warnings.warn( @@ -469,6 +480,152 @@ def draw_marked_line(self, **props): "coordinates!" ) + def _is_axes_reference_line(self, props): + """Check whether a line qualifies as an axes-coordinate reference line + (e.g. axhline, axvline, axline, or a 2-point segment drawn with + ``transform=ax.transAxes``) that can be drawn as a 2-point layout shape.""" + if not props.get("linestyle") or len(props.get("data", [])) != 2: + return False + if props["coordinates"] == "axes" and not self._processing_legend: + return True + if props["coordinates"] == "display" and isinstance( + props["mplobj"].get_transform(), transforms.BlendedGenericTransform + ): + return True + return False + + def _draw_axes_line(self, props): + """Draw an axes-coordinate reference line as a layout shape. + + axhline/axvline span their axes-fraction extent along one axis and sit + at a data value on the other. Segments in axes coordinates keep their + exact endpoints in axes domain coordinates. axline is extended in data + coordinates.""" + if not props.get("linestyle") or len(props.get("data", [])) != 2: + return + ax = self.current_mpl_ax + line = props["mplobj"] + trans = line.get_transform() + + axis_suffix = str(self.axis_ct) if self.axis_ct > 1 else "" + x_axis = "x{0}".format(axis_suffix) + y_axis = "y{0}".format(axis_suffix) + x_domain = "{0} domain".format(x_axis).strip() + y_domain = "{0} domain".format(y_axis).strip() + + if isinstance(trans, transforms.BlendedGenericTransform) and ( + trans == ax.get_yaxis_transform() or trans._x == ax.transAxes + ): + # axhline: x spans the axes domain [xmin, xmax], y is in data coordinates + x_data = line.get_xdata(orig=False) + y_data = line.get_ydata(orig=False) + x0, x1 = float(x_data[0]), float(x_data[1]) + y0, y1 = float(y_data[0]), float(y_data[1]) + xref = x_domain + yref = y_axis + elif isinstance(trans, transforms.BlendedGenericTransform) and ( + trans == ax.get_xaxis_transform() or trans._y == ax.transAxes + ): + # axvline: x is in data coordinates, y spans the axes domain [ymin, ymax] + x_data = line.get_xdata(orig=False) + y_data = line.get_ydata(orig=False) + x0, x1 = float(x_data[0]), float(x_data[1]) + y0, y1 = float(y_data[0]), float(y_data[1]) + if self.x_is_mpl_date: + x0, x1 = self._convert_x_dates([x0, x1]) + xref = x_axis + yref = y_domain + elif props["coordinates"] == "axes" and not isinstance(line, _AXLINE_CLASS): + # Segment fixed to the axes (e.g. transform=ax.transAxes): props["data"] + # holds its endpoints as axes fractions, which map directly onto the + # axes domain, so the segment keeps its extent and stays put on pan/zoom + (x0, y0), (x1, y1) = [(float(x), float(y)) for x, y in props["data"]] + xref = x_domain + yref = y_domain + else: + # general reference line (e.g. axline) + if props["coordinates"] == "display": + px_points = props["data"] + elif props["coordinates"] == "axes": + px_points = [ax.transAxes.transform(pt) for pt in props["data"]] + else: + px_points = [trans.transform(pt) for pt in props["data"]] + (x0, y0), (x1, y1) = [ + ax.transData.inverted().transform(pt) for pt in px_points + ] + + dx = x1 - x0 + dy = y1 - y0 + if math.isclose(dy, 0.0, abs_tol=1e-12): + # Horizontal line: use x domain so it spans the chart, y in data coordinates + x0, x1 = 0.0, 1.0 + y0, y1 = float(y0), float(y1) + xref = x_domain + yref = y_axis + elif math.isclose(dx, 0.0, abs_tol=1e-12): + # Vertical line: use y domain so it spans the chart, x in data coordinates + y0, y1 = 0.0, 1.0 + if self.x_is_mpl_date: + x0, x1 = self._convert_x_dates([x0, x1]) + else: + x0, x1 = float(x0), float(x1) + xref = x_axis + yref = y_domain + else: + # Diagonal line: extend endpoints in data coordinates so it spans + # across zoom levels while staying locked to data coordinates on pan/zoom + extension_factor = 100.0 + x0_ext = x0 - extension_factor * dx + y0_ext = y0 - extension_factor * dy + x1_ext = x1 + extension_factor * dx + y1_ext = y1 + extension_factor * dy + + if self.x_is_mpl_date: + min_date_num = float(mdates.date2num(datetime.datetime(1, 1, 1))) + max_date_num = float( + mdates.date2num(datetime.datetime(9999, 12, 31)) + ) + slope = dy / dx + if x0_ext < min_date_num: + y0_ext = y0 + slope * (min_date_num - x0) + x0_ext = min_date_num + elif x0_ext > max_date_num: + y0_ext = y0 + slope * (max_date_num - x0) + x0_ext = max_date_num + if x1_ext > max_date_num: + y1_ext = y1 + slope * (max_date_num - x1) + x1_ext = max_date_num + elif x1_ext < min_date_num: + y1_ext = y1 + slope * (min_date_num - x1) + x1_ext = min_date_num + x0, x1 = self._convert_x_dates([x0_ext, x1_ext]) + else: + x0, x1 = float(x0_ext), float(x1_ext) + y0, y1 = float(y0_ext), float(y1_ext) + xref = x_axis + yref = y_axis + + color = mpltools.merge_color_and_opacity( + props["linestyle"]["color"], props["linestyle"]["alpha"] + ) + shape = go.layout.Shape( + type="line", + x0=x0, + y0=y0, + x1=x1, + y1=y1, + xref=xref, + yref=yref, + line=go.layout.shape.Line( + color=color, + width=props["linestyle"]["linewidth"], + dash=mpltools.convert_dash(props["linestyle"]["dasharray"]), + ), + layer="above", + ) + self.plotly_fig["layout"]["shapes"] += (shape,) + self.msg += " Heck yeah, I drew that reference line\n" + def draw_image(self, **props): """Draw image. diff --git a/plotly/matplotlylib/tests/test_renderer.py b/plotly/matplotlylib/tests/test_renderer.py index 18ce2d02b3..1182b641f1 100644 --- a/plotly/matplotlylib/tests/test_renderer.py +++ b/plotly/matplotlylib/tests/test_renderer.py @@ -1,7 +1,10 @@ import datetime +import pytest +import matplotlib import numpy as np import matplotlib.pyplot as plt +from matplotlib import transforms import plotly.tools as tls @@ -315,6 +318,389 @@ def test_custom_date_xtickvals_are_converted(): ) +def test_axhline_converts(): + """axhline converts to a layout shape spanning the axes width using x domain.""" + fig, ax = plt.subplots() + ax.axhline(0.5) + + plotly_fig = tls.mpl_to_plotly(fig) + + assert len(plotly_fig.data) == 0 + assert len(plotly_fig.layout.shapes) == 1 + shape = plotly_fig.layout.shapes[0] + assert shape.type == "line" + assert shape.x0 == 0 + assert shape.x1 == 1 + assert abs(shape.y0 - 0.5) < 1e-9 + assert abs(shape.y1 - 0.5) < 1e-9 + assert shape.xref == "x domain" + assert shape.yref == "y" + assert shape.line.color == "rgba(31, 119, 180, 1)" + + +def test_axvline_converts(): + """axvline converts to a layout shape spanning the axes height using y domain.""" + fig, ax = plt.subplots() + ax.axvline(0.5) + + plotly_fig = tls.mpl_to_plotly(fig) + + assert len(plotly_fig.data) == 0 + assert len(plotly_fig.layout.shapes) == 1 + shape = plotly_fig.layout.shapes[0] + assert shape.type == "line" + assert abs(shape.x0 - 0.5) < 1e-9 + assert abs(shape.x1 - 0.5) < 1e-9 + assert shape.y0 == 0 + assert shape.y1 == 1 + assert shape.xref == "x" + assert shape.yref == "y domain" + + +def test_axline_converts(): + """axline converts to a layout shape extended in data coordinates.""" + fig, ax = plt.subplots() + ax.axline((0.5, 0.5), slope=1) + + plotly_fig = tls.mpl_to_plotly(fig) + + assert len(plotly_fig.data) == 0 + assert len(plotly_fig.layout.shapes) == 1 + shape = plotly_fig.layout.shapes[0] + assert shape.type == "line" + assert shape.xref == "x" + assert shape.yref == "y" + x_min, x_max = ax.get_xlim() + assert shape.x0 < x_min + assert shape.x1 > x_max + slope = (shape.y1 - shape.y0) / (shape.x1 - shape.x0) + assert abs(slope - 1.0) < 1e-9 + # Passes through (0.5, 0.5) + y_at_05 = shape.y0 + slope * (0.5 - shape.x0) + assert abs(y_at_05 - 0.5) < 1e-9 + + +def test_axline_arbitrary_slope_and_limits(): + """axline with non-trivial slopes and limits extends in data coordinates.""" + fig, ax = plt.subplots() + ax.scatter([1, 2, 4, 7, 9], [2, 5, 4, 8, 10]) + ax.axline((0, 1), slope=1.0) + ax.axline((1, 8), (8, 2)) + ax.set_xlim(0, 10) + ax.set_ylim(0, 12) + + plotly_fig = tls.mpl_to_plotly(fig) + shapes = plotly_fig.layout.shapes + assert len(shapes) == 2 + + # Line 1: (0, 1), slope 1, xlim [0, 10], ylim [0, 12] + assert shapes[0].xref == "x" + assert shapes[0].yref == "y" + assert shapes[0].x0 < -500 + assert shapes[0].x1 > 500 + slope1 = (shapes[0].y1 - shapes[0].y0) / (shapes[0].x1 - shapes[0].x0) + assert abs(slope1 - 1.0) < 1e-9 + y_at_0 = shapes[0].y0 + slope1 * (0.0 - shapes[0].x0) + assert abs(y_at_0 - 1.0) < 1e-9 + + # Line 2: (1, 8) to (8, 2) with slope -6/7 + assert shapes[1].xref == "x" + assert shapes[1].yref == "y" + assert shapes[1].x0 < -500 + assert shapes[1].x1 > 500 + slope2 = (shapes[1].y1 - shapes[1].y0) / (shapes[1].x1 - shapes[1].x0) + assert abs(slope2 - (-6.0 / 7.0)) < 1e-9 + y_at_1 = shapes[1].y0 + slope2 * (1.0 - shapes[1].x0) + assert abs(y_at_1 - 8.0) < 1e-9 + + +def _shape_x_datenums(shape): + """Return a shape's date-string x endpoints as matplotlib date numbers.""" + return [ + matplotlib.dates.date2num(datetime.datetime.fromisoformat(x)) + for x in (shape.x0, shape.x1) + ] + + +def test_axline_on_pre_1970_date_axis(): + """axline on a date axis before the matplotlib epoch (1970) extends on both + sides of the data and stays on the original line.""" + dates = [datetime.datetime(1950, 1, i) for i in range(1, 10)] + fig, ax = plt.subplots() + ax.plot(dates, range(len(dates))) + x_ref = matplotlib.dates.date2num(dates[0]) + ax.axline((x_ref, 0), slope=1) + + plotly_fig = tls.mpl_to_plotly(fig) + shapes = plotly_fig.layout.shapes + assert len(shapes) == 1 + shape = shapes[0] + assert shape.xref == "x" + assert shape.yref == "y" + + n0, n1 = _shape_x_datenums(shape) + x_min, x_max = ax.get_xlim() + assert n0 < x_min + assert n1 > x_max + slope = (shape.y1 - shape.y0) / (n1 - n0) + assert abs(slope - 1.0) < 1e-6 + assert abs(shape.y0 + slope * (x_ref - n0)) < 1e-6 + + +def test_axline_on_date_axis_clamps_to_matplotlib_date_range(): + """When extending a diagonal axline would leave matplotlib's supported date + range (years 0001-9999), its endpoints stop at the range limits while + staying on the original line.""" + dates = [datetime.datetime(1900, 1, 1), datetime.datetime(2000, 1, 1)] + fig, ax = plt.subplots() + ax.plot(dates, [0, 1]) + # Nearly flat line crossing the whole century, so the 100x extension of the + # visible segment reaches past both year 0001 and year 9999 + x_ref = matplotlib.dates.date2num(dates[0]) + ax.axline((x_ref, 0.5), slope=1e-6) + + plotly_fig = tls.mpl_to_plotly(fig) + shape = plotly_fig.layout.shapes[0] + assert isinstance(shape.x0, str) + assert isinstance(shape.x1, str) + assert shape.x0.startswith("0001-01-01") + assert shape.x1.startswith("9999-12-31") + + n0, n1 = _shape_x_datenums(shape) + slope = (shape.y1 - shape.y0) / (n1 - n0) + assert abs(slope - 1e-6) < 1e-12 + assert abs(shape.y0 + slope * (x_ref - n0) - 0.5) < 1e-6 + + +def test_axline_horizontal_and_vertical(): + """Horizontal and vertical axline use domain coordinates appropriately.""" + fig, ax = plt.subplots() + ax.axline((0, 5), slope=0) + ax.axline((3, 0), (3, 10)) + + plotly_fig = tls.mpl_to_plotly(fig) + shapes = plotly_fig.layout.shapes + assert len(shapes) == 2 + + # Horizontal axline: xref is domain, y is data + assert shapes[0].xref == "x domain" + assert shapes[0].yref == "y" + assert abs(shapes[0].x0 - 0.0) < 1e-9 + assert abs(shapes[0].x1 - 1.0) < 1e-9 + assert abs(shapes[0].y0 - 5.0) < 1e-9 + assert abs(shapes[0].y1 - 5.0) < 1e-9 + + # Vertical axline: xref is data, yref is domain + assert shapes[1].xref == "x" + assert shapes[1].yref == "y domain" + assert abs(shapes[1].x0 - 3.0) < 1e-9 + assert abs(shapes[1].x1 - 3.0) < 1e-9 + assert abs(shapes[1].y0 - 0.0) < 1e-9 + assert abs(shapes[1].y1 - 1.0) < 1e-9 + + +def test_axes_coordinate_segments_keep_their_endpoints(): + """Two-point lines in axes coordinates (not axline) convert to shapes in + axes domain coordinates with their exact endpoints, without extension.""" + fig, ax = plt.subplots() + ax.set_xlim(0, 10) + ax.set_ylim(0, 10) + ax.plot([0.2, 0.4], [0.5, 0.5], transform=ax.transAxes) + ax.plot([0.2, 0.4], [0.2, 0.6], transform=ax.transAxes) + + plotly_fig = tls.mpl_to_plotly(fig) + shapes = plotly_fig.layout.shapes + assert len(shapes) == 2 + + expected = [((0.2, 0.5), (0.4, 0.5)), ((0.2, 0.2), (0.4, 0.6))] + for shape, ((x0, y0), (x1, y1)) in zip(shapes, expected): + assert shape.type == "line" + assert shape.xref == "x domain" + assert shape.yref == "y domain" + assert abs(shape.x0 - x0) < 1e-9 + assert abs(shape.y0 - y0) < 1e-9 + assert abs(shape.x1 - x1) < 1e-9 + assert abs(shape.y1 - y1) < 1e-9 + + +def test_axes_coordinate_segment_on_subplot_uses_subplot_domain(): + """Axes-coordinate segments on a second subplot reference that subplot's domain.""" + fig, (ax1, ax2) = plt.subplots(1, 2) + ax2.plot([0.1, 0.3], [0.7, 0.9], transform=ax2.transAxes) + + plotly_fig = tls.mpl_to_plotly(fig) + shapes = plotly_fig.layout.shapes + assert len(shapes) == 1 + assert shapes[0].xref == "x2 domain" + assert shapes[0].yref == "y2 domain" + assert abs(shapes[0].x0 - 0.1) < 1e-9 + assert abs(shapes[0].y0 - 0.7) < 1e-9 + assert abs(shapes[0].x1 - 0.3) < 1e-9 + assert abs(shapes[0].y1 - 0.9) < 1e-9 + + +def test_axvline_and_axhline_on_date_xaxis(): + """axvline and axhline on a date x-axis use proper domain and data coordinate references.""" + dates = [datetime.datetime(2023, 1, i) for i in range(1, 10)] + fig, ax = plt.subplots() + ax.plot(dates, range(len(dates))) + ax.axvline(dates[4]) + ax.axhline(4) + + plotly_fig = tls.mpl_to_plotly(fig) + shapes = plotly_fig.layout.shapes + assert len(shapes) == 2 + + vline_shape = shapes[0] + assert isinstance(vline_shape.x0, str) + assert isinstance(vline_shape.x1, str) + assert vline_shape.x0.startswith("2023-01-05") + assert vline_shape.x1.startswith("2023-01-05") + assert vline_shape.y0 == 0 + assert vline_shape.y1 == 1 + assert vline_shape.xref == "x" + assert vline_shape.yref == "y domain" + + hline_shape = shapes[1] + assert hline_shape.x0 == 0 + assert hline_shape.x1 == 1 + assert hline_shape.y0 == 4 + assert hline_shape.y1 == 4 + assert hline_shape.xref == "x domain" + assert hline_shape.yref == "y" + + +def test_axhline_on_date_yaxis(): + """axhline with a datetime on a date y-axis converts without crashing.""" + dates = [datetime.datetime(2023, 1, i) for i in range(1, 10)] + fig, ax = plt.subplots() + ax.plot(range(len(dates)), dates) + ax.axhline(dates[3]) + + plotly_fig = tls.mpl_to_plotly(fig) + shapes = plotly_fig.layout.shapes + assert len(shapes) == 1 + shape = shapes[0] + assert shape.type == "line" + assert shape.xref == "x domain" + assert shape.yref == "y" + assert abs(shape.x0 - 0.0) < 1e-9 + assert abs(shape.x1 - 1.0) < 1e-9 + expected_y = float(matplotlib.dates.date2num(dates[3])) + assert abs(shape.y0 - expected_y) < 1e-9 + assert abs(shape.y1 - expected_y) < 1e-9 + + +def test_reference_lines_custom_limits_and_subplots(): + """axhline and axvline respect custom domain limits and subplot axis references.""" + fig, (ax1, ax2) = plt.subplots(1, 2) + ax1.axhline(0.5, xmin=0.2, xmax=0.8) + ax2.axvline(3.0, ymin=0.1, ymax=0.9) + + plotly_fig = tls.mpl_to_plotly(fig) + shapes = plotly_fig.layout.shapes + assert len(shapes) == 2 + + # First subplot + assert shapes[0].xref == "x domain" + assert shapes[0].yref == "y" + assert abs(shapes[0].x0 - 0.2) < 1e-9 + assert abs(shapes[0].x1 - 0.8) < 1e-9 + assert abs(shapes[0].y0 - 0.5) < 1e-9 + assert abs(shapes[0].y1 - 0.5) < 1e-9 + + # Second subplot + assert shapes[1].xref == "x2" + assert shapes[1].yref == "y2 domain" + assert abs(shapes[1].x0 - 3.0) < 1e-9 + assert abs(shapes[1].x1 - 3.0) < 1e-9 + assert abs(shapes[1].y0 - 0.1) < 1e-9 + assert abs(shapes[1].y1 - 0.9) < 1e-9 + + +def test_axes_line_with_more_than_two_points_does_not_crash(): + """Axes-coordinate line with != 2 points does not crash and is ignored with a warning.""" + fig, ax = plt.subplots() + ax.plot([0.1, 0.5, 0.9], [0.1, 0.5, 0.9], transform=ax.transAxes) + + with pytest.warns(UserWarning, match="Line2D objects from matplotlib"): + plotly_fig = tls.mpl_to_plotly(fig) + assert len(plotly_fig.layout.shapes) == 0 + + +def test_axes_line_two_points_markers_only_does_not_crash(): + """Axes-coordinate line with exactly two points but markers only (linestyle is None) does not crash.""" + fig, ax = plt.subplots() + # Exactly two points, marker-only (linestyle=None) + lines = ax.plot([0.1, 0.9], [0.1, 0.9], "o", transform=ax.transAxes) + assert len(lines[0].get_xydata()) == 2 + assert lines[0].get_linestyle() == "None" + + with pytest.warns(UserWarning, match="Line2D objects from matplotlib"): + plotly_fig = tls.mpl_to_plotly(fig) + assert len(plotly_fig.layout.shapes) == 0 + + +def test_blended_transform_line_invalid_does_not_crash(): + """Blended transform line with != 2 points or marker only does not crash.""" + fig, ax = plt.subplots() + trans = transforms.blended_transform_factory(ax.transData, ax.transAxes) + line1 = matplotlib.lines.Line2D([0.1, 0.5, 0.9], [0.1, 0.5, 0.9], transform=trans) + line2 = matplotlib.lines.Line2D( + [0.1, 0.9], [0.1, 0.9], linestyle="None", marker="o", transform=trans + ) + ax.add_line(line1) + ax.add_line(line2) + + with pytest.warns(UserWarning, match="Line2D objects from matplotlib"): + plotly_fig = tls.mpl_to_plotly(fig) + assert len(plotly_fig.layout.shapes) == 0 + + +def test_dotted_line_dash_converts_to_dot(): + """Dotted lines (linestyle=':') convert to dash='dot'.""" + fig, ax = plt.subplots() + ax.plot([0, 1], [0, 1], linestyle=":") + ax.axhline(0.5, linestyle=":") + + plotly_fig = tls.mpl_to_plotly(fig) + assert plotly_fig.data[0].line.dash == "dot" + assert plotly_fig.layout.shapes[0].line.dash == "dot" + + +def test_convert_dash_returns_valid_plotly_dash_styles(): + """convert_dash maps standard styles to Plotly names and custom patterns to px lists.""" + import plotly.graph_objs as go + from plotly.matplotlylib.mpltools import convert_dash + + expected_mappings = { + "10,0": "solid", + "6,6": "dash", + "2,2": "dot", + "4,4,2,4": "dashdot", + "none": "solid", + "7.4,3.2": "dash", + "2.0,3.3": "dot", + "12.8,3.2,2.0,3.2": "dashdot", + "5,5": "5px,5px", + "12,3,2,3": "12px,3px,2px,3px", + "dashed": "dash", + "dotted": "dot", + "--": "dash", + ":": "dot", + "": "solid", + None: "solid", + } + for inp, expected in expected_mappings.items(): + res = convert_dash(inp) + assert res == expected, f"Input {inp!r} produced {res!r}, expected {expected!r}" + # Confirm Plotly layout shape line and scatter line accept the converted dash + shape_line = go.layout.shape.Line(dash=res) + scatter_line = go.scatter.Line(dash=res) + assert shape_line.dash == res + assert scatter_line.dash == res + + def test_uneven_custom_date_xtickvals_are_converted(): """Unevenly spaced custom date ticks must be converted to date strings.""" dates = [datetime.datetime(2023, 1, i) for i in range(1, 11)]