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

199 statements  

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 

16 

17import asyncio 

18import types 

19from typing import Any, Generic, List, Optional, Sequence, TypeVar 

20import weakref 

21 

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 

28 

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 

53 

54ClientInterceptorT = TypeVar("ClientInterceptorT", bound=ClientInterceptor) 

55 

56_USER_AGENT = "grpc-python-asyncio/{}".format(_grpcio_metadata.__version__) 

57 

58 

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 ) 

76 

77 

78class _BaseMultiCallable( 

79 Generic[RequestType, ResponseType, ClientInterceptorT] 

80): 

81 """Base class of all multi callable objects. 

82 

83 Handles the initialization logic and stores common attributes. 

84 """ 

85 

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 

94 

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 

115 

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) 

128 

129 self._python_channel = self._references[0] 

130 

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 

153 

154 

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 ) 

199 

200 self._python_channel._register_call(call) 

201 

202 return call 

203 

204 

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) 

220 

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 ) 

250 

251 self._python_channel._register_call(call) 

252 

253 return call 

254 

255 

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) 

270 

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 ) 

300 

301 self._python_channel._register_call(call) 

302 

303 return call 

304 

305 

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) 

322 

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 ) 

352 

353 self._python_channel._register_call(call) 

354 

355 return call 

356 

357 

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] 

366 

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. 

376 

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 = [] 

390 

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 ) 

409 

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() 

418 

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) 

423 

424 async def __aenter__(self) -> Self: 

425 return self 

426 

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) 

434 

435 async def _close( 

436 self, grace: Optional[float] 

437 ) -> None: # pylint: disable=too-many-branches 

438 if self._channel.closed(): 

439 return 

440 

441 if grace and grace < 0: 

442 error_msg = f"grace must be non-negative, got {grace}." 

443 raise ValueError(error_msg) 

444 

445 # No new calls will be accepted by the Cython channel. 

446 self._channel.closing() 

447 

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 

456 

457 calls = list(self._active_calls) 

458 

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) 

467 

468 # Time to cancel existing calls. 

469 for call in calls: 

470 call.cancel() 

471 

472 calls.clear() 

473 self._active_calls.clear() 

474 

475 # Destroy the channel 

476 self._channel.close() 

477 

478 async def close(self, grace: Optional[float] = None) -> None: 

479 await self._close(grace) 

480 

481 def __del__(self): 

482 if hasattr(self, "_channel") and not self._channel.closed(): 

483 self._channel.close() 

484 

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] 

490 

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) 

514 

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) 

520 

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. 

526 

527 This is a semi-private method. It is intended for use only by gRPC generated code. 

528 

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. 

532 

533 Args: 

534 method: Required, the method name for the RPC. 

535 

536 Returns: 

537 The registered call handle pointer in the form of a Python Long. 

538 """ 

539 if not _registered_method: 

540 return 0 

541 

542 return self._channel.get_registered_call_handle(_common.encode(method)) 

543 

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 ) 

564 

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 ) 

585 

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 ) 

606 

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 ) 

627 

628 

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. 

636 

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. 

645 

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 ) 

656 

657 

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. 

666 

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. 

676 

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 )