Skip to content
Merged
13 changes: 13 additions & 0 deletions pyproject.toml
Original file line numberDiff line numberDiff line change
Expand Up@@ -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",
]
21 changes: 0 additions & 21 deletions src/zarr_sqlite/scratch.py

This file was deleted.

143 changes: 96 additions & 47 deletions src/zarr_sqlite/zarr_sqlite.py
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,57 @@
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,
RangeByteRequest,
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):
Expand DownExpand Up@@ -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)
Expand All@@ -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):
Expand All@@ -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:"):
Expand DownExpand Up@@ -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}")
Expand All@@ -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)
Expand DownExpand Up@@ -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()

Expand DownExpand Up@@ -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]
Expand All@@ -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(
Expand All@@ -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
Expand All@@ -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])
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
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;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Ensure file is created when it doesnt exist by auxym · Pull Request #5 · auxym/zarr-sqlite-python · GitHub
Skip to content
Merged
13 changes: 13 additions & 0 deletions pyproject.toml
Original file line numberDiff line numberDiff line change
Expand Up@@ -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",
]
21 changes: 0 additions & 21 deletions src/zarr_sqlite/scratch.py

This file was deleted.

143 changes: 96 additions & 47 deletions src/zarr_sqlite/zarr_sqlite.py
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,57 @@
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,
RangeByteRequest,
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):
Expand DownExpand Up@@ -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)
Expand All@@ -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):
Expand All@@ -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:"):
Expand DownExpand Up@@ -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}")
Expand All@@ -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)
Expand DownExpand Up@@ -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()

Expand DownExpand Up@@ -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]
Expand All@@ -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(
Expand All@@ -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
Expand All@@ -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])
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Ensure file is created when it doesnt exist by auxym · Pull Request #5 · auxym/zarr-sqlite-python · GitHub
Skip to content
Merged
13 changes: 13 additions & 0 deletions pyproject.toml
Original file line numberDiff line numberDiff line change
Expand Up@@ -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",
]
21 changes: 0 additions & 21 deletions src/zarr_sqlite/scratch.py

This file was deleted.

143 changes: 96 additions & 47 deletions src/zarr_sqlite/zarr_sqlite.py
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,57 @@
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,
RangeByteRequest,
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):
Expand DownExpand Up@@ -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)
Expand All@@ -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):
Expand All@@ -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:"):
Expand DownExpand Up@@ -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}")
Expand All@@ -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)
Expand DownExpand Up@@ -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()

Expand DownExpand Up@@ -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]
Expand All@@ -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(
Expand All@@ -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
Expand All@@ -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])
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Ensure file is created when it doesnt exist by auxym · Pull Request #5 · auxym/zarr-sqlite-python · GitHub
Skip to content
Merged
13 changes: 13 additions & 0 deletions pyproject.toml
Original file line numberDiff line numberDiff line change
Expand Up@@ -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",
]
21 changes: 0 additions & 21 deletions src/zarr_sqlite/scratch.py

This file was deleted.

143 changes: 96 additions & 47 deletions src/zarr_sqlite/zarr_sqlite.py
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,57 @@
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,
RangeByteRequest,
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):
Expand DownExpand Up@@ -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)
Expand All@@ -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):
Expand All@@ -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:"):
Expand DownExpand Up@@ -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}")
Expand All@@ -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)
Expand DownExpand Up@@ -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()

Expand DownExpand Up@@ -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]
Expand All@@ -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(
Expand All@@ -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
Expand All@@ -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])
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + ' Ensure file is created when it doesnt exist by auxym · Pull Request #5 · auxym/zarr-sqlite-python · GitHub
Skip to content
Merged
13 changes: 13 additions & 0 deletions pyproject.toml
Original file line numberDiff line numberDiff line change
Expand Up@@ -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",
]
21 changes: 0 additions & 21 deletions src/zarr_sqlite/scratch.py

This file was deleted.

143 changes: 96 additions & 47 deletions src/zarr_sqlite/zarr_sqlite.py
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,57 @@
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,
RangeByteRequest,
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):
Expand DownExpand Up@@ -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)
Expand All@@ -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):
Expand All@@ -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:"):
Expand DownExpand Up@@ -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}")
Expand All@@ -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)
Expand DownExpand Up@@ -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()

Expand DownExpand Up@@ -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]
Expand All@@ -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(
Expand All@@ -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
Expand All@@ -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])
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Ensure file is created when it doesnt exist by auxym · Pull Request #5 · auxym/zarr-sqlite-python · GitHub
Skip to content
Merged
13 changes: 13 additions & 0 deletions pyproject.toml
Original file line numberDiff line numberDiff line change
Expand Up@@ -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",
]
21 changes: 0 additions & 21 deletions src/zarr_sqlite/scratch.py

