Skip to content

Commit 51017e7

Browse files
committed
fix container memory cleanup
Signed-off-by: Mark Halka <mark.halka2001@gmail.com>
1 parent eeec259 commit 51017e7

5 files changed

Lines changed: 110 additions & 13 deletions

File tree

cpp/csp/python/PyNumbaNode.cpp

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -75,7 +75,8 @@ PyNumbaNode::~PyNumbaNode()
7575

7676
for( size_t i = 0; i < m_stateCount; ++i )
7777
{
78-
if( m_stateArgs[ i ] != nullptr )
78+
if( m_stateArgs[ i ] != nullptr &&
79+
m_containerStateIndices.find( i ) == m_containerStateIndices.end() )
7980
delete[] static_cast<char *>( m_stateArgs[ i ] );
8081
}
8182
delete[] m_stateArgs;
@@ -168,6 +169,7 @@ void PyNumbaNode::initStateArrays( PyObjectPtr stateVariables, PyObjectPtr nrtSt
168169
{
169170
size_t idxVal = static_cast<size_t>( PyLong_AsLongLong( idx ) );
170171
nrtSet.insert( idxVal );
172+
m_containerStateIndices.insert( idxVal );
171173
}
172174
}
173175
}

cpp/csp/python/PyNumbaNode.h

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
#include <csp/python/PyObjectPtr.h>
77
#include <Python.h>
88
#include <cstdint>
9+
#include <unordered_set>
910

1011
namespace csp::python
1112
{
@@ -103,6 +104,7 @@ class PyNumbaNode final : public csp::Node
103104
// State arrays
104105
void ** m_stateArgs = nullptr;
105106
size_t m_stateCount = 0;
107+
std::unordered_set<size_t> m_containerStateIndices;
106108

107109
// Keep Python objects alive
108110
PyObjectPtr m_dataReference;

csp/impl/wiring/csp_numba/input_handlers.py

Lines changed: 39 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
"""
99

1010
import inspect
11+
from dataclasses import dataclass
1112
from typing import Any, Optional, get_args, get_origin
1213

1314
from 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

6884
class 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

84100
class 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

108126
class 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

csp/impl/wiring/csp_numba/signal_set_support.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -332,7 +332,7 @@ def handle_signal_set_iteration(converter, node: ast.For) -> Optional[ast.AST]:
332332
else:
333333
filter_array = None
334334

335-
iter_var = f"_iter_{signal_set_name}"
335+
iter_var = f"_iter_{signal_set_name}_{converter.variable_factory.create_temporary_variable_name()}"
336336
range_call = ast.Call(
337337
func=ast.Name(id="range", ctx=ast.Load()),
338338
args=[ast.Constant(value=signal_set.length)],

csp/tests/test_numba_node.py

Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -541,6 +541,54 @@ def g():
541541

542542

543543
class TestFeatures(unittest.TestCase):
544+
def test_rejects_mismatched_signal_type(self):
545+
@numba_node
546+
def typed_signal(x: ts[int]) -> ts[int]:
547+
return 1
548+
549+
@csp.graph
550+
def g():
551+
csp.add_graph_output("result", typed_signal(csp.const(1.25)))
552+
553+
with self.assertRaisesRegex(TypeError, r"expected ts\[int\], got ts\[float\]"):
554+
csp.build_graph(g)
555+
556+
def test_allows_int_to_float_signal_upcast(self):
557+
@numba_node
558+
def typed_signal(x: ts[float]) -> ts[float]:
559+
return x + 0.5
560+
561+
@csp.graph
562+
def g():
563+
csp.add_graph_output("result", typed_signal(csp.const(1)))
564+
565+
results = csp.run(g, starttime=datetime(2024, 1, 1), endtime=timedelta(seconds=1))
566+
self.assertEqual([v for _, v in results["result"]], [1.5])
567+
568+
def test_rejects_mismatched_basket_element_type(self):
569+
@numba_node
570+
def typed_basket(xs: list[ts[int]]) -> ts[int]:
571+
return 1
572+
573+
@csp.graph
574+
def g():
575+
csp.add_graph_output("result", typed_basket([csp.const(1), csp.const(2.5)]))
576+
577+
with self.assertRaisesRegex(TypeError, r"expected ts\[int\], got ts\[float\]"):
578+
csp.build_graph(g)
579+
580+
def test_rejects_mismatched_dict_basket_key_type(self):
581+
@numba_node
582+
def typed_basket(xs: dict[int, ts[int]]) -> ts[int]:
583+
return 1
584+
585+
@csp.graph
586+
def g():
587+
csp.add_graph_output("result", typed_basket({"bad": csp.const(1)}))
588+
589+
with self.assertRaisesRegex(TypeError, r"key 'bad' must be int, got str"):
590+
csp.build_graph(g)
591+
544592
def test_lifecycle_start(self):
545593
@numba_node
546594
def with_start(x: ts[int]) -> ts[int]:
@@ -796,6 +844,23 @@ def g():
796844
self.assertEqual([v for _, v in results["tick_count"]], [1, 2])
797845
self.assertEqual([v for _, v in results["all_valid"]], [0, 1])
798846

847+
def test_nested_iteration_over_same_basket(self):
848+
@numba_node
849+
def cross_product(xs: list[ts[int]]) -> ts[int]:
850+
total = 0
851+
for i in xs.keys():
852+
for j in xs.keys():
853+
total = total + xs[i] * xs[j]
854+
return total
855+
856+
@csp.graph
857+
def g():
858+
result = cross_product([csp.const(2), csp.const(3)])
859+
csp.add_graph_output("result", result)
860+
861+
results = csp.run(g, starttime=datetime(2024, 1, 1), endtime=timedelta(seconds=1))
862+
self.assertEqual([v for _, v in results["result"]], [25])
863+
799864
def test_multiple_outputs(self):
800865
@numba_node
801866
def multi_out(x: ts[int]) -> csp.Outputs(doubled=ts[int], squared=ts[int], positive=ts[int]):

0 commit comments

Comments
 (0)