1+ import json
2+ from collections .abc import Mapping
13from dataclasses import dataclass
2- from typing import Any
4+ from types import UnionType
5+ from typing import Annotated , Any , Union , get_args , get_origin
6+
7+ TRUE_VALUES = {"true" , "1" , "yes" , "on" , "t" }
8+ FALSE_VALUES = {"false" , "0" , "no" , "off" , "f" }
9+
10+
11+ def _strip_annotated (cast : Any ) -> Any :
12+ while get_origin (cast ) is Annotated :
13+ cast = get_args (cast )[0 ]
14+ return cast
15+
16+
17+ def _is_union (cast : Any ) -> bool :
18+ return get_origin (cast ) in (Union , UnionType )
19+
20+
21+ def _as_list (value : Any ) -> list [Any ]:
22+ if isinstance (value , list ):
23+ return value
24+ if isinstance (value , tuple | set | frozenset ):
25+ return list (value )
26+ return [value ]
27+
28+
29+ def _as_mapping (value : Any ) -> Mapping [Any , Any ]:
30+ if isinstance (value , Mapping ):
31+ return value
32+
33+ if isinstance (value , bytes ):
34+ value = value .decode ("utf-8" )
35+
36+ if isinstance (value , str ):
37+ value = json .loads (value )
38+ if isinstance (value , Mapping ):
39+ return value
40+
41+ return dict (value )
42+
43+
44+ def get_cast_name (cast : Any ) -> str :
45+ return getattr (cast , "__name__" , str (cast ).replace ("typing." , "" ))
346
447
548@dataclass
649class BaseParam :
7- def __cast__ (self , value : Any , cast : type ) -> Any :
8- try :
9- if str (value ).lower () in ("true" , "1" , "yes" , "on" , "t" ) and cast is bool :
10- value = True
11- elif str (value ).lower () in ("false" , "0" , "no" , "off" , "f" ) and cast is bool :
12- value = False
13- else :
14- value = cast (value )
15- except Exception :
16- raise
17- return value
50+ def __cast__ (self , value : Any , cast : Any ) -> Any :
51+ cast = _strip_annotated (cast )
52+
53+ if cast is Any or cast is object :
54+ return value
55+
56+ if _is_union (cast ):
57+ union_types = [typ for typ in get_args (cast ) if typ is not type (None )]
58+ if value is None and len (union_types ) != len (get_args (cast )):
59+ return None
60+
61+ if isinstance (value , str ) and str in union_types and len (union_types ) > 1 :
62+ union_types = [typ for typ in union_types if typ is not str ] + [str ]
63+
64+ for typ in union_types :
65+ try :
66+ return self .__cast__ (value , typ )
67+ except (TypeError , ValueError ):
68+ continue
69+ raise ValueError (f"Cannot cast value { value !r} to { get_cast_name (cast )} " )
70+
71+ origin = get_origin (cast )
72+ args = get_args (cast )
73+
74+ if cast is bool :
75+ if isinstance (value , bool ):
76+ return value
77+ if isinstance (value , str ):
78+ value_lower = value .lower ()
79+ if value_lower in TRUE_VALUES :
80+ return True
81+ if value_lower in FALSE_VALUES :
82+ return False
83+ raise ValueError (f"Cannot cast value { value !r} to bool" )
84+ return bool (value )
85+
86+ if cast is list or origin is list :
87+ values = _as_list (value )
88+ item_cast = args [0 ] if args else Any
89+ return [self .__cast__ (item , item_cast ) for item in values ]
90+
91+ if cast is tuple or origin is tuple :
92+ values = _as_list (value )
93+ if not args :
94+ return tuple (values )
95+ if len (args ) == 2 and args [1 ] is Ellipsis :
96+ return tuple (self .__cast__ (item , args [0 ]) for item in values )
97+ return tuple (
98+ self .__cast__ (item , item_cast )
99+ for item , item_cast in zip (values , args , strict = False )
100+ )
101+
102+ if cast is set or origin is set :
103+ values = _as_list (value )
104+ item_cast = args [0 ] if args else Any
105+ return {self .__cast__ (item , item_cast ) for item in values }
106+
107+ if cast is frozenset or origin is frozenset :
108+ values = _as_list (value )
109+ item_cast = args [0 ] if args else Any
110+ return frozenset (self .__cast__ (item , item_cast ) for item in values )
111+
112+ if cast is dict or origin is dict :
113+ mapping = _as_mapping (value )
114+ key_cast , value_cast = args or (Any , Any )
115+ return {
116+ self .__cast__ (key , key_cast ): self .__cast__ (item , value_cast )
117+ for key , item in mapping .items ()
118+ }
119+
120+ return cast (value )
18121
19122
20123@dataclass
21124class Query (BaseParam ):
22125 default : Any | None = None
23126 alias : str | None = None
24127 required : bool = False
25- cast : type | None = None
128+ cast : Any | None = None
26129 description : str | None = None
27130
28- def resolve (self , value : Any , cast : type ) -> Any :
131+ def resolve (self , value : Any , cast : Any ) -> Any :
29132 return self .__cast__ (value , cast ) if cast else value
30133
31134
@@ -34,10 +137,10 @@ class Header(BaseParam):
34137 value : Any
35138 alias : str | None = None
36139 required : bool = False
37- cast : type | None = None
140+ cast : Any | None = None
38141 description : str | None = None
39142
40- def resolve (self , value : Any , cast : type ) -> Any :
143+ def resolve (self , value : Any , cast : Any ) -> Any :
41144 return self .__cast__ (value , cast ) if cast else value
42145
43146
@@ -46,8 +149,8 @@ class Cookie(BaseParam):
46149 value : Any
47150 alias : str | None = None
48151 required : bool = False
49- cast : type | None = None
152+ cast : Any | None = None
50153 description : str | None = None
51154
52- def resolve (self , value : Any , cast : type ) -> Any :
155+ def resolve (self , value : Any , cast : Any ) -> Any :
53156 return self .__cast__ (value , cast ) if cast else value
0 commit comments