Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_server.py1516 linesDownload Raw Back to grpc
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,

Showing the first 1,200 of 1516 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai