@pineforge/codegen-pyodide 0.8.0 → 0.9.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -66,8 +66,8 @@ from __future__ import annotations
66
66
  from typing import Any
67
67
 
68
68
  from ..ast_nodes import (
69
- ASTNode, BoolLiteral, FuncCall, Identifier, MemberAccess,
70
- NumberLiteral, StringLiteral, TupleLiteral,
69
+ ASTNode, BinOp, BoolLiteral, ExprStmt, FuncCall, Identifier, MemberAccess,
70
+ NumberLiteral, StringLiteral, Ternary, TupleLiteral, UnaryOp, VarDecl,
71
71
  )
72
72
  from ..symbols import PineType
73
73
  from .. import signatures as sigs
@@ -75,7 +75,7 @@ from .. import tv_input_choices as tv_in
75
75
  from .contracts import FixnanCallSite, FuncInfo, SecurityCallInfo, TACallSite
76
76
  from .tables import (
77
77
  BAR_FIELDS, TA_CLASS_MAP, TA_MULTI_CTOR, TA_NO_CTOR, TA_PERIOD_ARG,
78
- TA_TUPLE_RETURNS,
78
+ TA_TUPLE_RETURNS, TA_TUPLE_ELEMENT_COUNTS, TA_COMPUTE_ARGS,
79
79
  )
80
80
 
81
81
 
@@ -205,8 +205,10 @@ class CallHandlers:
205
205
  returns_tuple=returns_tuple,
206
206
  node=node,
207
207
  is_static=is_static,
208
+ owner_func=(self._enclosing_func_names[-1] if self._enclosing_func_names else None),
208
209
  )
209
210
  self._ta_call_sites.append(site)
211
+ self._ta_member_names.add(site.member_name)
210
212
  return PineType.FLOAT
211
213
 
212
214
  # Determine constructor args
@@ -234,9 +236,14 @@ class CallHandlers:
234
236
  elif func_name in TA_PERIOD_ARG:
235
237
  ctor_indices = {TA_PERIOD_ARG[func_name]}
236
238
 
237
- for i, arg in enumerate(all_args):
238
- if i not in ctor_indices and arg is not None:
239
- compute_args.append(arg)
239
+ if func_name in TA_COMPUTE_ARGS:
240
+ for i in TA_COMPUTE_ARGS[func_name]:
241
+ if i < len(all_args) and all_args[i] is not None:
242
+ compute_args.append(all_args[i])
243
+ else:
244
+ for i, arg in enumerate(all_args):
245
+ if i not in ctor_indices and arg is not None:
246
+ compute_args.append(arg)
240
247
 
241
248
  is_static = self._global_scope and all(self._is_static_expression(arg) for arg in compute_args)
242
249
  site = TACallSite(
@@ -247,11 +254,32 @@ class CallHandlers:
247
254
  returns_tuple=returns_tuple,
248
255
  node=node,
249
256
  is_static=is_static,
257
+ owner_func=(self._enclosing_func_names[-1] if self._enclosing_func_names else None),
250
258
  )
251
259
  self._ta_call_sites.append(site)
260
+ self._ta_member_names.add(site.member_name)
252
261
 
253
262
  return PineType.FLOAT
254
263
 
264
+ def _security_symbol_is_heikinashi(self, node, _seen=None) -> bool:
265
+ """True when a request.security symbol is ``ticker.heikinashi(<chart
266
+ symbol>)`` — directly, or via a global alias (``haTicker =
267
+ ticker.heikinashi(syminfo.tickerid)``). Name-cycle-guarded. The
268
+ support_checker has already rejected the cross-symbol HA case, so any HA
269
+ reaching here is the chart's own symbol."""
270
+ if _seen is None:
271
+ _seen = set()
272
+ if (isinstance(node, FuncCall) and isinstance(node.callee, MemberAccess)
273
+ and isinstance(node.callee.object, Identifier)
274
+ and node.callee.object.name == "ticker"
275
+ and node.callee.member == "heikinashi"):
276
+ return True
277
+ if (isinstance(node, Identifier) and node.name in self._global_expr_map
278
+ and node.name not in _seen):
279
+ _seen.add(node.name)
280
+ return self._security_symbol_is_heikinashi(self._global_expr_map[node.name], _seen)
281
+ return False
282
+
255
283
  def _handle_request_call(self, func_name: str, node: FuncCall) -> PineType:
