Skip to content

Commit c78da59

Browse files
committed
parser fix
Signed-off-by: Mark Halka <Mark.Halka@Point72.com>
1 parent d9dede6 commit c78da59

6 files changed

Lines changed: 291 additions & 433 deletions

File tree

csp/impl/wiring/csp_numba/csp_node_transformer.py

Lines changed: 17 additions & 83 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,9 @@
11
import ast
2-
import inspect
3-
import textwrap
42
from dataclasses import dataclass, field
5-
from typing import List, Optional, Set, Union
3+
from typing import TYPE_CHECKING, List, Optional, Union
4+
5+
if TYPE_CHECKING:
6+
from csp.impl.wiring.node_parser import ParsedNodeDefinition
67

78

89
@dataclass
@@ -23,6 +24,7 @@ class TransformedNode:
2324
transformed_source: str = ""
2425
start_body: List[ast.AST] = field(default_factory=list) # Code from with csp.start():
2526
stop_body: List[ast.AST] = field(default_factory=list) # Code from with csp.stop():
27+
transformed_ast: Optional[ast.FunctionDef] = None
2628

2729

2830
class CspNodeTransformer(ast.NodeTransformer):
@@ -41,7 +43,6 @@ class CspNodeTransformer(ast.NodeTransformer):
4143

4244
def __init__(self):
4345
self.state_variables: List[StateVariable] = []
44-
self.input_names: Set[str] = set()
4546
self.start_body: List[ast.AST] = []
4647
self.stop_body: List[ast.AST] = []
4748
self.csp_call_transformers = {
@@ -56,38 +57,6 @@ def __init__(self):
5657
op=ast.And(),
5758
),
5859
}
59-
self.special_block_transformers = {
60-
"state": self._transform_state_block,
61-
"start": lambda node: self.start_body.extend(self._flatten_transformed_statements(node.body)) or [],
62-
"stop": lambda node: self.stop_body.extend(self._flatten_transformed_statements(node.body)) or [],
63-
"alarms": lambda node: [],
64-
}
65-
66-
def _is_special_block(self, node: ast.AST, block_name: str) -> bool:
67-
"""Returns True if this is a block we need to handle"""
68-
if not isinstance(node, ast.With):
69-
return False
70-
71-
if len(node.items) != 1:
72-
return False
73-
74-
context_expr = node.items[0].context_expr
75-
76-
if isinstance(context_expr, ast.Call):
77-
func = context_expr.func
78-
if isinstance(func, ast.Name) and func.id == block_name:
79-
return True
80-
if isinstance(func, ast.Attribute):
81-
if isinstance(func.value, ast.Name) and func.value.id == "csp" and func.attr == block_name:
82-
return True
83-
84-
return False
85-
86-
def _get_special_block_name(self, node: ast.AST) -> Optional[str]:
87-
for block_name in self.special_block_transformers:
88-
if self._is_special_block(node, block_name):
89-
return block_name
90-
return None
9160

