Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions spatialdata/_core/_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

import numpy as np
import pyarrow as pa

from spatialdata._core.coordinate_system import CoordinateSystem, _get_spatial_axes

if TYPE_CHECKING:
pass


@dataclass(frozen=True)
class BaseSpatialRequest:
"""Base class for spatial queries."""

coordinate_system: CoordinateSystem

def __post_init__(self) -> None:
# validate the coordinate system
spatial_axes = _get_spatial_axes(self.coordinate_system)
if len(spatial_axes) == 0:
raise ValueError("No spatial axes in the requested coordinate system")


@dataclass(frozen=True)
class BoundingBoxRequest(BaseSpatialRequest):
"""Query with an axis-aligned bounding box.

Attributes
----------
coordinate_system : CoordinateSystem
The coordinate system the coordinates are expressed in.
min_coordinate : np.ndarray
The coordinate of the lower left hand corner (i.e., minimum values)
of the bounding box.
max_coordiate : np.ndarray
The coordinate of the upper right hand corner (i.e., maximum values)
of the bounding box
"""

min_coordinate: np.ndarray # type: ignore[type-arg]
max_coordinate: np.ndarray # type: ignore[type-arg]


def _bounding_box_query_points(points: pa.Table, request: BoundingBoxRequest) -> pa.Table:
"""Perform a spatial bounding box query on a points element.

Parameters
----------
points : pa.Table
The points element to perform the query on.
request : BoundingBoxRequest
The request for the query.

Returns
-------
query_result : pa.Table
The points contained within the specified bounding box.
"""
spatial_axes = _get_spatial_axes(request.coordinate_system)

for axis_index, axis_name in enumerate(spatial_axes):
# filter by lower bound
min_value = request.min_coordinate[axis_index]
points = points.filter(pa.compute.greater(points[axis_name], min_value))

# filter by upper bound
max_value = request.max_coordinate[axis_index]
points = points.filter(pa.compute.less(points[axis_name], max_value))

return points


def _bounding_box_query_points_dict(
points_dict: dict[str, pa.Table], request: BoundingBoxRequest
) -> dict[str, pa.Table]:
requested_points = {}
for points_name, points_data in points_dict.items():
points = _bounding_box_query_points(points_data, request)
if len(points) > 0:
# do not include elements with no data
requested_points[points_name] = points

return requested_points
42 changes: 42 additions & 0 deletions spatialdata/_core/_spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,11 @@
from ome_zarr.types import JSONDict
from spatial_image import SpatialImage

from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points_dict,
)
from spatialdata._core.coordinate_system import CoordinateSystem
from spatialdata._core.core_utils import SpatialElement, get_dims, get_transform
from spatialdata._core.models import (
Expand DownExpand Up@@ -142,6 +147,12 @@ def __init__(
Table_s.validate(table)
self._table = table

self._query = QueryManager(self)

@property
def query(self) -> QueryManager:
return self._query

def _add_image_in_memory(self, name: str, image: Union[SpatialImage, MultiscaleSpatialImage]) -> None:
if name in self._images:
raise ValueError(f"Image {name} already exists in the dataset.")
Expand DownExpand Up@@ -618,3 +629,34 @@ def _gen_elements(self) -> Generator[SpatialElement, None, None]:
for element_type in ["images", "labels", "points", "polygons", "shapes"]:
d = getattr(SpatialData, element_type).fget(self)
yield from d.values()


class QueryManager:
"""Perform queries on SpatialData objects"""

def __init__(self, sdata: SpatialData):
self._sdata = sdata

def bounding_box(self, request: BoundingBoxRequest) -> SpatialData:
"""Perform a bounding box query on the SpatialData object.

Parameters
----------
request : BoundingBoxRequest
The bounding box request.

Returns
-------
requested_sdata : SpatialData
The SpatialData object containing the requested data.
Elements with no valid data are omitted.
"""
requested_points = _bounding_box_query_points_dict(points_dict=self._sdata.points, request=request)

return SpatialData(points=requested_points)

def __call__(self, request: BaseSpatialRequest) -> SpatialData:
if isinstance(request, BoundingBoxRequest):
return self.bounding_box(request)
else:
raise TypeError("unknown request type")
24 changes: 22 additions & 2 deletions spatialdata/_core/coordinate_system.py
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
from __future__ import annotations

import json
from typing import TYPE_CHECKING, Any, Optional, Union

Expand DownExpand Up@@ -43,7 +45,7 @@ def __repr__(self) -> str:
return f"CoordinateSystem('{self.name}', {self._axes})"

@staticmethod
def from_dict(coord_sys: CoordSystem_t) -> "CoordinateSystem":
def from_dict(coord_sys: CoordSystem_t) -> CoordinateSystem:
if "name" not in coord_sys.keys():
raise ValueError("`coordinate_system` MUST have a name.")
if "axes" not in coord_sys.keys():
Expand DownExpand Up@@ -81,7 +83,7 @@ def from_array(self, array: Any) -> None:
raise NotImplementedError()

@staticmethod
def from_json(data: Union[str, bytes]) -> "CoordinateSystem":
def from_json(data: Union[str, bytes]) -> CoordinateSystem:
coord_sys = json.loads(data)
return CoordinateSystem.from_dict(coord_sys)

Expand All@@ -108,3 +110,21 @@ def axes_types(self) -> tuple[str, ...]:

def __hash__(self) -> int:
return hash(frozenset(self.to_dict()))


def _get_spatial_axes(
coordinate_system: CoordinateSystem,
) -> list[str]:
"""Get the names of the spatial axes in a coordinate system.

Parameters
----------
coordinate_system : CoordinateSystem
The coordinate system to get the spatial axes from.

Returns
-------
spatial_axis_names : List[str]
The names of the spatial axes.
"""
return [axis.name for axis in coordinate_system._axes if axis.type == "space"]
76 changes: 76 additions & 0 deletions tests/_core/test_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,76 @@
from dataclasses import FrozenInstanceError

import numpy as np
import pytest

from spatialdata import PointsModel
from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points,
)
from tests._core.conftest import c_cs, cyx_cs, czyx_cs, xy_cs


def _make_points_element():
"""Helper function to make a Points element."""
coordinates = np.array([[10, 10], [20, 20], [20, 30]], dtype=float)
return PointsModel.parse(coordinates)


def test_bounding_box_request_immutable():
"""Test that the bounding box request is immutable."""
request = BoundingBoxRequest(
coordinate_system=cyx_cs, min_coordinate=np.array([0, 0]), max_coordinate=np.array([10, 10])
)
isinstance(request, BaseSpatialRequest)

# fields should be immutable
with pytest.raises(FrozenInstanceError):
request.coordinate_system = czyx_cs
with pytest.raises(FrozenInstanceError):
request.min_coordinate = np.array([5, 5, 5])
with pytest.raises(FrozenInstanceError):
request.max_coordinate = np.array([5, 5, 5])


def test_bounding_box_request_no_spatial_axes():
"""Requests with no spatial axes should raise an error"""
with pytest.raises(ValueError):
_ = BoundingBoxRequest(coordinate_system=c_cs, min_coordinate=np.array([0]), max_coordinate=np.array([10]))


def test_bounding_box_points():
"""test the points bounding box_query"""
points_element = _make_points_element()
original_x = np.array(points_element["x"])
original_y = np.array(points_element["y"])

request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([18, 25]), max_coordinate=np.array([22, 35])
)
points_result = _bounding_box_query_points(points_element, request)
np.testing.assert_allclose(points_result["x"], [20])
np.testing.assert_allclose(points_result["y"], [30])

# result should be valid points element
PointsModel.validate(points_result)

# original element should be unchanged
np.testing.assert_allclose(points_element["x"], original_x)
np.testing.assert_allclose(points_element["y"], original_y)


def test_bounding_box_points_no_points():
"""Points bounding box query with no points in range should
return a points element with length 0.
"""
points_element = _make_points_element()
request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([40, 50]), max_coordinate=np.array([45, 55])
)
points_result = _bounding_box_query_points(points_element, request)
assert len(points_result) == 0

# result should be valid points element
PointsModel.validate(points_result)
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions spatialdata/_core/_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

