rsplot 0.3.1.dev3__py3-none-any.whl → 0.3.2.dev4__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,202 @@
1
+ """Area-overlap remapping and stable display support for gridded fields."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass
6
+ from typing import TYPE_CHECKING
7
+
8
+ import numpy as np
9
+
10
+ if TYPE_CHECKING:
11
+ from shapely.geometry.base import BaseGeometry
12
+
13
+
14
+ def inferred_corner_grid(centers: np.ndarray) -> np.ndarray:
15
+ """Estimate shared corners from the unfiltered 2-D native center lattice.
16
+
17
+ Midpoints use adjacent native rows/columns, never adjacent QA survivors.
18
+ Outer edges are linearly extrapolated. Missing geolocation stays missing.
19
+ """
20
+ if centers.ndim != 2 or min(centers.shape) < 2:
21
+ raise ValueError("估算像元边界需要至少 2×2 的原始定位网格。")
22
+ padded = np.pad(centers.astype(float), 1, mode="edge")
23
+ padded[0, 1:-1] = 2 * centers[0] - centers[1]
24
+ padded[-1, 1:-1] = 2 * centers[-1] - centers[-2]
25
+ padded[:, 0] = 2 * padded[:, 1] - padded[:, 2]
26
+ padded[:, -1] = 2 * padded[:, -2] - padded[:, -3]
27
+ return (
28
+ padded[:-1, :-1] + padded[1:, :-1] + padded[1:, 1:] + padded[:-1, 1:]
29
+ ) / 4
30
+
31
+
32
+ def pixel_corners(corners: np.ndarray) -> np.ndarray:
33
+ """Convert a shared (ny+1, nx+1) corner grid to (ny, nx, 4)."""
34
+ return np.stack(
35
+ [
36
+ corners[:-1, :-1],
37
+ corners[:-1, 1:],
38
+ corners[1:, 1:],
39
+ corners[1:, :-1],
40
+ ],
41
+ axis=-1,
42
+ )
43
+
44
+
45
+ def regular_grid(
46
+ extent: tuple[float, float, float, float],
47
+ res: float,
48
+ ) -> tuple[np.ndarray, np.ndarray]:
49
+ """Build globally aligned cell centers covering an extent."""
50
+ west, east, south, north = extent
51
+ x = (np.arange(np.floor(west / res), np.ceil(east / res)) + 0.5) * res
52
+ y = (np.arange(np.floor(south / res), np.ceil(north / res)) + 0.5) * res
53
+ return np.meshgrid(x, y)
54
+
55
+
56
+ def grid_edges(centers: np.ndarray, res: float) -> np.ndarray:
57
+ return np.concatenate((centers - res / 2, centers[-1:] + res / 2))
58
+
59
+
60
+ def rectangle_corners(
61
+ lon: np.ndarray,
62
+ lat: np.ndarray,
63
+ res: float,
64
+ ) -> tuple[np.ndarray, np.ndarray]:
65
+ x = grid_edges(lon[0], res)
66
+ y = grid_edges(lat[:, 0], res)
67
+ xx, yy = np.meshgrid(x, y)
68
+ return pixel_corners(xx).reshape(-1, 4), pixel_corners(yy).reshape(-1, 4)
69
+
70
+
71
+ def overlap_mean(
72
+ corner_lon: np.ndarray,
73
+ corner_lat: np.ndarray,
74
+ values: np.ndarray,
75
+ lon: np.ndarray,
76
+ lat: np.ndarray,
77
+ res: float,
78
+ ) -> np.ndarray:
79
+ """Remap valid source polygons with overlap-area weights in EPSG:6933.
80
+
81
+ No extrapolation: a target needs positive-area overlap with a valid source.
82
+ Chunked queries bound temporary target/intersection geometry allocations.
83
+ """
84
+ from pyproj import Transformer
85
+ from shapely import STRtree, area, intersection, is_valid, polygons
86
+
87
+ result = np.full(lon.size, np.nan)
88
+ if not len(values):
89
+ return result.reshape(lon.shape)
90
+ transform = Transformer.from_crs(4326, 6933, always_xy=True)
91
+ x, y = transform.transform(corner_lon, corner_lat)
92
+ finite = (
93
+ np.isfinite(values)
94
+ & np.isfinite(x).all(axis=1)
95
+ & np.isfinite(y).all(axis=1)
96
+ )
97
+ source = polygons(np.stack([x[finite], y[finite]], axis=-1))
98
+ source_values = values[finite]
99
+ valid = is_valid(source) & (area(source) > 0)
100
+ source, source_values = source[valid], source_values[valid]
101
+ if not len(source):
102
+ return result.reshape(lon.shape)
103
+ tree = STRtree(source)
104
+ xcorners, ycorners = rectangle_corners(lon, lat, res)
105
+ for start in range(0, lon.size, 8192):
106
+ stop = min(start + 8192, lon.size)
107
+ x, y = transform.transform(xcorners[start:stop], ycorners[start:stop])
108
+ target = polygons(np.stack([x, y], axis=-1))
109
+ ti, si = tree.query(target, predicate="intersects")
110
+ if not len(ti):
111
+ continue
112
+ weights = area(intersection(target[ti], source[si]))
113
+ # Shared edges can differ by a few floating-point ulps after projection.
114
+ # Exclude numerical slivers, without extending support to empty cells.
115
+ threshold = 1e-10 * np.minimum(area(target[ti]), area(source[si]))
116
+ weights = np.where(weights > threshold, weights, 0.0)
117
+ # Boundary-only touches do not supply an observation.
118
+ sums = np.bincount(
119
+ ti, weights=weights * source_values[si], minlength=len(target)
120
+ )
121
+ weight_sum = np.bincount(ti, weights=weights, minlength=len(target))
122
+ np.divide(
123
+ sums, weight_sum, out=result[start:stop], where=weight_sum > 0
124
+ )
125
+ return result.reshape(lon.shape)
126
+
127
+
128
+ def valid_grid_geometry(
129
+ lon: np.ndarray,
130
+ lat: np.ndarray,
131
+ valid: np.ndarray,
132
+ res: float,
133
+ ) -> BaseGeometry:
134
+ """Union valid analysis cells using row runs, preserving missing cells."""
135
+ from shapely import box, union_all
136
+
137
+ x = grid_edges(lon[0], res)
138
+ y = grid_edges(lat[:, 0], res)
139
+ rectangles = []
140
+ for row, mask in enumerate(valid):
141
+ changes = np.diff(np.r_[False, mask, False].astype(np.int8))
142
+ starts, stops = (
143
+ np.flatnonzero(changes == 1),
144
+ np.flatnonzero(changes == -1),
145
+ )
146
+ rectangles.extend(box(x[starts], y[row], x[stops], y[row + 1]))
147
+ return union_all(rectangles)
148
+
149
+
150
+ @dataclass
151
+ class FootprintDisplay:
152
+ lon: np.ndarray
153
+ lat: np.ndarray
154
+ values: np.ndarray
155
+ support: BaseGeometry
156
+
157
+
158
+ def prepare_footprint_display(
159
+ lon: np.ndarray,
160
+ lat: np.ndarray,
161
+ values: np.ndarray,
162
+ geometry: BaseGeometry,
163
+ *,
164
+ analysis_res: float,
165
+ display_res: float,
166
+ smooth_sigma: float | None = None,
167
+ ) -> FootprintDisplay:
168
+ """Resample a fixed field and return its invariant clipping geometry.
169
+
170
+ The renderer must clip the colored layer to ``support``. Coarse display
171
+ cells may span analysis holes; finite display values alone do not describe
172
+ valid coverage. Analysis spacing is supplied by the caller.
173
+ """
174
+ from rsplot.geo.processing import prepare_display_grid, region_cell_mask
175
+
176
+ finite = np.isfinite(values)
177
+ support = valid_grid_geometry(lon, lat, finite, analysis_res).intersection(
178
+ geometry
179
+ )
180
+ # Smooth only within the fixed analysis lattice, retaining QA/data holes.
181
+ prepared = prepare_display_grid(
182
+ lon, lat, values, geometry, smooth_sigma=smooth_sigma
183
+ )
184
+ x = grid_edges(lon[0], analysis_res)
185
+ y = grid_edges(lat[:, 0], analysis_res)
186
+ display_lon, display_lat = regular_grid(
187
+ (x[0], x[-1], y[0], y[-1]), display_res
188
+ )
189
+ cx, cy = rectangle_corners(lon, lat, analysis_res)
190
+ finite = np.isfinite(prepared).ravel()
191
+ display = overlap_mean(
192
+ cx[finite],
193
+ cy[finite],
194
+ prepared.ravel()[finite],
195
+ display_lon,
196
+ display_lat,
197
+ display_res,
198
+ )
199
+ display = np.where(
200
+ region_cell_mask(display_lon, display_lat, geometry), display, np.nan
201
+ )
202
+ return FootprintDisplay(display_lon, display_lat, display, support)
rsplot/geo/processing.py CHANGED
@@ -30,6 +30,8 @@ class PreparedGrid:
30
30
  coverage_pct: float
31
31
  n_valid_region: int
32
32
  total_region_pixels: int
33
+ # Prepared source before administrative center masking, for display only.
34
+ unmasked_values: np.ndarray | None = None
33
35
 
34
36
 
35
37
  def buffered_extent(
@@ -58,6 +60,7 @@ def prepare_existing_grid(
58
60
  prepared = (
59
61
  fill_nan_gaps(lon, lat, values) if fill_gaps else np.asarray(values)
60
62
  )
63
+ unmasked_values = prepared
61
64
  region_mask = region_grid_mask(lon, lat, geometry)
62
65
  prepared = mask_to_region(
63
66
  lon,
@@ -89,6 +92,7 @@ def prepare_existing_grid(
89
92
  coverage_pct=coverage,
90
93
  n_valid_region=n_valid,
91
94
  total_region_pixels=total,
95
+ unmasked_values=unmasked_values,
92
96
  )
93
97
 
94
98
 
@@ -118,3 +122,81 @@ def prepare_swath_grid(
118
122
  fill_gaps=True,
119
123
  smooth_sigma=smooth_sigma,
120
124
  )
125
+
126
+
127
+ def region_cell_mask(
128
+ lon: np.ndarray,
129
+ lat: np.ndarray,
130
+ geometry: BaseGeometry,
131
+ ) -> np.ndarray:
132
+ """Select rectilinear cells touching the region, including exterior centers.
133
+
134
+ Edges match pcolormesh's midpoint convention. A singleton axis has zero
135
+ width, matching the renderer when only one center is supplied.
136
+ """
137
+ from shapely import box, intersects
138
+
139
+ if lon.ndim != 2 or lat.shape != lon.shape:
140
+ raise ValueError("显示网格经纬度必须为同形状二维数组。")
141
+ if not (np.allclose(lon, lon[0:1, :]) and np.allclose(lat, lat[:, 0:1])):
142
+ raise ValueError("行政边界显示裁剪需要规则经纬度网格。")
143
+
144
+ def edges(centers: np.ndarray) -> np.ndarray:
145
+ if centers.size == 1:
146
+ return np.repeat(centers, 2)
147
+ half_steps = np.diff(centers) / 2
148
+ return np.concatenate(
149
+ (
150
+ centers[:1] - half_steps[:1],
151
+ centers[:-1] + half_steps,
152
+ centers[-1:] + half_steps[-1:],
153
+ )
154
+ )
155
+
156
+ x = edges(lon[0])
157
+ y = edges(lat[:, 0])
158
+ cells = box(x[:-1][None, :], y[:-1, None], x[1:][None, :], y[1:, None])
159
+ return intersects(geometry, cells)
160
+
161
+
162
+ def prepare_display_grid(
163
+ lon: np.ndarray,
164
+ lat: np.ndarray,
165
+ values: np.ndarray,
166
+ geometry: BaseGeometry,
167
+ *,
168
+ smooth_sigma: float | None = None,
169
+ ) -> np.ndarray:
170
+ """Keep valid intersecting cells; exact geometry clipping is render-only.
171
+
172
+ Input must precede the administrative center mask. Smoothing can use the
173
+ surrounding buffer but cannot create newly valid cells. Statistics must
174
+ continue to use PreparedGrid.values, never this display array.
175
+ """
176
+ visible = region_cell_mask(lon, lat, geometry)
177
+ valid = np.isfinite(values)
178
+ display = np.asarray(values)
179
+ if smooth_sigma is not None and smooth_sigma > 0:
180
+ # Normalized convolution avoids NaN-gap warnings on sparse scans.
181
+ # This display-only operation does not change legacy statistics.
182
+ from scipy.ndimage import gaussian_filter
183
+
184
+ numerator = gaussian_filter(
185
+ np.where(valid, values, 0.0),
186
+ smooth_sigma,
187
+ mode="constant",
188
+ cval=0.0,
189
+ )
190
+ denominator = gaussian_filter(
191
+ valid.astype(float),
192
+ smooth_sigma,
193
+ mode="constant",
194
+ cval=0.0,
195
+ )
196
+ display = np.divide(
197
+ numerator,
198
+ denominator,
199
+ out=np.full(values.shape, np.nan),
200
+ where=denominator > 0,
201
+ )
202
+ return np.where(visible & valid, display, np.nan)
@@ -24,7 +24,9 @@ from rsplot.tiles.tianditu import TianDiTuTiles
24
24
 
25
25
  if TYPE_CHECKING:
26
26
  from cartopy.mpl.geoaxes import GeoAxes
27
+ from matplotlib.artist import Artist
27
28
  from matplotlib.figure import Figure
29
+ from shapely.geometry.base import BaseGeometry
28
30
 
29
31
  from rsplot.geo.boundaries import RegionInfo
30
32
 
@@ -90,13 +92,27 @@ def add_basemap_tiles(frame: MapFrame, region: RegionInfo) -> None:
90
92
 
91
93
  def add_region_boundaries(ax: GeoAxes, region: RegionInfo) -> None:
92
94
  """Draw detail, subdivision, and main boundaries in established order."""
95
+ selected = {
96
+ member.geometry.normalize().wkb
97
+ for member in getattr(region, "members", ())
98
+ }
99
+
100
+ def visible_geometries(frame):
101
+ if not selected:
102
+ return frame.geometry
103
+ return [
104
+ geometry
105
+ for geometry in frame.geometry
106
+ if geometry.normalize().wkb not in selected
107
+ ]
108
+
93
109
  if (
94
110
  region.detail_boundary_gdf is not None
95
111
  and len(region.detail_boundary_gdf) > 0
96
112
  ):
97
113
  ax.add_feature(
98
114
  ShapelyFeature(
99
- region.detail_boundary_gdf.geometry,
115
+ visible_geometries(region.detail_boundary_gdf),
100
116
  ccrs.PlateCarree(),
101
117
  edgecolor="#aaa",
102
118
  facecolor="none",
@@ -112,7 +128,7 @@ def add_region_boundaries(ax: GeoAxes, region: RegionInfo) -> None:
112
128
  is_county_sub = region.level == "city"
113
129
  ax.add_feature(
114
130
  ShapelyFeature(
115
- region.sub_boundary_gdf.geometry,
131
+ visible_geometries(region.sub_boundary_gdf),
116
132
  ccrs.PlateCarree(),
117
133
  edgecolor="#888" if not is_county_sub else "#666",
118
134
  facecolor="none",
@@ -133,6 +149,9 @@ def add_region_boundaries(ax: GeoAxes, region: RegionInfo) -> None:
133
149
  )
134
150
  )
135
151
 
152
+ if getattr(region, "members", ()):
153
+ add_member_labels(ax, region)
154
+
136
155
 
137
156
  def add_subdivision_labels(ax: GeoAxes, region: RegionInfo) -> None:
138
157
  """Add name labels for a region's subdivision features."""
