Martmists

estimate_return_type.py

Jan 21st, 2019
254
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 7.06 KB | None | 0 0
  1. import ast
  2. import inspect
  3. from typing import Callable, Union, Tuple, List, Iterable, Dict, Any
  4.  
  5. METHOD_WRAPPER = type(1..__init__)
  6. WRAPPER_DESCRIPTOR = type(int.__init__)
  7.  
  8.  
  9. def l():
  10.     return [x for x in [1, 2, 3]]
  11.  
  12.  
  13. def m(x: List[int]):
  14.     return {y(): k for k in x}
  15.  
  16.  
  17. def n():
  18.     return {"x": "y", "z": 1}
  19.  
  20.  
  21. def o():
  22.     y = [1, 2, 3, 4]
  23.     return ["a"*(x+1) for x in y]
  24.  
  25.  
  26. def p():
  27.     return b""
  28.  
  29.  
  30. def q():  # float
  31.     x = 10 / 1.0
  32.     return x
  33.  
  34.  
  35. def r():  # str
  36.     return "abc"[1]
  37.  
  38.  
  39. def s():  # List[Union[str, int]]
  40.     return ["10", 15, "abc", 1, 2, 3]
  41.  
  42.  
  43. def t():  # List[int]
  44.     return [10, 20]
  45.  
  46.  
  47. def u():  # Tuple[int, bool]
  48.     return 10, True
  49.  
  50.  
  51. def v():  # float
  52.     return 10.0
  53.  
  54.  
  55. def w():  # str
  56.     return "abc"
  57.  
  58.  
  59. def y() -> int:  # int
  60.     z = 1
  61.     return z
  62.  
  63.  
  64. def x():  # int
  65.     return y()
  66.  
  67.  
  68. def z():  # None
  69.     pass
  70.  
  71.  
  72. def help(obj):
  73.     print("==========")
  74.     print(obj)
  75.     for k in dir(obj):
  76.         if k.startswith("__"):
  77.             continue
  78.         print(k, "->", getattr(obj, k))
  79.  
  80.  
  81. class Visitor:
  82.     def __init__(self, nodes):
  83.         self.nodes = nodes
  84.         self.args = {}
  85.         self.in_comp = False
  86.  
  87.     def set_function(self, f):
  88.         self.args = f.__annotations__
  89.  
  90.     def visit(self, node):
  91.         f = "visit_" + node.__class__.__name__
  92.         # print(f)
  93.         fun = getattr(self, f, self.generic_visit)
  94.         return fun(node)
  95.  
  96.     def unique(self,  items: Iterable):
  97.         res = []
  98.         for item in items:
  99.             if item not in res:
  100.                 res.append(item)
  101.  
  102.         return res
  103.  
  104.     def visit_Dict(self, node):
  105.         keys = self.unique(map(self.visit, node.keys))
  106.         values = self.unique(map(self.visit, node.values))
  107.         if len(keys) > 1:
  108.             keys = Union[tuple(keys)]
  109.         else:
  110.             keys = keys[0]
  111.  
  112.         if len(values) > 1:
  113.             values = Union[tuple(values)]
  114.         else:
  115.             values = values[0]
  116.  
  117.         return Dict[keys, values]
  118.  
  119.     def visit_comprehension(self, node):
  120.         types = self.visit(node.iter).__args__
  121.         if len(types) > 1:
  122.             return Union[types]
  123.         else:
  124.             return types[0]
  125.  
  126.     def visit_DictComp(self, node):
  127.         self.in_comp = True
  128.         self.nodes.append(node)
  129.         # node.key and node.value are the key/value code parts
  130.         # node.generators[0] is the iter
  131.         key = self.visit(node.key)
  132.         value = self.visit(node.value)
  133.         self.in_comp = False
  134.         return Dict[key, value]
  135.  
  136.     def visit_ListComp(self, node):
  137.         self.in_comp = True
  138.         self.nodes.append(node)
  139.         # node.elt is the action happening in the listcomp
  140.         # node.generators[0] is the iter
  141.         target = self.visit(node.elt)
  142.         self.in_comp = False
  143.         return List[target]
  144.  
  145.     def visit_Bytes(self, _):
  146.         return bytes
  147.  
  148.     def visit_Subscript(self, node):
  149.         type_ = self.visit(node.value)
  150.         if isinstance(type_.__getitem__, (METHOD_WRAPPER, WRAPPER_DESCRIPTOR)):
  151.             if hasattr(type_, "__args__"):
  152.                 return type_.__args__
  153.             return type_
  154.  
  155.         return type.__getitem__.__annotations__["return"]
  156.  
  157.     def visit_BinOp(self, node):
  158.         cls = type(node.op)
  159.         fname = {
  160.             ast.Mult: "__mul__",
  161.             ast.Div: "__truediv__",
  162.             ast.Add: "__add__",
  163.             ast.Sub: "__sub__"
  164.         }[cls]
  165.         rfname = {
  166.             ast.Mult: "__rmul__",
  167.             ast.Div: "__rdiv__",
  168.             ast.Add: "__radd__",
  169.             ast.Sub: "__rsub__"
  170.         }[cls]
  171.  
  172.         left = self.visit(node.left)
  173.         right = self.visit(node.right)
  174.  
  175.         real_func = getattr(left, fname, getattr(right, rfname, None))
  176.  
  177.         if left is None or left.__module__ == "typing":
  178.             return right
  179.  
  180.         if isinstance(real_func, (METHOD_WRAPPER, WRAPPER_DESCRIPTOR)):
  181.             # builtins/C, so make some guesses
  182.             if issubclass(left, str):
  183.                 return str
  184.             if issubclass(left, list):
  185.                 return left
  186.             if issubclass(right, float):
  187.                 return float
  188.             return left
  189.  
  190.         if real_func is not None:
  191.             return real_func.__annotations__["return"]
  192.  
  193.     def visit_List(self, node):
  194.         items = [self.visit(it) for it in node.elts]
  195.         res = self.unique(items)
  196.         if len(res) > 1:
  197.             return List[Union[tuple(res)]]
  198.         else:
  199.             return List[res[0]]
  200.  
  201.     def visit_NameConstant(self, node):
  202.         return type(node.value)
  203.  
  204.     def visit_Str(self, _):
  205.         return str
  206.  
  207.     def visit_Tuple(self, node):
  208.         nodes = []
  209.         for nnode in node.elts:
  210.             if isinstance(nnode, ast.Name):
  211.                 nodes.append(nnode.id)
  212.             else:
  213.                 nodes.append(self.visit(nnode))
  214.  
  215.         return Tuple[tuple(nodes)]
  216.  
  217.     def visit_Name(self, node):
  218.         var_name = node.id
  219.         for nnode in self.nodes:
  220.             if self.in_comp:
  221.                 if isinstance(nnode, ast.ListComp):
  222.                     e = nnode.elt
  223.                     if isinstance(e, ast.Name) and e == node:
  224.                         return self.visit(nnode.generators[0])
  225.                     # Do this too if this is a sub-sub-sub-sub element in the list
  226.                 if isinstance(nnode, ast.DictComp):
  227.                     if nnode.key == node or nnode.value == node:
  228.                         return self.visit(nnode.generators[0])
  229.             elif isinstance(nnode, ast.Assign):
  230.                 targets = nnode.targets[0]
  231.                 if isinstance(targets, ast.Name):
  232.                     if targets.id == var_name:
  233.                         return self.visit(nnode.value)
  234.  
  235.                 targets = self.visit(targets)
  236.                 value = self.visit(nnode.value)
  237.                 if not isinstance(targets, list):
  238.                     targets = [targets]
  239.                 if not isinstance(value, list):
  240.                     value = [value]
  241.                 for name, value in zip(targets, value):
  242.                     if name == var_name:
  243.                         return value
  244.         return self.args.get(var_name, Any)
  245.  
  246.     def visit_Num(self, node):
  247.         var_name = node._fields[0]
  248.         return type(getattr(node, var_name))
  249.  
  250.     def visit_Call(self, node):
  251.         f = node.func
  252.         func: Callable = eval(f.id)
  253.         # TODO: functions in call scope
  254.         return func.__annotations__['return']
  255.  
  256.     def generic_visit(self, node):
  257.         raise Exception(node.__class__.__name__)
  258.  
  259.  
  260. def get_return(f):
  261.     source = inspect.getsource(f)
  262.     module = ast.parse(source)
  263.     function = module.body[0]
  264.     nodes = function.body[::-1]
  265.     return_node = nodes[0]
  266.     vis = Visitor(nodes)
  267.     vis.set_function(f)
  268.     if not isinstance(return_node, ast.Return):
  269.         return None
  270.  
  271.     ret_val = return_node.value
  272.     return vis.visit(ret_val)
  273.  
  274.  
  275. for fun in (l, m, n, o, p, q, r, s, t, u, v, w, x, y, z):
  276.     print(fun.__name__, "->", get_return(fun))
Advertisement
Add Comment
Please, Sign In to add comment