opencode-pine2pyne 0.1.0__py3-none-any.whl

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.
@@ -0,0 +1,328 @@
1
+ """
2
+ Type inference for Pine Script variables.
3
+
4
+ Determines whether variables should be Series[T] or Persistent[T] in PyneCore.
5
+ """
6
+ from typing import Set
7
+ from .ast_nodes import *
8
+ from .symbol_table import SymbolTable, Symbol, VariableKind
9
+ from .pine_builtins import get_type_name
10
+
11
+
12
+ class TypeInference:
13
+ """Analyzes AST to infer variable types (Series vs Persistent vs simple)."""
14
+
15
+ def __init__(self, symbol_table: SymbolTable):
16
+ self.symbol_table = symbol_table
17
+ self.ta_functions = self._init_ta_functions()
18
+
19
+ def _init_ta_functions(self) -> Set[str]:
20
+ """Initialize complete set of ta.* functions/variables (Pine v5/v6)."""
21
+ return {
22
+ # Moving averages
23
+ 'ta.sma', 'ta.ema', 'ta.wma', 'ta.vwma', 'ta.swma', 'ta.alma',
24
+ 'ta.rma', 'ta.hma', 'ta.linreg',
25
+ # Oscillators & indicators
26
+ 'ta.rsi', 'ta.macd', 'ta.stoch', 'ta.cci', 'ta.cmo', 'ta.mfi',
27
+ 'ta.mom', 'ta.roc', 'ta.tsi', 'ta.wpr',
28
+ # Volatility
29
+ 'ta.atr', 'ta.tr', 'ta.bb', 'ta.bbw', 'ta.kc', 'ta.kcw',
30
+ # Trend
31
+ 'ta.dmi', 'ta.supertrend', 'ta.sar', 'ta.cog',
32
+ # Volume-based (functions and built-in variables)
33
+ 'ta.obv', 'ta.nvi', 'ta.pvi', 'ta.pvt', 'ta.vwap',
34
+ 'ta.accdist', 'ta.iii', 'ta.wad', 'ta.wvad',
35
+ # Statistical
36
+ 'ta.dev', 'ta.stdev', 'ta.variance', 'ta.correlation',
37
+ 'ta.median', 'ta.mode', 'ta.range', 'ta.cum',
38
+ 'ta.percentile_linear_interpolation', 'ta.percentile_nearest_rank',
39
+ 'ta.percentrank',
40
+ # Lookback
41
+ 'ta.highest', 'ta.lowest', 'ta.highestbars', 'ta.lowestbars',
42
+ 'ta.valuewhen', 'ta.barssince', 'ta.change',
43
+ # Pivot
44
+ 'ta.pivothigh', 'ta.pivotlow',
45
+ # Boolean series
46
+ 'ta.rising', 'ta.falling', 'ta.cross', 'ta.crossover', 'ta.crossunder',
47
+ }
48
+
49
+ def infer_types(self, script: Script) -> None:
50
+ """Infer types for all variables in the script."""
51
+ # First pass: mark variables with var/varip as Persistent
52
+ self._mark_persistent_variables(script)
53
+
54
+ # Second pass: mark variables from ta.* functions as Series
55
+ self._mark_ta_series_variables(script)
56
+
57
+ # Third pass: mark indexed variables as Series
58
+ self._mark_indexed_variables(script)
59
+
60
+ # Fourth pass: mark global non-var variables as Series (PRIMARY RULE)
61
+ self._mark_global_series_variables(script)
62
+
63
+ # Fifth pass: propagate Series type from assignments
64
+ self._propagate_series_types(script)
65
+
66
+ def _mark_persistent_variables(self, script: Script) -> None:
67
+ """Mark all var/varip variables as Persistent."""
68
+ for decl in script.declarations:
69
+ if isinstance(decl, (VarDecl, VaripDecl)):
70
+ symbol = self.symbol_table.lookup(decl.name)
71
+ if symbol:
72
+ # These are already marked as VAR/VARIP in symbol table
73
+ symbol.is_series = False # Not Series, will be Persistent
74
+
75
+ def _mark_ta_series_variables(self, script: Script) -> None:
76
+ """Mark variables assigned from ta.* functions as Series."""
77
+ # Check global assignments
78
+ for stmt in script.body:
79
+ if isinstance(stmt, Assignment):
80
+ if self._is_ta_function_call(stmt.value):
81
+ if isinstance(stmt.target, str):
82
+ self.symbol_table.mark_from_ta_function(stmt.target)
83
+
84
+ def _mark_indexed_variables(self, script: Script) -> None:
85
+ """Mark variables that are indexed (e.g., close[1]) as Series."""
86
+ self.indexed_var_names: set[str] = set()
87
+ self._find_indexed_in_node(script)
88
+
89
+ def _find_indexed_in_node(self, node: Any) -> None:
90
+ """Recursively find indexed access patterns."""
91
+ if isinstance(node, IndexAccess):
92
+ if isinstance(node.object, Identifier):
93
+ self.indexed_var_names.add(node.object.name)
94
+ self.symbol_table.mark_as_indexed(node.object.name)
95
+
96
+ # Recursively check all child nodes. dict must be traversed too:
97
+ # function-call keyword arguments are stored as a dict, so a history
98
+ # access like `x[1]` inside a kwarg (e.g. plot(color=(x >= x[1] ? ...)))
99
+ # would otherwise be invisible and x would not be promoted to Series.
100
+ if isinstance(node, dict):
101
+ for value in node.values():
102
+ self._find_indexed_in_node(value)
103
+ elif isinstance(node, (list, tuple)):
104
+ for item in node:
105
+ self._find_indexed_in_node(item)
106
+ elif hasattr(node, '__dict__'):
107
+ for attr_value in node.__dict__.values():
108
+ if isinstance(attr_value, (ASTNode, list, tuple, dict)):
109
+ self._find_indexed_in_node(attr_value)
110
+
111
+ def _mark_global_series_variables(self, script: Script) -> None:
112
+ """Mark all global non-var, non-input variables as Series (PRIMARY RULE)."""
113
+ for symbol in self.symbol_table.get_all_globals():
114
+ # Skip var/varip (they're Persistent)
115
+ if symbol.kind in (VariableKind.VAR, VariableKind.VARIP):
116
+ continue
117
+
118
+ # Skip inputs (they become function parameters)
119
+ if symbol.kind == VariableKind.INPUT:
120
+ continue
121
+
122
+ # Skip functions
123
+ if symbol.kind == VariableKind.FUNCTION:
124
+ continue
125
+
126
+ # Everything else is Series by default
127
+ if symbol.is_global:
128
+ symbol.is_series = True
129
+
130
+ def _propagate_series_types(self, script: Script) -> None:
131
+ """Propagate Series type from one variable to another in assignments."""
132
+ for stmt in script.body:
133
+ if isinstance(stmt, Assignment):
134
+ if isinstance(stmt.target, str):
135
+ # Check if value is a Series variable
136
+ if isinstance(stmt.value, Identifier):
137
+ source_symbol = self.symbol_table.lookup(stmt.value.name)
138
+ if source_symbol and source_symbol.is_series:
139
+ self.symbol_table.mark_as_series(stmt.target)
140
+
141
+ def _is_ta_function_call(self, expr: Expression) -> bool:
142
+ """Check if expression is a ta.* function call."""
143
+ if isinstance(expr, FunctionCall):
144
+ if isinstance(expr.func, str):
145
+ return expr.func in self.ta_functions
146
+ elif isinstance(expr.func, MemberAccess):
147
+ func_name = f"{expr.func.object}.{expr.func.member}"
148
+ return func_name in self.ta_functions
149
+ return False
150
+
151
+ def infer_type_hint(self, value: Expression) -> str:
152
+ """Infer the Python type hint for an expression."""
153
+ if isinstance(value, Literal):
154
+ type_map = {
155
+ 'int': 'int',
156
+ 'float': 'float',
157
+ 'string': 'str',
158
+ 'bool': 'bool',
159
+ 'color': 'Color',
160
+ }
161
+ return type_map.get(value.literal_type, 'Any')
162
+
163
+ if isinstance(value, NaLiteral):
164
+ return 'Any' # NA can be any type
165
+
166
+ if isinstance(value, UnaryOp):
167
+ if value.op == 'not':
168
+ return 'bool'
169
+ # Unary +/- preserves operand type
170
+ return self.infer_type_hint(value.operand)
171
+
172
+ if isinstance(value, BinaryOp):
173
+ # Boolean operations always return bool
174
+ if value.op in ('and', 'or', '==', '!=', '>', '<', '>=', '<='):
175
+ return 'bool'
176
+
177
+ # Infer from operands
178
+ left_type = self.infer_type_hint(value.left)
179
+ right_type = self.infer_type_hint(value.right)
180
+
181
+ # If either is float, result is float
182
+ if left_type == 'float' or right_type == 'float':
183
+ return 'float'
184
+
185
+ # If both are int, result is int (except for division)
186
+ if left_type == 'int' and right_type == 'int':
187
+ if value.op == '/':
188
+ return 'float'
189
+ return 'int'
190
+
191
+ return 'Any'
192
+
193
+ if isinstance(value, TernaryOp):
194
+ # Type is union of true and false branches
195
+ true_type = self.infer_type_hint(value.true_expr)
196
+ false_type = self.infer_type_hint(value.false_expr)
197
+ if true_type == false_type:
198
+ return true_type
199
+ # If mixed int/float, return float
200
+ if {true_type, false_type} == {'int', 'float'}:
201
+ return 'float'
202
+ return 'Any'
203
+
204
+ if isinstance(value, FunctionCall):
205
+ # Check if it's a ta.* function
206
+ func_name = None
207
+ if isinstance(value.func, str):
208
+ func_name = value.func
209
+ elif isinstance(value.func, MemberAccess):
210
+ func_name = f"{value.func.object}.{value.func.member}"
211
+
212
+ if func_name and func_name in self.ta_functions:
213
+ return 'float' # Most ta.* functions return float series
214
+
215
+ # input.* functions
216
+ if func_name and func_name.startswith('input.'):
217
+ input_type_map = {
218
+ 'input.int': 'int',
219
+ 'input.float': 'float',
220
+ 'input.bool': 'bool',
221
+ 'input.string': 'str',
222
+ 'input.color': 'Color',
223
+ 'input.source': 'float', # Sources are float series
224
+ }
225
+ return input_type_map.get(func_name, 'Any')
226
+
227
+ # PyneCore collection constructors
228
+ if func_name:
229
+ # matrix.new<float>() -> Matrix[float]
230
+ if func_name.startswith('matrix.new'):
231
+ # Check for generic type in func_name: matrix.new<float>
232
+ if '<' in func_name and '>' in func_name:
233
+ generic = func_name[func_name.index('<')+1:func_name.index('>')]
234
+ return f'Matrix[{generic}]'
235
+ return 'Matrix[float]' # Default to float
236
+
237
+ # matrix.* methods that return a NEW Matrix (so a variable assigned
238
+ # from them, e.g. `t = m.transpose()`, is typed as a Matrix and its
239
+ # own method calls resolve to the matrix module). Two spellings
240
+ # reach here: the module form `matrix.transpose(m)` (after method
241
+ # resolution), and the pre-resolution method form `m.transpose()`
242
+ # which the parser stores as FunctionCall(func='m.transpose') — for
243
+ # the latter the object must be a known Matrix so the same method
244
+ # names on arrays/maps are not mistyped.
245
+ _MATRIX_RETURNING = ('transpose', 'copy', 'submatrix', 'mult', 'kron',
246
+ 'pow', 'inv', 'pinv', 'reshape', 'diff')
247
+ if func_name.startswith('matrix.') and func_name.split('.', 1)[1] in _MATRIX_RETURNING:
248
+ return 'Matrix[float]'
249
+ if '.' in func_name:
250
+ _obj_name, _, _meth = func_name.rpartition('.')
251
+ if _meth in _MATRIX_RETURNING and '.' not in _obj_name:
252
+ _obj_symbol = self.symbol_table.lookup(_obj_name)
253
+ if _obj_symbol and _obj_symbol.type_hint and 'Matrix' in _obj_symbol.type_hint:
254
+ return 'Matrix[float]'
255
+
256
+ # array.new_*() -> list[type]
257
+ if func_name.startswith('array.new_'):
258
+ array_type_map = {
259
+ 'array.new_int': 'list[int]',
260
+ 'array.new_float': 'list[float]',
261
+ 'array.new_bool': 'list[bool]',
262
+ 'array.new_string': 'list[str]',
263
+ 'array.new_color': 'list[Color]',
264
+ 'array.new_line': 'list[Line]',
265
+ 'array.new_label': 'list[Label]',
266
+ 'array.new_box': 'list[Box]',
267
+ 'array.new_table': 'list[Table]',
268
+ }
269
+ return array_type_map.get(func_name, 'list[Any]')
270
+
271
+ # map.new<K,V>() -> dict[K, V]
272
+ if func_name.startswith('map.new'):
273
+ if '<' in func_name and '>' in func_name:
274
+ generic = func_name[func_name.index('<')+1:func_name.index('>')]
275
+ parts = [p.strip() for p in generic.split(',')]
276
+ if len(parts) == 2:
277
+ k = get_type_name(parts[0])
278
+ v = get_type_name(parts[1])
279
+ return f'dict[{k}, {v}]'
280
+ return 'dict[str, float]' # Default generic
281
+
282
+ # map.get -> float (most common map value type in Pine strategies)
283
+ if func_name == 'map.get':
284
+ return 'float'
285
+
286
+ # map.copy(m) / map.keys(m) / map.values(m) -> same type
287
+ if func_name == 'map.copy':
288
+ return 'dict[str, float]' # Copy returns a map
289
+ if func_name in ('map.keys', 'array.copy'):
290
+ return 'list[Any]'
291
+ if func_name == 'map.values':
292
+ return 'list[Any]'
293
+
294
+ # label.new(), line.new(), box.new(), table.new()
295
+ if func_name in ('label.new', 'line.new', 'box.new', 'table.new'):
296
+ type_name = func_name.split('.')[0].capitalize()
297
+ return type_name
298
+
299
+ # Check if it's a UDT constructor call (e.g., TradeStats())
300
+ if isinstance(value.func, str):
301
+ # Check if the function name is a known UDT type
302
+ symbol = self.symbol_table.lookup(value.func)
303
+ if symbol and symbol.kind == VariableKind.TYPE:
304
+ return value.func # Return the UDT type name
305
+ # Fallback: if function name starts with uppercase, assume it's a UDT
306
+ elif value.func and value.func[0].isupper():
307
+ return value.func
308
+
309
+ # Handle MethodCall (e.g., TradeStats.new() before transformation)
310
+ if isinstance(value, MethodCall):
311
+ # Check if it's UDT.new() pattern
312
+ if value.method == 'new' and isinstance(value.object, Identifier):
313
+ # Check if the object is a UDT type
314
+ if value.object.name[0].isupper():
315
+ return value.object.name # Return just the UDT name without .new suffix
316
+
317
+ # Handle FunctionCall for UDT constructors (after transformation: TradeStats())
318
+ if isinstance(value, FunctionCall):
319
+ if isinstance(value.func, str) and value.func and value.func[0].isupper():
320
+ # Already a constructor call - return the type name directly
321
+ return value.func
322
+
323
+ if isinstance(value, Identifier):
324
+ symbol = self.symbol_table.lookup(value.name)
325
+ if symbol and symbol.type_hint:
326
+ return symbol.type_hint
327
+
328
+ return 'Any'