codekingpro/portable-devtools
114k
1# Copyright 2016 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"""Service-side implementation of gRPC Python."""
15
16from __future__ import annotations
17
18import abc
19import collections
20from concurrent import futures
21import contextvars
22import enum
23import logging
24import threading
25import time
26import traceback
27from typing import (
28 Any,
29 Callable,
30 Dict,
31 Iterable,
32 Iterator,
33 List,
34 Mapping,
35 Optional,
36 Sequence,
37 Set,
38 Tuple,
39 Union,
40)
41
42import grpc
43from grpc import _common
44from grpc import _compression
45from grpc import _interceptor
46from grpc import _observability
47from grpc._cython import cygrpc
48from grpc._typing import ArityAgnosticMethodHandler
49from grpc._typing import ChannelArgumentType
50from grpc._typing import DeserializingFunction
51from grpc._typing import MetadataType
52from grpc._typing import NullaryCallbackType
53from grpc._typing import ResponseType
54from grpc._typing import SerializingFunction
55from grpc._typing import ServerCallbackTag
56from grpc._typing import ServerTagCallbackType
57from typing_extensions import override
58
59_LOGGER = logging.getLogger(__name__)
60
61_SHUTDOWN_TAG = "shutdown"
62_REQUEST_CALL_TAG = "request_call"
63
64_RECEIVE_CLOSE_ON_SERVER_TOKEN = "receive_close_on_server"
65_SEND_INITIAL_METADATA_TOKEN = "send_initial_metadata"
66_RECEIVE_MESSAGE_TOKEN = "receive_message"
67_SEND_MESSAGE_TOKEN = "send_message"
68_SEND_INITIAL_METADATA_AND_SEND_MESSAGE_TOKEN = (
69 "send_initial_metadata * send_message"
70)
71_SEND_STATUS_FROM_SERVER_TOKEN = "send_status_from_server"
72_SEND_INITIAL_METADATA_AND_SEND_STATUS_FROM_SERVER_TOKEN = (
73 "send_initial_metadata * send_status_from_server"
74)
75
76_OPEN = "open"
77_CLOSED = "closed"
78_CANCELLED = "cancelled"
79
80_EMPTY_FLAGS = 0
81
82_DEALLOCATED_SERVER_CHECK_PERIOD_S = 1.0
83_INF_TIMEOUT = 1e9
84
85
86def _serialized_request(request_event: cygrpc.BaseEvent) -> bytes:
87 return request_event.batch_operations[0].message()
88
89
90def _application_code(code: grpc.StatusCode) -> cygrpc.StatusCode:
91 cygrpc_code = _common.STATUS_CODE_TO_CYGRPC_STATUS_CODE.get(code)
92 return cygrpc.StatusCode.unknown if cygrpc_code is None else cygrpc_code
93
94
95def _completion_code(state: _RPCState) -> cygrpc.StatusCode:
96 if state.code is None:
97 return cygrpc.StatusCode.ok
98 return _application_code(state.code)
99
100
101def _abortion_code(
102 state: _RPCState, code: cygrpc.StatusCode
103) -> cygrpc.StatusCode:
104 if state.code is None:
105 return code
106 return _application_code(state.code)
107
108
109def _details(state: _RPCState) -> bytes:
110 return b"" if state.details is None else state.details
111
112
113class _HandlerCallDetails(
114 collections.namedtuple(
115 "_HandlerCallDetails",
116 (
117 "method",
118 "invocation_metadata",
119 ),
120 ),
121 grpc.HandlerCallDetails,
122):
123 pass
124
125
126class _Method(abc.ABC):
127 @abc.abstractmethod
128 def name(self) -> Optional[str]:
129 raise NotImplementedError()
130
131 @abc.abstractmethod
132 def handler(
133 self, handler_call_details: _HandlerCallDetails
134 ) -> Optional[grpc.RpcMethodHandler]:
135 raise NotImplementedError()
136
137
138class _RegisteredMethod(_Method):
139 def __init__(
140 self,
141 name: str,
142 registered_handler: Optional[grpc.RpcMethodHandler],
143 ):
144 self._name = name
145 self._registered_handler = registered_handler
146
147 @override
148 def name(self) -> Optional[str]:
149 return self._name
150
151 @override
152 def handler(
153 self, handler_call_details: _HandlerCallDetails
154 ) -> Optional[grpc.RpcMethodHandler]:
155 return self._registered_handler
156
157
158class _GenericMethod(_Method):
159 def __init__(
160 self,
161 generic_handlers: List[grpc.GenericRpcHandler],
162 ):
163 self._generic_handlers = generic_handlers
164
165 @override
166 def name(self) -> Optional[str]:
167 return None
168
169 @override
170 def handler(
171 self, handler_call_details: _HandlerCallDetails
172 ) -> Optional[grpc.RpcMethodHandler]:
173 # If the same method have both generic and registered handler,
174 # registered handler will take precedence.
175 for generic_handler in self._generic_handlers:
176 method_handler = generic_handler.service(handler_call_details)
177 if method_handler is not None:
178 return method_handler
179 return None
180
181
182class _RPCState(object):
183 context: contextvars.Context
184 condition: threading.Condition
185 due = Set[str]
186 request: Any
187 client: str
188 initial_metadata_allowed: bool
189 compression_algorithm: Optional[grpc.Compression]
190 disable_next_compression: bool
191 trailing_metadata: Optional[MetadataType]
192 code: Optional[grpc.StatusCode]
193 details: Optional[bytes]
194 statused: bool
195 rpc_errors: List[Exception]
196 callbacks: Optional[List[NullaryCallbackType]]
197 aborted: bool
198
199 def __init__(self):
200 self.context = contextvars.Context()
201 self.condition = threading.Condition()
202 self.due = set()
203 self.request = None
204 self.client = _OPEN
205 self.initial_metadata_allowed = True
206 self.compression_algorithm = None
207 self.disable_next_compression = False
208 self.trailing_metadata = None
209 self.code = None
210 self.details = None
211 self.statused = False
212 self.rpc_errors = []
213 self.callbacks = []
214 self.aborted = False
215
216
217def _raise_rpc_error(state: _RPCState) -> None:
218 rpc_error = grpc.RpcError()
219 state.rpc_errors.append(rpc_error)
220 raise rpc_error
221
222
223def _possibly_finish_call(
224 state: _RPCState, token: str
225) -> ServerTagCallbackType:
226 state.due.remove(token)
227 if not _is_rpc_state_active(state) and not state.due:
228 callbacks = state.callbacks
229 state.callbacks = None
230 return state, callbacks
231 return None, ()
232
233
234def _send_status_from_server(state: _RPCState, token: str) -> ServerCallbackTag:
235 def send_status_from_server(unused_send_status_from_server_event):
236 with state.condition:
237 return _possibly_finish_call(state, token)
238
239 return send_status_from_server
240
241
242def _get_initial_metadata(
243 state: _RPCState, metadata: Optional[MetadataType]
244) -> Optional[MetadataType]:
245 with state.condition:
246 if state.compression_algorithm:
247 compression_metadata = (
248 _compression.compression_algorithm_to_metadata(
249 state.compression_algorithm
250 ),
251 )
252 if metadata is None:
253 return compression_metadata
254 return compression_metadata + tuple(metadata)
255 return metadata
256
257
258def _get_initial_metadata_operation(
259 state: _RPCState, metadata: Optional[MetadataType]
260) -> cygrpc.Operation:
261 operation = cygrpc.SendInitialMetadataOperation(
262 _get_initial_metadata(state, metadata), _EMPTY_FLAGS
263 )
264 return operation
265
266
267def _abort(
268 state: _RPCState, call: cygrpc.Call, code: cygrpc.StatusCode, details: bytes
269) -> None:
270 if state.client is not _CANCELLED:
271 effective_code = _abortion_code(state, code)
272 effective_details = details if state.details is None else state.details
273 if state.initial_metadata_allowed:
274 operations = (
275 _get_initial_metadata_operation(state, None),
276 cygrpc.SendStatusFromServerOperation(
277 state.trailing_metadata,
278 effective_code,
279 effective_details,
280 _EMPTY_FLAGS,
281 ),
282 )
283 token = _SEND_INITIAL_METADATA_AND_SEND_STATUS_FROM_SERVER_TOKEN
284 else:
285 operations = (
286 cygrpc.SendStatusFromServerOperation(
287 state.trailing_metadata,
288 effective_code,
289 effective_details,
290 _EMPTY_FLAGS,
291 ),
292 )
293 token = _SEND_STATUS_FROM_SERVER_TOKEN
294 call.start_server_batch(
295 operations, _send_status_from_server(state, token)
296 )
297 state.statused = True
298 state.due.add(token)
299
300
301def _receive_close_on_server(state: _RPCState) -> ServerCallbackTag:
302 def receive_close_on_server(receive_close_on_server_event):
303 with state.condition:
304 if receive_close_on_server_event.batch_operations[0].cancelled():
305 state.client = _CANCELLED
306 elif state.client is _OPEN:
307 state.client = _CLOSED
308 state.condition.notify_all()
309 return _possibly_finish_call(state, _RECEIVE_CLOSE_ON_SERVER_TOKEN)
310
311 return receive_close_on_server
312
313
314def _receive_message(
315 state: _RPCState,
316 call: cygrpc.Call,
317 request_deserializer: Optional[DeserializingFunction],
318) -> ServerCallbackTag:
319 def receive_message(receive_message_event):
320 serialized_request = _serialized_request(receive_message_event)
321 if serialized_request is None:
322 with state.condition:
323 if state.client is _OPEN:
324 state.client = _CLOSED
325 state.condition.notify_all()
326 return _possibly_finish_call(state, _RECEIVE_MESSAGE_TOKEN)
327 else:
328 request = _common.deserialize(
329 serialized_request, request_deserializer
330 )
331 with state.condition:
332 if request is None:
333 _abort(
334 state,
335 call,
336 cygrpc.StatusCode.internal,
337 b"Exception deserializing request!",
338 )
339 else:
340 state.request = request
341 state.condition.notify_all()
342 return _possibly_finish_call(state, _RECEIVE_MESSAGE_TOKEN)
343
344 return receive_message
345
346
347def _send_initial_metadata(state: _RPCState) -> ServerCallbackTag:
348 def send_initial_metadata(unused_send_initial_metadata_event):
349 with state.condition:
350 return _possibly_finish_call(state, _SEND_INITIAL_METADATA_TOKEN)
351
352 return send_initial_metadata
353
354
355def _send_message(state: _RPCState, token: str) -> ServerCallbackTag:
356 def send_message(unused_send_message_event):
357 with state.condition:
358 state.condition.notify_all()
359 return _possibly_finish_call(state, token)
360
361 return send_message
362
363
364class _Context(grpc.ServicerContext):
365 _rpc_event: cygrpc.BaseEvent
366 _state: _RPCState
367 request_deserializer: Optional[DeserializingFunction]
368
369 def __init__(
370 self,
371 rpc_event: cygrpc.BaseEvent,
372 state: _RPCState,
373 request_deserializer: Optional[DeserializingFunction],
374 ):
375 self._rpc_event = rpc_event
376 self._state = state
377 self._request_deserializer = request_deserializer
378
379 def is_active(self) -> bool:
380 with self._state.condition:
381 return _is_rpc_state_active(self._state)
382
383 def time_remaining(self) -> float:
384 return max(self._rpc_event.call_details.deadline - time.time(), 0)
385
386 def cancel(self) -> None:
387 self._rpc_event.call.cancel()
388
389 def add_callback(self, callback: NullaryCallbackType) -> bool:
390 with self._state.condition:
391 if self._state.callbacks is None:
392 return False
393 self._state.callbacks.append(callback)
394 return True
395
396 def disable_next_message_compression(self) -> None:
397 with self._state.condition:
398 self._state.disable_next_compression = True
399
400 def invocation_metadata(self) -> Optional[MetadataType]:
401 return self._rpc_event.invocation_metadata
402
403 def peer(self) -> str:
404 return _common.decode(self._rpc_event.call.peer())
405
406 def peer_identities(self) -> Optional[Sequence[bytes]]:
407 return cygrpc.peer_identities(self._rpc_event.call)
408
409 def peer_identity_key(self) -> Optional[str]:
410 id_key = cygrpc.peer_identity_key(self._rpc_event.call)
411 return id_key if id_key is None else _common.decode(id_key)
412
413 def auth_context(self) -> Mapping[str, Sequence[bytes]]:
414 auth_context = cygrpc.auth_context(self._rpc_event.call)
415 auth_context_dict = {} if auth_context is None else auth_context
416 return {
417 _common.decode(key): value
418 for key, value in auth_context_dict.items()
419 }
420
421 def set_compression(self, compression: grpc.Compression) -> None:
422 with self._state.condition:
423 self._state.compression_algorithm = compression
424
425 def send_initial_metadata(self, initial_metadata: MetadataType) -> None:
426 with self._state.condition:
427 if self._state.client is _CANCELLED:
428 _raise_rpc_error(self._state)
429 if self._state.initial_metadata_allowed:
430 operation = _get_initial_metadata_operation(
431 self._state, initial_metadata
432 )
433 self._rpc_event.call.start_server_batch(
434 (operation,), _send_initial_metadata(self._state)
435 )
436 self._state.initial_metadata_allowed = False
437 self._state.due.add(_SEND_INITIAL_METADATA_TOKEN)
438 else:
439 error_msg = "Initial metadata no longer allowed!"
440 raise ValueError(error_msg)
441
442 def set_trailing_metadata(self, trailing_metadata: MetadataType) -> None:
443 with self._state.condition:
444 self._state.trailing_metadata = trailing_metadata
445
446 def trailing_metadata(self) -> Optional[MetadataType]:
447 return self._state.trailing_metadata
448
449 def abort(self, code: grpc.StatusCode, details: str) -> None:
450 # treat OK like other invalid arguments: fail the RPC
451 if code == grpc.StatusCode.OK:
452 _LOGGER.error(
453 "abort() called with StatusCode.OK; returning UNKNOWN"
454 )
455 code = grpc.StatusCode.UNKNOWN
456 details = ""
457 with self._state.condition:
458 self._state.code = code
459 self._state.details = _common.encode(details)
460 self._state.aborted = True
461 raise Exception() # noqa: TRY002
462
463 def abort_with_status(self, status: grpc.Status) -> None:
464 self._state.trailing_metadata = status.trailing_metadata
465 self.abort(status.code, status.details)
466
467 def set_code(self, code: grpc.StatusCode) -> None:
468 with self._state.condition:
469 self._state.code = code
470
471 def code(self) -> grpc.StatusCode:
472 return self._state.code
473
474 def set_details(self, details: str) -> None:
475 with self._state.condition:
476 self._state.details = _common.encode(details)
477
478 def details(self) -> bytes:
479 return self._state.details
480
481 def _finalize_state(self) -> None:
482 pass
483
484
485class _RequestIterator(object):
486 _state: _RPCState
487 _call: cygrpc.Call
488 _request_deserializer: Optional[DeserializingFunction]
489
490 def __init__(
491 self,
492 state: _RPCState,
493 call: cygrpc.Call,
494 request_deserializer: Optional[DeserializingFunction],
495 ):
496 self._state = state
497 self._call = call
498 self._request_deserializer = request_deserializer
499
500 def _raise_or_start_receive_message(self) -> None:
501 if self._state.client is _CANCELLED:
502 _raise_rpc_error(self._state)
503 elif not _is_rpc_state_active(self._state):
504 raise StopIteration()
505 else:
506 self._call.start_server_batch(
507 (cygrpc.ReceiveMessageOperation(_EMPTY_FLAGS),),
508 _receive_message(
509 self._state, self._call, self._request_deserializer
510 ),
511 )
512 self._state.due.add(_RECEIVE_MESSAGE_TOKEN)
513
514 def _look_for_request(self) -> Any:
515 if self._state.client is _CANCELLED:
516 _raise_rpc_error(self._state)
517 elif (
518 self._state.request is None
519 and _RECEIVE_MESSAGE_TOKEN not in self._state.due
520 ):
521 raise StopIteration()
522 else:
523 request = self._state.request
524 self._state.request = None
525 return request
526
527 raise AssertionError() # should never run
528
529 def _next(self) -> Any:
530 with self._state.condition:
531 self._raise_or_start_receive_message()
532 while True:
533 self._state.condition.wait()
534 request = self._look_for_request()
535 if request is not None:
536 return request
537
538 def __iter__(self) -> _RequestIterator:
539 return self
540
541 def __next__(self) -> Any:
542 return self._next()
543
544 def next(self) -> Any:
545 return self._next()
546
547
548def _unary_request(
549 rpc_event: cygrpc.BaseEvent,
550 state: _RPCState,
551 request_deserializer: Optional[DeserializingFunction],
552) -> Callable[[], Any]:
553 def unary_request():
554 with state.condition:
555 if not _is_rpc_state_active(state):
556 return None
557 rpc_event.call.start_server_batch(
558 (cygrpc.ReceiveMessageOperation(_EMPTY_FLAGS),),
559 _receive_message(state, rpc_event.call, request_deserializer),
560 )
561 state.due.add(_RECEIVE_MESSAGE_TOKEN)
562 while True:
563 state.condition.wait()
564 if state.request is None:
565 if state.client is _CLOSED:
566 details = (
567 '"{}" requires exactly one request message.'.format(
568 rpc_event.call_details.method
569 )
570 )
571 _abort(
572 state,
573 rpc_event.call,
574 cygrpc.StatusCode.unimplemented,
575 _common.encode(details),
576 )
577 return None
578 if state.client is _CANCELLED:
579 return None
580 else:
581 request = state.request
582 state.request = None
583 return request
584
585 return unary_request
586
587
588def _call_behavior(
589 rpc_event: cygrpc.BaseEvent,
590 state: _RPCState,
591 behavior: ArityAgnosticMethodHandler,
592 argument: Any,
593 request_deserializer: Optional[DeserializingFunction],
594 send_response_callback: Optional[Callable[[ResponseType], None]] = None,
595) -> Tuple[Union[ResponseType, Iterator[ResponseType]], bool]:
596 from grpc import _create_servicer_context
597
598 with _create_servicer_context(
599 rpc_event, state, request_deserializer
600 ) as context:
601 try:
602 response_or_iterator = None
603 if send_response_callback is not None:
604 response_or_iterator = behavior(
605 argument, context, send_response_callback
606 )
607 else:
608 response_or_iterator = behavior(argument, context)
609 return response_or_iterator, True
610 except Exception as exception: # pylint: disable=broad-except
611 with state.condition:
612 if state.aborted:
613 _abort(
614 state,
615 rpc_event.call,
616 cygrpc.StatusCode.unknown,
617 b"RPC Aborted",
618 )
619 elif exception not in state.rpc_errors:
620 try:
621 details = "Exception calling application: {}".format(
622 exception
623 )
624 except Exception: # pylint: disable=broad-except
625 details = (
626 "Calling application raised unprintable Exception!"
627 )
628 _LOGGER.exception(
629 traceback.format_exception(
630 type(exception),
631 exception,
632 exception.__traceback__,
633 )
634 )
635 traceback.print_exc()
636 _LOGGER.exception(details)
637 _abort(
638 state,
639 rpc_event.call,
640 cygrpc.StatusCode.unknown,
641 _common.encode(details),
642 )
643 return None, False
644
645
646def _take_response_from_response_iterator(
647 rpc_event: cygrpc.BaseEvent,
648 state: _RPCState,
649 response_iterator: Iterator[ResponseType],
650) -> Tuple[ResponseType, bool]:
651 try:
652 return next(response_iterator), True
653 except StopIteration:
654 return None, True
655 except Exception as exception: # pylint: disable=broad-except
656 with state.condition:
657 if state.aborted:
658 _abort(
659 state,
660 rpc_event.call,
661 cygrpc.StatusCode.unknown,
662 b"RPC Aborted",
663 )
664 elif exception not in state.rpc_errors:
665 details = "Exception iterating responses: {}".format(exception)
666 _LOGGER.exception(details)
667 _abort(
668 state,
669 rpc_event.call,
670 cygrpc.StatusCode.unknown,
671 _common.encode(details),
672 )
673 return None, False
674
675
676def _serialize_response(
677 rpc_event: cygrpc.BaseEvent,
678 state: _RPCState,
679 response: Any,
680 response_serializer: Optional[SerializingFunction],
681) -> Optional[bytes]:
682 serialized_response = _common.serialize(response, response_serializer)
683 if serialized_response is None:
684 with state.condition:
685 _abort(
686 state,
687 rpc_event.call,
688 cygrpc.StatusCode.internal,
689 b"Failed to serialize response!",
690 )
691 return None
692 return serialized_response
693
694
695def _get_send_message_op_flags_from_state(
696 state: _RPCState,
697) -> Union[int, cygrpc.WriteFlag]:
698 if state.disable_next_compression:
699 return cygrpc.WriteFlag.no_compress
700 return _EMPTY_FLAGS
701
702
703def _reset_per_message_state(state: _RPCState) -> None:
704 with state.condition:
705 state.disable_next_compression = False
706
707
708def _send_response(
709 rpc_event: cygrpc.BaseEvent, state: _RPCState, serialized_response: bytes
710) -> bool:
711 with state.condition:
712 if not _is_rpc_state_active(state):
713 return False
714 if state.initial_metadata_allowed:
715 operations = (
716 _get_initial_metadata_operation(state, None),
717 cygrpc.SendMessageOperation(
718 serialized_response,
719 _get_send_message_op_flags_from_state(state),
720 ),
721 )
722 state.initial_metadata_allowed = False
723 token = _SEND_INITIAL_METADATA_AND_SEND_MESSAGE_TOKEN
724 else:
725 operations = (
726 cygrpc.SendMessageOperation(
727 serialized_response,
728 _get_send_message_op_flags_from_state(state),
729 ),
730 )
731 token = _SEND_MESSAGE_TOKEN
732 rpc_event.call.start_server_batch(
733 operations, _send_message(state, token)
734 )
735 state.due.add(token)
736 _reset_per_message_state(state)
737 while True:
738 state.condition.wait()
739 if token not in state.due:
740 return _is_rpc_state_active(state)
741
742
743def _status(
744 rpc_event: cygrpc.BaseEvent,
745 state: _RPCState,
746 serialized_response: Optional[bytes],
747) -> None:
748 with state.condition:
749 if state.client is not _CANCELLED:
750 code = _completion_code(state)
751 details = _details(state)
752 operations = [
753 cygrpc.SendStatusFromServerOperation(
754 state.trailing_metadata, code, details, _EMPTY_FLAGS
755 ),
756 ]
757 if state.initial_metadata_allowed:
758 operations.append(_get_initial_metadata_operation(state, None))
759 if serialized_response is not None:
760 operations.append(
761 cygrpc.SendMessageOperation(
762 serialized_response,
763 _get_send_message_op_flags_from_state(state),
764 )
765 )
766 rpc_event.call.start_server_batch(
767 operations,
768 _send_status_from_server(state, _SEND_STATUS_FROM_SERVER_TOKEN),
769 )
770 state.statused = True
771 _reset_per_message_state(state)
772 state.due.add(_SEND_STATUS_FROM_SERVER_TOKEN)
773
774
775def _unary_response_in_pool(
776 rpc_event: cygrpc.BaseEvent,
777 state: _RPCState,
778 behavior: ArityAgnosticMethodHandler,
779 argument_thunk: Callable[[], Any],
780 request_deserializer: Optional[SerializingFunction],
781 response_serializer: Optional[SerializingFunction],
782) -> None:
783 cygrpc.install_context_from_request_call_event(rpc_event)
784
785 try:
786 argument = argument_thunk()
787 if argument is not None:
788 response, proceed = _call_behavior(
789 rpc_event, state, behavior, argument, request_deserializer
790 )
791 if proceed:
792 serialized_response = _serialize_response(
793 rpc_event, state, response, response_serializer
794 )
795 if serialized_response is not None:
796 _status(rpc_event, state, serialized_response)
797 except Exception: # pylint: disable=broad-except
798 traceback.print_exc()
799 finally:
800 cygrpc.uninstall_context()
801
802
803def _stream_response_in_pool(
804 rpc_event: cygrpc.BaseEvent,
805 state: _RPCState,
806 behavior: ArityAgnosticMethodHandler,
807 argument_thunk: Callable[[], Any],
808 request_deserializer: Optional[DeserializingFunction],
809 response_serializer: Optional[SerializingFunction],
810) -> None:
811 cygrpc.install_context_from_request_call_event(rpc_event)
812
813 def send_response(response: Any) -> None:
814 if response is None:
815 _status(rpc_event, state, None)
816 else:
817 serialized_response = _serialize_response(
818 rpc_event, state, response, response_serializer
819 )
820 if serialized_response is not None:
821 _send_response(rpc_event, state, serialized_response)
822
823 try:
824 argument = argument_thunk()
825 if argument is not None:
826 if (
827 hasattr(behavior, "experimental_non_blocking")
828 and behavior.experimental_non_blocking
829 ):
830 _call_behavior(
831 rpc_event,
832 state,
833 behavior,
834 argument,
835 request_deserializer,
836 send_response_callback=send_response,
837 )
838 else:
839 response_iterator, proceed = _call_behavior(
840 rpc_event, state, behavior, argument, request_deserializer
841 )
842 if proceed:
843 _send_message_callback_to_blocking_iterator_adapter(
844 rpc_event, state, send_response, response_iterator
845 )
846 except Exception: # pylint: disable=broad-except
847 traceback.print_exc()
848 finally:
849 cygrpc.uninstall_context()
850
851
852def _is_rpc_state_active(state: _RPCState) -> bool:
853 return state.client is not _CANCELLED and not state.statused
854
855
856def _send_message_callback_to_blocking_iterator_adapter(
857 rpc_event: cygrpc.BaseEvent,
858 state: _RPCState,
859 send_response_callback: Callable[[ResponseType], None],
860 response_iterator: Iterator[ResponseType],
861) -> None:
862 while True:
863 response, proceed = _take_response_from_response_iterator(
864 rpc_event, state, response_iterator
865 )
866 if proceed:
867 send_response_callback(response)
868 if not _is_rpc_state_active(state):
869 break
870 else:
871 break
872
873
874def _select_thread_pool_for_behavior(
875 behavior: ArityAgnosticMethodHandler,
876 default_thread_pool: futures.ThreadPoolExecutor,
877) -> futures.ThreadPoolExecutor:
878 if hasattr(behavior, "experimental_thread_pool") and isinstance(
879 behavior.experimental_thread_pool, futures.ThreadPoolExecutor
880 ):
881 return behavior.experimental_thread_pool
882 return default_thread_pool
883
884
885def _handle_unary_unary(
886 rpc_event: cygrpc.BaseEvent,
887 state: _RPCState,
888 method_handler: grpc.RpcMethodHandler,
889 default_thread_pool: futures.ThreadPoolExecutor,
890) -> futures.Future:
891 unary_request = _unary_request(
892 rpc_event, state, method_handler.request_deserializer
893 )
894 thread_pool = _select_thread_pool_for_behavior(
895 method_handler.unary_unary, default_thread_pool
896 )
897 return thread_pool.submit(
898 state.context.run,
899 _unary_response_in_pool,
900 rpc_event,
901 state,
902 method_handler.unary_unary,
903 unary_request,
904 method_handler.request_deserializer,
905 method_handler.response_serializer,
906 )
907
908
909def _handle_unary_stream(
910 rpc_event: cygrpc.BaseEvent,
911 state: _RPCState,
912 method_handler: grpc.RpcMethodHandler,
913 default_thread_pool: futures.ThreadPoolExecutor,
914) -> futures.Future:
915 unary_request = _unary_request(
916 rpc_event, state, method_handler.request_deserializer
917 )
918 thread_pool = _select_thread_pool_for_behavior(
919 method_handler.unary_stream, default_thread_pool
920 )
921 return thread_pool.submit(
922 state.context.run,
923 _stream_response_in_pool,
924 rpc_event,
925 state,
926 method_handler.unary_stream,
927 unary_request,
928 method_handler.request_deserializer,
929 method_handler.response_serializer,
930 )
931
932
933def _handle_stream_unary(
934 rpc_event: cygrpc.BaseEvent,
935 state: _RPCState,
936 method_handler: grpc.RpcMethodHandler,
937 default_thread_pool: futures.ThreadPoolExecutor,
938) -> futures.Future:
939 request_iterator = _RequestIterator(
940 state, rpc_event.call, method_handler.request_deserializer
941 )
942 thread_pool = _select_thread_pool_for_behavior(
943 method_handler.stream_unary, default_thread_pool
944 )
945 return thread_pool.submit(
946 state.context.run,
947 _unary_response_in_pool,
948 rpc_event,
949 state,
950 method_handler.stream_unary,
951 lambda: request_iterator,
952 method_handler.request_deserializer,
953 method_handler.response_serializer,
954 )
955
956
957def _handle_stream_stream(
958 rpc_event: cygrpc.BaseEvent,
959 state: _RPCState,
960 method_handler: grpc.RpcMethodHandler,
961 default_thread_pool: futures.ThreadPoolExecutor,
962) -> futures.Future:
963 request_iterator = _RequestIterator(
964 state, rpc_event.call, method_handler.request_deserializer
965 )
966 thread_pool = _select_thread_pool_for_behavior(
967 method_handler.stream_stream, default_thread_pool
968 )
969 return thread_pool.submit(
970 state.context.run,
971 _stream_response_in_pool,
972 rpc_event,
973 state,
974 method_handler.stream_stream,
975 lambda: request_iterator,
976 method_handler.request_deserializer,
977 method_handler.response_serializer,
978 )
979
980
981def _find_method_handler(
982 rpc_event: cygrpc.BaseEvent,
983 state: _RPCState,
984 method_with_handler: _Method,
985 interceptor_pipeline: Optional[_interceptor._ServicePipeline],
986) -> Optional[grpc.RpcMethodHandler]:
987 def query_handlers(
988 handler_call_details: _HandlerCallDetails,
989 ) -> Optional[grpc.RpcMethodHandler]:
990 return method_with_handler.handler(handler_call_details)
991
992 method_name = method_with_handler.name()
993 if not method_name:
994 method_name = _common.decode(rpc_event.call_details.method)
995
996 handler_call_details = _HandlerCallDetails(
997 method_name,
998 rpc_event.invocation_metadata,
999 )
1000
1001 if interceptor_pipeline is not None:
1002 return state.context.run(
1003 interceptor_pipeline.execute, query_handlers, handler_call_details
1004 )
1005 return state.context.run(query_handlers, handler_call_details)
1006
1007
1008def _reject_rpc(
1009 rpc_event: cygrpc.BaseEvent,
1010 rpc_state: _RPCState,
1011 status: cygrpc.StatusCode,
1012 details: bytes,
1013):
1014 operations = (
1015 _get_initial_metadata_operation(rpc_state, None),
1016 cygrpc.ReceiveCloseOnServerOperation(_EMPTY_FLAGS),
1017 cygrpc.SendStatusFromServerOperation(
1018 None, status, details, _EMPTY_FLAGS
1019 ),
1020 )
1021 rpc_event.call.start_server_batch(
1022 operations,
1023 lambda _ignored_event: (
1024 rpc_state,
1025 (),
1026 ),
1027 )
1028
1029
1030def _handle_with_method_handler(
1031 rpc_event: cygrpc.BaseEvent,
1032 state: _RPCState,
1033 method_handler: grpc.RpcMethodHandler,
1034 thread_pool: futures.ThreadPoolExecutor,
1035) -> futures.Future:
1036 with state.condition:
1037 rpc_event.call.start_server_batch(
1038 (cygrpc.ReceiveCloseOnServerOperation(_EMPTY_FLAGS),),
1039 _receive_close_on_server(state),
1040 )
1041 state.due.add(_RECEIVE_CLOSE_ON_SERVER_TOKEN)
1042 if method_handler.request_streaming:
1043 if method_handler.response_streaming:
1044 return _handle_stream_stream(
1045 rpc_event, state, method_handler, thread_pool
1046 )
1047 return _handle_stream_unary(
1048 rpc_event, state, method_handler, thread_pool
1049 )
1050 if method_handler.response_streaming:
1051 return _handle_unary_stream(
1052 rpc_event, state, method_handler, thread_pool
1053 )
1054 return _handle_unary_unary(
1055 rpc_event, state, method_handler, thread_pool
1056 )
1057
1058
1059def _handle_call(
1060 rpc_event: cygrpc.BaseEvent,
1061 method_with_handler: _Method,
1062 interceptor_pipeline: Optional[_interceptor._ServicePipeline],
1063 thread_pool: futures.ThreadPoolExecutor,
1064 concurrency_exceeded: bool,
1065) -> Tuple[Optional[_RPCState], Optional[futures.Future]]:
1066 """Handles RPC based on provided handlers.
1067
1068 When receiving a call event from Core, registered method will have its
1069 name as tag, we pass the tag as registered_method_name to this method,
1070 then we can find the handler in registered_method_handlers based on
1071 the method name.
1072
1073 For call event with unregistered method, the method name will be included
1074 in rpc_event.call_details.method and we need to query the generics handlers
1075 to find the actual handler.
1076 """
1077 if not rpc_event.success:
1078 return None, None
1079 if rpc_event.call_details.method or method_with_handler.name():
1080 rpc_state = _RPCState()
1081 try:
1082 method_handler = _find_method_handler(
1083 rpc_event,
1084 rpc_state,
1085 method_with_handler,
1086 interceptor_pipeline,
1087 )
1088 except Exception as exception: # pylint: disable=broad-except
1089 details = "Exception servicing handler: {}".format(exception)
1090 _LOGGER.exception(details)
1091 _reject_rpc(
1092 rpc_event,
1093 rpc_state,
1094 cygrpc.StatusCode.unknown,
1095 b"Error in service handler!",
1096 )
1097 return rpc_state, None
1098 if method_handler is None:
1099 _reject_rpc(
1100 rpc_event,
1101 rpc_state,
1102 cygrpc.StatusCode.unimplemented,
1103 b"Method not found!",
1104 )
1105 return rpc_state, None
1106 if concurrency_exceeded:
1107 _reject_rpc(
1108 rpc_event,
1109 rpc_state,
1110 cygrpc.StatusCode.resource_exhausted,
1111 b"Concurrent RPC limit exceeded!",
1112 )
1113 return rpc_state, None
1114 return (
1115 rpc_state,
1116 _handle_with_method_handler(
1117 rpc_event, rpc_state, method_handler, thread_pool
1118 ),
1119 )
1120 return None, None
1121
1122
1123@enum.unique
1124class _ServerStage(enum.Enum):
1125 STOPPED = "stopped"
1126 STARTED = "started"
1127 GRACE = "grace"
1128
1129
1130class _ServerState(object):
1131 lock: threading.RLock
1132 completion_queue: cygrpc.CompletionQueue
1133 server: cygrpc.Server
1134 generic_handlers: List[grpc.GenericRpcHandler]
1135 registered_method_handlers: Dict[str, grpc.RpcMethodHandler]
1136 interceptor_pipeline: Optional[_interceptor._ServicePipeline]
1137 thread_pool: futures.ThreadPoolExecutor
1138 stage: _ServerStage
1139 termination_event: threading.Event
1140 shutdown_events: List[threading.Event]
1141 maximum_concurrent_rpcs: Optional[int]
1142 active_rpc_count: int
1143 rpc_states: Set[_RPCState]
1144 due: Set[str]
1145 server_deallocated: bool
1146
1147 # pylint: disable=too-many-arguments
1148 def __init__(
1149 self,
1150 completion_queue: cygrpc.CompletionQueue,
1151 server: cygrpc.Server,
1152 generic_handlers: Sequence[grpc.GenericRpcHandler],
1153 interceptor_pipeline: Optional[_interceptor._ServicePipeline],
1154 thread_pool: futures.ThreadPoolExecutor,
1155 maximum_concurrent_rpcs: Optional[int],
1156 ):
1157 self.lock = threading.RLock()
1158 self.completion_queue = completion_queue
1159 self.server = server
1160 self.generic_handlers = list(generic_handlers)
1161 self.interceptor_pipeline = interceptor_pipeline
1162 self.thread_pool = thread_pool
1163 self.stage = _ServerStage.STOPPED
1164 self.termination_event = threading.Event()
1165 self.shutdown_events = [self.termination_event]
1166 self.maximum_concurrent_rpcs = maximum_concurrent_rpcs
1167 self.active_rpc_count = 0
1168 self.registered_method_handlers = {}
1169
1170 # TODO(https://github.com/grpc/grpc/issues/6597): eliminate these fields.
1171 self.rpc_states = set()
1172 self.due = set()
1173
1174 # A "volatile" flag to interrupt the daemon serving thread
1175 self.server_deallocated = False
1176
1177
1178def _add_generic_handlers(
1179 state: _ServerState, generic_handlers: Iterable[grpc.GenericRpcHandler]
1180) -> None:
1181 with state.lock:
1182 state.generic_handlers.extend(generic_handlers)
1183
1184
1185def _add_registered_method_handlers(
1186 state: _ServerState, method_handlers: Dict[str, grpc.RpcMethodHandler]
1187) -> None:
1188 with state.lock:
1189 state.registered_method_handlers.update(method_handlers)
1190
1191
1192def _add_insecure_port(state: _ServerState, address: bytes) -> int:
1193 with state.lock:
1194 return state.server.add_http2_port(address)
1195
1196
1197def _add_secure_port(
1198 state: _ServerState,
1199 address: bytes,
1200 server_credentials: grpc.ServerCredentials,
