File tree Expand file tree Collapse file tree
packages/google-cloud-bigtable
google/cloud/bigtable/data Expand file tree Collapse file tree Original file line number Diff line number Diff line change 1616
1717import atexit
1818import concurrent .futures
19+ import inspect
1920import logging
2021import time
2122import 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 } "
Original file line number Diff line number Diff line change 1919
2020import atexit
2121import concurrent .futures
22+ import inspect
2223import logging
2324import time
2425import 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 } "
Original file line number Diff line number Diff 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"""
Original file line number Diff line number Diff 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
You can’t perform that action at this time.
0 commit comments