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.
- redis_lua_py/__init__.py +87 -0
- redis_lua_py/_compile.py +787 -0
- redis_lua_py/_lua.py +308 -0
- redis_lua_py/_runtime.py +72 -0
- redis_lua_py/_script.py +151 -0
- redis_lua_py/errors.py +53 -0
- redis_lua_py/py.typed +0 -0
- redis_lua_py-0.1.0.dist-info/METADATA +320 -0
- redis_lua_py-0.1.0.dist-info/RECORD +11 -0
- redis_lua_py-0.1.0.dist-info/WHEEL +4 -0
- redis_lua_py-0.1.0.dist-info/licenses/LICENSE +21 -0
redis_lua_py/_compile.py
ADDED
|
@@ -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
|
+
)
|