@@ -143,7 +162,15 @@ def add_subdivision_labels(ax: GeoAxes, region: RegionInfo) -> None:
143
162
  name_field = region.sub_name_field
144
163
  fontsize = region.params.label_fontsize
145
164
 
165
+ selected_geometry = {
166
+ m.geometry.normalize().wkb for m in getattr(region, "members", ())
167
+ }
168
+ seen_labels = set()
146
169
  for _, row in gdf.iterrows():
170
+ key = row.geometry.normalize().wkb
171
+ if key in selected_geometry or (row[name_field], key) in seen_labels:
172
+ continue
173
+ seen_labels.add((row[name_field], key))
147
174
  point = row.geometry.representative_point()
148
175
  if region.level in (
149
176
  "key_region",
@@ -168,6 +195,31 @@ def add_subdivision_labels(ax: GeoAxes, region: RegionInfo) -> None:
168
195
  )
169
196
 
170
197
 
198
+ def add_member_labels(ax: GeoAxes, region: RegionInfo) -> None:
199
+ """Label selections separately from subdivisions; qualify duplicate names."""
200
+ from collections import Counter
201
+
202
+ names = Counter(member.name for member in region.members)
203
+ for member in region.members:
204
+ point = member.geometry.representative_point()
205
+ label = (
206
+ (member.qualified_name or member.name)
207
+ if names[member.name] > 1
208
+ else member.name
209
+ )
210
+ ax.text(
211
+ point.x,
212
+ point.y,
213
+ label,
214
+ transform=ccrs.PlateCarree(),
215
+ fontsize=region.params.label_fontsize,
216
+ ha="center",
217
+ va="center",
218
+ color="#222",
219
+ path_effects=[pe.withStroke(linewidth=2, foreground="white")],
220
+ )
221
+
222
+
171
223
  def configure_map_axes(
172
224
  ax: GeoAxes,
173
225
  region: RegionInfo,
@@ -195,8 +247,12 @@ def configure_map_axes(
195
247
  gridlines.top_labels = gridlines.right_labels = False
196
248
  gridlines.xformatter = LONGITUDE_FORMATTER
197
249
  gridlines.yformatter = LATITUDE_FORMATTER
198
- gridlines.xlocator = mticker.MultipleLocator(region.params.grid_step)
199
- gridlines.ylocator = mticker.MultipleLocator(region.params.grid_step)
250
+ if getattr(region, "members", ()):
251
+ gridlines.xlocator = mticker.MaxNLocator(nbins=5)
252
+ gridlines.ylocator = mticker.MaxNLocator(nbins=5)
253
+ else:
254
+ gridlines.xlocator = mticker.MultipleLocator(region.params.grid_step)
255
+ gridlines.ylocator = mticker.MultipleLocator(region.params.grid_step)
200
256
 
201
257
 
202
258
  def finalize_map_figure(
@@ -218,3 +274,63 @@ def finalize_map_figure(
218
274
  plt.close(fig)
219
275
  else:
220
276
  plt.show()
277
+
278
+
279
+ def clip_raster_to_region(
280
+ artist: Artist,
281
+ ax: GeoAxes,
282
+ region: RegionInfo,
283
+ *,
284
+ coverage_geometry: BaseGeometry | None = None,
285
+ ) -> None:
286
+ """Clip only the data layer, preserving holes and disconnected islands."""
287
+ import numpy as np
288
+ from matplotlib.path import Path as MplPath
289
+ from shapely.geometry import GeometryCollection, MultiPolygon
290
+ from shapely.geometry.polygon import orient
291
+
292
+ def orient_geometry(geometry: BaseGeometry) -> BaseGeometry:
293
+ if geometry.geom_type == "Polygon":
294
+ return orient(geometry, sign=1.0)
295
+ if geometry.geom_type == "MultiPolygon":
296
+ return MultiPolygon(
297
+ [orient(part, sign=1.0) for part in geometry.geoms]
298
+ )
299
+ if geometry.geom_type == "GeometryCollection":
300
+ return GeometryCollection(
301
+ [orient_geometry(part) for part in geometry.geoms]
302
+ )
303
+ return geometry
304
+
305
+ # Project the geometry into the axes coordinates before building a clip
306
+ # path, so both PlateCarree and satellite-basemap Mercator axes work.
307
+ clip_geometry = region.geometry
308
+ if coverage_geometry is not None:
309
+ clip_geometry = clip_geometry.intersection(coverage_geometry)
310
+ geometry = ax.projection.project_geometry(
311
+ orient_geometry(clip_geometry), ccrs.PlateCarree()
312
+ )
313
+ paths = []
314
+
315
+ def append_polygons(geometry: BaseGeometry) -> None:
316
+ if geometry.is_empty:
317
+ return
318
+ if geometry.geom_type == "Polygon":
319
+ polygon = orient(geometry, sign=1.0)
320
+ for ring in (polygon.exterior, *polygon.interiors):
321
+ vertices = np.asarray(ring.coords)[:, :2]
322
+ codes = np.full(len(vertices), MplPath.LINETO, dtype=np.uint8)
323
+ codes[0], codes[-1] = MplPath.MOVETO, MplPath.CLOSEPOLY
324
+ paths.append(MplPath(vertices, codes))
325
+ elif hasattr(geometry, "geoms"):
326
+ for part in geometry.geoms:
327
+ append_polygons(part)
328
+
329
+ append_polygons(geometry)
330
+ if not paths:
331
+ raise ValueError("行政区边界没有可用于裁剪的多边形。")
332
+ artist.set_clip_path(MplPath.make_compound_path(*paths), ax.transData)
333
+ # Cartopy can create an auxiliary collection for wrapped mesh cells.
334
+ wrapped = getattr(artist, "_wrapped_collection_fix", None)
335
+ if wrapped is not None:
336
+ wrapped.set_clip_path(MplPath.make_compound_path(*paths), ax.transData)
@@ -19,6 +19,7 @@ from rsplot.plotting.colormaps import (
19
19
  from rsplot.plotting.map_frame import (
20
20
  add_basemap_tiles,
21
21
  add_region_boundaries,
22
+ clip_raster_to_region,
22
23
  configure_map_axes,
23
24
  create_map_frame,
24
25
  finalize_map_figure,
@@ -29,6 +30,8 @@ from rsplot.plotting.station import (
29
30
  )
30
31
 
31
32
  if TYPE_CHECKING:
33
+ from shapely.geometry.base import BaseGeometry
34
+
32
35
  from rsplot.geo.boundaries import RegionInfo
33
36
  from rsplot.readers.guokong import StationData
34
37
 
@@ -55,6 +58,7 @@ def plot_overlay(
55
58
  hcho_cmap_file: str | None = None,
56
59
  title: str | None = None,
57
60
  output: str | None = None,
61
+ coverage_geometry: BaseGeometry | None = None,
58
62
  ) -> None:
59
63
  """Draw satellite raster + station scatter on the same axes.
60
64
 
@@ -89,6 +93,7 @@ def plot_overlay(
89
93
  alpha=raster_alpha,
90
94
  zorder=1,
91
95
  )
96
+ clip_raster_to_region(im, ax, region, coverage_geometry=coverage_geometry)
92
97
 
93
98
  # --- Basemap tiles ---
94
99
  add_basemap_tiles(frame, region)
rsplot/plotting/raster.py CHANGED
@@ -17,6 +17,7 @@ from rsplot.plotting.map_frame import (
17
17
  add_basemap_tiles,
18
18
  add_region_boundaries,
19
19
  add_subdivision_labels,
20
+ clip_raster_to_region,
20
21
  configure_map_axes,
21
22
  create_map_frame,
22
23
  figure_size_for_level,
@@ -24,6 +25,8 @@ from rsplot.plotting.map_frame import (
24
25
  )
25
26
 
26
27
  if TYPE_CHECKING:
28
+ from shapely.geometry.base import BaseGeometry
29
+
27
30
  from rsplot.geo.boundaries import RegionInfo
28
31
 
29
32
 
@@ -43,6 +46,7 @@ def plot_raster(
43
46
  hcho_cmap_file: str | None = None,
44
47
  title: str | None = None,
45
48
  output: str | None = None,
49
+ coverage_geometry: BaseGeometry | None = None,
46
50
  colorbar_label: str = "NO2 VCD (x10^15 molec/cm2)",
47
51
  ) -> None:
48
52
  """Draw a raster map for the given region.
@@ -79,6 +83,7 @@ def plot_raster(
79
83
  shading="auto",
80
84
  alpha=alpha,
81
85
  )
86
+ clip_raster_to_region(im, ax, region, coverage_geometry=coverage_geometry)
82
87
 
83
88
  # --- Basemap tiles (beneath boundaries, above raster) ---
84
89
  add_basemap_tiles(frame, region)
@@ -3,7 +3,7 @@
3
3
  from __future__ import annotations
4
4
 
5
5
  import math
6
- from dataclasses import dataclass, field
6
+ from dataclasses import dataclass, field, replace
7
7
 
8
8
  import numpy as np
9
9
 
@@ -40,6 +40,10 @@ class ProductInfo:
40
40
  default_factory=dict
41
41
  )
42
42
 
43
+ sensor: str = "tropomi"
44
+ data_key: str | None = None
45
+ quantity: str | None = None
46
+
43
47
  def defaults_for_level(self, level: str) -> ProductPlotDefaults:
44
48
  """Return product-specific defaults for an administrative level."""
45
49
  return self.level_defaults.get(level, ProductPlotDefaults())
@@ -91,15 +95,45 @@ PRODUCT_REGISTRY: dict[str, ProductInfo] = {
91
95
  }
92
96
 
93
97
 
94
- def get_product_info(product: str) -> ProductInfo:
95
- """Return ProductInfo for a product name."""
98
+ def get_product_info(product: str, *, sensor: str = "tropomi") -> ProductInfo:
99
+ """Resolve a product within a sensor, preserving legacy registry keys."""
96
100
  key = product.lower()
101
+ sensor = sensor.lower()
102
+ if sensor not in ("tropomi", "gems"):
103
+ raise ValueError(f"未知卫星来源 '{sensor}',可用: tropomi / gems")
97
104
  if key not in PRODUCT_REGISTRY:
98
105
  available = ", ".join(PRODUCT_REGISTRY)
99
106
  raise ValueError(
100
107
  f"Unsupported product '{product}'. Available: {available}"
101
108
  )
102
- return PRODUCT_REGISTRY[key]
109
+ if sensor == "tropomi":
110
+ return PRODUCT_REGISTRY[key]
111
+ from rsplot.readers.gems import (
112
+ GemsHCHOReader,
113
+ GemsNO2Reader,
114
+ GemsO3PRReader,
115
+ GemsO3Reader,
116
+ )
117
+
118
+ readers = {
119
+ "no2": GemsNO2Reader,
120
+ "hcho": GemsHCHOReader,
121
+ "o3": GemsO3Reader,
122
+ "o3pr": GemsO3PRReader,
123
+ }
124
+ quantities = {
125
+ "no2": "tropospheric_no2_column",
126
+ "hcho": "hcho_column",
127
+ "o3": "total_ozone_column",
128
+ "o3pr": "tropospheric_ozone_column",
129
+ }
130
+ return replace(
131
+ PRODUCT_REGISTRY[key],
132
+ reader_cls=readers[key],
133
+ sensor="gems",
134
+ data_key="gems_o3" if key == "o3pr" else f"gems_{key}",
135
+ quantity=quantities[key],
136
+ )
103
137
 
104
138
 
105
139
  def auto_vrange(
rsplot/readers/base.py CHANGED
@@ -3,7 +3,9 @@
3
3
  from __future__ import annotations
4
4
 
5
5
  from abc import ABC, abstractmethod
6
- from dataclasses import dataclass
6
+ from collections.abc import Iterator
7
+ from dataclasses import dataclass, field
8
+ from typing import Any
7
9
 
8
10
  import numpy as np
9
11
  import xarray as xr
@@ -19,11 +21,24 @@ class SwathData:
19
21
  n_pixels: int
20
22
  n_files: int
21
23
  n_broken: int
24
+ metadata: dict[str, Any] = field(default_factory=dict)
25
+ corner_lon: np.ndarray | None = None # (n_pixels, 4)
26
+ corner_lat: np.ndarray | None = None # (n_pixels, 4)
22
27
 
23
28
 
24
29
  class BaseReader(ABC):
25
30
  """Interface for reading a single-day satellite product."""
26
31
 
32
+ def iter_scans(
33
+ self,
34
+ data_dir: str,
35
+ date: str,
36
+ extent: tuple[float, float, float, float],
37
+ qa_threshold: float = 0.5,
38
+ ) -> Iterator[SwathData]:
39
+ """Yield temporal batches; legacy readers supply one daily batch."""
40
+ yield self.read(data_dir, date, extent, qa_threshold)
41
+
27
42
  @abstractmethod
28
43
  def read(
29
44
  self,