diff --git a/burr/integrations/persisters/b_redis.py b/burr/integrations/persisters/b_redis.py index 5f091dc43..79e57f796 100644 --- a/burr/integrations/persisters/b_redis.py +++ b/burr/integrations/persisters/b_redis.py @@ -54,6 +54,8 @@ class RedisBasePersister(persistence.BaseStatePersister): so this is an attempt to fix that in a backwards compatible way. """ + PARTITION_KEY_DEFAULT = "" + @classmethod def from_config(cls, config: dict) -> "RedisBasePersister": """Creates a new instance of the RedisBasePersister from a configuration dictionary.""" @@ -105,14 +107,16 @@ def set_serde_kwargs(self, serde_kwargs: dict): """Sets the serde_kwargs for the persister.""" self.serde_kwargs = serde_kwargs - def list_app_ids(self, partition_key: str, **kwargs) -> list[str]: + def list_app_ids(self, partition_key: Optional[str], **kwargs) -> list[str]: """List the app ids for a given partition key.""" + if partition_key is None: + partition_key = self.PARTITION_KEY_DEFAULT namespaced_partition_key = add_namespace_to_partition_key(partition_key, self.namespace) app_ids = self.connection.zrevrange(namespaced_partition_key, 0, -1) return [app_id.decode() for app_id in app_ids] def load( - self, partition_key: str, app_id: str, sequence_id: int = None, **kwargs + self, partition_key: Optional[str], app_id: str, sequence_id: int = None, **kwargs ) -> Optional[persistence.PersistedStateData]: """Load the state data for a given partition key, app id, and sequence id. @@ -124,6 +128,8 @@ def load( :param kwargs: :return: Value or None. """ + if partition_key is None: + partition_key = self.PARTITION_KEY_DEFAULT namespaced_partition_key = add_namespace_to_partition_key(partition_key, self.namespace) if sequence_id is None: sequence_id = self.connection.zscore(namespaced_partition_key, app_id) @@ -153,7 +159,7 @@ def create_key(self, app_id, partition_key, sequence_id): def save( self, - partition_key: str, + partition_key: Optional[str], app_id: str, sequence_id: int, position: str, @@ -172,6 +178,8 @@ def save( :param kwargs: :return: """ + if partition_key is None: + partition_key = self.PARTITION_KEY_DEFAULT key = self.create_key(app_id, partition_key, sequence_id) if self.connection.exists(key): raise ValueError(f"partition_key:app_id:sequence_id[{key}] already exists.") @@ -235,6 +243,8 @@ class AsyncRedisBasePersister(persistence.AsyncBaseStatePersister): It inherits from the AsyncBaseStatePersister class. """ + PARTITION_KEY_DEFAULT = "" + @classmethod def from_config(cls, config: dict) -> "AsyncRedisBasePersister": """Creates a new instance of the RedisBasePersister from a configuration dictionary.""" @@ -286,14 +296,16 @@ def set_serde_kwargs(self, serde_kwargs: dict): """Sets the serde_kwargs for the persister.""" self.serde_kwargs = serde_kwargs - async def list_app_ids(self, partition_key: str, **kwargs) -> list[str]: + async def list_app_ids(self, partition_key: Optional[str], **kwargs) -> list[str]: """List the app ids for a given partition key.""" + if partition_key is None: + partition_key = self.PARTITION_KEY_DEFAULT namespaced_partition_key = add_namespace_to_partition_key(partition_key, self.namespace) app_ids = await self.connection.zrevrange(namespaced_partition_key, 0, -1) return [app_id.decode() for app_id in app_ids] async def load( - self, partition_key: str, app_id: str, sequence_id: int = None, **kwargs + self, partition_key: Optional[str], app_id: str, sequence_id: int = None, **kwargs ) -> Optional[persistence.PersistedStateData]: """Load the state data for a given partition key, app id, and sequence id. @@ -305,6 +317,8 @@ async def load( :param kwargs: :return: Value or None. """ + if partition_key is None: + partition_key = self.PARTITION_KEY_DEFAULT namespaced_partition_key = add_namespace_to_partition_key(partition_key, self.namespace) if sequence_id is None: sequence_id = await self.connection.zscore(namespaced_partition_key, app_id) @@ -334,7 +348,7 @@ def create_key(self, app_id, partition_key, sequence_id): async def save( self, - partition_key: str, + partition_key: Optional[str], app_id: str, sequence_id: int, position: str, @@ -353,6 +367,8 @@ async def save( :param kwargs: :return: """ + if partition_key is None: + partition_key = self.PARTITION_KEY_DEFAULT key = self.create_key(app_id, partition_key, sequence_id) if await self.connection.exists(key): raise ValueError(f"partition_key:app_id:sequence_id[{key}] already exists.") diff --git a/tests/integrations/persisters/test_b_redis.py b/tests/integrations/persisters/test_b_redis.py index b41b92ca5..f6e35da35 100644 --- a/tests/integrations/persisters/test_b_redis.py +++ b/tests/integrations/persisters/test_b_redis.py @@ -51,6 +51,25 @@ def test_save_and_load_state(redis_persister): assert data["state"].get_all() == {"a": 1, "b": 2} +def test_save_and_load_with_default_partition_key(redis_persister): + """An application without a partition key saves and loads with partition_key=None.""" + redis_persister.save(None, "no_pk_app", 1, "pos", state.State({"a": 1}), "completed") + data = redis_persister.load(None, "no_pk_app") + assert data["state"].get_all() == {"a": 1} + # None is stored under the default partition key, so it reads back as "" as well + assert data["partition_key"] == "" + assert redis_persister.load("", "no_pk_app")["sequence_id"] == 1 + assert "no_pk_app" in redis_persister.list_app_ids(None) + + +def test_save_and_load_with_default_partition_key_with_ns(redis_persister_with_ns): + """An application without a partition key saves and loads with partition_key=None.""" + redis_persister_with_ns.save(None, "no_pk_app", 1, "pos", state.State({"a": 1}), "completed") + data = redis_persister_with_ns.load(None, "no_pk_app") + assert data["state"].get_all() == {"a": 1} + assert "no_pk_app" in redis_persister_with_ns.list_app_ids(None) + + def test_list_app_ids(redis_persister): redis_persister.save("pk", "app_id1", 1, "pos1", state.State({"a": 1}), "completed") redis_persister.save("pk", "app_id2", 2, "pos2", state.State({"b": 2}), "completed") @@ -136,6 +155,16 @@ async def test_async_save_and_load_state(async_redis_persister): assert data["state"].get_all() == {"a": 1, "b": 2} +async def test_async_save_and_load_with_default_partition_key(async_redis_persister): + """An application without a partition key saves and loads with partition_key=None.""" + await async_redis_persister.save( + None, "no_pk_app", 1, "pos", state.State({"a": 1}), "completed" + ) + data = await async_redis_persister.load(None, "no_pk_app") + assert data["state"].get_all() == {"a": 1} + assert "no_pk_app" in await async_redis_persister.list_app_ids(None) + + async def test_async_list_app_ids(async_redis_persister): await async_redis_persister.save("pk", "app_id1", 1, "pos1", state.State({"a": 1}), "completed") await async_redis_persister.save("pk", "app_id2", 2, "pos2", state.State({"b": 2}), "completed")