256
284
  """Handle request.* function calls."""
257
285
  if func_name == "security":
@@ -288,11 +316,41 @@ class CallHandlers:
288
316
 
289
317
  returns_tuple = isinstance(expr_node, TupleLiteral)
290
318
  tuple_size = len(expr_node.elements) if returns_tuple else 0
319
+ if not returns_tuple and isinstance(expr_node, FuncCall):
320
+ expr_func = None
321
+ expr_ns = None
322
+ if (isinstance(expr_node.callee, MemberAccess)
323
+ and isinstance(expr_node.callee.object, Identifier)):
324
+ expr_ns = expr_node.callee.object.name
325
+ expr_func = expr_node.callee.member
326
+ if expr_ns == "ta":
327
+ if expr_func == "vwap":
328
+ merged_v = list(expr_node.args)
329
+ for i, pname in enumerate(["source", "anchor", "stdev_mult"]):
330
+ if pname in expr_node.kwargs:
331
+ while len(merged_v) <= i:
332
+ merged_v.append(None)
333
+ if merged_v[i] is None:
334
+ merged_v[i] = expr_node.kwargs[pname]
335
+ if len(merged_v) >= 3:
336
+ expr_func = "vwap_bands"
337
+ if expr_func in TA_TUPLE_RETURNS:
338
+ returns_tuple = True
339
+ tuple_size = TA_TUPLE_ELEMENT_COUNTS.get(expr_func, 0)
291
340
 
292
341
  gaps_node = all_args[3] if len(all_args) > 3 else None
293
342
  lookahead_node = all_args[4] if len(all_args) > 4 else None
294
343
 
295
344
  mutable_globals = tuple(sorted(self._collect_security_mutable_globals(expr_node)))
345
+ # Heikin-Ashi same-symbol read: request.security(ticker.heikinashi(
346
+ # syminfo.tickerid), ...) (directly or via a global alias). The engine
347
+ # applies the HA candle transform inside the security eval.
348
+ symbol_node = all_args[0] if all_args else None
349
+ heikinashi = self._security_symbol_is_heikinashi(symbol_node)
350
+ # Capture the user function (if any) whose body contains this call,
351
+ # so the codegen can resolve a parameter ``tf`` via the call sites.
352
+ scope_name = self._symbols.current_scope.name
353
+ containing_func = scope_name[5:] if scope_name.startswith("func_") else ""
296
354
  self._security_calls.append(SecurityCallInfo(
297
355
  sec_id=sec_id,
298
356
  timeframe=tf_node,
@@ -302,8 +360,10 @@ class CallHandlers:
302
360
  gaps=gaps_node,
303
361
  lookahead=lookahead_node,
304
362
  ta_range=security_ta_range,
363
+ heikinashi=heikinashi,
305
364
  depends_on_mutable_globals=bool(mutable_globals),
306
365
  mutable_globals=mutable_globals,
366
+ containing_func=containing_func,
307
367
  ))
308
368
 
309
369
  return PineType.FLOAT
@@ -753,11 +813,18 @@ class CallHandlers:
753
813
  arg_type = self._visit(arg)
754
814
 
755
815
  self._fixnan_counter += 1
816
+ owner = self._enclosing_func_names[-1] if self._enclosing_func_names else None
756
817
  site = FixnanCallSite(
757
818
  member_name=f"_prev_fixnan_{self._fixnan_counter}",
758
819
  pine_type=arg_type,
820
+ node=node,
821
+ owner_func=owner,
759
822
  )
823
+ idx = len(self._fixnan_sites)
760
824
  self._fixnan_sites.append(site)
825
+ self._fixnan_member_names.add(site.member_name)
826
+ if owner is not None:
827
+ self._func_fixnan_indices.setdefault(owner, []).append(idx)
761
828
 
762
829
  return arg_type
763
830
 
@@ -765,6 +832,238 @@ class CallHandlers:
765
832
  # User-defined function calls
