|
34 | 34 | __CROSS_SYNC_OUTPUT__ = "tests.unit.data._sync_autogen.test_metrics_interceptor" |
35 | 35 |
|
36 | 36 |
|
37 | | -@CrossSync.drop |
38 | | -class _AsyncIterator: |
39 | | - """Helper class to wrap an iterator or async generator in an async iterator""" |
40 | 37 |
|
41 | | - def __init__(self, iterable): |
42 | | - if hasattr(iterable, "__anext__"): |
43 | | - self._iterator = iterable |
44 | | - else: |
45 | | - self._iterator = iter(iterable) |
46 | | - |
47 | | - def __aiter__(self): |
48 | | - return self |
49 | | - |
50 | | - async def __anext__(self): |
51 | | - if hasattr(self._iterator, "__anext__"): |
52 | | - return await self._iterator.__anext__() |
53 | | - try: |
54 | | - return next(self._iterator) |
55 | | - except StopIteration: |
56 | | - raise StopAsyncIteration |
57 | 38 |
|
58 | 39 |
|
59 | 40 | @CrossSync.convert_class(sync_name="TestMetricsInterceptor") |
@@ -273,7 +254,10 @@ async def test_unary_stream_interceptor_success(self): |
273 | 254 | continuation = CrossSync.Mock() |
274 | 255 | call = continuation.return_value |
275 | 256 | if CrossSync.is_async: |
276 | | - call.__aiter__ = mock.Mock(return_value=_AsyncIterator([1, 2])) |
| 257 | + async def gen(): |
| 258 | + yield 1 |
| 259 | + yield 2 |
| 260 | + call.__aiter__ = mock.Mock(return_value=gen()) |
277 | 261 | else: |
278 | 262 | call.__iter__ = mock.Mock(return_value=iter([1, 2])) |
279 | 263 | call.trailing_metadata = CrossSync.Mock(return_value=[("a", "b")]) |
@@ -311,7 +295,7 @@ async def test_unary_stream_interceptor_failure_mid_stream(self): |
311 | 295 | async def mock_generator(): |
312 | 296 | yield 1 |
313 | 297 | raise exc |
314 | | - call.__aiter__ = mock.Mock(return_value=_AsyncIterator(mock_generator())) |
| 298 | + call.__aiter__ = mock.Mock(return_value=mock_generator()) |
315 | 299 | else: |
316 | 300 | def mock_generator(): |
317 | 301 | yield 1 |
@@ -464,7 +448,10 @@ async def test_unary_stream_interceptor_start_operation(self, initial_state): |
464 | 448 | continuation = CrossSync.Mock() |
465 | 449 | call = continuation.return_value |
466 | 450 | if CrossSync.is_async: |
467 | | - call.__aiter__ = mock.Mock(return_value=_AsyncIterator([1, 2])) |
| 451 | + async def gen(): |
| 452 | + yield 1 |
| 453 | + yield 2 |
| 454 | + call.__aiter__ = mock.Mock(return_value=gen()) |
468 | 455 | else: |
469 | 456 | call.__iter__ = mock.Mock(return_value=iter([1, 2])) |
470 | 457 | call.trailing_metadata = CrossSync.Mock(return_value=[]) |
|
0 commit comments