From 1a189af21f511e04d9e660aaf9b19f41b809c05e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Thu, 16 Jul 2026 16:49:32 -0400 Subject: [PATCH 01/13] Ensure file is created when it doesnt exist - Change mode=rw to mode=rwc - Add unit test --- pyproject.toml | 5 +++ src/zarr_sqlite/scratch.py | 21 ------------ src/zarr_sqlite/zarr_sqlite.py | 2 +- test/{test_zarr.py => test_sqlitestore.py} | 38 ++++++++++++++++++---- uv.lock | 10 +++++- 5 files changed, 46 insertions(+), 30 deletions(-) delete mode 100644 src/zarr_sqlite/scratch.py rename test/{test_zarr.py => test_sqlitestore.py} (69%) diff --git a/pyproject.toml b/pyproject.toml index d035307..eb842e1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,3 +20,8 @@ test = [ [build-system] requires = ["uv_build>=0.8.3,<0.9.0"] build-backend = "uv_build" + +[dependency-groups] +dev = [ + "pytest>=8.4.1", +] diff --git a/src/zarr_sqlite/scratch.py b/src/zarr_sqlite/scratch.py deleted file mode 100644 index c757910..0000000 --- a/src/zarr_sqlite/scratch.py +++ /dev/null @@ -1,21 +0,0 @@ - -class SQLiteStore: - supports_writes: bool = True - supports_deletes: bool = True - supports_partial_writes: bool = True - supports_listing: bool = True - - root: str - - def __init__(self, root: str) -> None: - self.root = root - -a = SQLiteStore("aaa") -b = SQLiteStore("bbb") - -print(a.root) -print(b.root) - -a.root = "ccc" -print(a.root) -print(b.root) diff --git a/src/zarr_sqlite/zarr_sqlite.py b/src/zarr_sqlite/zarr_sqlite.py index 42e0074..f9d0b54 100644 --- a/src/zarr_sqlite/zarr_sqlite.py +++ b/src/zarr_sqlite/zarr_sqlite.py @@ -97,7 +97,7 @@ def _build_database_uri(database: Path | str, read_only: bool) -> str: Ref: https://sqlite.org/uri.html """ - query = {"mode": ["ro"] if read_only else ["rw"]} + query = {"mode": ["ro"] if read_only else ["rwc"]} uri_path = "" if isinstance(database, Path): diff --git a/test/test_zarr.py b/test/test_sqlitestore.py similarity index 69% rename from test/test_zarr.py rename to test/test_sqlitestore.py index 46fc089..00570e9 100644 --- a/test/test_zarr.py +++ b/test/test_sqlitestore.py @@ -5,9 +5,10 @@ import numpy as np import pytest - +import tempfile import zarr +from pathlib import Path from tempfile import NamedTemporaryFile from zarr_sqlite import SQLiteStore @@ -46,6 +47,7 @@ def test_store_array(sqlite_store): read_back = z[:] assert np.array_equal(read_back, data) + def test_create_group(sqlite_store): root = zarr.create_group(store=sqlite_store) group1 = root.create_group("group1") @@ -55,13 +57,12 @@ def test_create_group(sqlite_store): assert "group1" in root assert isinstance(root["group1"], zarr.Group) + def test_save_array_to_group(temp_db_file): with SQLiteStore(temp_db_file) as sqlite_store: root = zarr.create_group(store=sqlite_store) group1 = root.create_group("group1") - z = group1.create_array( - shape=(100, 100), chunks=(10, 10), dtype="f4", name="z" - ) + z = group1.create_array(shape=(100, 100), chunks=(10, 10), dtype="f4", name="z") data = random_array((100, 100), dtype=np.float32) z[:, :] = data @@ -72,13 +73,12 @@ def test_save_array_to_group(temp_db_file): root = zarr.open_group(store=sqlite_store) assert np.array_equal(data, root["group1/z"][:]) + def test_delete_array(temp_db_file): with SQLiteStore(temp_db_file) as sqlite_store: root = zarr.create_group(store=sqlite_store) group1 = root.create_group("group1") - group1.create_array( - shape=(100, 100), chunks=(10, 10), dtype="f4", name="z" - ) + group1.create_array(shape=(100, 100), chunks=(10, 10), dtype="f4", name="z") with SQLiteStore(temp_db_file) as sqlite_store: root = zarr.open_group(store=sqlite_store) @@ -94,6 +94,7 @@ def test_delete_array(temp_db_file): with pytest.raises(KeyError): root["group1/z"][:] + def test_append_array(sqlite_store): z = zarr.create_array( store=sqlite_store, shape=(100, 100), chunks=(10, 10), dtype="f4" @@ -106,3 +107,26 @@ def test_append_array(sqlite_store): assert z.shape == (200, 100) expected = np.tile(data, (2, 1)) assert np.array_equal(z[:], expected) + + +def test_store_array_creates_file_and_persists(): + with tempfile.TemporaryDirectory() as tmpdir: + db_path = Path(tmpdir) / "test.db" + + assert not db_path.exists(), "Database file should not exist before store creation" + + data = random_array((50, 50), dtype=np.float32) + + with SQLiteStore(db_path) as sqlite_store: + root = zarr.create_group(store=sqlite_store) + z = root.create_array(shape=(50, 50), chunks=(10, 10), dtype="f4", name="array") + z[:, :] = data + read_back = z[:] + assert np.array_equal(read_back, data) + + assert db_path.exists(), "Database file should be created by SQLiteStore" + + with SQLiteStore(db_path) as sqlite_store: + root = zarr.open_group(store=sqlite_store) + assert "array" in root + assert np.array_equal(root["array"][:], data) diff --git a/uv.lock b/uv.lock index a632cbb..5faa84b 100644 --- a/uv.lock +++ b/uv.lock @@ -1,5 +1,5 @@ version = 1 -revision = 2 +revision = 3 requires-python = ">=3.11" [[package]] @@ -364,6 +364,11 @@ test = [ { name = "pytest" }, ] +[package.dev-dependencies] +dev = [ + { name = "pytest" }, +] + [package.metadata] requires-dist = [ { name = "mypy", marker = "extra == 'test'", specifier = ">=1.17.0" }, @@ -371,3 +376,6 @@ requires-dist = [ { name = "zarr", specifier = ">=3.1.0" }, ] provides-extras = ["test"] + +[package.metadata.requires-dev] +dev = [{ name = "pytest", specifier = ">=8.4.1" }] From e95220d8b1cf4f69fa6ba62cd6dbf5c4a79c6fc2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Fri, 17 Jul 2026 10:57:08 -0400 Subject: [PATCH 02/13] Update unit test --- src/zarr_sqlite/zarr_sqlite.py | 15 ++++++++----- test/test_sqlitestore.py | 41 ++++++++++++++++++---------------- 2 files changed, 32 insertions(+), 24 deletions(-) diff --git a/src/zarr_sqlite/zarr_sqlite.py b/src/zarr_sqlite/zarr_sqlite.py index f9d0b54..640f310 100644 --- a/src/zarr_sqlite/zarr_sqlite.py +++ b/src/zarr_sqlite/zarr_sqlite.py @@ -14,13 +14,12 @@ Store, SuffixByteRequest, ) -from zarr.core.buffer import Buffer -from zarr.core.common import BytesLike if TYPE_CHECKING: from collections.abc import AsyncIterator, Iterable, Sequence - from zarr.core.buffer import BufferPrototype + from zarr.core.buffer import Buffer + from zarr.core.common import BytesLike class SQLiteStore(Store): @@ -75,7 +74,7 @@ def __init__( database: str | Path, *, read_only: bool = False, - journal_mode: str | None = 'WAL', + journal_mode: str | None = "WAL", ) -> None: super().__init__(read_only=read_only) self.database_uri = self._build_database_uri(database, read_only=read_only) @@ -144,7 +143,13 @@ async def _open(self) -> None: ) if not self._read_only: if self._journal_mode is not None: - if self._journal_mode not in ["DELETE", "TRUNCATE", "PERSIST", "WAL", "OFF"]: + if self._journal_mode not in [ + "DELETE", + "TRUNCATE", + "PERSIST", + "WAL", + "OFF", + ]: raise ValueError(f"Invalid journal_mode: {self._journal_mode}") self._con.autocommit = True self._con.execute(f"PRAGMA journal_mode={self._journal_mode}") diff --git a/test/test_sqlitestore.py b/test/test_sqlitestore.py index 00570e9..0001f77 100644 --- a/test/test_sqlitestore.py +++ b/test/test_sqlitestore.py @@ -110,23 +110,26 @@ def test_append_array(sqlite_store): def test_store_array_creates_file_and_persists(): + """Ensure file is created automatically when it doesn't exist""" with tempfile.TemporaryDirectory() as tmpdir: - db_path = Path(tmpdir) / "test.db" - - assert not db_path.exists(), "Database file should not exist before store creation" - - data = random_array((50, 50), dtype=np.float32) - - with SQLiteStore(db_path) as sqlite_store: - root = zarr.create_group(store=sqlite_store) - z = root.create_array(shape=(50, 50), chunks=(10, 10), dtype="f4", name="array") - z[:, :] = data - read_back = z[:] - assert np.array_equal(read_back, data) - - assert db_path.exists(), "Database file should be created by SQLiteStore" - - with SQLiteStore(db_path) as sqlite_store: - root = zarr.open_group(store=sqlite_store) - assert "array" in root - assert np.array_equal(root["array"][:], data) + fname = Path(tmpdir) / 'test1.zarrdb' + store = SQLiteStore(fname, read_only=False) + root = zarr.open(store=store, mode='a') + group1 = root.create_group('group1') + + i = np.arange(80000)/60.0 + x = i + np.random.normal(size=len(i), scale=1.5) + y = np.sin(x/13) + 0.7*np.cos(x/47) + np.random.normal(size=len(i), scale=0.05) + x.shape = (400, 200) + y.shape = (200, 400) + + group1.create_array(data=x, name='xdat') + group1.create_array(data=y, name='ydat') + store.close() + + read_root = zarr.open(store=SQLiteStore(fname, read_only=True), mode='r') + assert np.array_equal(read_root["group1/ydat"][:], y) + assert np.array_equal(read_root["group1/xdat"], x) + + # Ensure store is closed so tmpdir can be safely deleted + read_root.store.close() From b809f72a1375a23e5c244fc4331cf9d6c056b0ea Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Sat, 18 Jul 2026 15:03:56 -0400 Subject: [PATCH 03/13] Add numpy as explicit dep --- pyproject.toml | 1 + test/test_sqlitestore.py | 13 ++++--------- uv.lock | 6 +++++- 3 files changed, 10 insertions(+), 10 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index eb842e1..ca1a242 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -23,5 +23,6 @@ build-backend = "uv_build" [dependency-groups] dev = [ + "numpy>=2.3.2", "pytest>=8.4.1", ] diff --git a/test/test_sqlitestore.py b/test/test_sqlitestore.py index 0001f77..865c1b5 100644 --- a/test/test_sqlitestore.py +++ b/test/test_sqlitestore.py @@ -1,22 +1,17 @@ """Test integration with zarr library""" import os - -import numpy as np - -import pytest import tempfile -import zarr - from pathlib import Path -from tempfile import NamedTemporaryFile - +import pytest +import numpy as np +import zarr from zarr_sqlite import SQLiteStore @pytest.fixture def temp_db_file(): - tmp_db = NamedTemporaryFile(suffix=".db", delete=False, delete_on_close=False) + tmp_db = tempfile.NamedTemporaryFile(suffix=".db", delete=False, delete_on_close=False) tmp_db.close() yield tmp_db.name os.remove(tmp_db.name) diff --git a/uv.lock b/uv.lock index 5faa84b..8999618 100644 --- a/uv.lock +++ b/uv.lock @@ -366,6 +366,7 @@ test = [ [package.dev-dependencies] dev = [ + { name = "numpy" }, { name = "pytest" }, ] @@ -378,4 +379,7 @@ requires-dist = [ provides-extras = ["test"] [package.metadata.requires-dev] -dev = [{ name = "pytest", specifier = ">=8.4.1" }] +dev = [ + { name = "numpy", specifier = ">=2.3.2" }, + { name = "pytest", specifier = ">=8.4.1" }, +] From 0b02b9e0b751fdbcb97aaefd1cecea3390081bd1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Sat, 18 Jul 2026 15:04:39 -0400 Subject: [PATCH 04/13] Rename test_sqlitestore.py to test_zarr_integration.py --- test/{test_sqlitestore.py => test_zarr_integration.py} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename test/{test_sqlitestore.py => test_zarr_integration.py} (100%) diff --git a/test/test_sqlitestore.py b/test/test_zarr_integration.py similarity index 100% rename from test/test_sqlitestore.py rename to test/test_zarr_integration.py From 76258c3c5444ad35e29fbaf062fd098ba7ea57e7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Sat, 18 Jul 2026 17:01:37 -0400 Subject: [PATCH 05/13] Add integration test for group listing --- src/zarr_sqlite/zarr_sqlite.py | 23 +++++++++--------- test/test_zarr_integration.py | 44 ++++++++++++++++++++++++++++++++++ 2 files changed, 56 insertions(+), 11 deletions(-) diff --git a/src/zarr_sqlite/zarr_sqlite.py b/src/zarr_sqlite/zarr_sqlite.py index 640f310..b4142d1 100644 --- a/src/zarr_sqlite/zarr_sqlite.py +++ b/src/zarr_sqlite/zarr_sqlite.py @@ -1,12 +1,16 @@ from __future__ import annotations +from typing import override, cast +from collections.abc import Iterable, AsyncIterator, Sequence import asyncio import sqlite3 from pathlib import Path -from typing import TYPE_CHECKING, override, cast import urllib.parse import uuid +from zarr.core.buffer import BufferPrototype, Buffer +from zarr.core.common import BytesLike + from zarr.abc.store import ( ByteRequest, OffsetByteRequest, @@ -15,12 +19,6 @@ SuffixByteRequest, ) -if TYPE_CHECKING: - from collections.abc import AsyncIterator, Iterable, Sequence - from zarr.core.buffer import BufferPrototype - from zarr.core.buffer import Buffer - from zarr.core.common import BytesLike - class SQLiteStore(Store): """ @@ -302,20 +300,23 @@ async def list(self) -> AsyncIterator[str]: @override async def list_prefix(self, prefix: str) -> AsyncIterator[str]: - if not prefix.endswith("/"): - prefix += "/" - cur = await self._execute("SELECT k FROM zarr WHERE k GLOB ?", (prefix + "*",)) + glob = prefix + "*" + cur = await self._execute("SELECT k FROM zarr WHERE k GLOB ?", (glob,)) for row in cast(Iterable[tuple[str]], cur): yield row[0] @override async def list_dir(self, prefix: str) -> AsyncIterator[str]: + if prefix and not prefix.endswith("/"): + prefix += "/" + seen: set[str] = set() async for full_key in self.list_prefix(prefix): relative_parts = full_key.removeprefix(prefix).split("/") k = relative_parts[0] if len(relative_parts) > 1: - k = k + "/" # Is a prefix + # k is a prefix + k = k + "/" if k not in seen: seen.add(k) yield k diff --git a/test/test_zarr_integration.py b/test/test_zarr_integration.py index 865c1b5..4dd1901 100644 --- a/test/test_zarr_integration.py +++ b/test/test_zarr_integration.py @@ -104,6 +104,50 @@ def test_append_array(sqlite_store): assert np.array_equal(z[:], expected) +def test_group_listing_methods(sqlite_store): + root = zarr.create_group(store=sqlite_store) + + a = root.create_group("a") + a.create_array(shape=(10, 10), chunks=(10, 10), dtype="f4", name="array_a") + sub = a.create_group("sub") + sub.create_array(shape=(5, 5), chunks=(5, 5), dtype="f4", name="array_b") + + b = root.create_group("b") + b.create_array(shape=(3, 3), chunks=(3, 3), dtype="f4", name="array_c") + + root = zarr.open_group(store=sqlite_store) + + assert set(root.keys()) == {"a", "b"} + assert set(root.array_keys()) == set() + assert set(root.group_keys()) == {"a", "b"} + + assert set(k for k, _ in root.groups()) == {"a", "b"} + assert all(isinstance(g, zarr.Group) for _, g in root.groups()) + + a = root["a"] + assert set(a.keys()) == {"array_a", "sub"} + assert set(a.array_keys()) == {"array_a"} + assert set(a.group_keys()) == {"sub"} + assert len(list(a.array_values())) == 1 + assert set(k for k, _ in a.arrays()) == {"array_a"} + assert all(isinstance(arr, zarr.Array) for arr in a.array_values()) + + sub = a["sub"] + assert set(sub.array_keys()) == {"array_b"} + assert len(list(sub.array_values())) == 1 + + assert set(g.name for g in root.group_values()) == {"/a", "/b"} + assert all(isinstance(g, zarr.Group) for g in root.group_values()) + + assert set(k for k, _ in root.groups()) == {"a", "b"} + assert isinstance(root["a"], zarr.Group) + assert isinstance(root["b"], zarr.Group) + assert isinstance(root["a/array_a"], zarr.Array) + assert isinstance(root["a/sub/array_b"], zarr.Array) + + assert set(root) == {"a", "b"} + + def test_store_array_creates_file_and_persists(): """Ensure file is created automatically when it doesn't exist""" with tempfile.TemporaryDirectory() as tmpdir: From 2af4afbfdb86b50e3842963620ed1a918b141bae Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Sat, 18 Jul 2026 17:08:15 -0400 Subject: [PATCH 06/13] Raise valuerror on unknown byte range type --- src/zarr_sqlite/zarr_sqlite.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/zarr_sqlite/zarr_sqlite.py b/src/zarr_sqlite/zarr_sqlite.py index b4142d1..ea7d613 100644 --- a/src/zarr_sqlite/zarr_sqlite.py +++ b/src/zarr_sqlite/zarr_sqlite.py @@ -251,6 +251,8 @@ async def get( elif isinstance(byte_range, SuffixByteRequest): a = min(len(blob), byte_range.suffix) return prototype.buffer.from_bytes(blob[-a:]) + else: + raise ValueError(f"Unsupported byte range type: {type(byte_range)}") @override async def get_partial_values( From 035cd00735a6fe095cedf8348a183d8dec0f770c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Sat, 18 Jul 2026 17:30:28 -0400 Subject: [PATCH 07/13] Add unit tests that directly test the SqliteStore public API --- pyproject.toml | 7 + src/zarr_sqlite/zarr_sqlite.py | 80 ++++++--- test/test_sqlitestore.py | 288 +++++++++++++++++++++++++++++++++ test/test_zarr_integration.py | 27 ++-- uv.lock | 15 ++ 5 files changed, 386 insertions(+), 31 deletions(-) create mode 100644 test/test_sqlitestore.py diff --git a/pyproject.toml b/pyproject.toml index ca1a242..a3406d0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,4 +25,11 @@ build-backend = "uv_build" dev = [ "numpy>=2.3.2", "pytest>=8.4.1", + "pytest-asyncio>=1.1.0", +] + +[tool.pytest.ini_options] +asyncio_mode = "auto" +markers = [ + "asyncio: mark a test as an asyncio coroutine", ] diff --git a/src/zarr_sqlite/zarr_sqlite.py b/src/zarr_sqlite/zarr_sqlite.py index ea7d613..00718c7 100644 --- a/src/zarr_sqlite/zarr_sqlite.py +++ b/src/zarr_sqlite/zarr_sqlite.py @@ -20,6 +20,40 @@ ) +def _validate_key(key: str): + """Validates a key according to SQLiteStore specification + + From the Zarr core spec: + - a key is a Unicode string, where the final character is not a `/` character. + + Additional checks (not in the core spec): + - a key which starts with '/' is invalid, and + - a key that contains '//' is invalid. + + The empty string is a valid key: it addresses a store's root resource as a single blob. + """ + is_valid = not (key.startswith("/") or key.endswith("/") or "//" in key) + if not is_valid: + raise ValueError(f"Invalid key '{key}'") + + +def _normalize_prefix(prefix: str) -> str: + """Validate a prefix string and append trailing `/` if needed + + Validation is identical to key validation, except that a prefix may end in a `/` + character. A trailing `/` is appended to prefix if absent. + + The empty string is a valid prefix (root group). The string "/" is not a valid + prefix. + """ + is_valid = not (prefix.startswith("/") or "//" in prefix) + if not is_valid: + raise ValueError(f"Invalid prefix '{prefix}'") + if prefix != "" and not prefix.endswith("/"): + prefix += "/" + return prefix + + class SQLiteStore(Store): """ Store for the local file system. @@ -193,11 +227,9 @@ def close(self) -> None: @override async def is_empty(self, prefix: str) -> bool: - if not prefix.endswith("/"): - prefix += "/" - cur = await self._execute( - "SELECT COUNT(*) FROM zarr WHERE k GLOB ?", (prefix + "/*",) - ) + prefix = _normalize_prefix(prefix) + glob = prefix + "*" + cur = await self._execute("SELECT COUNT(*) FROM zarr WHERE k GLOB ?", (glob,)) return cast(tuple[int], cur.fetchone())[0] == 0 @override @@ -231,7 +263,10 @@ async def get( prototype: BufferPrototype, byte_range: ByteRequest | None = None, ) -> Buffer | None: + # TODO: use the blob API to select a byte range directly from SQLite if possible + + _validate_key(key) cur = await self._execute("SELECT v FROM zarr WHERE k = ?", (key,)) row = cast(tuple[object] | None, cur.fetchone()) if row is None: @@ -268,17 +303,20 @@ async def get_partial_values( @override async def exists(self, key: str) -> bool: + _validate_key(key) cur = await self._execute("SELECT v FROM zarr WHERE k = ?", (key,)) return cur.fetchone() is not None @override async def set(self, key: str, value: Buffer) -> None: + _validate_key(key) await self._execute_write( "INSERT OR REPLACE INTO zarr (k, v) VALUES (?, ?)", (key, value.to_bytes()) ) @override async def set_if_not_exists(self, key: str, value: Buffer) -> None: + _validate_key(key) await self._execute_write( "INSERT OR IGNORE INTO zarr (k, v) VALUES (?, ?)", (key, value.to_bytes()) ) @@ -302,6 +340,7 @@ async def list(self) -> AsyncIterator[str]: @override async def list_prefix(self, prefix: str) -> AsyncIterator[str]: + prefix = _normalize_prefix(prefix) glob = prefix + "*" cur = await self._execute("SELECT k FROM zarr WHERE k GLOB ?", (glob,)) for row in cast(Iterable[tuple[str]], cur): @@ -309,14 +348,13 @@ async def list_prefix(self, prefix: str) -> AsyncIterator[str]: @override async def list_dir(self, prefix: str) -> AsyncIterator[str]: - if prefix and not prefix.endswith("/"): - prefix += "/" - + prefix = _normalize_prefix(prefix) seen: set[str] = set() async for full_key in self.list_prefix(prefix): - relative_parts = full_key.removeprefix(prefix).split("/") - k = relative_parts[0] - if len(relative_parts) > 1: + rel_key = full_key.removeprefix(prefix) + parts = rel_key.split("/") + k = parts[0] + if len(parts) > 1: # k is a prefix k = k + "/" if k not in seen: @@ -325,30 +363,30 @@ async def list_dir(self, prefix: str) -> AsyncIterator[str]: @override async def delete_dir(self, prefix: str) -> None: - prefix = prefix.rstrip("/") - if await self.exists(prefix): + prefix = _normalize_prefix(prefix) + if await self.exists(prefix.rstrip("/")): raise ValueError( f"Cannot delete directory {prefix} as it is a key in the store." ) - else: - await self._execute_write( - "DELETE FROM zarr WHERE k GLOB ?", (prefix + "/*",) - ) + + glob = prefix + "*" + await self._execute_write("DELETE FROM zarr WHERE k GLOB ?", (glob,)) @override async def getsize(self, key: str) -> int: + _validate_key(key) cur = await self._execute("SELECT LENGTH(v) FROM zarr WHERE k = ?", (key,)) row = cast(tuple[int] | None, cur.fetchone()) if row is None: - raise FileNotFoundError(key) + raise ValueError(f"Key '{key}' does not exist in store.") return row[0] @override async def getsize_prefix(self, prefix: str) -> int: - if not prefix.endswith("/"): - prefix += "/" + prefix = _normalize_prefix(prefix) + glob = prefix + "*" cur = await self._execute( - "SELECT SUM(LENGTH(v)) FROM zarr WHERE k GLOB ?", (prefix + "*",) + "SELECT SUM(LENGTH(v)) FROM zarr WHERE k GLOB ?", (glob,) ) size = cast(tuple[int | None], cur.fetchone())[0] if size is None: diff --git a/test/test_sqlitestore.py b/test/test_sqlitestore.py new file mode 100644 index 0000000..143db7a --- /dev/null +++ b/test/test_sqlitestore.py @@ -0,0 +1,288 @@ +"""Unit tests for the public API of SQLiteStore.""" + +import pytest +from zarr.core.buffer import BufferPrototype, default_buffer_prototype +from zarr.abc.store import OffsetByteRequest, RangeByteRequest, SuffixByteRequest + +from zarr_sqlite import SQLiteStore + + +def make_buffer(data: bytes, prototype: BufferPrototype | None = None) -> object: + prototype = prototype or default_buffer_prototype() + return prototype.buffer.from_bytes(data) + + +async def collect(iterator): + return [item async for item in iterator] + + +@pytest.fixture +def store(): + s = SQLiteStore(":memory:") + yield s + s.close() + + +@pytest.mark.asyncio +async def test_set_and_get(store): + data = b"hello world" + await store.set("foo", make_buffer(data)) + buf = await store.get("foo", default_buffer_prototype()) + assert buf.to_bytes() == data + + +@pytest.mark.asyncio +async def test_get_nonexistent_returns_none(store): + buf = await store.get("missing", default_buffer_prototype()) + assert buf is None + + +@pytest.mark.asyncio +async def test_get_byte_range_none(store): + data = b"abcdefghij" + await store.set("k", make_buffer(data)) + buf = await store.get("k", default_buffer_prototype(), byte_range=None) + assert buf.to_bytes() == data + + +@pytest.mark.asyncio +async def test_get_offset_byte_request(store): + data = b"abcdefghij" + await store.set("k", make_buffer(data)) + buf = await store.get( + "k", default_buffer_prototype(), byte_range=OffsetByteRequest(3) + ) + assert buf.to_bytes() == b"defghij" + + +@pytest.mark.asyncio +async def test_get_offset_beyond_length(store): + data = b"abc" + await store.set("k", make_buffer(data)) + buf = await store.get( + "k", default_buffer_prototype(), byte_range=OffsetByteRequest(100) + ) + assert buf.to_bytes() == b"" + + +@pytest.mark.asyncio +async def test_get_range_byte_request(store): + data = b"abcdefghij" + await store.set("k", make_buffer(data)) + buf = await store.get( + "k", + default_buffer_prototype(), + byte_range=RangeByteRequest(start=2, end=5), + ) + assert buf.to_bytes() == b"cde" + + +@pytest.mark.asyncio +async def test_get_range_clamped(store): + data = b"abc" + await store.set("k", make_buffer(data)) + buf = await store.get( + "k", + default_buffer_prototype(), + byte_range=RangeByteRequest(start=-5, end=100), + ) + assert buf.to_bytes() == b"abc" + + +@pytest.mark.asyncio +async def test_get_suffix_byte_request(store): + data = b"abcdefghij" + await store.set("k", make_buffer(data)) + buf = await store.get( + "k", default_buffer_prototype(), byte_range=SuffixByteRequest(3) + ) + assert buf.to_bytes() == b"hij" + + +@pytest.mark.asyncio +async def test_get_suffix_larger_than_length(store): + data = b"abc" + await store.set("k", make_buffer(data)) + buf = await store.get( + "k", default_buffer_prototype(), byte_range=SuffixByteRequest(100) + ) + assert buf.to_bytes() == b"abc" + + +@pytest.mark.asyncio +async def test_get_non_bytes_raises(store): + await store._execute_write("INSERT INTO zarr (k, v) VALUES (?, ?)", ("k", 5)) + with pytest.raises(TypeError): + await store.get("k", default_buffer_prototype()) + + +@pytest.mark.asyncio +async def test_get_unsupported_byte_range(store): + data = b"abc" + await store.set("k", make_buffer(data)) + bad = object() + with pytest.raises(ValueError): + await store.get("k", default_buffer_prototype(), byte_range=bad) + + +@pytest.mark.asyncio +async def test_get_partial_values(store): + await store.set("a", make_buffer(b"0123456789")) + await store.set("b", make_buffer(b"ABCDEFGHIJ")) + results = await store.get_partial_values( + default_buffer_prototype(), + [ + ("a", None), + ("b", RangeByteRequest(start=0, end=3)), + ("missing", None), + ], + ) + assert results[0].to_bytes() == b"0123456789" + assert results[1].to_bytes() == b"ABC" + assert results[2] is None + + +@pytest.mark.asyncio +async def test_set_overwrites(store): + await store.set("k", make_buffer(b"first")) + await store.set("k", make_buffer(b"second")) + buf = await store.get("k", default_buffer_prototype()) + assert buf.to_bytes() == b"second" + + +@pytest.mark.asyncio +async def test_delete_erases_key(store): + await store.set("k", make_buffer(b"data")) + assert await store.exists("k") + await store.delete("k") + assert not await store.exists("k") + assert await store.get("k", default_buffer_prototype()) is None + + +@pytest.mark.asyncio +async def test_delete_missing_key_is_noop(store): + await store.delete("never_existed") + + +@pytest.mark.asyncio +async def test_delete_dir_erases_prefix(store): + await store.set("a/x", make_buffer(b"1")) + await store.set("a/y", make_buffer(b"2")) + await store.set("b/z", make_buffer(b"3")) + await store.delete_dir("a/") + assert not await store.exists("a/x") + assert not await store.exists("a/y") + assert await store.exists("b/z") + + +@pytest.mark.asyncio +async def test_delete_dir_raises_when_key_is_leaf(store): + await store.set("a", make_buffer(b"leaf")) + with pytest.raises(ValueError): + await store.delete_dir("a") + + +@pytest.mark.asyncio +async def test_delete_dir(store): + await store.set("a/b", make_buffer(b"1")) + await store.delete_dir("a/") + assert not await store.exists("a/b") + + +@pytest.mark.asyncio +async def test_list_returns_all_keys(store): + await store.set("a", make_buffer(b"1")) + await store.set("b/c", make_buffer(b"2")) + await store.set("b/d", make_buffer(b"3")) + keys = set(await collect(store.list())) + assert keys == {"a", "b/c", "b/d"} + + +@pytest.mark.asyncio +async def test_list_empty(store): + assert await collect(store.list()) == [] + + +@pytest.mark.asyncio +async def test_list_prefix(store): + await store.set("a/1", make_buffer(b"1")) + await store.set("a/2", make_buffer(b"2")) + await store.set("ab/3", make_buffer(b"3")) + await store.set("b/4", make_buffer(b"4")) + keys = set(await collect(store.list_prefix("a/"))) + assert keys == {"a/1", "a/2"} + + +@pytest.mark.asyncio +async def test_list_prefix_root(store): + await store.set("a/1", make_buffer(b"1")) + await store.set("b/2", make_buffer(b"2")) + keys = set(await collect(store.list_prefix(""))) + assert keys == {"a/1", "b/2"} + + +@pytest.mark.asyncio +async def test_list_prefix_no_match(store): + await store.set("a/1", make_buffer(b"1")) + assert await collect(store.list_prefix("zzz/")) == [] + + +@pytest.mark.asyncio +async def test_list_dir_root(store): + await store.set("a/1", make_buffer(b"1")) + await store.set("b/2", make_buffer(b"2")) + await store.set("c/d/3", make_buffer(b"3")) + + entries = set(await collect(store.list_dir(""))) + assert entries == {"a/", "b/", "c/"} + + +@pytest.mark.asyncio +async def test_list_dir_nested(store): + await store.set("a/x", make_buffer(b"1")) + await store.set("a/y", make_buffer(b"2")) + await store.set("a/sub/z", make_buffer(b"3")) + entries = set(await collect(store.list_dir("a/"))) + assert entries == {"x", "y", "sub/"} + + +@pytest.mark.asyncio +async def test_list_dir_returns_leaf_and_prefix(store): + await store.set("a/leaf", make_buffer(b"1")) + await store.set("a/group/child", make_buffer(b"2")) + entries = set(await collect(store.list_dir("a/"))) + assert entries == {"leaf", "group/"} + + +@pytest.mark.asyncio +async def test_list_dir_no_match(store): + await store.set("a/1", make_buffer(b"1")) + assert await collect(store.list_dir("zzz/")) == [] + + +@pytest.mark.asyncio +async def test_list_dir_empty_prefix_yields_nothing(store): + assert await collect(store.list_dir("nonexistent/")) == [] + + +@pytest.mark.asyncio +async def test_set_if_not_exists(store): + await store.set_if_not_exists("k", make_buffer(b"first")) + await store.set_if_not_exists("k", make_buffer(b"second")) + buf = await store.get("k", default_buffer_prototype()) + assert buf.to_bytes() == b"first" + + +@pytest.mark.asyncio +async def test_exists(store): + assert not await store.exists("k") + await store.set("k", make_buffer(b"data")) + assert await store.exists("k") + + +@pytest.mark.asyncio +async def test_is_empty(store): + assert await store.is_empty("a/") + await store.set("a/b", make_buffer(b"1")) + assert not await store.is_empty("a/") + assert await store.is_empty("c/") diff --git a/test/test_zarr_integration.py b/test/test_zarr_integration.py index 4dd1901..25aa2fb 100644 --- a/test/test_zarr_integration.py +++ b/test/test_zarr_integration.py @@ -11,7 +11,9 @@ @pytest.fixture def temp_db_file(): - tmp_db = tempfile.NamedTemporaryFile(suffix=".db", delete=False, delete_on_close=False) + tmp_db = tempfile.NamedTemporaryFile( + suffix=".db", delete=False, delete_on_close=False + ) tmp_db.close() yield tmp_db.name os.remove(tmp_db.name) @@ -69,6 +71,7 @@ def test_save_array_to_group(temp_db_file): assert np.array_equal(data, root["group1/z"][:]) +@pytest.mark.skip(reason="assumed bug in zarr-python") def test_delete_array(temp_db_file): with SQLiteStore(temp_db_file) as sqlite_store: root = zarr.create_group(store=sqlite_store) @@ -151,22 +154,26 @@ def test_group_listing_methods(sqlite_store): def test_store_array_creates_file_and_persists(): """Ensure file is created automatically when it doesn't exist""" with tempfile.TemporaryDirectory() as tmpdir: - fname = Path(tmpdir) / 'test1.zarrdb' - store = SQLiteStore(fname, read_only=False) - root = zarr.open(store=store, mode='a') - group1 = root.create_group('group1') + fname = Path(tmpdir) / "test1.zarrdb" + store = SQLiteStore(fname, read_only=False) + root = zarr.open(store=store, mode="a") + group1 = root.create_group("group1") - i = np.arange(80000)/60.0 + i = np.arange(80000) / 60.0 x = i + np.random.normal(size=len(i), scale=1.5) - y = np.sin(x/13) + 0.7*np.cos(x/47) + np.random.normal(size=len(i), scale=0.05) + y = ( + np.sin(x / 13) + + 0.7 * np.cos(x / 47) + + np.random.normal(size=len(i), scale=0.05) + ) x.shape = (400, 200) y.shape = (200, 400) - group1.create_array(data=x, name='xdat') - group1.create_array(data=y, name='ydat') + group1.create_array(data=x, name="xdat") + group1.create_array(data=y, name="ydat") store.close() - read_root = zarr.open(store=SQLiteStore(fname, read_only=True), mode='r') + read_root = zarr.open(store=SQLiteStore(fname, read_only=True), mode="r") assert np.array_equal(read_root["group1/ydat"][:], y) assert np.array_equal(read_root["group1/xdat"], x) diff --git a/uv.lock b/uv.lock index 8999618..3440569 100644 --- a/uv.lock +++ b/uv.lock @@ -290,6 +290,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/29/16/c8a903f4c4dffe7a12843191437d7cd8e32751d5de349d45d3fe69544e87/pytest-8.4.1-py3-none-any.whl", hash = "sha256:539c70ba6fcead8e78eebbf1115e8b589e7565830d7d006a8723f19ac8a0afb7", size = 365474, upload-time = "2025-06-18T05:48:03.955Z" }, ] +[[package]] +name = "pytest-asyncio" +version = "1.4.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pytest" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/43/7c/d36d04db312ecf4298932ef77e6e4a9e8ad017906e24e34f0b0c361a2473/pytest_asyncio-1.4.0.tar.gz", hash = "sha256:c6c0d2259945122819f171a32ecea2c349ead889ee28176caaf492143424be42", size = 58514, upload-time = "2026-05-26T09:56:04.083Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/03/e2/08a497ef684b88559c9cc5f4ad53a37e7b99e727094a86d6ea32536d5d3c/pytest_asyncio-1.4.0-py3-none-any.whl", hash = "sha256:933ca923a23075a87fb7070c0ec272a6848489824d887c85c812670932835aa1", size = 16930, upload-time = "2026-05-26T09:56:02.576Z" }, +] + [[package]] name = "pyyaml" version = "6.0.2" @@ -368,6 +381,7 @@ test = [ dev = [ { name = "numpy" }, { name = "pytest" }, + { name = "pytest-asyncio" }, ] [package.metadata] @@ -382,4 +396,5 @@ provides-extras = ["test"] dev = [ { name = "numpy", specifier = ">=2.3.2" }, { name = "pytest", specifier = ">=8.4.1" }, + { name = "pytest-asyncio", specifier = ">=1.1.0" }, ] From eae1e6826985ccb412e02d248c891168351d690b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Sun, 19 Jul 2026 14:26:32 -0400 Subject: [PATCH 08/13] Fix array delete test --- test/test_zarr_integration.py | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/test/test_zarr_integration.py b/test/test_zarr_integration.py index 25aa2fb..ea7faa7 100644 --- a/test/test_zarr_integration.py +++ b/test/test_zarr_integration.py @@ -71,26 +71,27 @@ def test_save_array_to_group(temp_db_file): assert np.array_equal(data, root["group1/z"][:]) -@pytest.mark.skip(reason="assumed bug in zarr-python") def test_delete_array(temp_db_file): with SQLiteStore(temp_db_file) as sqlite_store: root = zarr.create_group(store=sqlite_store) group1 = root.create_group("group1") - group1.create_array(shape=(100, 100), chunks=(10, 10), dtype="f4", name="z") + z = group1.create_array(shape=(100, 100), chunks=(10, 10), dtype="f4", name="z") + z[:] = 200 - with SQLiteStore(temp_db_file) as sqlite_store: - root = zarr.open_group(store=sqlite_store) assert "group1/z" in root assert root["group1/z"].shape == (100, 100) + + # Delete array del root["group1/z"] - with SQLiteStore(temp_db_file) as sqlite_store: - root = zarr.open_group(store=sqlite_store) + # Ensure group still exists assert "group1" in root assert isinstance(root["group1"], zarr.Group) + + # Ensure array deleted assert "group1/z" not in root with pytest.raises(KeyError): - root["group1/z"][:] + root["group1/z"] def test_append_array(sqlite_store): From 0ea52a8161d344e191a973e5b5745e92e24a73ad Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Sun, 19 Jul 2026 14:38:30 -0400 Subject: [PATCH 09/13] Remove type casts --- src/zarr_sqlite/zarr_sqlite.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/src/zarr_sqlite/zarr_sqlite.py b/src/zarr_sqlite/zarr_sqlite.py index 00718c7..6906f7c 100644 --- a/src/zarr_sqlite/zarr_sqlite.py +++ b/src/zarr_sqlite/zarr_sqlite.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import override, cast +from typing import override from collections.abc import Iterable, AsyncIterator, Sequence import asyncio import sqlite3 @@ -230,7 +230,7 @@ async def is_empty(self, prefix: str) -> bool: prefix = _normalize_prefix(prefix) glob = prefix + "*" cur = await self._execute("SELECT COUNT(*) FROM zarr WHERE k GLOB ?", (glob,)) - return cast(tuple[int], cur.fetchone())[0] == 0 + return cur.fetchone()[0] == 0 @override async def clear(self) -> None: @@ -268,7 +268,7 @@ async def get( _validate_key(key) cur = await self._execute("SELECT v FROM zarr WHERE k = ?", (key,)) - row = cast(tuple[object] | None, cur.fetchone()) + row = cur.fetchone() if row is None: return None blob = row[0] @@ -335,16 +335,16 @@ async def set_partial_values( @override async def list(self) -> AsyncIterator[str]: cur = await self._execute("SELECT k FROM zarr") - for row in cast(Iterable[tuple[str]], cur): - yield row[0] + for row in cur: + yield str(row[0]) @override async def list_prefix(self, prefix: str) -> AsyncIterator[str]: prefix = _normalize_prefix(prefix) glob = prefix + "*" cur = await self._execute("SELECT k FROM zarr WHERE k GLOB ?", (glob,)) - for row in cast(Iterable[tuple[str]], cur): - yield row[0] + for row in cur: + yield str(row[0]) @override async def list_dir(self, prefix: str) -> AsyncIterator[str]: @@ -376,10 +376,10 @@ async def delete_dir(self, prefix: str) -> None: async def getsize(self, key: str) -> int: _validate_key(key) cur = await self._execute("SELECT LENGTH(v) FROM zarr WHERE k = ?", (key,)) - row = cast(tuple[int] | None, cur.fetchone()) + row = cur.fetchone() if row is None: raise ValueError(f"Key '{key}' does not exist in store.") - return row[0] + return int(row[0]) @override async def getsize_prefix(self, prefix: str) -> int: @@ -388,7 +388,7 @@ async def getsize_prefix(self, prefix: str) -> int: cur = await self._execute( "SELECT SUM(LENGTH(v)) FROM zarr WHERE k GLOB ?", (glob,) ) - size = cast(tuple[int | None], cur.fetchone())[0] + size = cur.fetchone() if size is None: - size = 0 - return size + return 0 + return int(size) From a7b26476f6edd70a71091c30a46d58cc44cd360d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Sun, 19 Jul 2026 17:30:58 -0400 Subject: [PATCH 10/13] Add unit tests for getsize, getsize_prefix and others --- src/zarr_sqlite/zarr_sqlite.py | 6 +-- test/test_sqlitestore.py | 95 ++++++++++++++++++++++++++++++++++ 2 files changed, 98 insertions(+), 3 deletions(-) diff --git a/src/zarr_sqlite/zarr_sqlite.py b/src/zarr_sqlite/zarr_sqlite.py index 6906f7c..ae1b338 100644 --- a/src/zarr_sqlite/zarr_sqlite.py +++ b/src/zarr_sqlite/zarr_sqlite.py @@ -378,7 +378,7 @@ async def getsize(self, key: str) -> int: cur = await self._execute("SELECT LENGTH(v) FROM zarr WHERE k = ?", (key,)) row = cur.fetchone() if row is None: - raise ValueError(f"Key '{key}' does not exist in store.") + raise FileNotFoundError(key) return int(row[0]) @override @@ -389,6 +389,6 @@ async def getsize_prefix(self, prefix: str) -> int: "SELECT SUM(LENGTH(v)) FROM zarr WHERE k GLOB ?", (glob,) ) size = cur.fetchone() - if size is None: + if size is None or size[0] is None: return 0 - return int(size) + return int(size[0]) diff --git a/test/test_sqlitestore.py b/test/test_sqlitestore.py index 143db7a..18f83f6 100644 --- a/test/test_sqlitestore.py +++ b/test/test_sqlitestore.py @@ -5,6 +5,7 @@ from zarr.abc.store import OffsetByteRequest, RangeByteRequest, SuffixByteRequest from zarr_sqlite import SQLiteStore +from zarr_sqlite.zarr_sqlite import _validate_key, _normalize_prefix def make_buffer(data: bytes, prototype: BufferPrototype | None = None) -> object: @@ -286,3 +287,97 @@ async def test_is_empty(store): await store.set("a/b", make_buffer(b"1")) assert not await store.is_empty("a/") assert await store.is_empty("c/") + + +@pytest.mark.asyncio +async def test_clear(store): + await store.set("a", make_buffer(b"1")) + await store.set("b/c", make_buffer(b"2")) + assert await collect(store.list()) + await store.clear() + assert await collect(store.list()) == [] + + +@pytest.mark.asyncio +async def test_getsize(store): + await store.set("k", make_buffer(b"12345")) + assert await store.getsize("k") == 5 + + await store.set("empty", make_buffer(b"")) + assert await store.getsize("empty") == 0 + + with pytest.raises(FileNotFoundError): + await store.getsize("missing") + + +@pytest.mark.asyncio +async def test_getsize_prefix(store): + await store.set("a/1", make_buffer(b"abc")) + await store.set("a/2", make_buffer(b"de")) + await store.set("a/3", make_buffer(b"")) + assert await store.getsize_prefix("a/") == 5 + + # non-existent prefix, getsize_prefix should return 0 + assert await store.getsize_prefix("missing/") == 0 + + +def test_eq_same_path(tmp_path): + db = tmp_path / "eq.db" + s1 = SQLiteStore(db) + s2 = SQLiteStore(str(db)) + s3 = SQLiteStore(tmp_path / "other.db") + try: + assert s1 == s2 + assert s1 != s3 + assert s1 != "not a store" + finally: + s1.close() + s2.close() + s3.close() + + +def test_with_read_only(): + s = SQLiteStore(":memory:", read_only=False) + try: + ro = s.with_read_only(read_only=True) + assert ro._read_only is True + assert ro._is_open is False + finally: + s.close() + + +def test_validate_key_valid(): + for key in ["", "a", "a/b", "a/b/c.json", "0", "with-dash_and.dot"]: + _validate_key(key) + + +def test_validate_key_rejects_leading_slash(): + invalid_keys = ["/a", "a/", "a//b", "a/b/"] + valid_keys = ["", "a", "a/b", "a/b/c", "foo/bar/baz/qux"] + + for k in invalid_keys: + with pytest.raises(ValueError): + _validate_key(k) + + for k in valid_keys: + # Should not raise + _validate_key(k) + + +def test_normalize_prefix_valid_unchanged(): + assert _normalize_prefix("a/") == "a/" + assert _normalize_prefix("a/b/") == "a/b/" + + assert _normalize_prefix("a") == "a/" + assert _normalize_prefix("a/b") == "a/b/" + + assert _normalize_prefix("") == "" + + with pytest.raises(ValueError): + _normalize_prefix("/a") + + with pytest.raises(ValueError): + _normalize_prefix("a//b") + + with pytest.raises(ValueError): + _normalize_prefix("/") From b8b59d80e49d49bf1b72e19054ed55e169005827 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Mon, 20 Jul 2026 21:31:33 -0400 Subject: [PATCH 11/13] Edits to LLM-generated tests --- src/zarr_sqlite/zarr_sqlite.py | 1 + test/test_sqlitestore.py | 73 ++++++++++++++++++++-------------- 2 files changed, 44 insertions(+), 30 deletions(-) diff --git a/src/zarr_sqlite/zarr_sqlite.py b/src/zarr_sqlite/zarr_sqlite.py index ae1b338..b2edb4d 100644 --- a/src/zarr_sqlite/zarr_sqlite.py +++ b/src/zarr_sqlite/zarr_sqlite.py @@ -309,6 +309,7 @@ async def exists(self, key: str) -> bool: @override async def set(self, key: str, value: Buffer) -> None: + self._check_writable() _validate_key(key) await self._execute_write( "INSERT OR REPLACE INTO zarr (k, v) VALUES (?, ?)", (key, value.to_bytes()) diff --git a/test/test_sqlitestore.py b/test/test_sqlitestore.py index 18f83f6..daad65c 100644 --- a/test/test_sqlitestore.py +++ b/test/test_sqlitestore.py @@ -7,6 +7,8 @@ from zarr_sqlite import SQLiteStore from zarr_sqlite.zarr_sqlite import _validate_key, _normalize_prefix +from tempfile import NamedTemporaryFile + def make_buffer(data: bytes, prototype: BufferPrototype | None = None) -> object: prototype = prototype or default_buffer_prototype() @@ -17,6 +19,13 @@ async def collect(iterator): return [item async for item in iterator] +async def get_as_bytes(s: SQLiteStore, key: str) -> bytes | None: + buf = await s.get(key, default_buffer_prototype()) + if buf is None: + return None + return buf.to_bytes() + + @pytest.fixture def store(): s = SQLiteStore(":memory:") @@ -24,6 +33,15 @@ def store(): s.close() +@pytest.fixture +def tmpfile_store(): + fp = NamedTemporaryFile(suffix=".db") + fp.close + s = SQLiteStore(fp.name) + yield s + s.close() + + @pytest.mark.asyncio async def test_set_and_get(store): data = b"hello world" @@ -33,19 +51,11 @@ async def test_set_and_get(store): @pytest.mark.asyncio -async def test_get_nonexistent_returns_none(store): +async def test_get_nonexistent(store): buf = await store.get("missing", default_buffer_prototype()) assert buf is None -@pytest.mark.asyncio -async def test_get_byte_range_none(store): - data = b"abcdefghij" - await store.set("k", make_buffer(data)) - buf = await store.get("k", default_buffer_prototype(), byte_range=None) - assert buf.to_bytes() == data - - @pytest.mark.asyncio async def test_get_offset_byte_request(store): data = b"abcdefghij" @@ -80,14 +90,14 @@ async def test_get_range_byte_request(store): @pytest.mark.asyncio async def test_get_range_clamped(store): - data = b"abc" + data = b"abcdefghijkl" await store.set("k", make_buffer(data)) buf = await store.get( "k", default_buffer_prototype(), byte_range=RangeByteRequest(start=-5, end=100), ) - assert buf.to_bytes() == b"abc" + assert buf.to_bytes() == data @pytest.mark.asyncio @@ -97,6 +107,7 @@ async def test_get_suffix_byte_request(store): buf = await store.get( "k", default_buffer_prototype(), byte_range=SuffixByteRequest(3) ) + assert len(buf) == 3 assert buf.to_bytes() == b"hij" @@ -107,7 +118,7 @@ async def test_get_suffix_larger_than_length(store): buf = await store.get( "k", default_buffer_prototype(), byte_range=SuffixByteRequest(100) ) - assert buf.to_bytes() == b"abc" + assert buf.to_bytes() == data @pytest.mark.asyncio @@ -146,6 +157,9 @@ async def test_get_partial_values(store): @pytest.mark.asyncio async def test_set_overwrites(store): await store.set("k", make_buffer(b"first")) + buf = await store.get("k", default_buffer_prototype()) + assert buf.to_bytes() == b"first" + await store.set("k", make_buffer(b"second")) buf = await store.get("k", default_buffer_prototype()) assert buf.to_bytes() == b"second" @@ -233,9 +247,10 @@ async def test_list_dir_root(store): await store.set("a/1", make_buffer(b"1")) await store.set("b/2", make_buffer(b"2")) await store.set("c/d/3", make_buffer(b"3")) + await store.set("leaf", make_buffer(b"3")) entries = set(await collect(store.list_dir(""))) - assert entries == {"a/", "b/", "c/"} + assert entries == {"a/", "b/", "c/", "leaf"} @pytest.mark.asyncio @@ -247,14 +262,6 @@ async def test_list_dir_nested(store): assert entries == {"x", "y", "sub/"} -@pytest.mark.asyncio -async def test_list_dir_returns_leaf_and_prefix(store): - await store.set("a/leaf", make_buffer(b"1")) - await store.set("a/group/child", make_buffer(b"2")) - entries = set(await collect(store.list_dir("a/"))) - assert entries == {"leaf", "group/"} - - @pytest.mark.asyncio async def test_list_dir_no_match(store): await store.set("a/1", make_buffer(b"1")) @@ -283,8 +290,11 @@ async def test_exists(store): @pytest.mark.asyncio async def test_is_empty(store): + assert await store.is_empty("") assert await store.is_empty("a/") + await store.set("a/b", make_buffer(b"1")) + assert not await store.is_empty("") assert not await store.is_empty("a/") assert await store.is_empty("c/") @@ -293,7 +303,7 @@ async def test_is_empty(store): async def test_clear(store): await store.set("a", make_buffer(b"1")) await store.set("b/c", make_buffer(b"2")) - assert await collect(store.list()) + assert len(await collect(store.list())) == 2 await store.clear() assert await collect(store.list()) == [] @@ -336,14 +346,17 @@ def test_eq_same_path(tmp_path): s3.close() -def test_with_read_only(): - s = SQLiteStore(":memory:", read_only=False) - try: - ro = s.with_read_only(read_only=True) - assert ro._read_only is True - assert ro._is_open is False - finally: - s.close() +@pytest.mark.asyncio +async def test_with_read_only(tmpfile_store): + await tmpfile_store.set("a", make_buffer(b"data")) + + ro = tmpfile_store.with_read_only(read_only=True) + assert ro.read_only is True + + assert await get_as_bytes(ro, "a") == b"data" + + with pytest.raises(ValueError): + await ro.set("b", make_buffer(b"data")) def test_validate_key_valid(): From 713391102820909e9a14dfd4ee5eb1bb03e303c9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Thu, 23 Jul 2026 18:16:25 -0400 Subject: [PATCH 12/13] Use check_writable on all methods needing write access --- src/zarr_sqlite/zarr_sqlite.py | 4 ++++ test/test_sqlitestore.py | 18 ++++++++++++++++++ 2 files changed, 22 insertions(+) diff --git a/src/zarr_sqlite/zarr_sqlite.py b/src/zarr_sqlite/zarr_sqlite.py index b2edb4d..f43b36a 100644 --- a/src/zarr_sqlite/zarr_sqlite.py +++ b/src/zarr_sqlite/zarr_sqlite.py @@ -235,6 +235,7 @@ async def is_empty(self, prefix: str) -> bool: @override async def clear(self) -> None: """Clear the store.""" + self._check_writable() await self._execute_write("DROP TABLE IF EXISTS zarr") await self._create_schema() @@ -317,6 +318,7 @@ async def set(self, key: str, value: Buffer) -> None: @override async def set_if_not_exists(self, key: str, value: Buffer) -> None: + self._check_writable() _validate_key(key) await self._execute_write( "INSERT OR IGNORE INTO zarr (k, v) VALUES (?, ?)", (key, value.to_bytes()) @@ -324,6 +326,7 @@ async def set_if_not_exists(self, key: str, value: Buffer) -> None: @override async def delete(self, key: str) -> None: + self._check_writable() await self._execute_write("DELETE FROM zarr WHERE k = ?", (key,)) # TODO: Implement partial writes with blob API @@ -364,6 +367,7 @@ async def list_dir(self, prefix: str) -> AsyncIterator[str]: @override async def delete_dir(self, prefix: str) -> None: + self._check_writable() prefix = _normalize_prefix(prefix) if await self.exists(prefix.rstrip("/")): raise ValueError( diff --git a/test/test_sqlitestore.py b/test/test_sqlitestore.py index daad65c..d321311 100644 --- a/test/test_sqlitestore.py +++ b/test/test_sqlitestore.py @@ -359,6 +359,24 @@ async def test_with_read_only(tmpfile_store): await ro.set("b", make_buffer(b"data")) +@pytest.mark.asyncio +async def test_read_only_raises(tmpfile_store): + await tmpfile_store.set("a", make_buffer(b"data")) + tmpfile_store.close() + store = SQLiteStore(tmpfile_store.database_uri, read_only=True) + + with pytest.raises(ValueError): + await store.delete("a") + with pytest.raises(ValueError): + await store.set("b", make_buffer(b"data")) + with pytest.raises(ValueError): + await store.set_if_not_exists("b", make_buffer(b"data")) + with pytest.raises(ValueError): + await store.delete_dir("a/") + with pytest.raises(ValueError): + await store.clear() + + def test_validate_key_valid(): for key in ["", "a", "a/b", "a/b/c.json", "0", "with-dash_and.dot"]: _validate_key(key) From 623e81ca7e7e2debec46215de48434a583da0ae7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Francis=20Th=C3=A9rien?= Date: Thu, 23 Jul 2026 20:59:58 -0400 Subject: [PATCH 13/13] Formatting --- src/zarr_sqlite/zarr_sqlite.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/src/zarr_sqlite/zarr_sqlite.py b/src/zarr_sqlite/zarr_sqlite.py index f43b36a..dff1ae0 100644 --- a/src/zarr_sqlite/zarr_sqlite.py +++ b/src/zarr_sqlite/zarr_sqlite.py @@ -137,9 +137,8 @@ def _build_database_uri(database: Path | str, read_only: bool) -> str: # In-memory databases cannot be opened in read-only mode if read_only: raise ValueError("Cannot open an in-memory database in read-only mode.") - uri_path = "mem-" + str( - uuid.uuid4() - ) # Generate a unique ID for the in-memory database + # Generate a unique ID for the in-memory database + uri_path = "mem-" + str(uuid.uuid4()) query["mode"] = ["memory"] query["cache"] = ["shared"] elif not database.startswith("file:"): @@ -198,8 +197,7 @@ def with_read_only(self, read_only: bool = False) -> SQLiteStore: async def _execute_write(self, query: str, params: Sequence[object] = ()) -> None: """Execute a query with our lock and commit.""" await self._ensure_open() - if self._lock is None: - raise ValueError("Store is not open") + assert self._lock is not None async with self._lock: cursor = self._con.cursor() _ = cursor.execute(query, params)