766
833
  # ------------------------------------------------------------------
767
834
 
835
+ def _func_local_length_defs(self, func_def) -> dict[str, str]:
836
+ """Collect a user function's local scalar length-vars to their RHS
837
+ expression string, e.g. ``qqeCalc`` with ``wp = sf * 2 - 1`` returns
838
+ ``{"wp": "sf * 2 - 1"}``.
839
+
840
+ Only plain (non-``var``/``varip``) declarations whose RHS is a pure
841
+ arithmetic expression over identifiers/numbers (NumberLiteral, Identifier,
842
+ BinOp, UnaryOp, or a math.* FuncCall) qualify — these are the shapes that
843
+ can legitimately feed a TA constructor length. Series-valued locals (whose
844
+ RHS is a ta.* call, a subscript, a ternary, etc.) are skipped so we never
845
+ inline a price series into a ctor-length slot. Names reassigned with ``:=``
846
+ are also skipped (their value is not a stable compile-time length).
847
+ """
848
+ def _is_arith(n) -> bool:
849
+ if isinstance(n, (NumberLiteral, Identifier)):
850
+ return True
851
+ if isinstance(n, BinOp):
852
+ return _is_arith(n.left) and _is_arith(n.right)
853
+ if isinstance(n, UnaryOp):
854
+ return _is_arith(n.operand)
855
+ if isinstance(n, Ternary):
856
+ # ``cond ? a : b`` — expand when both branches are arith.
857
+ # The condition may be a comparison/logical over arith leaves;
858
+ # rely on the codegen stability gate to reject series deps.
859
+ return (_is_arith(n.true_val) and _is_arith(n.false_val)
860
+ and _is_arith(n.condition))
861
+ if isinstance(n, MemberAccess):
862
+ # ``timeframe.*`` / ``syminfo.*`` / ``math.pi`` etc. — stable
863
+ # per-run scalars that may appear inside a function-local
864
+ # derived length. Let the codegen stability classifier decide.
865
+ return True
866
+ if isinstance(n, FuncCall):
867
+ callee = n.callee
868
+ # Allow math.* helpers (math.round/sqrt/cos/...) over arith args.
869
+ if (isinstance(callee, MemberAccess)
870
+ and isinstance(callee.object, Identifier)
871
+ and callee.object.name == "math"):
872
+ return all(_is_arith(a) for a in n.args)
873
+ # Pine type-cast builtins ``int(x)`` / ``float(x)`` / ``bool(x)``
874
+ # / ``string(x)`` — transparent over arith args. Common in
875
+ # derived TA lengths (``int(math.round(2 / a))``).
876
+ if (isinstance(callee, Identifier)
877
+ and callee.name in ("int", "float", "bool", "string")):
878
+ return all(_is_arith(a) for a in n.args)
879
+ return False
880
+
881
+ reassigned: set[str] = set()
882
+ def _scan_reassign(stmts):
883
+ from ..ast_nodes import Assignment
884
+ for s in stmts or []:
885
+ if isinstance(s, Assignment) and isinstance(s.target, Identifier):
886
+ reassigned.add(s.target.name)
887
+ for attr in ("body", "else_body"):
888
+ sub = getattr(s, attr, None)
889
+ if isinstance(sub, list):
890
+ _scan_reassign(sub)
891
+ _scan_reassign(func_def.body)
892
+
893
+ defs: dict[str, str] = {}
894
+ for stmt in func_def.body or []:
895
+ if (isinstance(stmt, VarDecl)
896
+ and not stmt.is_var and not stmt.is_varip
897
+ and stmt.name not in reassigned
898
+ and stmt.value is not None
899
+ and _is_arith(stmt.value)):
900
+ defs[stmt.name] = self._expr_to_str(stmt.value)
901
+ return defs
902
+
903
+ def _materialize_user_func_call_site_state(
904
+ self, func_name: str, cs_idx: int, node: FuncCall,
905
+ *, reuse_existing_owner: str | None = None) -> None:
906
+ """Materialize TA/fixnan state for one UDF call-site variant.
907
+
908
+ Ordinary call sites are handled while walking the AST. A second class
909
+ is discovered only after that walk: a stateful helper reached through a
910
+ multi-call-site parent needs the parent's additional call-path indices
911
+ even though the helper has only one textual call. The late propagation
912
+ pass in ``Analyzer._propagate_call_site_counts`` calls this same helper
913
+ for those inherited variants so the exported count never references a
914
+ TA/fixnan clone that was not actually declared.
915
+
916
+ ``reuse_existing_owner`` is used only by late propagation. A
917
+ range-widened parent may already have materialized the default
918
+ ``{member}_cs{idx}`` clone for the borrowed callee site; when that clone
919
+ belongs to the parent currently being propagated, it is the desired
920
+ call-path state and must be reused rather than duplicated under a
921
+ disambiguated-but-unused name.
922
+ """
923
+ func_def = self._func_defs[func_name]
924
+
925
+ param_arg_map: dict[str, str] = {}
926
+ for p_idx, param_name in enumerate(func_def.params):
927
+ if p_idx < len(node.args):
928
+ param_arg_map[param_name] = self._expr_to_str(node.args[p_idx])
929
+
930
+ if func_name in self._func_ta_ranges:
931
+ start, end = self._func_ta_ranges[func_name]
932
+
933
+ # Map local derived length variables back to expressions over the
934
+ # function's parameters before substituting call-site arguments.
935
+ local_defs = self._func_local_length_defs(func_def)
936
+
937
+ def _subst_params(arg: str, pmap: dict[str, str]) -> str:
938
+ import re
939
+ result = arg
940
+ for param, value in sorted(
941
+ pmap.items(), key=lambda item: len(item[0]), reverse=True):
942
+ result = re.sub(rf'\b{re.escape(param)}\b', value, result)
943
+ return result
944
+
945
+ def _expand_locals(arg: str) -> str:
946
+ import re
947
+ if not local_defs:
948
+ return arg
949
+ for _ in range(32):
950
+ def _rep(match: re.Match) -> str:
951
+ name = match.group(0)
952
+ if name in local_defs:
953
+ return "(" + local_defs[name] + ")"
954
+ return name
955
+ expanded = re.sub(r"[A-Za-z_][A-Za-z_0-9]*", _rep, arg)
956
+ if expanded == arg:
957
+ break
958
+ arg = expanded
959
+ return arg
960
+
961
+ import re as _re
962
+ enclosing_params: set[str] = set()
963
+ for names in self._enclosing_func_params:
964
+ enclosing_params |= names
965
+
966
+ if cs_idx == 0:
967
+ # cs0 owns the source-level sites. Preserve their parameterized
968
+ # ctor args for every later direct or inherited clone.
969
+ for i in range(start, end):
970
+ site = self._ta_call_sites[i]
971
+ if not hasattr(site, '_orig_ctor_args'):
972
+ site._orig_ctor_args = [
973
+ _expand_locals(arg) for arg in site.ctor_args
974
+ ]
975
+ site.ctor_args = [
976
+ _subst_params(arg, param_arg_map)
977
+ for arg in site._orig_ctor_args
978
+ ]
979
+ # If a ctor is now expressed in an enclosing UDF's params,
980
+ # retain that expression so the enclosing call can resolve
981
+ # it and widen the enclosing TA range as before.
982
+ if enclosing_params and self._nested_ta_touched is not None:
983
+ for arg in site.ctor_args:
984
+ tokens = set(_re.findall(
985
+ r"[A-Za-z_][A-Za-z_0-9]*", arg))
986
+ if tokens & enclosing_params:
987
+ site._orig_ctor_args = list(site.ctor_args)
988
+ self._nested_ta_touched.add(i)
989
+ break
990
+ else:
991
+ clone_name_map: dict[str, str] = {}
992
+ for i in range(start, end):
993
+ orig = self._ta_call_sites[i]
994
+ orig_args = getattr(orig, '_orig_ctor_args', orig.ctor_args)
995
+ resolved_ctor = [
996
+ _subst_params(arg, param_arg_map) for arg in orig_args
997
+ ]
998
+ clone_name = f"{orig.member_name}_cs{cs_idx}"
999
+ existing = next(
1000
+ (site for site in self._ta_call_sites
1001
+ if site.member_name == clone_name),
1002
+ None,
1003
+ )
1004
+ if (reuse_existing_owner is not None
1005
+ and existing is not None
1006
+ and existing.owner_func == reuse_existing_owner):
1007
+ # The active parent's widened range already made the
1008
+ # exact member this inherited callee variant needs.
1009
+ continue
1010
+ if clone_name in self._ta_member_names:
1011
+ base = clone_name
1012
+ suffix = 2
1013
+ while clone_name in self._ta_member_names:
1014
+ clone_name = f"{base}_u{suffix}"
1015
+ suffix += 1
1016
+ clone_name_map[orig.member_name] = clone_name
1017
+ cloned = TACallSite(
1018
+ member_name=clone_name,
1019
+ class_name=orig.class_name,
1020
+ ctor_args=resolved_ctor,
1021
+ compute_args=orig.compute_args[:],
1022
+ returns_tuple=orig.returns_tuple,
1023
+ node=orig.node,
1024
+ is_static=orig.is_static,
1025
+ owner_func=func_name,
1026
+ )
1027
+ self._ta_call_sites.append(cloned)
1028
+ self._ta_member_names.add(clone_name)
1029
+ if clone_name_map:
1030
+ self._func_cs_ta_clone_names[(func_name, cs_idx)] = clone_name_map
1031
+
1032
+ # fixnan is stateful for the same reason as a rolling TA reducer: each
1033
+ # emitted function variant needs its own previous-value member.
1034
+ fn_indices = self._func_fixnan_indices.get(func_name, [])
1035
+ if cs_idx > 0 and fn_indices:
1036
+ clone_map: dict[str, str] = {}
1037
+ for fi in fn_indices:
1038
+ orig = self._fixnan_sites[fi]
1039
+ clone_name = f"{orig.member_name}_cs{cs_idx}"
1040
+ existing = next(
1041
+ (site for site in self._fixnan_sites
1042
+ if site.member_name == clone_name),
1043
+ None,
1044
+ )
1045
+ if (reuse_existing_owner is not None
1046
+ and existing is not None
1047
+ and existing.owner_func == reuse_existing_owner):
1048
+ continue
1049
+ if clone_name in self._fixnan_member_names:
1050
+ base = clone_name
1051
+ suffix = 2
1052
+ while clone_name in self._fixnan_member_names:
1053
+ clone_name = f"{base}_u{suffix}"
1054
+ suffix += 1
1055
+ clone_map[orig.member_name] = clone_name
1056
+ cloned = FixnanCallSite(
1057
+ member_name=clone_name,
1058
+ pine_type=orig.pine_type,
1059
+ node=orig.node,
1060
+ owner_func=func_name,
1061
+ )
1062
+ self._fixnan_sites.append(cloned)
1063
+ self._fixnan_member_names.add(clone_name)
1064
+ if clone_map:
1065
+ self._func_cs_fixnan_clone_names[(func_name, cs_idx)] = clone_map
1066
+
768
1067
  def _handle_user_func_call(self, func_name: str, node: FuncCall) -> PineType:
