use std::cell::RefCell;
use std::collections::HashMap;
use std::sync::Arc;
use serde::{Deserialize, Serialize};
use pluot_core::{maybe_timeout, FutureExt, Duration};
use pluot_core::log;
use pluot_core::wgpu;
use zarrs::storage::AsyncReadableStorageTraits;
use pluot_core::render_traits::{
DrawToRasterGpu, DrawToRasterCpu, DrawToSvg, MarginParams, PickableLayer, PreparedLayer, ViewParams, resolve_store_name,
};
use pluot_core::two::svg::SvgContext;
use pluot_core::multiscale_utils::{
ResolutionLevel, VisibleTile, get_visible_tiles, pick_visible_tile, select_resolution_level,
};
use pluot_core::render_types::{CpuContext, CpuRenderPass, PrepareResult};
use pluot_core::render_types::GpuContext;
use pluot_core::LayerPickingResult;
use pluot_core::viewport::{DataCoord, ScreenCoord};
use ome_zarr_metadata::v0_5::{
OmeFields, CoordinateTransform, CoordinateTransformScale,
Axis, AxisType, AxisUnit, AxisUnitSpace,
};
use crate::layers::ome_zarr_bitmap_layer::{OmeZarrBitmapLayer, OmeZarrBitmapLayerParams};
use crate::layers::ome_zarr_utils::{
OmeZarrChannelSetting, OmeDim, OmeDimensionOrder,
PhysicalRect, rects_overlap, bounding_box,
axis_unit_space_to_coefficient_and_exponent,
};
#[derive(Serialize, Deserialize, Debug, Clone)]
#[serde(default)]
pub struct OmeZarrBitmapMultiscaleLayerParams {
pub layer_id: String,
pub bounds: Option<MarginParams>,
pub store_name: Option<String>,
pub group_path: Option<String>,
pub multiscale_index: Option<usize>,
pub target_z: Option<u32>,
pub target_t: Option<u32>,
pub channel_settings: Vec<OmeZarrChannelSetting>,
pub opacity: f32,
}
impl Default for OmeZarrBitmapMultiscaleLayerParams {
fn default() -> Self {
Self {
layer_id: "".to_string(),
bounds: None,
store_name: None,
group_path: None,
multiscale_index: None,
target_z: None,
target_t: None,
channel_settings: vec![],
opacity: 1.0,
}
}
}
thread_local! {
static USE_MEMO_CACHE_MULTISCALE_METADATA: RefCell<Option<HashMap<Vec<String>, Arc<OmeZarrMultiscaleMetadata>>>> = const { RefCell::new(None) };
}
async fn use_memo_multiscale_metadata(
initializer: impl AsyncFnOnce() -> OmeZarrMultiscaleMetadata,
keys: &[String],
cache_enabled: bool,
) -> Arc<OmeZarrMultiscaleMetadata> {
if !cache_enabled {
return Arc::new(initializer().await);
}
let data_exists = USE_MEMO_CACHE_MULTISCALE_METADATA.with(|map| {
map.borrow()
.as_ref()
.and_then(|m| m.get(keys).cloned())
});
if let Some(data) = data_exists {
return data;
}
let data = Arc::new(initializer().await);
USE_MEMO_CACHE_MULTISCALE_METADATA.with(|map| {
let mut map_ref = map.borrow_mut();
if map_ref.is_none() {
*map_ref = Some(HashMap::new());
}
map_ref.as_mut().unwrap().insert(keys.to_vec(), data.clone());
});
data
}
struct OmeZarrMultiscaleMetadata {
resolution_levels: Vec<ResolutionLevel>,
dataset_paths: Vec<String>,
full_shapes: Vec<Vec<u64>>,
chunk_shapes: Vec<Vec<u64>>,
array_metadatas: Vec<zarrs::array::ArrayMetadata>,
dimension_order: OmeDimensionOrder,
}
struct LevelSublayers {
level_idx: usize,
sublayers: Vec<OmeZarrBitmapLayer>,
tiles: Vec<VisibleTile>,
tile_rects: Vec<PhysicalRect>,
model_matrix: [f32; 16],
prepare_results: Vec<PrepareResult>,
}
pub struct OmeZarrBitmapMultiscaleLayer {
view_params: ViewParams,
layer_params: OmeZarrBitmapMultiscaleLayerParams,
store: Arc<dyn AsyncReadableStorageTraits>,
store_name: String,
metadata: Option<Arc<OmeZarrMultiscaleMetadata>>,
level_sublayers: Vec<LevelSublayers>,
}
impl OmeZarrBitmapMultiscaleLayer {
pub fn new(view_params: ViewParams, layer_params: OmeZarrBitmapMultiscaleLayerParams) -> Self {
let store_name = resolve_store_name(&layer_params.store_name, &view_params);
let store = view_params.get_store(&store_name);
Self {
view_params,
layer_params,
store,
store_name,
metadata: None,
level_sublayers: Vec::new(),
}
}
async fn load_metadata(&self) -> Arc<OmeZarrMultiscaleMetadata> {
let store = self.store.clone();
let group_path = self
.layer_params
.group_path
.clone()
.unwrap_or_else(|| "/".to_string());
let multiscale_index = self.layer_params.multiscale_index.unwrap_or(0);
let cache_enabled = self.view_params.cache_enabled;
let keys = vec![
self.store_name.clone(),
group_path.clone(),
format!("multiscale_{}", multiscale_index),
];
let metadata_future = use_memo_multiscale_metadata(async || {
let group = zarrs::group::Group::async_open(store.clone(), &group_path)
.await
.expect("Failed to open zarr group for OME-Zarr metadata");
let attrs = group.attributes();
let ome_fields: OmeFields =
serde_json::from_value(attrs.get("ome").expect("OME attribute missing").clone())
.expect("Failed to parse OME attributes");
let multiscales = ome_fields
.multiscales
.expect("Expected OME-NGFF multiscales metadata");
let multiscale = &multiscales[multiscale_index];
let dimension_order_str: String = multiscale.axes.iter()
.map(|a| a.name.chars().next().unwrap_or('?'))
.collect();
let dimension_order = OmeDimensionOrder::try_from(dimension_order_str.as_str())
.unwrap_or_else(|e| panic!("Invalid OME-Zarr dimension order '{}': {}", dimension_order_str, e));
let x_dim_i = dimension_order.index_of(OmeDim::X).unwrap();
let y_dim_i = dimension_order.index_of(OmeDim::Y).unwrap();
let (x_unit_coeff, x_unit_exp) = match &multiscale.axes[x_dim_i].unit {
Some(AxisUnit::Space(unit)) => axis_unit_space_to_coefficient_and_exponent(unit),
None => (1.0, -6), _ => panic!("Expected space unit for X axis, got non-space unit: {:?}", multiscale.axes[x_dim_i].unit),
};
let (y_unit_coeff, y_unit_exp) = match &multiscale.axes[y_dim_i].unit {
Some(AxisUnit::Space(unit)) => axis_unit_space_to_coefficient_and_exponent(unit),
None => (1.0, -6), _ => panic!("Expected space unit for Y axis, got non-space unit: {:?}", multiscale.axes[y_dim_i].unit),
};
let mut resolution_levels = Vec::new();
let mut dataset_paths = Vec::new();
let mut full_shapes = Vec::new();
let mut chunk_shapes = Vec::new();
let mut array_metadatas = Vec::new();
for dataset in &multiscale.datasets {
let array_path = if group_path == "/" {
format!("/{}", dataset.path)
} else {
format!("{}/{}", group_path, dataset.path)
};
let array = zarrs::array::Array::async_open(store.clone(), &array_path)
.await
.unwrap_or_else(|e| panic!("Failed to open array at {}: {:?}", array_path, e));
let array_metadata = array.metadata().clone();
let shape = array.shape().to_vec();
let img_h = shape[y_dim_i];
let img_w = shape[x_dim_i];
let ndim = shape.len();
let origin = vec![0u64; ndim];
let chunk_shape_vec = array.chunk_shape(&origin)
.expect("Failed to get chunk shape for origin chunk");
let chunk_h = chunk_shape_vec[y_dim_i].get();
let chunk_w = chunk_shape_vec[x_dim_i].get();
let mut scale_x: f64 = 1.0;
let mut scale_y: f64 = 1.0;
for transform in &dataset.coordinate_transformations {
if let CoordinateTransform::Scale(CoordinateTransformScale::List { scale }) = transform {
scale_x = scale[x_dim_i] as f64;
scale_y = scale[y_dim_i] as f64;
}
}
let full_chunk_shape: Vec<u64> = chunk_shape_vec.iter().map(|s| s.get()).collect();
let scale_x_in_meters = scale_x * x_unit_coeff * 10_f64.powi(x_unit_exp);
let scale_y_in_meters = scale_y * y_unit_coeff * 10_f64.powi(y_unit_exp);
resolution_levels.push(ResolutionLevel {
shape: [img_h as u32, img_w as u32],
chunk_shape: [chunk_h as u32, chunk_w as u32],
scale: [scale_y_in_meters, scale_x_in_meters],
});
dataset_paths.push(array_path);
full_shapes.push(shape);
chunk_shapes.push(full_chunk_shape);
array_metadatas.push(array_metadata);
}
OmeZarrMultiscaleMetadata {
resolution_levels,
dataset_paths,
full_shapes,
chunk_shapes,
array_metadatas,
dimension_order,
}
}, &keys, cache_enabled).await;
return metadata_future;
}
fn build_sublayers(
&self,
metadata: &OmeZarrMultiscaleMetadata,
) -> Vec<LevelSublayers> {
let target_level = select_resolution_level(
&self.view_params,
&metadata.resolution_levels,
);
let num_levels = metadata.resolution_levels.len();
let target_z = self.layer_params.target_z.map(|v| v as u64);
let target_t = self.layer_params.target_t.map(|v| v as u64);
let mut all_level_sublayers = Vec::new();
let coarsest_idx = num_levels - 1;
let x_dim_i = metadata.dimension_order.index_of(OmeDim::X).unwrap();
let y_dim_i = metadata.dimension_order.index_of(OmeDim::Y).unwrap();
for level_idx in (target_level..=coarsest_idx).rev() {
let level = &metadata.resolution_levels[level_idx];
let scale_x = level.scale[1] as f32;
let scale_y = level.scale[0] as f32;
let model_matrix: [f32; 16] = [
scale_x, 0.0, 0.0, 0.0,
0.0, scale_y, 0.0, 0.0,
0.0, 0.0, 1.0, 0.0,
0.0, 0.0, 0.0, 1.0,
];
let tiles = get_visible_tiles(&self.view_params, level, Some(&model_matrix));
if tiles.is_empty() {
continue;
}
let dataset_path = &metadata.dataset_paths[level_idx];
let full_shape = &metadata.full_shapes[level_idx];
let chunk_shape = &metadata.chunk_shapes[level_idx];
let array_metadata = &metadata.array_metadatas[level_idx];
let mut sublayers = Vec::new();
let mut tile_rects = Vec::new();
for tile in &tiles {
sublayers.push(OmeZarrBitmapLayer::new(
self.view_params.clone(),
OmeZarrBitmapLayerParams {
store_name: Some(self.store_name.clone()),
array_path: dataset_path.clone(),
array_metadata: Some(array_metadata.clone()),
array_shape: full_shape.clone(),
array_chunk_shape: chunk_shape.clone(),
array_dimension_order: metadata.dimension_order.clone(),
target_z,
target_t,
model_matrix,
slice_x: Some((tile.tile_x_start, tile.tile_x_end)),
slice_y: Some((tile.tile_y_start, tile.tile_y_end)),
channel_settings: self.layer_params.channel_settings.clone(),
layer_id: format!(
"{}_level{}_tile_{}_{}",
self.layer_params.layer_id, level_idx, tile.row, tile.col
),
bounds: self.layer_params.bounds.clone(),
opacity: self.layer_params.opacity,
},
));
tile_rects.push(PhysicalRect {
x0: tile.phys_x0,
y0: tile.phys_y0,
x1: tile.phys_x1,
y1: tile.phys_y1,
});
}
all_level_sublayers.push(LevelSublayers {
level_idx,
sublayers,
tiles,
tile_rects,
model_matrix,
prepare_results: Vec::new(),
});
}
all_level_sublayers
}
fn is_tile_occluded(&self, coarse_rect: &PhysicalRect, from_group_idx: usize) -> bool {
for finer_group in &self.level_sublayers[from_group_idx..] {
let all_ready = finer_group
.prepare_results
.iter()
.all(|r| !r.bailed_early);
if !all_ready {
continue;
}
let overlapping_rects: Vec<&PhysicalRect> = finer_group
.tile_rects
.iter()
.enumerate()
.filter(|(i, rect)| {
!finer_group.prepare_results[*i].bailed_early
&& rects_overlap(coarse_rect, rect)
})
.map(|(_, rect)| rect)
.collect();
if overlapping_rects.is_empty() {
continue;
}
let union = bounding_box(&overlapping_rects);
if union.contains(coarse_rect) {
return true;
}
}
false
}
}
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
impl PreparedLayer for OmeZarrBitmapMultiscaleLayer {
async fn prepare(&mut self, _gpu_context: Option<&GpuContext<'_>>) -> PrepareResult {
let metadata_future = self.load_metadata();
let future_result = maybe_timeout!(metadata_future, self.view_params.timeout)
.await;
let metadata = match future_result {
Ok(metadata_result) => metadata_result,
Err(_) => {
return PrepareResult { bailed_early: true };
}
};
self.metadata = Some(metadata.clone());
let metadata = metadata.as_ref();
self.level_sublayers = self.build_sublayers(metadata);
let level_futures = self.level_sublayers.iter_mut().map(|level_group| async {
let futures = level_group.sublayers.iter_mut().map(|sublayer| sublayer.prepare(None));
let results = futures::future::join_all(futures).await;
let group_bailed = results.iter().any(|r| r.bailed_early);
level_group.prepare_results = results;
group_bailed
});
let level_results_future = futures::future::join_all(level_futures);
let level_results_result = maybe_timeout!(level_results_future, self.view_params.timeout)
.await;
match level_results_result {
Ok(level_results_vec) => {
let any_bailed = level_results_vec.into_iter().any(|b| b);
return PrepareResult { bailed_early: any_bailed };
},
Err(_) => {
return PrepareResult { bailed_early: true };
}
}
}
}
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
impl DrawToRasterGpu for OmeZarrBitmapMultiscaleLayer {
async fn draw(&self, gpu_context: &GpuContext<'_>, pass: &mut wgpu::RenderPass) {
let num_groups = self.level_sublayers.len();
for (group_i, level_group) in self.level_sublayers.iter().enumerate() {
let is_finest = group_i == num_groups - 1;
for (tile_i, sublayer) in level_group.sublayers.iter().enumerate() {
let should_draw = if is_finest {
true
} else {
let coarse_rect = &level_group.tile_rects[tile_i];
!self.is_tile_occluded(coarse_rect, group_i + 1)
};
if should_draw {
DrawToRasterGpu::draw(sublayer, gpu_context, pass).await;
}
}
}
}
}
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
impl DrawToRasterCpu for OmeZarrBitmapMultiscaleLayer {
async fn draw(&self, _cpu_context: &CpuContext<'_>, _pass: &mut CpuRenderPass) {}
}
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
impl DrawToSvg for OmeZarrBitmapMultiscaleLayer {
async fn draw(&self, ctx: &mut SvgContext) {
let num_groups = self.level_sublayers.len();
for (group_i, level_group) in self.level_sublayers.iter().enumerate() {
let is_finest = group_i == num_groups - 1;
for (tile_i, sublayer) in level_group.sublayers.iter().enumerate() {
let should_draw = if is_finest {
true
} else {
let coarse_rect = &level_group.tile_rects[tile_i];
!self.is_tile_occluded(coarse_rect, group_i + 1)
};
if should_draw {
DrawToSvg::draw(sublayer, ctx).await;
}
}
}
}
}
impl PickableLayer for OmeZarrBitmapMultiscaleLayer {
fn pick(&self, screen_coord: ScreenCoord, data_coord: Option<DataCoord>) -> Option<LayerPickingResult> {
let DataCoord::TwoD { x: cx, y: cy } = data_coord? else {
return None;
};
let metadata = self.metadata.as_ref()?;
for level_group in self.level_sublayers.iter().rev() {
let level = &metadata.resolution_levels[level_group.level_idx];
let Some(tile_i) = pick_visible_tile(
cx as f64,
cy as f64,
level,
&level_group.tiles,
Some(&level_group.model_matrix),
) else {
continue;
};
let ready = level_group
.prepare_results
.get(tile_i)
.is_some_and(|r| !r.bailed_early);
if !ready {
continue;
}
return PickableLayer::pick(
&level_group.sublayers[tile_i],
screen_coord,
data_coord,
);
}
None
}
}