diff --git a/packages/proto-plus/proto/message.py b/packages/proto-plus/proto/message.py index 1e5cd9cca623..73407472def2 100644 --- a/packages/proto-plus/proto/message.py +++ b/packages/proto-plus/proto/message.py @@ -17,7 +17,7 @@ import copy import re import warnings -from typing import Any, Dict, List, Optional, Type +from typing import Any, Dict, List, Mapping, Optional, Type, TypeVar, Union import google.protobuf from google.protobuf import descriptor_pb2, message @@ -36,6 +36,8 @@ _upb = has_upb() # Important to cache result here. +_MessageT = TypeVar("_MessageT", bound="Message") + class MessageMeta(type): """A metaclass for building and registering Message subclasses.""" @@ -347,7 +349,9 @@ def wrap(cls, pb): super(cls, instance).__setattr__("_pb", pb) return instance - def serialize(cls, instance) -> bytes: + def serialize( + cls, instance: Union["Message", message.Message, Mapping[str, Any]] + ) -> bytes: """Return the serialized proto. Args: @@ -359,7 +363,7 @@ def serialize(cls, instance) -> bytes: """ return cls.pb(instance, coerce=True).SerializeToString() - def deserialize(cls, payload: bytes) -> "Message": + def deserialize(cls: Type[_MessageT], payload: bytes) -> _MessageT: """Given a serialized proto, deserialize it into a Message instance. Args: