forcingkit 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,1199 @@
1
+ import os
2
+ from pathlib import Path
3
+ from fastapi import APIRouter, HTTPException, Query, Response as FastAPIResponse
4
+ from fastapi.responses import Response
5
+ from pydantic import BaseModel
6
+ from typing import List
7
+ import io
8
+ import logging
9
+ from forcingkit import settings
10
+
11
+ logger = logging.getLogger(__name__)
12
+
13
+ router = APIRouter(prefix="/api/v1/cache", tags=["Cache Viewer"])
14
+
15
+
16
+ class CacheItem(BaseModel):
17
+ id: str
18
+ type: str # 'zarr', 'grib', 'nc'
19
+ size_mb: float
20
+ modified_time: float
21
+ path: str
22
+ # Enriched metadata (populated for zarr datasets)
23
+ label: str | None = None # Human-readable label
24
+ variables: list[str] | None = None # Data variable names
25
+ grid_shape: str | None = None # e.g. "3×4" or "point"
26
+ time_steps: int | None = None # Number of time steps
27
+ time_range: str | None = None # e.g. "Mar 2 00:00 → Mar 3 23:00"
28
+ source: str | None = None # e.g. "ERA5", "HRRR", "NOAA Water Level"
29
+ donor_model: str | None = None # e.g. "HYCOM", "NERACOOS", "NECOFS"
30
+ resolution: str | None = None # e.g. "~7km", "~200m", "~31km"
31
+ spatial_extent: str | None = None # e.g. "-71.25, 42.75 → -70.25, 43.50"
32
+ target_date: str | None = None # ISO timestamp from sidecar metadata (e.g. ICs)
33
+
34
+
35
+ def _enrich_zarr_metadata(zarr_path: Path) -> dict:
36
+ label = None
37
+ label = None
38
+ try:
39
+ import xarray as xr
40
+ import numpy as np
41
+
42
+ # Open with decode_times=True for proper datetime parsing
43
+ try:
44
+ ds = xr.open_zarr(str(zarr_path), consolidated=False, decode_times=True)
45
+ except Exception:
46
+ ds = xr.open_zarr(str(zarr_path), consolidated=False, decode_times=False)
47
+ # Filter out scalar (0-d) variables - these are GRIB metadata artifacts
48
+ variables = [v for v in ds.data_vars if ds[v].ndim > 0]
49
+
50
+ # Grid shape, spatial extent, and resolution from coordinate arrays
51
+ lat_dim = next(
52
+ (
53
+ d
54
+ for d in ["latitude", "lat", "y", "eta", "eta_rho", "eta_u", "eta_v"]
55
+ if d in ds.sizes
56
+ ),
57
+ None,
58
+ )
59
+ lon_dim = next(
60
+ (
61
+ d
62
+ for d in ["longitude", "lon", "x", "xi", "xi_rho", "xi_u", "xi_v"]
63
+ if d in ds.sizes
64
+ ),
65
+ None,
66
+ )
67
+ # Separate dimensions from physical coordinates
68
+ lat_coord = next(
69
+ (
70
+ c
71
+ for c in ["lat_rho", "latitude", "lat", "lat_u", "lat_v"]
72
+ if c in ds.coords
73
+ ),
74
+ None,
75
+ )
76
+ lon_coord = next(
77
+ (
78
+ c
79
+ for c in ["lon_rho", "longitude", "lon", "lon_u", "lon_v"]
80
+ if c in ds.coords
81
+ ),
82
+ None,
83
+ )
84
+
85
+ if not lat_coord and lat_dim in ds.coords:
86
+ lat_coord = lat_dim
87
+ if not lon_coord and lon_dim in ds.coords:
88
+ lon_coord = lon_dim
89
+
90
+ spatial_extent = None
91
+ resolution = None
92
+ if lat_dim and lon_dim:
93
+ nlat, nlon = ds.sizes[lat_dim], ds.sizes[lon_dim]
94
+ grid_shape = f"{nlat}\u00d7{nlon}" if nlat > 0 and nlon > 0 else "point"
95
+ elif lat_coord and lon_coord and lat_coord in ds.coords:
96
+ # Curvilinear: infer shape from coord arrays
97
+ lat_arr = ds[lat_coord].values
98
+ nlat = lat_arr.shape[0] if lat_arr.ndim >= 1 else 0
99
+ nlon = (
100
+ lat_arr.shape[1]
101
+ if lat_arr.ndim >= 2
102
+ else (lat_arr.shape[0] if lat_arr.ndim == 1 else 0)
103
+ )
104
+ grid_shape = f"{nlat}\u00d7{nlon}" if nlat > 0 and nlon > 0 else "point"
105
+ else:
106
+ nlat = nlon = 0
107
+ grid_shape = "point"
108
+
109
+ # Extract coordinate bounds for spatial extent and resolution
110
+ if (
111
+ lat_coord
112
+ and lon_coord
113
+ and lat_coord in ds.coords
114
+ and lon_coord in ds.coords
115
+ ):
116
+ lat_vals = ds[lat_coord].values
117
+ lon_vals = ds[lon_coord].values
118
+ if lat_vals.size > 0 and lon_vals.size > 0:
119
+ spatial_extent = (
120
+ f"{float(np.min(lon_vals)):.2f}, {float(np.min(lat_vals)):.2f}"
121
+ f" \u2192 {float(np.max(lon_vals)):.2f}, {float(np.max(lat_vals)):.2f}"
122
+ )
123
+ # Compute median grid spacing in meters
124
+ try:
125
+ if lat_vals.ndim == 1 and len(lat_vals) > 1:
126
+ dlat = float(np.median(np.abs(np.diff(lat_vals))))
127
+ dlon = float(np.median(np.abs(np.diff(lon_vals))))
128
+ elif (
129
+ lat_vals.ndim == 2
130
+ and lat_vals.shape[0] > 1
131
+ and lat_vals.shape[1] > 1
132
+ ):
133
+ dlat = float(np.median(np.abs(np.diff(lat_vals, axis=0))))
134
+ dlon = float(np.median(np.abs(np.diff(lon_vals, axis=1))))
135
+ else:
136
+ dlat = dlon = 0
137
+ if dlat > 0 or dlon > 0:
138
+ mid_lat = float(np.mean(lat_vals))
139
+ dx_m = dlon * 111_320 * np.cos(np.radians(mid_lat))
140
+ dy_m = dlat * 111_320
141
+ avg_m = (dx_m + dy_m) / 2
142
+ if avg_m >= 1000:
143
+ resolution = f"~{avg_m / 1000:.0f}km"
144
+ else:
145
+ resolution = f"~{avg_m:.0f}m"
146
+ except Exception:
147
+ pass
148
+
149
+ # Time info - use decoded timestamps when available
150
+ time_dim = next(
151
+ (d for d in ["time", "t", "step", "nt", "timeseries"] if d in ds.sizes),
152
+ None,
153
+ )
154
+ time_steps = ds.sizes[time_dim] if time_dim else None
155
+ source = ds.attrs.get("type", None) or ds.attrs.get("source", None)
156
+ if source and "NECOFS" in source:
157
+ donor_model = "UMASS"
158
+ source = "NECOFS ~193m"
159
+ elif source and "HRRR" in source:
160
+ donor_model = "NOAA"
161
+ elif source and "ERA" in source:
162
+ donor_model = "ECMWF"
163
+ else:
164
+ donor_model = ds.attrs.get("donor_id", None)
165
+
166
+ time_range = None
167
+ delta_t_str = ""
168
+
169
+ if "target_date" in ds.attrs:
170
+ try:
171
+ from datetime import datetime as dt
172
+ import pandas as pd
173
+
174
+ d0 = pd.to_datetime(ds.attrs["target_date"])
175
+ time_range = f"{d0.strftime('%b %d %H:%M')} UTC"
176
+ except Exception:
177
+ time_range = f"{ds.attrs['target_date']} UTC"
178
+ if not time_steps:
179
+ time_steps = 1
180
+ elif time_dim and time_steps and time_steps > 0:
181
+ try:
182
+ from datetime import datetime as dt
183
+ import pandas as pd
184
+
185
+ t_vals = ds[time_dim].values
186
+ t0 = np.datetime64(t_vals[0], "ns")
187
+ t1 = np.datetime64(t_vals[-1], "ns")
188
+ d0 = t0.astype("datetime64[s]").astype(dt)
189
+ d1 = t1.astype("datetime64[s]").astype(dt)
190
+
191
+ if len(t_vals) > 1:
192
+ dt_secs = int(
193
+ np.timedelta64(np.datetime64(t_vals[1], "ns") - t0, "s").astype(
194
+ int
195
+ )
196
+ )
197
+ if dt_secs > 0:
198
+ if dt_secs >= 3600:
199
+ delta_t_str = f" (Δt={dt_secs // 3600}hr)"
200
+ else:
201
+ delta_t_str = f" (Δt={dt_secs // 60}m)"
202
+ time_range = f"{d0.strftime('%b %d %H:%M')} → {d1.strftime('%b %d %H:%M')} UTC{delta_t_str}"
203
+ else:
204
+ time_range = f"{d0.strftime('%b %d %H:%M')} UTC"
205
+ except Exception:
206
+ pass
207
+
208
+ label = None
209
+ return {
210
+ "label": locals().get("label", None),
211
+ "variables": locals().get("variables", []),
212
+ "grid_shape": locals().get("grid_shape", None),
213
+ "time_steps": locals().get("time_steps", None),
214
+ "time_range": locals().get("time_range", None),
215
+ "source": locals().get("source", None),
216
+ "donor_model": locals().get("donor_model", None),
217
+ "resolution": locals().get("resolution", None),
218
+ "spatial_extent": locals().get("spatial_extent", None),
219
+ "target_date": locals().get("target_date", None),
220
+ }
221
+ except Exception as e:
222
+ logger.warning(f"Failed to enrich metadata for {zarr_path.name}: {e}")
223
+ label = None
224
+ return {}
225
+
226
+
227
+ @router.get("/inventory", response_model=List[CacheItem])
228
+ async def get_cache_inventory(response: FastAPIResponse):
229
+ """Crawls the local data cache and returns an inventory of all processed forcing datasets."""
230
+ response.headers["Cache-Control"] = "no-store"
231
+ cache_dir = Path(Path(settings.cache_dir()).expanduser())
232
+ if not cache_dir.exists():
233
+ return []
234
+
235
+ inventory = []
236
+
237
+ # 1. Find Zarr datasets (directories)
238
+ for zarr_path in cache_dir.rglob("*.zarr"):
239
+ if zarr_path.is_dir():
240
+ size_bytes = sum(
241
+ f.stat().st_size for f in zarr_path.rglob("*") if f.is_file()
242
+ )
243
+ meta = _enrich_zarr_metadata(zarr_path)
244
+ inventory.append(
245
+ CacheItem(
246
+ id=zarr_path.stem,
247
+ type="zarr",
248
+ size_mb=round(size_bytes / (1024 * 1024), 2),
249
+ modified_time=zarr_path.stat().st_mtime,
250
+ path=str(zarr_path),
251
+ **meta,
252
+ )
253
+ )
254
+
255
+ # 2. Find GRIB/NC datasets (files)
256
+ for ext_glob in ["*.grib", "*.grib2", "*.nc"]:
257
+ for file_path in cache_dir.rglob(ext_glob):
258
+ if file_path.is_file():
259
+ stem = file_path.stem.lower()
260
+ label = None
261
+ source = None
262
+ donor_model = None
263
+ variables = None
264
+ grid_shape = None
265
+ spatial_extent = None
266
+ resolution = None
267
+
268
+ if "hrrr" in stem:
269
+ if "_f" in stem:
270
+ continue
271
+ source = "HRRR ~3km"
272
+ donor_model = "NOAA"
273
+ resolution = "~3km"
274
+ label = f"{source} Boundary Conditions"
275
+ elif "era5" in stem:
276
+ source = "ERA5T ~31km" if "era5t" in stem else "ERA5 ~31km"
277
+ donor_model = "ECMWF"
278
+ resolution = "~31km"
279
+ label = f"{source} Boundary Conditions"
280
+ elif "necofs" in stem:
281
+ source = "NECOFS ~193m"
282
+ donor_model = "UMASS"
283
+ resolution = "~193m"
284
+ label = "NECOFS Raw Validated"
285
+
286
+ # Enrich GRIB/NC with grid and variable metadata
287
+ try:
288
+ import xarray as xr
289
+ import numpy as np
290
+
291
+ # Suppress noisy expected grib warnings
292
+ import logging
293
+
294
+ logging.getLogger("cfgrib.messages").setLevel(logging.ERROR)
295
+
296
+ gds = None
297
+ for _inv_filter in [
298
+ {"typeOfLevel": "heightAboveGround", "stepType": "instant"},
299
+ {"typeOfLevel": "surface", "stepType": "instant"},
300
+ {"typeOfLevel": "heightAboveGround"},
301
+ {"typeOfLevel": "surface"},
302
+ ]:
303
+ try:
304
+ _candidate = xr.open_dataset(
305
+ str(file_path),
306
+ decode_times=True,
307
+ engine="cfgrib",
308
+ backend_kwargs={"filter_by_keys": _inv_filter},
309
+ )
310
+ if gds is None or len(list(_candidate.data_vars)) > len(
311
+ list(gds.data_vars)
312
+ ):
313
+ gds = _candidate
314
+ else:
315
+ _candidate.close()
316
+ except Exception:
317
+ continue
318
+ if gds is None:
319
+ try:
320
+ gds = xr.open_dataset(str(file_path), decode_times=True)
321
+ except Exception:
322
+ gds = None
323
+
324
+ if gds is not None:
325
+ variables = [str(v) for v in gds.data_vars if gds[v].ndim > 0]
326
+ _lat = next(
327
+ (c for c in gds.coords if "lat" in str(c).lower()), None
328
+ )
329
+ _lon = next(
330
+ (c for c in gds.coords if "lon" in str(c).lower()), None
331
+ )
332
+ time_dim = next(
333
+ (
334
+ d
335
+ for d in ["time", "t", "step", "valid_time"]
336
+ if d in gds.sizes or d in gds.coords
337
+ ),
338
+ None,
339
+ )
340
+ time_steps = 1
341
+ time_range = None
342
+ if time_dim:
343
+ try:
344
+ from datetime import datetime as dt
345
+ import numpy as np
346
+
347
+ t_vals = gds[time_dim].values
348
+ if getattr(t_vals, "ndim", 0) == 0:
349
+ t_vals = np.array([t_vals])
350
+ time_steps = len(t_vals)
351
+ t0 = np.datetime64(t_vals[0], "ns")
352
+ t1 = np.datetime64(t_vals[-1], "ns")
353
+ d0 = t0.astype("datetime64[s]").astype(dt)
354
+ d1 = t1.astype("datetime64[s]").astype(dt)
355
+ delta_t_str = ""
356
+ if len(t_vals) > 1:
357
+ dt_secs = int(
358
+ np.timedelta64(
359
+ np.datetime64(t_vals[1], "ns") - t0, "s"
360
+ ).astype(int)
361
+ )
362
+ if dt_secs > 0:
363
+ delta_t_str = (
364
+ f" (Δt={dt_secs // 3600}hr)"
365
+ if dt_secs >= 3600
366
+ else f" (Δt={dt_secs // 60}m)"
367
+ )
368
+ if d0 == d1:
369
+ time_range = f"{d0.strftime('%b %d %H:%M')} UTC"
370
+ else:
371
+ time_range = f"{d0.strftime('%b %d %H:%M')} → {d1.strftime('%b %d %H:%M')} UTC{delta_t_str}"
372
+ except Exception:
373
+ pass
374
+
375
+ if _lat and _lon:
376
+ if _lat in gds.coords and _lon in gds.coords:
377
+ lat_v = gds[_lat].values
378
+ lon_v = gds[_lon].values
379
+ if getattr(lat_v, "ndim", 0) == 2:
380
+ grid_shape = (
381
+ f"{lat_v.shape[0]}\u00d7{lat_v.shape[1]}"
382
+ )
383
+ else:
384
+ nlat = (
385
+ lat_v.shape[0]
386
+ if getattr(lat_v, "ndim", 0) > 0
387
+ else 1
388
+ )
389
+ nlon = (
390
+ lon_v.shape[0]
391
+ if getattr(lon_v, "ndim", 0) > 0
392
+ else 1
393
+ )
394
+ grid_shape = f"{nlat}\u00d7{nlon}"
395
+
396
+ if lat_v.size > 0 and lon_v.size > 0:
397
+ spatial_extent = (
398
+ f"{float(np.min(lon_v)):.2f}, {float(np.min(lat_v)):.2f}"
399
+ f" \u2192 {float(np.max(lon_v)):.2f}, {float(np.max(lat_v)):.2f}"
400
+ )
401
+ if (
402
+ getattr(lat_v, "ndim", 0) == 1
403
+ and len(lat_v) > 1
404
+ and not resolution
405
+ ):
406
+ dlat = float(np.median(np.abs(np.diff(lat_v))))
407
+ dlon = float(np.median(np.abs(np.diff(lon_v))))
408
+ mid_lat = float(np.mean(lat_v))
409
+ avg_m = (
410
+ (
411
+ dlon
412
+ * 111320
413
+ * np.cos(np.radians(mid_lat))
414
+ )
415
+ + (dlat * 111320)
416
+ ) / 2
417
+ resolution = (
418
+ f"~{avg_m / 1000:.0f}km"
419
+ if avg_m >= 1000
420
+ else f"~{avg_m:.0f}m"
421
+ )
422
+ gds.close()
423
+ except Exception as enrich_err:
424
+ logger.debug(
425
+ f"GRIB enrichment failed for {file_path.name}: {enrich_err}"
426
+ )
427
+
428
+ inventory.append(
429
+ CacheItem(
430
+ id=file_path.stem,
431
+ type=file_path.suffix.lstrip("."),
432
+ size_mb=round(file_path.stat().st_size / (1024 * 1024), 2),
433
+ modified_time=file_path.stat().st_mtime,
434
+ path=str(file_path),
435
+ label=label,
436
+ source=source,
437
+ donor_model=donor_model,
438
+ variables=variables,
439
+ grid_shape=grid_shape,
440
+ spatial_extent=spatial_extent,
441
+ resolution=resolution,
442
+ time_steps=time_steps,
443
+ time_range=time_range,
444
+ )
445
+ )
446
+
447
+ # Sort by descending modification time (newest first)
448
+ inventory.sort(key=lambda x: x.modified_time, reverse=True)
449
+ return inventory
450
+
451
+
452
+ @router.delete("/dataset/{dataset_id}")
453
+ async def delete_dataset(dataset_id: str):
454
+ """Deletes a single dataset from the cache by ID."""
455
+ import shutil
456
+
457
+ cache_dir = Path(settings.cache_dir()).expanduser()
458
+
459
+ # Use rglob to find the dataset ID anywhere in the cache
460
+ matches = list(cache_dir.rglob(f"{dataset_id}.zarr"))
461
+ if matches:
462
+ zarr_path = matches[0]
463
+ if zarr_path.is_dir():
464
+ shutil.rmtree(zarr_path)
465
+ # Also remove sidecar metadata if it exists (JSON companion may be named nicely)
466
+ sidecar = zarr_path.with_suffix(".json")
467
+ if sidecar.exists():
468
+ sidecar.unlink()
469
+ # Also remove old metadata pattern
470
+ old_sidecar = zarr_path.parent / f"{dataset_id}_metadata.json"
471
+ if old_sidecar.exists():
472
+ old_sidecar.unlink()
473
+ logger.info(f"Deleted Zarr dataset: {zarr_path}")
474
+ return {"status": "success", "message": f"Deleted {dataset_id}"}
475
+
476
+ for ext in ["grib", "grib2", "nc"]:
477
+ matches = list(cache_dir.rglob(f"{dataset_id}.{ext}"))
478
+ if matches:
479
+ file_path = matches[0]
480
+ if file_path.is_file():
481
+ file_path.unlink()
482
+ # Remove sidecar if exists
483
+ sidecar = file_path.with_suffix(".json")
484
+ if sidecar.exists():
485
+ sidecar.unlink()
486
+ logger.info(f"Deleted dataset: {file_path}")
487
+ return {"status": "success", "message": f"Deleted {dataset_id}"}
488
+
489
+ raise HTTPException(status_code=404, detail=f"Dataset {dataset_id} not found.")
490
+
491
+
492
+ @router.get("/preview")
493
+ async def get_dataset_preview(
494
+ dataset_id: str = Query(..., description="ID of the dataset (without extension)"),
495
+ ext: str = Query(..., description="Extension of dataset (zarr, grib, nc)"),
496
+ var_name: str = Query(..., description="Variable to plot (e.g., u10, v10, tp)"),
497
+ time_idx: int = Query(0, description="Time step index relative to data slice"),
498
+ ):
499
+ """
500
+ Renders a dynamic Matplotlib base64 preview of a cached multidimensional dataset.
501
+ """
502
+
503
+ # Deferred plotting imports for memory savings on edge devices
504
+ import xarray as xr
505
+ import matplotlib
506
+
507
+ matplotlib.use("Agg") # Non-interactive backend
508
+ import matplotlib.pyplot as plt
509
+ import numpy as np
510
+
511
+ cache_dir = Path(settings.cache_dir()).expanduser()
512
+ full_path = ""
513
+
514
+ if ext == "zarr":
515
+ full_path = os.path.join(cache_dir, f"{dataset_id}.zarr")
516
+ else:
517
+ full_path = os.path.join(cache_dir, f"{dataset_id}.{ext}")
518
+
519
+ if not os.path.exists(full_path):
520
+ raise HTTPException(
521
+ status_code=404, detail=f"Dataset file {full_path} not found."
522
+ )
523
+
524
+ try:
525
+ ds = None
526
+ if ext == "zarr":
527
+ ds = xr.open_zarr(full_path, consolidated=False, decode_times=True)
528
+ else:
529
+ # Handle ambiguous GRIB levels by attempting common filters
530
+ try:
531
+ ds = xr.open_dataset(full_path, decode_times=True)
532
+ except Exception as e:
533
+ # Try common level types in priority order for surface forcing
534
+ for filter_obj in [
535
+ {"typeOfLevel": "surface", "stepType": "instant"},
536
+ {"typeOfLevel": "heightAboveGround", "stepType": "instant"},
537
+ {"typeOfLevel": "surface"},
538
+ {"typeOfLevel": "heightAboveGround"},
539
+ {"typeOfLevel": "isobaricInhPa"},
540
+ {"typeOfLevel": "meanSea"},
541
+ {"typeOfLevel": "atmosphere"},
542
+ ]:
543
+ try:
544
+ candidate = xr.open_dataset(
545
+ full_path,
546
+ decode_times=True,
547
+ engine="cfgrib",
548
+ backend_kwargs={"filter_by_keys": filter_obj},
549
+ )
550
+ # Accept this filter if it contains the requested variable
551
+ if var_name in candidate.data_vars or ds is None:
552
+ ds = candidate
553
+ if var_name in candidate.data_vars:
554
+ break
555
+ except Exception:
556
+ continue
557
+ if ds is None:
558
+ raise HTTPException(
559
+ status_code=500, detail=f"Failed to parse GRIB: {e}"
560
+ )
561
+
562
+ # Map common variable names across naming conventions
563
+ if var_name not in ds.variables:
564
+ var_mapping = {
565
+ # Ocean IC aliases
566
+ "u": ["water_u", "u_current", "uo", "u10"],
567
+ "v": ["water_v", "v_current", "vo", "v10"],
568
+ "temp": ["water_temp", "temperature", "thetao"],
569
+ "salt": ["salinity", "so"],
570
+ "zeta": ["surf_el", "ssh", "zos"],
571
+ # Atmospheric aliases (ERA5 u10/v10 ↔ HRRR u/v at heightAboveGround)
572
+ "u10": ["u", "10u", "u_10m"],
573
+ "v10": ["v", "10v", "v_10m"],
574
+ "t2m": ["t", "2t", "t_2m"],
575
+ }
576
+ if var_name in var_mapping:
577
+ for alt_name in var_mapping[var_name]:
578
+ if alt_name in ds.data_vars:
579
+ var_name = alt_name
580
+ break
581
+
582
+ if var_name not in ds.variables:
583
+ available = list(ds.data_vars.keys())
584
+ raise HTTPException(
585
+ status_code=400,
586
+ detail=f"Variable '{var_name}' not found. Try: {available}",
587
+ )
588
+
589
+ data_var = ds[var_name]
590
+
591
+ if data_var.ndim == 0:
592
+ raise HTTPException(
593
+ status_code=400,
594
+ detail=f"Variable '{var_name}' is scalar (0-d) and cannot be rendered.",
595
+ )
596
+
597
+ # Detect time dimension
598
+ time_dim = next(
599
+ (
600
+ d
601
+ for d in ["time", "t", "step", "nt", "timeseries"]
602
+ if d in data_var.dims
603
+ ),
604
+ None,
605
+ )
606
+ if not time_dim and len(data_var.dims) == 1:
607
+ time_dim = str(data_var.dims[0])
608
+
609
+ # Detect spatial dimensions
610
+ lat_dim = next(
611
+ (
612
+ d
613
+ for d in ["latitude", "lat", "y", "eta", "eta_rho", "eta_u", "eta_v"]
614
+ if d in data_var.dims
615
+ ),
616
+ None,
617
+ )
618
+ lon_dim = next(
619
+ (
620
+ d
621
+ for d in ["longitude", "lon", "x", "xi", "xi_rho", "xi_u", "xi_v"]
622
+ if d in data_var.dims
623
+ ),
624
+ None,
625
+ )
626
+ has_spatial = (
627
+ lat_dim
628
+ and lon_dim
629
+ and ds.sizes.get(lat_dim, 0) > 0
630
+ and ds.sizes.get(lon_dim, 0) > 0
631
+ )
632
+
633
+ # Determine rendering mode:
634
+ # - "timeseries" for 1D time-only data or point grids (0×0 spatial)
635
+ # - "timeseries_grid" for very small grids (<=3×3) where overlaid timeseries are useful
636
+ # - "heatmap_small" for small grids (<= 20 cells per dim) with annotated values
637
+ # - "heatmap" for larger spatial grids
638
+ nlat = ds.sizes.get(lat_dim, 0) if lat_dim else 0
639
+ nlon = ds.sizes.get(lon_dim, 0) if lon_dim else 0
640
+
641
+ if not has_spatial and time_dim:
642
+ render_mode = "timeseries"
643
+ elif has_spatial and nlat <= 3 and nlon <= 3 and time_dim:
644
+ render_mode = "timeseries_grid"
645
+ elif has_spatial and nlat <= 20 and nlon <= 20:
646
+ render_mode = "heatmap_small"
647
+ else:
648
+ render_mode = "heatmap"
649
+
650
+ # Decode timestamp for the selected time index
651
+ time_label = f"t={time_idx}"
652
+ if time_dim and time_dim in ds.variables:
653
+ try:
654
+ t_vals = ds[time_dim].values
655
+ idx = min(time_idx, len(t_vals) - 1)
656
+ t_val = t_vals[idx]
657
+
658
+ # Robust decoding to numpy datetime64
659
+ if isinstance(t_val, np.datetime64):
660
+ d64 = t_val.astype("datetime64[us]")
661
+ else:
662
+ try:
663
+ d64 = np.datetime64(t_val, "us")
664
+ except Exception:
665
+ # Fallback for floats: assume seconds from epoch
666
+ d64 = np.datetime64(int(float(t_val)), "s").astype(
667
+ "datetime64[us]"
668
+ )
669
+
670
+ from datetime import datetime as dt
671
+
672
+ d = d64.astype(dt)
673
+ time_label = d.strftime("%Y%m%d-%H:%M:%S.%f")[:-3]
674
+ except Exception:
675
+ pass
676
+
677
+ fig, ax = plt.subplots(figsize=(7, 4.5))
678
+ fig.patch.set_facecolor("#0d1117")
679
+ ax.set_facecolor("#161b22")
680
+
681
+ # Subtitle with decoded timestamp
682
+ if time_label:
683
+ fig.text(
684
+ 0.5,
685
+ 0.97,
686
+ time_label,
687
+ ha="center",
688
+ va="top",
689
+ fontsize=8.5,
690
+ color="#8b949e",
691
+ fontstyle="italic",
692
+ )
693
+
694
+ # Style for dark theme
695
+ for spine in ax.spines.values():
696
+ spine.set_color("#30363d")
697
+ ax.tick_params(colors="#8b949e", labelsize=8)
698
+ ax.xaxis.label.set_color("#c9d1d9")
699
+ ax.yaxis.label.set_color("#c9d1d9")
700
+ ax.title.set_color("#f0f6fc")
701
+
702
+ if render_mode == "timeseries":
703
+ # Plot all time steps as a line chart
704
+ values = data_var.values.flatten()
705
+ valid = (
706
+ ~np.isnan(values)
707
+ if values.dtype.kind == "f"
708
+ else np.ones_like(values, dtype=bool)
709
+ )
710
+ ax.plot(
711
+ np.arange(len(values)),
712
+ values,
713
+ color="#58a6ff",
714
+ linewidth=1.5,
715
+ marker="o" if len(values) < 50 else None,
716
+ markersize=3,
717
+ )
718
+ if time_idx < len(values):
719
+ ax.axvline(
720
+ time_idx, color="#f85149", linewidth=1, linestyle="--", alpha=0.7
721
+ )
722
+ if valid[time_idx]:
723
+ ax.annotate(
724
+ f"{values[time_idx]:.3f}",
725
+ xy=(time_idx, values[time_idx]),
726
+ xytext=(5, 10),
727
+ textcoords="offset points",
728
+ color="#f0f6fc",
729
+ fontsize=9,
730
+ fontweight="bold",
731
+ bbox=dict(boxstyle="round,pad=0.3", fc="#21262d", ec="#30363d"),
732
+ )
733
+ ax.set_xlabel("Time Step", fontsize=9)
734
+ ax.set_ylabel(var_name, fontsize=9)
735
+ ax.set_title(f"{dataset_id} - {var_name} ({time_label})", fontsize=10)
736
+ ax.grid(True, alpha=0.15, color="#8b949e")
737
+
738
+ elif render_mode == "timeseries_grid":
739
+ # Small grid: plot timeseries for each cell as overlaid lines
740
+ if time_dim:
741
+ for ilat in range(nlat):
742
+ for ilon in range(nlon):
743
+ slice_opts = {}
744
+ if lat_dim:
745
+ slice_opts[lat_dim] = ilat
746
+ if lon_dim:
747
+ slice_opts[lon_dim] = ilon
748
+ ts = data_var.isel(slice_opts).values.flatten()
749
+ label = f"({ilat},{ilon})" if nlat * nlon <= 16 else None
750
+ ax.plot(ts, linewidth=1, alpha=0.7, label=label)
751
+ if time_idx is not None:
752
+ ax.axvline(
753
+ time_idx,
754
+ color="#f85149",
755
+ linewidth=1,
756
+ linestyle="--",
757
+ alpha=0.7,
758
+ )
759
+ ax.set_xlabel("Time Step", fontsize=9)
760
+ ax.set_ylabel(var_name, fontsize=9)
761
+ ax.set_title(
762
+ f"{dataset_id} \u2014 {var_name} ({nlat}\u00d7{nlon} grid)",
763
+ fontsize=10,
764
+ )
765
+ ax.grid(True, alpha=0.15, color="#8b949e")
766
+ if nlat * nlon <= 16:
767
+ ax.legend(fontsize=7, loc="upper right", ncol=2, framealpha=0.3)
768
+
769
+ elif render_mode == "heatmap_small":
770
+ # Small spatial grid with value annotations
771
+ slice_opts = {}
772
+ if time_dim:
773
+ max_t = ds.sizes[time_dim] - 1
774
+ slice_opts[time_dim] = min(time_idx, max_t)
775
+ for z_name in [
776
+ "depth",
777
+ "zC",
778
+ "level",
779
+ "s_rho",
780
+ "s_w",
781
+ "isobaricInhPa",
782
+ "number",
783
+ "surface",
784
+ "step",
785
+ ]:
786
+ if z_name in data_var.dims:
787
+ slice_opts[z_name] = -1 if z_name in ("s_rho", "s_w") else 0
788
+ valid_opts = {k: v for k, v in slice_opts.items() if k in data_var.dims}
789
+ array = data_var.isel(valid_opts).values
790
+
791
+ vmin = np.nanmin(array) if not np.all(np.isnan(array)) else 0.0
792
+ vmax = np.nanmax(array) if not np.all(np.isnan(array)) else 1.0
793
+ im = ax.imshow(
794
+ array, origin="lower", cmap="turbo", aspect="auto", vmin=vmin, vmax=vmax
795
+ )
796
+ # Annotate cell values
797
+ for yi in range(array.shape[0]):
798
+ for xi in range(array.shape[1]):
799
+ v = array[yi, xi]
800
+ if not np.isnan(v):
801
+ ax.text(
802
+ xi,
803
+ yi,
804
+ f"{v:.2f}",
805
+ ha="center",
806
+ va="center",
807
+ fontsize=7,
808
+ color="white",
809
+ fontweight="bold",
810
+ bbox=dict(boxstyle="round,pad=0.2", fc="black", alpha=0.4),
811
+ )
812
+ cbar = plt.colorbar(im, ax=ax, label=var_name)
813
+ cbar.ax.yaxis.set_tick_params(color="#8b949e")
814
+ cbar.ax.yaxis.label.set_color("#c9d1d9")
815
+ plt.setp(cbar.ax.yaxis.get_ticklabels(), color="#8b949e")
816
+ ax.set_title(f"{dataset_id} \u2014 {var_name} ({time_label})", fontsize=10)
817
+
818
+ else:
819
+ # Standard heatmap for larger grids
820
+ slice_opts = {}
821
+ if time_dim:
822
+ max_t = ds.sizes[time_dim] - 1
823
+ slice_opts[time_dim] = min(time_idx, max_t)
824
+ for z_name in [
825
+ "depth",
826
+ "zC",
827
+ "level",
828
+ "s_rho",
829
+ "s_w",
830
+ "isobaricInhPa",
831
+ "number",
832
+ "surface",
833
+ "step",
834
+ ]:
835
+ if z_name in data_var.dims:
836
+ slice_opts[z_name] = -1 if z_name in ("s_rho", "s_w") else 0
837
+ valid_opts = {k: v for k, v in slice_opts.items() if k in data_var.dims}
838
+ array = data_var.isel(valid_opts).values
839
+
840
+ vmin = np.nanmin(array) if not np.all(np.isnan(array)) else 0.0
841
+ vmax = np.nanmax(array) if not np.all(np.isnan(array)) else 1.0
842
+ im = ax.imshow(
843
+ array, origin="lower", cmap="turbo", aspect="auto", vmin=vmin, vmax=vmax
844
+ )
845
+ cbar = plt.colorbar(im, ax=ax, label=var_name)
846
+ cbar.ax.yaxis.set_tick_params(color="#8b949e")
847
+ cbar.ax.yaxis.label.set_color("#c9d1d9")
848
+ plt.setp(cbar.ax.yaxis.get_ticklabels(), color="#8b949e")
849
+ ax.set_title(f"{dataset_id} \u2014 {var_name} ({time_label})", fontsize=10)
850
+
851
+ buf = io.BytesIO()
852
+ plt.tight_layout(rect=(0, 0, 1, 0.95) if time_label else (0, 0, 1, 1))
853
+ plt.savefig(buf, format="png", dpi=100, facecolor=fig.get_facecolor())
854
+ buf.seek(0)
855
+ plt.close(fig)
856
+
857
+ return Response(content=buf.getvalue(), media_type="image/png")
858
+
859
+ except Exception as e:
860
+ logger.error(f"Failed to generate preview for {dataset_id}: {str(e)}")
861
+ raise HTTPException(status_code=500, detail=str(e))
862
+
863
+
864
+ @router.get("/preview3d")
865
+ async def get_dataset_preview_3d(
866
+ dataset_id: str = Query(..., description="ID of the dataset (without extension)"),
867
+ time_idx: int = Query(0, description="Time step index"),
868
+ ):
869
+ """
870
+ Returns downsampled 3D vector field JSON for Three.js rendering.
871
+ Designed for IC/BC datasets with u,v velocity fields.
872
+ """
873
+ import xarray as xr
874
+ import numpy as np
875
+ import json
876
+
877
+ cache_dir = Path(settings.cache_dir()).expanduser()
878
+ full_path = os.path.join(cache_dir, f"{dataset_id}.zarr")
879
+
880
+ if not os.path.exists(full_path):
881
+ raise HTTPException(status_code=404, detail=f"Dataset {dataset_id} not found.")
882
+
883
+ try:
884
+ ds = xr.open_zarr(full_path, consolidated=False, decode_times=True)
885
+ vars_list = list(ds.data_vars)
886
+
887
+ # Detect u/v variable pairs
888
+ u_var = v_var = None
889
+ for pair in [("u", "v"), ("water_u", "water_v"), ("u10", "v10")]:
890
+ if pair[0] in vars_list and pair[1] in vars_list:
891
+ u_var, v_var = pair
892
+ break
893
+
894
+ if u_var is None:
895
+ raise HTTPException(
896
+ status_code=400,
897
+ detail=f"No u,v vector pair found. Variables: {vars_list}",
898
+ )
899
+
900
+ u_da = ds[u_var]
901
+ v_da = ds[v_var]
902
+
903
+ # Detect depth coordinate (s_rho, depth, z)
904
+ depth_name = next(
905
+ (
906
+ n
907
+ for n in ["s_rho", "depth", "z", "level"]
908
+ if n in ds.coords or n in ds.dims
909
+ ),
910
+ None,
911
+ )
912
+
913
+ # Select time slice
914
+ time_name = next((n for n in ["time", "ocean_time", "t"] if n in ds.dims), None)
915
+ if time_name and time_name in u_da.dims:
916
+ t_idx = min(time_idx, u_da.sizes[time_name] - 1)
917
+ u_da = u_da.isel({time_name: t_idx})
918
+ v_da = v_da.isel({time_name: t_idx})
919
+
920
+ # Find the best lon/lat coordinate for each variable.
921
+ # Staggered ROMS grids have per-variable coords (lon_u/lat_u, lon_v/lat_v).
922
+ # We require the coord's dims to be a subset of the variable's dims so we
923
+ # don't accidentally pick up a rho-grid coord for a u-grid variable.
924
+ def _find_coord(da, candidates):
925
+ for n in candidates:
926
+ if n in ds.variables and set(ds[n].dims).issubset(set(da.dims)):
927
+ return n
928
+ return None
929
+
930
+ u_lon_name = _find_coord(
931
+ u_da, [f"lon_{u_var}", "longitude", "lon", "lon_rho", "x", "xi"]
932
+ )
933
+ u_lat_name = _find_coord(
934
+ u_da, [f"lat_{u_var}", "latitude", "lat", "lat_rho", "y", "eta"]
935
+ )
936
+
937
+ # Squeeze out broadcast/extra dims: keep only depth + the spatial dims that
938
+ # belong to the variable's own coordinate (handles the rho-dim bleed-in from
939
+ # ROMS curvilinear datasets).
940
+ def _squeeze_broadcast(da, lon_cname, lat_cname):
941
+ if lon_cname is None:
942
+ return da
943
+ spatial = set(ds[lon_cname].dims)
944
+ if lat_cname and lat_cname in ds.variables:
945
+ spatial |= set(ds[lat_cname].dims)
946
+ keep = spatial | ({depth_name} if depth_name else set())
947
+ extra = [d for d in da.dims if d not in keep]
948
+ # Mean across broadcast dims (skipna so masked corners don't pollute)
949
+ return da.mean(dim=extra, skipna=True) if extra else da
950
+
951
+ u_da = _squeeze_broadcast(u_da, u_lon_name, u_lat_name)
952
+ v_lon_name = _find_coord(
953
+ v_da, [f"lon_{v_var}", "longitude", "lon", "lon_rho", "x", "xi"]
954
+ )
955
+ v_lat_name = _find_coord(
956
+ v_da, [f"lat_{v_var}", "latitude", "lat", "lat_rho", "y", "eta"]
957
+ )
958
+ v_da = _squeeze_broadcast(v_da, v_lon_name, v_lat_name)
959
+
960
+ # Get coordinate arrays (2D for curvilinear grids)
961
+ lons = (
962
+ ds[u_lon_name].values
963
+ if u_lon_name
964
+ else np.arange(u_da.shape[-1], dtype=float)
965
+ )
966
+ lats = (
967
+ ds[u_lat_name].values
968
+ if u_lat_name
969
+ else np.arange(u_da.shape[-2] if u_da.ndim >= 2 else 1, dtype=float)
970
+ )
971
+ depths = ds[depth_name].values if depth_name else np.array([0.0])
972
+
973
+ lons_2d = lons.ndim == 2
974
+ lats_2d = lats.ndim == 2
975
+
976
+ u_arr = u_da.values
977
+ v_arr = v_da.values
978
+
979
+ # Replace NaN/missing with 0
980
+ u_arr = np.nan_to_num(u_arr, nan=0.0)
981
+ v_arr = np.nan_to_num(v_arr, nan=0.0)
982
+
983
+ def _lon(iy: int, ix: int) -> float:
984
+ if lons_2d:
985
+ return float(lons[iy, ix])
986
+ return float(lons[ix]) if ix < len(lons) else float(ix)
987
+
988
+ def _lat(iy: int, ix: int) -> float:
989
+ if lats_2d:
990
+ return float(lats[iy, ix])
991
+ return float(lats[iy]) if iy < len(lats) else float(iy)
992
+
993
+ # u and v may have different spatial shapes on staggered grids - clamp v indices
994
+ v_shape = v_arr.shape
995
+
996
+ def _v(iz_or_none, iy: int, ix: int) -> float:
997
+ if v_arr.ndim == 2:
998
+ return float(v_arr[min(iy, v_shape[0] - 1), min(ix, v_shape[1] - 1)])
999
+ if v_arr.ndim == 3:
1000
+ iz = iz_or_none if iz_or_none is not None else 0
1001
+ return float(
1002
+ v_arr[
1003
+ min(iz, v_shape[0] - 1),
1004
+ min(iy, v_shape[1] - 1),
1005
+ min(ix, v_shape[2] - 1),
1006
+ ]
1007
+ )
1008
+ return 0.0
1009
+
1010
+ vectors = []
1011
+ max_vectors = 5000 # cap for UI performance
1012
+
1013
+ if u_arr.ndim == 1:
1014
+ # 1D time-series (already time-sliced → scalar)
1015
+ vectors.append(
1016
+ {
1017
+ "lon": float(lons[0]) if len(lons) > 0 else 0,
1018
+ "lat": float(lats[0]) if len(lats) > 0 else 0,
1019
+ "depth": 0,
1020
+ "u": float(u_arr),
1021
+ "v": float(v_arr),
1022
+ }
1023
+ )
1024
+ elif u_arr.ndim == 2:
1025
+ # 2D: (lat, lon) or (eta, xi) - curvilinear or rectilinear
1026
+ ny, nx = u_arr.shape
1027
+ step = max(1, int(np.sqrt(ny * nx / max_vectors)))
1028
+ for iy in range(0, ny, step):
1029
+ for ix in range(0, nx, step):
1030
+ u_val = float(u_arr[iy, ix])
1031
+ v_val = _v(None, iy, ix)
1032
+ if np.isnan(u_val) or np.isnan(v_val):
1033
+ continue
1034
+ if abs(u_val) < 1e-10 and abs(v_val) < 1e-10:
1035
+ continue
1036
+ vectors.append(
1037
+ {
1038
+ "lon": _lon(iy, ix),
1039
+ "lat": _lat(iy, ix),
1040
+ "depth": 0,
1041
+ "u": u_val,
1042
+ "v": v_val,
1043
+ }
1044
+ )
1045
+ elif u_arr.ndim == 3:
1046
+ # 3D: (depth, lat, lon) or (s_rho, eta, xi)
1047
+ nz, ny, nx = u_arr.shape
1048
+ z_step = max(1, nz // 8) # ~8 depth slices like hydro viewer
1049
+ xy_step = max(
1050
+ 1, int(np.sqrt(ny * nx / (max_vectors // max(1, nz // z_step))))
1051
+ )
1052
+ for iz in range(0, nz, z_step):
1053
+ d_raw = float(depths[iz]) if iz < len(depths) else float(iz)
1054
+ # Map s_rho (-1 to 0) to physical depth in meters
1055
+ if depth_name == "s_rho" and -1.5 < d_raw < 0.5:
1056
+ depth_frac = float(np.clip(-d_raw, 0, 1))
1057
+ d = d_raw * 50.0 # approximate 50m water column
1058
+ else:
1059
+ # For regular depth coords, compute frac from min/max
1060
+ d = d_raw
1061
+ d_min = float(np.min(depths))
1062
+ d_max = float(np.max(depths))
1063
+ d_range = abs(d_max - d_min) if d_max != d_min else 1.0
1064
+
1065
+ # Surface is usually closest to 0.
1066
+ # If coords are negative (-50 to 0), max is surface.
1067
+ if d_max <= 0:
1068
+ depth_frac = float(np.clip(abs(d - d_max) / d_range, 0, 1))
1069
+ else:
1070
+ depth_frac = float(np.clip(abs(d - d_min) / d_range, 0, 1))
1071
+
1072
+ # Ensure physical plotting depth negative downwards
1073
+ if d_max > 0 and d_min >= 0:
1074
+ d = -d
1075
+
1076
+ for iy in range(0, ny, xy_step):
1077
+ for ix in range(0, nx, xy_step):
1078
+ u_val = float(u_arr[iz, iy, ix])
1079
+ v_val = _v(iz, iy, ix)
1080
+ if np.isnan(u_val) or np.isnan(v_val):
1081
+ continue
1082
+ if abs(u_val) < 1e-10 and abs(v_val) < 1e-10:
1083
+ continue
1084
+ vectors.append(
1085
+ {
1086
+ "lon": _lon(iy, ix),
1087
+ "lat": _lat(iy, ix),
1088
+ "depth": d,
1089
+ "depth_frac": depth_frac,
1090
+ "u": u_val,
1091
+ "v": v_val,
1092
+ }
1093
+ )
1094
+
1095
+ # Compute bounds for Three.js scene setup
1096
+ all_lons = [v["lon"] for v in vectors] if vectors else [0]
1097
+ all_lats = [v["lat"] for v in vectors] if vectors else [0]
1098
+
1099
+ result = {
1100
+ "vectors": vectors,
1101
+ "u_var": u_var,
1102
+ "v_var": v_var,
1103
+ "bounds": [min(all_lons), min(all_lats), max(all_lons), max(all_lats)],
1104
+ "depth_levels": sorted(set(v["depth"] for v in vectors)),
1105
+ "count": len(vectors),
1106
+ }
1107
+
1108
+ return Response(
1109
+ content=json.dumps(result),
1110
+ media_type="application/json",
1111
+ )
1112
+
1113
+ except HTTPException:
1114
+ raise
1115
+ except Exception as e:
1116
+ logger.error(f"Failed to generate 3D preview for {dataset_id}: {str(e)}")
1117
+ raise HTTPException(status_code=500, detail=str(e))
1118
+
1119
+
1120
+ @router.get("/point_value")
1121
+ async def get_point_value(
1122
+ dataset_id: str = Query(..., description="ID of the dataset (without extension)"),
1123
+ ext: str = Query(..., description="Extension of dataset (zarr, grib, nc)"),
1124
+ var_name: str = Query(..., description="Variable to query"),
1125
+ time_idx: int = Query(0, description="Time step index relative to data slice"),
1126
+ nx: float = Query(
1127
+ ..., ge=0.0, le=1.0, description="Normalized X coordinate (0.0 to 1.0)"
1128
+ ),
1129
+ ny: float = Query(
1130
+ ..., ge=0.0, le=1.0, description="Normalized Y coordinate (0.0 to 1.0)"
1131
+ ),
1132
+ ):
1133
+ """Returns a specific value for a dataset given a normalized X/Y hover coordinate."""
1134
+ import xarray as xr
1135
+ import numpy as np
1136
+
1137
+ cache_dir = Path(settings.cache_dir()).expanduser()
1138
+ full_path = os.path.join(
1139
+ cache_dir, f"{dataset_id}.zarr" if ext == "zarr" else f"{dataset_id}.{ext}"
1140
+ )
1141
+
1142
+ if not os.path.exists(full_path):
1143
+ raise HTTPException(status_code=404, detail="Dataset not found.")
1144
+
1145
+ try:
1146
+ ds = (
1147
+ xr.open_zarr(full_path)
1148
+ if ext == "zarr"
1149
+ else xr.open_dataset(
1150
+ full_path, engine="cfgrib" if ext == "grib" else "netcdf4"
1151
+ )
1152
+ )
1153
+ if var_name not in ds.variables:
1154
+ return {"value": None}
1155
+
1156
+ data_var = ds[var_name]
1157
+
1158
+ slice_opts = {}
1159
+ for time_cand in ["time", "valid_time", "step", "ocean_time"]:
1160
+ if time_cand in data_var.dims:
1161
+ max_t = ds.sizes[time_cand] - 1
1162
+ slice_opts[time_cand] = min(time_idx, max_t)
1163
+
1164
+ for z_name in [
1165
+ "depth",
1166
+ "zC",
1167
+ "level",
1168
+ "s_rho",
1169
+ "s_w",
1170
+ "isobaricInhPa",
1171
+ "number",
1172
+ "surface",
1173
+ "step",
1174
+ ]:
1175
+ if z_name in data_var.dims and z_name not in slice_opts:
1176
+ slice_opts[z_name] = -1 if z_name in ("s_rho", "s_w") else 0
1177
+
1178
+ valid_opts = {k: v for k, v in slice_opts.items() if k in data_var.dims}
1179
+ array = data_var.isel(valid_opts).values
1180
+
1181
+ # Image is rendered origin="lower", so Y goes from bottom (0) to top (1) in plotting space.
1182
+ # nx, ny from frontend will correspond to image dimensions.
1183
+ # Make sure we don't go out of bounds:
1184
+ h, w = array.shape
1185
+ xi = int(np.clip(nx * w, 0, w - 1))
1186
+ # Note: In matplotlib with origin="lower", ny=0 means the lowest array index
1187
+ yi = int(np.clip(ny * h, 0, h - 1))
1188
+
1189
+ val = array[yi, xi]
1190
+
1191
+ return {
1192
+ "value": float(val) if not np.isnan(val) else None,
1193
+ "xi": xi,
1194
+ "yi": yi,
1195
+ "shape": [h, w],
1196
+ }
1197
+ except Exception as e:
1198
+ logger.error(f"Failed to query point value: {str(e)}")
1199
+ raise HTTPException(status_code=500, detail=str(e))