769
1068
  """Handle calls to user-defined functions."""
770
1069
  func_def = self._func_defs[func_name]
@@ -783,14 +1082,28 @@ class CallHandlers:
783
1082
  # For now, use the cached return type from initial analysis
784
1083
  return_type = self._func_return_types.get(func_name, PineType.FLOAT)
785
1084
 
786
- # If the return type was UNKNOWN or VOID, infer from param types
1085
+ # If the return type was UNKNOWN or VOID, infer it ONLY when the body
1086
+ # is a single bare identifier that returns a parameter directly
1087
+ # (``f(s) => s``). Inferring from params for arbitrary bodies misfires
1088
+ # when a function merely HAS a string param but returns something else
1089
+ # (e.g. ``getLineStyle(s) => switch s ... => line.style_solid`` or a
1090
+ # body ending in ``label.new(...)``). Other cases rely on the cached
1091
+ # body type plus udt_return_type / tuple inference.
787
1092
  if return_type in (PineType.UNKNOWN, PineType.VOID):
788
- if any(t == PineType.STRING for t in param_types):
789
- return_type = PineType.STRING
790
- elif any(t == PineType.FLOAT for t in param_types):
791
- return_type = PineType.FLOAT
792
- elif any(t == PineType.INT for t in param_types):
793
- return_type = PineType.INT
1093
+ if (func_def.is_single_expr and func_def.body
1094
+ and isinstance(func_def.body[0], ExprStmt)
1095
+ and isinstance(func_def.body[0].expr, Identifier)):
1096
+ ret_name = func_def.body[0].expr.name
1097
+ for idx, pname in enumerate(func_def.params):
1098
+ if pname == ret_name and idx < len(param_types):
1099
+ pt = param_types[idx]
1100
+ if pt == PineType.STRING:
1101
+ return_type = PineType.STRING
1102
+ elif pt == PineType.INT:
1103
+ return_type = PineType.INT
1104
+ elif pt == PineType.FLOAT:
1105
+ return_type = PineType.FLOAT
1106
+ break
794
1107
 
