-
Notifications
You must be signed in to change notification settings - Fork 4
Expand file tree
/
Copy pathmisc_visitors.py
More file actions
225 lines (206 loc) · 10.4 KB
/
Copy pathmisc_visitors.py
File metadata and controls
225 lines (206 loc) · 10.4 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
import ast
from typing import Any
from .interpreter_core import ASTInterpreter
async def visit_Global(self: ASTInterpreter, node: ast.Global, wrap_exceptions: bool = True) -> None:
self.env_stack[-1].setdefault("__global_names__", set()).update(node.names)
async def visit_Nonlocal(self: ASTInterpreter, node: ast.Nonlocal, wrap_exceptions: bool = True) -> None:
self.env_stack[-1].setdefault("__nonlocal_names__", set()).update(node.names)
async def visit_Delete(self: ASTInterpreter, node: ast.Delete, wrap_exceptions: bool = True):
for target in node.targets:
if isinstance(target, ast.Name):
del self.env_stack[-1][target.id]
elif isinstance(target, ast.Subscript):
obj = await self.visit(target.value, wrap_exceptions=wrap_exceptions)
key = await self.visit(target.slice, wrap_exceptions=wrap_exceptions)
del obj[key]
else:
raise Exception(f"Unsupported del target: {type(target).__name__}")
async def visit_Assert(self: ASTInterpreter, node: ast.Assert, wrap_exceptions: bool = True) -> None:
test = await self.visit(node.test, wrap_exceptions=wrap_exceptions)
if not test:
msg = await self.visit(node.msg, wrap_exceptions=wrap_exceptions) if node.msg else "Assertion failed"
raise AssertionError(msg)
async def visit_Yield(self: ASTInterpreter, node: ast.Yield, wrap_exceptions: bool = True) -> Any:
# Check if we're in an async generator context
if hasattr(self, 'generator_context') and self.generator_context.get('active', False):
state = self.generator_context.get('state')
if state and state.collecting:
# Collect the yield value
if node.value:
value = await self.visit(node.value, wrap_exceptions=wrap_exceptions)
else:
value = None
state.yields.append(value)
return None
else:
# Get the current generator from the stack (most recent one)
generator_stack = self.generator_context.get('generator_stack', [])
current_generator_id = generator_stack[-1] if generator_stack else None
# Create a unique key for this specific yield based on generator ID and line number
yield_line = getattr(node, 'lineno', 0)
# First check if we're resuming from a previous asend/athrow
generator_resuming = self.generator_context.get('resuming_from_asend', False)
last_yield_key = self.generator_context.get('last_yield_key')
if generator_resuming and last_yield_key:
# We're resuming from asend/athrow, use the last yield key
yield_key = last_yield_key
resuming_key = f'resuming_{yield_key}'
# Clear the resuming flag to prevent infinite loops
self.generator_context['resuming_from_asend'] = False
else:
# Use execution count to differentiate yields on the same line in loops
yield_count_key = f'yield_count_{current_generator_id}_{yield_line}'
yield_count = self.generator_context.get(yield_count_key, 0)
self.generator_context[yield_count_key] = yield_count + 1
yield_key = f'yield_{current_generator_id}_{yield_line}_{yield_count}'
resuming_key = f'resuming_{yield_key}'
if self.generator_context.get(resuming_key, False):
# Reset the resuming flag for this specific yield
self.generator_context[resuming_key] = False
# Check if there's an exception to throw
if self.generator_context.get('exception_thrown', False):
self.generator_context['exception_thrown'] = False
thrown_exc = self.generator_context.get('thrown_exception')
if thrown_exc:
raise thrown_exc
# Return the sent value (instead of yielding again)
sent_value = self.generator_context.get('last_sent_value')
return sent_value
else:
# This is a fresh yield - evaluate the value and yield it
from .exceptions import YieldException
if node.value:
value = await self.visit(node.value, wrap_exceptions=wrap_exceptions)
else:
value = None
# Store this yield as the last one executed
self.generator_context['last_yield_key'] = yield_key
raise YieldException(value)
# Fallback for non-generator contexts
if node.value:
return await self.visit(node.value, wrap_exceptions=wrap_exceptions)
return None
async def visit_YieldFrom(self: ASTInterpreter, node: ast.YieldFrom, wrap_exceptions: bool = True) -> Any:
iterable = await self.visit(node.value, wrap_exceptions=wrap_exceptions)
if 'yield_queue' in self.generator_context and self.generator_context.get('active', False):
if hasattr(iterable, '__aiter__'):
async for val in iterable:
await self.generator_context['yield_queue'].put(val)
sent_value = await self.generator_context['sent_queue'].get()
if isinstance(sent_value, BaseException):
raise sent_value
else:
for val in iterable:
await self.generator_context['yield_queue'].put(val)
sent_value = await self.generator_context['sent_queue'].get()
if isinstance(sent_value, BaseException):
raise sent_value
return None
if hasattr(iterable, '__aiter__'):
async def async_gen():
async for value in iterable:
yield value
return async_gen()
else:
def sync_gen():
for value in iterable:
yield value
return sync_gen()
async def visit_Match(self: ASTInterpreter, node: ast.Match, wrap_exceptions: bool = True) -> Any:
subject = await self.visit(node.subject, wrap_exceptions=wrap_exceptions)
result = None
base_frame = self.env_stack[-1].copy()
for case in node.cases:
self.env_stack.append(base_frame.copy())
try:
if await self._match_pattern(subject, case.pattern):
if case.guard and not await self.visit(case.guard, wrap_exceptions=True):
continue
for stmt in case.body[:-1]:
await self.visit(stmt, wrap_exceptions=wrap_exceptions)
result = await self.visit(case.body[-1], wrap_exceptions=wrap_exceptions)
break
finally:
self.env_stack.pop()
return result
async def _match_pattern(self: ASTInterpreter, subject: Any, pattern: ast.AST) -> bool:
if isinstance(pattern, ast.MatchValue):
value = await self.visit(pattern.value, wrap_exceptions=True)
return subject == value
elif isinstance(pattern, ast.MatchSingleton):
return subject is pattern.value
elif isinstance(pattern, ast.MatchSequence):
if not isinstance(subject, (list, tuple)):
return False
if len(pattern.patterns) != len(subject) and not any(isinstance(p, ast.MatchStar) for p in pattern.patterns):
return False
star_idx = None
for i, pat in enumerate(pattern.patterns):
if isinstance(pat, ast.MatchStar):
if star_idx is not None:
return False
star_idx = i
if star_idx is None:
for sub, pat in zip(subject, pattern.patterns):
if not await self._match_pattern(sub, pat):
return False
return True
else:
before = pattern.patterns[:star_idx]
after = pattern.patterns[star_idx + 1:]
if len(before) + len(after) > len(subject):
return False
for sub, pat in zip(subject[:len(before)], before):
if not await self._match_pattern(sub, pat):
return False
for sub, pat in zip(subject[len(subject) - len(after):], after):
if not await self._match_pattern(sub, pat):
return False
star_pat = pattern.patterns[star_idx]
star_count = len(subject) - len(before) - len(after)
star_sub = subject[len(before):len(before) + star_count]
if star_pat.name:
self.set_variable(star_pat.name, star_sub)
return True
elif isinstance(pattern, ast.MatchMapping):
if not isinstance(subject, dict):
return False
keys = [await self.visit(k, wrap_exceptions=True) for k in pattern.keys]
if len(keys) != len(subject) and pattern.rest is None:
return False
for k, p in zip(keys, pattern.patterns):
if k not in subject or not await self._match_pattern(subject[k], p):
return False
if pattern.rest:
remaining = {k: v for k, v in subject.items() if k not in keys}
self.set_variable(pattern.rest, remaining)
return True
elif isinstance(pattern, ast.MatchClass):
cls = await self.visit(pattern.cls, wrap_exceptions=True)
if not isinstance(subject, cls):
return False
attrs = [getattr(subject, attr) for attr in pattern.attribute_names]
if len(attrs) != len(pattern.patterns):
return False
for attr_val, pat in zip(attrs, pattern.patterns):
if not await self._match_pattern(attr_val, pat):
return False
return True
elif isinstance(pattern, ast.MatchStar):
if pattern.name:
self.set_variable(pattern.name, subject)
return True
elif isinstance(pattern, ast.MatchAs):
if pattern.pattern:
if not await self._match_pattern(subject, pattern.pattern):
return False
if pattern.name:
self.set_variable(pattern.name, subject)
return True
elif isinstance(pattern, ast.MatchOr):
for pat in pattern.patterns:
if await self._match_pattern(subject, pat):
return True
return False
else:
raise Exception(f"Unsupported match pattern: {pattern.__class__.__name__}")