import numpy as np
import pyarrow as pa

from spatialdata._core.coordinate_system import CoordinateSystem, _get_spatial_axes

if TYPE_CHECKING:
pass


@dataclass(frozen=True)
class BaseSpatialRequest:
"""Base class for spatial queries."""

coordinate_system: CoordinateSystem

def __post_init__(self) -> None:
# validate the coordinate system
spatial_axes = _get_spatial_axes(self.coordinate_system)
if len(spatial_axes) == 0:
raise ValueError("No spatial axes in the requested coordinate system")


@dataclass(frozen=True)
class BoundingBoxRequest(BaseSpatialRequest):
"""Query with an axis-aligned bounding box.

Attributes
----------
coordinate_system : CoordinateSystem
The coordinate system the coordinates are expressed in.
min_coordinate : np.ndarray
The coordinate of the lower left hand corner (i.e., minimum values)
of the bounding box.
max_coordiate : np.ndarray
The coordinate of the upper right hand corner (i.e., maximum values)
of the bounding box
"""

min_coordinate: np.ndarray # type: ignore[type-arg]
max_coordinate: np.ndarray # type: ignore[type-arg]


def _bounding_box_query_points(points: pa.Table, request: BoundingBoxRequest) -> pa.Table:
"""Perform a spatial bounding box query on a points element.

Parameters
----------
points : pa.Table
The points element to perform the query on.
request : BoundingBoxRequest
The request for the query.

Returns
-------
query_result : pa.Table
The points contained within the specified bounding box.
"""
spatial_axes = _get_spatial_axes(request.coordinate_system)

for axis_index, axis_name in enumerate(spatial_axes):
# filter by lower bound
min_value = request.min_coordinate[axis_index]
points = points.filter(pa.compute.greater(points[axis_name], min_value))

# filter by upper bound
max_value = request.max_coordinate[axis_index]
points = points.filter(pa.compute.less(points[axis_name], max_value))

return points


def _bounding_box_query_points_dict(
points_dict: dict[str, pa.Table], request: BoundingBoxRequest
) -> dict[str, pa.Table]:
requested_points = {}
for points_name, points_data in points_dict.items():
points = _bounding_box_query_points(points_data, request)
if len(points) > 0:
# do not include elements with no data
requested_points[points_name] = points

return requested_points
42 changes: 42 additions & 0 deletions spatialdata/_core/_spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,11 @@
from ome_zarr.types import JSONDict
from spatial_image import SpatialImage

from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points_dict,
)
from spatialdata._core.coordinate_system import CoordinateSystem
from spatialdata._core.core_utils import SpatialElement, get_dims, get_transform
from spatialdata._core.models import (
Expand DownExpand Up@@ -142,6 +147,12 @@ def __init__(
Table_s.validate(table)
self._table = table

self._query = QueryManager(self)

@property
def query(self) -> QueryManager:
return self._query

def _add_image_in_memory(self, name: str, image: Union[SpatialImage, MultiscaleSpatialImage]) -> None:
if name in self._images:
raise ValueError(f"Image {name} already exists in the dataset.")
Expand DownExpand Up@@ -618,3 +629,34 @@ def _gen_elements(self) -> Generator[SpatialElement, None, None]:
for element_type in ["images", "labels", "points", "polygons", "shapes"]:
d = getattr(SpatialData, element_type).fget(self)
yield from d.values()


class QueryManager:
"""Perform queries on SpatialData objects"""

def __init__(self, sdata: SpatialData):
self._sdata = sdata

def bounding_box(self, request: BoundingBoxRequest) -> SpatialData:
"""Perform a bounding box query on the SpatialData object.

Parameters
----------
request : BoundingBoxRequest
The bounding box request.

Returns
-------
requested_sdata : SpatialData
The SpatialData object containing the requested data.
Elements with no valid data are omitted.
"""
requested_points = _bounding_box_query_points_dict(points_dict=self._sdata.points, request=request)

return SpatialData(points=requested_points)

def __call__(self, request: BaseSpatialRequest) -> SpatialData:
if isinstance(request, BoundingBoxRequest):
return self.bounding_box(request)
else:
raise TypeError("unknown request type")
24 changes: 22 additions & 2 deletions spatialdata/_core/coordinate_system.py
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
from __future__ import annotations

import json
from typing import TYPE_CHECKING, Any, Optional, Union

Expand DownExpand Up@@ -43,7 +45,7 @@ def __repr__(self) -> str:
return f"CoordinateSystem('{self.name}', {self._axes})"

@staticmethod
def from_dict(coord_sys: CoordSystem_t) -> "CoordinateSystem":
def from_dict(coord_sys: CoordSystem_t) -> CoordinateSystem:
if "name" not in coord_sys.keys():
raise ValueError("`coordinate_system` MUST have a name.")
if "axes" not in coord_sys.keys():
Expand DownExpand Up@@ -81,7 +83,7 @@ def from_array(self, array: Any) -> None:
raise NotImplementedError()

@staticmethod
def from_json(data: Union[str, bytes]) -> "CoordinateSystem":
def from_json(data: Union[str, bytes]) -> CoordinateSystem:
coord_sys = json.loads(data)
return CoordinateSystem.from_dict(coord_sys)

Expand All@@ -108,3 +110,21 @@ def axes_types(self) -> tuple[str, ...]:

def __hash__(self) -> int:
return hash(frozenset(self.to_dict()))


def _get_spatial_axes(
coordinate_system: CoordinateSystem,
) -> list[str]:
"""Get the names of the spatial axes in a coordinate system.

Parameters
----------
coordinate_system : CoordinateSystem
The coordinate system to get the spatial axes from.

Returns
-------
spatial_axis_names : List[str]
The names of the spatial axes.
"""
return [axis.name for axis in coordinate_system._axes if axis.type == "space"]
76 changes: 76 additions & 0 deletions tests/_core/test_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,76 @@
from dataclasses import FrozenInstanceError

import numpy as np
import pytest

from spatialdata import PointsModel
from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points,
)
from tests._core.conftest import c_cs, cyx_cs, czyx_cs, xy_cs


def _make_points_element():
"""Helper function to make a Points element."""
coordinates = np.array([[10, 10], [20, 20], [20, 30]], dtype=float)
return PointsModel.parse(coordinates)


def test_bounding_box_request_immutable():
"""Test that the bounding box request is immutable."""
request = BoundingBoxRequest(
coordinate_system=cyx_cs, min_coordinate=np.array([0, 0]), max_coordinate=np.array([10, 10])
)
isinstance(request, BaseSpatialRequest)

# fields should be immutable
with pytest.raises(FrozenInstanceError):
request.coordinate_system = czyx_cs
with pytest.raises(FrozenInstanceError):
request.min_coordinate = np.array([5, 5, 5])
with pytest.raises(FrozenInstanceError):
request.max_coordinate = np.array([5, 5, 5])


def test_bounding_box_request_no_spatial_axes():
"""Requests with no spatial axes should raise an error"""
with pytest.raises(ValueError):
_ = BoundingBoxRequest(coordinate_system=c_cs, min_coordinate=np.array([0]), max_coordinate=np.array([10]))


def test_bounding_box_points():
"""test the points bounding box_query"""
points_element = _make_points_element()
original_x = np.array(points_element["x"])
original_y = np.array(points_element["y"])

request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([18, 25]), max_coordinate=np.array([22, 35])
)
points_result = _bounding_box_query_points(points_element, request)
np.testing.assert_allclose(points_result["x"], [20])
np.testing.assert_allclose(points_result["y"], [30])

# result should be valid points element
PointsModel.validate(points_result)

# original element should be unchanged
np.testing.assert_allclose(points_element["x"], original_x)
np.testing.assert_allclose(points_element["y"], original_y)


def test_bounding_box_points_no_points():
"""Points bounding box query with no points in range should
return a points element with length 0.
"""
points_element = _make_points_element()
request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([40, 50]), max_coordinate=np.array([45, 55])
)
points_result = _bounding_box_query_points(points_element, request)
assert len(points_result) == 0

# result should be valid points element
PointsModel.validate(points_result)
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions spatialdata/_core/_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

