From 5c812c23a6f21282ba73bf72b835f09795fddb4a Mon Sep 17 00:00:00 2001 From: luisleo526 Date: Mon, 20 Jul 2026 21:43:56 +0800 Subject: [PATCH 1/8] fix: target type callable UDT na lifecycle --- pineforge_codegen/codegen/base.py | 67 ++++++++- pineforge_codegen/codegen/emit_top.py | 14 ++ pineforge_codegen/codegen/visit_call.py | 4 +- pineforge_codegen/codegen/visit_expr.py | 64 +++++++-- pineforge_codegen/codegen/visit_stmt.py | 26 ++++ tests/test_udt_na_lifecycle.py | 179 ++++++++++++++++++++++++ 6 files changed, 342 insertions(+), 12 deletions(-) create mode 100644 tests/test_udt_na_lifecycle.py diff --git a/pineforge_codegen/codegen/base.py b/pineforge_codegen/codegen/base.py index f130101..0be0c55 100644 --- a/pineforge_codegen/codegen/base.py +++ b/pineforge_codegen/codegen/base.py @@ -802,6 +802,12 @@ def func_var_storage(owner: str, raw_name: str) -> str: # scalar/non-drawing tombstone that prevents an outer same-named handle # from leaking inward. Block push/pop provides sibling isolation. self._lexical_drawing_types: dict[str, str | None] = {} + # Source-ordered arbitrary UDT bindings. Drawing handles have their + # specialized registry above, while this parallel map preserves the + # exact authored UDT name needed to target-type a later bare ``na`` + # reassignment. ``None`` is an ordinary scalar/collection tombstone; + # block push/pop keeps sibling declarations independent. + self._lexical_udt_types: dict[str, str | None] = {} # Source-ordered lexical Series status. A False tombstone prevents a # scalar local from inheriting a same-spelled global entry from the # legacy raw-name ``ctx.series_vars`` union. @@ -1468,6 +1474,48 @@ def _callable_var_collection_spec( candidates.append(spec) return candidates[0] if len(candidates) == 1 else None + def _callable_var_udt_spec( + self, name: str, owner_func: str | None = None) -> TypeSpec | None: + """Exact UDT TypeSpec for one callable persistent member identity. + + The analyzer's legacy ``_udt_var_types`` registry is keyed by raw Pine + spelling, so sibling declarations such as ``state`` / ``state__blk1`` + can overwrite each other. Declaration metadata retains both the exact + collision-safe member name and lexical owner; use it for base members + and every written-callsite clone. + """ + safe = self._safe_name(name) + candidates_by_owner: dict[str, list[TypeSpec]] = {} + metadata = getattr(self.ctx, "var_member_metadata_by_node", {}) or {} + specs = getattr(self.ctx, "var_member_type_specs_by_node", {}) or {} + owners = getattr(self.ctx, "var_member_owners_by_node", {}) or {} + for node_id, meta in metadata.items(): + _node, member_name, _ptype, _init_str, is_callable_scoped = meta + if (not is_callable_scoped + or self._safe_name(member_name) != safe): + continue + spec = specs.get(node_id) + owner = owners.get(node_id) + if (owner is None or spec is None + or spec.kind != "udt" + or spec.name not in self._udt_defs): + continue + owner_candidates = candidates_by_owner.setdefault(owner, []) + if spec not in owner_candidates: + owner_candidates.append(spec) + + if owner_func is not None: + owned = candidates_by_owner.get(owner_func, []) + if len(owned) == 1: + return owned[0] + + candidates: list[TypeSpec] = [] + for owned in candidates_by_owner.values(): + for spec in owned: + if spec not in candidates: + candidates.append(spec) + return candidates[0] if len(candidates) == 1 else None + def _emit_cloned_var_decl(self, orig_safe: str, cloned_safe: str, series_suffix: str, lines: list[str], owner_func: str | None = None) -> None: @@ -1483,6 +1531,7 @@ def _emit_cloned_var_decl(self, orig_safe: str, cloned_safe: str, collection_spec = self._callable_var_collection_spec( vname, owner_func ) + udt_spec = self._callable_var_udt_spec(vname, owner_func) drawing_cpp = self._drawing_var_member_cpp_types.get(vname) if (drawing_cpp is not None and orig_safe in self._series_var_member_names): @@ -1506,12 +1555,16 @@ def _emit_cloned_var_decl(self, orig_safe: str, cloned_safe: str, lines.append( f" {drawing_cpp} {cloned_safe} = {drawing_cpp}{{}};" ) - elif vname in self._udt_var_types: + elif udt_spec is not None or vname in self._udt_var_types: # Drawing handle / UDT var clone must match the original's # type (Line/Label/Box/), not the coarse PineType # default (double) — otherwise the clone can't hold the # handle and drawing access on it reads a garbage / na id. - udt_t = self._udt_var_types[vname] + udt_t = ( + udt_spec.name + if udt_spec is not None + else self._udt_var_types[vname] + ) handle_cpp = DRAWING_TYPE_TO_CPP.get(udt_t, udt_t) lines.append(f" {handle_cpp} {cloned_safe} = {handle_cpp}{{}};") elif vname in self._runtime_scalar_var_init_members: @@ -3234,7 +3287,9 @@ def generate(self) -> str: f.default, target_cpp_type=( cpp_type - if cpp_type.startswith("PineMap<") + if (cpp_type.startswith("PineMap<") + or cpp_type in self._udt_defs + or cpp_type in DRAWING_TYPE_TO_CPP.values()) else None ), ) @@ -3454,6 +3509,12 @@ def generate(self) -> str: f" {self._type_spec_to_cpp(callable_collection_spec)} {safe};" ) continue + callable_udt_spec = self._callable_var_udt_spec(name) + if callable_udt_spec is not None: + lines.append( + f" {self._type_spec_to_cpp(callable_udt_spec)} {safe};" + ) + continue # Detect array vars from init expression. Guard the substring # heuristic against a UDT constructor that merely WRAPS array.new / # array.from in its arguments — e.g. diff --git a/pineforge_codegen/codegen/emit_top.py b/pineforge_codegen/codegen/emit_top.py index 781edb8..ea5b7fa 100644 --- a/pineforge_codegen/codegen/emit_top.py +++ b/pineforge_codegen/codegen/emit_top.py @@ -775,6 +775,7 @@ def _emit_history_series_write( def _emit_on_bar(self, lines: list[str]) -> None: self._lexical_drawing_types = {} + self._lexical_udt_types = {} self._lexical_series_bindings = {} self._lexical_known_var_tombstones = set() lines.append(" void on_bar(const Bar& bar) override {") @@ -1421,6 +1422,7 @@ def _emit_func_def(self, fi: FuncInfo, lines: list[str], call_site_idx: int | No prev_func_locals = self._current_func_locals prev_func_local_types = self._current_func_local_types prev_lexical_drawing_types = self._lexical_drawing_types + prev_lexical_udt_types = self._lexical_udt_types prev_lexical_series_bindings = self._lexical_series_bindings prev_known_var_tombstones = self._lexical_known_var_tombstones prev_func_body = getattr(self, "_current_func_body", None) @@ -1438,6 +1440,17 @@ def _emit_func_def(self, fi: FuncInfo, lines: list[str], call_site_idx: int | No self._current_func_locals = {n for n, _, _ in self.ctx.func_var_members.get(fi.name, [])} self._current_func_local_types = {} self._lexical_drawing_types = {} + self._lexical_udt_types = { + param: ( + spec.name + if spec is not None + and spec.kind == "udt" + and spec.name in self._udt_defs + else None + ) + for param in node.params + for spec in (self._current_func_param_specs.get(param),) + } self._lexical_series_bindings = { param: param in self._current_func_series_params for param in node.params @@ -1527,6 +1540,7 @@ def _emit_func_def(self, fi: FuncInfo, lines: list[str], call_site_idx: int | No self._current_func_locals = prev_func_locals self._current_func_local_types = prev_func_local_types self._lexical_drawing_types = prev_lexical_drawing_types + self._lexical_udt_types = prev_lexical_udt_types self._lexical_series_bindings = prev_lexical_series_bindings self._lexical_known_var_tombstones = prev_known_var_tombstones self._current_func_body = prev_func_body diff --git a/pineforge_codegen/codegen/visit_call.py b/pineforge_codegen/codegen/visit_call.py index 2962451..d6b525e 100644 --- a/pineforge_codegen/codegen/visit_call.py +++ b/pineforge_codegen/codegen/visit_call.py @@ -1898,7 +1898,9 @@ def _visit_func_call(self, node: FuncCall) -> str: value_node, target_cpp_type=( f_cpp_type - if f_cpp_type.startswith("PineMap<") + if (f_cpp_type.startswith("PineMap<") + or f_cpp_type in self._udt_defs + or f_cpp_type in DRAWING_TYPE_TO_CPP.values()) else None ), ) diff --git a/pineforge_codegen/codegen/visit_expr.py b/pineforge_codegen/codegen/visit_expr.py index 0ca3bc3..8843faa 100644 --- a/pineforge_codegen/codegen/visit_expr.py +++ b/pineforge_codegen/codegen/visit_expr.py @@ -295,16 +295,60 @@ def _drawing_target_cpp_type( # raw-name registry can be overwritten by a same-named local. return self._global_drawing_cpp_types.get(target_name) + def _udt_target_cpp_type( + self, + *, + target_name: str | None = None, + target_node=None, + type_hint: str | None = None, + ) -> str | None: + """Return the exact authored UDT type for a contextual RHS target. + + Pine's bare ``na`` is target typed. Arbitrary UDTs represent their na + ID with a default-constructed ``T{}``, just as generated drawing + handles do, but the drawing-only registry cannot safely answer for + same-spelled ordinary UDT declarations in sibling blocks/functions. + Prefer explicit/node context, then the source-ordered lexical map, and + consult legacy raw-name metadata only when no lexical binding exists. + """ + spec = self._type_spec_from_hint_name(type_hint) if type_hint else None + if spec is None and target_node is not None: + spec = self._type_spec_from_expr(target_node) + if spec is not None and spec.kind == "udt" and spec.name in self._udt_defs: + return spec.name + if not target_name: + return None + lexical = getattr(self, "_lexical_udt_types", {}) + if target_name in lexical: + udt_name = lexical[target_name] + return udt_name if udt_name in self._udt_defs else None + param_spec = getattr(self, "_current_func_param_specs", {}).get( + target_name + ) + if (param_spec is not None + and param_spec.kind == "udt" + and param_spec.name in self._udt_defs): + return param_spec.name + local_cpp = getattr(self, "_current_func_local_types", {}).get( + target_name + ) + if local_cpp is not None: + local_cpp = local_cpp.removesuffix("&") + return local_cpp if local_cpp in self._udt_defs else None + udt_name = self._udt_var_types.get(target_name) + return udt_name if udt_name in self._udt_defs else None + def _visit_rhs_value(self, value_node, target_name: str | None = None, target_cpp_type: str | None = None) -> str: """Visit an assignment / declaration RHS. A bare ``na`` lowers to a type-appropriate initializer for the target - instead of ``na()``: drawing handles brace-init to their na - handle (``Box{}`` / ``Line{}`` / …); ``std::string``/``int``/``int64_t``/ - ``bool`` use ``na()``. Without this, ``Box b = na;`` and - ``string s = na;`` would both emit ``na()`` and fail to compile - (no viable ``operator=`` / conversion). Every other RHS lowers unchanged. + instead of ``na()``: drawing and authored UDT handles + brace-init to their na object (``Box{}`` / ``State{}``), while + ``std::string``/``int``/``int64_t``/``bool`` use ``na()``. Without + this, ``State s = na;`` and ``string x = na;`` would both emit + ``na()`` and fail to compile (no viable assignment/conversion). + Every other RHS lowers unchanged. """ drawing_target = self._drawing_target_cpp_type( target_name, @@ -313,6 +357,8 @@ def _visit_rhs_value(self, value_node, target_name: str | None = None, if self._is_na_expr(value_node): if drawing_target is not None: return f"{drawing_target}{{}}" + if target_cpp_type in self._udt_defs: + return f"{target_cpp_type}{{}}" if target_cpp_type and target_cpp_type.startswith("PineMap<"): # PineMap's default constructor is the typed ``na`` ID. A # map.new call is the only operation that allocates storage. @@ -322,12 +368,14 @@ def _visit_rhs_value(self, value_node, target_name: str | None = None, if (isinstance(value_node, Ternary) and ((target_cpp_type and target_cpp_type.startswith("PineMap<")) - or drawing_target is not None)): + or drawing_target is not None + or target_cpp_type in self._udt_defs)): # C++ cannot deduce a common type for ``na()`` and a # PineMap/drawing handle. Pine's ternary is target typed, so # propagate the exact declared/reassignment target into both arms. - # Arbitrary UDTs, arrays, matrices, and scalar ternaries retain the - # established generic expression path. + # Arrays, matrices, and scalar ternaries retain the established + # generic expression path; collection IDs require a nullable + # runtime representation before target typing alone can be safe. branch_target = drawing_target or target_cpp_type condition = self._visit_expr(value_node.condition) true_value = self._visit_rhs_value( diff --git a/pineforge_codegen/codegen/visit_stmt.py b/pineforge_codegen/codegen/visit_stmt.py index f10d70e..ad8513c 100644 --- a/pineforge_codegen/codegen/visit_stmt.py +++ b/pineforge_codegen/codegen/visit_stmt.py @@ -256,6 +256,13 @@ def _visit_stmt(self, node: ASTNode, lines: list[str], indent: int) -> None: if decl_spec is not None and decl_spec.kind == "udt" else None ) + self._lexical_udt_types[node.name] = ( + decl_spec.name + if (decl_spec is not None + and decl_spec.kind == "udt" + and decl_spec.name in self._udt_defs) + else None + ) if ( node.name == "map" and getattr(self, "_block_map_visibility_depth", 0) > 0 @@ -856,6 +863,11 @@ def _visit_assignment(self, node: Assignment, lines: list[str], pad: str) -> Non name=target_name, target_node=node.target if target_name is None else None, ) + if selection_cpp_type is None: + selection_cpp_type = self._udt_target_cpp_type( + target_name=target_name, + target_node=node.target if target_name is None else None, + ) if selection_cpp_type is None: selection_cpp_type = self._drawing_target_cpp_type( target_name, @@ -894,6 +906,10 @@ def _visit_assignment(self, node: Assignment, lines: list[str], pad: str) -> Non target_cpp_type = self._map_target_cpp_type( target_node=node.target, ) + if target_cpp_type is None: + target_cpp_type = self._udt_target_cpp_type( + target_node=node.target, + ) val_cpp = self._visit_rhs_value( node.value, target_cpp_type=target_cpp_type ) @@ -972,6 +988,8 @@ def _visit_assignment(self, node: Assignment, lines: list[str], pad: str) -> Non # is_na()). Only computed for bare na — every other RHS is # unaffected. tct = self._map_target_cpp_type(name=target_name) + if tct is None: + tct = self._udt_target_cpp_type(target_name=target_name) if tct is None and self._is_na_expr(node.value): tct = self._na_reassign_cpp_type(target_name) val_cpp = self._visit_rhs_value(node.value, target_name, target_cpp_type=tct) @@ -985,6 +1003,8 @@ def _visit_assignment(self, node: Assignment, lines: list[str], pad: str) -> Non lines.append(f"{pad}{safe} {node.op} {val_cpp};") else: tct = self._map_target_cpp_type(name=target_name) + if tct is None: + tct = self._udt_target_cpp_type(target_name=target_name) if tct is None and self._is_na_expr(node.value): tct = self._na_reassign_cpp_type(target_name) val_cpp = self._visit_rhs_value(node.value, target_name, target_cpp_type=tct) @@ -1231,6 +1251,8 @@ def _push_block_var_remap(self, owner): self._block_map_visibility_depth = previous_map_depth + 1 saved_drawing_types = self._lexical_drawing_types self._lexical_drawing_types = dict(saved_drawing_types) + saved_udt_types = self._lexical_udt_types + self._lexical_udt_types = dict(saved_udt_types) saved_series_bindings = self._lexical_series_bindings self._lexical_series_bindings = dict(saved_series_bindings) saved_known_tombstones = self._lexical_known_var_tombstones @@ -1247,6 +1269,7 @@ def _push_block_var_remap(self, owner): previous_map_visible, previous_map_depth, saved_drawing_types, + saved_udt_types, saved_series_bindings, saved_known_tombstones, ) @@ -1282,6 +1305,7 @@ def _push_block_var_remap(self, owner): previous_map_visible, previous_map_depth, saved_drawing_types, + saved_udt_types, saved_series_bindings, saved_known_tombstones, ) @@ -1293,6 +1317,7 @@ def _pop_block_var_remap(self, saved) -> None: previous_map_visible, previous_map_depth, saved_drawing_types, + saved_udt_types, saved_series_bindings, saved_known_tombstones, ) = saved @@ -1310,6 +1335,7 @@ def _pop_block_var_remap(self, saved) -> None: ) = saved_collections finally: self._lexical_drawing_types = saved_drawing_types + self._lexical_udt_types = saved_udt_types self._lexical_series_bindings = saved_series_bindings self._lexical_known_var_tombstones = saved_known_tombstones self._block_map_binding_visible = previous_map_visible diff --git a/tests/test_udt_na_lifecycle.py b/tests/test_udt_na_lifecycle.py new file mode 100644 index 0000000..9fbb419 --- /dev/null +++ b/tests/test_udt_na_lifecycle.py @@ -0,0 +1,179 @@ +"""Target-typed Pine ``na`` for UDT declarations and lifecycle resets.""" + +from __future__ import annotations + +import re + +from pineforge_codegen import transpile +from tests._compile import compile_cpp +from tests.test_runtime_var_initialization import _compile_and_run + + +_MULTICALL_SOURCE = r'''//@version=6 +strategy("UDT na multicall lifecycle") +type State + float value +probe(bool gate, bool reset, float seed) => + if gate + var State state = na + bool startedNull = na(state) + if startedNull + state := State.new(seed) + state.value += 1.0 + if reset + state := na + (startedNull ? 1000.0 : 0.0) + (na(state) ? 100.0 : state.value) + else + -1.0 +never = probe(false, false, 1.0) +early = probe(bar_index >= 1, bar_index == 2, 10.0) +late = probe(bar_index >= 3, false, 100.0) +''' + + +_METHOD_SOURCE = r'''//@version=6 +strategy("UDT na method lifecycle") +type State + float value +type Carrier + float seed +method probe(Carrier self, bool gate, bool reset) => + if gate + var State state = na + bool startedNull = na(state) + if startedNull + state := State.new(self.seed) + state.value += 1.0 + if reset + state := na + (startedNull ? 1000.0 : 0.0) + (na(state) ? 100.0 : state.value) + else + -1.0 +var Carrier firstCarrier = Carrier.new(1.0) +var Carrier secondCarrier = Carrier.new(10.0) +never = firstCarrier.probe(false, false) +early = secondCarrier.probe(bar_index >= 1, bar_index == 2) +late = firstCarrier.probe(bar_index >= 3, false) +''' + + +_DRIVER = r''' +#include +#include +int main() { + Bar bars[] = { + Bar{1.0, 2.0, 0.0, 1.0, 1.0, 0}, + Bar{2.0, 3.0, 1.0, 2.0, 1.0, 60000}, + Bar{3.0, 4.0, 2.0, 3.0, 1.0, 120000}, + Bar{4.0, 5.0, 3.0, 4.0, 1.0, 180000}, + }; + GeneratedStrategy strategy; + strategy.run(bars, 4); + if (!strategy.last_error().empty()) return 2; + std::cout << std::fixed << std::setprecision(1) + << strategy.never << " " + << strategy.early << " " + << strategy.late << "\n"; +} +''' + + +def test_callable_udt_na_reset_keeps_written_callsites_independent() -> None: + cpp = transpile(_MULTICALL_SOURCE) + assert not re.search(r"state(?:_cs\d+)? = na\(\);", cpp) + for target in ("state", "state_cs1", "state_cs2"): + assert re.search(rf"^\s+(?:this->)?{target} = State\{{\}};$", cpp, re.M) + compile_cpp(cpp, label="udt-na-multicall") + assert _compile_and_run(cpp + _DRIVER) == "-1.0 1011.0 1101.0\n" + + +def test_method_udt_na_reset_keeps_written_callsites_independent() -> None: + cpp = transpile(_METHOD_SOURCE) + assert not re.search(r"state(?:_cs\d+)? = na\(\);", cpp) + for target in ("state", "state_cs1", "state_cs2"): + assert re.search(rf"^\s+(?:this->)?{target} = State\{{\}};$", cpp, re.M) + compile_cpp(cpp, label="udt-na-method") + assert _compile_and_run(cpp + _DRIVER) == "-1.0 1011.0 1002.0\n" + + +def test_plain_same_raw_name_with_different_udt_owners_stays_lexical() -> None: + source = r'''//@version=6 +strategy("UDT na owner collision") +type Left + float value +type Right + int value +left(bool reset) => + Left state = Left.new(1.0) + if reset + state := na + na(state) +right(bool reset) => + Right state = Right.new(2) + if reset + state := na + na(state) +leftNull = left(true) +rightNull = right(true) + ''' + cpp = transpile(source) + assert "state = Left{};" in cpp + assert "state = Right{};" in cpp + compile_cpp(cpp, label="udt-na-owner-collision") + + +def test_sibling_branch_udt_na_targets_follow_exact_storage_names() -> None: + source = r'''//@version=6 +strategy("UDT na sibling declarations") +type Left + float value +type Right + int value +probe(bool chooseLeft) => + if chooseLeft + var Left state = Left.new(1.0) + state := na + na(state) + else + var Right state = Right.new(2) + state := na + na(state) +leftNull = probe(true) +rightNull = probe(false) +''' + cpp = transpile(source) + assert "state = Left{};" in cpp + assert "state__blk1 = Right{};" in cpp + compile_cpp(cpp, label="udt-na-sibling-branches") + + +def test_plain_and_nested_udt_na_contexts_are_target_typed() -> None: + source = r'''//@version=6 +strategy("UDT contextual na") +type Inner + float value +type Outer + Inner inner = na + Inner other +make(bool choose) => + Inner local = na + local := choose ? Inner.new(1.0) : na + local := if choose + local + else + na + local := switch choose + true => local + => na + Outer.new(na, local) +var Outer holder = make(true) +holder.inner := na +ok = na(holder.inner) and not na(holder.other) +''' + cpp = transpile(source) + assert "Inner inner = Inner{};" in cpp + assert "Inner local = Inner{};" in cpp + assert cpp.count("local = Inner{};") >= 3 + assert "holder.inner = Inner{};" in cpp + assert "Outer{.inner = Inner{}" in cpp + compile_cpp(cpp, label="udt-na-contexts") From 478b406e697827d484983d8f5220724254d880bb Mon Sep 17 00:00:00 2001 From: luisleo526 Date: Mon, 20 Jul 2026 22:42:19 +0800 Subject: [PATCH 2/8] fix: complete target-typed UDT na lifecycle --- pineforge_codegen/analyzer/base.py | 71 +++++++++++++++++++ pineforge_codegen/analyzer/types.py | 30 +++++++++ pineforge_codegen/codegen/emit_top.py | 3 +- pineforge_codegen/codegen/types.py | 3 +- pineforge_codegen/codegen/visit_stmt.py | 13 +++- tests/test_udt_na_lifecycle.py | 90 +++++++++++++++++++++++++ 6 files changed, 207 insertions(+), 3 deletions(-) diff --git a/pineforge_codegen/analyzer/base.py b/pineforge_codegen/analyzer/base.py index 8fcca2f..2c39a82 100644 --- a/pineforge_codegen/analyzer/base.py +++ b/pineforge_codegen/analyzer/base.py @@ -3201,6 +3201,63 @@ def _udt_name_from_ctor(self, value: ASTNode) -> str | None: return None return owner + def _udt_name_from_nullable_ctor_selection( + self, value: ASTNode | None + ) -> str | None: + """Exact user-UDT type for ctor-only nullable selections. + + Keep this deliberately narrower than generic UDT expression + inference. In particular, a terminal ``array.get(...UDT...)`` is a + reference-identity surface with its own fail-closed rules; treating + every UDT-valued expression selected against ``na`` as a by-value + return would accidentally bypass those rules. + """ + nullable = object() + + def terminal(body: list[ASTNode]) -> ASTNode | None: + if not body: + return None + node = body[-1] + return node.expr if isinstance(node, ExprStmt) else node + + def resolve(node: ASTNode | None) -> str | None | object: + if node is None: + return nullable + if isinstance(node, ExprStmt): + return resolve(node.expr) + if isinstance(node, NaLiteral): + return nullable + direct = self._udt_name_from_ctor(node) + if direct in self._udt_fields: + return direct + if isinstance(node, Ternary): + return merge((resolve(node.true_val), resolve(node.false_val))) + if isinstance(node, IfStmt): + return merge(( + resolve(terminal(node.body)), + resolve(terminal(node.else_body)), + )) + if isinstance(node, SwitchStmt): + results = [ + resolve(terminal(branch)) + for _case, branch in node.cases + ] + results.append(resolve(terminal(node.default_body))) + return merge(results) + return None + + def merge(results) -> str | None | object: + resolved = list(results) + if any(item is None for item in resolved): + return None + concrete = {item for item in resolved if item is not nullable} + if not concrete: + return nullable + return next(iter(concrete)) if len(concrete) == 1 else None + + result = resolve(value) + return result if isinstance(result, str) else None + def _func_terminal_drawing_type(self, func_node: FuncDef) -> str | None: """Resolve the drawing-handle / UDT type of a function's terminal (return) expression for cases the direct ``_udt_name_from_ctor`` on the @@ -3933,6 +3990,8 @@ def _visit_FuncDef(self, node: FuncDef) -> PineType: if node.body: ret_expr = terminal_ret_expr udt_ret = self._udt_name_from_ctor(ret_expr) if ret_expr is not None else None + if udt_ret is None: + udt_ret = self._udt_name_from_nullable_ctor_selection(ret_expr) if (udt_ret is None and terminal_direct_return_spec is not None and terminal_direct_return_spec.kind == "udt"): @@ -4097,6 +4156,7 @@ def _visit_MethodDef(self, node) -> PineType: self._nested_ta_touched = set() terminal_ret_expr = self._direct_terminal_return_expr(node) return_type_spec = None + method_udt_return = None try: for stmt in node.body: ret_type = self._visit(stmt) @@ -4104,6 +4164,14 @@ def _visit_MethodDef(self, node) -> PineType: terminal_spec = self._type_spec_from_expr(terminal_ret_expr) if terminal_spec is not None and terminal_spec.kind == "map": return_type_spec = terminal_spec + method_udt_return = ( + self._udt_name_from_ctor(terminal_ret_expr) + or self._udt_name_from_nullable_ctor_selection( + terminal_ret_expr + ) + ) + if method_udt_return is not None: + return_type_spec = TypeSpec.udt(method_udt_return) finally: self._global_scope = old_global self._collection_scope_stack.pop() @@ -4123,6 +4191,8 @@ def _visit_MethodDef(self, node) -> PineType: if hi > lo: self._func_ta_ranges[method_key] = (lo, hi) self._symbols.exit_scope() + if method_udt_return is not None: + self._func_udt_return_types[method_key] = method_udt_return # Detect tuple return on UDT methods (mirrors the regular FuncDef logic # earlier in this file). Without this, codegen emits the method with a @@ -4168,6 +4238,7 @@ def _visit_MethodDef(self, node) -> PineType: param_defaults=param_defaults, param_type_specs=param_specs, return_type_spec=return_type_spec, + udt_return_type=method_udt_return, ) self._func_infos.append(fi) return PineType.VOID diff --git a/pineforge_codegen/analyzer/types.py b/pineforge_codegen/analyzer/types.py index dc2348b..f4857bc 100644 --- a/pineforge_codegen/analyzer/types.py +++ b/pineforge_codegen/analyzer/types.py @@ -194,6 +194,20 @@ def _type_spec_from_expr(self, value: ASTNode | None) -> TypeSpec | None: if isinstance(value, Ternary): true_spec = self._type_spec_from_expr(value.true_val) false_spec = self._type_spec_from_expr(value.false_val) + + def direct_user_udt_ctor_name(node: ASTNode) -> str | None: + if not isinstance(node, FuncCall): + return None + callee = node.callee + if not ( + isinstance(callee, MemberAccess) + and isinstance(callee.object, Identifier) + and callee.member == "new" + ): + return None + name = callee.object.name + return name if name in self._udt_fields else None + # Selecting between two values of the same user-defined type # preserves that receiver type. Codegen already applies this # rule; the analyzer must agree so stateful method calls on a UDT @@ -217,6 +231,22 @@ def _type_spec_from_expr(self, value: ASTNode | None) -> TypeSpec | None: and false_spec.kind == "map" and isinstance(value.true_val, NaLiteral)): return false_spec + # A direct user-UDT constructor selected against bare ``na`` has + # one unambiguous value type. Require the constructor AST itself, + # not merely an inferred UDT expression, so temporary array-element + # identity returns continue to fail closed on their own surface. + true_ctor = direct_user_udt_ctor_name(value.true_val) + if (true_spec is not None + and true_spec.kind == "udt" + and true_spec.name == true_ctor + and isinstance(value.false_val, NaLiteral)): + return true_spec + false_ctor = direct_user_udt_ctor_name(value.false_val) + if (false_spec is not None + and false_spec.kind == "udt" + and false_spec.name == false_ctor + and isinstance(value.true_val, NaLiteral)): + return false_spec # Drawing handles are nullable reference-like values in Pine. A # bare ``na`` arm therefore acquires the other arm's exact handle # type, just like the established PineMap path above. Keep this diff --git a/pineforge_codegen/codegen/emit_top.py b/pineforge_codegen/codegen/emit_top.py index ea5b7fa..f7850f1 100644 --- a/pineforge_codegen/codegen/emit_top.py +++ b/pineforge_codegen/codegen/emit_top.py @@ -1383,7 +1383,8 @@ def _emit_func_def(self, fi: FuncInfo, lines: list[str], call_site_idx: int | No rhs_return_cpp_type = ( ret_type if (ret_type.startswith("PineMap<") - or ret_type in DRAWING_TYPE_TO_CPP.values()) + or ret_type in DRAWING_TYPE_TO_CPP.values() + or ret_type in self._udt_defs) else None ) diff --git a/pineforge_codegen/codegen/types.py b/pineforge_codegen/codegen/types.py index 2b9aff7..663544c 100644 --- a/pineforge_codegen/codegen/types.py +++ b/pineforge_codegen/codegen/types.py @@ -492,7 +492,8 @@ def terminal_expr(body): if return_spec is not None: return return_spec udt_return = getattr(func_info, "udt_return_type", None) - if udt_return in DRAWING_TYPE_TO_CPP: + if (udt_return in DRAWING_TYPE_TO_CPP + or udt_return in self._udt_defs): return TypeSpec.udt(udt_return) # ticker.* constructors (inherit/standard/heikinashi) return a symbol # string; without this the member-type inference defaults to double diff --git a/pineforge_codegen/codegen/visit_stmt.py b/pineforge_codegen/codegen/visit_stmt.py index ad8513c..1a556a6 100644 --- a/pineforge_codegen/codegen/visit_stmt.py +++ b/pineforge_codegen/codegen/visit_stmt.py @@ -724,12 +724,18 @@ def remember_local_type(cpp_type: str | None) -> None: cpp_type if (cpp_type is not None and (cpp_type.startswith("PineMap<") - or cpp_type in DRAWING_TYPE_TO_CPP.values())) + or cpp_type in DRAWING_TYPE_TO_CPP.values() + or cpp_type in self._udt_defs)) else self._map_target_cpp_type( name=node.name, type_hint=node.type_hint, ) ) + if selection_cpp_type is None: + selection_cpp_type = self._udt_target_cpp_type( + target_name=node.name, + type_hint=node.type_hint, + ) if selection_cpp_type is None: selection_cpp_type = self._drawing_target_cpp_type( node.name, @@ -821,6 +827,11 @@ def remember_local_type(cpp_type: str | None) -> None: name=node.name, type_hint=node.type_hint, ) + if target_cpp_type is None: + target_cpp_type = self._udt_target_cpp_type( + target_name=node.name, + type_hint=node.type_hint, + ) cpp_val = self._visit_rhs_value( node.value, node.name, diff --git a/tests/test_udt_na_lifecycle.py b/tests/test_udt_na_lifecycle.py index 9fbb419..20e6197 100644 --- a/tests/test_udt_na_lifecycle.py +++ b/tests/test_udt_na_lifecycle.py @@ -177,3 +177,93 @@ def test_plain_and_nested_udt_na_contexts_are_target_typed() -> None: assert "holder.inner = Inner{};" in cpp assert "Outer{.inner = Inner{}" in cpp compile_cpp(cpp, label="udt-na-contexts") + + +def test_chart_scope_typed_udt_na_is_target_typed_each_bar() -> None: + source = r'''//@version=6 +strategy("chart UDT na lifecycle") +type State + float value +State state = na +int observed = na(state) ? 1 : 0 +''' + cpp = transpile(source) + assert "state = State{};" in cpp + assert "state = na();" not in cpp + compile_cpp(cpp, label="udt-na-chart-scope") + + +def test_direct_udt_ctor_na_selection_return_is_target_typed() -> None: + source = r'''//@version=6 +strategy("UDT nullable selection return") +type State + float value +choose(bool enabled, float seed) => enabled ? State.new(seed) : na +State missing = choose(false, 1.0) +State present = choose(true, 7.5) +float observed = (na(missing) ? 100.0 : 0.0) + + (na(present) ? 0.0 : present.value) +''' + cpp = transpile(source) + assert re.search(r"State choose\(bool enabled, double seed\)", cpp) + assert "State{.value = seed, .__pf_na = false}) : (State{})" in cpp + compile_cpp(cpp, label="udt-na-selection-return") + driver = r''' +#include +int main() { + Bar bars[] = {Bar{1.0, 2.0, 0.0, 1.0, 1.0, 1700000000000LL}}; + GeneratedStrategy strategy; + strategy.run(bars, 1); + if (!strategy.last_error().empty()) return 7; + std::cout << strategy.observed << "\n"; +} +''' + assert _compile_and_run(cpp + driver) == "107.5\n" + + +def test_if_and_switch_udt_ctor_na_returns_are_target_typed() -> None: + bodies = { + "if": r'''choose(bool enabled) => + if enabled + State.new(2.5) + else + na''', + "switch": r'''choose(bool enabled) => + switch enabled + true => State.new(2.5) + => na''', + } + for label, body in bodies.items(): + source = f'''//@version=6 +strategy("UDT {label} nullable return") +type State + float value +{body} +State observed = choose(false) +''' + cpp = transpile(source) + assert re.search(r"State choose\(bool enabled\)", cpp) + assert "_func_ret = State{};" in cpp + compile_cpp(cpp, label=f"udt-na-{label}-return") + + +def test_udt_method_nullable_udt_return_carries_exact_type() -> None: + source = r'''//@version=6 +strategy("UDT method nullable return") +type State + float value +type Factory + float seed +method choose(Factory self, bool enabled) => + enabled ? State.new(self.seed) : na +var Factory factory = Factory.new(9.25) +State missing = factory.choose(false) +State present = factory.choose(true) +float observed = (na(missing) ? 100.0 : 0.0) + present.value +''' + cpp = transpile(source) + assert re.search( + r"State _udt_Factory_choose(?:_cs\d+)?\(Factory& self, bool enabled\)", + cpp, + ) + compile_cpp(cpp, label="udt-method-na-return") From c6bdf6012511b15a44319e7a4616865a0ce11f21 Mon Sep 17 00:00:00 2001 From: luisleo526 Date: Mon, 20 Jul 2026 23:09:48 +0800 Subject: [PATCH 3/8] fix: isolate global UDT target types --- pineforge_codegen/codegen/base.py | 36 +++++++++++++++++++++---- pineforge_codegen/codegen/emit_top.py | 6 ++--- pineforge_codegen/codegen/types.py | 8 ++++++ pineforge_codegen/codegen/visit_expr.py | 4 +++ tests/test_udt_na_lifecycle.py | 26 ++++++++++++++++++ 5 files changed, 72 insertions(+), 8 deletions(-) diff --git a/pineforge_codegen/codegen/base.py b/pineforge_codegen/codegen/base.py index 0be0c55..fdda429 100644 --- a/pineforge_codegen/codegen/base.py +++ b/pineforge_codegen/codegen/base.py @@ -1025,7 +1025,7 @@ def _is_runtime_scalar_var_initializer( ) if is_series and name in self.ctx.series_vars: return False - udt_type = self._udt_var_types.get(name) + udt_type = self._member_udt_type(name) if udt_type in self._udt_defs: return False type_spec = self._collection_types.get(name) @@ -1057,6 +1057,24 @@ def _prepare_runtime_scalar_var_initializers(self) -> None: self._drawing_var_member_cpp_types: dict[str, str] = {} self._drawing_var_decl_info_by_node: dict[int, dict] = {} self._global_drawing_cpp_types: dict[str, str] = {} + # Exact direct-program UDT identity, including ``None`` tombstones for + # primitive globals. The analyzer's legacy UDT registry is keyed only + # by raw spelling and can be overwritten by an unrelated callable + # local with the same name. + self._global_udt_types: dict[str, str | None] = {} + for stmt in self.ctx.ast.body: + if not isinstance(stmt, VarDecl): + continue + spec = ( + self._type_spec_from_hint_name(stmt.type_hint) + if stmt.type_hint + else self._type_spec_from_expr(stmt.value) + ) + self._global_udt_types[stmt.name] = ( + spec.name + if spec is not None and spec.kind == "udt" + else None + ) self._runtime_var_init_flags: dict[tuple[int, str], str] = {} used_names = set(self._all_member_names) @@ -1516,6 +1534,13 @@ def _callable_var_udt_spec( candidates.append(spec) return candidates[0] if len(candidates) == 1 else None + def _member_udt_type(self, name: str) -> str | None: + """Exact UDT type for class-member storage, with global tombstones.""" + global_types = getattr(self, "_global_udt_types", {}) + if name in global_types: + return global_types[name] + return self._udt_var_types.get(name) + def _emit_cloned_var_decl(self, orig_safe: str, cloned_safe: str, series_suffix: str, lines: list[str], owner_func: str | None = None) -> None: @@ -3583,9 +3608,10 @@ def generate(self) -> str: # C++ handle struct (Series when also history-referenced). # Drawing names are NOT in _udt_defs, so the udt branch below would # self-zero them to double; handle them first. + member_udt_type = self._member_udt_type(name) _draw_cpp = ( self._drawing_var_member_cpp_types.get(name) - or DRAWING_TYPE_TO_CPP.get(self._udt_var_types.get(name)) + or DRAWING_TYPE_TO_CPP.get(member_udt_type) ) if _draw_cpp is not None: if safe in self._series_var_member_names: @@ -3593,7 +3619,7 @@ def generate(self) -> str: else: lines.append(f" {_draw_cpp} {safe};") continue - udt_type = self._udt_var_types.get(name) + udt_type = member_udt_type if udt_type not in self._udt_defs: udt_type = None if udt_type is None: @@ -3699,12 +3725,12 @@ def generate(self) -> str: "localPivots", "securityPivotPointsArray", "pivotPointsArray", ): lines.append(f" std::vector {safe} = std::vector();") - elif name in self._udt_var_types: + elif self._member_udt_type(name) is not None: # Non-var global of UDT type — declare as the struct so # downstream method dispatch works. Probes: # data/validation/udt-method-probe-19-array-of-udt-method, # data/validation/udt-method-probe-20-udt-return-from-func. - udt_t = self._udt_var_types[name] + udt_t = self._member_udt_type(name) # Drawing handle global (L-N6 / U): map line/box/label/linefill # to the C++ handle struct (the default is na, id=-1). _draw_cpp = DRAWING_TYPE_TO_CPP.get(udt_t) diff --git a/pineforge_codegen/codegen/emit_top.py b/pineforge_codegen/codegen/emit_top.py index f7850f1..9d1ef8e 100644 --- a/pineforge_codegen/codegen/emit_top.py +++ b/pineforge_codegen/codegen/emit_top.py @@ -579,15 +579,15 @@ def _emit_constructor(self, lines: list[str]) -> None: # UDT-typed var members (``var SDZone z = na``) default-construct to # na via the struct's in-class ``__pf_na = true``; a ctor init like # ``z(na())`` would not type-match the struct member. - if name in self._udt_var_types and self._udt_var_types[name] in self._udt_defs: + member_udt_type = self._member_udt_type(name) + if member_udt_type in self._udt_defs: continue # Drawing handle var member (L-N3): ``var line x`` / ``var box b`` # default-construct to {-1} (na). A ``b(na())`` ctor init # would not type-match the handle struct — skip it (the in-class # member default is the once-only persistent na init). if (name in self._drawing_var_member_cpp_types - or (name in self._udt_var_types - and self._udt_var_types[name] in DRAWING_TYPE_TO_CPP)): + or member_udt_type in DRAWING_TYPE_TO_CPP): continue if safe not in self._series_var_member_names: cpp_val = self._resolve_known(init_expr) diff --git a/pineforge_codegen/codegen/types.py b/pineforge_codegen/codegen/types.py index 663544c..83be366 100644 --- a/pineforge_codegen/codegen/types.py +++ b/pineforge_codegen/codegen/types.py @@ -396,6 +396,14 @@ def _type_spec_from_expr(self, node) -> TypeSpec | None: for pine_name, cpp_name in DRAWING_TYPE_TO_CPP.items(): if cpp_name == global_cpp: return TypeSpec.udt(pine_name) + global_udt_types = getattr(self, "_global_udt_types", {}) + if node.name in global_udt_types: + global_udt = global_udt_types[node.name] + return ( + TypeSpec.udt(global_udt) + if global_udt is not None + else None + ) if node.name in self._udt_var_types: return TypeSpec.udt(self._udt_var_types[node.name]) # Drawing-typed method/function parameter (L.6d / U.5): a ``line ln`` diff --git a/pineforge_codegen/codegen/visit_expr.py b/pineforge_codegen/codegen/visit_expr.py index 8843faa..971786e 100644 --- a/pineforge_codegen/codegen/visit_expr.py +++ b/pineforge_codegen/codegen/visit_expr.py @@ -335,6 +335,10 @@ def _udt_target_cpp_type( if local_cpp is not None: local_cpp = local_cpp.removesuffix("&") return local_cpp if local_cpp in self._udt_defs else None + global_types = getattr(self, "_global_udt_types", {}) + if target_name in global_types: + udt_name = global_types[target_name] + return udt_name if udt_name in self._udt_defs else None udt_name = self._udt_var_types.get(target_name) return udt_name if udt_name in self._udt_defs else None diff --git a/tests/test_udt_na_lifecycle.py b/tests/test_udt_na_lifecycle.py index 20e6197..12b23fb 100644 --- a/tests/test_udt_na_lifecycle.py +++ b/tests/test_udt_na_lifecycle.py @@ -193,6 +193,32 @@ def test_chart_scope_typed_udt_na_is_target_typed_each_bar() -> None: compile_cpp(cpp, label="udt-na-chart-scope") +def test_callable_udt_locals_do_not_retype_same_named_global_scalars() -> None: + source = r'''//@version=6 +strategy("UDT raw-name global isolation") +type State + float value +float state = 2.0 +var float tracker = 3.0 +probe() => + State state = State.new(9.0) + State tracker = State.new(10.0) + state := na + tracker := na + na(state) and na(tracker) +bool localNulls = probe() +float retained = state + tracker +''' + cpp = transpile(source) + assert re.search(r"^\s+double state(?: = [^;]+)?;$", cpp, re.M) + assert re.search(r"^\s+double tracker(?: = [^;]+)?;$", cpp, re.M) + assert "State state = State{.value = 9.0" in cpp + assert "State tracker = State{.value = 10.0" in cpp + assert "state = State{};" in cpp + assert "tracker = State{};" in cpp + compile_cpp(cpp, label="udt-na-raw-name-global-isolation") + + def test_direct_udt_ctor_na_selection_return_is_target_typed() -> None: source = r'''//@version=6 strategy("UDT nullable selection return") From adacabf4c362d23fe94831bf1927800cefd0988e Mon Sep 17 00:00:00 2001 From: luisleo526 Date: Mon, 20 Jul 2026 23:43:18 +0800 Subject: [PATCH 4/8] fix: resolve lexical UDT identities exactly --- pineforge_codegen/codegen/base.py | 12 +-- pineforge_codegen/codegen/emit_top.py | 6 +- pineforge_codegen/codegen/types.py | 116 ++++++++++++++++-------- pineforge_codegen/codegen/visit_call.py | 10 +- pineforge_codegen/codegen/visit_expr.py | 23 +---- pineforge_codegen/codegen/visit_stmt.py | 48 +++++++++- tests/test_udt_na_lifecycle.py | 83 +++++++++++++++++ 7 files changed, 222 insertions(+), 76 deletions(-) diff --git a/pineforge_codegen/codegen/base.py b/pineforge_codegen/codegen/base.py index fdda429..f6f4371 100644 --- a/pineforge_codegen/codegen/base.py +++ b/pineforge_codegen/codegen/base.py @@ -4084,15 +4084,9 @@ def _is_omitted_udt_field(self, node) -> bool: """ if not isinstance(node, MemberAccess): return False - # Cheap path: receiver is a bare identifier we already track in - # ``_udt_var_types`` (the common case — ``m.tag``, ``s.ln``). - if isinstance(node.object, Identifier): - udt_name = self._udt_var_types.get(node.object.name) - if udt_name is None: - return False - return node.member in self._udt_omitted_fields.get(udt_name, ()) - # General path: try to infer the receiver's UDT type via the same - # spec-resolver visit_expr uses for fallback member access. + # Resolve even bare identifiers through the lexical/exact TypeSpec + # path. The legacy raw-name UDT registry can describe an unrelated + # callable local and therefore cannot safely drive field omission. recv_spec = self._type_spec_from_expr(node.object) if recv_spec is not None and recv_spec.kind == "udt" and recv_spec.name: return node.member in self._udt_omitted_fields.get(recv_spec.name, ()) diff --git a/pineforge_codegen/codegen/emit_top.py b/pineforge_codegen/codegen/emit_top.py index 9d1ef8e..b5b3111 100644 --- a/pineforge_codegen/codegen/emit_top.py +++ b/pineforge_codegen/codegen/emit_top.py @@ -1447,7 +1447,11 @@ def _emit_func_def(self, fi: FuncInfo, lines: list[str], call_site_idx: int | No if spec is not None and spec.kind == "udt" and spec.name in self._udt_defs - else None + else ( + self._udt_param_udt.get(param) + if self._udt_param_udt.get(param) in self._udt_defs + else None + ) ) for param in node.params for spec in (self._current_func_param_specs.get(param),) diff --git a/pineforge_codegen/codegen/types.py b/pineforge_codegen/codegen/types.py index 83be366..250cbf0 100644 --- a/pineforge_codegen/codegen/types.py +++ b/pineforge_codegen/codegen/types.py @@ -370,48 +370,18 @@ def _type_spec_from_expr(self, node) -> TypeSpec | None: return TypeSpec.primitive("string") if isinstance(node, Identifier): collection_spec = self._collection_spec_for_name(node.name) - if collection_spec is not None: + if (collection_spec is not None + and collection_spec.kind in {"array", "map", "matrix"}): return collection_spec - if node.name in getattr(self, "_lexical_drawing_types", {}): - lexical_cpp = self._lexical_drawing_types[node.name] - if lexical_cpp is None: - return None - for pine_name, cpp_name in DRAWING_TYPE_TO_CPP.items(): - if cpp_name == lexical_cpp: - return TypeSpec.udt(pine_name) - return None - local_cpp = getattr(self, "_current_func_local_types", {}).get( - node.name - ) - if local_cpp is not None: - local_cpp = local_cpp.removesuffix("&") - for pine_name, cpp_name in DRAWING_TYPE_TO_CPP.items(): - if cpp_name == local_cpp: - return TypeSpec.udt(pine_name) - return None - global_cpp = getattr(self, "_global_drawing_cpp_types", {}).get( + found_udt_binding, exact_udt = self._identifier_udt_binding( node.name ) - if global_cpp is not None: - for pine_name, cpp_name in DRAWING_TYPE_TO_CPP.items(): - if cpp_name == global_cpp: - return TypeSpec.udt(pine_name) - global_udt_types = getattr(self, "_global_udt_types", {}) - if node.name in global_udt_types: - global_udt = global_udt_types[node.name] + if found_udt_binding: return ( - TypeSpec.udt(global_udt) - if global_udt is not None + TypeSpec.udt(exact_udt) + if exact_udt is not None else None ) - if node.name in self._udt_var_types: - return TypeSpec.udt(self._udt_var_types[node.name]) - # Drawing-typed method/function parameter (L.6d / U.5): a ``line ln`` - # method receiver registers in _udt_param_udt so its body getters - # resolve to the drawing udt and dispatch through the §4.3 path. - _pu = getattr(self, "_udt_param_udt", None) - if _pu and node.name in _pu and _pu[node.name] in DRAWING_TYPE_TO_CPP: - return TypeSpec.udt(_pu[node.name]) if self._collection_name_is_lexically_shadowed(node.name): return None sym = self.ctx.symbols.resolve(node.name) @@ -637,6 +607,76 @@ def terminal_expr(body): return TypeSpec.udt("line") return None + def _identifier_udt_binding( + self, name: str + ) -> tuple[bool, str | None]: + """Return ``(binding_known, exact_udt_or_none)`` for ``name``. + + ``ctx.udt_var_types`` is a legacy raw-name registry: a declaration in + an unrelated callable can overwrite the type of a same-named global. + Every correctness-sensitive consumer must prefer the active lexical + binding (including scalar tombstones), then exact parameter/local and + direct-program metadata, before consulting that registry. + """ + known_udts = set(self._udt_defs) | set(DRAWING_TYPE_TO_CPP) + + lexical_drawing = getattr(self, "_lexical_drawing_types", {}) + if name in lexical_drawing and lexical_drawing[name] is not None: + cpp_type = lexical_drawing[name] + for pine_name, candidate_cpp in DRAWING_TYPE_TO_CPP.items(): + if candidate_cpp == cpp_type: + return True, pine_name + + lexical = getattr(self, "_lexical_udt_types", {}) + if name in lexical: + candidate = lexical[name] + return True, candidate if candidate in known_udts else None + + param_specs = getattr(self, "_current_func_param_specs", {}) + param_spec = param_specs.get(name) or param_specs.get( + self._safe_name(name) + ) + if (param_spec is not None + and param_spec.kind == "udt" + and param_spec.name in known_udts): + return True, param_spec.name + param_types = getattr(self, "_current_func_param_types", {}) + if name in param_types or self._safe_name(name) in param_types: + return True, None + + local_types = getattr(self, "_current_func_local_types", {}) + if name in local_types: + local_cpp = local_types[name] + candidate = local_cpp.removesuffix("&").removesuffix("*") + return True, candidate if candidate in known_udts else None + + param_udts = getattr(self, "_udt_param_udt", {}) + candidate = param_udts.get(name) or param_udts.get( + self._safe_name(name) + ) + if candidate in known_udts: + return True, candidate + + global_types = getattr(self, "_global_udt_types", {}) + if name in global_types: + candidate = global_types[name] + return True, candidate if candidate in known_udts else None + + member_resolver = getattr(self, "_member_udt_type", None) + candidate = ( + member_resolver(name) + if callable(member_resolver) + else self._udt_var_types.get(name) + ) + if candidate is not None: + return True, candidate if candidate in known_udts else None + return False, None + + def _identifier_udt_type(self, name: str) -> str | None: + """Exact UDT type for ``name``; scalar tombstones return ``None``.""" + _found, udt_type = self._identifier_udt_binding(name) + return udt_type + # ------------------------------------------------------------------ # Method lowering for collection types (used by visit_call paths) # ------------------------------------------------------------------ @@ -1303,7 +1343,7 @@ def _na_reassign_cpp_type(self, name: str) -> str | None: return self._type_spec_to_cpp(collection_spec) if ((collection_spec is not None and collection_spec.kind in {"array", "matrix"}) - or name in self._udt_var_types): + or self._identifier_udt_type(name) is not None): return None cpp_type: str | None = None # 1. ``var`` member (class-scope OR function-local: both are recorded in @@ -1410,7 +1450,7 @@ def _is_udt_lvalue(self, expr) -> str | None: return None if not isinstance(expr, Identifier): return None - udt_t = self._udt_var_types.get(expr.name) + udt_t = self._identifier_udt_type(expr.name) if udt_t is None or udt_t not in self._udt_defs: return None if udt_t in DRAWING_TYPE_TO_CPP: diff --git a/pineforge_codegen/codegen/visit_call.py b/pineforge_codegen/codegen/visit_call.py index d6b525e..2800db8 100644 --- a/pineforge_codegen/codegen/visit_call.py +++ b/pineforge_codegen/codegen/visit_call.py @@ -1125,10 +1125,12 @@ def _visit_func_call(self, node: FuncCall) -> str: f"matrix.{meth_raw}: wrong number of arguments", hint="Check Pine v6 matrix method signature (positional vs keyword).", ) - safe_o = self._safe_name(oname) - udt_t = self._udt_var_types.get(oname) or self._udt_var_types.get(safe_o) - if udt_t is None: - udt_t = self._udt_param_udt.get(oname) or self._udt_param_udt.get(safe_o) + recv_spec = self._type_spec_from_expr(obj) + udt_t = ( + recv_spec.name + if recv_spec is not None and recv_spec.kind == "udt" + else None + ) if udt_t is not None: mk = f"{udt_t}.{meth_raw}" fi_u = self._func_info_map.get(mk) diff --git a/pineforge_codegen/codegen/visit_expr.py b/pineforge_codegen/codegen/visit_expr.py index 971786e..4872420 100644 --- a/pineforge_codegen/codegen/visit_expr.py +++ b/pineforge_codegen/codegen/visit_expr.py @@ -318,28 +318,7 @@ def _udt_target_cpp_type( return spec.name if not target_name: return None - lexical = getattr(self, "_lexical_udt_types", {}) - if target_name in lexical: - udt_name = lexical[target_name] - return udt_name if udt_name in self._udt_defs else None - param_spec = getattr(self, "_current_func_param_specs", {}).get( - target_name - ) - if (param_spec is not None - and param_spec.kind == "udt" - and param_spec.name in self._udt_defs): - return param_spec.name - local_cpp = getattr(self, "_current_func_local_types", {}).get( - target_name - ) - if local_cpp is not None: - local_cpp = local_cpp.removesuffix("&") - return local_cpp if local_cpp in self._udt_defs else None - global_types = getattr(self, "_global_udt_types", {}) - if target_name in global_types: - udt_name = global_types[target_name] - return udt_name if udt_name in self._udt_defs else None - udt_name = self._udt_var_types.get(target_name) + udt_name = self._identifier_udt_type(target_name) return udt_name if udt_name in self._udt_defs else None def _visit_rhs_value(self, value_node, target_name: str | None = None, diff --git a/pineforge_codegen/codegen/visit_stmt.py b/pineforge_codegen/codegen/visit_stmt.py index 1a556a6..7d82fe8 100644 --- a/pineforge_codegen/codegen/visit_stmt.py +++ b/pineforge_codegen/codegen/visit_stmt.py @@ -276,6 +276,7 @@ def _visit_stmt(self, node: ASTNode, lines: list[str], indent: int) -> None: self._visit_assignment(node, lines, pad) elif isinstance(node, TupleAssign): self._visit_tuple_assign(node, lines, pad) + tuple_cpp_types = self._tuple_binding_cpp_types(node) if getattr(self, "_active_func_name", None) is not None: for name in node.names: if name and name != "_": @@ -299,8 +300,23 @@ def _visit_stmt(self, node: ASTNode, lines: list[str], indent: int) -> None: self._active_var_remap = dict(self._active_var_remap) for name in scalar_names: self._active_var_remap.pop(self._safe_name(name), None) - for name in node.names: + for index, name in enumerate(node.names): if name and name != "_": + # A supported tuple destructure creates fresh lexical + # bindings. Tuple elements are primitive on the current + # supported surface, so install an explicit tombstone: an + # unrelated same-named global/callable UDT must not target- + # type a later ``name := na`` as ``State{}``. + self._lexical_udt_types[name] = None + if (getattr(self, "_active_func_name", None) is not None + and not self._decl_binding_is_series( + id(node), name + )): + self._current_func_local_types[name] = ( + tuple_cpp_types[index] + if index < len(tuple_cpp_types) + else "double" + ) self._lexical_series_bindings[name] = ( self._decl_binding_is_series(id(node), name) ) @@ -816,7 +832,12 @@ def remember_local_type(cpp_type: str | None) -> None: # global scope, so a function-local sharing the name keeps its own path. if (node.name in self._udt_array_get_ref_locals and getattr(self, "_current_func_body", None) is None): - udt_t = self._udt_var_types.get(node.name) + udt_t = self._is_udt_lvalue(node.value) + if udt_t is None: + self._codegen_error( + node, + "UDT array-element alias lost its exact element type.", + ) cpp_val = self._visit_rhs_value(node.value, node.name, target_cpp_type=udt_t) lines.append(f"{pad}{udt_t}& {safe} = {cpp_val};") return @@ -1245,6 +1266,29 @@ def emit_call_tuple(call_expr: str) -> None: lines.append(f"{pad}/* unsupported tuple assignment */") + def _tuple_binding_cpp_types(self, node: TupleAssign) -> list[str]: + """Exact supported tuple element types for later lexical operations.""" + count = len(node.names) + if not isinstance(node.value, FuncCall): + return ["double"] * count + func_name, namespace = self._resolve_callee(node.value.callee) + fi = None + if namespace is None: + fi = self._func_info_map.get(func_name) + elif isinstance(node.value.callee, MemberAccess): + recv_spec = self._type_spec_from_expr(node.value.callee.object) + if (recv_spec is not None + and recv_spec.kind == "udt" + and recv_spec.name): + fi = self._func_info_map.get( + f"{recv_spec.name}.{node.value.callee.member}" + ) + if (fi is not None + and fi.node is not None + and getattr(fi, "returns_tuple", False)): + return self._infer_tuple_types(fi.node, count) + return ["double"] * count + def _push_block_var_remap(self, owner): """Activate exact lexical metadata for one branch/loop body. diff --git a/tests/test_udt_na_lifecycle.py b/tests/test_udt_na_lifecycle.py index 12b23fb..4e758a4 100644 --- a/tests/test_udt_na_lifecycle.py +++ b/tests/test_udt_na_lifecycle.py @@ -219,6 +219,89 @@ def test_callable_udt_locals_do_not_retype_same_named_global_scalars() -> None: compile_cpp(cpp, label="udt-na-raw-name-global-isolation") +def test_exact_global_udt_drives_methods_fields_and_lvalue_aliases() -> None: + source = r'''//@version=6 +strategy("exact global UDT identity") +type Left + float value + float tag +type Right + float value + table tag +method read(Left self) => self.value +method read(Right self) => self.value + 100.0 +Left state = Left.new(1.0, 7.0) +shadow() => + Right state = Right.new(2.0) + na(state) +mutate() => + Left alias = state + alias.value := 3.0 + alias.value +bool ignored = shadow() +float methodValue = state.read() +float aliasValue = mutate() +float tagValue = state.tag +''' + cpp = transpile(source) + assert "_udt_Left_read(state)" in cpp + assert "Left& alias = state;" in cpp + assert "tagValue = state.tag;" in cpp + assert "tagValue = /* drawing field omitted */ 0;" not in cpp + compile_cpp(cpp, label="udt-exact-global-identity") + + +def test_global_udt_array_alias_uses_rhs_element_type_not_raw_name() -> None: + source = r'''//@version=6 +strategy("exact global UDT array alias") +type Left + float value +type Right + float value +var array items = array.new() +if barstate.isfirst + items.push(Left.new(1.0)) +for i = 0 to items.size() - 1 + Left state = items.get(i) + state.value := 3.0 +shadow() => + Right state = Right.new(2.0) + na(state) +bool ignored = shadow() +''' + cpp = transpile(source) + assert "Left& state =" in cpp + assert "Right& state =" not in cpp + compile_cpp(cpp, label="udt-exact-global-array-alias") + + +def test_scalar_and_tuple_tombstones_keep_typed_na_reassignment() -> None: + source = r'''//@version=6 +strategy("exact scalar tombstones") +type State + float value +State state = State.new(5.0) +int tracker = 2 +pair() => [1, 2] +shadow() => + State tracker = State.new(9.0) + tracker := na + na(tracker) +tupleProbe() => + [state, other] = pair() + state := na + state +bool ignored = shadow() +int tupleValue = tupleProbe() +tracker := na +bool trackerNull = na(tracker) +''' + cpp = transpile(source) + assert "state = na();" in cpp + assert "tracker = na();" in cpp + compile_cpp(cpp, label="udt-exact-scalar-tuple-tombstones") + + def test_direct_udt_ctor_na_selection_return_is_target_typed() -> None: source = r'''//@version=6 strategy("UDT nullable selection return") From c4552e6fe422af2a08025da3bd425ccf920204e9 Mon Sep 17 00:00:00 2001 From: luisleo526 Date: Mon, 20 Jul 2026 23:50:32 +0800 Subject: [PATCH 5/8] fix: preserve exact UDT member ownership --- pineforge_codegen/analyzer/base.py | 8 ++- pineforge_codegen/codegen/base.py | 33 ++++++++++- tests/test_udt_na_lifecycle.py | 92 ++++++++++++++++++++++++++++++ 3 files changed, 129 insertions(+), 4 deletions(-) diff --git a/pineforge_codegen/analyzer/base.py b/pineforge_codegen/analyzer/base.py index 2c39a82..d107f25 100644 --- a/pineforge_codegen/analyzer/base.py +++ b/pineforge_codegen/analyzer/base.py @@ -2170,12 +2170,16 @@ def _resolved_user_call_name(call: FuncCall, owner: str | None) -> str | None: spec = specs[param_idx] if param_idx < len(specs) else None if spec is not None and spec.kind == "udt": udt_name = spec.name - if udt_name is None: - udt_name = self._udt_var_types.get(recv.name) + # Resolve the surviving exact global/lexical symbol before the + # flat raw-name registry. A later callable-local declaration can + # overwrite that registry and otherwise attach the wrong stateful + # method edge to wrappers and their written call sites. if udt_name is None: spec = self._type_spec_from_expr(recv) if spec is not None and spec.kind == "udt": udt_name = spec.name + if udt_name is None and isinstance(recv, Identifier): + udt_name = self._udt_var_types.get(recv.name) key = f"{udt_name}.{method}" if udt_name else "" return key if key in func_defs else None diff --git a/pineforge_codegen/codegen/base.py b/pineforge_codegen/codegen/base.py index f6f4371..131ce00 100644 --- a/pineforge_codegen/codegen/base.py +++ b/pineforge_codegen/codegen/base.py @@ -1075,6 +1075,31 @@ def _prepare_runtime_scalar_var_initializers(self) -> None: if spec is not None and spec.kind == "udt" else None ) + # Analyzer metadata preserves the collision-safe storage identity for + # every persistent declaration, including chart-scope siblings that + # are not direct Program children. Keep primitive tombstones too: the + # raw-name UDT union must not swap ``state`` and ``state__blk1`` types. + self._member_udt_types: dict[str, str | None] = {} + metadata_by_node = getattr( + self.ctx, "var_member_metadata_by_node", {} + ) or {} + type_specs_by_node = getattr( + self.ctx, "var_member_type_specs_by_node", {} + ) or {} + for node_id, meta in metadata_by_node.items(): + stmt, member_name, _ptype, _init_str, _callable = meta + spec = type_specs_by_node.get(node_id) + if spec is None and isinstance(stmt, VarDecl): + spec = ( + self._type_spec_from_hint_name(stmt.type_hint) + if stmt.type_hint + else self._type_spec_from_expr(stmt.value) + ) + self._member_udt_types[member_name] = ( + spec.name + if spec is not None and spec.kind == "udt" + else None + ) self._runtime_var_init_flags: dict[tuple[int, str], str] = {} used_names = set(self._all_member_names) @@ -1539,6 +1564,9 @@ def _member_udt_type(self, name: str) -> str | None: global_types = getattr(self, "_global_udt_types", {}) if name in global_types: return global_types[name] + member_types = getattr(self, "_member_udt_types", {}) + if name in member_types: + return member_types[name] return self._udt_var_types.get(name) def _emit_cloned_var_decl(self, orig_safe: str, cloned_safe: str, @@ -1557,6 +1585,7 @@ def _emit_cloned_var_decl(self, orig_safe: str, cloned_safe: str, vname, owner_func ) udt_spec = self._callable_var_udt_spec(vname, owner_func) + member_udt_type = self._member_udt_type(vname) drawing_cpp = self._drawing_var_member_cpp_types.get(vname) if (drawing_cpp is not None and orig_safe in self._series_var_member_names): @@ -1580,7 +1609,7 @@ def _emit_cloned_var_decl(self, orig_safe: str, cloned_safe: str, lines.append( f" {drawing_cpp} {cloned_safe} = {drawing_cpp}{{}};" ) - elif udt_spec is not None or vname in self._udt_var_types: + elif udt_spec is not None or member_udt_type is not None: # Drawing handle / UDT var clone must match the original's # type (Line/Label/Box/), not the coarse PineType # default (double) — otherwise the clone can't hold the @@ -1588,7 +1617,7 @@ def _emit_cloned_var_decl(self, orig_safe: str, cloned_safe: str, udt_t = ( udt_spec.name if udt_spec is not None - else self._udt_var_types[vname] + else member_udt_type ) handle_cpp = DRAWING_TYPE_TO_CPP.get(udt_t, udt_t) lines.append(f" {handle_cpp} {cloned_safe} = {handle_cpp}{{}};") diff --git a/tests/test_udt_na_lifecycle.py b/tests/test_udt_na_lifecycle.py index 4e758a4..d0905bb 100644 --- a/tests/test_udt_na_lifecycle.py +++ b/tests/test_udt_na_lifecycle.py @@ -147,6 +147,32 @@ def test_sibling_branch_udt_na_targets_follow_exact_storage_names() -> None: compile_cpp(cpp, label="udt-na-sibling-branches") +def test_chart_scope_sibling_udt_vars_keep_exact_member_types() -> None: + source = r'''//@version=6 +strategy("chart UDT sibling declarations") +type Left + float value +type Right + int value +if close > 0 + var Left state = na + if na(state) + state := Left.new(1.0) + state := na +else + var Right state = na + if na(state) + state := Right.new(2) + state := na +''' + cpp = transpile(source) + assert re.search(r"^\s+Left state;$", cpp, re.M) + assert re.search(r"^\s+Right state__blk1;$", cpp, re.M) + assert "state = Left{};" in cpp + assert "state__blk1 = Right{};" in cpp + compile_cpp(cpp, label="udt-na-chart-sibling-members") + + def test_plain_and_nested_udt_na_contexts_are_target_typed() -> None: source = r'''//@version=6 strategy("UDT contextual na") @@ -302,6 +328,72 @@ def test_scalar_and_tuple_tombstones_keep_typed_na_reassignment() -> None: compile_cpp(cpp, label="udt-exact-scalar-tuple-tombstones") +def test_callable_primitive_clones_ignore_unrelated_udt_raw_name() -> None: + source = r'''//@version=6 +strategy("UDT callable primitive clone isolation") +type State + float value +counter(bool reset) => + var int state = 1 + if reset + state := na + state +shadow() => + State state = State.new(4.0) + na(state) +int first = counter(true) +int second = counter(false) +bool ignored = shadow() +''' + cpp = transpile(source) + assert re.search(r"^\s+int state(?: = [^;]+)?;$", cpp, re.M) + assert re.search(r"^\s+int state_cs1(?: = [^;]+)?;$", cpp, re.M) + assert "State state_cs1" not in cpp + compile_cpp(cpp, label="udt-na-callable-primitive-clone-isolation") + + +def test_stateful_method_edge_uses_exact_global_udt_owner() -> None: + source = r'''//@version=6 +strategy("stateful method exact edge") +type Left + float value +type Right + float value +method accumulate(Left self) => + var float leftTotal = 0.0 + leftTotal += self.value + leftTotal +method accumulate(Right self) => + var float rightTotal = 0.0 + rightTotal += self.value + rightTotal +Left state = Left.new(2.0) +wrapped() => state.accumulate() +shadow() => + Right state = Right.new(3.0) + na(state) +bool ignored = shadow() +float first = wrapped() +float second = wrapped() +''' + cpp = transpile(source) + assert "_udt_Left_accumulate_cs0(state)" in cpp + assert "_udt_Left_accumulate_cs1(state)" in cpp + assert "_udt_Right_accumulate_cs0(state)" not in cpp + compile_cpp(cpp, label="udt-na-stateful-method-exact-edge") + driver = r''' +#include +int main() { + Bar bars[] = {Bar{1.0, 2.0, 0.0, 1.0, 1.0, 1700000000000LL}}; + GeneratedStrategy strategy; + strategy.run(bars, 1); + if (!strategy.last_error().empty()) return 8; + std::cout << strategy.first << " " << strategy.second << "\n"; +} +''' + assert _compile_and_run(cpp + driver) == "2 2\n" + + def test_direct_udt_ctor_na_selection_return_is_target_typed() -> None: source = r'''//@version=6 strategy("UDT nullable selection return") From 559948faa9ba7d1c8d16f63bda078f07e531535a Mon Sep 17 00:00:00 2001 From: luisleo526 Date: Tue, 21 Jul 2026 00:09:18 +0800 Subject: [PATCH 6/8] fix: preserve exact lexical binder types --- pineforge_codegen/codegen/emit_top.py | 14 ++++++++- pineforge_codegen/codegen/types.py | 6 +++- pineforge_codegen/codegen/visit_stmt.py | 38 ++++++++++++++++++++++++- tests/test_udt_na_lifecycle.py | 34 ++++++++++++++++++++++ 4 files changed, 89 insertions(+), 3 deletions(-) diff --git a/pineforge_codegen/codegen/emit_top.py b/pineforge_codegen/codegen/emit_top.py index b5b3111..18a74cf 100644 --- a/pineforge_codegen/codegen/emit_top.py +++ b/pineforge_codegen/codegen/emit_top.py @@ -1440,7 +1440,19 @@ def _emit_func_def(self, fi: FuncInfo, lines: list[str], call_site_idx: int | No self._udt_ptr_alias_locals = set() self._current_func_locals = {n for n, _, _ in self.ctx.func_var_members.get(fi.name, [])} self._current_func_local_types = {} - self._lexical_drawing_types = {} + self._lexical_drawing_types = { + param: DRAWING_TYPE_TO_CPP[drawing_name] + for param in node.params + for drawing_name in ( + ( + self._current_func_param_specs[param].name + if (param in self._current_func_param_specs + and self._current_func_param_specs[param].kind == "udt") + else self._udt_param_udt.get(param) + ), + ) + if drawing_name in DRAWING_TYPE_TO_CPP + } self._lexical_udt_types = { param: ( spec.name diff --git a/pineforge_codegen/codegen/types.py b/pineforge_codegen/codegen/types.py index 250cbf0..9446320 100644 --- a/pineforge_codegen/codegen/types.py +++ b/pineforge_codegen/codegen/types.py @@ -370,8 +370,12 @@ def _type_spec_from_expr(self, node) -> TypeSpec | None: return TypeSpec.primitive("string") if isinstance(node, Identifier): collection_spec = self._collection_spec_for_name(node.name) + # This resolver also carries exact primitive loop/parameter/global + # metadata despite its historical name. Preserve those scalar + # families here (notably bool for array.from); only arbitrary UDTs + # need the collision-safe lexical resolver below. if (collection_spec is not None - and collection_spec.kind in {"array", "map", "matrix"}): + and collection_spec.kind != "udt"): return collection_spec found_udt_binding, exact_udt = self._identifier_udt_binding( node.name diff --git a/pineforge_codegen/codegen/visit_stmt.py b/pineforge_codegen/codegen/visit_stmt.py index 7d82fe8..b49161d 100644 --- a/pineforge_codegen/codegen/visit_stmt.py +++ b/pineforge_codegen/codegen/visit_stmt.py @@ -1505,6 +1505,10 @@ def _visit_for(self, node: ForStmt, lines: list[str], indent: int) -> None: self._current_loop_var_specs[node.var] = TypeSpec.primitive("int") _blk_saved = self._push_block_var_remap(node) if node.var: + # The loop counter is a fresh primitive lexical binding. Keep it + # from inheriting a same-spelled outer/raw UDT or drawing type. + self._lexical_drawing_types[node.var] = None + self._lexical_udt_types[node.var] = None self._lexical_series_bindings[node.var] = False self._lexical_known_var_tombstones.add(node.var) try: @@ -1635,8 +1639,40 @@ def _visit_for_in(self, node, lines: list[str], indent: int) -> None: loop_binding_names = ( [node.var] if node.var else list(node.vars or []) ) - for name in loop_binding_names: + loop_binding_specs = ( + [elem_spec] + if node.var + else [ + tuple_specs[index] if index < len(tuple_specs) else None + for index in range(len(loop_binding_names)) + ] + ) + for index, name in enumerate(loop_binding_names): if name and name != "_": + spec = ( + loop_binding_specs[index] + if index < len(loop_binding_specs) + else None + ) + drawing_name = ( + spec.name + if (spec is not None + and spec.kind == "udt" + and spec.name in DRAWING_TYPE_TO_CPP) + else None + ) + self._lexical_drawing_types[name] = ( + DRAWING_TYPE_TO_CPP[drawing_name] + if drawing_name is not None + else None + ) + self._lexical_udt_types[name] = ( + spec.name + if (spec is not None + and spec.kind == "udt" + and spec.name in self._udt_defs) + else None + ) self._lexical_series_bindings[name] = False self._lexical_known_var_tombstones.add(name) try: diff --git a/tests/test_udt_na_lifecycle.py b/tests/test_udt_na_lifecycle.py index d0905bb..e3f1cdf 100644 --- a/tests/test_udt_na_lifecycle.py +++ b/tests/test_udt_na_lifecycle.py @@ -352,6 +352,40 @@ def test_callable_primitive_clones_ignore_unrelated_udt_raw_name() -> None: compile_cpp(cpp, label="udt-na-callable-primitive-clone-isolation") +def test_primitive_metadata_survives_global_udt_tombstone_registry() -> None: + source = r'''//@version=6 +strategy("exact primitive metadata") +bool enabled = true +array flags = array.from(enabled, false) +''' + cpp = transpile(source) + assert "std::vector flags;" in cpp + assert "flags = std::vector{" in cpp + compile_cpp(cpp, label="udt-exact-primitive-metadata") + + +def test_drawing_method_parameter_keeps_exact_lexical_type() -> None: + source = r'''//@version=6 +strategy("exact drawing method parameter") +type Shadow + float value +method span(line ln) => + ln.get_x2() - ln.get_x1() +poison() => + Shadow ln = Shadow.new(1.0) + na(ln) +var line segment = line.new(bar_index, close, bar_index + 1, close) +float measured = segment.span() +bool ignored = poison() +''' + cpp = transpile(source) + assert "pf_line_get_x2(_pf_lines_, ln)" in cpp + assert "pf_line_get_x1(_pf_lines_, ln)" in cpp + assert "ln.get_x2()" not in cpp + assert "ln.get_x1()" not in cpp + compile_cpp(cpp, label="udt-exact-drawing-method-parameter") + + def test_stateful_method_edge_uses_exact_global_udt_owner() -> None: source = r'''//@version=6 strategy("stateful method exact edge") From 284a196beaefa7c2cf88aa5912b71b6e527726c0 Mon Sep 17 00:00:00 2001 From: luisleo526 Date: Tue, 21 Jul 2026 00:24:31 +0800 Subject: [PATCH 7/8] fix: prioritize lexical UDT parameters --- pineforge_codegen/codegen/types.py | 17 ++++++++++++++++- tests/test_udt_na_lifecycle.py | 23 ++++++++++++++++------- 2 files changed, 32 insertions(+), 8 deletions(-) diff --git a/pineforge_codegen/codegen/types.py b/pineforge_codegen/codegen/types.py index 9446320..6e9c017 100644 --- a/pineforge_codegen/codegen/types.py +++ b/pineforge_codegen/codegen/types.py @@ -217,7 +217,7 @@ def _collection_spec_for_name(self, name: str) -> TypeSpec | None: # historically mask a same-named top-level collection registry. # Declared scalar/UDT parameters do shadow it, while inferred or # declared collection parameters always carry their exact kind. - if (param_spec.kind in {"array", "map", "matrix"} + if (param_spec.kind in {"array", "map", "matrix", "udt"} or name in getattr( self, "_current_func_declared_param_names", set() )): @@ -369,6 +369,21 @@ def _type_spec_from_expr(self, node) -> TypeSpec | None: if isinstance(node, StringLiteral): return TypeSpec.primitive("string") if isinstance(node, Identifier): + # An active lexical UDT/drawing binding must beat a same-spelled + # top-level primitive or collection. This is intentionally the + # exact lexical prefix of _identifier_udt_binding(), not its raw- + # name fallback: method receivers such as ``line ln`` otherwise + # inherit an unrelated global ``ln`` before dispatch is resolved. + lexical_drawing = getattr(self, "_lexical_drawing_types", {}) + drawing_cpp = lexical_drawing.get(node.name) + if drawing_cpp is not None: + for pine_name, cpp_name in DRAWING_TYPE_TO_CPP.items(): + if cpp_name == drawing_cpp: + return TypeSpec.udt(pine_name) + lexical_udts = getattr(self, "_lexical_udt_types", {}) + lexical_udt = lexical_udts.get(node.name) + if lexical_udt in self._udt_defs: + return TypeSpec.udt(lexical_udt) collection_spec = self._collection_spec_for_name(node.name) # This resolver also carries exact primitive loop/parameter/global # metadata despite its historical name. Preserve those scalar diff --git a/tests/test_udt_na_lifecycle.py b/tests/test_udt_na_lifecycle.py index e3f1cdf..3358aa5 100644 --- a/tests/test_udt_na_lifecycle.py +++ b/tests/test_udt_na_lifecycle.py @@ -365,10 +365,16 @@ def test_primitive_metadata_survives_global_udt_tombstone_registry() -> None: def test_drawing_method_parameter_keeps_exact_lexical_type() -> None: - source = r'''//@version=6 + top_level_collisions = { + "primitive": "bool ln = true", + "collection": "array ln = array.from(1)", + } + for label, top_level in top_level_collisions.items(): + source = rf'''//@version=6 strategy("exact drawing method parameter") type Shadow float value +{top_level} method span(line ln) => ln.get_x2() - ln.get_x1() poison() => @@ -378,12 +384,15 @@ def test_drawing_method_parameter_keeps_exact_lexical_type() -> None: float measured = segment.span() bool ignored = poison() ''' - cpp = transpile(source) - assert "pf_line_get_x2(_pf_lines_, ln)" in cpp - assert "pf_line_get_x1(_pf_lines_, ln)" in cpp - assert "ln.get_x2()" not in cpp - assert "ln.get_x1()" not in cpp - compile_cpp(cpp, label="udt-exact-drawing-method-parameter") + cpp = transpile(source) + assert "pf_line_get_x2(_pf_lines_, ln)" in cpp + assert "pf_line_get_x1(_pf_lines_, ln)" in cpp + assert "ln.get_x2()" not in cpp + assert "ln.get_x1()" not in cpp + compile_cpp( + cpp, + label=f"udt-exact-drawing-method-parameter-{label}", + ) def test_stateful_method_edge_uses_exact_global_udt_owner() -> None: From 858562d37b0a6bbf186db432306e390046dc7879 Mon Sep 17 00:00:00 2001 From: luisleo526 Date: Tue, 21 Jul 2026 00:30:22 +0800 Subject: [PATCH 8/8] fix: keep collection parameters out of UDT aliases --- pineforge_codegen/codegen/emit_top.py | 12 ++++++++++-- tests/test_udt_lvalue_alias.py | 23 +++++++++++++++++++++++ tests/test_udt_na_lifecycle.py | 22 ++++++++++++++++++++++ 3 files changed, 55 insertions(+), 2 deletions(-) diff --git a/pineforge_codegen/codegen/emit_top.py b/pineforge_codegen/codegen/emit_top.py index 18a74cf..adc1375 100644 --- a/pineforge_codegen/codegen/emit_top.py +++ b/pineforge_codegen/codegen/emit_top.py @@ -1451,7 +1451,11 @@ def _emit_func_def(self, fi: FuncInfo, lines: list[str], call_site_idx: int | No else self._udt_param_udt.get(param) ), ) - if drawing_name in DRAWING_TYPE_TO_CPP + if (drawing_name in DRAWING_TYPE_TO_CPP + and self._current_func_param_types.get( + param, "" + ).removesuffix("&").removesuffix("*") + == DRAWING_TYPE_TO_CPP[drawing_name]) } self._lexical_udt_types = { param: ( @@ -1461,7 +1465,11 @@ def _emit_func_def(self, fi: FuncInfo, lines: list[str], call_site_idx: int | No and spec.name in self._udt_defs else ( self._udt_param_udt.get(param) - if self._udt_param_udt.get(param) in self._udt_defs + if (self._udt_param_udt.get(param) in self._udt_defs + and self._current_func_param_types.get( + param, "" + ).removesuffix("&").removesuffix("*") + == self._udt_param_udt.get(param)) else None ) ) diff --git a/tests/test_udt_lvalue_alias.py b/tests/test_udt_lvalue_alias.py index 081e5c4..1166d73 100644 --- a/tests/test_udt_lvalue_alias.py +++ b/tests/test_udt_lvalue_alias.py @@ -134,6 +134,29 @@ def test_array_get_udt_local_mutation_emits_reference_alias(): assert "p.crossed = true;" in body +def test_array_param_get_udt_local_mutation_emits_reference_alias(): + src = PROLOGUE + """ +var array pivots = array.new() +upd(array items, int i) => + pivot p = array.get(items, i) + p.currentLevel := close + p.crossed := true + 0 +if array.size(pivots) == 0 + array.push(pivots, pivot.new(na, false)) +if close > open + upd(pivots, 0) +plot(close) +""" + cpp = transpile(src) + body = _func_body(cpp, "upd") + assert "pivot& p =" in body + assert "}((i)); }((items));" in body + assert "pivot p =" not in body + assert "p.currentLevel = " in body + assert "p.crossed = true;" in body + + _GLOBAL_ARRAY_PROLOGUE = PROLOGUE + """ var array pivots = array.new() if array.size(pivots) == 0 diff --git a/tests/test_udt_na_lifecycle.py b/tests/test_udt_na_lifecycle.py index 3358aa5..8dc55c0 100644 --- a/tests/test_udt_na_lifecycle.py +++ b/tests/test_udt_na_lifecycle.py @@ -395,6 +395,28 @@ def test_drawing_method_parameter_keeps_exact_lexical_type() -> None: ) +def test_drawing_array_parameter_keeps_collection_identity() -> None: + expressions = { + "method": "items.copy().get(0).get_x1()", + "functional": "array.get(array.copy(items), 0).get_x1()", + } + for label, expression in expressions.items(): + source = rf'''//@version=6 +strategy("exact drawing array parameter") +probe(array items) => + {expression} +var array segments = array.new_line() +if barstate.isfirst + segments.push(line.new(bar_index, close, bar_index + 1, close)) +float observed = probe(segments) +''' + cpp = transpile(source) + assert "double probe(std::vector& items)" in cpp + assert "return None();" not in cpp + assert "pf_line_get_x1(_pf_lines_," in cpp + compile_cpp(cpp, label=f"udt-exact-drawing-array-{label}") + + def test_stateful_method_edge_uses_exact_global_udt_owner() -> None: source = r'''//@version=6 strategy("stateful method exact edge")