From 50444af982ec09989daee71c8621aeed8a2727c0 Mon Sep 17 00:00:00 2001 From: Brian Bjarke Jensen Date: Wed, 8 Jul 2026 21:06:13 +0200 Subject: [PATCH] Add Redis scan_keys iterator API. Provide a SCAN-based key iterator for Redis adapters so callers can enumerate large keyspaces without relying on blocking KEYS lookups. Co-authored-by: Cursor --- README.md | 12 +++++ python_repositories/adapters/redis_adapter.py | 35 +++++++++++-- .../interfaces/json_repository_interface.py | 11 ++++ tests/integration/redis_adapter_test.py | 29 +++++++++++ tests/unit/json_repository_interface_test.py | 21 ++++++++ tests/unit/redis_adapter_test.py | 50 +++++++++++++++++++ 6 files changed, 155 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index 4542d32..6149a2c 100644 --- a/README.md +++ b/README.md @@ -34,6 +34,18 @@ Requires Redis with the RedisJSON module (e.g. redis-stack). | -------------------- | ---------------------------------------------------- | | `REDIS_URI` | Redis connection URL (e.g. `redis://localhost:6379`) | +For key discovery: + +- `list_keys(pattern)` is simple and returns a `list[str]`, but it uses Redis `KEYS` and may block on large datasets. +- `scan_keys(pattern, *, count=None)` is preferred for production use and yields keys incrementally via Redis `SCAN`. + +Example: + +```python +for key in repo.scan_keys("user:*"): + print(key) +``` + ### MinIO (`ObjectRepositoryInterface`) | Environment variable | Description | diff --git a/python_repositories/adapters/redis_adapter.py b/python_repositories/adapters/redis_adapter.py index 47a88a4..4091c98 100644 --- a/python_repositories/adapters/redis_adapter.py +++ b/python_repositories/adapters/redis_adapter.py @@ -1,6 +1,7 @@ """Definition of RedisAdapter class.""" from __future__ import annotations +from collections.abc import Iterator from typing import cast from python_repositories.adapters.connection_aware_adapter import ( @@ -136,11 +137,13 @@ class RedisAdapter(JsonRepositoryInterface, ConnectionAwareAdapter): self._client.json().delete(key) self.logger.debug(f"Deleted {key}") - def list_keys(self, pattern: str) -> list[str]: - """List keys in Redis matching a pattern.""" - # Check input + def _validate_pattern(self, pattern: str) -> None: if not isinstance(pattern, str) or len(pattern) == 0: raise ValueError("Pattern must be a non-empty string") + + def list_keys(self, pattern: str) -> list[str]: + """List keys in Redis using KEYS; may block on large datasets.""" + self._validate_pattern(pattern) # Check connection self._require_connected() assert self._client is not None @@ -152,3 +155,29 @@ class RedisAdapter(JsonRepositoryInterface, ConnectionAwareAdapter): keys: list[str] = [key.decode(self.encoding) for key in keys_raw] self.logger.debug(f"Got {keys} matching {pattern}") return keys + + def scan_keys( + self, + pattern: str, + *, + count: int | None = None, + ) -> Iterator[str]: + """Yield keys in Redis using SCAN to avoid blocking large datasets.""" + self._validate_pattern(pattern) + self._require_connected() + assert self._client is not None + client = self._client + + def _decode(key: bytes | str) -> str: + return key if isinstance(key, str) else key.decode(self.encoding) + + def _iter() -> Iterator[str]: + scan_iter = ( + client.scan_iter(match=pattern, count=count) + if count is not None + else client.scan_iter(match=pattern) + ) + for key_raw in scan_iter: + yield _decode(key_raw) + + return _iter() diff --git a/python_repositories/interfaces/json_repository_interface.py b/python_repositories/interfaces/json_repository_interface.py index 1e271e4..79cb4f2 100644 --- a/python_repositories/interfaces/json_repository_interface.py +++ b/python_repositories/interfaces/json_repository_interface.py @@ -1,5 +1,6 @@ """Definition of JsonRepositoryInterface abstract base class.""" +from collections.abc import Iterator from abc import ABC, abstractmethod @@ -25,3 +26,13 @@ class JsonRepositoryInterface(ABC): def list_keys(self, pattern: str) -> list[str]: """List keys matching a glob pattern.""" ... + + def scan_keys( + self, + pattern: str, + *, + count: int | None = None, + ) -> Iterator[str]: + """Yield keys matching a glob pattern incrementally.""" + del count + yield from self.list_keys(pattern) diff --git a/tests/integration/redis_adapter_test.py b/tests/integration/redis_adapter_test.py index 5185204..408c00f 100644 --- a/tests/integration/redis_adapter_test.py +++ b/tests/integration/redis_adapter_test.py @@ -254,5 +254,34 @@ def test_should_raise_connection_error_on_list_keys_when_not_connected( adapter.list_keys("some_pattern") +def test_should_scan_keys( + redis_adapter: RedisAdapter, +) -> None: + """Test scanning keys matching a pattern returns correct keys.""" + redis_adapter.set("key1", {"a": 1}) + redis_adapter.set("key2", {"b": 2}) + keys = set(redis_adapter.scan_keys("key*")) + assert keys == {"key1", "key2"} + + +def test_should_raise_value_error_on_invalid_scan_keys_pattern( + redis_adapter: RedisAdapter, +) -> None: + """Test that the RedisAdapter raises ValueError when scanning keys with an invalid pattern.""" + invalid_patterns = ["", 123, None] + for pattern in invalid_patterns: + with pytest.raises(ValueError): + list(redis_adapter.scan_keys(pattern)) # type: ignore[arg-type] + + +def test_should_raise_connection_error_on_scan_keys_when_not_connected( + redis_config: RedisConfig, +) -> None: + """Test that the RedisAdapter raises ConnectionError when scanning keys while not connected.""" + adapter = RedisAdapter(config=redis_config) + with pytest.raises(ConnectionError): + list(adapter.scan_keys("some_pattern")) + + if __name__ == "__main__": pytest.main(["-s", "-v", __file__]) diff --git a/tests/unit/json_repository_interface_test.py b/tests/unit/json_repository_interface_test.py index 88b98d8..f026f2c 100644 --- a/tests/unit/json_repository_interface_test.py +++ b/tests/unit/json_repository_interface_test.py @@ -80,3 +80,24 @@ def test_instantiation_fails_when_list_keys_not_implemented() -> None: with pytest.raises(TypeError): _ = Incomplete() # type: ignore + + +def test_scan_keys_defaults_to_list_keys() -> None: + """Test that the default scan_keys implementation delegates to list_keys.""" + + class Complete(JsonRepositoryInterface): + def get(self, key: str) -> dict | None: + return None + + def set(self, key: str, data: dict) -> None: + pass + + def delete(self, key: str) -> None: + pass + + def list_keys(self, pattern: str) -> list[str]: + return [f"{pattern}-1", f"{pattern}-2"] + + repository = Complete() + + assert list(repository.scan_keys("user")) == ["user-1", "user-2"] diff --git a/tests/unit/redis_adapter_test.py b/tests/unit/redis_adapter_test.py index 140d45e..d9bf888 100644 --- a/tests/unit/redis_adapter_test.py +++ b/tests/unit/redis_adapter_test.py @@ -114,3 +114,53 @@ def test_connect_closes_existing_non_injected_client( stale_client.close.assert_called_once() assert adapter._client is new_client + + +def test_scan_keys_yields_decoded_keys() -> None: + mock_client = MagicMock(spec=redis.Redis) + mock_client.scan_iter.return_value = iter([b"key1", b"key2"]) + adapter = RedisAdapter(config=TEST_REDIS_CONFIG, client=mock_client) + + keys = list(adapter.scan_keys("key*")) + + assert keys == ["key1", "key2"] + mock_client.scan_iter.assert_called_once_with(match="key*") + mock_client.keys.assert_not_called() + + +def test_scan_keys_forwards_count() -> None: + mock_client = MagicMock(spec=redis.Redis) + mock_client.scan_iter.return_value = iter([b"key1"]) + adapter = RedisAdapter(config=TEST_REDIS_CONFIG, client=mock_client) + + keys = list(adapter.scan_keys("key*", count=50)) + + assert keys == ["key1"] + mock_client.scan_iter.assert_called_once_with(match="key*", count=50) + + +def test_list_keys_raises_value_error_on_invalid_pattern() -> None: + mock_client = MagicMock(spec=redis.Redis) + adapter = RedisAdapter(config=TEST_REDIS_CONFIG, client=mock_client) + + invalid_patterns = ["", 123, None] + for pattern in invalid_patterns: + with pytest.raises(ValueError, match="Pattern must be a non-empty string"): + adapter.list_keys(pattern) # type: ignore[arg-type] + + +def test_scan_keys_raises_value_error_on_invalid_pattern() -> None: + mock_client = MagicMock(spec=redis.Redis) + adapter = RedisAdapter(config=TEST_REDIS_CONFIG, client=mock_client) + + invalid_patterns = ["", 123, None] + for pattern in invalid_patterns: + with pytest.raises(ValueError, match="Pattern must be a non-empty string"): + list(adapter.scan_keys(pattern)) # type: ignore[arg-type] + + +def test_scan_keys_raises_connection_error_when_not_connected() -> None: + adapter = RedisAdapter(config=TEST_REDIS_CONFIG) + + with pytest.raises(ConnectionError): + list(adapter.scan_keys("key*"))