Check which optional parameters are supplied to a function call

Viewed 57

My goal is to run through all the *.py files in a directory and look at each call to a specific function test_func. This function has some optional parameters and I need to audit when the function is called with the optional parameters. My thought is to use the ast library (specifically ast.walk()).

I suppose this is a static analysis problem.

# function definition
def test_func(
    name: str,
    *,
    user: Optional['User'] = None,
    request: Optional[WebRequest] = None,
    **kwargs
) -> bool:
    pass

# somewhere in another file ...

test_func('name0')
test_func('name1', request=request)
test_func('name1')
test_func('name2', user=user)

# figure out something like below:
# name0 is never given any optional parameters
# name1 is sometimes given request
# name2 is always given user
2 Answers

Here is a POC :

import typing
from typing import Optional
class User: pass
class WebRequest: pass


# function definition
def test_func(
    name: str,
    *,
    user: Optional['User'] = None,
    request: Optional[WebRequest] = None,
    **kwargs
) -> bool:
    pass

# somewhere in another file ...

test_func('name0')
test_func('name1', request=WebRequest())
test_func('name1')
test_func('name2', user=User())

# figure out something like below:
# name0 is never given any optional parameters
# name1 is sometimes given request
# name2 is always given user


with open(__file__, "rt") as py_file:
    py_code = py_file.read()

import collections
each_call_kwargs_names_by_arg0_value: typing.Dict[str, typing.List[typing.Tuple[str, ...]]] = collections.defaultdict(list)

import ast
tree = ast.parse(py_code)
for node in ast.walk(tree):
    if isinstance(node, ast.Call):
        if hasattr(node.func, "id"):
            name = node.func.id
        elif hasattr(node.func, "attr"):
            name = node.func.attr
        elif hasattr(node.func, "value"):
            name = node.func.value.id
        else:
            raise NotImplementedError
        print(name)
        if name == "test_func":
            arg0_value = typing.cast(ast.Str, node.args[0]).s
            each_call_kwargs_names_by_arg0_value[arg0_value].append(
                tuple(keyword.arg for keyword in node.keywords)
            )

for arg0_value, each_call_kwargs_names in each_call_kwargs_names_by_arg0_value.items():
    frequency = "NEVER" if all(len(call_args) == 0 for call_args in each_call_kwargs_names) else \
                "ALWAYS" if all(len(call_args) != 0 for call_args in each_call_kwargs_names) else \
                "SOMETIMES"
    print(f"{arg0_value!r} {frequency}: {each_call_kwargs_names}")
# Output :
# 'name0' NEVER: [()]
# 'name1' SOMETIMES: [('request',), ()]
# 'name2' ALWAYS: [('user',)]

You can use a recursive generator function to traverse an ast of your Python code:

import ast
def get_calls(d, f = ['test_func']):
   if isinstance(d, ast.Call) and d.func.id in f:
      yield None if not d.args else d.args[0].value, [i.arg for i in d.keywords]
   for i in getattr(d, '_fields', []):
      vals = (m if isinstance((m:=getattr(d, i)), list) else [m])
      yield from [j for k in vals for j in get_calls(k, f = f)]

Putting it all together:

import os, collections
d = collections.defaultdict(list)
for f in os.listdir(os.getcwd()):
   if f.endswith('.py'):
      with open(f) as f:
          for a, b in get_calls(ast.parse(f.read())):
             d[a].append(b)

r = {a:{'verdict':'never' if not any(b) else 'always' if all(b) else 'sometimes', 'params':[i[0] for i in b if i]}
     for a, b in d.items()}

Output:

{'name0': {'verdict': 'never', 'params': []}, 
 'name1': {'verdict': 'sometimes', 'params': ['request']}, 
 'name2': {'verdict': 'always', 'params': ['user']}}
Related