Skip to content

Commit 628cfde

Browse files
committed
support async callbacks
1 parent 0b4a65f commit 628cfde

4 files changed

Lines changed: 70 additions & 2 deletions

File tree

‎packages/google-cloud-bigtable/google/cloud/bigtable/data/_async/mutations_batcher.py‎

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616

1717
import atexit
1818
import concurrent.futures
19+
import inspect
1920
import logging
2021
import time
2122
import warnings
@@ -444,9 +445,21 @@ async def _execute_mutate_rows(
444445
await self._flow_control.remove_from_flow(batch)
445446

446447
# Call batch done callback with list of statuses.
448+
# This is an internal callback used by the legacy synchronous shim,
449+
# though async callbacks are awaited when running in async mode.
447450
if self._user_batch_completed_callback:
448451
try:
449-
self._user_batch_completed_callback(statuses)
452+
result = self._user_batch_completed_callback(statuses)
453+
if CrossSync.is_async:
454+
if inspect.isawaitable(result):
455+
await result
456+
else:
457+
if inspect.isawaitable(result):
458+
if inspect.iscoroutine(result):
459+
result.close()
460+
raise TypeError(
461+
"_user_batch_completed_callback must be a synchronous callable"
462+
)
450463
except Exception as exc:
451464
_LOGGER.warning(
452465
f"Exception raised in user batch completion callback: {exc}"

‎packages/google-cloud-bigtable/google/cloud/bigtable/data/_sync_autogen/mutations_batcher.py‎

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919

2020
import atexit
2121
import concurrent.futures
22+
import inspect
2223
import logging
2324
import time
2425
import warnings
@@ -386,7 +387,13 @@ def _execute_mutate_rows(
386387
self._flow_control.remove_from_flow(batch)
387388
if self._user_batch_completed_callback:
388389
try:
389-
self._user_batch_completed_callback(statuses)
390+
result = self._user_batch_completed_callback(statuses)
391+
if inspect.isawaitable(result):
392+
if inspect.iscoroutine(result):
393+
result.close()
394+
raise TypeError(
395+
"_user_batch_completed_callback must be a synchronous callable"
396+
)
390397
except Exception as exc:
391398
_LOGGER.warning(
392399
f"Exception raised in user batch completion callback: {exc}"

‎packages/google-cloud-bigtable/tests/unit/data/_async/test_mutations_batcher.py‎

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1075,6 +1075,31 @@ async def test__execute_mutate_rows_batch_completed_callback_exception(self):
10751075
callback.assert_called_once()
10761076
assert result == []
10771077

1078+
@CrossSync.pytest
1079+
async def test__execute_mutate_rows_batch_completed_callback_coroutine(self):
1080+
from google.rpc import code_pb2, status_pb2
1081+
1082+
with mock.patch.object(CrossSync, "_MutateRowsOperation") as mutate_rows:
1083+
mutate_rows.return_value = CrossSync.Mock()
1084+
table = mock.Mock()
1085+
table.default_mutate_rows_operation_timeout = 17
1086+
table.default_mutate_rows_attempt_timeout = 13
1087+
table.default_mutate_rows_retryable_errors = ()
1088+
called_with = []
1089+
async_callback = mock.AsyncMock(
1090+
side_effect=lambda statuses: called_with.append(statuses)
1091+
)
1092+
1093+
async with self._make_one(table) as instance:
1094+
instance._user_batch_completed_callback = async_callback
1095+
batch = [self._make_mutation()]
1096+
result = await instance._execute_mutate_rows(batch, mock.Mock())
1097+
assert result == []
1098+
if CrossSync.is_async:
1099+
assert called_with == [[status_pb2.Status(code=code_pb2.OK)]]
1100+
else:
1101+
assert called_with == []
1102+
10781103
@CrossSync.pytest
10791104
async def test__raise_exceptions(self):
10801105
"""Raise exceptions and reset error state"""

‎packages/google-cloud-bigtable/tests/unit/data/_sync_autogen/test_mutations_batcher.py‎

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -957,6 +957,29 @@ def test__execute_mutate_rows_batch_completed_callback_exception(self):
957957
callback.assert_called_once()
958958
assert result == []
959959

960+
def test__execute_mutate_rows_batch_completed_callback_coroutine(self):
961+
from google.rpc import code_pb2, status_pb2
962+
963+
with mock.patch.object(
964+
CrossSync._Sync_Impl, "_MutateRowsOperation"
965+
) as mutate_rows:
966+
mutate_rows.return_value = CrossSync._Sync_Impl.Mock()
967+
table = mock.Mock()
968+
table.default_mutate_rows_operation_timeout = 17
969+
table.default_mutate_rows_attempt_timeout = 13
970+
table.default_mutate_rows_retryable_errors = ()
971+
called_with = []
972+
async_callback = mock.AsyncMock(
973+
side_effect=lambda statuses: called_with.append(statuses)
974+
)
975+
976+
with self._make_one(table) as instance:
977+
instance._user_batch_completed_callback = async_callback
978+
batch = [self._make_mutation()]
979+
result = instance._execute_mutate_rows(batch, mock.Mock())
980+
assert result == []
981+
assert called_with == []
982+
960983
def test__raise_exceptions(self):
961984
"""Raise exceptions and reset error state"""
962985
from google.cloud.bigtable.data import exceptions

0 commit comments

Comments
 (0)