This file was deleted.

143 changes: 96 additions & 47 deletions src/zarr_sqlite/zarr_sqlite.py
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,57 @@
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,
RangeByteRequest,
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):
Expand DownExpand Up@@ -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)
Expand All@@ -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):
Expand All@@ -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:"):
Expand DownExpand Up@@ -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}")
Expand All@@ -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)
Expand DownExpand Up@@ -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()

Expand DownExpand Up@@ -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]
Expand All@@ -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(
Expand All@@ -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
Expand All@@ -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])
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Ensure file is created when it doesnt exist by auxym · Pull Request #5 · auxym/zarr-sqlite-python · GitHub
Skip to content
Merged
13 changes: 13 additions & 0 deletions pyproject.toml
Original file line numberDiff line numberDiff line change
Expand Up@@ -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",
]
21 changes: 0 additions & 21 deletions src/zarr_sqlite/scratch.py

This file was deleted.

143 changes: 96 additions & 47 deletions src/zarr_sqlite/zarr_sqlite.py
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,57 @@
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,
RangeByteRequest,
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):
Expand DownExpand Up@@ -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)
Expand All@@ -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):
Expand All@@ -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:"):
Expand DownExpand Up@@ -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}")
Expand All@@ -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)
Expand DownExpand Up@@ -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()

Expand DownExpand Up@@ -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]
Expand All@@ -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(
Expand All@@ -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
Expand All@@ -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])
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Universal Dark Mode - works on any site (function() { var enabled = true; function applyDarkMode() { if (!enabled) return; // Create style element if it doesn't exist var style = document.getElementById('universal-dark-mode-style'); if (!style) { style = document.createElement('style'); style.id = 'universal-dark-mode-style'; document.head.appendChild(style); } // Dark mode CSS - inverts colors but preserves images/video style.textContent = ' /* Invert everything except media */ html { filter: invert(1) hue-rotate(180deg) !important; background: #1a1a2e !important; } /* Restore images, videos, iframes, canvas */ img, video, iframe, canvas, svg, picture, [style*="background-image"] { filter: invert(1) hue-rotate(180deg) !important; } /* Preserve specific elements that should not be inverted */ .no-dark-mode, .no-dark-mode *, [data-theme="light"], [data-theme="light"], .ace_editor, .ace_editor *, .CodeMirror, .CodeMirror *, .monaco-editor, .monaco-editor *, .markdown-body pre, .markdown-body pre *, .highlight, .highlight *, pre code, pre code * { filter: none !important; } /* Fix common UI elements */ .modal, .popup, .dropdown-menu, .tooltip, .popover { filter: invert(1) hue-rotate(180deg) !important; background: #2d2d44 !important; border-color: #444 !important; } /* Scrollbars */ ::-webkit-scrollbar { background: #1a1a2e !important; } ::-webkit-scrollbar-thumb { background: #444 !important; } ::-webkit-scrollbar-thumb:hover { background: #555 !important; } /* Selection */ ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; } ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; } '; } function removeDarkMode() { var style = document.getElementById('universal-dark-mode-style'); if (style) style.remove(); } // Toggle with Alt+Shift+D document.addEventListener('keydown', function(e) { if (e.altKey && e.shiftKey && e.key === 'D') { e.preventDefault(); enabled = !enabled; if (enabled) { applyDarkMode(); console.log('[Universal Dark Mode] Enabled'); } else { removeDarkMode(); console.log('[Universal Dark Mode] Disabled'); } } }); // Apply on load applyDarkMode(); // Re-apply on dynamic content var observer = new MutationObserver(function(mutations) { if (enabled && !document.getElementById('universal-dark-mode-style')) { applyDarkMode(); } }); observer.observe(document.head, { childList: true }); console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle'); })(); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })(); Ensure file is created when it doesnt exist by auxym · Pull Request #5 · auxym/zarr-sqlite-python · GitHub
Skip to content
Merged
13 changes: 13 additions & 0 deletions pyproject.toml
Original file line numberDiff line numberDiff line change
Expand Up@@ -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",
]
21 changes: 0 additions & 21 deletions src/zarr_sqlite/scratch.py

This file was deleted.

143 changes: 96 additions & 47 deletions src/zarr_sqlite/zarr_sqlite.py
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,57 @@
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,
RangeByteRequest,
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):
Expand DownExpand Up@@ -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)
Expand All@@ -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):
Expand All@@ -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:"):
Expand DownExpand Up@@ -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}")
Expand All@@ -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)
Expand DownExpand Up@@ -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()

Expand DownExpand Up@@ -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]
Expand All@@ -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(
Expand All@@ -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
Expand All@@ -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])
Loading