Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/grpc/aio/_channel.py: 42%
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
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"""Invocation-side implementation of gRPC Asyncio Python."""
15# pyright: reportPrivateUsage = false
17import asyncio
18import types
19from typing import Any, Generic, List, Optional, Sequence, TypeVar
20import weakref
22import grpc
23from grpc import _common
24from grpc import _compression
25from grpc import _grpcio_metadata
26from grpc._cython import cygrpc
27from typing_extensions import Self
29from . import _base_call
30from . import _base_channel
31from ._call import StreamStreamCall
32from ._call import StreamUnaryCall
33from ._call import UnaryStreamCall
34from ._call import UnaryUnaryCall
35from ._interceptor import ClientInterceptor
36from ._interceptor import InterceptedStreamStreamCall
37from ._interceptor import InterceptedStreamUnaryCall
38from ._interceptor import InterceptedUnaryStreamCall
39from ._interceptor import InterceptedUnaryUnaryCall
40from ._interceptor import StreamStreamClientInterceptor
41from ._interceptor import StreamUnaryClientInterceptor
42from ._interceptor import UnaryStreamClientInterceptor
43from ._interceptor import UnaryUnaryClientInterceptor
44from ._metadata import Metadata
45from ._typing import ChannelArgumentType
46from ._typing import DeserializingFunction
47from ._typing import MetadataType
48from ._typing import RequestIterableType
49from ._typing import RequestType
50from ._typing import ResponseType
51from ._typing import SerializingFunction
52from ._utils import _timeout_to_deadline
54ClientInterceptorT = TypeVar("ClientInterceptorT", bound=ClientInterceptor)
56_USER_AGENT = "grpc-python-asyncio/{}".format(_grpcio_metadata.__version__)
59def _augment_channel_arguments(
60 base_options: ChannelArgumentType, compression: Optional[grpc.Compression]
61) -> ChannelArgumentType:
62 compression_channel_argument = _compression.create_channel_option(
63 compression
64 )
65 user_agent_channel_argument = (
66 (
67 cygrpc.ChannelArgKey.primary_user_agent_string.decode(),
68 _USER_AGENT,
69 ),
70 )
71 return (
72 tuple(base_options)
73 + compression_channel_argument
74 + user_agent_channel_argument
75 )
78class _BaseMultiCallable(
79 Generic[RequestType, ResponseType, ClientInterceptorT]
80):
81 """Base class of all multi callable objects.
83 Handles the initialization logic and stores common attributes.
84 """
86 _loop: asyncio.AbstractEventLoop
87 _channel: cygrpc.AioChannel
88 _method: bytes
89 _request_serializer: Optional[SerializingFunction[RequestType]]
90 _response_deserializer: Optional[DeserializingFunction[ResponseType]]
91 _interceptors: Optional[Sequence[ClientInterceptorT]]
92 _references: List[Any]
93 _registered_call_handle: int
95 # pylint: disable=too-many-arguments
96 def __init__(
97 self,
98 channel: cygrpc.AioChannel,
99 method: bytes,
100 request_serializer: Optional[SerializingFunction[RequestType]],
101 response_deserializer: Optional[DeserializingFunction[ResponseType]],
102 interceptors: Optional[Sequence[ClientInterceptorT]],
103 references: List[Any],
104 loop: asyncio.AbstractEventLoop,
105 _registered_call_handle: int = 0,
106 ) -> None:
107 self._loop = loop
108 self._channel = channel
109 self._method = method
110 self._request_serializer = request_serializer
111 self._response_deserializer = response_deserializer
112 self._interceptors = interceptors
113 self._references = references
114 self._registered_call_handle = _registered_call_handle
116 if not self._references:
117 error_msg = (
118 "MultiCallable must be attached to a Channel, unexpectedly"
119 " found no references."
120 )
121 raise ValueError(error_msg)
122 if not isinstance(self._references[0], Channel):
123 error_msg = (
124 "Invalid reference type. MultiCallable must be attached to a"
125 " Channel."
126 )
127 raise TypeError(error_msg)
129 self._python_channel = self._references[0]
131 @staticmethod
132 def _init_metadata(
133 metadata: Optional[MetadataType] = None,
134 compression: Optional[grpc.Compression] = None,
135 ) -> Metadata:
136 """Based on the provided values for <metadata> or <compression> initialise the final
137 metadata, as it should be used for the current call.
138 """
139 metadata = metadata or Metadata()
140 if not isinstance(
141 metadata, Metadata
142 ) and isinstance( # pyright: ignore[reportUnnecessaryIsInstance]
143 metadata, Sequence
144 ):
145 metadata = Metadata.from_tuple(tuple(metadata))
146 if compression:
147 augmented_metadata = _compression.augment_metadata(
148 metadata, compression
149 )
150 if augmented_metadata is not None:
151 metadata = Metadata(*augmented_metadata)
152 return metadata
155class UnaryUnaryMultiCallable(
156 _BaseMultiCallable[RequestType, ResponseType, UnaryUnaryClientInterceptor],
157 _base_channel.UnaryUnaryMultiCallable[RequestType, ResponseType],
158):
159 def __call__(
160 self,
161 request: RequestType,
162 *,
163 timeout: Optional[float] = None,
164 metadata: Optional[MetadataType] = None,
165 credentials: Optional[grpc.CallCredentials] = None,
166 wait_for_ready: Optional[bool] = None,
167 compression: Optional[grpc.Compression] = None,
168 ) -> _base_call.UnaryUnaryCall[RequestType, ResponseType]:
169 metadata = self._init_metadata(metadata, compression)
170 if not self._interceptors:
171 call = UnaryUnaryCall(
172 request,
173 _timeout_to_deadline(timeout),
174 metadata,
175 credentials,
176 wait_for_ready,
177 self._channel,
178 self._method,
179 self._request_serializer,
180 self._response_deserializer,
181 self._loop,
182 self._registered_call_handle,
183 )
184 else:
185 call = InterceptedUnaryUnaryCall(
186 self._interceptors,
187 request,
188 timeout,
189 metadata,
190 credentials,
191 wait_for_ready,
192 self._channel,
193 self._method,
194 self._request_serializer,
195 self._response_deserializer,
196 self._loop,
197 self._registered_call_handle,
198 )
200 self._python_channel._register_call(call)
202 return call
205class UnaryStreamMultiCallable(
206 _BaseMultiCallable[RequestType, ResponseType, UnaryStreamClientInterceptor],
207 _base_channel.UnaryStreamMultiCallable[RequestType, ResponseType],
208):
209 def __call__(
210 self,
211 request: RequestType,
212 *,
213 timeout: Optional[float] = None,
214 metadata: Optional[MetadataType] = None,
215 credentials: Optional[grpc.CallCredentials] = None,
216 wait_for_ready: Optional[bool] = None,
217 compression: Optional[grpc.Compression] = None,
218 ) -> _base_call.UnaryStreamCall[RequestType, ResponseType]:
219 metadata = self._init_metadata(metadata, compression)
221 if not self._interceptors:
222 call = UnaryStreamCall(
223 request,
224 _timeout_to_deadline(timeout),
225 metadata,
226 credentials,
227 wait_for_ready,
228 self._channel,
229 self._method,
230 self._request_serializer,
231 self._response_deserializer,
232 self._loop,
233 self._registered_call_handle,
234 )
235 else:
236 call = InterceptedUnaryStreamCall(
237 self._interceptors,
238 request,
239 timeout,
240 metadata,
241 credentials,
242 wait_for_ready,
243 self._channel,
244 self._method,
245 self._request_serializer,
246 self._response_deserializer,
247 self._loop,
248 self._registered_call_handle,
249 )
251 self._python_channel._register_call(call)
253 return call
256class StreamUnaryMultiCallable(
257 _BaseMultiCallable[RequestType, ResponseType, StreamUnaryClientInterceptor],
258 _base_channel.StreamUnaryMultiCallable[RequestType, ResponseType],
259):
260 def __call__(
261 self,
262 request_iterator: Optional[RequestIterableType[RequestType]] = None,
263 timeout: Optional[float] = None,
264 metadata: Optional[MetadataType] = None,
265 credentials: Optional[grpc.CallCredentials] = None,
266 wait_for_ready: Optional[bool] = None,
267 compression: Optional[grpc.Compression] = None,
268 ) -> _base_call.StreamUnaryCall[RequestType, ResponseType]:
269 metadata = self._init_metadata(metadata, compression)
271 if not self._interceptors:
272 call = StreamUnaryCall(
273 request_iterator,
274 _timeout_to_deadline(timeout),
275 metadata,
276 credentials,
277 wait_for_ready,
278 self._channel,
279 self._method,
280 self._request_serializer,
281 self._response_deserializer,
282 self._loop,
283 self._registered_call_handle,
284 )
285 else:
286 call = InterceptedStreamUnaryCall(
287 self._interceptors,
288 request_iterator,
289 timeout,
290 metadata,
291 credentials,
292 wait_for_ready,
293 self._channel,
294 self._method,
295 self._request_serializer,
296 self._response_deserializer,
297 self._loop,
298 self._registered_call_handle,
299 )
301 self._python_channel._register_call(call)
303 return call
306class StreamStreamMultiCallable(
307 _BaseMultiCallable[
308 RequestType, ResponseType, StreamStreamClientInterceptor
309 ],
310 _base_channel.StreamStreamMultiCallable[RequestType, ResponseType],
311):
312 def __call__(
313 self,
314 request_iterator: Optional[RequestIterableType[RequestType]] = None,
315 timeout: Optional[float] = None,
316 metadata: Optional[MetadataType] = None,
317 credentials: Optional[grpc.CallCredentials] = None,
318 wait_for_ready: Optional[bool] = None,
319 compression: Optional[grpc.Compression] = None,
320 ) -> _base_call.StreamStreamCall[RequestType, ResponseType]:
321 metadata = self._init_metadata(metadata, compression)
323 if not self._interceptors:
324 call = StreamStreamCall(
325 request_iterator,
326 _timeout_to_deadline(timeout),
327 metadata,
328 credentials,
329 wait_for_ready,
330 self._channel,
331 self._method,
332 self._request_serializer,
333 self._response_deserializer,
334 self._loop,
335 self._registered_call_handle,
336 )
337 else:
338 call = InterceptedStreamStreamCall(
339 self._interceptors,
340 request_iterator,
341 timeout,
342 metadata,
343 credentials,
344 wait_for_ready,
345 self._channel,
346 self._method,
347 self._request_serializer,
348 self._response_deserializer,
349 self._loop,
350 self._registered_call_handle,
351 )
353 self._python_channel._register_call(call)
355 return call
358class Channel(_base_channel.Channel):
359 _loop: asyncio.AbstractEventLoop
360 _channel: cygrpc.AioChannel
361 _unary_unary_interceptors: List[UnaryUnaryClientInterceptor]
362 _unary_stream_interceptors: List[UnaryStreamClientInterceptor]
363 _stream_unary_interceptors: List[StreamUnaryClientInterceptor]
364 _stream_stream_interceptors: List[StreamStreamClientInterceptor]
365 _active_calls: weakref.WeakSet[_base_call.Call]
367 def __init__(
368 self,
369 target: str,
370 options: ChannelArgumentType,
371 credentials: Optional[cygrpc.ChannelCredentials],
372 compression: Optional[grpc.Compression],
373 interceptors: Optional[Sequence[ClientInterceptor]],
374 ):
375 """Constructor.
377 Args:
378 target: The target to which to connect.
379 options: Configuration options for the channel.
380 credentials: A cygrpc.ChannelCredentials or None.
381 compression: An optional value indicating the compression method to be
382 used over the lifetime of the channel.
383 interceptors: An optional list of interceptors that would be used for
384 intercepting any RPC executed with that channel.
385 """
386 self._unary_unary_interceptors = []
387 self._unary_stream_interceptors = []
388 self._stream_unary_interceptors = []
389 self._stream_stream_interceptors = []
391 if interceptors is not None:
392 for interceptor in interceptors:
393 if isinstance(interceptor, UnaryUnaryClientInterceptor):
394 self._unary_unary_interceptors.append(interceptor)
395 elif isinstance(interceptor, UnaryStreamClientInterceptor):
396 self._unary_stream_interceptors.append(interceptor)
397 elif isinstance(interceptor, StreamUnaryClientInterceptor):
398 self._stream_unary_interceptors.append(interceptor)
399 elif isinstance(interceptor, StreamStreamClientInterceptor):
400 self._stream_stream_interceptors.append(interceptor)
401 else:
402 raise ValueError( # noqa: TRY004
403 "Interceptor {} must be ".format(interceptor)
404 + "{} or ".format(UnaryUnaryClientInterceptor.__name__)
405 + "{} or ".format(UnaryStreamClientInterceptor.__name__)
406 + "{} or ".format(StreamUnaryClientInterceptor.__name__)
407 + "{}. ".format(StreamStreamClientInterceptor.__name__)
408 )
410 self._loop = cygrpc.get_working_loop()
411 self._channel = cygrpc.AioChannel(
412 _common.encode(target),
413 _augment_channel_arguments(options, compression),
414 credentials,
415 self._loop,
416 )
417 self._active_calls = weakref.WeakSet()
419 def _register_call(self, call: _base_call.Call) -> None:
420 """Register a call to be tracked by the channel."""
421 self._active_calls.add(call)
422 call.add_done_callback(self._active_calls.discard)
424 async def __aenter__(self) -> Self:
425 return self
427 async def __aexit__(
428 self,
429 exc_type: Optional[type[BaseException]],
430 exc_val: Optional[BaseException],
431 exc_tb: Optional[types.TracebackType],
432 ) -> Optional[bool]:
433 await self._close(None)
435 async def _close(
436 self, grace: Optional[float]
437 ) -> None: # pylint: disable=too-many-branches
438 if self._channel.closed():
439 return
441 if grace and grace < 0:
442 error_msg = f"grace must be non-negative, got {grace}."
443 raise ValueError(error_msg)
445 # No new calls will be accepted by the Cython channel.
446 self._channel.closing()
448 async def _wait_for_call_to_complete(call: _base_call.Call) -> None:
449 try:
450 await call.code()
451 except Exception: # pylint: disable=broad-except
452 # Ignore exceptions here as true RPC errors bubble up via
453 # standard application paths. Silencing prevents channel close
454 # from failing and suppresses asyncio noise warnings.
455 pass
457 calls = list(self._active_calls)
459 if grace:
460 call_tasks = [
461 self._loop.create_task(_wait_for_call_to_complete(call))
462 for call in calls
463 if not call.done()
464 ]
465 if call_tasks:
466 await asyncio.wait(call_tasks, timeout=grace)
468 # Time to cancel existing calls.
469 for call in calls:
470 call.cancel()
472 calls.clear()
473 self._active_calls.clear()
475 # Destroy the channel
476 self._channel.close()
478 async def close(self, grace: Optional[float] = None) -> None:
479 await self._close(grace)
481 def __del__(self):
482 if hasattr(self, "_channel") and not self._channel.closed():
483 self._channel.close()
485 def get_state(
486 self, try_to_connect: bool = False
487 ) -> grpc.ChannelConnectivity:
488 result = self._channel.check_connectivity_state(try_to_connect)
489 return _common.CYGRPC_CONNECTIVITY_STATE_TO_CHANNEL_CONNECTIVITY[result]
491 async def wait_for_state_change(
492 self,
493 last_observed_state: grpc.ChannelConnectivity,
494 ) -> None:
495 # We raise a RuntimeError if watch_connectivity_state returns False.
496 #
497 # The watch_connectivity_state method returns True when it observes a state change
498 # and False when it times out (which shouldn't happen since no timeout is specified).
499 # A channel close triggers a transition to SHUTDOWN, which resolves all pending watch
500 # calls and makes them return True. Thus, watch_connectivity_state should only return
501 # True under normal operation; returning False indicates an implementation issue.
502 #
503 # We do not use an assert statement here because asserts
504 # can be optimized out under python -O.
505 # See https://github.com/grpc/grpc/issues/42393 for context.
506 resolved = await self._channel.watch_connectivity_state(
507 last_observed_state.value[0], None
508 )
509 if not resolved:
510 error_msg = (
511 "gRPC channel connectivity state watch failed unexpectedly."
512 )
513 raise RuntimeError(error_msg)
515 async def channel_ready(self) -> None:
516 state = self.get_state(try_to_connect=True)
517 while state != grpc.ChannelConnectivity.READY:
518 await self.wait_for_state_change(state)
519 state = self.get_state(try_to_connect=True)
521 def _get_registered_call_handle(
522 self, method: str, _registered_method: Optional[bool]
523 ) -> int:
524 """
525 Get the registered call handle for a registered method or None.
527 This is a semi-private method. It is intended for use only by gRPC generated code.
529 This method is not thread-safe. It is acceptable since method is only called
530 during multicallable construction, not during RPC exeution. Moreover there
531 are no `await` suspension points, which can interleave.
533 Args:
534 method: Required, the method name for the RPC.
536 Returns:
537 The registered call handle pointer in the form of a Python Long.
538 """
539 if not _registered_method:
540 return 0
542 return self._channel.get_registered_call_handle(_common.encode(method))
544 # pylint: disable=arguments-differ
545 def unary_unary(
546 self,
547 method: str,
548 request_serializer: Optional[SerializingFunction[RequestType]] = None,
549 response_deserializer: Optional[
550 DeserializingFunction[ResponseType]
551 ] = None,
552 _registered_method: Optional[bool] = False,
553 ) -> UnaryUnaryMultiCallable[RequestType, ResponseType]:
554 return UnaryUnaryMultiCallable(
555 self._channel,
556 _common.encode(method),
557 request_serializer,
558 response_deserializer,
559 self._unary_unary_interceptors,
560 [self],
561 self._loop,
562 self._get_registered_call_handle(method, _registered_method),
563 )
565 # pylint: disable=arguments-differ
566 def unary_stream(
567 self,
568 method: str,
569 request_serializer: Optional[SerializingFunction[RequestType]] = None,
570 response_deserializer: Optional[
571 DeserializingFunction[ResponseType]
572 ] = None,
573 _registered_method: Optional[bool] = False,
574 ) -> UnaryStreamMultiCallable[RequestType, ResponseType]:
575 return UnaryStreamMultiCallable(
576 self._channel,
577 _common.encode(method),
578 request_serializer,
579 response_deserializer,
580 self._unary_stream_interceptors,
581 [self],
582 self._loop,
583 self._get_registered_call_handle(method, _registered_method),
584 )
586 # pylint: disable=arguments-differ
587 def stream_unary(
588 self,
589 method: str,
590 request_serializer: Optional[SerializingFunction[RequestType]] = None,
591 response_deserializer: Optional[
592 DeserializingFunction[ResponseType]
593 ] = None,
594 _registered_method: Optional[bool] = False,
595 ) -> StreamUnaryMultiCallable[RequestType, ResponseType]:
596 return StreamUnaryMultiCallable(
597 self._channel,
598 _common.encode(method),
599 request_serializer,
600 response_deserializer,
601 self._stream_unary_interceptors,
602 [self],
603 self._loop,
604 self._get_registered_call_handle(method, _registered_method),
605 )
607 # pylint: disable=arguments-differ
608 def stream_stream(
609 self,
610 method: str,
611 request_serializer: Optional[SerializingFunction[RequestType]] = None,
612 response_deserializer: Optional[
613 DeserializingFunction[ResponseType]
614 ] = None,
615 _registered_method: Optional[bool] = False,
616 ) -> StreamStreamMultiCallable[RequestType, ResponseType]:
617 return StreamStreamMultiCallable(
618 self._channel,
619 _common.encode(method),
620 request_serializer,
621 response_deserializer,
622 self._stream_stream_interceptors,
623 [self],
624 self._loop,
625 self._get_registered_call_handle(method, _registered_method),
626 )
629def insecure_channel(
630 target: str,
631 options: Optional[ChannelArgumentType] = None,
632 compression: Optional[grpc.Compression] = None,
633 interceptors: Optional[Sequence[ClientInterceptor]] = None,
634) -> _base_channel.Channel:
635 """Creates an insecure asynchronous Channel to a server.
637 Args:
638 target: The server address
639 options: An optional list of key-value pairs (:term:`channel_arguments`
640 in gRPC Core runtime) to configure the channel.
641 compression: An optional value indicating the compression method to be
642 used over the lifetime of the channel.
643 interceptors: An optional sequence of interceptors that will be executed for
644 any call executed with this channel.
646 Returns:
647 A Channel.
648 """
649 return Channel(
650 target,
651 () if options is None else options,
652 None,
653 compression,
654 interceptors,
655 )
658def secure_channel(
659 target: str,
660 credentials: grpc.ChannelCredentials,
661 options: Optional[ChannelArgumentType] = None,
662 compression: Optional[grpc.Compression] = None,
663 interceptors: Optional[Sequence[ClientInterceptor]] = None,
664) -> _base_channel.Channel:
665 """Creates a secure asynchronous Channel to a server.
667 Args:
668 target: The server address.
669 credentials: A ChannelCredentials instance.
670 options: An optional list of key-value pairs (:term:`channel_arguments`
671 in gRPC Core runtime) to configure the channel.
672 compression: An optional value indicating the compression method to be
673 used over the lifetime of the channel.
674 interceptors: An optional sequence of interceptors that will be executed for
675 any call executed with this channel.
677 Returns:
678 An aio.Channel.
679 """
680 return Channel(
681 target,
682 () if options is None else options,
683 credentials._credentials,
684 compression,
685 interceptors,
686 )