import numpy as np
import pyarrow as pa

from spatialdata._core.coordinate_system import CoordinateSystem, _get_spatial_axes

if TYPE_CHECKING:
pass


@dataclass(frozen=True)
class BaseSpatialRequest:
"""Base class for spatial queries."""

coordinate_system: CoordinateSystem

def __post_init__(self) -> None:
# validate the coordinate system
spatial_axes = _get_spatial_axes(self.coordinate_system)
if len(spatial_axes) == 0:
raise ValueError("No spatial axes in the requested coordinate system")


@dataclass(frozen=True)
class BoundingBoxRequest(BaseSpatialRequest):
"""Query with an axis-aligned bounding box.

Attributes
----------
coordinate_system : CoordinateSystem
The coordinate system the coordinates are expressed in.
min_coordinate : np.ndarray
The coordinate of the lower left hand corner (i.e., minimum values)
of the bounding box.
max_coordiate : np.ndarray
The coordinate of the upper right hand corner (i.e., maximum values)
of the bounding box
"""

min_coordinate: np.ndarray # type: ignore[type-arg]
max_coordinate: np.ndarray # type: ignore[type-arg]


def _bounding_box_query_points(points: pa.Table, request: BoundingBoxRequest) -> pa.Table:
"""Perform a spatial bounding box query on a points element.

Parameters
----------
points : pa.Table
The points element to perform the query on.
request : BoundingBoxRequest
The request for the query.

Returns
-------
query_result : pa.Table
The points contained within the specified bounding box.
"""
spatial_axes = _get_spatial_axes(request.coordinate_system)

for axis_index, axis_name in enumerate(spatial_axes):
# filter by lower bound
min_value = request.min_coordinate[axis_index]
points = points.filter(pa.compute.greater(points[axis_name], min_value))

# filter by upper bound
max_value = request.max_coordinate[axis_index]
points = points.filter(pa.compute.less(points[axis_name], max_value))

return points


def _bounding_box_query_points_dict(
points_dict: dict[str, pa.Table], request: BoundingBoxRequest
) -> dict[str, pa.Table]:
requested_points = {}
for points_name, points_data in points_dict.items():
points = _bounding_box_query_points(points_data, request)
if len(points) > 0:
# do not include elements with no data
requested_points[points_name] = points

return requested_points
42 changes: 42 additions & 0 deletions spatialdata/_core/_spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,11 @@
from ome_zarr.types import JSONDict
from spatial_image import SpatialImage

from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points_dict,
)
from spatialdata._core.coordinate_system import CoordinateSystem
from spatialdata._core.core_utils import SpatialElement, get_dims, get_transform
from spatialdata._core.models import (
Expand DownExpand Up@@ -142,6 +147,12 @@ def __init__(
Table_s.validate(table)
self._table = table

self._query = QueryManager(self)

@property
def query(self) -> QueryManager:
return self._query

def _add_image_in_memory(self, name: str, image: Union[SpatialImage, MultiscaleSpatialImage]) -> None:
if name in self._images:
raise ValueError(f"Image {name} already exists in the dataset.")
Expand DownExpand Up@@ -618,3 +629,34 @@ def _gen_elements(self) -> Generator[SpatialElement, None, None]:
for element_type in ["images", "labels", "points", "polygons", "shapes"]:
d = getattr(SpatialData, element_type).fget(self)
yield from d.values()


class QueryManager:
"""Perform queries on SpatialData objects"""

def __init__(self, sdata: SpatialData):
self._sdata = sdata

def bounding_box(self, request: BoundingBoxRequest) -> SpatialData:
"""Perform a bounding box query on the SpatialData object.

Parameters
----------
request : BoundingBoxRequest
The bounding box request.

Returns
-------
requested_sdata : SpatialData
The SpatialData object containing the requested data.
Elements with no valid data are omitted.
"""
requested_points = _bounding_box_query_points_dict(points_dict=self._sdata.points, request=request)

return SpatialData(points=requested_points)

def __call__(self, request: BaseSpatialRequest) -> SpatialData:
if isinstance(request, BoundingBoxRequest):
return self.bounding_box(request)
else:
raise TypeError("unknown request type")
24 changes: 22 additions & 2 deletions spatialdata/_core/coordinate_system.py
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
from __future__ import annotations

import json
from typing import TYPE_CHECKING, Any, Optional, Union

Expand DownExpand Up@@ -43,7 +45,7 @@ def __repr__(self) -> str:
return f"CoordinateSystem('{self.name}', {self._axes})"

@staticmethod
def from_dict(coord_sys: CoordSystem_t) -> "CoordinateSystem":
def from_dict(coord_sys: CoordSystem_t) -> CoordinateSystem:
if "name" not in coord_sys.keys():
raise ValueError("`coordinate_system` MUST have a name.")
if "axes" not in coord_sys.keys():
Expand DownExpand Up@@ -81,7 +83,7 @@ def from_array(self, array: Any) -> None:
raise NotImplementedError()

@staticmethod
def from_json(data: Union[str, bytes]) -> "CoordinateSystem":
def from_json(data: Union[str, bytes]) -> CoordinateSystem:
coord_sys = json.loads(data)
return CoordinateSystem.from_dict(coord_sys)

Expand All@@ -108,3 +110,21 @@ def axes_types(self) -> tuple[str, ...]:

def __hash__(self) -> int:
return hash(frozenset(self.to_dict()))


def _get_spatial_axes(
coordinate_system: CoordinateSystem,
) -> list[str]:
"""Get the names of the spatial axes in a coordinate system.

Parameters
----------
coordinate_system : CoordinateSystem
The coordinate system to get the spatial axes from.

Returns
-------
spatial_axis_names : List[str]
The names of the spatial axes.
"""
return [axis.name for axis in coordinate_system._axes if axis.type == "space"]
76 changes: 76 additions & 0 deletions tests/_core/test_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,76 @@
from dataclasses import FrozenInstanceError

import numpy as np
import pytest

from spatialdata import PointsModel
from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points,
)
from tests._core.conftest import c_cs, cyx_cs, czyx_cs, xy_cs


def _make_points_element():
"""Helper function to make a Points element."""
coordinates = np.array([[10, 10], [20, 20], [20, 30]], dtype=float)
return PointsModel.parse(coordinates)


def test_bounding_box_request_immutable():
"""Test that the bounding box request is immutable."""
request = BoundingBoxRequest(
coordinate_system=cyx_cs, min_coordinate=np.array([0, 0]), max_coordinate=np.array([10, 10])
)
isinstance(request, BaseSpatialRequest)

# fields should be immutable
with pytest.raises(FrozenInstanceError):
request.coordinate_system = czyx_cs
with pytest.raises(FrozenInstanceError):
request.min_coordinate = np.array([5, 5, 5])
with pytest.raises(FrozenInstanceError):
request.max_coordinate = np.array([5, 5, 5])


def test_bounding_box_request_no_spatial_axes():
"""Requests with no spatial axes should raise an error"""
with pytest.raises(ValueError):
_ = BoundingBoxRequest(coordinate_system=c_cs, min_coordinate=np.array([0]), max_coordinate=np.array([10]))


def test_bounding_box_points():
"""test the points bounding box_query"""
points_element = _make_points_element()
original_x = np.array(points_element["x"])
original_y = np.array(points_element["y"])

request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([18, 25]), max_coordinate=np.array([22, 35])
)
points_result = _bounding_box_query_points(points_element, request)
np.testing.assert_allclose(points_result["x"], [20])
np.testing.assert_allclose(points_result["y"], [30])

# result should be valid points element
PointsModel.validate(points_result)

# original element should be unchanged
np.testing.assert_allclose(points_element["x"], original_x)
np.testing.assert_allclose(points_element["y"], original_y)


def test_bounding_box_points_no_points():
"""Points bounding box query with no points in range should
return a points element with length 0.
"""
points_element = _make_points_element()
request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([40, 50]), max_coordinate=np.array([45, 55])
)
points_result = _bounding_box_query_points(points_element, request)
assert len(points_result) == 0

# result should be valid points element
PointsModel.validate(points_result)
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions spatialdata/_core/_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

import numpy as np
import pyarrow as pa

