88"""
99
1010import inspect
11+ from dataclasses import dataclass
1112from typing import Any , Optional , get_args , get_origin
1213
1314from numba_cfunc_compiler .function_analyzer import InputTypeHandler
@@ -41,28 +42,43 @@ def _parse_list_basket_annotation(ann: Any) -> tuple[bool, Optional[type]]:
4142 return False , None
4243
4344
44- def _parse_dict_basket_annotation (ann : Any ) -> tuple [bool , Optional [type ]]:
45+ @dataclass (frozen = True )
46+ class _DictBasketExpectedType :
47+ key_type : type
48+ element_type : type
49+
50+
51+ def _parse_dict_basket_annotation (ann : Any ) -> tuple [bool , Optional [type ], Optional [type ]]:
4552 key_type = None
4653 value_ann = None
4754
4855 # Legacy syntax: {str: ts[T]} / {int: ts[T]}
4956 if isinstance (ann , dict ):
5057 if len (ann ) != 1 :
51- return True , None
58+ return True , None , None
5259 key_type , value_ann = next (iter (ann .items ()))
5360 # Modern syntax: dict[str, ts[T]] / typing.Dict[str, ts[T]]
5461 elif get_origin (ann ) is dict :
5562 args = get_args (ann )
5663 if len (args ) != 2 :
57- return True , None
64+ return True , None , None
5865 key_type , value_ann = args
5966 else :
60- return False , None
67+ return False , None , None
6168
6269 if key_type not in (str , int ):
63- return False , None
70+ return False , None , None
6471
65- return True , _extract_ts_inner_type (value_ann )
72+ return True , key_type , _extract_ts_inner_type (value_ann )
73+
74+
75+ def _validate_edge_type (param_name : str , edge : Edge , expected_type : type ) -> Edge :
76+ actual_type = edge .tstype .typ
77+ if actual_type is not expected_type and not (expected_type is float and actual_type is int ):
78+ raise TypeError (
79+ f"Argument '{ param_name } ' expected ts[{ expected_type .__name__ } ], got ts[{ actual_type .__name__ } ]"
80+ )
81+ return edge
6682
6783
6884class TsInputHandler (InputTypeHandler ):
@@ -78,7 +94,7 @@ def try_parse(self, param: inspect.Parameter, ann: Any) -> Optional[ParameterInf
7894 def validate_value (self , param_name : str , value : Any , expected_type : Any ) -> Any :
7995 if not isinstance (value , Edge ):
8096 raise TypeError (f"Argument '{ param_name } ' must be an Edge, got { type (value ).__name__ } " )
81- return value
97+ return _validate_edge_type ( param_name , value , expected_type )
8298
8399
84100class ListBasketInputHandler (InputTypeHandler ):
@@ -96,36 +112,48 @@ def try_parse(self, param: inspect.Parameter, ann: Any) -> Optional[ParameterInf
96112 def validate_value (self , param_name : str , value : Any , expected_type : Any ) -> Any :
97113 if not isinstance (value , (list , tuple )):
98114 raise TypeError (f"Argument '{ param_name } ' must be a list, got { type (value ).__name__ } " )
115+ if not value :
116+ raise ValueError (f"List basket '{ param_name } ' must not be empty" )
99117
100118 result = {}
101119 for i , edge in enumerate (value ):
102120 if not isinstance (edge , Edge ):
103121 raise TypeError (f"List basket '{ param_name } [{ i } ]' must be an Edge, got { type (edge ).__name__ } " )
104- result [i ] = edge
122+ result [i ] = _validate_edge_type ( f" { param_name } [ { i } ]" , edge , expected_type )
105123 return result
106124
107125
108126class DictBasketInputHandler (InputTypeHandler ):
109127 """Handles {key_type: ts[type]} and dict[key_type, ts[type]] basket annotations."""
110128
111129 def try_parse (self , param : inspect .Parameter , ann : Any ) -> Optional [ParameterInfo ]:
112- matched , inner_type = _parse_dict_basket_annotation (ann )
130+ matched , key_type , inner_type = _parse_dict_basket_annotation (ann )
113131 if not matched :
114132 return None
115133 if inner_type is None :
116134 raise TypeError (f"Dict basket '{ param .name } ' element ts[type] is missing type argument" )
117135
118- return ParameterInfo (expected_type = inner_type , category = "signal_set" )
136+ return ParameterInfo (
137+ expected_type = _DictBasketExpectedType (key_type = key_type , element_type = inner_type ),
138+ category = "signal_set" ,
139+ )
119140
120141 def validate_value (self , param_name : str , value : Any , expected_type : Any ) -> Any :
121142 if not isinstance (value , dict ):
122143 raise TypeError (f"Argument '{ param_name } ' must be a dict, got { type (value ).__name__ } " )
144+ if not value :
145+ raise ValueError (f"Dict basket '{ param_name } ' must not be empty" )
123146
124147 result = {}
125148 for key , edge in value .items ():
149+ if not isinstance (key , expected_type .key_type ):
150+ raise TypeError (
151+ f"Dict basket '{ param_name } ' key { key !r} must be "
152+ f"{ expected_type .key_type .__name__ } , got { type (key ).__name__ } "
153+ )
126154 if not isinstance (edge , Edge ):
127155 raise TypeError (f"Dict basket '{ param_name } [{ key !r} ]' must be an Edge, got { type (edge ).__name__ } " )
128- result [key ] = edge
156+ result [key ] = _validate_edge_type ( f" { param_name } [ { key !r } ]" , edge , expected_type . element_type )
129157 return result
130158
131159
0 commit comments