9261
def _infer_type(self, name: str, value: ast.AST) -> ast.AST:
9362
if isinstance(value, ast.Constant):
@@ -203,21 +172,10 @@ def _create_state_assignment(
203172
self.state_variables.append(StateVariable(name=name, type_annotation=type_annotation, initial_value=value))
204173
return ast.AnnAssign(target=target, annotation=state_type, value=value, simple=simple)
205174

206-
def _transform_state_block(self, node: ast.With) -> List[ast.AST]:
207-
"""
208-
Transform state block variables to have State[type] annotations.
209-
210-
with csp.state():
211-
x = 0
212-
y: float = 1.0
213-
214-
becomes:
215-
x: State[int] = 0
216-
y: State[float] = 1.0
217-
"""
175+
def _transform_state_statements(self, statements) -> List[ast.AST]:
218176
transformed = []
219177

220-
for stmt in node.body:
178+
for stmt in statements:
221179
if isinstance(stmt, ast.AnnAssign):
222180
# Already annotated: y: float = 1.0
223181
name = stmt.target.id if isinstance(stmt.target, ast.Name) else None
@@ -263,31 +221,12 @@ def _flatten_transformed_statements(self, statements: List[ast.AST]) -> List[ast
263221
flattened.append(transformed_stmt)
264222
return flattened
265223

266-
def _transform_body(self, body: List[ast.AST]) -> List[ast.AST]:
267-
transformed = []
268-
269-
for stmt in body:
270-
block_name = self._get_special_block_name(stmt)
271-
if block_name is not None:
272-
transformed.extend(self.special_block_transformers[block_name](stmt))
273-
continue
274-
275-
# Transform the statement
276-
transformed.extend(self._flatten_transformed_statements([stmt]))
277-
278-
return transformed
279-
280-
def _transform_func_def(self, func_def: ast.FunctionDef) -> TransformedNode:
224+
def _reset(self):
281225
self.state_variables = []
282-
self.input_names = set()
283226
self.start_body = []
284227
self.stop_body = []
285228

286-
# Collect input names (don't transform annotations)
287-
for arg in func_def.args.args:
288-
self.input_names.add(arg.arg)
289-
290-
transformed_body = self._transform_body(func_def.body)
229+
def _build_transformed_node(self, func_def: ast.FunctionDef, transformed_body: List[ast.AST]) -> TransformedNode:
291230
new_func_def = ast.FunctionDef(
292231
name=func_def.name,
293232
args=func_def.args, # Keep original args with ts[type] annotations
@@ -304,21 +243,16 @@ def _transform_func_def(self, func_def: ast.FunctionDef) -> TransformedNode:
304243
state_variables=self.state_variables,
305244
transformed_body=transformed_body,
306245
original_ast=func_def,
246+
transformed_ast=new_func_def,
307247
transformed_source=transformed_source,
308248
start_body=self.start_body,
309249
stop_body=self.stop_body,
310250
)
311251

312-
def transform_csp_node(self, func) -> TransformedNode:
313-
if callable(func):
314-
source = textwrap.dedent(inspect.getsource(func))
315-
else:
316-
source = textwrap.dedent(func)
317-
318-
tree = ast.parse(source)
319-
func_def = tree.body[0]
320-
321-
if not isinstance(func_def, ast.FunctionDef):
322-
raise ValueError("Expected a function definition")
323-
324-
return self._transform_func_def(func_def)
252+
def transform_csp_node(self, definition: "ParsedNodeDefinition") -> TransformedNode:
253+
self._reset()
254+
self.start_body = self._flatten_transformed_statements(definition.blocks.start)
255+
self.stop_body = self._flatten_transformed_statements(definition.blocks.stop)
256+
transformed_body = self._transform_state_statements(definition.blocks.state)
257+
transformed_body.extend(self._flatten_transformed_statements(definition.blocks.body))
258+
return self._build_transformed_node(definition.funcdef, transformed_body)
Lines changed: 9 additions & 145 deletions
Original file line numberDiff line numberDiff line change
@@ -1,165 +1,29 @@
1-
"""
2-
CSP timeseries input handlers for numba_node.
3-
4-
Provides:
5-
- TsInputHandler for ts[type] parameters
6-
- ListBasketInputHandler for [ts[type]] list basket parameters
7-
- DictBasketInputHandler for {key: ts[type]} dict basket parameters
8-
"""
1+
"""Pass CSP-parsed timeseries input metadata to numba_cfunc_compiler."""
92

103
import inspect
114
from dataclasses import dataclass
12-
from typing import Any, Optional, get_args, get_origin
5+
from typing import Any, Optional
136

147
from numba_cfunc_compiler.function_analyzer import InputTypeHandler
158
from numba_cfunc_compiler.models import ParameterInfo
169

17-
from csp.impl.types.tstype import TsType
18-
from csp.impl.wiring.edge import Edge
19-
20-
21-
def _extract_ts_inner_type(ann: Any) -> Optional[type]:
22-
if get_origin(ann) is not TsType:
23-
return None
24-
args = get_args(ann)
25-
return args[0] if args else getattr(ann, "typ", None)
26-
27-
28-
def _parse_list_basket_annotation(ann: Any) -> tuple[bool, Optional[type]]:
29-
# Legacy syntax: [ts[T]]
30-
if isinstance(ann, list):
31-
if len(ann) != 1:
32-
return True, None
33-
return True, _extract_ts_inner_type(ann[0])
34-
35-
# Modern syntax: list[ts[T]] / typing.List[ts[T]]
36-
if get_origin(ann) is list:
37-
args = get_args(ann)
38-
if len(args) != 1:
39-
return True, None
40-
return True, _extract_ts_inner_type(args[0])
41-
42-
return False, None
43-
4410

4511
@dataclass(frozen=True)
46-
class _DictBasketExpectedType:
47-
key_type: type
48-
element_type: type
49-
12+
class CspInputMetadata:
13+
category: str
5014

51-
def _parse_dict_basket_annotation(ann: Any) -> tuple[bool, Optional[type], Optional[type]]:
52-
key_type = None
53-
value_ann = None
54-
55-
# Legacy syntax: {str: ts[T]} / {int: ts[T]}
56-
if isinstance(ann, dict):
57-
if len(ann) != 1:
58-
return True, None, None
59-
key_type, value_ann = next(iter(ann.items()))
60-
# Modern syntax: dict[str, ts[T]] / typing.Dict[str, ts[T]]
61-
elif get_origin(ann) is dict:
62-
args = get_args(ann)
63-
if len(args) != 2:
64-
return True, None, None
65-
key_type, value_ann = args
66-
else:
67-
return False, None, None
68-
69-
if key_type not in (str, int):
70-
return False, None, None
71-
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
82-
83-
84-
class TsInputHandler(InputTypeHandler):
85-
"""Handles ts[type] input annotations."""
8615

16+
class CspInputHandler(InputTypeHandler):
8717
def try_parse(self, param: inspect.Parameter, ann: Any) -> Optional[ParameterInfo]:
88-
inner_type = _extract_ts_inner_type(ann)
89-
if inner_type is None:
18+
if not isinstance(ann, CspInputMetadata):
9019
return None
91-
92-
return ParameterInfo(expected_type=inner_type, category="signal")
20+
return ParameterInfo(expected_type=ann, category=ann.category)
9321

9422
def validate_value(self, param_name: str, value: Any, expected_type: Any) -> Any:
95-
if not isinstance(value, Edge):
96-
raise TypeError(f"Argument '{param_name}' must be an Edge, got {type(value).__name__}")
97-
return _validate_edge_type(param_name, value, expected_type)
98-
99-
100-
class ListBasketInputHandler(InputTypeHandler):
101-
"""Handles [ts[type]] and list[ts[type]] basket input annotations."""
102-
103-
def try_parse(self, param: inspect.Parameter, ann: Any) -> Optional[ParameterInfo]:
104-
matched, inner_type = _parse_list_basket_annotation(ann)
105-
if not matched:
106-
return None
107-
if inner_type is None:
108-
raise TypeError(f"List basket '{param.name}' element ts[type] is missing type argument")
109-
110-
return ParameterInfo(expected_type=inner_type, category="signal_set")
111-
112-
def validate_value(self, param_name: str, value: Any, expected_type: Any) -> Any:
113-
if not isinstance(value, (list, tuple)):
114-
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")
117-
118-
result = {}
119-
for i, edge in enumerate(value):
120-
if not isinstance(edge, Edge):
121-
raise TypeError(f"List basket '{param_name}[{i}]' must be an Edge, got {type(edge).__name__}")
122-
result[i] = _validate_edge_type(f"{param_name}[{i}]", edge, expected_type)
123-
return result
124-
125-
126-
class DictBasketInputHandler(InputTypeHandler):
127-
"""Handles {key_type: ts[type]} and dict[key_type, ts[type]] basket annotations."""
128-
129-
def try_parse(self, param: inspect.Parameter, ann: Any) -> Optional[ParameterInfo]:
130-
matched, key_type, inner_type = _parse_dict_basket_annotation(ann)
131-
if not matched:
132-
return None
133-
if inner_type is None:
134-
raise TypeError(f"Dict basket '{param.name}' element ts[type] is missing type argument")
135-
136-
return ParameterInfo(
137-
expected_type=_DictBasketExpectedType(key_type=key_type, element_type=inner_type),
138-
category="signal_set",
139-
)
140-
141-
def validate_value(self, param_name: str, value: Any, expected_type: Any) -> Any:
142-
if not isinstance(value, dict):
143-
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")
146-
147-
result = {}
148-
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-
)
154-
if not isinstance(edge, Edge):
155-
raise TypeError(f"Dict basket '{param_name}[{key!r}]' must be an Edge, got {type(edge).__name__}")
156-
result[key] = _validate_edge_type(f"{param_name}[{key!r}]", edge, expected_type.element_type)
157-
return result
23+
return value
15824

15925

16026
def register():
16127
from numba_cfunc_compiler.function_analyzer import FunctionAnalyzer
16228

163-
FunctionAnalyzer.register_input_handler(TsInputHandler())
164-
FunctionAnalyzer.register_input_handler(ListBasketInputHandler())
165-
FunctionAnalyzer.register_input_handler(DictBasketInputHandler())
29+
FunctionAnalyzer.register_input_handler(CspInputHandler())

0 commit comments

Comments
 (0)