795
1108
  # If this function has series params, ensure bar-field arguments
796
1109
  # passed at the call site are registered as series_bar_fields so that
@@ -802,66 +1115,29 @@ class CallHandlers:
802
1115
  arg = node.args[p_idx]
803
1116
  if isinstance(arg, Identifier) and arg.name in BAR_FIELDS:
804
1117
  self._series_bar_fields.add(arg.name)
805
-
806
- # Per-call-site cloning: if this function has TA calls or series vars,
807
- # track call sites so codegen can create per-call-site variants.
808
- # This prevents shared state corruption when the function is called
809
- # multiple times per bar.
1118
+ elif isinstance(arg, Identifier):
1119
+ sym = self._symbols.resolve(arg.name)
1120
+ spec = getattr(sym, "type_spec", None) if sym is not None else None
1121
+ if spec is not None and spec.kind in ("array", "map", "matrix"):
1122
+ continue
1123
+ if sym is not None:
1124
+ sym.is_series = True
1125
+ if sym.scope and sym.scope.startswith("func_"):
1126
+ caller_name = sym.scope[5:]
1127
+ self._func_series_vars.setdefault(caller_name, set()).add(arg.name)
1128
+ else:
1129
+ self._series_vars.add(arg.name)
1130
+
1131
+ # Per-call-site cloning: TA, series/var, and fixnan state all advance
1132
+ # across bars/calls and therefore require isolated UDF variants.
810
1133
  has_ta = func_name in self._func_ta_ranges
