From accbc0b89d4998bf0ab7f9f97d72f19e7f7dc91c Mon Sep 17 00:00:00 2001 From: Valentin Lorentz Date: Mon, 18 Jan 2016 11:23:58 +0100 Subject: [PATCH] Add Python 3 support. --- .travis.yml | 1 + py14/analysis.py | 16 +++++++++ py14/clike.py | 8 +++++ py14/context.py | 6 ++-- py14/scope.py | 7 ++-- py14/tests/test_transpiler.py | 16 ++++++--- py14/tracer.py | 12 ++++--- py14/transpiler.py | 61 +++++++++++++++++++++++++++-------- regtests/test_range.py | 6 ++-- 9 files changed, 99 insertions(+), 34 deletions(-) diff --git a/.travis.yml b/.travis.yml index 68e04d5..337474d 100644 --- a/.travis.yml +++ b/.travis.yml @@ -1,6 +1,7 @@ language: python python: - "2.7" + - "3.4" install: - "pip install -r requirements.txt --use-mirrors" - "pip install coveralls" diff --git a/py14/analysis.py b/py14/analysis.py index ecba837..0a6d012 100644 --- a/py14/analysis.py +++ b/py14/analysis.py @@ -1,3 +1,4 @@ +import sys import ast @@ -11,6 +12,21 @@ def is_void_function(fun): finder.visit(fun) return not finder.returns +if sys.version_info[0] >= 3: + def get_id(var): + if isinstance(var, ast.alias): + return var.name + elif isinstance(var, ast.Name): + return var.id + elif isinstance(var, ast.arg): + return var.arg +else: + def get_id(var): + if isinstance(var, ast.alias): + return var.name + elif isinstance(var, ast.Name): + return var.id + class ReturnFinder(ast.NodeVisitor): returns = False diff --git a/py14/clike.py b/py14/clike.py index 5900f44..cf1fb53 100644 --- a/py14/clike.py +++ b/py14/clike.py @@ -47,6 +47,14 @@ def visit_Name(self, node): if node.id in self.builtin_constants: return node.id.lower() return node.id + + def visit_NameConstant(self, node): + if node.value is True: + return "true" + elif node.value is False: + return "false" + else: + return node.value def visit_Num(self, node): return str(node.n) diff --git a/py14/context.py b/py14/context.py index bbcd3c1..007ab7f 100644 --- a/py14/context.py +++ b/py14/context.py @@ -1,5 +1,5 @@ import ast -from scope import ScopeMixin +from .scope import ScopeMixin def add_list_calls(node): @@ -55,11 +55,11 @@ def visit_Import(self, node): def visit_If(self, node): node.vars = [] - map(self.visit, node.body) + list(map(self.visit, node.body)) node.body_vars = node.vars node.vars = [] - map(self.visit, node.orelse) + list(map(self.visit, node.orelse)) node.orelse_vars = node.vars node.vars = [] diff --git a/py14/scope.py b/py14/scope.py index b6bc8b5..8483da3 100644 --- a/py14/scope.py +++ b/py14/scope.py @@ -1,6 +1,7 @@ import ast from contextlib import contextmanager +from .analysis import get_id def add_scope_context(node): """Provide to scope context to all nodes""" @@ -42,13 +43,9 @@ class ScopeList(list): """ def find(self, lookup): """Find definition of variable lookup.""" - def is_match(var): - return ((isinstance(var, ast.alias) and var.name == lookup) or - (isinstance(var, ast.Name) and var.id == lookup)) - def find_definition(scope, var_attr="vars"): for var in getattr(scope, var_attr): - if is_match(var): + if get_id(var) == lookup: return var for scope in self: diff --git a/py14/tests/test_transpiler.py b/py14/tests/test_transpiler.py index 128e226..14edefa 100644 --- a/py14/tests/test_transpiler.py +++ b/py14/tests/test_transpiler.py @@ -1,5 +1,6 @@ -from py14.transpiler import transpile +import sys import pytest +from py14.transpiler import transpile def parse(*args): @@ -26,7 +27,10 @@ def test_empty_return(): def test_print_multiple_vars(): - source = parse('print("hi", "there" )') + if sys.version_info[0] >= 3: + source = parse('print(("hi", "there" ))') + else: + source = parse('print("hi", "there" )') cpp = transpile(source) assert cpp == ('std::cout << std::string {"hi"} ' '<< std::string {"there"} << std::endl;') @@ -228,12 +232,13 @@ def test_bubble_sort(): " return seq", ) cpp = transpile(source) + range_f = "range" if sys.version_info[0] < 3 else "xrange" assert cpp == parse( "template ", "auto sort(T1 seq) {", "auto L = seq.size();", - "for(auto _ : rangepp::range(L)) {", - "for(auto n : rangepp::range(1, L)) {", + "for(auto _ : rangepp::{0}(L)) {{".format(range_f), + "for(auto n : rangepp::{0}(1, L)) {{".format(range_f), "if(seq[n] < seq[n - 1]) {", "std::tie(seq[n - 1], seq[n]) = " "std::make_tuple(seq[n], seq[n - 1]);", @@ -287,6 +292,7 @@ def test_comb_sort(): " return seq", ) cpp = transpile(source) + range_f = "range" if sys.version_info[0] < 3 else "xrange" assert cpp == parse( "template ", "auto sort(T1 seq) {", @@ -295,7 +301,7 @@ def test_comb_sort(): "while (gap > 1||swap) {", "gap = std::max(1, py14::to_int(gap / 1.25));", "swap = false;", - "for(auto i : rangepp::range(seq.size() - gap)) {", + "for(auto i : rangepp::{0}(seq.size() - gap)) {{".format(range_f), "if(seq[i] > seq[i + gap]) {", "std::tie(seq[i], seq[i + gap]) = " "std::make_tuple(seq[i + gap], seq[i]);", diff --git a/py14/tracer.py b/py14/tracer.py index db245d8..178e5da 100644 --- a/py14/tracer.py +++ b/py14/tracer.py @@ -2,7 +2,8 @@ Trace object types that are inserted into Python list. """ import ast -from clike import CLikeTranspiler +from .analysis import get_id +from .clike import CLikeTranspiler def decltype(node): @@ -24,7 +25,7 @@ def is_list(node): elif isinstance(node, ast.Assign): return is_list(node.value) elif isinstance(node, ast.Name): - var = node.scopes.find(node.id) + var = node.scopes.find(get_id(node)) return (hasattr(var, "assigned_from") and not isinstance(var.assigned_from, ast.FunctionDef) and is_list(var.assigned_from.value)) @@ -59,13 +60,13 @@ def visit_Str(self, node): return node.s def visit_Name(self, node): - var = node.scopes.find(node.id) + var = node.scopes.find(get_id(node)) if isinstance(var.assigned_from, ast.For): it = var.assigned_from.iter return "std::declval()".format( self.visit(it)) elif isinstance(var.assigned_from, ast.FunctionDef): - return var.id + return get_id(var) else: return self.visit(var.assigned_from.value) @@ -99,6 +100,9 @@ def visit_Name(self, node): else: return self.visit(var.assigned_from.value) + def visit_NameConstant(self, node): + return CLikeTranspiler().visit(node) + def visit_Call(self, node): params = ",".join([self.visit(arg) for arg in node.args]) return "{0}({1})".format(node.func.id, params) diff --git a/py14/transpiler.py b/py14/transpiler.py index 22b5c42..814d8ee 100644 --- a/py14/transpiler.py +++ b/py14/transpiler.py @@ -1,9 +1,10 @@ +import sys import ast -from clike import CLikeTranspiler -from scope import add_scope_context -from analysis import add_imports, is_void_function -from context import add_variable_context, add_list_calls -from tracer import decltype, is_list, is_builtin_import, defined_before +from .clike import CLikeTranspiler +from .scope import add_scope_context +from .context import add_variable_context, add_list_calls +from .analysis import add_imports, is_void_function, get_id +from .tracer import decltype, is_list, is_builtin_import, defined_before def transpile(source, headers=False, testing=False): @@ -43,7 +44,7 @@ def generate_catch_test_case(node, body): def generate_template_fun(node, body): params = [] for idx, arg in enumerate(node.args.args): - params.append(("T" + str(idx + 1), arg.id)) + params.append(("T" + str(idx + 1), get_id(arg))) typenames = ["typename " + arg[0] for arg in params] template = "inline " @@ -90,9 +91,10 @@ def visit_FunctionDef(self, node): def visit_Attribute(self, node): attr = node.attr - if is_builtin_import(node.value.id): - return "py14::" + node.value.id + "::" + attr - elif node.value.id == "math": + value_id = get_id(node.value) + if is_builtin_import(value_id): + return "py14::" + value_id + "::" + attr + elif value_id == "math": if node.attr == "asin": return "std::asin" elif node.attr == "atan": @@ -103,7 +105,7 @@ def visit_Attribute(self, node): if is_list(node.value): if node.attr == "append": attr = "push_back" - return node.value.id + "." + attr + return value_id + "." + attr def visit_Call(self, node): fname = self.visit(node.func) @@ -120,11 +122,24 @@ def visit_Call(self, node): elif fname == "max": return "std::max({0})".format(args) elif fname == "range": - return "rangepp::range({0})".format(args) + if sys.version_info[0] >= 3: + return "rangepp::xrange({0})".format(args) + else: + return "rangepp::range({0})".format(args) elif fname == "xrange": return "rangepp::xrange({0})".format(args) elif fname == "len": return "{0}.size()".format(self.visit(node.args[0])) + elif fname == "print": + buf = [] + for n in node.args: + value = self.visit(n) + if isinstance(n, ast.List) or isinstance(n, ast.Tuple): + buf.append("std::cout << {0} << std::endl;".format( + " << ".join([self.visit(el) for el in n.elts]))) + else: + buf.append('std::cout << {0} << std::endl;'.format(value)) + return '\n'.join(buf) return '{0}({1})'.format(fname, args) @@ -157,9 +172,17 @@ def visit_Name(self, node): else: return super(CppTranspiler, self).visit_Name(node) + def visit_NameConstant(self, node): + if node.value is True: + return "true" + elif node.value is False: + return "false" + else: + return super(CppTranspiler, self).visit_NameConstant(node) + def visit_If(self, node): - body_vars = set([v.id for v in node.scopes[-1].body_vars]) - orelse_vars = set([v.id for v in node.scopes[-1].orelse_vars]) + body_vars = set([get_id(v) for v in node.scopes[-1].body_vars]) + orelse_vars = set([get_id(v) for v in node.scopes[-1].orelse_vars]) node.common_vars = body_vars.intersection(orelse_vars) var_definitions = [] @@ -179,6 +202,16 @@ def visit_If(self, node): return ("".join(var_definitions) + super(CppTranspiler, self).visit_If(node)) + def visit_UnaryOp(self, node): + if isinstance(node.op, ast.USub): + if isinstance(node.operand, (ast.Call, ast.Num)): + # Shortcut if parenthesis are not needed + return "-{0}".format(self.visit(node.operand)) + else: + return "-({0})".format(self.visit(node.operand)) + else: + return super(CppTranspiler, self).visit_UnaryOp(node) + def visit_BinOp(self, node): if (isinstance(node.left, ast.List) and isinstance(node.op, ast.Mult) @@ -197,7 +230,7 @@ def visit_alias(self, node): def visit_Import(self, node): imports = [self.visit(n) for n in node.names] - return "\n".join(filter(None, imports)) + return "\n".join(i for i in imports if i) def visit_List(self, node): if len(node.elts) > 0: diff --git a/regtests/test_range.py b/regtests/test_range.py index dd42734..8e83094 100644 --- a/regtests/test_range.py +++ b/regtests/test_range.py @@ -1,16 +1,16 @@ def test_range(): - simple = range(0, 10) + simple = list(range(0, 10)) results = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9] assert simple == results def test_range_with_steps(): - with_steps = range(0, 10, 2) + with_steps = list(range(0, 10, 2)) results = [0, 2, 4, 6, 8] assert with_steps == results def test_range_with_negative_steps(): - with_steps = range(10, 0, -2) + with_steps = list(range(10, 0, -2)) results = [10, 8, 6, 4, 2] assert with_steps == results