Kaydet (Commit) c51d586a authored tarafından Berker Peksag's avatar Berker Peksag

Merge pull request #40 from berkerpeksag/pm_develop

pm_develop contains all the work I have been doing over the last couple of weeks
......@@ -15,10 +15,13 @@ astor is designed to allow easy manipulation of Python source via the AST.
There are some other similar libraries, but astor focuses on the following areas:
- Round-trip back to Python via Armin Ronacher's codegen.py module:
- Round-trip an AST back to Python:
- Modified AST doesn't need linenumbers, ctx, etc. or otherwise be directly compileable
- Modified AST doesn't need linenumbers, ctx, etc. or otherwise
be directly compileable for the round-trip to work.
- Easy to read generated code as, well, code
- Can round-trip two different source trees to compare for functional
differences, using the astor.rtrip tool (for example, after PEP8 edits).
- Dump pretty-printing of AST
......
......@@ -9,8 +9,6 @@ Copyright 2013 (c) Berker Peksag
"""
__version__ = '0.6'
from .code_gen import to_source # NOQA
from .node_util import iter_node, strip_tree, dump_tree
from .node_util import ExplicitNodeVisitor
......@@ -19,8 +17,10 @@ from .op_util import get_op_symbol, get_op_precedence # NOQA
from .op_util import symbol_data
from .tree_walk import TreeWalk # NOQA
__version__ = '0.6'
#DEPRECATED!!!
# DEPRECATED!!!
# These aliases support old programs. Please do not use in future.
......@@ -30,7 +30,7 @@ from .tree_walk import TreeWalk # NOQA
# things could be accessed from their submodule.
get_boolop = get_binop = get_cmpop = get_unaryop = get_op_symbol # NOQA
get_boolop = get_binop = get_cmpop = get_unaryop = get_op_symbol # NOQA
get_anyop = get_op_symbol
parsefile = code_to_ast.parse_file
codetoast = code_to_ast
......
......@@ -20,11 +20,14 @@ this code came from here (in 2012):
import ast
import sys
from .op_util import get_op_symbol
from .op_util import get_op_symbol, get_op_precedence, Precedence
from .node_util import ExplicitNodeVisitor
from .string_repr import pretty_string
from .source_repr import pretty_source
def to_source(node, indent_with=' ' * 4, add_line_information=False):
def to_source(node, indent_with=' ' * 4, add_line_information=False,
pretty_string=pretty_string, pretty_source=pretty_source):
"""This function can convert a node tree back into python sourcecode.
This is useful for debugging purposes, especially if you're dealing with
custom asts not generated by python itself.
......@@ -43,19 +46,70 @@ def to_source(node, indent_with=' ' * 4, add_line_information=False):
number information of statement nodes.
"""
generator = SourceGenerator(indent_with, add_line_information)
generator = SourceGenerator(indent_with, add_line_information,
pretty_string)
generator.visit(node)
return ''.join(str(s) for s in generator.result)
generator.result.append('\n')
return pretty_source(str(s) for s in generator.result)
def enclose(enclosure):
def decorator(func):
def newfunc(self, node):
self.write(enclosure[0])
func(self, node)
self.write(enclosure[-1])
return newfunc
return decorator
def set_precedence(value, *nodes):
"""Set the precedence (of the parent) into the children.
"""
if isinstance(value, ast.AST):
value = get_op_precedence(value)
for node in nodes:
if isinstance(node, ast.AST):
node._pp = value
elif isinstance(node, list):
set_precedence(value, *node)
else:
assert node is None, node
class Delimit(object):
"""A context manager that can add enclosing
delimiters around the output of a
SourceGenerator method. By default, the
parentheses are added, but the enclosed code
may set discard=True to get rid of them.
"""
discard = False
def __init__(self, tree, *args):
""" use write instead of using result directly
for initial data, because it may flush
preceding data into result.
"""
delimiters = '()'
node = None
op = None
for arg in args:
if isinstance(arg, ast.AST):
if node is None:
node = arg
else:
op = arg
else:
delimiters = arg
tree.write(delimiters[0])
result = self.result = tree.result
self.index = len(result)
self.closing = delimiters[1]
if node is not None:
self.p = p = get_op_precedence(op or node)
self.pp = pp = tree.get__pp(node)
self.discard = p >= pp
def __enter__(self):
return self
def __exit__(self, *exc_info):
if self.discard:
self.result[self.index - 1] = ''
else:
self.result.append(self.closing)
class SourceGenerator(ExplicitNodeVisitor):
......@@ -67,12 +121,34 @@ class SourceGenerator(ExplicitNodeVisitor):
"""
def __init__(self, indent_with, add_line_information=False):
def __init__(self, indent_with, add_line_information=False,
pretty_string=pretty_string):
self.result = []
self.indent_with = indent_with
self.add_line_information = add_line_information
self.indentation = 0
self.new_lines = 0
self.pretty_string = pretty_string
def __getattr__(self, name, defaults=dict(keywords=(),
_pp=Precedence.highest).get):
""" Get an attribute of the node.
like dict.get (returns None if doesn't exist)
"""
if not name.startswith('get_'):
raise AttributeError
geta = getattr
shortname = name[4:]
default = defaults(shortname)
def getter(node):
return geta(node, shortname, default)
setattr(self, name, getter)
return getter
def delimit(self, *args):
return Delimit(self, *args)
def write(self, *params):
for item in params:
......@@ -93,6 +169,7 @@ class SourceGenerator(ExplicitNodeVisitor):
def conditional_write(self, *stuff):
if stuff[-1] is not None:
self.write(*stuff)
# Inform the caller that we wrote
return True
def newline(self, node=None, extra=0):
......@@ -125,6 +202,7 @@ class SourceGenerator(ExplicitNodeVisitor):
want_comma.append(True)
def loop_args(args, defaults):
set_precedence(Precedence.Comma, defaults)
padding = [None] * (len(args) - len(defaults))
for arg, default in zip(args, padding + defaults):
self.write(write_comma, arg)
......@@ -133,7 +211,7 @@ class SourceGenerator(ExplicitNodeVisitor):
loop_args(node.args, node.defaults)
self.conditional_write(write_comma, '*', node.vararg)
kwonlyargs = getattr(node, 'kwonlyargs', None)
kwonlyargs = self.get_kwonlyargs(node)
if kwonlyargs:
if node.vararg is None:
self.write(write_comma, '*')
......@@ -150,6 +228,7 @@ class SourceGenerator(ExplicitNodeVisitor):
self.statement(decorator, '@', decorator)
def comma_list(self, items, trailing=False):
set_precedence(Precedence.Comma, *items)
for idx, item in enumerate(items):
self.write(', ' if idx else '', item)
self.write(',' if trailing else '')
......@@ -157,18 +236,20 @@ class SourceGenerator(ExplicitNodeVisitor):
# Statements
def visit_Assign(self, node):
set_precedence(node, node.value, *node.targets)
self.newline(node)
for target in node.targets:
self.write(target, ' = ')
self.visit(node.value)
def visit_AugAssign(self, node):
set_precedence(node, node.value, node.target)
self.statement(node, node.target, get_op_symbol(node.op, ' %s= '),
node.value)
def visit_ImportFrom(self, node):
self.statement(node, 'from ', node.level * '.',
node.module or '', ' import ')
node.module or '', ' import ')
self.comma_list(node.names)
def visit_Import(self, node):
......@@ -176,16 +257,17 @@ class SourceGenerator(ExplicitNodeVisitor):
self.comma_list(node.names)
def visit_Expr(self, node):
set_precedence(node, node.value)
self.statement(node)
self.generic_visit(node)
def visit_FunctionDef(self, node, async=False):
prefix = 'async ' if async else ''
self.decorators(node, 1)
self.statement(node, '%sdef %s(' % (prefix, node.name))
self.statement(node, '%sdef %s' % (prefix, node.name), '(')
self.visit_arguments(node.args)
self.write(')')
self.conditional_write(' ->', getattr(node, 'returns', None))
self.conditional_write(' ->', self.get_returns(node))
self.write(':')
self.body(node.body)
......@@ -207,24 +289,24 @@ class SourceGenerator(ExplicitNodeVisitor):
self.statement(node, 'class %s' % node.name)
for base in node.bases:
self.write(paren_or_comma, base)
#keywords not available in early version
for keyword in getattr(node, 'keywords', ()):
# keywords not available in early version
for keyword in self.get_keywords(node):
self.write(paren_or_comma, keyword.arg or '',
'=' if keyword.arg else '**', keyword.value)
self.conditional_write(paren_or_comma, '*',
getattr(node, 'starargs', None))
self.conditional_write(paren_or_comma, '**',
getattr(node, 'kwargs', None))
self.conditional_write(paren_or_comma, '*', self.get_starargs(node))
self.conditional_write(paren_or_comma, '**', self.get_kwargs(node))
self.write(have_args and '):' or ':')
self.body(node.body)
def visit_If(self, node):
set_precedence(node, node.test)
self.statement(node, 'if ', node.test, ':')
self.body(node.body)
while True:
else_ = node.orelse
if len(else_) == 1 and isinstance(else_[0], ast.If):
node = else_[0]
set_precedence(node, node.test)
self.write('\n', 'elif ', node.test, ':')
self.body(node.body)
else:
......@@ -232,6 +314,7 @@ class SourceGenerator(ExplicitNodeVisitor):
break
def visit_For(self, node, async=False):
set_precedence(node, node.target)
prefix = 'async ' if async else ''
self.statement(node, '%sfor ' % prefix,
node.target, ' in ', node.iter, ':')
......@@ -242,6 +325,7 @@ class SourceGenerator(ExplicitNodeVisitor):
self.visit_For(node, async=True)
def visit_While(self, node):
set_precedence(node, node.test)
self.statement(node, 'while ', node.test, ':')
self.body_or_else(node)
......@@ -320,6 +404,7 @@ class SourceGenerator(ExplicitNodeVisitor):
self.conditional_write(', ', dicts[1])
def visit_Assert(self, node):
set_precedence(node, node.test, node.msg)
self.statement(node, 'assert ', node.test)
self.conditional_write(', ', node.msg)
......@@ -330,6 +415,7 @@ class SourceGenerator(ExplicitNodeVisitor):
self.statement(node, 'nonlocal ', ', '.join(node.names))
def visit_Return(self, node):
set_precedence(node, node.value)
self.statement(node, 'return')
self.conditional_write(' ', node.value)
......@@ -342,19 +428,17 @@ class SourceGenerator(ExplicitNodeVisitor):
def visit_Raise(self, node):
# XXX: Python 2.6 / 3.0 compatibility
self.statement(node, 'raise')
if self.conditional_write(' ', getattr(node, 'exc', None)):
if self.conditional_write(' ', self.get_exc(node)):
self.conditional_write(' from ', node.cause)
elif self.conditional_write(' ', getattr(node, 'type', None)):
elif self.conditional_write(' ', self.get_type(node)):
set_precedence(node, node.inst)
self.conditional_write(', ', node.inst)
self.conditional_write(', ', node.tback)
# Expressions
def visit_Attribute(self, node):
if isinstance(node.value, ast.Num):
self.write('(', node.value, ')', '.', node.attr)
else:
self.write(node.value, '.', node.attr)
self.write(node.value, '.', node.attr)
def visit_Call(self, node):
want_comma = []
......@@ -365,86 +449,135 @@ class SourceGenerator(ExplicitNodeVisitor):
else:
want_comma.append(True)
args = node.args
keywords = node.keywords
starargs = self.get_starargs(node)
kwargs = self.get_kwargs(node)
numargs = len(args) + len(keywords)
numargs += starargs is not None
numargs += kwargs is not None
p = Precedence.Comma if numargs > 1 else Precedence.call_one_arg
set_precedence(p, *args)
self.visit(node.func)
self.write('(')
for arg in node.args:
for arg in args:
self.write(write_comma, arg)
for keyword in node.keywords:
set_precedence(Precedence.Comma, *(x.value for x in keywords))
for keyword in keywords:
# a keyword.arg of None indicates dictionary unpacking
# (Python >= 3.5)
arg = keyword.arg or ''
self.write(write_comma, arg, '=' if arg else '**',keyword.value)
self.write(write_comma, arg, '=' if arg else '**', keyword.value)
# 3.5 no longer has these
self.conditional_write(write_comma, '*',
getattr(node, 'starargs', None))
self.conditional_write(write_comma, '**',
getattr(node, 'kwargs', None))
self.conditional_write(write_comma, '*', starargs)
self.conditional_write(write_comma, '**', kwargs)
self.write(')')
def visit_Name(self, node):
self.write(node.id)
def visit_Str(self, node):
self.write(repr(node.s))
result = self.result
# Cheesy way to force a flush
self.write('foo')
result.pop()
result.append(self.pretty_string(node.s, result))
def visit_Bytes(self, node):
self.write(repr(node.s))
def visit_Num(self, node):
# Hack because ** binds more closely than '-'
s = repr(node.n)
signed = s.startswith('-')
if s[signed].isalpha():
im = s[-1] == 'j' and 'j' or ''
assert s[signed:signed+3] == 'inf', s
s = '%s1e1000%s' % ('-' if signed else '', im)
if signed:
s = '(%s)' % s
self.write(s)
@enclose('()')
def visit_Num(self, node,
# constants
new=sys.version_info >= (3, 0)):
with self.delimit(node) as delimiters:
s = repr(node.n)
# Deal with infinities -- if detected, we can
# generate them with 1e1000.
signed = s.startswith('-')
if s[signed].isalpha():
im = s[-1] == 'j' and 'j' or ''
assert s[signed:signed + 3] == 'inf', s
s = '%s1e1000%s' % ('-' if signed else '', im)
self.write(s)
# The Python 2.x compiler merges a unary minus
# with a number. This is a premature optimization
# that we deal with here...
if not new and delimiters.discard:
if signed:
pow_lhs = Precedence.Pow + 1
delimiters.discard = delimiters.pp != pow_lhs
else:
op = self.get__p_op(node)
delimiters.discard = not isinstance(op, ast.USub)
def visit_Tuple(self, node):
self.comma_list(node.elts, len(node.elts) == 1)
with self.delimit(node) as delimiters:
# Two things are special about tuples:
# 1) We cannot discard the enclosing parentheses if empty
# 2) We need the trailing comma if only one item
elts = node.elts
delimiters.discard = delimiters.discard and elts
self.comma_list(elts, len(elts) == 1)
@enclose('[]')
def visit_List(self, node):
self.comma_list(node.elts)
with self.delimit('[]'):
self.comma_list(node.elts)
@enclose('{}')
def visit_Set(self, node):
self.comma_list(node.elts)
with self.delimit('{}'):
self.comma_list(node.elts)
@enclose('{}')
def visit_Dict(self, node):
for idx, (key, value) in enumerate(zip(node.keys, node.values)):
self.write(', ' if idx else '',
key if key else '',
': ' if key else '**', value)
set_precedence(Precedence.Comma, *node.values)
with self.delimit('{}'):
for idx, (key, value) in enumerate(zip(node.keys, node.values)):
self.write(', ' if idx else '',
key if key else '',
': ' if key else '**', value)
@enclose('()')
def visit_BinOp(self, node):
self.write(node.left, get_op_symbol(node.op, ' %s '), node.right)
op, left, right = node.op, node.left, node.right
with self.delimit(node, op) as delimiters:
ispow = isinstance(op, ast.Pow)
p = delimiters.p
set_precedence((Precedence.Pow + 1) if ispow else p, left)
set_precedence(Precedence.PowRHS if ispow else (p + 1), right)
self.write(left, get_op_symbol(op, ' %s '), right)
@enclose('()')
def visit_BoolOp(self, node):
op = get_op_symbol(node.op, ' %s ')
for idx, value in enumerate(node.values):
self.write(idx and op or '', value)
with self.delimit(node, node.op) as delimiters:
op = get_op_symbol(node.op, ' %s ')
set_precedence(delimiters.p + 1, *node.values)
for idx, value in enumerate(node.values):
self.write(idx and op or '', value)
@enclose('()')
def visit_Compare(self, node):
self.visit(node.left)
for op, right in zip(node.ops, node.comparators):
self.write(get_op_symbol(op, ' %s '), right)
with self.delimit(node, node.ops[0]) as delimiters:
set_precedence(delimiters.p + 1, node.left, *node.comparators)
self.visit(node.left)
for op, right in zip(node.ops, node.comparators):
self.write(get_op_symbol(op, ' %s '), right)
@enclose('()')
def visit_UnaryOp(self, node):
self.write(get_op_symbol(node.op), '(', node.operand, ')')
with self.delimit(node, node.op) as delimiters:
set_precedence(delimiters.p, node.operand)
# In Python 2.x, a unary negative of a literal
# number is merged into the number itself. This
# bit of ugliness means it is useful to know
# what the parent operation was...
node.operand._p_op = node.op
sym = get_op_symbol(node.op)
self.write(sym, ' ' if sym.isalpha() else '', node.operand)
def visit_Subscript(self, node):
set_precedence(node, node.slice)
self.write(node.value, '[', node.slice, ']')
def visit_Slice(self, node):
set_precedence(node, node.lower, node.upper, node.step)
self.conditional_write(node.lower)
self.write(':')
self.conditional_write(node.upper)
......@@ -455,62 +588,72 @@ class SourceGenerator(ExplicitNodeVisitor):
self.visit(node.step)
def visit_Index(self, node):
self.visit(node.value)
with self.delimit(node) as delimiters:
set_precedence(delimiters.p, node.value)
self.visit(node.value)
def visit_ExtSlice(self, node):
self.comma_list(node.dims, len(node.dims) == 1)
dims = node.dims
set_precedence(node, *dims)
self.comma_list(dims, len(dims) == 1)
@enclose('()')
def visit_Yield(self, node):
self.write('yield')
self.conditional_write(' ', node.value)
with self.delimit(node):
set_precedence(get_op_precedence(node) + 1, node.value)
self.write('yield')
self.conditional_write(' ', node.value)
# new for Python 3.3
@enclose('()')
def visit_YieldFrom(self, node):
self.write('yield from')
self.conditional_write(' ', node.value)
with self.delimit(node):
self.write('yield from ', node.value)
# new for Python 3.5
def visit_Await(self, node):
self.write('await ', node.value)
@enclose('()')
def visit_Lambda(self, node):
self.write('lambda ')
self.visit_arguments(node.args)
self.write(': ', node.body)
with self.delimit(node) as delimiters:
set_precedence(delimiters.p, node.body)
self.write('lambda ')
self.visit_arguments(node.args)
self.write(': ', node.body)
def visit_Ellipsis(self, node):
self.write('...')
@enclose('[]')
def visit_ListComp(self, node):
self.write(node.elt, *node.generators)
with self.delimit('[]'):
self.write(node.elt, *node.generators)
@enclose('()')
def visit_GeneratorExp(self, node):
self.write(node.elt, *node.generators)
with self.delimit(node) as delimiters:
if delimiters.pp == Precedence.call_one_arg:
delimiters.discard = True
set_precedence(Precedence.Comma, node.elt)
self.write(node.elt, *node.generators)
@enclose('{}')
def visit_SetComp(self, node):
self.write(node.elt, *node.generators)
with self.delimit('{}'):
self.write(node.elt, *node.generators)
@enclose('{}')
def visit_DictComp(self, node):
self.write(node.key, ': ', node.value, *node.generators)
with self.delimit('{}'):
self.write(node.key, ': ', node.value, *node.generators)
@enclose('()')
def visit_IfExp(self, node):
self.write(node.body, ' if ', node.test, ' else ', node.orelse)
with self.delimit(node) as delimiters:
set_precedence(delimiters.p + 1, node.body, node.test)
set_precedence(delimiters.p, node.orelse)
self.write(node.body, ' if ', node.test, ' else ', node.orelse)
def visit_Starred(self, node):
self.write('*', node.value)
@enclose('``')
def visit_Repr(self, node):
# XXX: python 2.6 only
self.visit(node.value)
with self.delimit('``'):
self.visit(node.value)
def visit_Module(self, node):
self.write(*node.body)
......@@ -526,6 +669,8 @@ class SourceGenerator(ExplicitNodeVisitor):
self.conditional_write(' as ', node.asname)
def visit_comprehension(self, node):
set_precedence(node, node.iter, *node.ifs)
set_precedence(Precedence.comprehension_target, node.target)
self.write(' for ', node.target, ' in ', node.iter)
for if_ in node.ifs:
self.write(' if ', if_)
......@@ -4,8 +4,8 @@ Part of the astor library for Python AST manipulation.
License: 3-clause BSD
Copyright 2012-2015 (c) Patrick Maupin
Copyright 2013-2015 (c) Berker Peksag
Copyright (c) 2012-2015 Patrick Maupin
Copyright (c) 2013-2015 Berker Peksag
Functions that interact with the filesystem go here.
......@@ -37,6 +37,8 @@ class CodeToAst(object):
designed to be used in code that uses this class.
"""
if not os.path.isdir(srctree):
yield os.path.split(srctree)
for srcpath, _, fnames in os.walk(srctree):
# Avoid infinite recursion for silly users
if ignore is not None and ignore in srcpath:
......@@ -52,14 +54,19 @@ class CodeToAst(object):
TODO: Handle encodings other than the default (issue #26)
"""
with open(fname, 'r') as f:
fstr = f.read()
try:
with open(fname, 'r') as f:
fstr = f.read()
except IOError:
if fname != 'stdin':
raise
sys.stdout.write('\nReading from stdin:\n\n')
fstr = sys.stdin.read()
fstr = fstr.replace('\r\n', '\n').replace('\r', '\n')
if not fstr.endswith('\n'):
fstr += '\n'
return ast.parse(fstr, filename=fname)
@staticmethod
def get_file_info(codeobj):
"""Returns the file and line number of a code object.
......
......@@ -53,10 +53,10 @@ def iter_node(node, name='', unknown=None,
def dump_tree(node, name=None, initial_indent='', indentation=' ',
maxline=120, maxmerged=80,
#Runtime optimization
iter_node=iter_node, special=ast.AST,
list=list, isinstance=isinstance, type=type, len=len):
maxline=120, maxmerged=80,
# Runtime optimization
iter_node=iter_node, special=ast.AST,
list=list, isinstance=isinstance, type=type, len=len):
"""Dumps an AST or similar structure:
- Pretty-prints with indentation
......@@ -87,9 +87,9 @@ def dump_tree(node, name=None, initial_indent='', indentation=' ',
def strip_tree(node,
#Runtime optimization
iter_node=iter_node, special=ast.AST,
list=list, isinstance=isinstance, type=type, len=len):
# Runtime optimization
iter_node=iter_node, special=ast.AST,
list=list, isinstance=isinstance, type=type, len=len):
"""Strips an AST by removing all attributes not in _fields.
Returns a set of the names of all attributes stripped.
......@@ -97,6 +97,7 @@ def strip_tree(node,
This canonicalizes two trees for comparison purposes.
"""
stripped = set()
def strip(node, indent):
unknown = set()
leaf = True
......@@ -134,3 +135,31 @@ class ExplicitNodeVisitor(ast.NodeVisitor):
method = 'visit_' + node.__class__.__name__
visitor = getattr(self, method, abort)
return visitor(node)
def allow_ast_comparison():
"""This ugly little monkey-patcher adds in a helper class
to all the AST node types. This helper class allows
eq/ne comparisons to work, so that entire trees can
be easily compared by Python's comparison machinery.
Used by the anti8 functions to compare old and new ASTs.
Could also be used by the test library.
"""
class CompareHelper(object):
def __eq__(self, other):
return type(self) == type(other) and vars(self) == vars(other)
def __ne__(self, other):
return type(self) != type(other) or vars(self) != vars(other)
for item in vars(ast).values():
if type(item) != type:
continue
if issubclass(item, ast.AST):
try:
item.__bases__ = tuple(list(item.__bases__) + [CompareHelper])
except TypeError:
pass
......@@ -4,7 +4,7 @@ Part of the astor library for Python AST manipulation.
License: 3-clause BSD
Copyright (c) 2012-2015 Patrick Maupin
Copyright (c) 2015 Patrick Maupin
This module provides data and functions for mapping
AST nodes to symbols and precedences.
......@@ -14,48 +14,91 @@ AST nodes to symbols and precedences.
import ast
op_data = """
Or or 4
And and 6
Not not 8
Eq == 10
Gt > 10
GtE >= 10
In in 10
Is is 10
NotEq != 10
Lt < 10
LtE <= 10
NotIn not in 10
IsNot is not 10
BitOr | 12
BitXor ^ 14
BitAnd & 16
LShift << 18
RShift >> 18
Add + 20
Sub - 20
Mult * 22
Div / 22
Mod % 22
FloorDiv // 22
MatMult @ 22
UAdd + 24
USub - 24
Invert ~ 24
Pow ** 26
GeneratorExp 1
Assign 1
AugAssign 0
Expr 0
Yield 1
YieldFrom 0
If 1
For 0
While 0
Return 1
Slice 1
Subscript 0
Index 1
ExtSlice 1
comprehension_target 1
Tuple 0
Comma 1
Assert 0
Raise 0
call_one_arg 1
Lambda 1
IfExp 0
comprehension 1
Or or 1
And and 1
Not not 1
Eq == 1
Gt > 0
GtE >= 0
In in 0
Is is 0
NotEq != 0
Lt < 0
LtE <= 0
NotIn not in 0
IsNot is not 0
BitOr | 1
BitXor ^ 1
BitAnd & 1
LShift << 1
RShift >> 0
Add + 1
Sub - 0
Mult * 1
Div / 0
Mod % 0
FloorDiv // 0
MatMult @ 0
PowRHS 1
Invert ~ 1
UAdd + 0
USub - 0
Pow ** 1
Num 1
"""
op_data = [x.split() for x in op_data.splitlines()]
op_data = [(x[0], ' '.join(x[1:-1]), int(x[-1])) for x in op_data if x]
op_data = [[x[0], ' '.join(x[1:-1]), int(x[-1])] for x in op_data if x]
for index in range(1, len(op_data)):
op_data[index][2] *= 2
op_data[index][2] += op_data[index - 1][2]
precedence_data = dict((getattr(ast, x, None), z) for x, y, z in op_data)
symbol_data = dict((getattr(ast, x, None), y) for x, y, z in op_data)
def get_op_symbol(obj, fmt='%s', symbol_data=symbol_data, type=type):
"""Given an AST node object, returns a string containing the symbol.
"""
return fmt % symbol_data[type(obj)]
def get_op_precedence(obj, precedence_data=precedence_data, type=type):
"""Given an AST node object, returns the precedence.
"""
return precedence_data[type(obj)]
class Precedence(object):
vars().update((x, z) for x, y, z in op_data)
highest = max(z for x, y, z in op_data) + 2
#! /usr/bin/env python
# -*- coding: utf-8 -*-
"""
Part of the astor library for Python AST manipulation.
License: 3-clause BSD
Copyright (c) 2015 Patrick Maupin
Usage:
python -m astor.rtrip [readonly] [<source>]
This utility tests round-tripping of Python source to AST
and back to source.
.. versionadded:: 0.6
If readonly is specified, then the source will be tested,
but no files will be written.
if the source is specified to be "stdin" (without quotes)
then any source entered at the command line will be compiled
into an AST, converted back to text, and then compiled to
an AST again, and the results will be displayed to stdout.
If neither readonly nor stdin is specified, then rtrip
will create a mirror directory named tmp_rtrip and will
recursively round-trip all the Python source from the source
into the tmp_rtrip dir, after compiling it and then reconstituting
it through code_gen.to_source.
If the source is not specified, the entire Python library will be used.
The purpose of rtrip is to place Python code into a canonical form.
This is useful both for functional testing of astor, and for
validating code edits.
For example, if you make manual edits for PEP8 compliance,
you can diff the rtrip output of the original code against
the rtrip output of the edited code, to insure that you
didn't make any functional changes.
For testing astor itself, it is useful to point to a big codebase,
e.g::
python -m astor.rtrip
to roundtrip the standard library.
If any round-tripped files fail to be built or to match, the
tmp_rtrip directory will also contain fname.srcdmp and fname.dstdmp,
which are textual representations of the ASTs.
Note 1:
The canonical form is only canonical for a given version of
this module and the astor toolbox. It is not guaranteed to
be stable. The only desired guarantee is that two source modules
that parse to the same AST will be converted back into the same
canonical form.
Note 2:
This tool WILL TRASH the tmp_rtrip directory (unless readonly
is specified) -- as far as it is concerned, it OWNS that directory.
Note 3: Why is it "readonly" and not "-r"? Because python -m slurps
all the thingies starting with the dash.
"""
import sys
import os
import ast
import shutil
import logging
from astor.code_gen import to_source
from astor.file_util import code_to_ast
from astor.node_util import allow_ast_comparison, dump_tree, strip_tree
dsttree = 'tmp_rtrip'
def convert(srctree, dsttree=dsttree, readonly=False, dumpall=False):
"""Walk the srctree, and convert/copy all python files
into the dsttree
"""
allow_ast_comparison()
parse_file = code_to_ast.parse_file
find_py_files = code_to_ast.find_py_files
srctree = os.path.normpath(srctree)
if not readonly:
dsttree = os.path.normpath(dsttree)
logging.info('')
logging.info('Trashing ' + dsttree)
shutil.rmtree(dsttree, True)
unknown_src_nodes = set()
unknown_dst_nodes = set()
badfiles = set()
broken = []
# TODO: When issue #26 resolved, remove UnicodeDecodeError
handled_exceptions = SyntaxError, UnicodeDecodeError
oldpath = None
allfiles = find_py_files(srctree, None if readonly else dsttree)
for srcpath, fname in allfiles:
# Create destination directory
if not readonly and srcpath != oldpath:
oldpath = srcpath
if srcpath >= srctree:
dstpath = srcpath.replace(srctree, dsttree, 1)
if not dstpath.startswith(dsttree):
raise ValueError("%s not a subdirectory of %s" %
(dstpath, dsttree))
else:
assert srctree.startswith(srcpath)
dstpath = dsttree
os.makedirs(dstpath)
srcfname = os.path.join(srcpath, fname)
logging.info('Converting %s' % srcfname)
try:
srcast = parse_file(srcfname)
except handled_exceptions:
badfiles.add(srcfname)
continue
dsttxt = to_source(srcast)
if not readonly:
dstfname = os.path.join(dstpath, fname)
try:
with open(dstfname, 'w') as f:
f.write(dsttxt)
except UnicodeEncodeError:
badfiles.add(dstfname)
# As a sanity check, make sure that ASTs themselves
# round-trip OK
try:
dstast = ast.parse(dsttxt) if readonly else parse_file(dstfname)
except SyntaxError:
dstast = []
unknown_src_nodes.update(strip_tree(srcast))
unknown_dst_nodes.update(strip_tree(dstast))
if dumpall or srcast != dstast:
srcdump = dump_tree(srcast)
dstdump = dump_tree(dstast)
bad = srcdump != dstdump
logging.warning(' calculating dump -- %s' %
('bad' if bad else 'OK'))
if bad:
broken.append(srcfname)
if dumpall or bad:
if not readonly:
try:
with open(dstfname[:-3] + '.srcdmp', 'w') as f:
f.write(srcdump)
except UnicodeEncodeError:
badfiles.add(dstfname[:-3] + '.srcdmp')
try:
with open(dstfname[:-3] + '.dstdmp', 'w') as f:
f.write(dstdump)
except UnicodeEncodeError:
badfiles.add(dstfname[:-3] + '.dstdmp')
elif dumpall:
sys.stdout.write('\n\nAST:\n\n ')
sys.stdout.write(srcdump.replace('\n', '\n '))
sys.stdout.write('\n\nDecompile:\n\n ')
sys.stdout.write(dsttxt.replace('\n', '\n '))
sys.stdout.write('\n\nNew AST:\n\n ')
sys.stdout.write('(same as old)' if dstdump == srcdump
else dstdump.replace('\n', '\n '))
sys.stdout.write('\n')
if badfiles:
logging.warning('\nFiles not processed due to syntax errors:')
for fname in sorted(badfiles):
logging.warning(' %s' % fname)
if broken:
logging.warning('\nFiles failed to round-trip to AST:')
for srcfname in broken:
logging.warning(' %s' % srcfname)
ok_to_strip = 'col_offset _precedence _use_parens lineno _p_op _pp'
ok_to_strip = set(ok_to_strip.split())
bad_nodes = (unknown_dst_nodes | unknown_src_nodes) - ok_to_strip
if bad_nodes:
logging.error('\nERROR -- UNKNOWN NODES STRIPPED: %s' % bad_nodes)
logging.info('\n')
def usage(msg):
raise SystemExit(textwrap.dedent("""
Error: %s
Usage:
python -m astor.rtrip [readonly] [<source>]
This utility tests round-tripping of Python source to AST
and back to source.
If readonly is specified, then the source will be tested,
but no files will be written.
if the source is specified to be "stdin" (without quotes)
then any source entered at the command line will be compiled
into an AST, converted back to text, and then compiled to
an AST again, and the results will be displayed to stdout.
If neither readonly nor stdin is specified, then rtrip
will create a mirror directory named tmp_rtrip and will
recursively round-trip all the Python source from the source
into the tmp_rtrip dir, after compiling it and then reconstituting
it through code_gen.to_source.
If the source is not specified, the entire Python library will be used.
""") % msg)
if __name__ == '__main__':
import textwrap
args = sys.argv[1:]
readonly = 'readonly' in args
if readonly:
args.remove('readonly')
if not args:
args = [os.path.dirname(textwrap.__file__)]
if len(args) > 1:
usage("Too many arguments")
fname, = args
dumpall = False
if not os.path.exists(fname):
dumpall = fname == 'stdin' or usage("Cannot find directory %s" % fname)
logging.basicConfig(format='%(msg)s', level=logging.INFO)
convert(fname, readonly=readonly or dumpall, dumpall=dumpall)
# -*- coding: utf-8 -*-
"""
Part of the astor library for Python AST manipulation.
License: 3-clause BSD
Copyright (c) 2015 Patrick Maupin
Pretty-print source -- post-process for the decompiler
The goals of the initial cut of this engine are:
1) Do a passable, if not PEP8, job of line-wrapping.
2) Serve as an example of an interface to the decompiler
for anybody who wants to do a better job. :)
"""
def pretty_source(source):
""" Prettify the source.
"""
return ''.join(flatten(split_lines(source)))
def flatten(source, list=list, isinstance=isinstance):
""" Deal with nested lists
"""
def flatten_iter(source):
for item in source:
if isinstance(item, list):
for item in flatten_iter(item):
yield item
else:
yield item
return flatten_iter(source)
def split_lines(source, maxline=79):
"""Split inputs according to lines.
If a line is short enough, just yield it.
Otherwise, fix it.
"""
line = []
multiline = False
count = 0
for item in source:
if item.startswith('\n'):
if line:
if count <= maxline or multiline:
yield line
else:
for item2 in wrap_line(line, maxline):
yield item2
count = 0
multiline = False
line = []
yield item
else:
line.append(item)
multiline = '\n' in item
count += len(item)
def count(group):
return sum(len(x) for x in group)
def wrap_line(line, maxline=79, count=count):
""" We have a line that is too long,
so we're going to try to wrap it.
"""
# Extract the indentation
indentation = line[0]
lenfirst = len(indentation)
indent = lenfirst - len(indentation.strip())
assert indent in (0, lenfirst)
indentation = line.pop(0) if indent else ''
# Get splittable/non-splittable groups
dgroups = list(delimiter_groups(line))
unsplittable = dgroups[::2]
splittable = dgroups[1::2]
# If the largest non-splittable group won't fit
# on a line, try to add parentheses to the line.
if max(count(x) for x in unsplittable) > maxline - indent:
line = add_parens(line, maxline, indent)
dgroups = list(delimiter_groups(line))
unsplittable = dgroups[::2]
splittable = dgroups[1::2]
# Deal with the first (always unsplittable) group, and
# then set up to deal with the remainder in pairs.
first = unsplittable[0]
yield indentation
yield first
if not splittable:
return
pos = indent + count(first)
indentation += ' '
indent += 4
if indent >= maxline/2:
maxline = maxline/2 + indent
for sg, nsg in zip(splittable, unsplittable[1:]):
if sg:
# If we already have stuff on the line and even
# the very first item won't fit, start a new line
if pos > indent and pos + len(sg[0]) > maxline:
yield '\n'
yield indentation
pos = indent
# Dump lines out of the splittable group
# until the entire thing fits
csg = count(sg)
while pos + csg > maxline:
ready, sg = split_group(sg, pos, maxline)
if ready[-1].endswith(' '):
ready[-1] = ready[-1][:-1]
yield ready
yield '\n'
yield indentation
pos = indent
csg = count(sg)
# Dump the remainder of the splittable group
if sg:
yield sg
pos += csg
# Dump the unsplittable group, optionally
# preceded by a linefeed.
cnsg = count(nsg)
if pos > indent and pos + cnsg > maxline:
yield '\n'
yield indentation
pos = indent
yield nsg
pos += cnsg
def split_group(source, pos, maxline):
""" Split a group into two subgroups. The
first will be appended to the current
line, the second will start the new line.
Note that the first group must always
contain at least one item.
The original group may be destroyed.
"""
first = []
source.reverse()
while source:
tok = source.pop()
first.append(tok)
pos += len(tok)
if source:
tok = source[-1]
allowed = (maxline + 1) if tok.endswith(' ') else (maxline - 4)
if pos + len(tok) > allowed:
break
source.reverse()
return first, source
begin_delim = set('([{')
end_delim = set(')]}')
end_delim.add('):')
def delimiter_groups(line, begin_delim=begin_delim,
end_delim=end_delim):
"""Split a line into alternating groups.
The first group cannot have a line feed inserted,
the next one can, etc.
"""
text = []
line = iter(line)
while True:
# First build and yield an unsplittable group
for item in line:
text.append(item)
if item in begin_delim:
break
if not text:
break
yield text
# Now build and yield a splittable group
level = 0
text = []
for item in line:
if item in begin_delim:
level += 1
elif item in end_delim:
level -= 1
if level < 0:
yield text
text = [item]
break
text.append(item)
else:
assert not text, text
break
statements = set(['del ', 'return', 'yield ', 'if ', 'while '])
def add_parens(line, maxline, indent, statements=statements, count=count):
"""Attempt to add parentheses around the line
in order to make it splittable.
"""
if line[0] in statements:
index = 1
if not line[0].endswith(' '):
index = 2
assert line[1] == ' '
line.insert(index, '(')
if line[-1] == ':':
line.insert(-1, ')')
else:
line.append(')')
# That was the easy stuff. Now for assignments.
groups = list(get_assign_groups(line))
if len(groups) == 1:
# So sad, too bad
return line
counts = list(count(x) for x in groups)
didwrap = False
# If the LHS is large, wrap it first
if sum(counts[:-1]) >= maxline - indent - 4:
for group in groups[:-1]:
didwrap = False # Only want to know about last group
if len(group) > 1:
group.insert(0, '(')
group.insert(-1, ')')
didwrap = True
# Might not need to wrap the RHS if wrapped the LHS
if not didwrap or counts[-1] > maxline - indent - 10:
groups[-1].insert(0, '(')
groups[-1].append(')')
return [item for group in groups for item in group]
# Assignment operators
ops = list('|^&+-*/%@~') + '<< >> // **'.split() + ['']
ops = set(' %s= ' % x for x in ops)
def get_assign_groups(line, ops=ops):
""" Split a line into groups by assignment (including
augmented assignment)
"""
group = []
for item in line:
group.append(item)
if item in ops:
yield group
group = []
yield group
# -*- coding: utf-8 -*-
"""
Part of the astor library for Python AST manipulation.
License: 3-clause BSD
Copyright (c) 2015 Patrick Maupin
Pretty-print strings for the decompiler
We either return the repr() of the string,
or try to format it as a triple-quoted string.
This is a lot harder than you would think.
This has lots of Python 2 / Python 3 ugliness.
"""
import re
import logging
try:
special_unicode = unicode
except NameError:
class special_unicode(object):
pass
try:
basestring = basestring
except NameError:
basestring = str
def _get_line(current_output):
""" Back up in the output buffer to
find the start of the current line,
and return the entire line.
"""
myline = []
index = len(current_output)
while index:
index -= 1
try:
s = str(current_output[index])
except:
raise
myline.append(s)
if '\n' in s:
break
myline = ''.join(reversed(myline))
return myline.rsplit('\n', 1)[-1]
def _properly_indented(s, current_line):
line_indent = len(current_line) - len(current_line.lstrip())
mylist = s.split('\n')[1:]
mylist = [x.rstrip() for x in mylist]
mylist = [x for x in mylist if x]
if not s:
return False
counts = [(len(x) - len(x.lstrip())) for x in mylist]
return counts and min(counts) >= line_indent
mysplit = re.compile(r'(\\|\"\"\"|\"$)').split
replacements = {'\\': '\\\\', '"""': '""\\"', '"': '\\"'}
def _prep_triple_quotes(s, mysplit=mysplit, replacements=replacements):
""" Split the string up and force-feed some replacements
to make sure it will round-trip OK
"""
s = mysplit(s)
s[1::2] = (replacements[x] for x in s[1::2])
return ''.join(s)
def pretty_string(s, current_output, min_trip_str=20, max_line=100):
"""There are a lot of reasons why we might not want to or
be able to return a triple-quoted string. We can always
punt back to the default normal string.
"""
default = repr(s)
# Punt on abnormal strings
if (isinstance(s, special_unicode) or not isinstance(s, basestring)):
return default
len_s = len(default)
current_line = _get_line(current_output)
if current_line.strip():
if len_s < min_trip_str:
return default
total_len = len(current_line) + len_s
if total_len < max_line and not _properly_indented(s, current_line):
return default
fancy = '"""%s"""' % _prep_triple_quotes(s)
# Sometimes this doesn't work. One reason is that
# the AST has no understanding of whether \r\n was
# entered that way in the string or was a cr/lf in the
# file. So we punt just so we can round-trip properly.
try:
if eval(fancy) == s and '\r' not in fancy:
return fancy
except:
pass
"""
logging.warning("***String conversion did not work\n")
#print (eval(fancy), s)
print
print (fancy, repr(s))
print
"""
return default
......@@ -7,6 +7,9 @@ License: 3-clause BSD
Copyright 2012 (c) Patrick Maupin
Copyright 2013 (c) Berker Peksag
This file contains a TreeWalk class that views a node tree
as a unified whole and allows several modes of traversal.
"""
from .node_util import iter_node
......@@ -76,9 +79,9 @@ class TreeWalk(MetaFlatten):
methods can be written. They will be called in alphabetical order.
"""
nodestack = None
def __init__(self, node=None):
self.nodestack = []
self.setup()
if node is not None:
self.walk(node)
......@@ -106,11 +109,11 @@ class TreeWalk(MetaFlatten):
"""
pre_handlers = self.pre_handlers.get
post_handlers = self.post_handlers.get
oldstack = self.nodestack
self.nodestack = nodestack = []
nodestack = self.nodestack
emptystack = len(nodestack)
append, pop = nodestack.append, nodestack.pop
append([node, name, list(iter_node(node, name + '_item')), -1])
while nodestack:
while len(nodestack) > emptystack:
node, name, subnodes, index = nodestack[-1]
if index >= len(subnodes):
handler = (post_handlers(type(node).__name__) or
......@@ -138,7 +141,6 @@ class TreeWalk(MetaFlatten):
else:
node, name = subnodes[index]
append([node, name, list(iter_node(node, name + '_item')), -1])
self.nodestack = oldstack
@property
def parent(self):
......
......@@ -208,4 +208,69 @@ Functions
get_unaryop, and get_anyop.
Command line utilities
--------------------------
rtrip
''''''
There is currently one command-line utility::
python -m astor.rtrip [readonly] [<source>]
This utility tests round-tripping of Python source to AST
and back to source.
.. versionadded:: 0.6
If readonly is specified, then the source will be tested,
but no files will be written.
if the source is specified to be "stdin" (without quotes)
then any source entered at the command line will be compiled
into an AST, converted back to text, and then compiled to
an AST again, and the results will be displayed to stdout.
If neither readonly nor stdin is specified, then rtrip
will create a mirror directory named tmp_rtrip and will
recursively round-trip all the Python source from the source
into the tmp_rtrip dir, after compiling it and then reconstituting
it through code_gen.to_source.
If the source is not specified, the entire Python library will be used.
The purpose of rtrip is to place Python code into a canonical form.
This is useful both for functional testing of astor, and for
validating code edits.
For example, if you make manual edits for PEP8 compliance,
you can diff the rtrip output of the original code against
the rtrip output of the edited code, to insure that you
didn't make any functional changes.
For testing astor itself, it is useful to point to a big codebase,
e.g::
python -m astor.rtrip
to round-trip the standard library.
If any round-tripped files fail to be built or to match, the
tmp_rtrip directory will also contain fname.srcdmp and fname.dstdmp,
which are textual representations of the ASTs.
Note 1:
The canonical form is only canonical for a given version of
this module and the astor toolbox. It is not guaranteed to
be stable. The only desired guarantee is that two source modules
that parse to the same AST will be converted back into the same
canonical form.
Note 2:
This tool WILL TRASH the tmp_rtrip directory (unless readonly
is specified) -- as far as it is concerned, it OWNS that directory.
.. _GitHub: https://github.com/berkerpeksag/astor/
......@@ -36,5 +36,5 @@ setup(
'Topic :: Software Development :: Code Generators',
'Topic :: Software Development :: Compilers',
],
keywords='ast, codegen',
keywords='ast, codegen, PEP8',
)
#! /usr/bin/env python
# -*- coding: utf-8 -*-
"""
Part of the astor library for Python AST manipulation.
License: 3-clause BSD
Copyright (c) 2015 Patrick Maupin
This module generates a lot of permutations of Python
expressions, and dumps them into a python module
all_expr_x_y.py (where x and y are the python version tuple)
as a string.
This string is later used by check_expressions.
This module takes a loooooooooong time to execute.
"""
import sys
import collections
import itertools
import textwrap
import ast
import astor
all_operators = (
# Selected special operands
'3 -3 () yield',
# operators with one parameter
'yield lambda_: not + - ~ $, yield_from',
# operators with two parameters
'or and == != > >= < <= in not_in is is_not '
'| ^ & << >> + - * / % // @ ** for$in$ $($) $[$] . '
'$,$ ',
# operators with 3 parameters
'$if$else$ $for$in$'
)
select_operators = (
# Selected special operands -- remove
# some at redundant precedence levels
'-3',
# operators with one parameter
'yield lambda_: not - ~ $,',
# operators with two parameters
'or and == in is '
'| ^ & >> - % ** for$in$ $($) . ',
# operators with 3 parameters
'$if$else$ $for$in$'
)
def get_primitives(base):
"""Attempt to return formatting strings for all operators,
and selected operands.
Here, I use the term operator loosely to describe anything
that accepts an expression and can be used in an additional
expression.
"""
operands = []
operators = []
for nparams, s in enumerate(base):
s = s.replace('%', '%%').split()
for s in (x.replace('_', ' ') for x in s):
if nparams and '$' not in s:
assert nparams in (1, 2)
s = '%s%s$' % ('$' if nparams == 2 else '', s)
assert nparams == s.count('$'), (nparams, s)
s = s.replace('$', ' %s ').strip()
# Normalize the spacing
s = s.replace(' ,', ',')
s = s.replace(' . ', '.')
s = s.replace(' [ ', '[').replace(' ]', ']')
s = s.replace(' ( ', '(').replace(' )', ')')
if nparams == 1:
s = s.replace('+ ', '+')
s = s.replace('- ', '-')
s = s.replace('~ ', '~')
if nparams:
operators.append((s, nparams))
else:
operands.append(s)
return operators, operands
def get_sub_combinations(maxop):
"""Return a dictionary of lists of combinations suitable
for recursively building expressions.
Each dictionary key is a tuple of (numops, numoperands),
where:
numops is the number of operators we
should build an expression for
numterms is the number of operands required
by the current operator.
Each list contains all permutations of the number
of operators that the recursively called function
should use for each operand.
"""
combo = collections.defaultdict(list)
for numops in range(maxop+1):
if numops:
combo[numops, 1].append((numops-1,))
for op1 in range(numops):
combo[numops, 2].append((op1, numops - op1 - 1))
for op2 in range(numops - op1):
combo[numops, 3].append((op1, op2, numops - op1 - op2 - 1))
return combo
def get_paren_combos():
"""This function returns a list of lists.
The first list is indexed by the number of operands
the current operator has.
Each sublist contains all permutations of wrapping
the operands in parentheses or not.
"""
results = [None] * 4
options = [('%s', '(%s)')]
for i in range(1, 4):
results[i] = list(itertools.product(*(i * options)))
return results
def operand_combo(expressions, operands, max_operand=13):
op_combos = []
operands = list(operands)
operands.append('%s')
for n in range(max_operand):
this_combo = []
op_combos.append(this_combo)
for i in range(n):
for op in operands:
mylist = ['%s'] * n
mylist[i] = op
this_combo.append(tuple(mylist))
for expr in expressions:
expr = expr.replace('%%', '%%%%')
for op in op_combos[expr.count('%s')]:
yield expr % op
def build(numops=2, all_operators=all_operators, use_operands=False,
# Runtime optimization
tuple=tuple):
operators, operands = get_primitives(all_operators)
combo = get_sub_combinations(numops)
paren_combos = get_paren_combos()
product = itertools.product
try:
izip = itertools.izip
except AttributeError:
izip = zip
def recurse_build(numops):
if not numops:
yield '%s'
for myop, nparams in operators:
myop = myop.replace('%%', '%%%%')
myparens = paren_combos[nparams]
# print combo[numops, nparams]
for mycombo in combo[numops, nparams]:
# print mycombo
call_again = (recurse_build(x) for x in mycombo)
for subexpr in product(*call_again):
for parens in myparens:
wrapped = tuple(x % y for (x, y)
in izip(parens, subexpr))
yield myop % wrapped
result = recurse_build(numops)
return operand_combo(result, operands) if use_operands else result
def makelib():
parse = ast.parse
dump_tree = astor.dump_tree
def default_value(): return 1000000, ''
mydict = collections.defaultdict(default_value)
allparams = [tuple('abcdefghijklmnop'[:x]) for x in range(13)]
alltxt = itertools.chain(build(1, use_operands=True),
build(2, use_operands=True),
build(3, select_operators))
yieldrepl = list(('yield %s %s' % (operator, operand),
'yield %s%s' % (operator, operand))
for operator in '+-' for operand in '(ab')
yieldrepl.append(('yield[', 'yield ['))
# alltxt = itertools.chain(build(1), build(2))
badexpr = 0
goodexpr = 0
silly = '3( 3.( 3[ 3.['.split()
for expr in alltxt:
params = allparams[expr.count('%s')]
expr %= params
try:
myast = parse(expr)
except:
badexpr += 1
continue
goodexpr += 1
key = dump_tree(myast)
expr = expr.replace(', - ', ', -')
ignore = [x for x in silly if x in expr]
if ignore:
continue
if 'yield' in expr:
for x in yieldrepl:
expr = expr.replace(*x)
mydict[key] = min(mydict[key], (len(expr), expr))
print(badexpr, goodexpr)
stuff = [x[1] for x in mydict.values()]
stuff.sort()
lineend = '\n'.encode('utf-8')
with open('all_expr_%s_%s.py' % sys.version_info[:2], 'wb') as f:
f.write(textwrap.dedent('''
# AUTOMAGICALLY GENERATED!!! DO NOT MODIFY!!
#
all_expr = """
''').encode('utf-8'))
for item in stuff:
f.write(item.encode('utf-8'))
f.write(lineend)
f.write('"""\n'.encode('utf-8'))
if __name__ == '__main__':
makelib()
#! /usr/bin/env python
# -*- coding: utf-8 -*-
"""
Part of the astor library for Python AST manipulation.
License: 3-clause BSD
Copyright (c) 2015 Patrick Maupin
This module reads the strings generated by build_expressions,
and runs them through the Python interpreter.
For strings that are suboptimal (too many spaces, etc.),
it simply dumps them to a miscompare file.
For strings that seem broken (do not parse after roundtrip)
or are maybe too compressed, it dumps information to the console.
This module does not take too long to execute; however, the
underlying build_expressions module takes forever, so this
should not be part of the automated regressions.
"""
import sys
import collections
import itertools
import textwrap
import hashlib
import ast
import astor
try:
import importlib
except ImportError:
try:
import all_expr_2_6 as mymod
except ImportError:
print("Expression list does not exist -- building")
import build_expressions
build_expressions.makelib()
print("Expression list built")
import all_expr_2_6 as mymod
else:
mymodname = 'all_expr_%s_%s' % sys.version_info[:2]
try:
mymod = importlib.import_module(mymodname)
except ImportError:
print("Expression list does not exist -- building")
import build_expressions
build_expressions.makelib()
print("Expression list built")
mymod = importlib.import_module(mymodname)
def checklib():
print("Checking expressions")
parse = ast.parse
dump_tree = astor.dump_tree
to_source = astor.to_source
with open('mismatch_%s_%s.txt' % sys.version_info[:2], 'wb') as f:
for srctxt in mymod.all_expr.strip().splitlines():
srcast = parse(srctxt)
dsttxt = to_source(srcast)
if dsttxt != srctxt:
srcdmp = dump_tree(srcast)
try:
dstast = parse(dsttxt)
except SyntaxError:
bad = True
dstdmp = 'aborted'
else:
dstdmp = dump_tree(dstast)
bad = srcdmp != dstdmp
if bad or len(dsttxt) < len(srctxt):
print(srctxt, dsttxt)
if bad:
print('****************** Original')
print(srcdmp)
print('****************** Extra Crispy')
print(dstdmp)
print('******************')
print()
print()
f.write(('%s %s\n' % (repr(srctxt),
repr(dsttxt))).encode('utf-8'))
if __name__ == '__main__':
checklib()
......@@ -3,7 +3,8 @@ Part of the astor library for Python AST manipulation
License: 3-clause BSD
Copyright 2014 (c) Berker Peksag
Copyright (c) 2014 Berker Peksag
Copyright (c) 2015 Patrick Maupin
"""
import ast
......@@ -21,6 +22,7 @@ import astor
def canonical(srctxt):
return textwrap.dedent(srctxt).strip()
class CodegenTestCase(unittest.TestCase):
def assertAstEqual(self, srctxt):
......@@ -37,7 +39,7 @@ class CodegenTestCase(unittest.TestCase):
self.assertEqual(dstdmp, srcdmp)
def assertAstEqualIfAtLeastVersion(self, source, min_should_work,
max_should_error=None):
max_should_error=None):
if max_should_error is None:
max_should_error = min_should_work[0], min_should_work[1] - 1
if sys.version_info >= min_should_work:
......@@ -52,10 +54,10 @@ class CodegenTestCase(unittest.TestCase):
which may not always be appropriate.
"""
srctxt = canonical(srctxt)
self.assertEqual(astor.to_source(ast.parse(srctxt)), srctxt)
self.assertEqual(astor.to_source(ast.parse(srctxt)).rstrip(), srctxt)
def assertAstSourceEqualIfAtLeastVersion(self, source, min_should_work,
max_should_error=None):
max_should_error=None):
if max_should_error is None:
max_should_error = min_should_work[0], min_should_work[1] - 1
if sys.version_info >= min_should_work:
......@@ -83,28 +85,37 @@ class CodegenTestCase(unittest.TestCase):
def test_try_expect(self):
source = """
try:
'spam'[10]
except IndexError:
pass"""
try:
'spam'[10]
except IndexError:
pass"""
self.assertAstEqual(source)
source = """
try:
'spam'[10]
except IndexError as exc:
sys.stdout.write(exc)"""
try:
'spam'[10]
except IndexError as exc:
sys.stdout.write(exc)"""
self.assertAstEqual(source)
source = """
try:
'spam'[10]
except IndexError as exc:
sys.stdout.write(exc)
else:
pass
finally:
pass"""
try:
'spam'[10]
except IndexError as exc:
sys.stdout.write(exc)
else:
pass
finally:
pass"""
self.assertAstEqual(source)
source = """
try:
size = len(iterable)
except (TypeError, AttributeError):
pass
else:
if n >= size:
return sorted(iterable, key=key, reverse=True)[:n]"""
self.assertAstEqual(source)
def test_del_statement(self):
......@@ -115,57 +126,57 @@ class CodegenTestCase(unittest.TestCase):
def test_arguments(self):
source = """
j = [1, 2, 3]
j = [1, 2, 3]
def test(a1, a2, b1=j, b2='123', b3={}, b4=[]):
pass"""
def test(a1, a2, b1=j, b2='123', b3={}, b4=[]):
pass"""
self.assertAstSourceEqual(source)
def test_pass_arguments_node(self):
source = canonical("""
j = [1, 2, 3]
j = [1, 2, 3]
def test(a1, a2, b1=j, b2='123', b3={}, b4=[]):
pass""")
def test(a1, a2, b1=j, b2='123', b3={}, b4=[]):
pass""")
root_node = ast.parse(source)
arguments_node = [n for n in ast.walk(root_node)
if isinstance(n, ast.arguments)][0]
self.assertEqual(astor.to_source(arguments_node),
self.assertEqual(astor.to_source(arguments_node).rstrip(),
"a1, a2, b1=j, b2='123', b3={}, b4=[]")
source = """
def call(*popenargs, timeout=None, **kwargs):
pass"""
def call(*popenargs, timeout=None, **kwargs):
pass"""
# Probably also works on < 3.4, but doesn't work on 2.7...
self.assertAstSourceEqualIfAtLeastVersion(source, (3, 4), (2, 7))
def test_matrix_multiplication(self):
for source in ("(a @ b)", "a @= b"):
self.assertAstSourceEqualIfAtLeastVersion(source, (3, 5))
self.assertAstEqualIfAtLeastVersion(source, (3, 5))
def test_multiple_call_unpackings(self):
source = """
my_function(*[1], *[2], **{'three': 3}, **{'four': 'four'})"""
my_function(*[1], *[2], **{'three': 3}, **{'four': 'four'})"""
self.assertAstSourceEqualIfAtLeastVersion(source, (3, 5))
def test_right_hand_side_dictionary_unpacking(self):
source = """
our_dict = {'a': 1, **{'b': 2, 'c': 3}}"""
our_dict = {'a': 1, **{'b': 2, 'c': 3}}"""
self.assertAstSourceEqualIfAtLeastVersion(source, (3, 5))
def test_async_def_with_for(self):
source = """
async def read_data(db):
async with connect(db) as db_cxn:
data = await db_cxn.fetch('SELECT foo FROM bar;')
async for datum in data:
if quux(datum):
return datum"""
async def read_data(db):
async with connect(db) as db_cxn:
data = await db_cxn.fetch('SELECT foo FROM bar;')
async for datum in data:
if quux(datum):
return datum"""
self.assertAstSourceEqualIfAtLeastVersion(source, (3, 5))
def test_class_definition_with_starbases_and_kwargs(self):
source = """
class TreeFactory(*[FactoryMixin, TreeBase], **{'metaclass': Foo}):
pass"""
class TreeFactory(*[FactoryMixin, TreeBase], **{'metaclass': Foo}):
pass"""
self.assertAstSourceEqualIfAtLeastVersion(source, (3, 0))
def test_yield(self):
......@@ -179,8 +190,22 @@ class CodegenTestCase(unittest.TestCase):
self.assertAstEqual(source)
source = "(yield bar)()"
self.assertAstEqual(source)
source = "return (yield 1)"
self.assertAstEqual(source)
source = "return (yield from sam())"
self.assertAstEqualIfAtLeastVersion(source, (3, 3))
source = "((yield a) for b in c)"
self.assertAstEqual(source)
source = "[(yield)]"
self.assertAstEqual(source)
source = "if (yield): pass"
self.assertAstEqual(source)
source = "if (yield from foo): pass"
self.assertAstEqualIfAtLeastVersion(source, (3, 3))
source = "(yield from (a, b))"
self.assertAstEqualIfAtLeastVersion(source, (3, 3))
source = "yield from sam()"
self.assertAstSourceEqualIfAtLeastVersion(source, (3, 3))
def test_with(self):
source = """
......@@ -197,7 +222,7 @@ class CodegenTestCase(unittest.TestCase):
with foo as bar, mary, william as bill:
pass
"""
self.assertAstSourceEqualIfAtLeastVersion(source, (3, 3), (1, 0))
self.assertAstEqualIfAtLeastVersion(source, (2, 7))
def test_inf(self):
source = """
......@@ -205,6 +230,99 @@ class CodegenTestCase(unittest.TestCase):
"""
self.assertAstEqual(source)
def test_unary(self):
source = """
-(1) + ~(2) + +(3)
"""
self.assertAstEqual(source)
def test_pow(self):
source = """
(-2) ** (-3)
"""
self.assertAstEqual(source)
source = """
(+2) ** (+3)
"""
self.assertAstEqual(source)
source = """
2 ** 3 ** 4
"""
self.assertAstEqual(source)
source = """
-2 ** -3
"""
self.assertAstEqual(source)
source = """
-2 ** -3 ** -4
"""
self.assertAstEqual(source)
source = """
-((-1) ** other._sign)
(-1) ** self._sign
"""
self.assertAstEqual(source)
def test_comprehension(self):
source = """
((x,y) for x,y in zip(a,b))
"""
self.assertAstEqual(source)
source = """
fields = [(a, _format(b)) for (a, b) in iter_fields(node)]
"""
self.assertAstEqual(source)
source = """
ra = np.fromiter(((i * 3, i * 2) for i in range(10)),
n, dtype='i8,f8')
"""
self.assertAstEqual(source)
def test_tuple_corner_cases(self):
source = """
a = ()
"""
self.assertAstEqual(source)
source = """
assert (a, b), (c, d)
"""
self.assertAstEqual(source)
source = """
return UUID(fields=(time_low, time_mid, time_hi_version,
clock_seq_hi_variant, clock_seq_low, node), version=1)
"""
self.assertAstEqual(source)
source = """
raise(os.error, ('multiple errors:', errors))
"""
self.assertAstEqual(source)
source = """
exec(expr, global_dict, local_dict)
"""
self.assertAstEqual(source)
source = """
with (a, b) as (c, d):
pass
"""
self.assertAstEqual(source)
self.assertAstEqual(source)
source = """
with (a, b) as (c, d), (e,f) as (h,g):
pass
"""
self.assertAstEqualIfAtLeastVersion(source, (2, 7))
source = """
Pxx[..., (0,-1)] = xft[..., (0,-1)]**2
"""
self.assertAstEqualIfAtLeastVersion(source, (2, 7))
source = """
responses = {
v: (v.phrase, v.description)
for v in HTTPStatus.__members__.values()
}
"""
self.assertAstEqualIfAtLeastVersion(source, (2, 7))
if __name__ == '__main__':
unittest.main()
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment