diff --git a/pyproject.toml b/pyproject.toml index d035307..a3406d0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,3 +20,16 @@ test = [ [build-system] requires = ["uv_build>=0.8.3,<0.9.0"] build-backend = "uv_build" + +[dependency-groups] +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/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..dff1ae0 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 +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, @@ -14,13 +18,40 @@ 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 +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): @@ -75,7 +106,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) @@ -97,7 +128,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): @@ -106,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:"): @@ -144,7 +174,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}") @@ -161,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) @@ -190,16 +225,15 @@ 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 + "/*",) - ) - return cast(tuple[int], cur.fetchone())[0] == 0 + prefix = _normalize_prefix(prefix) + glob = prefix + "*" + cur = await self._execute("SELECT COUNT(*) FROM zarr WHERE k GLOB ?", (glob,)) + return cur.fetchone()[0] == 0 @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() @@ -228,9 +262,12 @@ 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()) + row = cur.fetchone() if row is None: return None blob = row[0] @@ -248,6 +285,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( @@ -263,23 +302,29 @@ 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: + self._check_writable() + _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: + self._check_writable() + _validate_key(key) await self._execute_write( "INSERT OR IGNORE INTO zarr (k, v) VALUES (?, ?)", (key, value.to_bytes()) ) @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 @@ -292,57 +337,61 @@ 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]: - if not prefix.endswith("/"): - prefix += "/" - cur = await self._execute("SELECT k FROM zarr WHERE k GLOB ?", (prefix + "*",)) - for row in cast(Iterable[tuple[str]], cur): - yield row[0] + prefix = _normalize_prefix(prefix) + glob = prefix + "*" + cur = await self._execute("SELECT k FROM zarr WHERE k GLOB ?", (glob,)) + for row in cur: + yield str(row[0]) @override async def list_dir(self, prefix: str) -> AsyncIterator[str]: + 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: - k = k + "/" # Is a prefix + 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: seen.add(k) yield k @override async def delete_dir(self, prefix: str) -> None: - prefix = prefix.rstrip("/") - if await self.exists(prefix): + self._check_writable() + 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()) + row = cur.fetchone() if row is None: raise FileNotFoundError(key) - return row[0] + return int(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: - size = 0 - return size + size = cur.fetchone() + if size is None or size[0] is None: + return 0 + return int(size[0]) diff --git a/test/test_sqlitestore.py b/test/test_sqlitestore.py new file mode 100644 index 0000000..d321311 --- /dev/null +++ b/test/test_sqlitestore.py @@ -0,0 +1,414 @@ +"""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 +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() + return prototype.buffer.from_bytes(data) + + +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:") + yield s + 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" + 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(store): + buf = await store.get("missing", default_buffer_prototype()) + assert buf is None + + +@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"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() == data + + +@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 len(buf) == 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() == data + + +@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")) + 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" + + +@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")) + await store.set("leaf", make_buffer(b"3")) + + entries = set(await collect(store.list_dir(""))) + assert entries == {"a/", "b/", "c/", "leaf"} + + +@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_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("") + 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/") + + +@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 len(await collect(store.list())) == 2 + 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() + + +@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")) + + +@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) + + +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("/") diff --git a/test/test_zarr.py b/test/test_zarr.py deleted file mode 100644 index 46fc089..0000000 --- a/test/test_zarr.py +++ /dev/null @@ -1,108 +0,0 @@ -"""Test integration with zarr library""" - -import os - -import numpy as np - -import pytest - -import zarr - -from tempfile import NamedTemporaryFile - -from zarr_sqlite import SQLiteStore - - -@pytest.fixture -def temp_db_file(): - tmp_db = NamedTemporaryFile(suffix=".db", delete=False, delete_on_close=False) - tmp_db.close() - yield tmp_db.name - os.remove(tmp_db.name) - - -@pytest.fixture -def sqlite_store(temp_db_file): - store = SQLiteStore(temp_db_file) - yield store - store.close() - - -def random_array(shape, dtype=np.float64): - return np.random.default_rng().random(shape, dtype) - - -def test_open_close(temp_db_file): - store = SQLiteStore.open(temp_db_file) - store.close() - - -def test_store_array(sqlite_store): - z = zarr.create_array( - store=sqlite_store, shape=(100, 100), chunks=(10, 10), dtype="f4" - ) - data = random_array((100, 100), dtype=np.float32) - z[:, :] = data - 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") - - assert isinstance(root, zarr.Group) - assert isinstance(group1, zarr.Group) - 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" - ) - - data = random_array((100, 100), dtype=np.float32) - z[:, :] = data - - del root, group1, z, sqlite_store - - with SQLiteStore(temp_db_file) as sqlite_store: - 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" - ) - - 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) - del root["group1/z"] - - with SQLiteStore(temp_db_file) as sqlite_store: - root = zarr.open_group(store=sqlite_store) - assert "group1" in root - assert isinstance(root["group1"], zarr.Group) - assert "group1/z" not in root - 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" - ) - data = random_array((100, 100), dtype=np.float32) - z[:, :] = data - - z.append(data, axis=0) - - assert z.shape == (200, 100) - expected = np.tile(data, (2, 1)) - assert np.array_equal(z[:], expected) diff --git a/test/test_zarr_integration.py b/test/test_zarr_integration.py new file mode 100644 index 0000000..ea7faa7 --- /dev/null +++ b/test/test_zarr_integration.py @@ -0,0 +1,182 @@ +"""Test integration with zarr library""" + +import os +import tempfile +from pathlib import Path +import pytest +import numpy as np +import zarr +from zarr_sqlite import SQLiteStore + + +@pytest.fixture +def temp_db_file(): + tmp_db = tempfile.NamedTemporaryFile( + suffix=".db", delete=False, delete_on_close=False + ) + tmp_db.close() + yield tmp_db.name + os.remove(tmp_db.name) + + +@pytest.fixture +def sqlite_store(temp_db_file): + store = SQLiteStore(temp_db_file) + yield store + store.close() + + +def random_array(shape, dtype=np.float64): + return np.random.default_rng().random(shape, dtype) + + +def test_open_close(temp_db_file): + store = SQLiteStore.open(temp_db_file) + store.close() + + +def test_store_array(sqlite_store): + z = zarr.create_array( + store=sqlite_store, shape=(100, 100), chunks=(10, 10), dtype="f4" + ) + data = random_array((100, 100), dtype=np.float32) + z[:, :] = data + 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") + + assert isinstance(root, zarr.Group) + assert isinstance(group1, zarr.Group) + 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") + + data = random_array((100, 100), dtype=np.float32) + z[:, :] = data + + del root, group1, z, sqlite_store + + with SQLiteStore(temp_db_file) as sqlite_store: + 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") + z = group1.create_array(shape=(100, 100), chunks=(10, 10), dtype="f4", name="z") + z[:] = 200 + + assert "group1/z" in root + assert root["group1/z"].shape == (100, 100) + + # Delete array + del root["group1/z"] + + # 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"] + + +def test_append_array(sqlite_store): + z = zarr.create_array( + store=sqlite_store, shape=(100, 100), chunks=(10, 10), dtype="f4" + ) + data = random_array((100, 100), dtype=np.float32) + z[:, :] = data + + z.append(data, axis=0) + + assert z.shape == (200, 100) + expected = np.tile(data, (2, 1)) + 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: + 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() diff --git a/uv.lock b/uv.lock index a632cbb..3440569 100644 --- a/uv.lock +++ b/uv.lock @@ -1,5 +1,5 @@ version = 1 -revision = 2 +revision = 3 requires-python = ">=3.11" [[package]] @@ -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" @@ -364,6 +377,13 @@ test = [ { name = "pytest" }, ] +[package.dev-dependencies] +dev = [ + { name = "numpy" }, + { name = "pytest" }, + { name = "pytest-asyncio" }, +] + [package.metadata] requires-dist = [ { name = "mypy", marker = "extra == 'test'", specifier = ">=1.17.0" }, @@ -371,3 +391,10 @@ requires-dist = [ { name = "zarr", specifier = ">=3.1.0" }, ] provides-extras = ["test"] + +[package.metadata.requires-dev] +dev = [ + { name = "numpy", specifier = ">=2.3.2" }, + { name = "pytest", specifier = ">=8.4.1" }, + { name = "pytest-asyncio", specifier = ">=1.1.0" }, +]