diff --git a/burr/integrations/persisters/b_asyncpg.py b/burr/integrations/persisters/b_asyncpg.py index c694350d5..c47e34e39 100644 --- a/burr/integrations/persisters/b_asyncpg.py +++ b/burr/integrations/persisters/b_asyncpg.py @@ -370,7 +370,7 @@ async def load( async def save( self, - partition_key: str, + partition_key: Optional[str], app_id: str, sequence_id: int, position: str, @@ -395,6 +395,8 @@ async def save( before the action was applied. :return: None """ + if partition_key is None: + partition_key = self.PARTITION_KEY_DEFAULT logger.debug( "saving %s, %s, %s, %s, %s, %s", partition_key, diff --git a/burr/integrations/persisters/b_psycopg2.py b/burr/integrations/persisters/b_psycopg2.py index 26425f805..7c0430a7f 100644 --- a/burr/integrations/persisters/b_psycopg2.py +++ b/burr/integrations/persisters/b_psycopg2.py @@ -216,7 +216,7 @@ def load( def save( self, - partition_key: str, + partition_key: Optional[str], app_id: str, sequence_id: int, position: str, @@ -241,6 +241,8 @@ def save( before the action was applied. :return: None """ + if partition_key is None: + partition_key = self.PARTITION_KEY_DEFAULT logger.debug( "saving %s, %s, %s, %s, %s, %s", partition_key, diff --git a/tests/integrations/persisters/test_postgresql.py b/tests/integrations/persisters/test_postgresql.py index 0e8cd700a..0100cd93e 100644 --- a/tests/integrations/persisters/test_postgresql.py +++ b/tests/integrations/persisters/test_postgresql.py @@ -49,6 +49,13 @@ def test_save_and_load_state(postgresql_persister): assert data["state"].get_all() == {"a": 1, "b": 2} +def test_save_and_load_with_default_partition_key(postgresql_persister): + """An application without a partition key saves and loads with partition_key=None.""" + postgresql_persister.save(None, "no_pk_app", 1, "pos", state.State({"a": 1}), "completed") + data = postgresql_persister.load(None, "no_pk_app") + assert data["state"].get_all() == {"a": 1} + + def test_list_app_ids(postgresql_persister): postgresql_persister.save("pk", "app_id1", 1, "pos1", state.State({"a": 1}), "completed") postgresql_persister.save("pk", "app_id2", 2, "pos2", state.State({"b": 2}), "completed") @@ -146,6 +153,15 @@ async def test_async_save_and_load_state(asyncpostgresql_persister): assert data["state"].get_all() == {"a": 1, "b": 2} +async def test_async_save_and_load_with_default_partition_key(asyncpostgresql_persister): + """An application without a partition key saves and loads with partition_key=None.""" + await asyncpostgresql_persister.save( + None, "no_pk_app", 1, "pos", state.State({"a": 1}), "completed" + ) + data = await asyncpostgresql_persister.load(None, "no_pk_app") + assert data["state"].get_all() == {"a": 1} + + async def test_async_list_app_ids(asyncpostgresql_persister): await asyncpostgresql_persister.save( "pk", "app_id1", 1, "pos1", state.State({"a": 1}), "completed"