implement 'compile_expression()'

This commit is contained in:
Mike Fährmann
2021-06-03 22:34:58 +02:00
parent e39c4633ba
commit 0abad8bc12
2 changed files with 43 additions and 13 deletions

View File

@@ -21,6 +21,7 @@ import sqlite3
import binascii import binascii
import datetime import datetime
import operator import operator
import functools
import itertools import itertools
import urllib.parse import urllib.parse
from http.cookiejar import Cookie from http.cookiejar import Cookie
@@ -346,8 +347,6 @@ CODES = {
"zh": "Chinese", "zh": "Chinese",
} }
SPECIAL_EXTRACTORS = {"oauth", "recursive", "test"}
class UniversalNone(): class UniversalNone():
"""None-style object that supports more operations than None itself""" """None-style object that supports more operations than None itself"""
@@ -373,6 +372,20 @@ class UniversalNone():
NONE = UniversalNone() NONE = UniversalNone()
WINDOWS = (os.name == "nt") WINDOWS = (os.name == "nt")
SENTINEL = object() SENTINEL = object()
SPECIAL_EXTRACTORS = {"oauth", "recursive", "test"}
GLOBALS = {
"parse_int": text.parse_int,
"urlsplit" : urllib.parse.urlsplit,
"datetime" : datetime.datetime,
"abort" : raises(exception.StopExtraction),
"terminate": raises(exception.TerminateExtraction),
"re" : re,
}
def compile_expression(expr, name="<expr>", globals=GLOBALS):
code_object = compile(expr, name, "eval")
return functools.partial(eval, code_object, globals)
def build_predicate(predicates): def build_predicate(predicates):
@@ -472,20 +485,13 @@ class UniquePredicate():
class FilterPredicate(): class FilterPredicate():
"""Predicate; True if evaluating the given expression returns True""" """Predicate; True if evaluating the given expression returns True"""
def __init__(self, filterexpr, target="image"): def __init__(self, expr, target="image"):
name = "<{} filter>".format(target) name = "<{} filter>".format(target)
self.codeobj = compile(filterexpr, name, "eval") self.expr = compile_expression(expr, name)
self.globals = {
"parse_int": text.parse_int,
"urlsplit" : urllib.parse.urlsplit,
"datetime" : datetime.datetime,
"abort" : raises(exception.StopExtraction),
"re" : re,
}
def __call__(self, url, kwds): def __call__(self, _, kwdict):
try: try:
return eval(self.codeobj, self.globals, kwds) return self.expr(kwdict)
except exception.GalleryDLException: except exception.GalleryDLException:
raise raise
except Exception as exc: except Exception as exc:

View File

@@ -493,6 +493,30 @@ class TestOther(unittest.TestCase):
def test_noop(self): def test_noop(self):
self.assertEqual(util.noop(), None) self.assertEqual(util.noop(), None)
def test_compile_expression(self):
expr = util.compile_expression("1 + 2 * 3")
self.assertEqual(expr(), 7)
self.assertEqual(expr({"a": 1, "b": 2, "c": 3}), 7)
self.assertEqual(expr({"a": 9, "b": 9, "c": 9}), 7)
expr = util.compile_expression("a + b * c")
self.assertEqual(expr({"a": 1, "b": 2, "c": 3}), 7)
self.assertEqual(expr({"a": 9, "b": 9, "c": 9}), 90)
with self.assertRaises(NameError):
expr()
with self.assertRaises(NameError):
expr({"a": 2})
with self.assertRaises(SyntaxError):
util.compile_expression("")
with self.assertRaises(SyntaxError):
util.compile_expression("x++")
expr = util.compile_expression("1 and abort()")
with self.assertRaises(exception.StopExtraction):
expr()
def test_generate_token(self): def test_generate_token(self):
tokens = set() tokens = set()
for _ in range(100): for _ in range(100):