redis-lua-py 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,787 @@
1
+ """Turn a Python function into Lua.
2
+
3
+ The supported subset is deliberately small. Anything outside it raises
4
+ :class:`UnsupportedSyntax` pointing at the offending line, because a script
5
+ body that *looks* like Python but is never executed by Python is exactly the
6
+ place where a silent mistranslation would be most expensive.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import ast
12
+ import inspect
13
+ import textwrap
14
+ from collections.abc import Callable
15
+ from types import ModuleType
16
+ from typing import Any, NoReturn
17
+
18
+ from . import _lua as lua
19
+ from ._runtime import _Namespace
20
+ from ._script import CompiledScript
21
+ from .errors import CompileError, UnsupportedSyntax
22
+
23
+ # Members of the `redis` table that are not commands and keep their own name.
24
+ _REDIS_DIRECT = frozenset(
25
+ {
26
+ "call",
27
+ "pcall",
28
+ "error_reply",
29
+ "status_reply",
30
+ "sha1hex",
31
+ "log",
32
+ "replicate_commands",
33
+ "setresp",
34
+ "breakpoint",
35
+ "debug",
36
+ }
37
+ )
38
+
39
+ # Receivers are normally resolved by value. These spellings are the fallback
40
+ # for a name that is not bound in the defining module at all.
41
+ _RECEIVER_FALLBACK = {"redis": "redis", "call": "redis", "cjson": "cjson"}
42
+
43
+ _UNBOUND = object()
44
+
45
+
46
+ def _is_redis_py(value: object) -> bool:
47
+ """True for the redis-py package itself or one of its client objects."""
48
+ if isinstance(value, ModuleType):
49
+ return (value.__name__ or "").split(".")[0] == "redis"
50
+ return (type(value).__module__ or "").split(".")[0] == "redis"
51
+
52
+
53
+ _COMPARE_OPS: dict[type[ast.cmpop], str] = {
54
+ ast.Eq: "==",
55
+ ast.NotEq: "~=",
56
+ ast.Lt: "<",
57
+ ast.LtE: "<=",
58
+ ast.Gt: ">",
59
+ ast.GtE: ">=",
60
+ }
61
+
62
+ _BIN_OPS: dict[type[ast.operator], str] = {
63
+ ast.Add: "+",
64
+ ast.Sub: "-",
65
+ ast.Mult: "*",
66
+ ast.Div: "/",
67
+ ast.Mod: "%",
68
+ ast.Pow: "^",
69
+ }
70
+
71
+ # Python builtins with an exact Lua counterpart.
72
+ _BUILTIN_FUNCS: dict[str, str] = {
73
+ "int": "tonumber",
74
+ "float": "tonumber",
75
+ "str": "tostring",
76
+ "tonumber": "tonumber",
77
+ "tostring": "tostring",
78
+ "abs": "math.abs",
79
+ "min": "math.min",
80
+ "max": "math.max",
81
+ }
82
+
83
+ _TRUTHY_HELPER = """\
84
+ -- Python truthiness: 0, '', empty tables and nil are all false.
85
+ local function __truthy(v)
86
+ if v == nil or v == false then return false end
87
+ if v == 0 or v == '' then return false end
88
+ if type(v) == 'table' and next(v) == nil then return false end
89
+ return true
90
+ end
91
+ """
92
+
93
+ _ISNIL_HELPER = """\
94
+ -- A Redis command with nothing to return hands Lua false, not nil, so an
95
+ -- `is None` test has to accept both. Taking v as an argument also means the
96
+ -- operand is evaluated once, not once per comparison.
97
+ local function __isnil(v)
98
+ return v == nil or v == false
99
+ end
100
+ """
101
+
102
+
103
+ class _Compiler:
104
+ def __init__(
105
+ self,
106
+ func: ast.FunctionDef,
107
+ *,
108
+ filename: str,
109
+ first_lineno: int,
110
+ lines: list[str],
111
+ globalns: dict[str, Any],
112
+ ) -> None:
113
+ self.func = func
114
+ self.globalns = globalns
115
+ self.filename = filename
116
+ self.first_lineno = first_lineno
117
+ self.lines = lines
118
+ self.known: set[str] = set()
119
+ self.keys: list[str] = []
120
+ self.args: list[str] = []
121
+ self.params: list[str] = []
122
+ self.numeric_args: set[str] = set()
123
+ self.hoisted: list[str] = []
124
+ self.simple_assigns: set[int] = set()
125
+ self.needs_truthy = False
126
+ self.needs_isnil = False
127
+ self._temp = 0
128
+
129
+ # ----------------------------------------------------------------- errors
130
+
131
+ def fail(self, node: ast.AST, message: str, hint: str | None = None) -> NoReturn:
132
+ lineno = getattr(node, "lineno", 1)
133
+ source_line = self.lines[lineno - 1] if 0 < lineno <= len(self.lines) else None
134
+ raise UnsupportedSyntax(
135
+ message,
136
+ filename=self.filename,
137
+ lineno=self.first_lineno + lineno - 1,
138
+ col=getattr(node, "col_offset", 0),
139
+ source_line=source_line,
140
+ hint=hint,
141
+ )
142
+
143
+ # ---------------------------------------------------------------- scoping
144
+
145
+ def collect_assigned(self) -> None:
146
+ """Decide which names need hoisting to function scope.
147
+
148
+ Python scopes assignments to the whole function; Lua's ``local`` scopes
149
+ them to the enclosing block. A name assigned inside an ``if`` and read
150
+ after it must therefore be declared up front, or the read sees ``nil``.
151
+
152
+ A name whose *first* assignment sits at the top level of the body needs
153
+ no hoist: a ``local`` there is already in scope for everything that
154
+ follows, nested blocks included. Only names that first appear inside a
155
+ block get declared up front.
156
+ """
157
+ nodes: dict[str, list[ast.stmt]] = {}
158
+
159
+ def record(name: str, node: ast.stmt) -> None:
160
+ nodes.setdefault(name, []).append(node)
161
+
162
+ for node in ast.walk(self.func):
163
+ if isinstance(node, ast.Assign):
164
+ for target in node.targets:
165
+ if isinstance(target, ast.Name):
166
+ record(target.id, node)
167
+ elif isinstance(node, ast.AnnAssign | ast.AugAssign) and isinstance(
168
+ node.target, ast.Name
169
+ ):
170
+ record(node.target.id, node)
171
+ elif isinstance(node, ast.For) and isinstance(node.target, ast.Name):
172
+ # Loop targets get Lua's own loop scope; only record them so
173
+ # that reads of the name resolve.
174
+ self.known.add(node.target.id)
175
+
176
+ top_level = {id(stmt) for stmt in self.func.body}
177
+ for name, assignments in nodes.items():
178
+ self.known.add(name)
179
+ # ast.walk is breadth-first, so ask for source order explicitly.
180
+ first = min(assignments, key=lambda n: (n.lineno, n.col_offset))
181
+ if id(first) in top_level:
182
+ self.simple_assigns.add(id(first))
183
+ else:
184
+ self.hoisted.append(name)
185
+
186
+ # -------------------------------------------------------------- signature
187
+
188
+ def compile_signature(self) -> list[lua.Stat]:
189
+ sig = self.func.args
190
+ if sig.vararg or sig.kwarg:
191
+ self.fail(self.func, "*args and **kwargs are not supported in a script signature")
192
+ if sig.posonlyargs:
193
+ self.fail(self.func, "positional-only parameters are not supported")
194
+
195
+ prelude: list[lua.Stat] = []
196
+ for arg in [*sig.args, *sig.kwonlyargs]:
197
+ name = arg.arg
198
+ if not lua.is_identifier(name):
199
+ self.fail(arg, f"{name!r} is a reserved word in Lua")
200
+ self.params.append(name)
201
+ self.known.add(name)
202
+
203
+ if self._is_key(arg.annotation):
204
+ self.keys.append(name)
205
+ source: lua.Expr = lua.Index(lua.Name("KEYS"), lua.Num(len(self.keys)))
206
+ else:
207
+ self.args.append(name)
208
+ source = lua.Index(lua.Name("ARGV"), lua.Num(len(self.args)))
209
+ if self._is_numeric(arg.annotation):
210
+ # ARGV always arrives as strings; an int/float annotation
211
+ # is the author asking for the conversion.
212
+ self.numeric_args.add(name)
213
+ source = lua.Call(lua.Name("tonumber"), (source,))
214
+ prelude.append(lua.Local([name], [source]))
215
+
216
+ if sig.defaults or any(d is not None for d in sig.kw_defaults):
217
+ self.fail(
218
+ self.func,
219
+ "default values are not supported",
220
+ hint="Redis has no notion of an absent ARGV entry; pass the value explicitly.",
221
+ )
222
+ return prelude
223
+
224
+ @staticmethod
225
+ def _annotation_name(node: ast.expr | None) -> str | None:
226
+ if isinstance(node, ast.Name):
227
+ return node.id
228
+ if isinstance(node, ast.Attribute):
229
+ return node.attr
230
+ if isinstance(node, ast.Subscript):
231
+ return _Compiler._annotation_name(node.value)
232
+ if isinstance(node, ast.Constant) and isinstance(node.value, str):
233
+ return node.value.rsplit(".", 1)[-1].split("[", 1)[0]
234
+ return None
235
+
236
+ def _is_key(self, node: ast.expr | None) -> bool:
237
+ return self._annotation_name(node) == "Key"
238
+
239
+ def _is_numeric(self, node: ast.expr | None) -> bool:
240
+ return self._annotation_name(node) in {"int", "float"}
241
+
242
+ # ------------------------------------------------------------ expressions
243
+
244
+ def expr(self, node: ast.expr) -> lua.Expr:
245
+ match node:
246
+ case ast.Constant(value=None):
247
+ return lua.Nil()
248
+ case ast.Constant(value=bool() as v):
249
+ return lua.Bool(v)
250
+ case ast.Constant(value=int() | float() as v):
251
+ return lua.Num(v)
252
+ case ast.Constant(value=str() as v):
253
+ return lua.Str(v)
254
+ case ast.Constant(value=bytes() as v):
255
+ return lua.Str(v.decode("utf-8", "surrogateescape"))
256
+ case ast.Name(id=name):
257
+ if name not in self.known:
258
+ self.fail(
259
+ node,
260
+ f"undefined name {name!r}",
261
+ hint="A script can only use its parameters and names it assigns; "
262
+ "values from the enclosing Python scope are not available.",
263
+ )
264
+ return lua.Name(name)
265
+ case ast.BinOp():
266
+ return self.binop(node)
267
+ case ast.UnaryOp(op=ast.USub(), operand=operand):
268
+ return lua.UnOp("-", self.expr(operand))
269
+ case ast.UnaryOp(op=ast.UAdd(), operand=operand):
270
+ return self.expr(operand)
271
+ case ast.UnaryOp(op=ast.Not()):
272
+ return self.condition(node)
273
+ case ast.Compare():
274
+ return self.compare(node)
275
+ case ast.BoolOp():
276
+ self.fail(
277
+ node,
278
+ "'and'/'or' are only supported in an if or while condition",
279
+ hint="In Python these return an operand, which does not survive the "
280
+ "difference in truthiness. Use an if statement instead.",
281
+ )
282
+ case ast.IfExp():
283
+ self.fail(
284
+ node,
285
+ "conditional expressions (a if c else b) are not supported",
286
+ hint="Lua's `c and a or b` is wrong when a is false or nil. "
287
+ "Use an if statement.",
288
+ )
289
+ case ast.Call():
290
+ return self.call(node)
291
+ case ast.Subscript():
292
+ return self.subscript(node)
293
+ case ast.List(elts=elts) | ast.Tuple(elts=elts):
294
+ return lua.Table(array=tuple(self.expr(e) for e in elts))
295
+ case ast.Dict(keys=keys, values=values):
296
+ pairs = []
297
+ for key_node, value_node in zip(keys, values, strict=True):
298
+ if key_node is None:
299
+ self.fail(node, "dict unpacking (**) is not supported")
300
+ pairs.append((self.expr(key_node), self.expr(value_node)))
301
+ return lua.Table(hash=tuple(pairs))
302
+ case ast.JoinedStr(values=values):
303
+ return self.fstring(node, values)
304
+ case ast.Attribute(attr=attr):
305
+ self.fail(node, f"attribute access .{attr} is not supported here")
306
+ case _:
307
+ self.fail(node, f"{type(node).__name__} expressions are not supported")
308
+
309
+ def binop(self, node: ast.BinOp) -> lua.Expr:
310
+ left, right = self.expr(node.left), self.expr(node.right)
311
+ if isinstance(node.op, ast.FloorDiv):
312
+ return lua.Call(lua.Name("math.floor"), (lua.BinOp("/", left, right),))
313
+ op = _BIN_OPS.get(type(node.op))
314
+ if op is None:
315
+ self.fail(node, f"the {type(node.op).__name__} operator is not supported")
316
+ if op == "+" and (isinstance(node.left, ast.Constant) and isinstance(node.left.value, str)):
317
+ self.fail(
318
+ node,
319
+ "'+' is arithmetic in Lua and will not concatenate strings",
320
+ hint="Use an f-string, which compiles to Lua's .. operator.",
321
+ )
322
+ return lua.BinOp(op, left, right)
323
+
324
+ def compare(self, node: ast.Compare) -> lua.Expr:
325
+ if len(node.ops) != 1:
326
+ self.fail(
327
+ node,
328
+ "chained comparisons are not supported",
329
+ hint="Split 'a < b < c' into 'a < b and b < c'.",
330
+ )
331
+ op_node, right_node = node.ops[0], node.comparators[0]
332
+ left = self.expr(node.left)
333
+
334
+ if isinstance(op_node, ast.Is | ast.IsNot):
335
+ if not (isinstance(right_node, ast.Constant) and right_node.value is None):
336
+ self.fail(node, "'is' is only supported against None")
337
+ # Not `== nil`: Redis reports a missing value to Lua as false.
338
+ self.needs_isnil = True
339
+ check: lua.Expr = lua.Call(lua.Name("__isnil"), (left,))
340
+ return check if isinstance(op_node, ast.Is) else lua.UnOp("not", check)
341
+
342
+ op = _COMPARE_OPS.get(type(op_node))
343
+ if op is None:
344
+ self.fail(
345
+ node,
346
+ f"the {type(op_node).__name__} comparison is not supported",
347
+ hint="Lua 5.1 has no 'in' operator; loop over the table instead."
348
+ if isinstance(op_node, ast.In | ast.NotIn)
349
+ else None,
350
+ )
351
+ return lua.BinOp(op, left, self.expr(right_node))
352
+
353
+ def fstring(self, node: ast.expr, values: list[ast.expr]) -> lua.Expr:
354
+ parts: list[lua.Expr] = []
355
+ for value in values:
356
+ if isinstance(value, ast.Constant) and isinstance(value.value, str):
357
+ parts.append(lua.Str(value.value))
358
+ elif isinstance(value, ast.FormattedValue):
359
+ if value.format_spec is not None or value.conversion not in (-1, 115):
360
+ self.fail(value, "format specs and conversions are not supported in f-strings")
361
+ parts.append(lua.Call(lua.Name("tostring"), (self.expr(value.value),)))
362
+ else: # pragma: no cover - JoinedStr only holds these two kinds
363
+ self.fail(value, "unsupported f-string component")
364
+ if not parts:
365
+ return lua.Str("")
366
+ # Lua's .. is right-associative; folding the same way avoids emitting a
367
+ # nest of parentheses that mean nothing.
368
+ result = parts[-1]
369
+ for part in reversed(parts[:-1]):
370
+ result = lua.BinOp("..", part, result)
371
+ return result
372
+
373
+ def subscript(self, node: ast.Subscript) -> lua.Expr:
374
+ if isinstance(node.slice, ast.Slice):
375
+ self.fail(
376
+ node,
377
+ "slicing is not supported",
378
+ hint="Loop over the table, or slice with a Redis command such as LRANGE.",
379
+ )
380
+ obj = self.expr(node.value)
381
+ return lua.Index(obj, self.index(node.slice))
382
+
383
+ def index(self, node: ast.expr) -> lua.Expr:
384
+ """Translate a 0-based Python index to Lua's 1-based one."""
385
+ if isinstance(node, ast.Constant) and isinstance(node.value, str):
386
+ return lua.Str(node.value) # string key: no offset
387
+ if isinstance(node, ast.Constant) and isinstance(node.value, int):
388
+ if node.value < 0:
389
+ self.fail(
390
+ node,
391
+ "negative indexing is not supported",
392
+ hint="Lua tables have no negative indices; use t[len(t) - 1] instead.",
393
+ )
394
+ return lua.Num(node.value + 1)
395
+ if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.USub):
396
+ self.fail(node, "negative indexing is not supported")
397
+ return lua.BinOp("+", self.expr(node), lua.Num(1))
398
+
399
+ # ----------------------------------------------------------------- calls
400
+
401
+ def call(self, node: ast.Call) -> lua.Expr:
402
+ if node.keywords:
403
+ self.fail(node, "keyword arguments are not supported in a script body")
404
+ args = tuple(self.expr(a) for a in node.args)
405
+
406
+ match node.func:
407
+ case ast.Name(id="len"):
408
+ if len(args) != 1:
409
+ self.fail(node, "len() takes exactly one argument")
410
+ return lua.UnOp("#", args[0])
411
+ case ast.Name(id=name) if name in _BUILTIN_FUNCS:
412
+ return lua.Call(lua.Name(_BUILTIN_FUNCS[name]), args)
413
+ case ast.Attribute(value=ast.Name(id=recv), attr=attr):
414
+ return self.namespace_call(node, recv, attr, args)
415
+ case ast.Attribute(attr=attr):
416
+ self.fail(
417
+ node,
418
+ f"method call .{attr}() is not supported",
419
+ hint="Only the redis and cjson namespaces, and list.append(), are available.",
420
+ )
421
+ case ast.Name(id=name):
422
+ self.fail(
423
+ node,
424
+ f"{name}() is not available inside a script",
425
+ hint="A script cannot call Python functions; only Redis commands "
426
+ "and a small set of builtins.",
427
+ )
428
+ case _:
429
+ self.fail(node, "unsupported call target")
430
+
431
+ def namespace_call(
432
+ self, node: ast.Call, recv: str, attr: str, args: tuple[lua.Expr, ...]
433
+ ) -> lua.Expr:
434
+ kind = self.receiver_kind(node, recv)
435
+ if kind == "cjson":
436
+ if attr not in {"encode", "decode"}:
437
+ self.fail(node, f"cjson has no {attr!r} function")
438
+ return lua.Call(lua.Index(lua.Name("cjson"), lua.Str(attr)), args)
439
+ if kind == "redis":
440
+ return self.redis_call(node, attr, args)
441
+ self.fail(
442
+ node,
443
+ f"method call .{attr}() is not supported",
444
+ hint="Only the redis and cjson namespaces, and list.append(), are available.",
445
+ )
446
+
447
+ def receiver_kind(self, node: ast.Call, name: str) -> str | None:
448
+ """Work out what the receiver of an attribute call refers to.
449
+
450
+ Resolution is by value, through the globals of the module that defined
451
+ the script, so the namespace works under any alias. Only a name bound
452
+ to nothing at all falls back to the conventional spellings.
453
+ """
454
+ value = self.globalns.get(name, _UNBOUND)
455
+ if isinstance(value, _Namespace):
456
+ return value.kind
457
+ if value is _UNBOUND:
458
+ return _RECEIVER_FALLBACK.get(name)
459
+ if _is_redis_py(value):
460
+ # Silently compiling this would aim the script at the client
461
+ # library, which is the one mistake this whole design invites.
462
+ self.fail(
463
+ node,
464
+ f"{name!r} is bound to redis-py here, not to the script namespace",
465
+ hint="Import the namespace under another name "
466
+ "(from redis_lua_py import redis as r), or the client under "
467
+ "another name (import redis as redis_client).",
468
+ )
469
+ return None
470
+
471
+ def redis_call(self, node: ast.Call, attr: str, args: tuple[lua.Expr, ...]) -> lua.Expr:
472
+ if attr in _REDIS_DIRECT:
473
+ return lua.Call(lua.Index(lua.Name("redis"), lua.Str(attr)), args)
474
+ if attr.startswith("_"):
475
+ self.fail(node, f"redis.{attr} is not a Redis command")
476
+ # `zrangebyscore` -> ZRANGEBYSCORE; `script_load` -> SCRIPT LOAD.
477
+ tokens = tuple(lua.Str(part.upper()) for part in attr.split("_") if part)
478
+ return lua.Call(lua.Index(lua.Name("redis"), lua.Str("call")), tokens + args)
479
+
480
+ # ------------------------------------------------------------ conditions
481
+
482
+ def condition(self, node: ast.expr) -> lua.Expr:
483
+ """Compile an expression used for its truth value.
484
+
485
+ Lua treats 0 and '' as true, so anything that is not already a boolean
486
+ gets routed through the __truthy helper.
487
+ """
488
+ match node:
489
+ case ast.Compare():
490
+ return self.compare(node)
491
+ case ast.BoolOp(op=op, values=values):
492
+ lua_op = "and" if isinstance(op, ast.And) else "or"
493
+ result = self.condition(values[0])
494
+ for value in values[1:]:
495
+ result = lua.BinOp(lua_op, result, self.condition(value))
496
+ return result
497
+ case ast.UnaryOp(op=ast.Not(), operand=operand):
498
+ return lua.UnOp("not", self.condition(operand))
499
+ case ast.Constant(value=bool() as v):
500
+ return lua.Bool(v)
501
+ case _:
502
+ self.needs_truthy = True
503
+ return lua.Call(lua.Name("__truthy"), (self.expr(node),))
504
+
505
+ # ----------------------------------------------------------- statements
506
+
507
+ def block(self, body: list[ast.stmt]) -> list[lua.Stat]:
508
+ out: list[lua.Stat] = []
509
+ for stmt in body:
510
+ out.extend(self.stmt(stmt))
511
+ return out
512
+
513
+ def stmt(self, node: ast.stmt) -> list[lua.Stat]:
514
+ match node:
515
+ case ast.Pass():
516
+ return []
517
+ case ast.Expr(value=ast.Constant(value=str())):
518
+ return [] # a stray string literal, e.g. a docstring
519
+ case ast.Expr(value=ast.Call() as inner):
520
+ return self.call_statement(inner)
521
+ case ast.Expr():
522
+ self.fail(node, "this expression has no effect in Lua")
523
+ case ast.Assign(targets=targets, value=value):
524
+ if len(targets) != 1:
525
+ self.fail(node, "chained assignment (a = b = c) is not supported")
526
+ return self.assign(node, targets[0], value)
527
+ case ast.AnnAssign(target=target, value=value):
528
+ if value is None:
529
+ self.fail(node, "a bare annotation declares nothing in Lua")
530
+ return self.assign(node, target, value)
531
+ case ast.AugAssign(target=target, op=op, value=value):
532
+ synthetic = ast.BinOp(left=target, op=op, right=value)
533
+ ast.copy_location(synthetic, node)
534
+ return self.assign(node, target, synthetic, augmented=True)
535
+ case ast.Return(value=value):
536
+ return [lua.Return(None if value is None else self.expr(value))]
537
+ case ast.If():
538
+ return [self.if_stmt(node)]
539
+ case ast.For():
540
+ return [self.for_stmt(node)]
541
+ case ast.While(test=test, orelse=orelse, body=body):
542
+ if orelse:
543
+ self.fail(node, "while/else is not supported")
544
+ return [lua.While(self.condition(test), self.block(body))]
545
+ case ast.Break():
546
+ return [lua.Break()]
547
+ case ast.Continue():
548
+ self.fail(
549
+ node,
550
+ "Lua 5.1 has no 'continue' statement",
551
+ hint="Invert the condition and put the rest of the loop body inside the if.",
552
+ )
553
+ case ast.Assert():
554
+ self.fail(
555
+ node,
556
+ "assert is not supported",
557
+ hint="Return redis.error_reply('...') to signal failure to the caller.",
558
+ )
559
+ case ast.Try() | ast.Raise():
560
+ self.fail(
561
+ node,
562
+ "exception handling is not supported",
563
+ hint="Use redis.pcall() and check the result for an 'err' field.",
564
+ )
565
+ case ast.FunctionDef() | ast.AsyncFunctionDef() | ast.ClassDef() | ast.Lambda():
566
+ self.fail(node, "a script cannot define nested functions or classes")
567
+ case ast.Import() | ast.ImportFrom():
568
+ self.fail(node, "a script cannot import anything")
569
+ case ast.With() | ast.AsyncWith():
570
+ self.fail(node, "'with' is not supported")
571
+ case ast.Global() | ast.Nonlocal():
572
+ self.fail(node, "'global' and 'nonlocal' are not supported")
573
+ case _:
574
+ self.fail(node, f"{type(node).__name__} statements are not supported")
575
+
576
+ def call_statement(self, node: ast.Call) -> list[lua.Stat]:
577
+ # `items.append(x)` is the one method call worth special-casing: it is
578
+ # how you build a return value, and Lua spells it t[#t + 1] = x.
579
+ if (
580
+ isinstance(node.func, ast.Attribute)
581
+ and node.func.attr == "append"
582
+ and isinstance(node.func.value, ast.Name)
583
+ and node.func.value.id in self.known
584
+ ):
585
+ if len(node.args) != 1:
586
+ self.fail(node, "append() takes exactly one argument")
587
+ target = self.expr(node.func.value)
588
+ slot = lua.Index(target, lua.BinOp("+", lua.UnOp("#", target), lua.Num(1)))
589
+ return [lua.Assign([slot], [self.expr(node.args[0])])]
590
+ return [lua.ExprStat(self.call(node))]
591
+
592
+ def assign(
593
+ self, node: ast.stmt, target: ast.expr, value: ast.expr, *, augmented: bool = False
594
+ ) -> list[lua.Stat]:
595
+ rhs = self.expr(value)
596
+ match target:
597
+ case ast.Name(id=name):
598
+ if not lua.is_identifier(name):
599
+ self.fail(target, f"{name!r} is a reserved word in Lua")
600
+ if not augmented and id(node) in self.simple_assigns:
601
+ return [lua.Local([name], [rhs])]
602
+ return [lua.Assign([lua.Name(name)], [rhs])]
603
+ case ast.Subscript():
604
+ return [lua.Assign([self.subscript(target)], [rhs])]
605
+ case ast.Tuple() | ast.List():
606
+ self.fail(target, "tuple unpacking is not supported")
607
+ case _:
608
+ self.fail(target, "unsupported assignment target")
609
+
610
+ def if_stmt(self, node: ast.If) -> lua.If:
611
+ branches = [(self.condition(node.test), self.block(node.body))]
612
+ orelse = node.orelse
613
+ # Collapse `else: if ...` chains into Lua's elseif.
614
+ while len(orelse) == 1 and isinstance(orelse[0], ast.If):
615
+ nested = orelse[0]
616
+ branches.append((self.condition(nested.test), self.block(nested.body)))
617
+ orelse = nested.orelse
618
+ return lua.If(branches, self.block(orelse))
619
+
620
+ def for_stmt(self, node: ast.For) -> lua.Stat:
621
+ if node.orelse:
622
+ self.fail(node, "for/else is not supported")
623
+ if not isinstance(node.target, ast.Name):
624
+ self.fail(node.target, "only a single loop variable is supported")
625
+ var = node.target.id
626
+ self.known.add(var)
627
+
628
+ if (
629
+ isinstance(node.iter, ast.Call)
630
+ and isinstance(node.iter.func, ast.Name)
631
+ and node.iter.func.id == "range"
632
+ ):
633
+ return self.range_loop(node, var, node.iter)
634
+
635
+ # Iterating a table: bind the sequence once, then walk it by index so
636
+ # that a call in the iterable is not re-evaluated every step.
637
+ iterable = self.expr(node.iter)
638
+ self._temp += 1
639
+ idx = f"__i{self._temp}"
640
+
641
+ # Bind the iterable to a temporary so a call is not re-evaluated on
642
+ # every step. A plain name is already stable, so leave it alone.
643
+ prefix: list[lua.Stat] = []
644
+ if isinstance(iterable, lua.Name):
645
+ seq: lua.Expr = iterable
646
+ else:
647
+ seq = lua.Name(f"__seq{self._temp}")
648
+ prefix = [lua.Local([f"__seq{self._temp}"], [iterable])]
649
+
650
+ body: list[lua.Stat] = [
651
+ lua.Local([var], [lua.Index(seq, lua.Name(idx))]),
652
+ *self.block(node.body),
653
+ ]
654
+ return _Block([*prefix, lua.NumericFor(idx, lua.Num(1), lua.UnOp("#", seq), None, body)])
655
+
656
+ def range_loop(self, node: ast.For, var: str, call: ast.Call) -> lua.Stat:
657
+ args = call.args
658
+ if not 1 <= len(args) <= 3:
659
+ self.fail(call, "range() takes one to three arguments")
660
+
661
+ step: lua.Expr | None = None
662
+ descending = False
663
+ if len(args) == 3:
664
+ step_value = _literal_int(args[2])
665
+ if step_value is None:
666
+ self.fail(args[2], "range() step must be an integer literal")
667
+ if step_value == 0:
668
+ self.fail(args[2], "range() step cannot be zero")
669
+ descending = step_value < 0
670
+ step = lua.Num(step_value)
671
+
672
+ if len(args) == 1:
673
+ start: lua.Expr = lua.Num(0)
674
+ stop_node = args[0]
675
+ else:
676
+ start = self.expr(args[0])
677
+ stop_node = args[1]
678
+
679
+ # Python's range excludes the stop value; Lua's numeric for includes it.
680
+ stop = self._offset(self.expr(stop_node), 1 if descending else -1)
681
+ return lua.NumericFor(var, start, stop, step, self.block(node.body))
682
+
683
+ @staticmethod
684
+ def _offset(expr: lua.Expr, delta: int) -> lua.Expr:
685
+ if isinstance(expr, lua.Num) and isinstance(expr.value, int):
686
+ return lua.Num(expr.value + delta)
687
+ return lua.BinOp("+" if delta > 0 else "-", expr, lua.Num(abs(delta)))
688
+
689
+
690
+ class _Block(lua.Stat):
691
+ """Several statements where the grammar expects one."""
692
+
693
+ __slots__ = ("body",)
694
+
695
+ def __init__(self, body: list[lua.Stat]) -> None:
696
+ self.body = body
697
+
698
+
699
+ def _literal_int(node: ast.expr) -> int | None:
700
+ if isinstance(node, ast.Constant) and isinstance(node.value, int):
701
+ return node.value
702
+ if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.USub):
703
+ inner = _literal_int(node.operand)
704
+ return None if inner is None else -inner
705
+ return None
706
+
707
+
708
+ def _flatten(body: list[lua.Stat]) -> list[lua.Stat]:
709
+ out: list[lua.Stat] = []
710
+ for stat in body:
711
+ if isinstance(stat, _Block):
712
+ out.extend(_flatten(stat.body))
713
+ elif isinstance(stat, lua.If):
714
+ stat.branches = [(t, _flatten(b)) for t, b in stat.branches]
715
+ stat.orelse = _flatten(stat.orelse)
716
+ out.append(stat)
717
+ elif isinstance(stat, lua.While | lua.NumericFor):
718
+ stat.body = _flatten(stat.body)
719
+ out.append(stat)
720
+ else:
721
+ out.append(stat)
722
+ return out
723
+
724
+
725
+ def parse_function(func: Callable[..., Any]) -> tuple[ast.FunctionDef, str, int, list[str]]:
726
+ try:
727
+ lines, first_lineno = inspect.getsourcelines(func)
728
+ except (OSError, TypeError) as exc: # pragma: no cover - needs an exotic environment
729
+ raise CompileError(
730
+ f"cannot read the source of {func.__name__!r}. A script must be defined in a "
731
+ "file on disk, not in a REPL or an exec() string."
732
+ ) from exc
733
+
734
+ source = textwrap.dedent("".join(lines))
735
+ module = ast.parse(source)
736
+ node = module.body[0]
737
+ if not isinstance(node, ast.FunctionDef):
738
+ raise CompileError(f"@script can only be applied to a function, got {type(node).__name__}")
739
+ node.decorator_list = []
740
+ filename = inspect.getsourcefile(func) or "<unknown>"
741
+ return node, filename, first_lineno, source.splitlines()
742
+
743
+
744
+ def compile_function(func: Callable[..., Any], *, name: str | None = None) -> CompiledScript:
745
+ node, filename, first_lineno, lines = parse_function(func)
746
+ compiler = _Compiler(
747
+ node,
748
+ filename=filename,
749
+ first_lineno=first_lineno,
750
+ lines=lines,
751
+ # Receivers are resolved against the defining module, so the namespace
752
+ # is recognised under whatever name it was imported as.
753
+ globalns=getattr(func, "__globals__", {}),
754
+ )
755
+
756
+ prelude = compiler.compile_signature()
757
+ compiler.collect_assigned()
758
+
759
+ body = node.body
760
+ doc = ast.get_docstring(node)
761
+ if doc is not None:
762
+ body = body[1:]
763
+
764
+ statements = _flatten(compiler.block(body))
765
+ if compiler.hoisted:
766
+ prelude.append(lua.Local(sorted(compiler.hoisted), []))
767
+
768
+ header = [
769
+ f"-- {name or func.__name__}",
770
+ f"-- Generated by redis-lua-py from {filename}:{first_lineno}. Do not edit.",
771
+ ]
772
+ parts = ["\n".join(header)]
773
+ if compiler.needs_truthy:
774
+ parts.append(_TRUTHY_HELPER.rstrip())
775
+ if compiler.needs_isnil:
776
+ parts.append(_ISNIL_HELPER.rstrip())
777
+ parts.append(lua.emit(prelude + statements).rstrip())
778
+
779
+ return CompiledScript(
780
+ name=name or func.__name__,
781
+ lua="\n".join(parts) + "\n",
782
+ params=tuple(compiler.params),
783
+ keys=tuple(compiler.keys),
784
+ args=tuple(compiler.args),
785
+ doc=doc,
786
+ source=f"{filename}:{first_lineno}",
787
+ )