mathtext2doc 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.
@@ -0,0 +1,775 @@
1
+ """函数图绘制器。
2
+
3
+ 输入:parser.Plot(一组 PlotItem)
4
+ 输出:PNG 文件路径
5
+
6
+ 支持:
7
+ - 2D 显函数 y = f(x)
8
+ - 2D 隐函数 F(x, y) = 0
9
+ - 多函数同图(显隐混合)
10
+ - 自动分配颜色
11
+ - 标签默认标注在曲线可见部分的几何中点偏上
12
+ - 标签之间自动避让(基于 bbox 重叠检测的迭代位移)
13
+ - 中文 label(matplotlib 使用 Noto Sans SC)
14
+ - 常见初等函数:+ - * / ^、sin cos tan log ln exp sqrt abs pi e
15
+
16
+ 实现说明:
17
+ - 显函数:sympy 解析表达式 → lambdify → numpy linspace 采样
18
+ - 隐函数:sympy 解析 F(x, y) = 0 → 移项为 F(x, y) → numpy meshgrid
19
+ + matplotlib contour(level=0) 绘制零等高线
20
+ - 标签避让:先按曲线可见中点初始化每个标签位置;然后用迭代算法
21
+ 检查所有标签 bbox 是否重叠,重叠则沿曲线方向滑动或上移,最多迭代 N 次。
22
+ """
23
+
24
+ from __future__ import annotations
25
+
26
+ import math
27
+ import os
28
+ from dataclasses import dataclass
29
+ from typing import List, Optional, Tuple
30
+
31
+ import matplotlib
32
+ matplotlib.use("Agg") # 无显示设备
33
+ import matplotlib.font_manager as fm
34
+
35
+ # 注册中文字体(自动探测常见路径)
36
+ _FONT_CANDIDATES = [
37
+ # Linux 常见路径
38
+ "/usr/share/fonts/truetype/chinese/NotoSansSC[wght].ttf",
39
+ "/usr/share/fonts/truetype/chinese/NotoSansSC-Regular.ttf",
40
+ "/usr/share/fonts/truetype/noto/NotoSansCJK-Regular.ttc",
41
+ "/usr/share/fonts/opentype/noto/NotoSansCJK-Regular.ttc",
42
+ "/usr/share/fonts/truetype/wqy/wqy-zenhei.ttc",
43
+ "/usr/share/fonts/truetype/wqy/wqy-microhei.ttc",
44
+ "/usr/share/fonts/truetype/lxgw-wenkai/LXGWWenKai-Regular.ttf",
45
+ # macOS 常见路径
46
+ "/System/Library/Fonts/PingFang.ttc",
47
+ "/Library/Fonts/Songti.ttc",
48
+ # DejaVu Sans 兜底(拉丁 + 符号)
49
+ "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf",
50
+ ]
51
+ for _f in _FONT_CANDIDATES:
52
+ if os.path.exists(_f):
53
+ try:
54
+ fm.fontManager.addfont(_f)
55
+ except Exception:
56
+ pass
57
+
58
+ # 探测可用的中文字体名(按优先级)
59
+ _AVAILABLE_CN_FONTS = []
60
+ for _name in [
61
+ "Noto Sans SC", "Noto Sans CJK SC", "WenQuanYi Zen Hei",
62
+ "WenQuanYi Micro Hei", "LXGW WenKai", "PingFang SC", "Songti SC",
63
+ "Sarasa Mono SC",
64
+ ]:
65
+ try:
66
+ if fm.findfont(_name, fallback_to_default=False) != fm.findfont("DejaVu Sans"):
67
+ _AVAILABLE_CN_FONTS.append(_name)
68
+ except Exception:
69
+ pass
70
+
71
+ import matplotlib.pyplot as plt
72
+ import numpy as np
73
+ import sympy as sp
74
+
75
+ # 优先中文字体,然后 DejaVu Sans 兜底
76
+ plt.rcParams["font.sans-serif"] = (_AVAILABLE_CN_FONTS or ["Noto Sans SC"]) + [
77
+ "DejaVu Sans",
78
+ ]
79
+ plt.rcParams["axes.unicode_minus"] = False
80
+
81
+ from .parser import Plot, PlotItem
82
+
83
+
84
+ # ---------------------------------------------------------------------------
85
+ # 颜色循环
86
+ # ---------------------------------------------------------------------------
87
+
88
+ # 一组对比明显、色盲友好的颜色
89
+ _COLOR_CYCLE = [
90
+ "#1f77b4", # 蓝
91
+ "#d62728", # 红
92
+ "#2ca02c", # 绿
93
+ "#ff7f0e", # 橙
94
+ "#9467bd", # 紫
95
+ "#17becf", # 青
96
+ "#8c564b", # 棕
97
+ "#e377c2", # 粉
98
+ ]
99
+
100
+
101
+ # ---------------------------------------------------------------------------
102
+ # 表达式解析
103
+ # ---------------------------------------------------------------------------
104
+
105
+ # sympy 默认支持 log (=ln)、sin、cos、tan、exp、sqrt、abs、pi、E。
106
+ # 我们额外提供 `ln` 作为 `log` 的别名,`e` 作为 E 的别名。
107
+ _LOCAL_NAMES = {
108
+ "pi": sp.pi,
109
+ "e": sp.E,
110
+ "E": sp.E,
111
+ "ln": sp.log,
112
+ "log": sp.log, # sympy 中 log 默认是自然对数
113
+ "log10": sp.log,
114
+ "sin": sp.sin,
115
+ "cos": sp.cos,
116
+ "tan": sp.tan,
117
+ "exp": sp.exp,
118
+ "sqrt": sp.sqrt,
119
+ "abs": sp.Abs,
120
+ }
121
+
122
+
123
+ def _to_sympy_expr(s: str) -> sp.Expr:
124
+ """把用户表达式转成 sympy 表达式。把 ^ 替换为 **。"""
125
+ s = s.replace("^", "**")
126
+ try:
127
+ return sp.sympify(s, locals=_LOCAL_NAMES)
128
+ except Exception as ex:
129
+ raise PlotRenderError(f"无法解析表达式 {s!r}:{ex}")
130
+
131
+
132
+ # ---------------------------------------------------------------------------
133
+ # 显函数 / 隐函数绘制
134
+ # ---------------------------------------------------------------------------
135
+
136
+ @dataclass
137
+ class _PlottedCurve:
138
+ """已经绘制好的曲线信息,用于标签避让。"""
139
+ label: Optional[str]
140
+ color: str
141
+ # 曲线上的采样点(用于找标签锚点)。numpy 数组,可能是 NaN。
142
+ xs: np.ndarray
143
+ ys: np.ndarray
144
+ # 标签候选锚点(在数据坐标里)
145
+ anchor: Optional[Tuple[float, float]] = None
146
+ # 最终标签锚点(经过避让后)
147
+ final_anchor: Optional[Tuple[float, float]] = None
148
+
149
+
150
+ class PlotRenderError(Exception):
151
+ pass
152
+
153
+
154
+ def _plot_explicit(ax, item: PlotItem, color: str) -> _PlottedCurve:
155
+ """绘制 y = f(x)。"""
156
+ # 从 "y = sin(x)" 中取出右边的 f(x)
157
+ if "=" not in item.expr:
158
+ raise PlotRenderError(f"显函数缺少 '=':{item.expr!r}")
159
+ lhs, rhs = item.expr.split("=", 1)
160
+ lhs = lhs.strip()
161
+ rhs = rhs.strip()
162
+ if lhs != "y":
163
+ raise PlotRenderError(
164
+ f"显函数左边必须是 'y',实际为 {lhs!r}(在 {item.expr!r} 中)"
165
+ )
166
+ expr = _to_sympy_expr(rhs)
167
+ x = sp.symbols("x")
168
+ try:
169
+ f = sp.lambdify(x, expr, modules=["numpy", _LOCAL_NAMES])
170
+ except Exception as ex:
171
+ raise PlotRenderError(f"显函数 lambdify 失败 {rhs!r}:{ex}")
172
+
173
+ a, b = item.x_range
174
+ a = float(a)
175
+ b = float(b)
176
+ n = max(400, int((b - a) * 80))
177
+ n = min(n, 4000)
178
+ xs = np.linspace(a, b, n)
179
+ try:
180
+ ys_raw = f(xs)
181
+ # 常数函数 / 标量返回:广播到与 xs 同形状
182
+ if np.isscalar(ys_raw) or (hasattr(ys_raw, "shape") and ys_raw.shape == ()):
183
+ ys = np.full_like(xs, float(ys_raw))
184
+ else:
185
+ ys = np.asarray(ys_raw, dtype=float)
186
+ # 形状不匹配(如某些 sympy 函数返回 (1, N))→ flatten
187
+ if ys.shape != xs.shape:
188
+ ys = ys.reshape(xs.shape) if ys.size == xs.size else np.broadcast_to(ys, xs.shape).astype(float)
189
+ except Exception as ex:
190
+ raise PlotRenderError(f"显函数求值失败 {rhs!r}:{ex}")
191
+
192
+ # 屏蔽 NaN / inf
193
+ mask = np.isfinite(ys)
194
+ xs_plot = xs[mask]
195
+ ys_plot = ys[mask]
196
+
197
+ # 限制 ys 在合理范围内(避免 tan 等函数在大值处画穿)
198
+ # 我们把 |y| > 1e6 的点设为 NaN,让 matplotlib 断开线段
199
+ ys_limited = ys.copy()
200
+ ys_limited[np.abs(ys_limited) > 1e6] = np.nan
201
+ ax.plot(xs, ys_limited, color=color, linewidth=1.6, label=item.label or "")
202
+
203
+ # 计算可见部分几何中点
204
+ if len(xs_plot) > 0:
205
+ # 用弧长加权中点:找累计弧长 50% 处的点
206
+ dx = np.diff(xs_plot)
207
+ dy = np.diff(ys_plot)
208
+ seglen = np.sqrt(dx * dx + dy * dy)
209
+ cum = np.concatenate([[0], np.cumsum(seglen)])
210
+ if cum[-1] > 0:
211
+ half = cum[-1] / 2.0
212
+ idx = int(np.searchsorted(cum, half))
213
+ idx = min(idx, len(xs_plot) - 1)
214
+ anchor = (float(xs_plot[idx]), float(ys_plot[idx]))
215
+ else:
216
+ anchor = (float(xs_plot[0]), float(ys_plot[0]))
217
+ else:
218
+ anchor = ((a + b) / 2.0, 0.0)
219
+
220
+ return _PlottedCurve(
221
+ label=item.label, color=color, xs=xs_plot, ys=ys_plot, anchor=anchor
222
+ )
223
+
224
+
225
+ def _plot_implicit(ax, item: PlotItem, color: str) -> _PlottedCurve:
226
+ """绘制 F(x, y) = 0。"""
227
+ if "=" not in item.expr:
228
+ raise PlotRenderError(f"隐函数缺少 '=':{item.expr!r}")
229
+ lhs, rhs = item.expr.split("=", 1)
230
+ lhs_expr = _to_sympy_expr(lhs.strip())
231
+ rhs_expr = _to_sympy_expr(rhs.strip())
232
+ F = lhs_expr - rhs_expr
233
+ x, y = sp.symbols("x y")
234
+ try:
235
+ f = sp.lambdify((x, y), F, modules=["numpy", _LOCAL_NAMES])
236
+ except Exception as ex:
237
+ raise PlotRenderError(f"隐函数 lambdify 失败 {item.expr!r}:{ex}")
238
+
239
+ a, b = item.x_range
240
+ c, d = item.y_range # 已在 parser 中校验非 None
241
+ a, b = float(a), float(b)
242
+ c, d = float(c), float(d)
243
+
244
+ n = 400
245
+ xs = np.linspace(a, b, n)
246
+ ys = np.linspace(c, d, n)
247
+ X, Y = np.meshgrid(xs, ys)
248
+ try:
249
+ Z = np.asarray(f(X, Y), dtype=float)
250
+ except Exception as ex:
251
+ raise PlotRenderError(f"隐函数求值失败 {item.expr!r}:{ex}")
252
+
253
+ # 屏蔽 NaN / inf
254
+ Z = np.where(np.isfinite(Z), Z, np.nan)
255
+
256
+ # 用 contour 画 0 等高线
257
+ try:
258
+ cs = ax.contour(X, Y, Z, levels=[0], colors=[color], linewidths=1.6)
259
+ except Exception as ex:
260
+ raise PlotRenderError(f"隐函数 contour 绘制失败 {item.expr!r}:{ex}")
261
+
262
+ # 取出 contour 的路径段,计算几何中点
263
+ # matplotlib 3.9+:ContourSet 自身就是 Collection;旧版本有 .collections
264
+ if hasattr(cs, "get_paths"):
265
+ try:
266
+ paths = cs.get_paths()
267
+ except Exception:
268
+ paths = []
269
+ elif hasattr(cs, "collections") and cs.collections:
270
+ try:
271
+ paths = cs.collections[0].get_paths()
272
+ except Exception:
273
+ paths = []
274
+ else:
275
+ paths = []
276
+
277
+ all_pts = []
278
+ for p in paths:
279
+ v = p.vertices
280
+ if len(v) > 0:
281
+ all_pts.append(v)
282
+ if all_pts:
283
+ pts = np.vstack(all_pts)
284
+ # 弧长中点
285
+ dx = np.diff(pts[:, 0])
286
+ dy = np.diff(pts[:, 1])
287
+ seglen = np.sqrt(dx * dx + dy * dy)
288
+ cum = np.concatenate([[0], np.cumsum(seglen)])
289
+ if cum[-1] > 0:
290
+ half = cum[-1] / 2.0
291
+ idx = int(np.searchsorted(cum, half))
292
+ idx = min(idx, len(pts) - 1)
293
+ anchor = (float(pts[idx, 0]), float(pts[idx, 1]))
294
+ else:
295
+ anchor = (float(pts[0, 0]), float(pts[0, 1]))
296
+ xs_plot = pts[:, 0]
297
+ ys_plot = pts[:, 1]
298
+ else:
299
+ anchor = ((a + b) / 2.0, (c + d) / 2.0)
300
+ xs_plot = np.array([])
301
+ ys_plot = np.array([])
302
+
303
+ return _PlottedCurve(
304
+ label=item.label, color=color, xs=xs_plot, ys=ys_plot, anchor=anchor
305
+ )
306
+
307
+
308
+ # ---------------------------------------------------------------------------
309
+ # 几何图形渲染
310
+ # ---------------------------------------------------------------------------
311
+
312
+ def _plot_shape(ax, item: PlotItem, color: str) -> _PlottedCurve:
313
+ """渲染几何图形。"""
314
+ shape = item.shape
315
+ params = item.params
316
+ label = item.label
317
+ anchor = None
318
+ xs_plot = np.array([])
319
+ ys_plot = np.array([])
320
+
321
+ if shape == "point":
322
+ x, y = params["at"]
323
+ ax.plot([x], [y], marker="o", color=color, markersize=6,
324
+ markeredgecolor=color, markerfacecolor=color)
325
+ anchor = (float(x), float(y))
326
+ xs_plot = np.array([float(x)])
327
+ ys_plot = np.array([float(y)])
328
+
329
+ elif shape == "segment":
330
+ x1, y1 = params["from"]
331
+ x2, y2 = params["to"]
332
+ ax.plot([x1, x2], [y1, y2], color=color, linewidth=1.8,
333
+ solid_capstyle="round")
334
+ # 端点小圆点
335
+ ax.plot([x1, x2], [y1, y2], marker="o", color=color, markersize=4,
336
+ linestyle="None")
337
+ anchor = (float((x1 + x2) / 2), float((y1 + y2) / 2))
338
+ xs_plot = np.array([float(x1), float(x2)])
339
+ ys_plot = np.array([float(y1), float(y2)])
340
+
341
+ elif shape == "line":
342
+ # 直线:过两点,但要延伸到 axes 边界
343
+ x1, y1 = params["from"]
344
+ x2, y2 = params["to"]
345
+ # 计算方向向量,延伸到当前 xlim 的两端
346
+ dx, dy = x2 - x1, y2 - y1
347
+ if dx == 0 and dy == 0:
348
+ raise PlotRenderError(f"直线 from 和 to 不能重合:{item.params}")
349
+ # 用参数 t 延伸:t=0 在 from,t=1 在 to,延伸到 t=-100..100 兜底
350
+ # 实际延伸长度由 set_xlim 后的 clip 决定
351
+ ts = np.array([-1000, 1000])
352
+ xs = x1 + ts * dx
353
+ ys = y1 + ts * dy
354
+ ax.plot(xs, ys, color=color, linewidth=1.4)
355
+ anchor = (float((x1 + x2) / 2), float((y1 + y2) / 2))
356
+ xs_plot = np.array([float(x1), float(x2)])
357
+ ys_plot = np.array([float(y1), float(y2)])
358
+
359
+ elif shape == "circle":
360
+ cx, cy = params["center"]
361
+ r = float(params["r"])
362
+ if r <= 0:
363
+ raise PlotRenderError(f"圆半径必须为正:r={r}")
364
+ theta = np.linspace(0, 2 * np.pi, 200)
365
+ xs = cx + r * np.cos(theta)
366
+ ys = cy + r * np.sin(theta)
367
+ ax.plot(xs, ys, color=color, linewidth=1.6)
368
+ # 圆心小点
369
+ ax.plot([cx], [cy], marker="o", color=color, markersize=3)
370
+ anchor = (float(cx), float(cy + r)) # 顶部
371
+ xs_plot = xs
372
+ ys_plot = ys
373
+
374
+ elif shape == "ellipse":
375
+ cx, cy = params["center"]
376
+ a = float(params["a"])
377
+ b = float(params["b"])
378
+ if a <= 0 or b <= 0:
379
+ raise PlotRenderError(f"椭圆半轴必须为正:a={a}, b={b}")
380
+ theta = np.linspace(0, 2 * np.pi, 200)
381
+ xs = cx + a * np.cos(theta)
382
+ ys = cy + b * np.sin(theta)
383
+ ax.plot(xs, ys, color=color, linewidth=1.6)
384
+ ax.plot([cx], [cy], marker="o", color=color, markersize=3)
385
+ anchor = (float(cx), float(cy + b))
386
+ xs_plot = xs
387
+ ys_plot = ys
388
+
389
+ elif shape == "polygon":
390
+ pts = params["points"]
391
+ if len(pts) < 3:
392
+ raise PlotRenderError(f"多边形至少需要 3 个顶点,得到 {len(pts)}")
393
+ xs = [p[0] for p in pts] + [pts[0][0]]
394
+ ys = [p[1] for p in pts] + [pts[0][1]]
395
+ ax.plot(xs, ys, color=color, linewidth=1.6)
396
+ ax.plot([p[0] for p in pts], [p[1] for p in pts],
397
+ marker="o", color=color, markersize=4, linestyle="None")
398
+ # 几何中心
399
+ cx = sum(p[0] for p in pts) / len(pts)
400
+ cy = sum(p[1] for p in pts) / len(pts)
401
+ anchor = (float(cx), float(cy))
402
+ xs_plot = np.array([float(x) for x in xs])
403
+ ys_plot = np.array([float(y) for y in ys])
404
+
405
+ elif shape == "rectangle":
406
+ x0, y0 = params["origin"]
407
+ w = float(params["w"])
408
+ h = float(params["h"])
409
+ if w <= 0 or h <= 0:
410
+ raise PlotRenderError(f"矩形宽高必须为正:w={w}, h={h}")
411
+ xs = [x0, x0 + w, x0 + w, x0, x0]
412
+ ys = [y0, y0, y0 + h, y0 + h, y0]
413
+ ax.plot(xs, ys, color=color, linewidth=1.6)
414
+ anchor = (float(x0 + w / 2), float(y0 + h / 2))
415
+ xs_plot = np.array([float(x) for x in xs])
416
+ ys_plot = np.array([float(y) for y in ys])
417
+
418
+ elif shape == "vector":
419
+ x1, y1 = params["from"]
420
+ x2, y2 = params["to"]
421
+ # 用 ax.annotate 画带箭头的向量
422
+ ax.annotate(
423
+ "",
424
+ xy=(x2, y2), xytext=(x1, y1),
425
+ arrowprops=dict(arrowstyle="->", color=color, lw=1.8),
426
+ )
427
+ anchor = (float((x1 + x2) / 2), float((y1 + y2) / 2))
428
+ xs_plot = np.array([float(x1), float(x2)])
429
+ ys_plot = np.array([float(y1), float(y2)])
430
+
431
+ elif shape == "parabola":
432
+ # 标准方程(vertex=(h,k), p=焦距, direction=up/down/left/right):
433
+ # up: (x-h)^2 = 4p(y-k) → y = (x-h)^2/(4p) + k
434
+ # down: (x-h)^2 = -4p(y-k) → y = -(x-h)^2/(4p) + k
435
+ # right: (y-k)^2 = 4p(x-h) → x = (y-k)^2/(4p) + h
436
+ # left: (y-k)^2 = -4p(x-h) → x = -(y-k)^2/(4p) + h
437
+ h, k = params["vertex"]
438
+ p = float(params["p"])
439
+ direction = str(params.get("direction", "up")).lower().strip()
440
+ if p == 0:
441
+ raise PlotRenderError("parabola 的 p 不能为 0")
442
+ # 在 vertex 周围画 4 倍 p 的范围(够看清形状又不会太远)
443
+ span = abs(p) * 4 + 1.5
444
+ if direction in ("up", "down"):
445
+ sign = 1 if direction == "up" else -1
446
+ xs = np.linspace(float(h) - span, float(h) + span, 400)
447
+ ys = sign * (xs - float(h)) ** 2 / (4 * p) + float(k)
448
+ ax.plot(xs, ys, color=color, linewidth=1.6)
449
+ # 顶点
450
+ ax.plot([h], [k], marker="o", color=color, markersize=4)
451
+ anchor = (float(h), float(k) + sign * abs(p))
452
+ xs_plot = xs
453
+ ys_plot = ys
454
+ elif direction in ("left", "right"):
455
+ sign = 1 if direction == "right" else -1
456
+ ys = np.linspace(float(k) - span, float(k) + span, 400)
457
+ xs = sign * (ys - float(k)) ** 2 / (4 * p) + float(h)
458
+ ax.plot(xs, ys, color=color, linewidth=1.6)
459
+ ax.plot([h], [k], marker="o", color=color, markersize=4)
460
+ anchor = (float(h) + sign * abs(p), float(k))
461
+ xs_plot = xs
462
+ ys_plot = ys
463
+ else:
464
+ raise PlotRenderError(
465
+ f"parabola 的 direction 必须是 up/down/left/right,得到 {direction!r}"
466
+ )
467
+
468
+ else:
469
+ raise PlotRenderError(f"未知几何图形:{shape!r}")
470
+
471
+ return _PlottedCurve(
472
+ label=label, color=color, xs=xs_plot, ys=ys_plot, anchor=anchor
473
+ )
474
+
475
+
476
+ # ---------------------------------------------------------------------------
477
+ # 标签避让
478
+ # ---------------------------------------------------------------------------
479
+
480
+ def _resolve_label_anchor(
481
+ ax, curves: List[_PlottedCurve], renderer
482
+ ) -> None:
483
+ """对每条有 label 的曲线计算最终锚点。
484
+
485
+ 算法:
486
+ 1. 初始锚点 = 曲线几何中点,向上偏移 0.04 * (y_max - y_min)。
487
+ 2. 把数据坐标转成显示坐标(pixels),用文本 bbox 检测重叠。
488
+ 3. 重叠时,对后插入的标签,沿 8 个方向(上、右上、右、右下、下、左下、左、左上)
489
+ 螺旋外扩,直到不重叠或达到最大迭代次数。
490
+ 4. 若仍重叠,则把标签放在轴外上方,作为兜底。
491
+ """
492
+ if not curves:
493
+ return
494
+
495
+ # 数据范围
496
+ all_xmin = min(np.min(c.xs) if len(c.xs) else math.inf for c in curves if c.anchor)
497
+ all_xmax = max(np.max(c.xs) if len(c.xs) else -math.inf for c in curves if c.anchor)
498
+ all_ymin = min(np.min(c.ys) if len(c.ys) else math.inf for c in curves if c.anchor)
499
+ all_ymax = max(np.max(c.ys) if len(c.ys) else -math.inf for c in curves if c.anchor)
500
+ if not all(math.isfinite(v) for v in [all_xmin, all_xmax, all_ymin, all_ymax]):
501
+ # 兜底
502
+ y_offset = 0.1
503
+ else:
504
+ y_offset = 0.04 * (all_ymax - all_ymin) if all_ymax > all_ymin else 0.1
505
+
506
+ placed_boxes = [] # 已放置标签的 bbox(display coords)
507
+
508
+ for c in curves:
509
+ if not c.label or c.anchor is None:
510
+ c.final_anchor = c.anchor
511
+ continue
512
+ ax_pt = (c.anchor[0], c.anchor[1] + y_offset)
513
+ # 转 display 坐标
514
+ disp = ax.transData.transform(ax_pt)
515
+
516
+ # 估算文本 bbox
517
+ # 用 renderer 测量
518
+ text_obj = ax.text(
519
+ ax_pt[0], ax_pt[1], c.label,
520
+ color=c.color, fontsize=10,
521
+ ha="center", va="bottom",
522
+ )
523
+ try:
524
+ bb = text_obj.get_window_extent(renderer=renderer)
525
+ except Exception:
526
+ bb = None
527
+
528
+ # 螺旋避让
529
+ if bb is not None:
530
+ best_bb = bb
531
+ best_disp = disp
532
+ best_pt = ax_pt
533
+ directions = [
534
+ (0, 1), (1, 1), (1, 0), (1, -1),
535
+ (0, -1), (-1, -1), (-1, 0), (-1, 1),
536
+ ]
537
+ step = 8 # 像素
538
+ max_iter = 40
539
+ ok = False
540
+ for it in range(max_iter):
541
+ overlap = False
542
+ for prev in placed_boxes:
543
+ if _bbox_overlap(best_bb, prev):
544
+ overlap = True
545
+ break
546
+ if not overlap:
547
+ ok = True
548
+ break
549
+ # 沿外扩方向移动
550
+ d = directions[it % len(directions)]
551
+ radius = step * (1 + it // len(directions))
552
+ new_disp = (best_disp[0] + d[0] * radius, best_disp[1] + d[1] * radius)
553
+ # 回到数据坐标
554
+ new_pt = ax.transData.inverted().transform(new_disp)
555
+ text_obj.set_position((float(new_pt[0]), float(new_pt[1])))
556
+ try:
557
+ best_bb = text_obj.get_window_extent(renderer=renderer)
558
+ except Exception:
559
+ pass
560
+ best_disp = new_disp
561
+ best_pt = (float(new_pt[0]), float(new_pt[1]))
562
+
563
+ placed_boxes.append(best_bb)
564
+ c.final_anchor = best_pt
565
+ if not ok:
566
+ # 已经尽力,保留当前位置
567
+ pass
568
+ else:
569
+ c.final_anchor = ax_pt
570
+ placed_boxes.append(None)
571
+
572
+
573
+ def _bbox_overlap(a, b) -> bool:
574
+ if a is None or b is None:
575
+ return False
576
+ # matplotlib Bbox:x0,y0,x1,y1
577
+ return not (a.x1 < b.x0 or b.x1 < a.x0 or a.y1 < b.y0 or b.y1 < a.y0)
578
+
579
+
580
+ # ---------------------------------------------------------------------------
581
+ # 主入口
582
+ # ---------------------------------------------------------------------------
583
+
584
+ def render_plot(
585
+ plot: Plot,
586
+ out_path: str,
587
+ dpi: int = 150,
588
+ display_width: Optional[float] = None,
589
+ ) -> str:
590
+ r"""渲染一张 plot 到 PNG。
591
+
592
+ 参数:
593
+ plot: Plot 节点
594
+ out_path: 输出 PNG 路径
595
+ dpi: 目标有效 DPI(每英寸显示长度的像素数)。若 display_width 给定,
596
+ 实际 savefig dpi 会自动调整以保持此有效分辨率一致。
597
+ display_width: 图在文档中的显示宽度(相对于 \paperwidth,0~1)。
598
+ 给定时启用 auto-DPI;为 None 时直接用 dpi 作为 savefig dpi。
599
+
600
+ 返回 out_path。
601
+ """
602
+ if not plot.items:
603
+ raise PlotRenderError("@plot 指令没有任何绘制项")
604
+
605
+ # ---------- 计算 savefig dpi(auto-DPI)----------
606
+ # 思路:无论图显示多大,保证"每英寸显示长度的像素数"≈ dpi,
607
+ # 这样大图小图都有相同的视觉清晰度,不浪费像素也不糊。
608
+ # display_inches = display_width × paperwidth_inches
609
+ # pixel_width = display_inches × dpi
610
+ # savefig_dpi = pixel_width / figsize_width
611
+ _FIGSIZE_WIDTH = 5.0 # 当前 figsize=(5, 4) 的宽度
612
+ _PAPERWIDTH_INCHES = 8.27 # A4 纸宽 21cm
613
+ _MIN_DPI = 80 # 下限:避免小图文字锯齿
614
+ _MAX_DPI = 400 # 上限:避免大图文件过大
615
+ if display_width is not None:
616
+ display_inches = display_width * _PAPERWIDTH_INCHES
617
+ auto_dpi = display_inches * dpi / _FIGSIZE_WIDTH
618
+ savefig_dpi = max(_MIN_DPI, min(_MAX_DPI, auto_dpi))
619
+ else:
620
+ savefig_dpi = float(dpi)
621
+
622
+ # ---------- 决定坐标范围 ----------
623
+ # 函数曲线项必须有 x_range(parser 已校验),几何图形项的 x_range 可选
624
+ # 收集:用户显式指定的范围 + 几何图形的图形边界
625
+ user_xmin = math.inf
626
+ user_xmax = -math.inf
627
+ user_ymin = math.inf
628
+ user_ymax = -math.inf
629
+ has_user_x = False
630
+ has_user_y = False
631
+ has_implicit = False
632
+ has_shape = False
633
+
634
+ for it in plot.items:
635
+ if it.kind == "implicit":
636
+ has_implicit = True
637
+ if it.kind == "shape":
638
+ has_shape = True
639
+ if it.x_range is not None:
640
+ a, b = it.x_range
641
+ user_xmin = min(user_xmin, float(a))
642
+ user_xmax = max(user_xmax, float(b))
643
+ has_user_x = True
644
+ if it.y_range is not None:
645
+ c, d = it.y_range
646
+ user_ymin = min(user_ymin, float(c))
647
+ user_ymax = max(user_ymax, float(d))
648
+ has_user_y = True
649
+
650
+ fig, ax = plt.subplots(figsize=(5, 4), constrained_layout=True)
651
+
652
+ # ---------- 渲染所有项 ----------
653
+ curves: List[_PlottedCurve] = []
654
+ for i, item in enumerate(plot.items):
655
+ color = _COLOR_CYCLE[i % len(_COLOR_CYCLE)]
656
+ if item.kind == "explicit":
657
+ cur = _plot_explicit(ax, item, color)
658
+ elif item.kind == "implicit":
659
+ cur = _plot_implicit(ax, item, color)
660
+ elif item.kind == "shape":
661
+ cur = _plot_shape(ax, item, color)
662
+ else:
663
+ raise PlotRenderError(f"未知绘图类型:{item.kind!r}")
664
+ curves.append(cur)
665
+
666
+ # ---------- 设置坐标范围 ----------
667
+ if has_user_x and has_user_y:
668
+ # 用户显式指定了 x、y 范围
669
+ xmin, xmax = user_xmin, user_xmax
670
+ ymin, ymax = user_ymin, user_ymax
671
+ elif has_user_x and not has_user_y:
672
+ # 用户只指定了 x 范围,y 由 autoscale 决定
673
+ xmin, xmax = user_xmin, user_xmax
674
+ ax.relim()
675
+ ax.autoscale_view()
676
+ ymin, ymax = ax.get_ylim()
677
+ ax.set_xlim(xmin, xmax)
678
+ elif has_shape and not has_user_x:
679
+ # 纯几何图形,无任何范围指定:从图形数据自动估算
680
+ ax.relim()
681
+ ax.autoscale_view()
682
+ xmin, xmax = ax.get_xlim()
683
+ ymin, ymax = ax.get_ylim()
684
+ else:
685
+ # 兜底(不应该到这)
686
+ ax.relim()
687
+ ax.autoscale_view()
688
+ xmin, xmax = ax.get_xlim()
689
+ ymin, ymax = ax.get_ylim()
690
+
691
+ # 给范围留 8% 边距(避免图形贴边)
692
+ x_span = xmax - xmin
693
+ y_span = ymax - ymin
694
+ if x_span > 0:
695
+ pad = x_span * 0.08
696
+ xmin -= pad
697
+ xmax += pad
698
+ if y_span > 0:
699
+ pad = y_span * 0.08
700
+ ymin -= pad
701
+ ymax += pad
702
+
703
+ ax.set_xlim(xmin, xmax)
704
+ ax.set_ylim(ymin, ymax)
705
+
706
+ # ---------- 决定 aspect ----------
707
+ # 隐函数图、几何图形、显隐混合 → equal(几何意义正确)
708
+ # 纯显函数图 → equal 除非 y 跨度 > 3 倍 x 跨度(避免 tan 等压扁)
709
+ if has_implicit or has_shape:
710
+ ax.set_aspect("equal", adjustable="box")
711
+ else:
712
+ x_span_final = xmax - xmin
713
+ y_span_final = ymax - ymin
714
+ if x_span_final > 0 and y_span_final / x_span_final <= 3.0:
715
+ ax.set_aspect("equal", adjustable="box")
716
+
717
+ # ---------- 增强坐标系绘制 ----------
718
+ _draw_axes(ax, xmin, xmax, ymin, ymax)
719
+
720
+ # 标签避让(必须先 draw 一次才能拿到 renderer)
721
+ fig.canvas.draw()
722
+ renderer = fig.canvas.get_renderer()
723
+ _resolve_label_anchor(ax, curves, renderer)
724
+
725
+ fig.savefig(out_path, dpi=savefig_dpi)
726
+ plt.close(fig)
727
+ return out_path
728
+
729
+
730
+ def _draw_axes(ax, xmin, xmax, ymin, ymax) -> None:
731
+ """绘制增强坐标系:箭头轴线、原点 O、x/y 轴标签。
732
+
733
+ 与 matplotlib 默认 spines 不同,这里用 axhline/axvline + annotate 画带箭头的轴,
734
+ 更接近中学/大学数学教材的坐标系画法。
735
+ """
736
+ # 隐藏默认 spines
737
+ for spine in ax.spines.values():
738
+ spine.set_visible(False)
739
+
740
+ # 网格
741
+ ax.grid(True, linewidth=0.4, alpha=0.4, color="#ccc")
742
+
743
+ # 坐标轴:用浅色线穿过原点(如果原点在视野内)
744
+ ax.axhline(0, color="#444", linewidth=1.0, zorder=1)
745
+ ax.axvline(0, color="#444", linewidth=1.0, zorder=1)
746
+
747
+ # x 轴箭头(右端)
748
+ ax.annotate(
749
+ "",
750
+ xy=(xmax, 0), xytext=(xmax - (xmax - xmin) * 0.04, 0),
751
+ arrowprops=dict(arrowstyle="->", color="#444", lw=1.2),
752
+ annotation_clip=False,
753
+ )
754
+ # y 轴箭头(上端)
755
+ ax.annotate(
756
+ "",
757
+ xy=(0, ymax), xytext=(0, ymax - (ymax - ymin) * 0.04),
758
+ arrowprops=dict(arrowstyle="->", color="#444", lw=1.2),
759
+ annotation_clip=False,
760
+ )
761
+
762
+ # 轴标签:x 在右端下方,y 在上端左侧
763
+ ax.text(xmax, 0, " x", ha="left", va="bottom", fontsize=11,
764
+ color="#222", clip_on=False)
765
+ ax.text(0, ymax, "y ", ha="right", va="top", fontsize=11,
766
+ color="#222", clip_on=False)
767
+
768
+ # 原点 O(仅当原点在视野内且不在边缘时显示)
769
+ if xmin < 0 < xmax and ymin < 0 < ymax:
770
+ ax.text(0, 0, " O", ha="left", va="top", fontsize=9,
771
+ color="#222", clip_on=False)
772
+
773
+ # 刻度
774
+ ax.tick_params(axis="both", which="both", direction="out",
775
+ top=False, right=False, labelsize=8, colors="#444")