from spatialdata._core.coordinate_system import CoordinateSystem, _get_spatial_axes

if TYPE_CHECKING:
pass


@dataclass(frozen=True)
class BaseSpatialRequest:
"""Base class for spatial queries."""

coordinate_system: CoordinateSystem

def __post_init__(self) -> None:
# validate the coordinate system
spatial_axes = _get_spatial_axes(self.coordinate_system)
if len(spatial_axes) == 0:
raise ValueError("No spatial axes in the requested coordinate system")


@dataclass(frozen=True)
class BoundingBoxRequest(BaseSpatialRequest):
"""Query with an axis-aligned bounding box.

Attributes
----------
coordinate_system : CoordinateSystem
The coordinate system the coordinates are expressed in.
min_coordinate : np.ndarray
The coordinate of the lower left hand corner (i.e., minimum values)
of the bounding box.
max_coordiate : np.ndarray
The coordinate of the upper right hand corner (i.e., maximum values)
of the bounding box
"""

min_coordinate: np.ndarray # type: ignore[type-arg]
max_coordinate: np.ndarray # type: ignore[type-arg]


def _bounding_box_query_points(points: pa.Table, request: BoundingBoxRequest) -> pa.Table:
"""Perform a spatial bounding box query on a points element.

Parameters
----------
points : pa.Table
The points element to perform the query on.
request : BoundingBoxRequest
The request for the query.

Returns
-------
query_result : pa.Table
The points contained within the specified bounding box.
"""
spatial_axes = _get_spatial_axes(request.coordinate_system)

for axis_index, axis_name in enumerate(spatial_axes):
# filter by lower bound
min_value = request.min_coordinate[axis_index]
points = points.filter(pa.compute.greater(points[axis_name], min_value))

# filter by upper bound
max_value = request.max_coordinate[axis_index]
points = points.filter(pa.compute.less(points[axis_name], max_value))

return points


def _bounding_box_query_points_dict(
points_dict: dict[str, pa.Table], request: BoundingBoxRequest
) -> dict[str, pa.Table]:
requested_points = {}
for points_name, points_data in points_dict.items():
points = _bounding_box_query_points(points_data, request)
if len(points) > 0:
# do not include elements with no data
requested_points[points_name] = points

return requested_points
42 changes: 42 additions & 0 deletions spatialdata/_core/_spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,11 @@
from ome_zarr.types import JSONDict
from spatial_image import SpatialImage

from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points_dict,
)
from spatialdata._core.coordinate_system import CoordinateSystem
from spatialdata._core.core_utils import SpatialElement, get_dims, get_transform
from spatialdata._core.models import (
Expand DownExpand Up@@ -142,6 +147,12 @@ def __init__(
Table_s.validate(table)
self._table = table

self._query = QueryManager(self)

@property
def query(self) -> QueryManager:
return self._query

def _add_image_in_memory(self, name: str, image: Union[SpatialImage, MultiscaleSpatialImage]) -> None:
if name in self._images:
raise ValueError(f"Image {name} already exists in the dataset.")
Expand DownExpand Up@@ -618,3 +629,34 @@ def _gen_elements(self) -> Generator[SpatialElement, None, None]:
for element_type in ["images", "labels", "points", "polygons", "shapes"]:
d = getattr(SpatialData, element_type).fget(self)
yield from d.values()


class QueryManager:
"""Perform queries on SpatialData objects"""

def __init__(self, sdata: SpatialData):
self._sdata = sdata

def bounding_box(self, request: BoundingBoxRequest) -> SpatialData:
"""Perform a bounding box query on the SpatialData object.

Parameters
----------
request : BoundingBoxRequest
The bounding box request.

Returns
-------
requested_sdata : SpatialData
The SpatialData object containing the requested data.
Elements with no valid data are omitted.
"""
requested_points = _bounding_box_query_points_dict(points_dict=self._sdata.points, request=request)

return SpatialData(points=requested_points)

def __call__(self, request: BaseSpatialRequest) -> SpatialData:
if isinstance(request, BoundingBoxRequest):
return self.bounding_box(request)
else:
raise TypeError("unknown request type")
24 changes: 22 additions & 2 deletions spatialdata/_core/coordinate_system.py
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
from __future__ import annotations

import json
from typing import TYPE_CHECKING, Any, Optional, Union

Expand DownExpand Up@@ -43,7 +45,7 @@ def __repr__(self) -> str:
return f"CoordinateSystem('{self.name}', {self._axes})"

@staticmethod
def from_dict(coord_sys: CoordSystem_t) -> "CoordinateSystem":
def from_dict(coord_sys: CoordSystem_t) -> CoordinateSystem:
if "name" not in coord_sys.keys():
raise ValueError("`coordinate_system` MUST have a name.")
if "axes" not in coord_sys.keys():
Expand DownExpand Up@@ -81,7 +83,7 @@ def from_array(self, array: Any) -> None:
raise NotImplementedError()

@staticmethod
def from_json(data: Union[str, bytes]) -> "CoordinateSystem":
def from_json(data: Union[str, bytes]) -> CoordinateSystem:
coord_sys = json.loads(data)
return CoordinateSystem.from_dict(coord_sys)

Expand All@@ -108,3 +110,21 @@ def axes_types(self) -> tuple[str, ...]:

def __hash__(self) -> int:
return hash(frozenset(self.to_dict()))


def _get_spatial_axes(
coordinate_system: CoordinateSystem,
) -> list[str]:
"""Get the names of the spatial axes in a coordinate system.

Parameters
----------
coordinate_system : CoordinateSystem
The coordinate system to get the spatial axes from.

Returns
-------
spatial_axis_names : List[str]
The names of the spatial axes.
"""
return [axis.name for axis in coordinate_system._axes if axis.type == "space"]
76 changes: 76 additions & 0 deletions tests/_core/test_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,76 @@
from dataclasses import FrozenInstanceError

import numpy as np
import pytest

from spatialdata import PointsModel
from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points,
)
from tests._core.conftest import c_cs, cyx_cs, czyx_cs, xy_cs


def _make_points_element():
"""Helper function to make a Points element."""
coordinates = np.array([[10, 10], [20, 20], [20, 30]], dtype=float)
return PointsModel.parse(coordinates)


def test_bounding_box_request_immutable():
"""Test that the bounding box request is immutable."""
request = BoundingBoxRequest(
coordinate_system=cyx_cs, min_coordinate=np.array([0, 0]), max_coordinate=np.array([10, 10])
)
isinstance(request, BaseSpatialRequest)

# fields should be immutable
with pytest.raises(FrozenInstanceError):
request.coordinate_system = czyx_cs
with pytest.raises(FrozenInstanceError):
request.min_coordinate = np.array([5, 5, 5])
with pytest.raises(FrozenInstanceError):
request.max_coordinate = np.array([5, 5, 5])


def test_bounding_box_request_no_spatial_axes():
"""Requests with no spatial axes should raise an error"""
with pytest.raises(ValueError):
_ = BoundingBoxRequest(coordinate_system=c_cs, min_coordinate=np.array([0]), max_coordinate=np.array([10]))


def test_bounding_box_points():
"""test the points bounding box_query"""
points_element = _make_points_element()
original_x = np.array(points_element["x"])
original_y = np.array(points_element["y"])

request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([18, 25]), max_coordinate=np.array([22, 35])
)
points_result = _bounding_box_query_points(points_element, request)
np.testing.assert_allclose(points_result["x"], [20])
np.testing.assert_allclose(points_result["y"], [30])

# result should be valid points element
PointsModel.validate(points_result)

# original element should be unchanged
np.testing.assert_allclose(points_element["x"], original_x)
np.testing.assert_allclose(points_element["y"], original_y)


def test_bounding_box_points_no_points():
"""Points bounding box query with no points in range should
return a points element with length 0.
"""
points_element = _make_points_element()
request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([40, 50]), max_coordinate=np.array([45, 55])
)
points_result = _bounding_box_query_points(points_element, request)
assert len(points_result) == 0

# result should be valid points element
PointsModel.validate(points_result)
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions spatialdata/_core/_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

import numpy as np
import pyarrow as pa

from spatialdata._core.coordinate_system import CoordinateSystem, _get_spatial_axes

if TYPE_CHECKING:
pass


@dataclass(frozen=True)
class BaseSpatialRequest:
"""Base class for spatial queries."""

coordinate_system: CoordinateSystem

def __post_init__(self) -> None:
# validate the coordinate system
spatial_axes = _get_spatial_axes(self.coordinate_system)
if len(spatial_axes) == 0:
raise ValueError("No spatial axes in the requested coordinate system")


@dataclass(frozen=True)
class BoundingBoxRequest(BaseSpatialRequest):
"""Query with an axis-aligned bounding box.

Attributes
----------
coordinate_system : CoordinateSystem
The coordinate system the coordinates are expressed in.
min_coordinate : np.ndarray
The coordinate of the lower left hand corner (i.e., minimum values)
of the bounding box.
max_coordiate : np.ndarray
The coordinate of the upper right hand corner (i.e., maximum values)
of the bounding box
"""

min_coordinate: np.ndarray # type: ignore[type-arg]
max_coordinate: np.ndarray # type: ignore[type-arg]


def _bounding_box_query_points(points: pa.Table, request: BoundingBoxRequest) -> pa.Table:
"""Perform a spatial bounding box query on a points element.

Parameters
----------
points : pa.Table
The points element to perform the query on.
request : BoundingBoxRequest
The request for the query.

Returns
-------
query_result : pa.Table
The points contained within the specified bounding box.
"""
spatial_axes = _get_spatial_axes(request.coordinate_system)

for axis_index, axis_name in enumerate(spatial_axes):
# filter by lower bound
min_value = request.min_coordinate[axis_index]
points = points.filter(pa.compute.greater(points[axis_name], min_value))

# filter by upper bound
max_value = request.max_coordinate[axis_index]
points = points.filter(pa.compute.less(points[axis_name], max_value))

return points


def _bounding_box_query_points_dict(
points_dict: dict[str, pa.Table], request: BoundingBoxRequest
) -> dict[str, pa.Table]:
requested_points = {}
for points_name, points_data in points_dict.items():
points = _bounding_box_query_points(points_data, request)
if len(points) > 0:
# do not include elements with no data
requested_points[points_name] = points

return requested_points
42 changes: 42 additions & 0 deletions spatialdata/_core/_spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,11 @@
from ome_zarr.types import JSONDict
from spatial_image import SpatialImage

from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points_dict,
)
from spatialdata._core.coordinate_system import CoordinateSystem
from spatialdata._core.core_utils import SpatialElement, get_dims, get_transform
from spatialdata._core.models import (
Expand DownExpand Up@@ -142,6 +147,12 @@ def __init__(
Table_s.validate(table)
self._table = table

self._query = QueryManager(self)

@property
def query(self) -> QueryManager:
return self._query

def _add_image_in_memory(self, name: str, image: Union[SpatialImage, MultiscaleSpatialImage]) -> None:
if name in self._images:
raise ValueError(f"Image {name} already exists in the dataset.")
Expand DownExpand Up@@ -618,3 +629,34 @@ def _gen_elements(self) -> Generator[SpatialElement, None, None]:
for element_type in ["images", "labels", "points", "polygons", "shapes"]:
d = getattr(SpatialData, element_type).fget(self)
yield from d.values()


class QueryManager:
"""Perform queries on SpatialData objects"""

def __init__(self, sdata: SpatialData):
self._sdata = sdata

def bounding_box(self, request: BoundingBoxRequest) -> SpatialData:
"""Perform a bounding box query on the SpatialData object.

Parameters
----------
request : BoundingBoxRequest
The bounding box request.

Returns
-------
requested_sdata : SpatialData
The SpatialData object containing the requested data.
Elements with no valid data are omitted.
"""
requested_points = _bounding_box_query_points_dict(points_dict=self._sdata.points, request=request)

return SpatialData(points=requested_points)

def __call__(self, request: BaseSpatialRequest) -> SpatialData:
if isinstance(request, BoundingBoxRequest):
return self.bounding_box(request)
else:
raise TypeError("unknown request type")
24 changes: 22 additions & 2 deletions spatialdata/_core/coordinate_system.py
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
from __future__ import annotations

import json
from typing import TYPE_CHECKING, Any, Optional, Union

Expand DownExpand Up@@ -43,7 +45,7 @@ def __repr__(self) -> str:
return f"CoordinateSystem('{self.name}', {self._axes})"

@staticmethod
def from_dict(coord_sys: CoordSystem_t) -> "CoordinateSystem":
def from_dict(coord_sys: CoordSystem_t) -> CoordinateSystem:
if "name" not in coord_sys.keys():
raise ValueError("`coordinate_system` MUST have a name.")
if "axes" not in coord_sys.keys():
Expand DownExpand Up@@ -81,7 +83,7 @@ def from_array(self, array: Any) -> None:
raise NotImplementedError()

@staticmethod
def from_json(data: Union[str, bytes]) -> "CoordinateSystem":
def from_json(data: Union[str, bytes]) -> CoordinateSystem:
coord_sys = json.loads(data)
return CoordinateSystem.from_dict(coord_sys)

Expand All@@ -108,3 +110,21 @@ def axes_types(self) -> tuple[str, ...]:

def __hash__(self) -> int:
return hash(frozenset(self.to_dict()))


def _get_spatial_axes(
coordinate_system: CoordinateSystem,
) -> list[str]:
"""Get the names of the spatial axes in a coordinate system.

Parameters
----------
coordinate_system : CoordinateSystem
The coordinate system to get the spatial axes from.

Returns
-------
spatial_axis_names : List[str]
The names of the spatial axes.
"""
return [axis.name for axis in coordinate_system._axes if axis.type == "space"]
76 changes: 76 additions & 0 deletions tests/_core/test_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,76 @@
from dataclasses import FrozenInstanceError

import numpy as np
import pytest

from spatialdata import PointsModel
from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points,
)
from tests._core.conftest import c_cs, cyx_cs, czyx_cs, xy_cs


def _make_points_element():
"""Helper function to make a Points element."""
coordinates = np.array([[10, 10], [20, 20], [20, 30]], dtype=float)
return PointsModel.parse(coordinates)


def test_bounding_box_request_immutable():
"""Test that the bounding box request is immutable."""
request = BoundingBoxRequest(
coordinate_system=cyx_cs, min_coordinate=np.array([0, 0]), max_coordinate=np.array([10, 10])
)
isinstance(request, BaseSpatialRequest)

# fields should be immutable
with pytest.raises(FrozenInstanceError):
request.coordinate_system = czyx_cs
with pytest.raises(FrozenInstanceError):
request.min_coordinate = np.array([5, 5, 5])
with pytest.raises(FrozenInstanceError):
request.max_coordinate = np.array([5, 5, 5])


def test_bounding_box_request_no_spatial_axes():
"""Requests with no spatial axes should raise an error"""
with pytest.raises(ValueError):
_ = BoundingBoxRequest(coordinate_system=c_cs, min_coordinate=np.array([0]), max_coordinate=np.array([10]))


def test_bounding_box_points():
"""test the points bounding box_query"""
points_element = _make_points_element()
original_x = np.array(points_element["x"])
original_y = np.array(points_element["y"])

request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([18, 25]), max_coordinate=np.array([22, 35])
)
points_result = _bounding_box_query_points(points_element, request)
np.testing.assert_allclose(points_result["x"], [20])
np.testing.assert_allclose(points_result["y"], [30])

# result should be valid points element
PointsModel.validate(points_result)

# original element should be unchanged
np.testing.assert_allclose(points_element["x"], original_x)
np.testing.assert_allclose(points_element["y"], original_y)


def test_bounding_box_points_no_points():
"""Points bounding box query with no points in range should
return a points element with length 0.
"""
points_element = _make_points_element()
request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([40, 50]), max_coordinate=np.array([45, 55])
)
points_result = _bounding_box_query_points(points_element, request)
assert len(points_result) == 0

# result should be valid points element
PointsModel.validate(points_result)
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions spatialdata/_core/_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

import numpy as np
import pyarrow as pa

from spatialdata._core.coordinate_system import CoordinateSystem, _get_spatial_axes