811
1134
  has_series = func_name in self._func_series_vars or func_name in self._func_var_members
812
- if has_ta or has_series:
1135
+ has_fixnan = func_name in self._func_fixnan_indices
1136
+ if has_ta or has_series or has_fixnan:
813
1137
  cs_idx = self._func_call_site_count.get(func_name, 0)
814
1138
  self._func_call_site_count[func_name] = cs_idx + 1
815
1139
  self._func_call_cs_map[id(node)] = (func_name, cs_idx)
816
-
817
- # Build parameter -> call-site argument string mapping
818
- param_arg_map: dict[str, str] = {}
819
- for p_idx, param_name in enumerate(func_def.params):
820
- if p_idx < len(node.args):
821
- param_arg_map[param_name] = self._expr_to_str(node.args[p_idx])
822
-
823
- # Clone TA call sites (only if function has TA ranges)
824
- if has_ta:
825
- start, end = self._func_ta_ranges[func_name]
826
-
827
- def _subst_params(arg: str, pmap: dict[str, str]) -> str:
828
- """Substitute parameter names in an expression string.
829
-
830
- Handles both exact matches ('len' -> 'len3') and parameter
831
- names within expressions ('len / 2' -> 'len3 / 2').
832
- """
833
- import re
834
- result = arg
835
- # Sort by length descending to avoid partial replacements
836
- for param, value in sorted(pmap.items(), key=lambda x: len(x[0]), reverse=True):
837
- result = re.sub(rf'\b{re.escape(param)}\b', value, result)
838
- return result
839
-
840
- if cs_idx == 0:
841
- # First call site: save original param-based ctor_args for future cloning,
842
- # then resolve to actual call-site values
843
- for i in range(start, end):
844
- site = self._ta_call_sites[i]
845
- if not hasattr(site, '_orig_ctor_args'):
846
- site._orig_ctor_args = site.ctor_args[:]
847
- site.ctor_args = [_subst_params(a, param_arg_map) for a in site._orig_ctor_args]
848
- else:
849
- # Subsequent call sites: clone using saved original param names,
850
- # substituted with this call site's arguments
851
- for i in range(start, end):
852
- orig = self._ta_call_sites[i]
853
- orig_args = getattr(orig, '_orig_ctor_args', orig.ctor_args)
854
- resolved_ctor = [_subst_params(a, param_arg_map) for a in orig_args]
855
- cloned = TACallSite(
856
- member_name=f"{orig.member_name}_cs{cs_idx}",
857
- class_name=orig.class_name,
858
- ctor_args=resolved_ctor,
859
- compute_args=orig.compute_args[:],
860
- returns_tuple=orig.returns_tuple,
861
- node=orig.node,
862
- is_static=orig.is_static,
863
- )
864
- self._ta_call_sites.append(cloned)
1140
+ self._materialize_user_func_call_site_state(func_name, cs_idx, node)
865
1141
 
