1919from functools import lru_cache , reduce
2020from types import GenericAlias
2121from typing import Any , Dict , List , NamedTuple , Optional , Type , cast
22- from typing import Literal as TypingLiteral
2322
2423import msgpack
2524from dataclasses_json import DataClassJsonMixin , dataclass_json
@@ -1088,9 +1087,9 @@ def assert_type(self, t: Type[enum.Enum], v: T):
10881087 raise TypeTransformerFailedError (f"Value { v } is not in Enum { t } " )
10891088
10901089
1091- class LiteralTypeTransformer (TypeTransformer [TypingLiteral ]):
1090+ class LiteralTypeTransformer (TypeTransformer [object ]):
10921091 def __init__ (self ):
1093- super ().__init__ ("LiteralTypeTransformer" , TypingLiteral )
1092+ super ().__init__ ("LiteralTypeTransformer" , object )
10941093
10951094 def get_literal_type (self , t : Type ) -> LiteralType :
10961095 args = get_args (t )
@@ -1116,23 +1115,23 @@ def get_literal_type(self, t: Type) -> LiteralType:
11161115 else :
11171116 raise TypeTransformerFailedError (f"Unsupported Literal base type: { base_type } " )
11181117
1119- def to_literal (self , ctx : FlyteContext , python_val : T , python_type : Type [ T ] , expected : LiteralType ) -> Literal :
1120- if expected .simple == SimpleType .STRING :
1121- return StrTransformer .to_literal (ctx , python_val , python_type , expected )
1122- elif expected .simple == SimpleType .INTEGER :
1123- return IntTransformer .to_literal (ctx , python_val , python_type , expected )
1124- elif expected .simple == SimpleType .FLOAT :
1125- return FloatTransformer .to_literal (ctx , python_val , python_type , expected )
1126- elif expected .simple == SimpleType .BOOLEAN :
1127- return BoolTransformer .to_literal (ctx , python_val , python_type , expected )
1128- elif expected .simple == SimpleType .DATETIME :
1129- return DatetimeTransformer .to_literal (ctx , python_val , python_type , expected )
1130- elif expected .simple == SimpleType .DURATION :
1131- return TimedeltaTransformer .to_literal (ctx , python_val , python_type , expected )
1118+ def to_literal (self , ctx : FlyteContext , python_val : T , python_type : Type , expected : LiteralType ) -> Literal :
1119+ if expected .simple == SimpleType .STRING and isinstance ( python_val , str ) :
1120+ return StrTransformer .to_literal (ctx , python_val , str , expected )
1121+ elif expected .simple == SimpleType .INTEGER and isinstance ( python_val , int ) :
1122+ return IntTransformer .to_literal (ctx , python_val , int , expected )
1123+ elif expected .simple == SimpleType .FLOAT and isinstance ( python_val , float ) :
1124+ return FloatTransformer .to_literal (ctx , python_val , float , expected )
1125+ elif expected .simple == SimpleType .BOOLEAN and isinstance ( python_val , bool ) :
1126+ return BoolTransformer .to_literal (ctx , python_val , bool , expected )
1127+ elif expected .simple == SimpleType .DATETIME and isinstance ( python_val , datetime . datetime ) :
1128+ return DatetimeTransformer .to_literal (ctx , python_val , datetime . datetime , expected )
1129+ elif expected .simple == SimpleType .DURATION and isinstance ( python_val , datetime . timedelta ) :
1130+ return TimedeltaTransformer .to_literal (ctx , python_val , datetime . timedelta , expected )
11321131 else :
11331132 raise TypeError (f"Unsupported LiteralType for LiteralTypeTransformer: { expected .simple } " )
11341133
1135- def to_python_value (self , ctx : FlyteContext , lv : Literal , expected_python_type : Type [ T ] ) -> T :
1134+ def to_python_value (self , ctx : FlyteContext , lv : Literal , expected_python_type : Type ) -> object :
11361135 if lv .scalar .primitive .string_value is not None :
11371136 return StrTransformer .to_python_value (ctx , lv , str )
11381137 elif lv .scalar .primitive .integer is not None :
@@ -1301,7 +1300,7 @@ def _get_transformer(cls, python_type: Type) -> Optional[TypeTransformer[T]]:
13011300 # Special case: prevent that for a type `FooEnum(str, Enum)`, the str transformer is used.
13021301 return cls ._ENUM_TRANSFORMER
13031302
1304- if get_origin (python_type ) == TypingLiteral :
1303+ if get_origin (python_type ) == typing . Literal :
13051304 return cls ._LITERAL_TYPE_TRANSFORMER
13061305
13071306 if hasattr (python_type , "__origin__" ):
@@ -2688,7 +2687,6 @@ def _register_default_type_transformers():
26882687 TypeEngine .register (BinaryIOTransformer ())
26892688 TypeEngine .register (EnumTransformer ())
26902689 TypeEngine .register (ProtobufTransformer ())
2691- TypeEngine .register (LiteralTypeTransformer ())
26922690
26932691 # inner type is. Also unsupported are typing's Tuples. Even though you can look inside them, Flyte's type system
26942692 # doesn't support these currently.
0 commit comments