if TYPE_CHECKING:
pass


@dataclass(frozen=True)
class BaseSpatialRequest:
"""Base class for spatial queries."""

coordinate_system: CoordinateSystem

def __post_init__(self) -> None:
# validate the coordinate system
spatial_axes = _get_spatial_axes(self.coordinate_system)
if len(spatial_axes) == 0:
raise ValueError("No spatial axes in the requested coordinate system")


@dataclass(frozen=True)
class BoundingBoxRequest(BaseSpatialRequest):
"""Query with an axis-aligned bounding box.

Attributes
----------
coordinate_system : CoordinateSystem
The coordinate system the coordinates are expressed in.
min_coordinate : np.ndarray
The coordinate of the lower left hand corner (i.e., minimum values)
of the bounding box.
max_coordiate : np.ndarray
The coordinate of the upper right hand corner (i.e., maximum values)
of the bounding box
"""

min_coordinate: np.ndarray # type: ignore[type-arg]
max_coordinate: np.ndarray # type: ignore[type-arg]


def _bounding_box_query_points(points: pa.Table, request: BoundingBoxRequest) -> pa.Table:
"""Perform a spatial bounding box query on a points element.

Parameters
----------
points : pa.Table
The points element to perform the query on.
request : BoundingBoxRequest
The request for the query.

Returns
-------
query_result : pa.Table
The points contained within the specified bounding box.
"""
spatial_axes = _get_spatial_axes(request.coordinate_system)

for axis_index, axis_name in enumerate(spatial_axes):
# filter by lower bound
min_value = request.min_coordinate[axis_index]
points = points.filter(pa.compute.greater(points[axis_name], min_value))

# filter by upper bound
max_value = request.max_coordinate[axis_index]
points = points.filter(pa.compute.less(points[axis_name], max_value))

return points


def _bounding_box_query_points_dict(
points_dict: dict[str, pa.Table], request: BoundingBoxRequest
) -> dict[str, pa.Table]:
requested_points = {}
for points_name, points_data in points_dict.items():
points = _bounding_box_query_points(points_data, request)
if len(points) > 0:
# do not include elements with no data
requested_points[points_name] = points

return requested_points
42 changes: 42 additions & 0 deletions spatialdata/_core/_spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,11 @@
from ome_zarr.types import JSONDict
from spatial_image import SpatialImage

from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points_dict,
)
from spatialdata._core.coordinate_system import CoordinateSystem
from spatialdata._core.core_utils import SpatialElement, get_dims, get_transform
from spatialdata._core.models import (
Expand DownExpand Up@@ -142,6 +147,12 @@ def __init__(
Table_s.validate(table)
self._table = table

self._query = QueryManager(self)

@property
def query(self) -> QueryManager:
return self._query

def _add_image_in_memory(self, name: str, image: Union[SpatialImage, MultiscaleSpatialImage]) -> None:
if name in self._images:
raise ValueError(f"Image {name} already exists in the dataset.")
Expand DownExpand Up@@ -618,3 +629,34 @@ def _gen_elements(self) -> Generator[SpatialElement, None, None]:
for element_type in ["images", "labels", "points", "polygons", "shapes"]:
d = getattr(SpatialData, element_type).fget(self)
yield from d.values()


class QueryManager:
"""Perform queries on SpatialData objects"""

def __init__(self, sdata: SpatialData):
self._sdata = sdata

def bounding_box(self, request: BoundingBoxRequest) -> SpatialData:
"""Perform a bounding box query on the SpatialData object.

Parameters
----------
request : BoundingBoxRequest
The bounding box request.

Returns
-------
requested_sdata : SpatialData
The SpatialData object containing the requested data.
Elements with no valid data are omitted.
"""
requested_points = _bounding_box_query_points_dict(points_dict=self._sdata.points, request=request)

return SpatialData(points=requested_points)

def __call__(self, request: BaseSpatialRequest) -> SpatialData:
if isinstance(request, BoundingBoxRequest):
return self.bounding_box(request)
else:
raise TypeError("unknown request type")
24 changes: 22 additions & 2 deletions spatialdata/_core/coordinate_system.py
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
from __future__ import annotations

import json
from typing import TYPE_CHECKING, Any, Optional, Union

Expand DownExpand Up@@ -43,7 +45,7 @@ def __repr__(self) -> str:
return f"CoordinateSystem('{self.name}', {self._axes})"

@staticmethod
def from_dict(coord_sys: CoordSystem_t) -> "CoordinateSystem":
def from_dict(coord_sys: CoordSystem_t) -> CoordinateSystem:
if "name" not in coord_sys.keys():
raise ValueError("`coordinate_system` MUST have a name.")
if "axes" not in coord_sys.keys():
Expand DownExpand Up@@ -81,7 +83,7 @@ def from_array(self, array: Any) -> None:
raise NotImplementedError()

@staticmethod
def from_json(data: Union[str, bytes]) -> "CoordinateSystem":
def from_json(data: Union[str, bytes]) -> CoordinateSystem:
coord_sys = json.loads(data)
return CoordinateSystem.from_dict(coord_sys)

Expand All@@ -108,3 +110,21 @@ def axes_types(self) -> tuple[str, ...]:

def __hash__(self) -> int:
return hash(frozenset(self.to_dict()))


def _get_spatial_axes(
coordinate_system: CoordinateSystem,
) -> list[str]:
"""Get the names of the spatial axes in a coordinate system.

Parameters
----------
coordinate_system : CoordinateSystem
The coordinate system to get the spatial axes from.

Returns
-------
spatial_axis_names : List[str]
The names of the spatial axes.
"""
return [axis.name for axis in coordinate_system._axes if axis.type == "space"]
76 changes: 76 additions & 0 deletions tests/_core/test_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,76 @@
from dataclasses import FrozenInstanceError

import numpy as np
import pytest

from spatialdata import PointsModel
from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points,
)
from tests._core.conftest import c_cs, cyx_cs, czyx_cs, xy_cs


def _make_points_element():
"""Helper function to make a Points element."""
coordinates = np.array([[10, 10], [20, 20], [20, 30]], dtype=float)
return PointsModel.parse(coordinates)


def test_bounding_box_request_immutable():
"""Test that the bounding box request is immutable."""
request = BoundingBoxRequest(
coordinate_system=cyx_cs, min_coordinate=np.array([0, 0]), max_coordinate=np.array([10, 10])
)
isinstance(request, BaseSpatialRequest)

# fields should be immutable
with pytest.raises(FrozenInstanceError):
request.coordinate_system = czyx_cs
with pytest.raises(FrozenInstanceError):
request.min_coordinate = np.array([5, 5, 5])
with pytest.raises(FrozenInstanceError):
request.max_coordinate = np.array([5, 5, 5])


def test_bounding_box_request_no_spatial_axes():
"""Requests with no spatial axes should raise an error"""
with pytest.raises(ValueError):
_ = BoundingBoxRequest(coordinate_system=c_cs, min_coordinate=np.array([0]), max_coordinate=np.array([10]))


def test_bounding_box_points():
"""test the points bounding box_query"""
points_element = _make_points_element()
original_x = np.array(points_element["x"])
original_y = np.array(points_element["y"])

request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([18, 25]), max_coordinate=np.array([22, 35])
)
points_result = _bounding_box_query_points(points_element, request)
np.testing.assert_allclose(points_result["x"], [20])
np.testing.assert_allclose(points_result["y"], [30])

# result should be valid points element
PointsModel.validate(points_result)

# original element should be unchanged
np.testing.assert_allclose(points_element["x"], original_x)
np.testing.assert_allclose(points_element["y"], original_y)


def test_bounding_box_points_no_points():
"""Points bounding box query with no points in range should
return a points element with length 0.
"""
points_element = _make_points_element()
request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([40, 50]), max_coordinate=np.array([45, 55])
)
points_result = _bounding_box_query_points(points_element, request)
assert len(points_result) == 0

# result should be valid points element
PointsModel.validate(points_result)
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions spatialdata/_core/_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

import numpy as np
import pyarrow as pa

from spatialdata._core.coordinate_system import CoordinateSystem, _get_spatial_axes

if TYPE_CHECKING:
pass


@dataclass(frozen=True)
class BaseSpatialRequest:
"""Base class for spatial queries."""

coordinate_system: CoordinateSystem

def __post_init__(self) -> None:
# validate the coordinate system
spatial_axes = _get_spatial_axes(self.coordinate_system)
if len(spatial_axes) == 0:
raise ValueError("No spatial axes in the requested coordinate system")


@dataclass(frozen=True)
class BoundingBoxRequest(BaseSpatialRequest):
"""Query with an axis-aligned bounding box.

Attributes
----------
coordinate_system : CoordinateSystem
The coordinate system the coordinates are expressed in.
min_coordinate : np.ndarray
The coordinate of the lower left hand corner (i.e., minimum values)
of the bounding box.
max_coordiate : np.ndarray
The coordinate of the upper right hand corner (i.e., maximum values)
of the bounding box
"""

min_coordinate: np.ndarray # type: ignore[type-arg]
max_coordinate: np.ndarray # type: ignore[type-arg]


def _bounding_box_query_points(points: pa.Table, request: BoundingBoxRequest) -> pa.Table:
"""Perform a spatial bounding box query on a points element.

Parameters
----------
points : pa.Table
The points element to perform the query on.
request : BoundingBoxRequest
The request for the query.

Returns
-------
query_result : pa.Table
The points contained within the specified bounding box.
"""
spatial_axes = _get_spatial_axes(request.coordinate_system)

for axis_index, axis_name in enumerate(spatial_axes):
# filter by lower bound
min_value = request.min_coordinate[axis_index]
points = points.filter(pa.compute.greater(points[axis_name], min_value))

# filter by upper bound
max_value = request.max_coordinate[axis_index]
points = points.filter(pa.compute.less(points[axis_name], max_value))

return points


def _bounding_box_query_points_dict(
points_dict: dict[str, pa.Table], request: BoundingBoxRequest
) -> dict[str, pa.Table]:
requested_points = {}
for points_name, points_data in points_dict.items():
points = _bounding_box_query_points(points_data, request)
if len(points) > 0:
# do not include elements with no data
requested_points[points_name] = points

return requested_points
42 changes: 42 additions & 0 deletions spatialdata/_core/_spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,11 @@
from ome_zarr.types import JSONDict
from spatial_image import SpatialImage

from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points_dict,
)
from spatialdata._core.coordinate_system import CoordinateSystem
from spatialdata._core.core_utils import SpatialElement, get_dims, get_transform
from spatialdata._core.models import (
Expand DownExpand Up@@ -142,6 +147,12 @@ def __init__(
Table_s.validate(table)
self._table = table

self._query = QueryManager(self)

@property
def query(self) -> QueryManager:
return self._query

def _add_image_in_memory(self, name: str, image: Union[SpatialImage, MultiscaleSpatialImage]) -> None:
if name in self._images:
raise ValueError(f"Image {name} already exists in the dataset.")
Expand DownExpand Up@@ -618,3 +629,34 @@ def _gen_elements(self) -> Generator[SpatialElement, None, None]:
for element_type in ["images", "labels", "points", "polygons", "shapes"]:
d = getattr(SpatialData, element_type).fget(self)
yield from d.values()


class QueryManager:
"""Perform queries on SpatialData objects"""

def __init__(self, sdata: SpatialData):
self._sdata = sdata

def bounding_box(self, request: BoundingBoxRequest) -> SpatialData:
"""Perform a bounding box query on the SpatialData object.

Parameters
----------
request : BoundingBoxRequest
The bounding box request.

Returns
-------
requested_sdata : SpatialData
The SpatialData object containing the requested data.
Elements with no valid data are omitted.
"""
requested_points = _bounding_box_query_points_dict(points_dict=self._sdata.points, request=request)

return SpatialData(points=requested_points)

def __call__(self, request: BaseSpatialRequest) -> SpatialData:
if isinstance(request, BoundingBoxRequest):
return self.bounding_box(request)
else:
raise TypeError("unknown request type")
24 changes: 22 additions & 2 deletions spatialdata/_core/coordinate_system.py
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
from __future__ import annotations

import json
from typing import TYPE_CHECKING, Any, Optional, Union

Expand DownExpand Up@@ -43,7 +45,7 @@ def __repr__(self) -> str:
return f"CoordinateSystem('{self.name}', {self._axes})"

@staticmethod
def from_dict(coord_sys: CoordSystem_t) -> "CoordinateSystem":
def from_dict(coord_sys: CoordSystem_t) -> CoordinateSystem:
if "name" not in coord_sys.keys():
raise ValueError("`coordinate_system` MUST have a name.")
if "axes" not in coord_sys.keys():
Expand DownExpand Up@@ -81,7 +83,7 @@ def from_array(self, array: Any) -> None:
raise NotImplementedError()

@staticmethod
def from_json(data: Union[str, bytes]) -> "CoordinateSystem":
def from_json(data: Union[str, bytes]) -> CoordinateSystem:
coord_sys = json.loads(data)
return CoordinateSystem.from_dict(coord_sys)

Expand All@@ -108,3 +110,21 @@ def axes_types(self) -> tuple[str, ...]:

def __hash__(self) -> int:
return hash(frozenset(self.to_dict()))


def _get_spatial_axes(
coordinate_system: CoordinateSystem,
) -> list[str]:
"""Get the names of the spatial axes in a coordinate system.

Parameters
----------
coordinate_system : CoordinateSystem
The coordinate system to get the spatial axes from.

Returns
-------
spatial_axis_names : List[str]
The names of the spatial axes.
"""
return [axis.name for axis in coordinate_system._axes if axis.type == "space"]
76 changes: 76 additions & 0 deletions tests/_core/test_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,76 @@
from dataclasses import FrozenInstanceError

import numpy as np
import pytest

from spatialdata import PointsModel
from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points,
)
from tests._core.conftest import c_cs, cyx_cs, czyx_cs, xy_cs


def _make_points_element():
"""Helper function to make a Points element."""
coordinates = np.array([[10, 10], [20, 20], [20, 30]], dtype=float)
return PointsModel.parse(coordinates)


def test_bounding_box_request_immutable():
"""Test that the bounding box request is immutable."""
request = BoundingBoxRequest(
coordinate_system=cyx_cs, min_coordinate=np.array([0, 0]), max_coordinate=np.array([10, 10])
)
isinstance(request, BaseSpatialRequest)

# fields should be immutable
with pytest.raises(FrozenInstanceError):
request.coordinate_system = czyx_cs
with pytest.raises(FrozenInstanceError):
request.min_coordinate = np.array([5, 5, 5])
with pytest.raises(FrozenInstanceError):
request.max_coordinate = np.array([5, 5, 5])


def test_bounding_box_request_no_spatial_axes():
"""Requests with no spatial axes should raise an error"""
with pytest.raises(ValueError):
_ = BoundingBoxRequest(coordinate_system=c_cs, min_coordinate=np.array([0]), max_coordinate=np.array([10]))


def test_bounding_box_points():
"""test the points bounding box_query"""
points_element = _make_points_element()
original_x = np.array(points_element["x"])
original_y = np.array(points_element["y"])

request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([18, 25]), max_coordinate=np.array([22, 35])
)
points_result = _bounding_box_query_points(points_element, request)
np.testing.assert_allclose(points_result["x"], [20])
np.testing.assert_allclose(points_result["y"], [30])

# result should be valid points element
PointsModel.validate(points_result)

# original element should be unchanged
np.testing.assert_allclose(points_element["x"], original_x)
np.testing.assert_allclose(points_element["y"], original_y)


def test_bounding_box_points_no_points():
"""Points bounding box query with no points in range should
return a points element with length 0.
"""
points_element = _make_points_element()
request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([40, 50]), max_coordinate=np.array([45, 55])
)
points_result = _bounding_box_query_points(points_element, request)
assert len(points_result) == 0

# result should be valid points element
PointsModel.validate(points_result)
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions spatialdata/_core/_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

import numpy as np
import pyarrow as pa

from spatialdata._core.coordinate_system import CoordinateSystem, _get_spatial_axes

if TYPE_CHECKING:
pass