866
1142
  # Create or update FuncInfo
867
1143
  is_tuple = self._func_returns_tuple.get(func_name, False)
@@ -869,6 +1145,18 @@ class CallHandlers:
869
1145
  # Forward UDT-return inference (set in _visit_FuncDef) so codegen can
870
1146
  # emit the struct return type. Probe: udt-method-probe-20.
871
1147
  udt_ret = self._func_udt_return_types.get(func_name)
1148
+ ret_spec = getattr(self, "_func_return_type_specs", {}).get(func_name)
1149
+ # Per-param TypeSpec: declared hints are authoritative; for untyped
1150
+ # params, infer from the call-site argument's type_spec (so an untyped
1151
+ # ``s`` used as a string, or a UDT passed by value, emits correctly).
1152
+ param_specs = list(
1153
+ getattr(self, "_func_param_type_specs", {}).get(func_name)
1154
+ or self._param_type_specs_from_def(func_def)
1155
+ )
1156
+ arg_specs = [self._type_spec_from_expr(arg) for arg in node.args]
1157
+ for i in range(len(param_specs)):
1158
+ if param_specs[i] is None and i < len(arg_specs):
1159
+ param_specs[i] = arg_specs[i]
872
1160
  existing = [fi for fi in self._func_infos if fi.name == func_name]
873
1161
  if not existing:
874
1162
  fi = FuncInfo(
@@ -879,6 +1167,8 @@ class CallHandlers:
879
1167
  returns_tuple=is_tuple,
880
1168
  tuple_element_count=tuple_count,
881
1169
  udt_return_type=udt_ret,
1170
+ param_type_specs=param_specs,
1171
+ return_type_spec=ret_spec,
882
1172
  )
883
1173
  self._func_infos.append(fi)
884
1174
  else:
@@ -889,7 +1179,17 @@ class CallHandlers:
889
1179
  for i, pt in enumerate(param_types):
890
1180
  if i < len(fi.param_types) and fi.param_types[i] == PineType.UNKNOWN:
891
1181
  fi.param_types[i] = pt
1182
+ # Merge per-param TypeSpecs: keep declared hints (authoritative),
1183
+ # fill untyped slots from this call site if still unknown.
1184
+ if not fi.param_type_specs:
1185
+ fi.param_type_specs = list(param_specs)
1186
+ else:
1187
+ for i in range(len(param_specs)):
1188
+ if i < len(fi.param_type_specs) and fi.param_type_specs[i] is None:
1189
+ fi.param_type_specs[i] = param_specs[i]
892
1190
  if fi.udt_return_type is None and udt_ret is not None:
893
1191
  fi.udt_return_type = udt_ret
1192
+ if fi.return_type_spec is None and ret_spec is not None:
1193
+ fi.return_type_spec = ret_spec
894
1194
 
895
1195
  return return_type