1# Copyright 2019 gRPC authors.
2#
3# Licensed under the Apache License, Version 2.0 (the "License");
4# you may not use this file except in compliance with the License.
5# You may obtain a copy of the License at
6#
7# http://www.apache.org/licenses/LICENSE-2.0
8#
9# Unless required by applicable law or agreed to in writing, software
10# distributed under the License is distributed on an "AS IS" BASIS,
11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12# See the License for the specific language governing permissions and
13# limitations under the License.
14"""Interceptors implementation of gRPC Asyncio Python."""
15
16from __future__ import annotations
17
18from abc import ABCMeta
19from abc import abstractmethod
20import asyncio
21import collections
22import functools
23from typing import (
24 AsyncIterable,
25 AsyncIterator,
26 Awaitable,
27 Callable,
28 List,
29 Optional,
30 Sequence,
31 Union,
32)
33
34import grpc
35from grpc._cython import cygrpc
36
37from . import _base_call
38from ._call import AioRpcError
39from ._call import StreamStreamCall
40from ._call import StreamUnaryCall
41from ._call import UnaryStreamCall
42from ._call import UnaryUnaryCall
43from ._call import _API_STYLE_ERROR
44from ._call import _RPC_ALREADY_FINISHED_DETAILS
45from ._call import _RPC_HALF_CLOSED_DETAILS
46from ._metadata import Metadata
47from ._typing import DeserializingFunction
48from ._typing import DoneCallbackType
49from ._typing import EOFType
50from ._typing import RequestIterableType
51from ._typing import RequestType
52from ._typing import ResponseIterableType
53from ._typing import ResponseType
54from ._typing import SerializingFunction
55from ._utils import _timeout_to_deadline
56
57_LOCAL_CANCELLATION_DETAILS = "Locally cancelled by application!"
58
59
60class ServerInterceptor(metaclass=ABCMeta):
61 """Affords intercepting incoming RPCs on the service-side.
62
63 This is an EXPERIMENTAL API.
64 """
65
66 @abstractmethod
67 async def intercept_service(
68 self,
69 continuation: Callable[
70 [grpc.HandlerCallDetails], Awaitable[grpc.RpcMethodHandler]
71 ],
72 handler_call_details: grpc.HandlerCallDetails,
73 ) -> grpc.RpcMethodHandler:
74 """Intercepts incoming RPCs before handing them over to a handler.
75
76 State can be passed from an interceptor to downstream interceptors
77 via contextvars. The first interceptor is called from an empty
78 contextvars.Context, and the same Context is used for downstream
79 interceptors and for the final handler call. Note that there are no
80 guarantees that interceptors and handlers will be called from the
81 same thread.
82
83 Args:
84 continuation: A function that takes a HandlerCallDetails and
85 proceeds to invoke the next interceptor in the chain, if any,
86 or the RPC handler lookup logic, with the call details passed
87 as an argument, and returns an RpcMethodHandler instance if
88 the RPC is considered serviced, or None otherwise.
89 handler_call_details: A HandlerCallDetails describing the RPC.
90
91 Returns:
92 An RpcMethodHandler with which the RPC may be serviced if the
93 interceptor chooses to service this RPC, or None otherwise.
94 """
95
96
97class ClientCallDetails(
98 collections.namedtuple(
99 "ClientCallDetails",
100 ("method", "timeout", "metadata", "credentials", "wait_for_ready"),
101 ),
102 grpc.ClientCallDetails,
103):
104 """Describes an RPC to be invoked.
105
106 This is an EXPERIMENTAL API.
107
108 Args:
109 method: The method name of the RPC.
110 timeout: An optional duration of time in seconds to allow for the RPC.
111 metadata: Optional metadata to be transmitted to the service-side of
112 the RPC.
113 credentials: An optional CallCredentials for the RPC.
114 wait_for_ready: An optional flag to enable :term:`wait_for_ready` mechanism.
115 """
116
117 method: bytes
118 timeout: Optional[float]
119 metadata: Optional[Metadata]
120 credentials: Optional[grpc.CallCredentials]
121 wait_for_ready: Optional[bool]
122
123
124class ClientInterceptor(metaclass=ABCMeta):
125 """Base class used for all Aio Client Interceptor classes"""
126
127
128class UnaryUnaryClientInterceptor(ClientInterceptor, metaclass=ABCMeta):
129 """Affords intercepting unary-unary invocations."""
130
131 @abstractmethod
132 async def intercept_unary_unary(
133 self,
134 continuation: Callable[
135 [ClientCallDetails, RequestType], UnaryUnaryCall
136 ],
137 client_call_details: ClientCallDetails,
138 request: RequestType,
139 ) -> Union[UnaryUnaryCall, ResponseType]:
140 """Intercepts a unary-unary invocation asynchronously.
141
142 Args:
143 continuation: A coroutine that proceeds with the invocation by
144 executing the next interceptor in the chain or invoking the
145 actual RPC on the underlying Channel. It is the interceptor's
146 responsibility to call it if it decides to move the RPC forward.
147 The interceptor can use
148 `call = await continuation(client_call_details, request)`
149 to continue with the RPC. `continuation` returns the call to the
150 RPC.
151 client_call_details: A ClientCallDetails object describing the
152 outgoing RPC.
153 request: The request value for the RPC.
154
155 Returns:
156 An object with the RPC response.
157
158 Raises:
159 AioRpcError: Indicating that the RPC terminated with non-OK status.
160 asyncio.CancelledError: Indicating that the RPC was canceled.
161 """
162
163
164class UnaryStreamClientInterceptor(ClientInterceptor, metaclass=ABCMeta):
165 """Affords intercepting unary-stream invocations."""
166
167 @abstractmethod
168 async def intercept_unary_stream(
169 self,
170 continuation: Callable[
171 [ClientCallDetails, RequestType], UnaryStreamCall
172 ],
173 client_call_details: ClientCallDetails,
174 request: RequestType,
175 ) -> Union[ResponseIterableType, UnaryStreamCall]:
176 """Intercepts a unary-stream invocation asynchronously.
177
178 The function could return the call object or an asynchronous
179 iterator, in case of being an asyncrhonous iterator this will
180 become the source of the reads done by the caller.
181
182 Args:
183 continuation: A coroutine that proceeds with the invocation by
184 executing the next interceptor in the chain or invoking the
185 actual RPC on the underlying Channel. It is the interceptor's
186 responsibility to call it if it decides to move the RPC forward.
187 The interceptor can use
188 `call = await continuation(client_call_details, request)`
189 to continue with the RPC. `continuation` returns the call to the
190 RPC.
191 client_call_details: A ClientCallDetails object describing the
192 outgoing RPC.
193 request: The request value for the RPC.
194
195 Returns:
196 The RPC Call or an asynchronous iterator.
197
198 Raises:
199 AioRpcError: Indicating that the RPC terminated with non-OK status.
200 asyncio.CancelledError: Indicating that the RPC was canceled.
201 """
202
203
204class StreamUnaryClientInterceptor(ClientInterceptor, metaclass=ABCMeta):
205 """Affords intercepting stream-unary invocations."""
206
207 @abstractmethod
208 async def intercept_stream_unary(
209 self,
210 continuation: Callable[
211 [ClientCallDetails, RequestType], StreamUnaryCall
212 ],
213 client_call_details: ClientCallDetails,
214 request_iterator: RequestIterableType,
215 ) -> StreamUnaryCall:
216 """Intercepts a stream-unary invocation asynchronously.
217
218 Within the interceptor the usage of the call methods like `write` or
219 even awaiting the call should be done carefully, since the caller
220 could be expecting an untouched call, for example for start writing
221 messages to it.
222
223 Args:
224 continuation: A coroutine that proceeds with the invocation by
225 executing the next interceptor in the chain or invoking the
226 actual RPC on the underlying Channel. It is the interceptor's
227 responsibility to call it if it decides to move the RPC forward.
228 The interceptor can use
229 `call = await continuation(client_call_details, request_iterator)`
230 to continue with the RPC. `continuation` returns the call to the
231 RPC.
232 client_call_details: A ClientCallDetails object describing the
233 outgoing RPC.
234 request_iterator: The request iterator that will produce requests
235 for the RPC.
236
237 Returns:
238 The RPC Call.
239
240 Raises:
241 AioRpcError: Indicating that the RPC terminated with non-OK status.
242 asyncio.CancelledError: Indicating that the RPC was canceled.
243 """
244
245
246class StreamStreamClientInterceptor(ClientInterceptor, metaclass=ABCMeta):
247 """Affords intercepting stream-stream invocations."""
248
249 @abstractmethod
250 async def intercept_stream_stream(
251 self,
252 continuation: Callable[
253 [ClientCallDetails, RequestType], StreamStreamCall
254 ],
255 client_call_details: ClientCallDetails,
256 request_iterator: RequestIterableType,
257 ) -> Union[ResponseIterableType, StreamStreamCall]:
258 """Intercepts a stream-stream invocation asynchronously.
259
260 Within the interceptor the usage of the call methods like `write` or
261 even awaiting the call should be done carefully, since the caller
262 could be expecting an untouched call, for example for start writing
263 messages to it.
264
265 The function could return the call object or an asynchronous
266 iterator, in case of being an asyncrhonous iterator this will
267 become the source of the reads done by the caller.
268
269 Args:
270 continuation: A coroutine that proceeds with the invocation by
271 executing the next interceptor in the chain or invoking the
272 actual RPC on the underlying Channel. It is the interceptor's
273 responsibility to call it if it decides to move the RPC forward.
274 The interceptor can use
275 `call = await continuation(client_call_details, request_iterator)`
276 to continue with the RPC. `continuation` returns the call to the
277 RPC.
278 client_call_details: A ClientCallDetails object describing the
279 outgoing RPC.
280 request_iterator: The request iterator that will produce requests
281 for the RPC.
282
283 Returns:
284 The RPC Call or an asynchronous iterator.
285
286 Raises:
287 AioRpcError: Indicating that the RPC terminated with non-OK status.
288 asyncio.CancelledError: Indicating that the RPC was canceled.
289 """
290
291
292class InterceptedCall:
293 """Base implementation for all intercepted call arities.
294
295 Interceptors might have some work to do before the RPC invocation with
296 the capacity of changing the invocation parameters, and some work to do
297 after the RPC invocation with the capacity for accessing to the wrapped
298 `UnaryUnaryCall`.
299
300 It handles also early and later cancellations, when the RPC has not even
301 started and the execution is still held by the interceptors or when the
302 RPC has finished but again the execution is still held by the interceptors.
303
304 Once the RPC is finally executed, all methods are finally done against the
305 intercepted call, being at the same time the same call returned to the
306 interceptors.
307
308 As a base class for all of the interceptors implements the logic around
309 final status, metadata and cancellation.
310 """
311
312 _interceptors_task: asyncio.Task
313 _pending_add_done_callbacks: Sequence[DoneCallbackType]
314
315 def __init__(self, interceptors_task: asyncio.Task) -> None:
316 self._interceptors_task = interceptors_task
317 self._pending_add_done_callbacks = []
318 self._interceptors_task.add_done_callback(
319 self._fire_or_add_pending_done_callbacks
320 )
321
322 def __del__(self):
323 self.cancel()
324
325 def _fire_or_add_pending_done_callbacks(
326 self, interceptors_task: asyncio.Task
327 ) -> None:
328 if not self._pending_add_done_callbacks:
329 return
330
331 call_completed = False
332
333 if interceptors_task.cancelled() or (
334 interceptors_task.done()
335 and interceptors_task.exception() is not None
336 ):
337 call_completed = True
338 else:
339 call = interceptors_task.result()
340 if call.done():
341 call_completed = True
342
343 if call_completed:
344 for callback in self._pending_add_done_callbacks:
345 callback(self)
346 else:
347 for callback in self._pending_add_done_callbacks:
348 callback = functools.partial(
349 self._wrap_add_done_callback, callback
350 )
351 call.add_done_callback(callback)
352
353 self._pending_add_done_callbacks = []
354
355 def _wrap_add_done_callback(
356 self, callback: DoneCallbackType, unused_call: _base_call.Call
357 ) -> None:
358 callback(self)
359
360 def cancel(self) -> bool:
361 if self._interceptors_task.cancelled():
362 return False
363
364 if not self._interceptors_task.done():
365 # There is no yet the intercepted call available,
366 # Trying to cancel it by using the generic Asyncio
367 # cancellation method.
368 return self._interceptors_task.cancel()
369
370 if self._interceptors_task.exception() is not None:
371 return False
372
373 call = self._interceptors_task.result()
374
375 return call.cancel()
376
377 def cancelled(self) -> bool:
378 if self._interceptors_task.cancelled():
379 return True
380 if not self._interceptors_task.done():
381 return False
382
383 exc = self._interceptors_task.exception()
384 if exc is not None:
385 if isinstance(exc, AioRpcError):
386 return exc.code() == grpc.StatusCode.CANCELLED
387 return False
388
389 call = self._interceptors_task.result()
390
391 return call.cancelled()
392
393 def done(self) -> bool:
394 if self._interceptors_task.cancelled():
395 return True
396 if not self._interceptors_task.done():
397 return False
398 if self._interceptors_task.exception() is not None:
399 return True
400
401 call = self._interceptors_task.result()
402
403 return call.done()
404
405 def add_done_callback(self, callback: DoneCallbackType) -> None:
406 if not self._interceptors_task.done():
407 self._pending_add_done_callbacks.append(callback)
408 return
409
410 if (
411 self._interceptors_task.cancelled()
412 or self._interceptors_task.exception() is not None
413 ):
414 callback(self)
415 return
416
417 call = self._interceptors_task.result()
418
419 if call.done():
420 callback(self)
421 else:
422 callback = functools.partial(self._wrap_add_done_callback, callback)
423 call.add_done_callback(callback)
424
425 def time_remaining(self) -> Optional[float]:
426 raise NotImplementedError()
427
428 async def initial_metadata(self) -> Optional[Metadata]:
429 try:
430 call = await self._interceptors_task
431 except AioRpcError as err:
432 return err.initial_metadata()
433 except asyncio.CancelledError:
434 return None
435
436 return await call.initial_metadata()
437
438 async def trailing_metadata(self) -> Optional[Metadata]:
439 try:
440 call = await self._interceptors_task
441 except AioRpcError as err:
442 return err.trailing_metadata()
443 except asyncio.CancelledError:
444 return None
445
446 return await call.trailing_metadata()
447
448 async def code(self) -> grpc.StatusCode:
449 try:
450 call = await self._interceptors_task
451 except AioRpcError as err:
452 return err.code()
453 except asyncio.CancelledError:
454 return grpc.StatusCode.CANCELLED
455
456 return await call.code()
457
458 async def details(self) -> str:
459 try:
460 call = await self._interceptors_task
461 except AioRpcError as err:
462 return err.details()
463 except asyncio.CancelledError:
464 return _LOCAL_CANCELLATION_DETAILS
465
466 return await call.details()
467
468 async def debug_error_string(self) -> Optional[str]:
469 try:
470 call = await self._interceptors_task
471 except AioRpcError as err:
472 return err.debug_error_string()
473 except asyncio.CancelledError:
474 return ""
475
476 return await call.debug_error_string()
477
478 async def wait_for_connection(self) -> None:
479 call = await self._interceptors_task
480 return await call.wait_for_connection()
481
482
483class _InterceptedUnaryResponseMixin:
484 def __await__(self):
485 call = yield from self._interceptors_task.__await__()
486 response = yield from call.__await__()
487 return response
488
489
490class _InterceptedStreamResponseMixin:
491 _response_aiter: Optional[AsyncIterable[ResponseType]]
492
493 def _init_stream_response_mixin(self) -> None:
494 # Is initialized later, otherwise if the iterator is not finally
495 # consumed a logging warning is emitted by Asyncio.
496 self._response_aiter = None
497
498 async def _wait_for_interceptor_task_response_iterator(
499 self,
500 ) -> ResponseType:
501 call = await self._interceptors_task
502 async for response in call:
503 yield response
504
505 def __aiter__(self) -> AsyncIterator[ResponseType]:
506 if self._response_aiter is None:
507 self._response_aiter = (
508 self._wait_for_interceptor_task_response_iterator()
509 )
510 return self._response_aiter
511
512 async def read(self) -> Union[EOFType, ResponseType]:
513 if self._response_aiter is None:
514 self._response_aiter = (
515 self._wait_for_interceptor_task_response_iterator()
516 )
517 try:
518 return await self._response_aiter.asend(None)
519 except StopAsyncIteration:
520 return cygrpc.EOF
521
522
523class _InterceptedStreamRequestMixin:
524 _write_to_iterator_async_gen: Optional[AsyncIterable[RequestType]]
525 _write_to_iterator_queue: Optional[asyncio.Queue]
526 _status_code_task: Optional[asyncio.Task]
527
528 _FINISH_ITERATOR_SENTINEL = object()
529
530 def _init_stream_request_mixin(
531 self, request_iterator: Optional[RequestIterableType]
532 ) -> RequestIterableType:
533 if request_iterator is None:
534 # We provide our own request iterator which is a proxy
535 # of the futures writes that will be done by the caller.
536 self._write_to_iterator_queue = asyncio.Queue(maxsize=1)
537 self._write_to_iterator_async_gen = (
538 self._proxy_writes_as_request_iterator()
539 )
540 self._status_code_task = None
541 request_iterator = self._write_to_iterator_async_gen
542 else:
543 self._write_to_iterator_queue = None
544
545 return request_iterator
546
547 async def _proxy_writes_as_request_iterator(self):
548 await self._interceptors_task
549
550 while True:
551 value = await self._write_to_iterator_queue.get()
552 if (
553 value
554 is _InterceptedStreamRequestMixin._FINISH_ITERATOR_SENTINEL
555 ):
556 break
557 yield value
558
559 async def _write_to_iterator_queue_interruptible(
560 self,
561 request: RequestType,
562 call: _base_call.Call,
563 ):
564 # Write the specified 'request' to the request iterator queue using the
565 # specified 'call' to allow for interruption of the write in the case
566 # of abrupt termination of the call.
567 if self._status_code_task is None:
568 self._status_code_task = self._loop.create_task(call.code())
569
570 await asyncio.wait(
571 (
572 self._loop.create_task(
573 self._write_to_iterator_queue.put(request)
574 ),
575 self._status_code_task,
576 ),
577 return_when=asyncio.FIRST_COMPLETED,
578 )
579
580 async def write(self, request: RequestType) -> None:
581 # If no queue was created it means that requests
582 # should be expected through an iterators provided
583 # by the caller.
584 if self._write_to_iterator_queue is None:
585 raise cygrpc.UsageError(_API_STYLE_ERROR)
586
587 try:
588 call = await self._interceptors_task
589 except (asyncio.CancelledError, AioRpcError):
590 raise asyncio.InvalidStateError(_RPC_ALREADY_FINISHED_DETAILS)
591
592 if call.done():
593 raise asyncio.InvalidStateError(_RPC_ALREADY_FINISHED_DETAILS)
594 if call._done_writing_flag:
595 raise asyncio.InvalidStateError(_RPC_HALF_CLOSED_DETAILS)
596
597 await self._write_to_iterator_queue_interruptible(request, call)
598
599 if call.done():
600 raise asyncio.InvalidStateError(_RPC_ALREADY_FINISHED_DETAILS)
601
602 async def done_writing(self) -> None:
603 """Signal peer that client is done writing.
604
605 This method is idempotent.
606 """
607 # If no queue was created it means that requests
608 # should be expected through an iterators provided
609 # by the caller.
610 if self._write_to_iterator_queue is None:
611 raise cygrpc.UsageError(_API_STYLE_ERROR)
612
613 try:
614 call = await self._interceptors_task
615 except asyncio.CancelledError:
616 raise asyncio.InvalidStateError(_RPC_ALREADY_FINISHED_DETAILS)
617
618 await self._write_to_iterator_queue_interruptible(
619 _InterceptedStreamRequestMixin._FINISH_ITERATOR_SENTINEL, call
620 )
621
622
623def _resolve_registered_call_handle(
624 channel: cygrpc.AioChannel,
625 original_method: bytes,
626 new_method: bytes,
627 registered_call_handle: int,
628) -> int:
629 """Returns the registered call handle matching the outgoing method."""
630 resolved_registered_call_handle = registered_call_handle
631 if new_method != original_method and registered_call_handle != 0:
632 resolved_registered_call_handle = channel.get_registered_call_handle(
633 new_method,
634 )
635 return resolved_registered_call_handle
636
637
638class InterceptedUnaryUnaryCall(
639 _InterceptedUnaryResponseMixin, InterceptedCall, _base_call.UnaryUnaryCall
640):
641 """Used for running a `UnaryUnaryCall` wrapped by interceptors.
642
643 For the `__await__` method is it is proxied to the intercepted call only when
644 the interceptor task is finished.
645 """
646
647 _loop: asyncio.AbstractEventLoop
648 _channel: cygrpc.AioChannel
649
650 def __init__(
651 self,
652 interceptors: Sequence[UnaryUnaryClientInterceptor],
653 request: RequestType,
654 timeout: Optional[float],
655 metadata: Metadata,
656 credentials: Optional[grpc.CallCredentials],
657 wait_for_ready: Optional[bool],
658 channel: cygrpc.AioChannel,
659 method: bytes,
660 request_serializer: Optional[SerializingFunction],
661 response_deserializer: Optional[DeserializingFunction],
662 loop: asyncio.AbstractEventLoop,
663 registered_call_handle: int = 0,
664 ) -> None:
665 self._loop = loop
666 self._channel = channel
667 self._registered_call_handle = registered_call_handle
668 interceptors_task = loop.create_task(
669 self._invoke(
670 interceptors,
671 method,
672 timeout,
673 metadata,
674 credentials,
675 wait_for_ready,
676 request,
677 request_serializer,
678 response_deserializer,
679 )
680 )
681 super().__init__(interceptors_task)
682
683 async def _invoke(
684 self,
685 interceptors: Sequence[UnaryUnaryClientInterceptor],
686 method: bytes,
687 timeout: Optional[float],
688 metadata: Optional[Metadata],
689 credentials: Optional[grpc.CallCredentials],
690 wait_for_ready: Optional[bool],
691 request: RequestType,
692 request_serializer: Optional[SerializingFunction],
693 response_deserializer: Optional[DeserializingFunction],
694 ) -> Union[UnaryUnaryCall, UnaryUnaryCallResponse]:
695 """Run the RPC call wrapped in interceptors"""
696 client_call_details = ClientCallDetails(
697 method, timeout, metadata, credentials, wait_for_ready
698 )
699 return await self._run_interceptor(
700 list(interceptors),
701 method,
702 request_serializer,
703 response_deserializer,
704 client_call_details,
705 request,
706 )
707
708 async def _run_interceptor(
709 self,
710 interceptors: List[UnaryUnaryClientInterceptor],
711 method: bytes,
712 request_serializer: Optional[SerializingFunction],
713 response_deserializer: Optional[DeserializingFunction],
714 client_call_details: ClientCallDetails,
715 request: RequestType,
716 ) -> Union[UnaryUnaryCall, UnaryUnaryCallResponse]:
717 if interceptors:
718 continuation = functools.partial(
719 self._run_interceptor,
720 interceptors[1:],
721 method,
722 request_serializer,
723 response_deserializer,
724 )
725 call_or_response = await interceptors[0].intercept_unary_unary(
726 continuation, client_call_details, request
727 )
728
729 if isinstance(call_or_response, _base_call.UnaryUnaryCall):
730 return call_or_response
731 return UnaryUnaryCallResponse(call_or_response)
732
733 registered_call_handle = _resolve_registered_call_handle(
734 self._channel,
735 method,
736 client_call_details.method,
737 self._registered_call_handle,
738 )
739
740 return UnaryUnaryCall(
741 request,
742 _timeout_to_deadline(client_call_details.timeout),
743 client_call_details.metadata,
744 client_call_details.credentials,
745 client_call_details.wait_for_ready,
746 self._channel,
747 client_call_details.method,
748 request_serializer,
749 response_deserializer,
750 self._loop,
751 registered_call_handle,
752 )
753
754 def time_remaining(self) -> Optional[float]:
755 raise NotImplementedError()
756
757
758class InterceptedUnaryStreamCall(
759 _InterceptedStreamResponseMixin, InterceptedCall, _base_call.UnaryStreamCall
760):
761 """Used for running a `UnaryStreamCall` wrapped by interceptors."""
762
763 _loop: asyncio.AbstractEventLoop
764 _channel: cygrpc.AioChannel
765 _last_returned_call_from_interceptors = Optional[_base_call.UnaryStreamCall]
766
767 def __init__(
768 self,
769 interceptors: Sequence[UnaryStreamClientInterceptor],
770 request: RequestType,
771 timeout: Optional[float],
772 metadata: Metadata,
773 credentials: Optional[grpc.CallCredentials],
774 wait_for_ready: Optional[bool],
775 channel: cygrpc.AioChannel,
776 method: bytes,
777 request_serializer: Optional[SerializingFunction],
778 response_deserializer: Optional[DeserializingFunction],
779 loop: asyncio.AbstractEventLoop,
780 registered_call_handle: int = 0,
781 ) -> None:
782 self._loop = loop
783 self._channel = channel
784 self._registered_call_handle = registered_call_handle
785 self._init_stream_response_mixin()
786 self._last_returned_call_from_interceptors = None
787 interceptors_task = loop.create_task(
788 self._invoke(
789 interceptors,
790 method,
791 timeout,
792 metadata,
793 credentials,
794 wait_for_ready,
795 request,
796 request_serializer,
797 response_deserializer,
798 )
799 )
800 super().__init__(interceptors_task)
801
802 async def _invoke(
803 self,
804 interceptors: Sequence[UnaryStreamClientInterceptor],
805 method: bytes,
806 timeout: Optional[float],
807 metadata: Optional[Metadata],
808 credentials: Optional[grpc.CallCredentials],
809 wait_for_ready: Optional[bool],
810 request: RequestType,
811 request_serializer: Optional[SerializingFunction],
812 response_deserializer: Optional[DeserializingFunction],
813 ) -> Union[UnaryStreamCall, UnaryStreamCallResponseIterator]:
814 """Run the RPC call wrapped in interceptors"""
815 client_call_details = ClientCallDetails(
816 method, timeout, metadata, credentials, wait_for_ready
817 )
818 return await self._run_interceptor(
819 list(interceptors),
820 method,
821 request_serializer,
822 response_deserializer,
823 client_call_details,
824 request,
825 )
826
827 async def _run_interceptor(
828 self,
829 interceptors: List[UnaryStreamClientInterceptor],
830 method: bytes,
831 request_serializer: Optional[SerializingFunction],
832 response_deserializer: Optional[DeserializingFunction],
833 client_call_details: ClientCallDetails,
834 request: RequestType,
835 ) -> Union[UnaryStreamCall, UnaryStreamCallResponseIterator]:
836 if interceptors:
837 continuation = functools.partial(
838 self._run_interceptor,
839 interceptors[1:],
840 method,
841 request_serializer,
842 response_deserializer,
843 )
844
845 call_or_response_iterator = await interceptors[
846 0
847 ].intercept_unary_stream(continuation, client_call_details, request)
848
849 if isinstance(
850 call_or_response_iterator, _base_call.UnaryStreamCall
851 ):
852 self._last_returned_call_from_interceptors = (
853 call_or_response_iterator
854 )
855 else:
856 self._last_returned_call_from_interceptors = (
857 UnaryStreamCallResponseIterator(
858 self._last_returned_call_from_interceptors,
859 call_or_response_iterator,
860 )
861 )
862 return self._last_returned_call_from_interceptors
863
864 registered_call_handle = _resolve_registered_call_handle(
865 self._channel,
866 method,
867 client_call_details.method,
868 self._registered_call_handle,
869 )
870
871 self._last_returned_call_from_interceptors = UnaryStreamCall(
872 request,
873 _timeout_to_deadline(client_call_details.timeout),
874 client_call_details.metadata,
875 client_call_details.credentials,
876 client_call_details.wait_for_ready,
877 self._channel,
878 client_call_details.method,
879 request_serializer,
880 response_deserializer,
881 self._loop,
882 registered_call_handle,
883 )
884
885 return self._last_returned_call_from_interceptors
886
887 def time_remaining(self) -> Optional[float]:
888 raise NotImplementedError()
889
890
891class InterceptedStreamUnaryCall(
892 _InterceptedUnaryResponseMixin,
893 _InterceptedStreamRequestMixin,
894 InterceptedCall,
895 _base_call.StreamUnaryCall,
896):
897 """Used for running a `StreamUnaryCall` wrapped by interceptors.
898
899 For the `__await__` method is it is proxied to the intercepted call only when
900 the interceptor task is finished.
901 """
902
903 _loop: asyncio.AbstractEventLoop
904 _channel: cygrpc.AioChannel
905
906 def __init__(
907 self,
908 interceptors: Sequence[StreamUnaryClientInterceptor],
909 request_iterator: Optional[RequestIterableType],
910 timeout: Optional[float],
911 metadata: Metadata,
912 credentials: Optional[grpc.CallCredentials],
913 wait_for_ready: Optional[bool],
914 channel: cygrpc.AioChannel,
915 method: bytes,
916 request_serializer: Optional[SerializingFunction],
917 response_deserializer: Optional[DeserializingFunction],
918 loop: asyncio.AbstractEventLoop,
919 registered_call_handle: int = 0,
920 ) -> None:
921 self._loop = loop
922 self._channel = channel
923 self._registered_call_handle = registered_call_handle
924 request_iterator = self._init_stream_request_mixin(request_iterator)
925 interceptors_task = loop.create_task(
926 self._invoke(
927 interceptors,
928 method,
929 timeout,
930 metadata,
931 credentials,
932 wait_for_ready,
933 request_iterator,
934 request_serializer,
935 response_deserializer,
936 )
937 )
938 super().__init__(interceptors_task)
939
940 async def _invoke(
941 self,
942 interceptors: Sequence[StreamUnaryClientInterceptor],
943 method: bytes,
944 timeout: Optional[float],
945 metadata: Optional[Metadata],
946 credentials: Optional[grpc.CallCredentials],
947 wait_for_ready: Optional[bool],
948 request_iterator: RequestIterableType,
949 request_serializer: Optional[SerializingFunction],
950 response_deserializer: Optional[DeserializingFunction],
951 ) -> StreamUnaryCall:
952 """Run the RPC call wrapped in interceptors"""
953 client_call_details = ClientCallDetails(
954 method, timeout, metadata, credentials, wait_for_ready
955 )
956 return await self._run_interceptor(
957 list(interceptors),
958 method,
959 request_serializer,
960 response_deserializer,
961 client_call_details,
962 request_iterator,
963 )
964
965 async def _run_interceptor(
966 self,
967 interceptors: Sequence[StreamUnaryClientInterceptor],
968 method: bytes,
969 request_serializer: Optional[SerializingFunction],
970 response_deserializer: Optional[DeserializingFunction],
971 client_call_details: ClientCallDetails,
972 request_iterator: RequestIterableType,
973 ) -> _base_call.StreamUnaryCall:
974 if interceptors:
975 continuation = functools.partial(
976 self._run_interceptor,
977 interceptors[1:],
978 method,
979 request_serializer,
980 response_deserializer,
981 )
982
983 return await interceptors[0].intercept_stream_unary(
984 continuation, client_call_details, request_iterator
985 )
986
987 registered_call_handle = _resolve_registered_call_handle(
988 self._channel,
989 method,
990 client_call_details.method,
991 self._registered_call_handle,
992 )
993
994 return StreamUnaryCall(
995 request_iterator,
996 _timeout_to_deadline(client_call_details.timeout),
997 client_call_details.metadata,
998 client_call_details.credentials,
999 client_call_details.wait_for_ready,
1000 self._channel,
1001 client_call_details.method,
1002 request_serializer,
1003 response_deserializer,
1004 self._loop,
1005 registered_call_handle,
1006 )
1007
1008 def time_remaining(self) -> Optional[float]:
1009 raise NotImplementedError()
1010
1011
1012class InterceptedStreamStreamCall(
1013 _InterceptedStreamResponseMixin,
1014 _InterceptedStreamRequestMixin,
1015 InterceptedCall,
1016 _base_call.StreamStreamCall,
1017):
1018 """Used for running a `StreamStreamCall` wrapped by interceptors."""
1019
1020 _loop: asyncio.AbstractEventLoop
1021 _channel: cygrpc.AioChannel
1022 _last_returned_call_from_interceptors = Optional[
1023 _base_call.StreamStreamCall
1024 ]
1025
1026 def __init__(
1027 self,
1028 interceptors: Sequence[StreamStreamClientInterceptor],
1029 request_iterator: Optional[RequestIterableType],
1030 timeout: Optional[float],
1031 metadata: Metadata,
1032 credentials: Optional[grpc.CallCredentials],
1033 wait_for_ready: Optional[bool],
1034 channel: cygrpc.AioChannel,
1035 method: bytes,
1036 request_serializer: Optional[SerializingFunction],
1037 response_deserializer: Optional[DeserializingFunction],
1038 loop: asyncio.AbstractEventLoop,
1039 registered_call_handle: int = 0,
1040 ) -> None:
1041 self._loop = loop
1042 self._channel = channel
1043 self._registered_call_handle = registered_call_handle
1044 self._init_stream_response_mixin()
1045 request_iterator = self._init_stream_request_mixin(request_iterator)
1046 self._last_returned_call_from_interceptors = None
1047 interceptors_task = loop.create_task(
1048 self._invoke(
1049 interceptors,
1050 method,
1051 timeout,
1052 metadata,
1053 credentials,
1054 wait_for_ready,
1055 request_iterator,
1056 request_serializer,
1057 response_deserializer,
1058 )
1059 )
1060 super().__init__(interceptors_task)
1061
1062 async def _invoke(
1063 self,
1064 interceptors: Sequence[StreamStreamClientInterceptor],
1065 method: bytes,
1066 timeout: Optional[float],
1067 metadata: Optional[Metadata],
1068 credentials: Optional[grpc.CallCredentials],
1069 wait_for_ready: Optional[bool],
1070 request_iterator: RequestIterableType,
1071 request_serializer: Optional[SerializingFunction],
1072 response_deserializer: Optional[DeserializingFunction],
1073 ) -> Union[StreamStreamCall, StreamStreamCallResponseIterator]:
1074 """Run the RPC call wrapped in interceptors"""
1075 client_call_details = ClientCallDetails(
1076 method, timeout, metadata, credentials, wait_for_ready
1077 )
1078 return await self._run_interceptor(
1079 list(interceptors),
1080 method,
1081 request_serializer,
1082 response_deserializer,
1083 client_call_details,
1084 request_iterator,
1085 )
1086
1087 async def _run_interceptor(
1088 self,
1089 interceptors: List[StreamStreamClientInterceptor],
1090 method: bytes,
1091 request_serializer: Optional[SerializingFunction],
1092 response_deserializer: Optional[DeserializingFunction],
1093 client_call_details: ClientCallDetails,
1094 request_iterator: RequestIterableType,
1095 ) -> Union[StreamStreamCall, StreamStreamCallResponseIterator]:
1096 if interceptors:
1097 continuation = functools.partial(
1098 self._run_interceptor,
1099 interceptors[1:],
1100 method,
1101 request_serializer,
1102 response_deserializer,
1103 )
1104
1105 call_or_response_iterator = await interceptors[
1106 0
1107 ].intercept_stream_stream(
1108 continuation, client_call_details, request_iterator
1109 )
1110
1111 if isinstance(
1112 call_or_response_iterator, _base_call.StreamStreamCall
1113 ):
1114 self._last_returned_call_from_interceptors = (
1115 call_or_response_iterator
1116 )
1117 else:
1118 self._last_returned_call_from_interceptors = (
1119 StreamStreamCallResponseIterator(
1120 self._last_returned_call_from_interceptors,
1121 call_or_response_iterator,
1122 )
1123 )
1124 return self._last_returned_call_from_interceptors
1125
1126 registered_call_handle = _resolve_registered_call_handle(
1127 self._channel,
1128 method,
1129 client_call_details.method,
1130 self._registered_call_handle,
1131 )
1132
1133 self._last_returned_call_from_interceptors = StreamStreamCall(
1134 request_iterator,
1135 _timeout_to_deadline(client_call_details.timeout),
1136 client_call_details.metadata,
1137 client_call_details.credentials,
1138 client_call_details.wait_for_ready,
1139 self._channel,
1140 client_call_details.method,
1141 request_serializer,
1142 response_deserializer,
1143 self._loop,
1144 registered_call_handle,
1145 )
1146 return self._last_returned_call_from_interceptors
1147
1148 def time_remaining(self) -> Optional[float]:
1149 raise NotImplementedError()
1150
1151
1152class UnaryUnaryCallResponse(_base_call.UnaryUnaryCall):
1153 """Final UnaryUnaryCall class finished with a response."""
1154
1155 _response: ResponseType
1156
1157 def __init__(self, response: ResponseType) -> None:
1158 self._response = response
1159
1160 def cancel(self) -> bool:
1161 return False
1162
1163 def cancelled(self) -> bool:
1164 return False
1165
1166 def done(self) -> bool:
1167 return True
1168
1169 def add_done_callback(self, unused_callback) -> None:
1170 raise NotImplementedError()
1171
1172 def time_remaining(self) -> Optional[float]:
1173 raise NotImplementedError()
1174
1175 async def initial_metadata(self) -> Optional[Metadata]:
1176 return None
1177
1178 async def trailing_metadata(self) -> Optional[Metadata]:
1179 return None
1180
1181 async def code(self) -> grpc.StatusCode:
1182 return grpc.StatusCode.OK
1183
1184 async def details(self) -> str:
1185 return ""
1186
1187 async def debug_error_string(self) -> Optional[str]:
1188 return None
1189
1190 def __await__(self):
1191 if False: # pylint: disable=using-constant-test
1192 # This code path is never used, but a yield statement is needed
1193 # for telling the interpreter that __await__ is a generator.
1194 yield None
1195 return self._response
1196
1197 async def wait_for_connection(self) -> None:
1198 pass
1199
1200
1201class _StreamCallResponseIterator:
1202 _call: Union[_base_call.UnaryStreamCall, _base_call.StreamStreamCall]
1203 _response_iterator: AsyncIterable[ResponseType]
1204
1205 def __init__(
1206 self,
1207 call: Union[_base_call.UnaryStreamCall, _base_call.StreamStreamCall],
1208 response_iterator: AsyncIterable[ResponseType],
1209 ) -> None:
1210 self._response_iterator = response_iterator
1211 self._call = call
1212
1213 def cancel(self) -> bool:
1214 return self._call.cancel()
1215
1216 def cancelled(self) -> bool:
1217 return self._call.cancelled()
1218
1219 def done(self) -> bool:
1220 return self._call.done()
1221
1222 def add_done_callback(self, callback) -> None:
1223 self._call.add_done_callback(callback)
1224
1225 def time_remaining(self) -> Optional[float]:
1226 return self._call.time_remaining()
1227
1228 async def initial_metadata(self) -> Optional[Metadata]:
1229 return await self._call.initial_metadata()
1230
1231 async def trailing_metadata(self) -> Optional[Metadata]:
1232 return await self._call.trailing_metadata()
1233
1234 async def code(self) -> grpc.StatusCode:
1235 return await self._call.code()
1236
1237 async def details(self) -> str:
1238 return await self._call.details()
1239
1240 async def debug_error_string(self) -> Optional[str]:
1241 return await self._call.debug_error_string()
1242
1243 def __aiter__(self):
1244 return self._response_iterator.__aiter__()
1245
1246 async def wait_for_connection(self) -> None:
1247 return await self._call.wait_for_connection()
1248
1249
1250class UnaryStreamCallResponseIterator(
1251 _StreamCallResponseIterator, _base_call.UnaryStreamCall
1252):
1253 """UnaryStreamCall class which uses an alternative response iterator."""
1254
1255 async def read(self) -> Union[EOFType, ResponseType]:
1256 # Behind the scenes everything goes through the
1257 # async iterator. So this path should not be reached.
1258 raise NotImplementedError()
1259
1260
1261class StreamStreamCallResponseIterator(
1262 _StreamCallResponseIterator, _base_call.StreamStreamCall
1263):
1264 """StreamStreamCall class which uses an alternative response iterator."""
1265
1266 async def read(self) -> Union[EOFType, ResponseType]:
1267 # Behind the scenes everything goes through the
1268 # async iterator. So this path should not be reached.
1269 raise NotImplementedError()
1270
1271 async def write(self, request: RequestType) -> None:
1272 # Behind the scenes everything goes through the
1273 # async iterator provided by the InterceptedStreamStreamCall.
1274 # So this path should not be reached.
1275 raise NotImplementedError()
1276
1277 async def done_writing(self) -> None:
1278 # Behind the scenes everything goes through the
1279 # async iterator provided by the InterceptedStreamStreamCall.
1280 # So this path should not be reached.
1281 raise NotImplementedError()
1282
1283 @property
1284 def _done_writing_flag(self) -> bool:
1285 return self._call._done_writing_flag