@dataclass(frozen=True)
class BaseSpatialRequest:
"""Base class for spatial queries."""

coordinate_system: CoordinateSystem

def __post_init__(self) -> None:
# validate the coordinate system
spatial_axes = _get_spatial_axes(self.coordinate_system)
if len(spatial_axes) == 0:
raise ValueError("No spatial axes in the requested coordinate system")


@dataclass(frozen=True)
class BoundingBoxRequest(BaseSpatialRequest):
"""Query with an axis-aligned bounding box.

Attributes
----------
coordinate_system : CoordinateSystem
The coordinate system the coordinates are expressed in.
min_coordinate : np.ndarray
The coordinate of the lower left hand corner (i.e., minimum values)
of the bounding box.
max_coordiate : np.ndarray
The coordinate of the upper right hand corner (i.e., maximum values)
of the bounding box
"""

min_coordinate: np.ndarray # type: ignore[type-arg]
max_coordinate: np.ndarray # type: ignore[type-arg]


def _bounding_box_query_points(points: pa.Table, request: BoundingBoxRequest) -> pa.Table:
"""Perform a spatial bounding box query on a points element.

Parameters
----------
points : pa.Table
The points element to perform the query on.
request : BoundingBoxRequest
The request for the query.

Returns
-------
query_result : pa.Table
The points contained within the specified bounding box.
"""
spatial_axes = _get_spatial_axes(request.coordinate_system)

for axis_index, axis_name in enumerate(spatial_axes):
# filter by lower bound
min_value = request.min_coordinate[axis_index]
points = points.filter(pa.compute.greater(points[axis_name], min_value))

# filter by upper bound
max_value = request.max_coordinate[axis_index]
points = points.filter(pa.compute.less(points[axis_name], max_value))

return points


def _bounding_box_query_points_dict(
points_dict: dict[str, pa.Table], request: BoundingBoxRequest
) -> dict[str, pa.Table]:
requested_points = {}
for points_name, points_data in points_dict.items():
points = _bounding_box_query_points(points_data, request)
if len(points) > 0:
# do not include elements with no data
requested_points[points_name] = points

return requested_points
42 changes: 42 additions & 0 deletions spatialdata/_core/_spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,11 @@
from ome_zarr.types import JSONDict
from spatial_image import SpatialImage

from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points_dict,
)
from spatialdata._core.coordinate_system import CoordinateSystem
from spatialdata._core.core_utils import SpatialElement, get_dims, get_transform
from spatialdata._core.models import (
Expand DownExpand Up@@ -142,6 +147,12 @@ def __init__(
Table_s.validate(table)
self._table = table

self._query = QueryManager(self)

@property
def query(self) -> QueryManager:
return self._query

def _add_image_in_memory(self, name: str, image: Union[SpatialImage, MultiscaleSpatialImage]) -> None:
if name in self._images:
raise ValueError(f"Image {name} already exists in the dataset.")
Expand DownExpand Up@@ -618,3 +629,34 @@ def _gen_elements(self) -> Generator[SpatialElement, None, None]:
for element_type in ["images", "labels", "points", "polygons", "shapes"]:
d = getattr(SpatialData, element_type).fget(self)
yield from d.values()


class QueryManager:
"""Perform queries on SpatialData objects"""

def __init__(self, sdata: SpatialData):
self._sdata = sdata

def bounding_box(self, request: BoundingBoxRequest) -> SpatialData:
"""Perform a bounding box query on the SpatialData object.

Parameters
----------
request : BoundingBoxRequest
The bounding box request.

Returns
-------
requested_sdata : SpatialData
The SpatialData object containing the requested data.
Elements with no valid data are omitted.
"""
requested_points = _bounding_box_query_points_dict(points_dict=self._sdata.points, request=request)

return SpatialData(points=requested_points)

def __call__(self, request: BaseSpatialRequest) -> SpatialData:
if isinstance(request, BoundingBoxRequest):
return self.bounding_box(request)
else:
raise TypeError("unknown request type")
24 changes: 22 additions & 2 deletions spatialdata/_core/coordinate_system.py
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
from __future__ import annotations

import json
from typing import TYPE_CHECKING, Any, Optional, Union

Expand DownExpand Up@@ -43,7 +45,7 @@ def __repr__(self) -> str:
return f"CoordinateSystem('{self.name}', {self._axes})"

@staticmethod
def from_dict(coord_sys: CoordSystem_t) -> "CoordinateSystem":
def from_dict(coord_sys: CoordSystem_t) -> CoordinateSystem:
if "name" not in coord_sys.keys():
raise ValueError("`coordinate_system` MUST have a name.")
if "axes" not in coord_sys.keys():
Expand DownExpand Up@@ -81,7 +83,7 @@ def from_array(self, array: Any) -> None:
raise NotImplementedError()

@staticmethod
def from_json(data: Union[str, bytes]) -> "CoordinateSystem":
def from_json(data: Union[str, bytes]) -> CoordinateSystem:
coord_sys = json.loads(data)
return CoordinateSystem.from_dict(coord_sys)

Expand All@@ -108,3 +110,21 @@ def axes_types(self) -> tuple[str, ...]:

def __hash__(self) -> int:
return hash(frozenset(self.to_dict()))


def _get_spatial_axes(
coordinate_system: CoordinateSystem,
) -> list[str]:
"""Get the names of the spatial axes in a coordinate system.

Parameters
----------
coordinate_system : CoordinateSystem
The coordinate system to get the spatial axes from.

Returns
-------
spatial_axis_names : List[str]
The names of the spatial axes.
"""
return [axis.name for axis in coordinate_system._axes if axis.type == "space"]
76 changes: 76 additions & 0 deletions tests/_core/test_spatial_query.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,76 @@
from dataclasses import FrozenInstanceError

import numpy as np
import pytest

from spatialdata import PointsModel
from spatialdata._core._spatial_query import (
BaseSpatialRequest,
BoundingBoxRequest,
_bounding_box_query_points,
)
from tests._core.conftest import c_cs, cyx_cs, czyx_cs, xy_cs


def _make_points_element():
"""Helper function to make a Points element."""
coordinates = np.array([[10, 10], [20, 20], [20, 30]], dtype=float)
return PointsModel.parse(coordinates)


def test_bounding_box_request_immutable():
"""Test that the bounding box request is immutable."""
request = BoundingBoxRequest(
coordinate_system=cyx_cs, min_coordinate=np.array([0, 0]), max_coordinate=np.array([10, 10])
)
isinstance(request, BaseSpatialRequest)

# fields should be immutable
with pytest.raises(FrozenInstanceError):
request.coordinate_system = czyx_cs
with pytest.raises(FrozenInstanceError):
request.min_coordinate = np.array([5, 5, 5])
with pytest.raises(FrozenInstanceError):
request.max_coordinate = np.array([5, 5, 5])


def test_bounding_box_request_no_spatial_axes():
"""Requests with no spatial axes should raise an error"""
with pytest.raises(ValueError):
_ = BoundingBoxRequest(coordinate_system=c_cs, min_coordinate=np.array([0]), max_coordinate=np.array([10]))


def test_bounding_box_points():
"""test the points bounding box_query"""
points_element = _make_points_element()
original_x = np.array(points_element["x"])
original_y = np.array(points_element["y"])

request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([18, 25]), max_coordinate=np.array([22, 35])
)
points_result = _bounding_box_query_points(points_element, request)
np.testing.assert_allclose(points_result["x"], [20])
np.testing.assert_allclose(points_result["y"], [30])

# result should be valid points element
PointsModel.validate(points_result)

# original element should be unchanged
np.testing.assert_allclose(points_element["x"], original_x)
np.testing.assert_allclose(points_element["y"], original_y)


def test_bounding_box_points_no_points():
"""Points bounding box query with no points in range should
return a points element with length 0.
"""
points_element = _make_points_element()
request = BoundingBoxRequest(
coordinate_system=xy_cs, min_coordinate=np.array([40, 50]), max_coordinate=np.array([45, 55])
)
points_result = _bounding_box_query_points(points_element, request)
assert len(points_result) == 0

# result should be valid points element
PointsModel.validate(points_result)