diff --git a/CHANGELOG.md b/CHANGELOG.md index b3e7311ec..06631b477 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,8 @@ # **Upcoming release** +- Preserve positional-only and keyword-only parameter bindings, call inference, + and default-expression scope; refuse unsupported signature rewrites + (@yangfan-yf-yf) - #895 patchedast Starred and keyword now consumes their syntactically expected and ** (@lieryan) - #896 patchedast cleanup and refactoring (@lieryan) diff --git a/docs/library.rst b/docs/library.rst index 35d497337..717e09d71 100644 --- a/docs/library.rst +++ b/docs/library.rst @@ -234,6 +234,43 @@ Note, however, that the use of ``automatic_soa`` is discouraged, because it may slow down saving considerably. +Function Parameter Information +------------------------------ + +For ordinary functions and methods, ``PyFunction.get_parameters()`` includes +positional-only and keyword-only bindings. A ``ParameterName.index`` is an +object slot, not a position in the source signature. The regular slots are +ordered as positional-only, positional-or-keyword, then keyword-only; +``*args`` and ``**kwargs`` bindings follow them. Signature separators do not +occupy slots. + +``get_param_names(special_args=False)`` returns the regular names in this +order. ``get_parameter_layout()`` returns a tuple of ``(kind, name)`` pairs +for all bindings, with kinds ``posonly``, ``positional``, ``kwonly``, +``vararg``, and ``kwarg``. Use ``get_positional_param_names()`` and +``get_keyword_param_names()`` when matching positional and keyword +arguments. In particular, a keyword with the same name as a positional-only +parameter can belong to ``**kwargs`` and is not a reference to that parameter. + +``get_parameter_defaults()`` maps names to their default-expression AST +nodes. Required parameters are absent from this mapping; an explicit +``=None`` has an AST node. Defaults belong to the definition's enclosing +scope. These APIs describe bindings and do not serialize a modified signature. +They do not provide a general annotation-expression scope model; existing +same-name annotation shadowing remains a separate limitation. + +Change Signature, Introduce Parameter, Inline Parameter, and Move Method +raise ``RefactoringError`` for signatures containing positional-only or +keyword-only parameters before generating changes. Use Function refuses +keyword-only parameters. Existing positional-only Inline Method calls remain +supported, while keyword-only Inline Method calls are refused. Call inference +keeps unknown ``*args`` and ``**kwargs`` expansions conservative. + +Local to Field requires a named positional receiver. It refuses methods with +no positional parameters and static methods, while retaining instance-method +and class-method field conversions. + + Closing The Project ------------------- diff --git a/rope/base/arguments.py b/rope/base/arguments.py index 12a626816..52c962f93 100644 --- a/rope/base/arguments.py +++ b/rope/base/arguments.py @@ -10,9 +10,10 @@ class Arguments: """ - def __init__(self, args, scope): + def __init__(self, args, scope, pyfunction=None): self.args = args self.scope = scope + self.pyfunction = pyfunction self.instance = None def get_arguments(self, parameters): @@ -25,6 +26,8 @@ def get_arguments(self, parameters): return result def get_pynames(self, parameters): + if isinstance(self.pyfunction, rope.base.pyobjects.PyFunction): + return self._get_function_pynames(parameters) result = [None] * max(len(parameters), len(self.args)) for index, arg in enumerate(self.args): if isinstance(arg, ast.keyword) and arg.arg in parameters: @@ -33,6 +36,44 @@ def get_pynames(self, parameters): result[index] = self._evaluate(arg) return result + def _get_function_pynames(self, parameters): + positional = [ + name + for name in self.pyfunction.get_positional_param_names() + if name in parameters + ] + keywords = set(self.pyfunction.get_keyword_param_names()) & set(parameters) + result = dict.fromkeys(parameters) + provided = set() + position = 0 + unknown_positions = False + unknown_keywords = False + for arg in self.args: + if isinstance(arg, ast.Starred): + unknown_positions = True + elif isinstance(arg, ast.keyword): + if arg.arg is None: + unknown_keywords = True + elif arg.arg in keywords: + result[arg.arg] = self._evaluate(arg.value) + provided.add(arg.arg) + elif not unknown_positions: + if position < len(positional): + result[positional[position]] = self._evaluate(arg) + provided.add(positional[position]) + position += 1 + for name, default in self.pyfunction.get_parameter_defaults().items(): + if name not in result or name in provided: + continue + if unknown_positions and name in positional: + continue + if unknown_keywords and name in keywords: + continue + result[name] = rope.base.evaluate.eval_node( + self.pyfunction.parent.get_scope(), default + ) + return [result[name] for name in parameters] + def get_instance_pyname(self): if self.args: return self._evaluate(self.args[0]) @@ -41,15 +82,28 @@ def _evaluate(self, ast_node): return rope.base.evaluate.eval_node(self.scope, ast_node) -def create_arguments(primary, pyfunction, call_node, scope): +def create_arguments(primary, pyfunction, call_node, scope, ignore_instance=False): """A factory for creating `Arguments`""" args = list(call_node.args) args.extend(call_node.keywords) called = call_node.func - # XXX: Handle constructors - if _is_method_call(primary, pyfunction) and isinstance(called, ast.Attribute): + result = Arguments(args, scope, pyfunction) + if ignore_instance or not isinstance(called, ast.Attribute): + return result + if isinstance(pyfunction, rope.base.pyobjects.PyFunction) and primary is not None: + kind = pyfunction.get_kind() + receiver = primary.get_object() + if kind == "classmethod": + if not isinstance(receiver, rope.base.pyobjects.AbstractClass): + receiver = receiver.get_type() + return MixedArguments( + rope.base.pynames.UnboundName(receiver), result, scope + ) + if kind == "method" and _is_method_call(primary, pyfunction): + return MixedArguments(primary, result, scope) + elif _is_method_call(primary, pyfunction): args.insert(0, called.value) - return Arguments(args, scope) + return result class ObjectArguments: @@ -79,6 +133,13 @@ def __init__(self, pyname, arguments, scope): self.args = arguments def get_pynames(self, parameters): + if not parameters: + return [] + function = getattr(self.args, "pyfunction", None) + if isinstance(function, rope.base.pyobjects.PyFunction): + positional = function.get_positional_param_names() + if not positional or parameters[0] != positional[0]: + return self.args.get_pynames(parameters) return [self.pyname] + self.args.get_pynames(parameters[1:]) def get_arguments(self, parameters): diff --git a/rope/base/evaluate.py b/rope/base/evaluate.py index 4c7580cce..7f5289052 100644 --- a/rope/base/evaluate.py +++ b/rope/base/evaluate.py @@ -93,12 +93,17 @@ def get_primary_and_pyname_at( ) -> Tuple[Optional[rope.base.pynames.PyName], Optional[rope.base.pynames.PyName]]: lineno = self.lines.get_line_number(offset) holding_scope = self.module_scope.get_inner_scope_for_offset(offset) + parameter = self._get_parameter_at(holding_scope, offset) + if parameter is not None: + return (None, parameter) # function keyword parameter if self.worder.is_function_keyword_parameter(offset): keyword_name = self.worder.get_word_at(offset) pyobject = self.get_enclosing_function(offset) if isinstance(pyobject, pyobjectsdef.PyFunction): - parameter_name = pyobject.get_parameters().get(keyword_name, None) + parameter_name = None + if keyword_name in pyobject.get_keyword_param_names(): + parameter_name = pyobject.get_parameters().get(keyword_name) return (None, parameter_name) elif isinstance(pyobject, pyobjects.AbstractFunction): parameter_name = rope.base.pynames.ParameterName() @@ -130,6 +135,24 @@ def get_primary_and_pyname_at( name = self.worder.get_primary_at(offset) return eval_str2(holding_scope, name) + def _get_parameter_at(self, scope, offset): + if scope.get_kind() != "Function": + return None + function = scope.pyobject + args = function.arguments + nodes = ( + args.posonlyargs + args.args + args.kwonlyargs + [args.vararg, args.kwarg] + ) + for node in nodes: + if node is None: + continue + prefix = self.lines.get_line(node.lineno).encode("utf-8")[: node.col_offset] + start = self.lines.get_line_start(node.lineno) + len(prefix.decode("utf-8")) + _, end = self.worder.get_word_range(start) + if start <= offset < end: + return function.get_parameters()[node.arg] + return None + def get_enclosing_function(self, offset): function_parens = self.worder.find_parens_start_from_inside(offset) try: @@ -182,8 +205,16 @@ def _Call(self, node): if pyobject is None: return - def _get_returned(pyobject): - args = arguments.create_arguments(primary, pyobject, node, self.scope) + def _get_returned(pyobject, receiver=None): + args = arguments.create_arguments( + primary, + pyobject, + node, + self.scope, + ignore_instance=receiver is not None, + ) + if receiver is not None: + args = arguments.MixedArguments(receiver, args, self.scope) return pyobject.get_returned_object(args) if isinstance(pyobject, rope.base.pyobjects.AbstractClass): @@ -197,13 +228,15 @@ def _get_returned(pyobject): return pyfunction = None + receiver = None if isinstance(pyobject, rope.base.pyobjects.AbstractFunction): pyfunction = pyobject elif "__call__" in pyobject: pyfunction = pyobject["__call__"].get_object() + receiver = rope.base.pynames.UnboundName(pyobject) if pyfunction is not None: self.result = rope.base.pynames.UnboundName( - pyobject=_get_returned(pyfunction) + pyobject=_get_returned(pyfunction, receiver) ) def _Str(self, node): @@ -350,7 +383,7 @@ def _call_function(self, node, function_name, other_args=None): args = [node] if other_args: args += other_args - arguments_ = arguments.Arguments(args, self.scope) + arguments_ = arguments.Arguments(args, self.scope, called) self.result = rope.base.pynames.UnboundName( pyobject=called.get_returned_object(arguments_) ) diff --git a/rope/base/oi/objectinfo.py b/rope/base/oi/objectinfo.py index c6457044b..ca86eafff 100644 --- a/rope/base/oi/objectinfo.py +++ b/rope/base/oi/objectinfo.py @@ -1,8 +1,29 @@ +import json import warnings -from rope.base import exceptions, resourceobserver +from rope.base import exceptions, pyobjects, resourceobserver from rope.base.oi import memorydb, objectdb, transform +_CALL_LAYOUT_PREFIX = "!parameter-layout-v1:" + + +def _is_layout_key(key): + return isinstance(key, str) and key.startswith(_CALL_LAYOUT_PREFIX) + + +def _call_scope_key(pyfunction, key): + if isinstance(pyfunction, pyobjects.PyFunction): + layout = pyfunction.get_parameter_layout() + kind = pyfunction.get_kind() + if kind != "function" or any( + parameter_kind in ("posonly", "kwonly") for parameter_kind, _ in layout + ): + # Scope keys also appear as JSON object keys in the persisted database. + return _CALL_LAYOUT_PREFIX + json.dumps( + (key, kind, layout), separators=(",", ":") + ) + return key + class ObjectInfoManager: """Stores object information @@ -71,10 +92,14 @@ def get_returned(self, pyobject, args): result = self.get_exact_returned(pyobject, args) if result is not None: return result - path, key = self._get_scope(pyobject) + path, key = self._get_call_scope(pyobject) if path is None: return None for call_info in self.objectdb.get_callinfos(path, key): + if _is_layout_key(key) and len(call_info.get_parameters()) != len( + pyobject.get_param_names(special_args=False) + ): + continue returned = call_info.get_returned() if returned and returned[0] not in ("unknown", "none"): result = returned @@ -85,8 +110,10 @@ def get_returned(self, pyobject, args): return self.to_pyobject(result) def get_exact_returned(self, pyobject, args): - path, key = self._get_scope(pyobject) + path, key = self._get_call_scope(pyobject) if path is not None: + if args is None: + return None returned = self.objectdb.get_returned( path, key, self._args_to_textual(pyobject, args) ) @@ -100,7 +127,7 @@ def _args_to_textual(self, pyfunction, args): return textual_args def get_parameter_objects(self, pyobject): - path, key = self._get_scope(pyobject) + path, key = self._get_call_scope(pyobject) if path is None: return None arg_count = len(pyobject.get_param_names(special_args=False)) @@ -108,6 +135,8 @@ def get_parameter_objects(self, pyobject): parameters = [None] * arg_count for call_info in self.objectdb.get_callinfos(path, key): args = call_info.get_parameters() + if _is_layout_key(key) and len(args) != arg_count: + continue for index, arg in enumerate(args[:arg_count]): old = parameters[index] if self.validation.is_more_valid(arg, old): @@ -120,12 +149,17 @@ def get_parameter_objects(self, pyobject): return [self.to_pyobject(parameter) for parameter in parameters] def get_passed_objects(self, pyfunction, parameter_index): - path, key = self._get_scope(pyfunction) + path, key = self._get_call_scope(pyfunction) if path is None: return [] result = [] + count = len(pyfunction.get_param_names(special_args=False)) + if parameter_index >= count: + return result for call_info in self.objectdb.get_callinfos(path, key): args = call_info.get_parameters() + if _is_layout_key(key) and len(args) != count: + continue if len(args) > parameter_index: parameter = self.to_pyobject(args[parameter_index]) if parameter is not None: @@ -145,7 +179,8 @@ def doi_to_normal(textual): def function_called(self, pyfunction, params, returned=None): function_text = self.to_textual(pyfunction) - params_text = tuple(self.to_textual(param) for param in params) + count = len(pyfunction.get_param_names(special_args=False)) + params_text = tuple(self.to_textual(param) for param in params[:count]) returned_text = ("unknown",) if returned is not None: returned_text = self.to_textual(returned) @@ -164,7 +199,20 @@ def get_per_name(self, scope, name): return self.to_pyobject(result) def _save_data(self, function, args, returned=("unknown",)): - self.objectdb.add_callinfo(function[1], function[2], args, returned) + pyfunction = self.to_pyobject(function) + if not isinstance(pyfunction, pyobjects.PyFunction): + return + path, key = self._get_call_scope(pyfunction) + count = len(pyfunction.get_param_names(special_args=False)) + if path is None or (_is_layout_key(key) and len(args) < count): + return + self.objectdb.add_callinfo(path, key, args[:count], returned) + + def _get_call_scope(self, pyfunction): + path, key = self._get_scope(pyfunction) + if path is not None: + key = _call_scope_key(pyfunction, key) + return path, key def _get_scope(self, pyobject): resource = pyobject.get_module().get_resource() @@ -206,11 +254,43 @@ def is_file_valid(self, path): return self.to_pyobject.path_to_resource(path) is not None def is_scope_valid(self, path, key): + call_key = None + if _is_layout_key(key): + call_key = key + try: + encoded = json.loads(key[len(_CALL_LAYOUT_PREFIX) :]) + except (ValueError, TypeError): + return False + if not ( + isinstance(encoded, list) + and len(encoded) == 3 + and isinstance(encoded[0], str) + and encoded[1] in ("function", "method", "staticmethod", "classmethod") + and isinstance(encoded[2], list) + and all( + isinstance(parameter, list) + and len(parameter) == 2 + and parameter[0] + in ("posonly", "positional", "kwonly", "vararg", "kwarg") + and isinstance(parameter[1], str) + for parameter in encoded[2] + ) + ): + return False + key = encoded[0] + elif not isinstance(key, str): + return False if key == "": textual = ("defined", path) else: textual = ("defined", path, key) - return self.to_pyobject(textual) is not None + pyobject = self.to_pyobject(textual) + if call_key is not None: + return ( + isinstance(pyobject, pyobjects.PyFunction) + and _call_scope_key(pyobject, key) == call_key + ) + return pyobject is not None class _FileListObserver: diff --git a/rope/base/oi/runmod.py b/rope/base/oi/runmod.py index fc6bc90c9..df30da88a 100644 --- a/rope/base/oi/runmod.py +++ b/rope/base/oi/runmod.py @@ -84,7 +84,9 @@ def on_function_call(self, frame, event, arg): args = [] returned = ("unknown",) code = frame.f_code - for argname in code.co_varnames[: code.co_argcount]: + for argname in code.co_varnames[ + : code.co_argcount + code.co_kwonlyargcount + ]: try: argvalue = self._object_to_persisted_form(frame.f_locals[argname]) args.append(argvalue) diff --git a/rope/base/oi/soa.py b/rope/base/oi/soa.py index df60c6fbb..45638b686 100644 --- a/rope/base/oi/soa.py +++ b/rope/base/oi/soa.py @@ -71,7 +71,9 @@ def _Call(self, node): self._call(pyfunction, args) def _args_with_self(self, primary, self_pyname, pyfunction, node): - base_args = arguments.create_arguments(primary, pyfunction, node, self.scope) + base_args = arguments.create_arguments( + primary, pyfunction, node, self.scope, ignore_instance=True + ) return arguments.MixedArguments(self_pyname, base_args, self.scope) def _call(self, pyfunction, args): @@ -79,7 +81,7 @@ def _call(self, pyfunction, args): if self.follow is not None: before = self._parameter_objects(pyfunction) self.pycore.object_info.function_called( - pyfunction, args.get_arguments(pyfunction.get_param_names()) + pyfunction, args.get_arguments(pyfunction.get_param_names(False)) ) pyfunction._set_parameter_pyobjects(None) if self.follow is not None: diff --git a/rope/base/oi/soi.py b/rope/base/oi/soi.py index 090b10337..2968be370 100644 --- a/rope/base/oi/soi.py +++ b/rope/base/oi/soi.py @@ -49,10 +49,12 @@ def infer_parameter_objects(pyfunction): def _handle_first_parameter(pyobject, parameters): kind = pyobject.get_kind() + if not pyobject.get_positional_param_names(): + return if not parameters: - if not pyobject.get_param_names(special_args=False): - return parameters.append(pyobjects.get_unknown()) + if parameters[0] is not None and parameters[0] != pyobjects.get_unknown(): + return if kind == "method": parameters[0] = pyobjects.PyObject(pyobject.parent) if kind == "classmethod": diff --git a/rope/base/pyobjects.py b/rope/base/pyobjects.py index fbf1b34df..f2245d904 100644 --- a/rope/base/pyobjects.py +++ b/rope/base/pyobjects.py @@ -147,6 +147,12 @@ def get_doc(self): def get_param_names(self, special_args=True): return [] + def get_positional_param_names(self): + return self.get_param_names(special_args=False) + + def get_keyword_param_names(self): + return self.get_param_names(special_args=False) + def get_returned_object(self, args): return get_unknown() diff --git a/rope/base/pyobjectsdef.py b/rope/base/pyobjectsdef.py index 95bd15691..99a323af4 100644 --- a/rope/base/pyobjectsdef.py +++ b/rope/base/pyobjectsdef.py @@ -45,11 +45,13 @@ def _infer_returned(self, args=None): return rope.base.oi.soi.infer_returned_object(self, args) def _handle_special_args(self, pyobjects): - if len(pyobjects) == len(self.arguments.args): - if self.arguments.vararg: - pyobjects.append(rope.base.builtins.get_list()) - if self.arguments.kwarg: - pyobjects.append(rope.base.builtins.get_dict()) + count = len(self.get_param_names(special_args=False)) + pyobjects[:] = pyobjects[:count] + pyobjects.extend([None] * (count - len(pyobjects))) + if self.arguments.vararg: + pyobjects.append(rope.base.builtins.get_list()) + if self.arguments.kwarg: + pyobjects.append(rope.base.builtins.get_dict()) def _set_parameter_pyobjects(self, pyobjects): if pyobjects is not None: @@ -76,13 +78,53 @@ def get_name(self): return self.get_ast().name def get_param_names(self, special_args=True): - # TODO: handle tuple parameters - result = [node.arg for node in self.arguments.args if isinstance(node, ast.arg)] - if special_args: - if self.arguments.vararg: - result.append(self.arguments.vararg.arg) - if self.arguments.kwarg: - result.append(self.arguments.kwarg.arg) + return [ + name + for kind, name in self.get_parameter_layout() + if special_args or kind not in ("vararg", "kwarg") + ] + + def get_parameter_layout(self): + """Return binding kinds in the order used by parameter object slots.""" + result = [("posonly", node.arg) for node in self.arguments.posonlyargs] + result.extend(("positional", node.arg) for node in self.arguments.args) + result.extend(("kwonly", node.arg) for node in self.arguments.kwonlyargs) + if self.arguments.vararg: + result.append(("vararg", self.arguments.vararg.arg)) + if self.arguments.kwarg: + result.append(("kwarg", self.arguments.kwarg.arg)) + return tuple(result) + + def get_positional_param_names(self): + return [ + name + for kind, name in self.get_parameter_layout() + if kind in ("posonly", "positional") + ] + + def get_keyword_param_names(self): + return [ + name + for kind, name in self.get_parameter_layout() + if kind in ("positional", "kwonly") + ] + + def get_parameter_defaults(self): + positional = self.arguments.posonlyargs + self.arguments.args + defaults = self.arguments.defaults + result = { + node.arg: default + for node, default in zip( + positional[len(positional) - len(defaults) :], defaults + ) + } + result.update( + (node.arg, default) + for node, default in zip( + self.arguments.kwonlyargs, self.arguments.kw_defaults + ) + if default is not None + ) return result def get_kind(self): @@ -96,9 +138,14 @@ def get_kind(self): if isinstance(self.parent, PyClass): for decorator in self.decorators: pyname = rope.base.evaluate.eval_node(scope, decorator) - if pyname == rope.base.builtins.builtins["staticmethod"]: + if pyname is None: + continue + decorator_object = pyname.get_object() + if not isinstance(decorator_object, rope.base.builtins.BuiltinClass): + continue + if decorator_object.builtin is staticmethod: return "staticmethod" - if pyname == rope.base.builtins.builtins["classmethod"]: + if decorator_object.builtin is classmethod: return "classmethod" return "method" return "function" @@ -133,6 +180,7 @@ def __init__(self, pycore, ast_node, parent): rope.base.pyobjects.PyDefinedObject.__init__(self, pycore, ast_node, parent) self.parent = parent self._superclasses = self.get_module()._get_concluded_data() + self._receiver_attributes_collected = False def get_superclasses(self): if self._superclasses.get() is None: @@ -142,6 +190,30 @@ def get_superclasses(self): def get_name(self): return self.get_ast().name + def _get_structural_attributes(self): + attributes = super()._get_structural_attributes() + if ( + self.structural_attributes is not None + and not self._receiver_attributes_collected + ): + # Resolve decorators only after all class bindings exist and the + # structural visitor's recursion guard has been released. + self._receiver_attributes_collected = True + self.attributes.set(None) + visitor = self.visitor_class(self.pycore, self) + visitor.names = attributes + try: + for defined in self.defineds: + if ( + isinstance(defined, PyFunction) + and defined.get_kind() != "staticmethod" + ): + visitor.visit_method_attributes(defined.get_ast()) + finally: + # get_kind() can cache a view before the receiver fields exist. + self.attributes.set(None) + return attributes + def _create_concluded_attributes(self): result = {} for base in reversed(self.get_superclasses()): @@ -587,10 +659,10 @@ class _GlobalVisitor(_ScopeVisitor): class _ClassVisitor(_ScopeVisitor): - def _FunctionDef(self, node): - _ScopeVisitor._FunctionDef(self, node) - if len(node.args.args) > 0: - first = node.args.args[0] + def visit_method_attributes(self, node): + positional = node.args.posonlyargs + node.args.args + if positional: + first = positional[0] new_visitor = None if isinstance(first, ast.arg): new_visitor = _ClassInitVisitor(self, first.arg) diff --git a/rope/base/pyscopes.py b/rope/base/pyscopes.py index 565541903..fb9fc40a1 100644 --- a/rope/base/pyscopes.py +++ b/rope/base/pyscopes.py @@ -308,11 +308,31 @@ def _get_body_indents(self, scope): def get_holding_scope_for_offset(scope, offset): for inner_scope in scope.get_scopes(): if inner_scope.in_region(offset): + if isinstance(inner_scope, FunctionScope): + if _HoldingScopeFinder._is_default_offset(inner_scope, offset): + return scope return _HoldingScopeFinder.get_holding_scope_for_offset( inner_scope, offset ) return scope + @staticmethod + def _is_default_offset(scope, offset): + lines = scope.pyobject.get_module().lines + + def source_offset(lineno, column): + # AST columns count UTF-8 bytes; source offsets count characters. + prefix = lines.get_line(lineno).encode("utf-8")[:column] + return lines.get_line_start(lineno) + len(prefix.decode("utf-8")) + + defaults = scope.pyobject.get_parameter_defaults().values() + for default in defaults: + start = source_offset(default.lineno, default.col_offset) + end = source_offset(default.end_lineno, default.end_col_offset) + if start <= offset < end: + return True + return False + def find_scope_end(self, scope): if not scope.parent: return self.lines.length() diff --git a/rope/contrib/codeassist.py b/rope/contrib/codeassist.py index 32547891e..6e855e6d6 100644 --- a/rope/contrib/codeassist.py +++ b/rope/contrib/codeassist.py @@ -1,9 +1,13 @@ +import io import keyword import sys +import tokenize import warnings from rope.base import ( + ast, builtins, + codeanalyze, evaluate, exceptions, libutils, @@ -335,6 +339,13 @@ def get_default(self): Returns None if there is no default value for this param. """ + if isinstance(self._function, pyobjects.PyFunction): + default = self._function.get_parameter_defaults().get(self.argname) + if default is None: + return None + return ast.get_source_segment( + self._function.get_module().source_code, default + ) definfo = functionutils.DefinitionInfo.read(self._function) for arg, default in definfo.args_with_defaults: if self.argname == arg: @@ -526,7 +537,7 @@ def _keyword_parameters(self, pymodule, scope): pyobject = pyobject["__call__"].get_object() if isinstance(pyobject, pyobjects.AbstractFunction): param_names = [] - param_names.extend(pyobject.get_param_names(special_args=False)) + param_names.extend(pyobject.get_keyword_param_names()) result = {} for name in param_names: if name.startswith(self.starting): @@ -609,8 +620,14 @@ def get_calltip(self, pyobject, ignore_unknown=False, remove_self=False): if ignore_unknown and not isinstance(pyobject, pyobjects.PyFunction): return if isinstance(pyobject, pyobjects.AbstractFunction): - result = self._get_function_signature(pyobject, add_module=True) - if remove_self and self._is_method(pyobject): + result = self._get_function_signature( + pyobject, add_module=True, remove_self=remove_self + ) + if ( + remove_self + and self._is_method(pyobject) + and not self._has_parameter_kinds(pyobject) + ): return result.replace("(self)", "()").replace("(self, ", "(") return result @@ -663,17 +680,90 @@ def _get_super_methods(self, pyclass, name): result.extend(self._get_super_methods(super_class, name)) return result - def _get_function_signature(self, pyfunction, add_module=False): + def _get_function_signature(self, pyfunction, add_module=False, remove_self=False): location = self._location(pyfunction, add_module) if isinstance(pyfunction, pyobjects.PyFunction): - info = functionutils.DefinitionInfo.read(pyfunction) - return location + info.to_string() + if self._has_parameter_kinds(pyfunction): + signature = functionutils._get_function_signature(pyfunction) + if signature.endswith(":"): + signature = signature[:-1] + if remove_self and self._is_method(pyfunction): + signature = self._remove_self(pyfunction, signature) + return location + signature + return location + functionutils.DefinitionInfo.read(pyfunction).to_string() else: return "{}({})".format( location + pyfunction.get_name(), ", ".join(pyfunction.get_param_names()), ) + @staticmethod + def _has_parameter_kinds(pyfunction): + args = pyfunction.arguments + return bool(args.posonlyargs or args.kwonlyargs) + + @staticmethod + def _remove_self(pyfunction, signature): + args = pyfunction.arguments + positional = args.posonlyargs + args.args + if not positional or positional[0].arg != "self": + return signature + receiver = positional[0] + defaults = pyfunction.get_parameter_defaults() + last = defaults.get("self", receiver) + pymodule = pyfunction.get_module() + lines = pymodule.lines + + def offset(lineno, column): + prefix = lines.get_line(lineno).encode("utf-8")[:column] + return lines.get_line_start(lineno) + len(prefix.decode("utf-8")) + + origin = pymodule.source_code.index( + signature, lines.get_line_start(pyfunction.get_ast().lineno) + ) + start = offset(receiver.lineno, receiver.col_offset) - origin + end = offset(last.end_lineno, last.end_col_offset) - origin + signature_lines = codeanalyze.SourceLinesAdapter(signature) + + def token_offset(token): + return signature_lines.get_line_start(token.start[0]) + token.start[1] + + tokens = [ + token + for token in tokenize.generate_tokens(io.StringIO(signature).readline) + if token.type + not in (tokenize.COMMENT, tokenize.NL, tokenize.NEWLINE, tokenize.ENDMARKER) + and token_offset(token) >= start + ] + depth = 0 + boundary = None + for index, token in enumerate(tokens): + if token.string in ("(", "[", "{"): + depth += 1 + elif token.string in (")", "]", "}"): + if depth: + depth -= 1 + else: + boundary = index + break + elif token.string == "," and depth == 0 and token_offset(token) >= end: + boundary = index + break + if boundary is None: + return signature + index = boundary + (tokens[boundary].string == ",") + if ( + len(args.posonlyargs) == 1 + and index < len(tokens) + and tokens[index].string == "/" + ): + index += 1 + if index < len(tokens) and tokens[index].string == ",": + index += 1 + if index < len(tokens): + end = token_offset(tokens[index]) + return signature[:start] + signature[end:] + def _location(self, pyobject, add_module=False): location = [] parent = pyobject.parent diff --git a/rope/refactor/change_signature.py b/rope/refactor/change_signature.py index c6671b084..b2eaadea9 100644 --- a/rope/refactor/change_signature.py +++ b/rope/refactor/change_signature.py @@ -98,6 +98,7 @@ def _definfo(self): @utils.deprecated() def normalize(self): + functionutils._check_signature_parameters(self.pyname.get_object()) changer = _FunctionChangers( self.pyname.get_object(), self.get_definition_info(), [ArgumentNormalizer()] ) @@ -105,6 +106,7 @@ def normalize(self): @utils.deprecated() def remove(self, index): + functionutils._check_signature_parameters(self.pyname.get_object()) changer = _FunctionChangers( self.pyname.get_object(), self.get_definition_info(), @@ -114,6 +116,7 @@ def remove(self, index): @utils.deprecated() def add(self, index, name, default=None, value=None): + functionutils._check_signature_parameters(self.pyname.get_object()) changer = _FunctionChangers( self.pyname.get_object(), self.get_definition_info(), @@ -123,6 +126,7 @@ def add(self, index, name, default=None, value=None): @utils.deprecated() def inline_default(self, index): + functionutils._check_signature_parameters(self.pyname.get_object()) changer = _FunctionChangers( self.pyname.get_object(), self.get_definition_info(), @@ -132,6 +136,7 @@ def inline_default(self, index): @utils.deprecated() def reorder(self, new_ordering): + functionutils._check_signature_parameters(self.pyname.get_object()) changer = _FunctionChangers( self.pyname.get_object(), self.get_definition_info(), @@ -156,6 +161,7 @@ def get_changes( in the project are searched. """ + functionutils._check_signature_parameters(self.pyname.get_object()) function_changer = _FunctionChangers( self.pyname.get_object(), self._definfo(), changers ) @@ -326,10 +332,19 @@ def get_changed_module(self): for occurrence in self.occurrence_finder.find_occurrences(self.resource): if not occurrence.is_called() and not occurrence.is_defined(): continue + primary, pyname = occurrence.get_primary_and_pyname() + if pyname is not None: + pyfunction = pyname.get_object() + if ( + isinstance(pyfunction, pyobjects.PyClass) + and "__init__" in pyfunction + ): + pyfunction = pyfunction["__init__"].get_object() + if isinstance(pyfunction, pyobjects.PyFunction): + functionutils._check_signature_parameters(pyfunction) start, end = occurrence.get_primary_range() begin_parens, end_parens = word_finder.get_word_parens_range(end - 1) if occurrence.is_called(): - primary, pyname = occurrence.get_primary_and_pyname() changed_call = self.call_changer.change_call( primary, pyname, self.source[start:end_parens] ) diff --git a/rope/refactor/functionutils.py b/rope/refactor/functionutils.py index 406f102dc..c1617d177 100644 --- a/rope/refactor/functionutils.py +++ b/rope/refactor/functionutils.py @@ -1,11 +1,30 @@ import ast from typing import List, Tuple -from rope.base import pyobjects, worder +from rope.base import exceptions, pyobjects, worder from rope.base.builtins import Lambda from rope.base.codeanalyze import SourceLinesAdapter +def _get_function_signature(pyfunction): + """Read the original signature without using the legacy editable model.""" + pymodule = pyfunction.get_module() + word_finder = worder.Worder(pymodule.source_code) + start = pymodule.lines.get_line_start(pyfunction.get_ast().lineno) + if isinstance(pyfunction, Lambda): + return word_finder.get_lambda_and_args(start) + return word_finder.get_function_and_args_in_header(start) + + +def _check_signature_parameters(pyfunction): + """Refuse parameter kinds the legacy signature writer cannot preserve.""" + arguments = pyfunction.get_ast().args + if arguments.posonlyargs or arguments.kwonlyargs: + raise exceptions.RefactoringError( + "This refactoring does not support positional-only or keyword-only parameters." + ) + + class DefinitionInfo: def __init__( self, function_name, is_method, args_with_defaults, args_arg, keywords_arg @@ -59,15 +78,7 @@ def _read(pyfunction, code): @staticmethod def read(pyfunction): - pymodule = pyfunction.get_module() - word_finder = worder.Worder(pymodule.source_code) - lineno = pyfunction.get_ast().lineno - start = pymodule.lines.get_line_start(lineno) - if isinstance(pyfunction, Lambda): - call = word_finder.get_lambda_and_args(start) - else: - call = word_finder.get_function_and_args_in_header(start) - return DefinitionInfo._read(pyfunction, call) + return DefinitionInfo._read(pyfunction, _get_function_signature(pyfunction)) class CallInfo: diff --git a/rope/refactor/inline.py b/rope/refactor/inline.py index 768d30b7d..2e8056108 100644 --- a/rope/refactor/inline.py +++ b/rope/refactor/inline.py @@ -239,6 +239,10 @@ def get_kind(self): class InlineVariable(_Inliner): def __init__(self, *args, **kwds): super().__init__(*args, **kwds) + if not isinstance(self.pyname, pynames.AssignedName): + raise exceptions.RefactoringError( + "Inline variable should be performed on an assigned variable." + ) self.pymodule = self.pyname.get_definition_location()[0] self.resource = self.pymodule.get_resource() self._check_exceptional_conditions() @@ -328,8 +332,28 @@ def get_kind(self): class InlineParameter(_Inliner): def __init__(self, *args, **kwds): super().__init__(*args, **kwds) + if not isinstance(self.pyname, pynames.ParameterName): + raise exceptions.RefactoringError( + "Inline parameter should be performed on a parameter." + ) + pyfunction = self.pyname.pyfunction + functionutils._check_signature_parameters(pyfunction) + if self.name not in [ + argument.arg for argument in pyfunction.get_ast().args.args + ]: + raise exceptions.RefactoringError( + "Cannot inline the default of a list or keyword argument." + ) + definition_info = functionutils.DefinitionInfo.read(pyfunction) + indices = [ + index + for index, (name, default) in enumerate(definition_info.args_with_defaults) + if name == self.name + ] + if len(indices) != 1: + raise exceptions.RefactoringError("Cannot resolve the parameter default.") resource, offset = self._function_location() - index = self.pyname.index + index = indices[0] self.changers = [change_signature.ArgumentDefaultInliner(index)] self.signature = change_signature.ChangeSignature( self.project, resource, offset @@ -393,6 +417,11 @@ def __init__(self, project, pyfunction, body=None): self.body = sourceutils.get_body(self.pyfunction) def _get_definition_info(self): + arguments = self.pyfunction.get_ast().args + if arguments.kwonlyargs or arguments.vararg or arguments.kwarg: + raise exceptions.RefactoringError( + "Cannot inline functions with list and keyword arguments." + ) return functionutils.DefinitionInfo.read(self.pyfunction) def _get_definition_params(self): @@ -426,6 +455,13 @@ def _calculate_header(self, primary, pyname, call): call_info = functionutils.CallInfo.read( primary, pyname, self.definition_info, call ) + positional_only = { + argument.arg for argument in self.pyfunction.get_ast().args.posonlyargs + } + if any(name in positional_only for name, value in call_info.keywords): + raise exceptions.RefactoringError( + "Cannot inline a positional-only parameter passed as a keyword." + ) paramdict = self.definition_params mapping = functionutils.ArgumentMapping(self.definition_info, call_info) for param_name, value in mapping.param_dict.items(): diff --git a/rope/refactor/introduce_parameter.py b/rope/refactor/introduce_parameter.py index 93b30d0be..0072bba03 100644 --- a/rope/refactor/introduce_parameter.py +++ b/rope/refactor/introduce_parameter.py @@ -62,6 +62,7 @@ def _get_name_and_pyname(self): ) def get_changes(self, new_parameter): + functionutils._check_signature_parameters(self.pyfunction) definition_info = functionutils.DefinitionInfo.read(self.pyfunction) definition_info.args_with_defaults.append((new_parameter, self._get_primary())) collector = codeanalyze.ChangeCollector(self.resource.read()) diff --git a/rope/refactor/localtofield.py b/rope/refactor/localtofield.py index e936d5161..bfa8d1ebc 100644 --- a/rope/refactor/localtofield.py +++ b/rope/refactor/localtofield.py @@ -35,7 +35,16 @@ def _check_redefinition(self, name, function_scope): raise exceptions.RefactoringError("The field %s already exists" % name) def _get_field_name(self, pyfunction, name): - self_name = pyfunction.get_param_names()[0] + if pyfunction.get_kind() == "staticmethod": + raise exceptions.RefactoringError( + "Cannot convert a local variable to a field without a method receiver." + ) + positional = pyfunction.get_positional_param_names() + if not positional: + raise exceptions.RefactoringError( + "Cannot convert a local variable to a field without a method receiver." + ) + self_name = positional[0] new_name = self_name + "." + name return new_name diff --git a/rope/refactor/move.py b/rope/refactor/move.py index f13aca6d1..e697cb5fb 100644 --- a/rope/refactor/move.py +++ b/rope/refactor/move.py @@ -99,6 +99,7 @@ def get_changes( will be applied to all python files. """ + functionutils._check_signature_parameters(self.pyfunction) changes = ChangeSet("Moving method <%s>" % self.method_name) if resources is None: resources = self.project.get_python_files() @@ -188,6 +189,7 @@ def _get_changes_made_by_new_class(self, dest_attr, new_name): return resource, start, end, body def get_new_method(self, name): + functionutils._check_signature_parameters(self.pyfunction) return "{}\n{}".format( self._get_new_header(name), sourceutils.fix_indentation( diff --git a/rope/refactor/usefunction.py b/rope/refactor/usefunction.py index cf1330e80..6603a312b 100644 --- a/rope/refactor/usefunction.py +++ b/rope/refactor/usefunction.py @@ -48,6 +48,10 @@ def _check_returns(self): ) def get_changes(self, resources=None, task_handle=taskhandle.DEFAULT_TASK_HANDLE): + if self.pyfunction.get_ast().args.kwonlyargs: + raise exceptions.RefactoringError( + "Use function does not support keyword-only parameters." + ) if resources is None: resources = self.project.get_python_files() changes = change.ChangeSet("Using function <%s>" % self.pyfunction.get_name()) diff --git a/ropetest/parameter_kinds_signature_test.py b/ropetest/parameter_kinds_signature_test.py new file mode 100644 index 000000000..74d526cd5 --- /dev/null +++ b/ropetest/parameter_kinds_signature_test.py @@ -0,0 +1,350 @@ +"""Runtime and atomicity checks for parameter-kind signature consumers.""" + +import contextlib +import io +from textwrap import dedent + +import pytest + +from rope.base import exceptions +from rope.base.project import Project +from rope.refactor import change_signature, introduce_parameter, move, usefunction +from rope.refactor.inline import InlineVariable, create_inline +from rope.refactor.localtofield import LocalToField +from rope.refactor.method_object import MethodObject + + +@pytest.fixture +def project(tmp_path): + result = Project( + str(tmp_path), save_objectdb=False, save_history=False, automatic_soa=False + ) + yield result + result.close() + + +def module(project, source, name="case"): + result = project.root.create_file(name + ".py") + result.write(source) + return result + + +def execute(source): + output = io.StringIO() + with contextlib.redirect_stdout(output): + exec(compile(source, "", "exec"), {}) + return output.getvalue() + + +def refused(resources, operation): + sources = [resource.read() for resource in resources] + outputs = [execute(source) for source in sources] + try: + with pytest.raises(exceptions.RefactoringError): + operation() + finally: + assert [resource.read() for resource in resources] == sources + assert [execute(resource.read()) for resource in resources] == outputs + + +@pytest.mark.parametrize( + "parameters", ["value, /", "*, value", "value=1, *, required=2"] +) +@pytest.mark.parametrize( + "api", ["get_changes", "normalize", "remove", "add", "inline_default", "reorder"] +) +def test_signature_writers_refuse_new_kinds_atomically(project, parameters, api): + source = dedent(f"""\ + def f({parameters}): + return value + """) + resource = module(project, source) + other = module(project, "print('unrelated')\n", "other") + signature = change_signature.ChangeSignature(project, resource, source.index("f(")) + operations = { + "get_changes": lambda: signature.get_changes( + [change_signature.ArgumentNormalizer()] + ), + "normalize": signature.normalize, + "remove": lambda: signature.remove(0), + "add": lambda: signature.add(0, "added", "0", "0"), + "inline_default": lambda: signature.inline_default(0), + "reorder": lambda: signature.reorder([0]), + } + refused([resource, other], operations[api]) + + +@pytest.mark.parametrize("kind", ["posonly", "kwonly"]) +def test_signature_hierarchy_checks_actual_override(project, kind): + parameters = "self, value, /" if kind == "posonly" else "self, *, value" + source = ( + dedent("""\ + class Base: + def f(self, value): + return value + """) + + dedent(f"""\ + class Child(Base): + def f({parameters}): + return value + 1 + """) + ) + source += ( + "print(Child().f(value=3))\n" if kind == "kwonly" else "print(Child().f(3))\n" + ) + resource = module(project, source) + other = module(project, "print('unrelated')\n", "other") + signature = change_signature.ChangeSignature(project, resource, source.index("f(")) + refused( + [resource, other], + lambda: signature.get_changes( + [change_signature.ArgumentNormalizer()], in_hierarchy=True + ), + ) + + +@pytest.mark.parametrize("selection", ["class", "initializer", "call"]) +def test_constructor_signature_refuses_posonly(project, selection): + source = dedent("""\ + class C: + def __init__(self, /, value): + self.value = value + item = C(3) + print(item.value) + """) + resource = module(project, source) + offsets = { + "class": source.index("C:"), + "initializer": source.index("__init__"), + "call": source.index("C(3)"), + } + signature = change_signature.ChangeSignature(project, resource, offsets[selection]) + refused( + [resource], + lambda: signature.get_changes([change_signature.ArgumentNormalizer()]), + ) + + +@pytest.mark.parametrize( + "header", ["def f(value, /)", "def f(*, value)", "async def f(value, /)"] +) +def test_introduce_parameter_refuses_header_rewrite(project, header): + source = dedent(f"""\ + base = 3 + {header}: + return value + base + """) + resource = module(project, source) + operation = introduce_parameter.IntroduceParameter( + project, resource, source.rindex("base") + 1 + ) + refused([resource], lambda: operation.get_changes("added")) + + +@pytest.mark.parametrize( + "parameters", ["value=1, *, required", "value=1, /", "*values", "**values"] +) +def test_parameter_inline_refuses_unsupported_whole_signature(project, parameters): + source = dedent(f"""\ + def f({parameters}): + pass + """) + resource = module(project, source) + name = "values" if "values" in parameters else "value" + refused( + [resource], + lambda: create_inline(project, resource, source.index(name) + 1).get_changes(), + ) + + +def test_direct_variable_inliner_refuses_parameter(project): + source = dedent("""\ + def f(value): + return value + print(f(3)) + """) + resource = module(project, source) + refused( + [resource], lambda: InlineVariable(project, resource, source.index("value") + 1) + ) + + +@pytest.mark.parametrize("parameters", ["self, value, /", "self, *, value"]) +@pytest.mark.parametrize("direct", [False, True]) +def test_move_method_refuses_unsupported_header(project, parameters, direct): + source = dedent(f"""\ + class Target: + pass + class Owner: + def __init__(self): + self.dest = Target() + def act({parameters}): + return value + 1 + """) + resource = module(project, source) + operation = move.MoveMethod(project, resource, source.index("act(")) + refused( + [resource], + lambda: ( + operation.get_new_method("act") if direct else operation.get_changes("dest") + ), + ) + + +def test_use_function_refuses_kwonly_generated_calls(project): + source = dedent("""\ + def f(*, value): + return value + 1 + answer = 3 + 1 + print(answer) + """) + resource = module(project, source) + operation = usefunction.UseFunction(project, resource, source.index("f(")) + refused([resource], operation.get_changes) + + +def test_ordinary_signature_normalization_keeps_runtime(project): + source = dedent("""\ + def f(value): + return value + 1 + print(f(value=3)) + """) + resource = module(project, source) + signature = change_signature.ChangeSignature(project, resource, source.index("f(")) + project.do(signature.get_changes([change_signature.ArgumentNormalizer()])) + assert execute(source) == execute(resource.read()) == "4\n" + + +def test_ordinary_parameter_default_inline_keeps_runtime(project): + source = dedent("""\ + def f(value=1): + return value + 1 + print(f()) + """) + resource = module(project, source) + project.do( + create_inline(project, resource, source.index("value") + 1).get_changes() + ) + assert "f(1)" in resource.read() + assert execute(source) == execute(resource.read()) == "2\n" + + +def test_posonly_method_inline_keeps_existing_success(project): + source = dedent("""\ + def f(value, /, other): + return value + other + 1 + answer = f(1, 2) + print(answer) + """) + resource = module(project, source) + project.do(create_inline(project, resource, source.index("f(")).get_changes()) + assert execute(source) == execute(resource.read()) == "4\n" + + +def test_posonly_method_inline_refuses_invalid_keyword_call(project): + source = dedent("""\ + def f(value, /): + return value + try: + print(f(value=3)) + except TypeError: + print("keyword rejected") + """) + resource = module(project, source) + refused( + [resource], + lambda: create_inline(project, resource, source.index("f(")).get_changes(), + ) + + +def test_kwonly_method_inline_keeps_existing_refusal(project): + source = dedent("""\ + def f(value, *, other): + return value + other + print(f(1, other=2)) + """) + resource = module(project, source) + refused( + [resource], + lambda: create_inline(project, resource, source.index("f(")).get_changes(), + ) + + +@pytest.mark.parametrize( + "parameters, call", [("value, /", "3"), ("*, value", "value=3")] +) +def test_method_object_preserves_original_parameter_contract(project, parameters, call): + source = dedent(f"""\ + def f({parameters}): + return value + 1 + print(f({call})) + """) + resource = module(project, source) + operation = MethodObject(project, resource, source.index("f(")) + project.do(operation.get_changes(classname="Callable")) + assert f"def f({parameters}):" in resource.read() + assert execute(source) == execute(resource.read()) == "4\n" + + +def test_local_to_field_preserves_posonly_receiver(project): + source = dedent("""\ + class C: + def f(self, /, *, value): + local = value + 1 + return local + print(C().f(value=3)) + """) + resource = module(project, source) + project.do(LocalToField(project, resource, source.index("local") + 1).get_changes()) + assert "self.local" in resource.read() + assert execute(source) == execute(resource.read()) == "4\n" + + +@pytest.mark.parametrize( + "staticmethod, parameters", + [(False, "*, value"), (True, "value"), (True, "*, value")], +) +def test_local_to_field_refuses_missing_receiver(project, staticmethod, parameters): + decorator = "@staticmethod" if staticmethod else "" + source = dedent(f"""\ + class Value: + number = 3 + obj = Value() + class C: + {decorator} + def f({parameters}): + local = value.number + 1 + return local + print(C.f(value=obj)) + print(hasattr(obj, 'local')) + """) + resource = module(project, source) + assert execute(source) == "4\nFalse\n" + refused( + [resource], + lambda: LocalToField( + project, resource, source.index("local") + 1 + ).get_changes(), + ) + + +@pytest.mark.parametrize("parameters", ["cls, *, value", "cls, /, *, value"]) +def test_local_to_field_preserves_classmethod_receiver(project, parameters): + source = dedent(f"""\ + class Value: + number = 3 + obj = Value() + class C: + @classmethod + def f({parameters}): + local = value.number + 1 + return local + print(C.f(value=obj)) + print(hasattr(obj, 'local')) + print(hasattr(C, 'local')) + """) + resource = module(project, source) + project.do(LocalToField(project, resource, source.index("local") + 1).get_changes()) + assert "cls.local" in resource.read() + assert execute(source) == "4\nFalse\nFalse\n" + assert execute(resource.read()) == "4\nFalse\nTrue\n" diff --git a/ropetest/parameter_kinds_test.py b/ropetest/parameter_kinds_test.py new file mode 100644 index 000000000..df4b76a0e --- /dev/null +++ b/ropetest/parameter_kinds_test.py @@ -0,0 +1,1247 @@ +import subprocess +import sys +from textwrap import dedent + +import pytest + +from rope.base import arguments, builtins, evaluate, exceptions, pynames +from rope.base.project import Project +from rope.contrib.codeassist import code_assist, get_calltip, get_doc +from rope.refactor import change_signature, inline, rename + + +@pytest.fixture +def project(tmp_path): + project = Project( + str(tmp_path), save_objectdb=False, save_history=False, automatic_soa=False + ) + yield project + project.close() + + +def module(project, code, name="source"): + source = project.root.create_file(name + ".py") + source.write(dedent(code)) + return source + + +def execute(project, source): + result = subprocess.run( + [sys.executable, source.real_path], + cwd=project.address, + capture_output=True, + text=True, + timeout=10, + ) + assert result.returncode == 0, result.stderr + return result.stdout + + +@pytest.mark.parametrize("definition", ["def", "async def"]) +@pytest.mark.parametrize( + "parameters,call", [("target, /", "int"), ("*, target", "target=int")] +) +def test_outer_inline_preserves_parameter_binding( + project, definition, parameters, call +): + invocation = f"func({call})" + if definition == "async def": + invocation = f"asyncio.run({invocation})" + source = module( + project, + dedent(f"""\ + import asyncio + target = 42 + {definition} func({parameters}): + return target + print({invocation} is int) + print(target) + """), + ) + before = source.read() + assert execute(project, source) == "True\n42\n" + project.do( + inline.create_inline( + project, source, before.index("target =") + 1 + ).get_changes() + ) + assert source.read() != before + assert f"func({parameters})" in source.read() + assert "return target" in source.read() + assert execute(project, source) == "True\n42\n" + + +@pytest.mark.parametrize( + "parameters,call", [("target, /", "int"), ("*, target", "target=int")] +) +def test_parameter_identity_is_independent_of_outer_assignment( + project, parameters, call +): + source = module( + project, + dedent(f"""\ + target = 42 + def func({parameters}): + return target + print(func({call}) is int) + """), + ) + code = source.read() + pymodule = project.get_pymodule(source) + formal = evaluate.eval_location( + pymodule, code.index(parameters) + parameters.index("target") + 1 + ) + body = evaluate.eval_location( + pymodule, code.index("return target") + len("return ") + 1 + ) + assert isinstance(formal, pynames.ParameterName) + assert formal is body + assert formal is not pymodule["target"] + assert formal.get_definition_location() == (pymodule, 2) + assert execute(project, source) == "True\n" + with pytest.raises(exceptions.RefactoringError): + inline.InlineVariable( + project, source, code.index("return target") + len("return ") + 1 + ) + assert source.read() == code + assert execute(project, source) == "True\n" + + +def test_all_parameter_slots_keep_distinct_objects(project): + source = module( + project, + dedent("""\ + class Pos: + pass + class Normal: + pass + class Key: + pass + class Default: + pass + def func(pos, /, normal, *rest, key, default, **extras): + return pos + result = func(Pos(), Normal(), 42, default=Default(), key=Key(), extra="value") + print(isinstance(result, Pos)) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + assert pymodule["result"].get_object().get_type() is pymodule["Pos"].get_object() + function = pymodule["func"].get_object() + assert function.get_param_names(False) == ["pos", "normal", "key", "default"] + assert function.get_param_names() == [ + "pos", + "normal", + "key", + "default", + "rest", + "extras", + ] + parameters = function.get_parameters() + for index, (name, class_name) in enumerate( + [("pos", "Pos"), ("normal", "Normal"), ("key", "Key"), ("default", "Default")] + ): + assert parameters[name].index == index + assert ( + parameters[name].get_object().get_type() + is pymodule[class_name].get_object() + ) + assert parameters["rest"].index == 4 + assert isinstance(parameters["rest"].get_object().get_type(), builtins.List) + assert parameters["extras"].index == 5 + assert isinstance(parameters["extras"].get_object().get_type(), builtins.Dict) + + +def test_positional_only_keyword_does_not_replace_position_slot(project): + source = module( + project, + dedent("""\ + class Pos: + pass + class Extra: + pass + def func(target, /, **extras): + return target + result = func(Pos(), target=Extra()) + print(isinstance(result, Pos)) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + assert pymodule["result"].get_object().get_type() is pymodule["Pos"].get_object() + offset = source.read().index("target=Extra") + 1 + assert evaluate.eval_location(pymodule, offset) is None + + +def test_keyword_only_call_infers_actual_object(project): + source = module( + project, + dedent("""\ + class Value: + pass + def func(*, target): + return target + result = func(target=Value()) + print(isinstance(result, Value)) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + assert pymodule["result"].get_object().get_type() is pymodule["Value"].get_object() + + +@pytest.mark.parametrize("parameters", ["target=target, /", "*, target=target"]) +def test_default_reference_keeps_outer_identity_and_evaluation_time( + project, parameters +): + override = "func(target=int)" if parameters.startswith("*") else "func(int)" + source = module( + project, + dedent(f"""\ + events = [] + def make(): + events.append("created") + return 42 + target = make() + def func({parameters}): + return target + print(func(), func(), {override} is int, target, len(events)) + """), + ) + before = source.read() + assert execute(project, source) == "42 42 True 42 1\n" + pymodule = project.get_pymodule(source) + default_offset = before.index("target=target") + len("target=") + 1 + assert evaluate.eval_location(pymodule, default_offset) is pymodule["target"] + project.do( + rename.Rename(project, source, before.index("target =") + 1).get_changes( + "outer" + ) + ) + assert execute(project, source) == "42 42 True 42 1\n" + assert "target=outer" in source.read() + assert "return target" in source.read() + + +def test_parameter_rename_preserves_positional_only_kwargs_key(project): + source = module( + project, + dedent("""\ + target = 42 + def func(target, /, **extras): + return target, extras["target"] + result = func(int, target=str) + print(result[0] is int, result[1] is str, target) + """), + ) + before = source.read() + assert execute(project, source) == "True True 42\n" + project.do( + rename.Rename( + project, source, before.index("func(target") + len("func(") + 1 + ).get_changes("value") + ) + assert "target = 42" in source.read() + assert "target=str" in source.read() + assert 'extras["target"]' in source.read() + assert execute(project, source) == "True True 42\n" + + +def test_parameter_rename_updates_keyword_only_call(project): + source = module( + project, + dedent("""\ + target = 42 + def func(*, target): + return target + print(func(target=int) is int, target) + """), + ) + before = source.read() + assert execute(project, source) == "True 42\n" + project.do( + rename.Rename( + project, source, before.index("*, target") + len("*, ") + 1 + ).get_changes("value") + ) + assert "target = 42" in source.read() + assert "func(value=int)" in source.read() + assert execute(project, source) == "True 42\n" + + +def test_keyword_completion_uses_keyword_capable_parameters(project): + source = module( + project, + dedent("""\ + def func(pos, /, normal, *rest, key, **extras): + pass + func( + """), + ) + proposals = { + proposal.name + for proposal in code_assist( + project, source.read(), len(source.read()), resource=source + ) + } + assert {"normal=", "key="} <= proposals + assert not {"pos=", "rest=", "extras="} & proposals + + +@pytest.mark.parametrize("parameters,call", [("target=1, /", ""), ("*, target=1", "")]) +def test_unsupported_signature_rewrite_is_atomic(project, parameters, call): + source = module( + project, + dedent(f"""\ + def func({parameters}): + return target + print(func({call})) + """), + ) + before = source.read() + assert execute(project, source) == "1\n" + with pytest.raises(exceptions.RefactoringError): + changes = change_signature.ChangeSignature( + project, source, before.index("func") + 1 + ).get_changes([change_signature.ArgumentDefaultInliner(0)]) + project.do(changes) + assert source.read() == before + assert execute(project, source) == "1\n" + + +@pytest.mark.parametrize( + "decorator,parameters,call,receiver_name,receiver_is_class", + [ + ( + "", + "self, /, value, *, key", + "Child().func(Value(), key=Key())", + "self", + False, + ), + ( + "", + "self, /, value, *, key", + "Child.func(Child(), Value(), key=Key())", + "self", + False, + ), + ( + "@staticmethod", + "value, *, key", + "Child.func(Value(), key=Key())", + None, + False, + ), + ( + "@staticmethod", + "value, *, key", + "Child().func(Value(), key=Key())", + None, + False, + ), + ( + "@classmethod", + "cls, value, *, key", + "Child.func(Value(), key=Key())", + "cls", + True, + ), + ( + "@classmethod", + "cls, value, *, key", + "Child().func(Value(), key=Key())", + "cls", + True, + ), + ], +) +def test_method_call_slots_use_actual_receiver( + project, decorator, parameters, call, receiver_name, receiver_is_class +): + source = module( + project, + dedent(f"""\ + class Value: + pass + class Key: + pass + class Base: + {decorator} + def func({parameters}): + return value + class Child(Base): + pass + result = {call} + print(isinstance(result, Value)) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + assert pymodule["result"].get_object().get_type() is pymodule["Value"].get_object() + project.pycore.analyze_module(source) + function = pymodule["Base"].get_object()["func"].get_object() + bindings = function.get_parameters() + for name, expected in [("value", "Value"), ("key", "Key")]: + assert bindings[name].get_object().get_type() is pymodule[expected].get_object() + assert all( + value.get_type() is pymodule[expected].get_object() + for value in bindings[name].get_objects() + ) + if receiver_name is not None: + expected = pymodule["Child"].get_object() + receiver = bindings[receiver_name].get_object() + assert (receiver if receiver_is_class else receiver.get_type()) is expected + assert all( + (value if receiver_is_class else value.get_type()) is expected + for value in bindings[receiver_name].get_objects() + ) + + +def test_callable_attribute_adds_its_receiver_once(project): + source = module( + project, + dedent("""\ + class Value: + pass + class Key: + pass + class Callable: + def __call__(self, value, *, key): + return value + class Holder: + callback = Callable() + result = Holder().callback(Value(), key=Key()) + print(isinstance(result, Value)) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + assert pymodule["result"].get_object().get_type() is pymodule["Value"].get_object() + project.pycore.analyze_module(source) + function = pymodule["Callable"].get_object()["__call__"].get_object() + for name, expected in [("self", "Callable"), ("value", "Value"), ("key", "Key")]: + binding = function.get_parameters()[name] + assert binding.get_object().get_type() is pymodule[expected].get_object() + assert all( + value.get_type() is pymodule[expected].get_object() + for value in binding.get_objects() + ) + + +def test_defaults_follow_positional_tail_and_keyword_names(project): + source = module( + project, + dedent("""\ + class Pos: + pass + class Normal: + pass + class Key: + pass + def func(pos=Pos(), /, normal=Normal(), *rest, key=Key(), **extras): + return normal + result = func() + print(isinstance(result, Normal)) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + assert pymodule["result"].get_object().get_type() is pymodule["Normal"].get_object() + project.pycore.analyze_module(source) + for name, expected in [("pos", "Pos"), ("normal", "Normal"), ("key", "Key")]: + binding = pymodule["func"].get_object().get_parameters()[name] + assert binding.get_object().get_type() is pymodule[expected].get_object() + + +def test_keyword_completion_distinguishes_required_and_none_default(project): + source = module( + project, + dedent("""\ + def func(*, required, optional=None): + return required + print(func(required=int) is int) + """), + ) + assert execute(project, source) == "True\n" + code = source.read() + "func(" + proposals = { + proposal.name: proposal + for proposal in code_assist(project, code, len(code), resource=source) + } + assert proposals["required="].get_default() is None + assert proposals["optional="].get_default() == "None" + + +@pytest.mark.parametrize("unpacking", ["*supply()", "**supply()"]) +def test_unknown_unpacking_does_not_invent_default_bindings(project, unpacking): + source = module( + project, + dedent(f"""\ + class Pos: + pass + class Key: + pass + def supply(): + return eval("{'[]' if unpacking.startswith('*s') else '{}'}") + def func(pos=Pos(), /, *, key=Key()): + return pos + result = func({unpacking}) + print(isinstance(result, Pos)) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + function = pymodule["func"].get_object() + call = pymodule.get_ast().body[-2].value + actual = arguments.create_arguments(None, function, call, pymodule.get_scope()) + objects = actual.get_arguments(function.get_param_names(False)) + if unpacking.startswith("*s"): + assert objects[0] is None + assert objects[1].get_type() is pymodule["Key"].get_object() + else: + assert objects[0].get_type() is pymodule["Pos"].get_object() + assert objects[1] is None + + +def test_variadic_receiver_does_not_occupy_keyword_only_slot(project): + source = module( + project, + dedent("""\ + class Key: + pass + class Owner: + def func(*items, key): + return key + result = Owner().func(key=Key()) + print(isinstance(result, Key)) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + assert pymodule["result"].get_object().get_type() is pymodule["Key"].get_object() + project.pycore.analyze_module(source) + function = pymodule["Owner"].get_object()["func"].get_object() + assert ( + function.get_parameters()["key"].get_object().get_type() + is pymodule["Key"].get_object() + ) + assert isinstance( + function.get_parameters()["items"].get_object().get_type(), builtins.List + ) + + +@pytest.mark.parametrize( + "parameters,call", + [("target, /, **extras", "Value(), target=int"), ("*, target", "target=Value()")], +) +def test_dynamic_calls_record_ordinary_parameter_slots(project, parameters, call): + source = module( + project, + dedent(f"""\ + class Value: + pass + def func({parameters}): + return eval("target") + result = func({call}) + print(isinstance(result, Value)) + """), + ) + assert execute(project, source) == "True\n" + project.pycore.run_module(source).wait_process() + pymodule = project.get_pymodule(source) + function = pymodule["func"].get_object() + expected = pymodule["Value"].get_object() + assert function.get_parameters()["target"].get_object().get_type() is expected + assert function.get_parameters()["target"].get_objects() + assert all( + value.get_type() is expected + for value in function.get_parameters()["target"].get_objects() + ) + assert pymodule["result"].get_object().get_type() is expected + manager = project.pycore.object_info + path, key = manager._get_call_scope(function) + calls = list(manager.objectdb.get_callinfos(path, key)) + assert calls + for call_info in calls: + recorded = call_info.get_parameters() + assert len(recorded) == 1 + assert manager.to_pyobject(recorded[0]).get_type() is expected + + +def test_default_offsets_count_characters_after_non_ascii_prefix(project): + source = module( + project, + dedent("""\ + target = int + def func( + *, 中文="字", + value=("字", target)[1], + ): + return value + print(func() is int, func(value=str) is str) + """), + ) + before = source.read() + assert execute(project, source) == "True True\n" + pymodule = project.get_pymodule(source) + default_offset = before.index('"字", target') + len('"字", ') + 1 + assert evaluate.eval_location(pymodule, default_offset) is pymodule["target"] + project.do( + rename.Rename(project, source, before.index("target") + 1).get_changes("outer") + ) + assert execute(project, source) == "True True\n" + assert 'value=("字", outer)' in source.read() + + +@pytest.mark.parametrize("legacy", ["same_length", "short"]) +def test_changed_layout_isolates_legacy_calls_and_preserves_name_data(tmp_path, legacy): + folder = str(tmp_path) + project = Project( + folder, save_objectdb=True, save_history=False, automatic_soa=False + ) + source = module( + project, + dedent("""\ + class B: + pass + class Q: + pass + class K: + pass + def func(b, q, k): + return q + result = func(B(), Q(), K()) + print(isinstance(result, Q)) + """), + ) + if legacy == "short": + source.write( + source.read() + .replace("func(b, q, k)", "func(b)") + .replace("return q", "return b") + .replace("func(B(), Q(), K())", "func(B())") + .replace("result, Q", "result, B") + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + function = pymodule["func"].get_object() + previous_class = "Q" if legacy == "same_length" else "B" + assert ( + pymodule["result"].get_object().get_type() + is pymodule[previous_class].get_object() + ) + project.pycore.analyze_module(source) + manager = project.pycore.object_info + raw_scope = manager._get_scope(function) + assert manager._get_call_scope(function) == raw_scope + assert ( + manager.get_returned(function, None).get_type() + is pymodule[previous_class].get_object() + ) + manager.save_per_name( + function.get_scope(), "remembered", pymodule["B"].get_object() + ) + if legacy == "same_length": + replacement = ( + source.read() + .replace("func(b, q, k)", "func(b, *rest, k, q, **extras)") + .replace("return q", "return k") + .replace("func(B(), Q(), K())", "func(B(), q=Q(), k=K())") + .replace("result, Q", "result, K") + ) + else: + replacement = ( + source.read() + .replace("func(b)", "func(b, /, *, k)") + .replace("return b", "return k") + .replace("func(B())", "func(B(), k=K())") + .replace("result, B", "result, K") + ) + source.write(replacement) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + function = pymodule["func"].get_object() + assert manager.get_parameter_objects(function) is None + assert manager.get_returned(function, None) is None + assert ( + manager.get_per_name(function.get_scope(), "remembered") + is pymodule["B"].get_object() + ) + assert pymodule["result"].get_object().get_type() is pymodule["K"].get_object() + project.pycore.analyze_module(source) + manager.objectdb.validate_file(raw_scope[0]) + project.close() + reopened = Project( + folder, save_objectdb=True, save_history=False, automatic_soa=False + ) + try: + pymodule = reopened.get_pymodule(reopened.get_file("source.py")) + function = pymodule["func"].get_object() + manager = reopened.pycore.object_info + assert manager._get_call_scope(function) != manager._get_scope(function) + assert ( + manager.get_returned(function, None).get_type() + is pymodule["K"].get_object() + ) + assert ( + manager.get_per_name(function.get_scope(), "remembered") + is pymodule["B"].get_object() + ) + expected_parameters = [("b", "B"), ("k", "K")] + if legacy == "same_length": + expected_parameters.append(("q", "Q")) + for name, expected in expected_parameters: + assert ( + function.get_parameters()[name].get_object().get_type() + is pymodule[expected].get_object() + ) + finally: + reopened.close() + + +def test_nested_default_uses_enclosing_parameter_identity(project): + source = module( + project, + dedent("""\ + def outer(target, /): + def inner(*, target=target): + return target + return inner() + print(outer(int) is int, outer(str) is str) + """), + ) + before = source.read() + assert execute(project, source) == "True True\n" + pymodule = project.get_pymodule(source) + outer = pymodule["outer"].get_object() + default_offset = before.index("target=target") + len("target=") + 1 + assert ( + evaluate.eval_location(pymodule, default_offset) + is outer.get_parameters()["target"] + ) + project.do( + rename.Rename(project, source, before.index("target") + 1).get_changes("value") + ) + assert execute(project, source) == "True True\n" + assert "target=value" in source.read() + assert "return target" in source.read() + + +def test_positional_only_self_registers_instance_attributes(project): + source = module( + project, + dedent("""\ + class Value: + pass + class Owner: + def __init__(self, /, value): + self.value = value + result = Owner(Value()).value + print(isinstance(result, Value)) + """), + ) + assert execute(project, source) == "True\n" + project.pycore.analyze_module(source) + pymodule = project.get_pymodule(source) + assert pymodule["result"].get_object().get_type() is pymodule["Value"].get_object() + assert "value" in pymodule["Owner"].get_object() + + +def test_persisted_malformed_call_scopes_are_discarded(tmp_path): + folder = str(tmp_path) + project = Project( + folder, save_objectdb=True, save_history=False, automatic_soa=False + ) + source = module( + project, + dedent("""\ + class Value: + pass + def func(*, value): + return value + result = func(value=Value) + print(result is Value) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + assert pymodule["result"].get_object() is pymodule["Value"].get_object() + manager = project.pycore.object_info + function = pymodule["func"].get_object() + path, valid_key = manager._get_call_scope(function) + invalid_keys = [ + "!parameter-layout-v1:", + "!parameter-layout-v1:{broken", + "!parameter-layout-v1:[]", + '!parameter-layout-v1:["func","invalid",[]]', + '!parameter-layout-v1:["func","function",[["kwonly","previous"]]]', + ] + for key in invalid_keys: + manager.objectdb.add_pername( + path, key, "discarded", manager.to_textual(pymodule["Value"].get_object()) + ) + manager.save_per_name( + function.get_scope(), "remembered", pymodule["Value"].get_object() + ) + project.close() + reopened = Project( + folder, save_objectdb=True, save_history=False, automatic_soa=False + ) + try: + manager = reopened.pycore.object_info + assert set(invalid_keys) <= set(manager.objectdb.files[path]) + manager.objectdb.validate_file(path) + assert not set(invalid_keys) & set(manager.objectdb.files[path]) + assert valid_key in manager.objectdb.files[path] + pymodule = reopened.get_pymodule(reopened.get_file("source.py")) + function = pymodule["func"].get_object() + assert manager.get_returned(function, None) is pymodule["Value"].get_object() + assert ( + manager.get_per_name(function.get_scope(), "remembered") + is pymodule["Value"].get_object() + ) + finally: + reopened.close() + + +def test_ordinary_method_kinds_and_callable_use_correct_receivers(project): + source = module( + project, + dedent("""\ + class Value: + pass + class Base: + def normal(self, value): + return value + @staticmethod + def static(value): + return value + @classmethod + def class_method(cls, value): + return cls + class Child(Base): + pass + class Callable: + def __call__(self, value): + return value + class Holder: + callback = Callable() + normal_result = Child().normal(Value()) + static_result = Child().static(Value()) + class_result = Child.class_method(Value()) + callable_result = Holder().callback(Value()) + print(isinstance(normal_result, Value), isinstance(static_result, Value), class_result is Child, isinstance(callable_result, Value)) + """), + ) + assert execute(project, source) == "True True True True\n" + pymodule = project.get_pymodule(source) + for result in ["normal_result", "static_result", "callable_result"]: + assert ( + pymodule[result].get_object().get_type() is pymodule["Value"].get_object() + ) + assert pymodule["class_result"].get_object() is pymodule["Child"].get_object() + project.pycore.analyze_module(source) + manager = project.pycore.object_info + for class_name, method, receiver, receiver_type in [ + ("Base", "normal", "self", "Child"), + ("Base", "static", None, None), + ("Base", "class_method", "cls", "Child"), + ("Callable", "__call__", "self", "Callable"), + ]: + function = pymodule[class_name].get_object()[method].get_object() + assert manager._get_call_scope(function) != manager._get_scope(function) + for passed in function.get_parameters()["value"].get_objects(): + assert passed.get_type() is pymodule["Value"].get_object() + if receiver is not None: + for passed in function.get_parameters()[receiver].get_objects(): + actual = passed if receiver == "cls" else passed.get_type() + assert actual is pymodule[receiver_type].get_object() + + +@pytest.mark.parametrize("selected", ["formal", "body"]) +def test_defaulted_positional_parameter_rename_keeps_its_declaration(project, selected): + source = module( + project, + dedent("""\ + target = int + def func(target=target, /, **extras): + return target, extras["target"] + result = func(str, target=bytes) + print(result[0] is str, result[1] is bytes, target is int) + """), + ) + before = source.read() + assert execute(project, source) == "True True True\n" + pymodule = project.get_pymodule(source) + function = pymodule["func"].get_object() + parameter = function.get_parameters()["target"] + formal_offset = before.index("target=target") + for delta in [0, 1, 5]: + assert evaluate.eval_location(pymodule, formal_offset + delta) is parameter + body_offset = before.index("return target") + len("return ") + 1 + assert evaluate.eval_location(pymodule, body_offset) is parameter + assert ( + evaluate.eval_location(pymodule, formal_offset + len("target=") + 1) + is pymodule["target"] + ) + assert evaluate.eval_location(pymodule, before.index("target=bytes") + 1) is None + offset = formal_offset if selected == "formal" else body_offset + project.do(rename.Rename(project, source, offset).get_changes("local")) + assert execute(project, source) == "True True True\n" + assert "func(local=target, /, **extras)" in source.read() + assert "target=bytes" in source.read() + + +def test_default_and_same_line_body_call_keywords_keep_callee_identity(project): + source = module( + project, + dedent("""\ + def build(*, target): + return target + def func(target=build(target=int), /): return build(target=target) + print(func() is int, func(str) is str) + """), + ) + before = source.read() + assert execute(project, source) == "True True\n" + pymodule = project.get_pymodule(source) + parameter = pymodule["build"].get_object().get_parameters()["target"] + for keyword in ["target=int", "target=target)"]: + assert evaluate.eval_location(pymodule, before.index(keyword) + 1) is parameter + project.do( + rename.Rename( + project, source, before.index("*, target") + len("*, ") + 1 + ).get_changes("value") + ) + assert execute(project, source) == "True True\n" + assert ( + "func(target=build(value=int), /): return build(value=target)" in source.read() + ) + + +@pytest.mark.parametrize( + "parameters,display,without_self", + [ + ('self: "C", value: int = 1', "self, value=1", "value=1"), + ('self: "C", /', 'self: "C", /', ""), + ( + 'self: "C", /, value: int = 1', + 'self: "C", /, value: int = 1', + "value: int = 1", + ), + ( + 'self: "C", value: int = 1, /, *, key: str = "ok"', + 'self: "C", value: int = 1, /, *, key: str = "ok"', + 'value: int = 1, /, *, key: str = "ok"', + ), + ( + 'self: "C", *, value: int = 1', + 'self: "C", *, value: int = 1', + "*, value: int = 1", + ), + ], +) +def test_calltip_preserves_ordinary_display_and_parameter_separators( + project, parameters, display, without_self +): + result = "value" if "value" in parameters else "1" + source = module( + project, + dedent(f'''\ + class C: + def f({parameters}): + """Details.""" + return {result} + print(C().f()) + '''), + ) + code = source.read() + assert execute(project, source) == "1\n" + offset = code.rindex(".f(") + 1 + assert ( + get_calltip(project, code, offset, resource=source) == f"source.C.f({display})" + ) + assert ( + get_calltip(project, code, offset, resource=source, remove_self=True) + == f"source.C.f({without_self})" + ) + doc = get_doc(project, code, offset, resource=source) + assert doc.splitlines()[0] == f"C.f({display}):" + + +@pytest.mark.parametrize( + "parameters", + [ + 'self: "C"=((None)), /, *, payload=2', + 'self: ("C")=((lambda a, b: None)(None, None)), /, *, payload=2', + ], +) +def test_calltip_removes_complete_parenthesized_receiver(project, parameters): + source = module( + project, + dedent(f'''\ + class C: + def f({parameters}): + """Details.""" + return payload + print(C().f()) + '''), + ) + code = source.read() + assert execute(project, source) == "2\n" + offset = code.rindex(".f(") + 1 + assert ( + get_calltip(project, code, offset, resource=source) + == f"source.C.f({parameters})" + ) + assert ( + get_calltip(project, code, offset, resource=source, remove_self=True) + == "source.C.f(*, payload=2)" + ) + assert ( + get_doc(project, code, offset, resource=source).splitlines()[0] + == f"C.f({parameters}):" + ) + + +@pytest.mark.parametrize( + "declaration,decorator,parameters", + [ + ("", "staticmethod", "value, /"), + ("", "staticmethod", "value"), + ("static = staticmethod", "static", "value, /"), + ("from builtins import staticmethod as static", "static", "value, /"), + ("first = staticmethod; static = first", "static", "value, /"), + ("import builtins as namespace", "namespace.staticmethod", "value, /"), + ], +) +def test_static_parameter_attributes_do_not_belong_to_owner( + project, declaration, decorator, parameters +): + source = module( + project, + dedent(f"""\ + class Value: + pass + class Owner: + {declaration} + @{decorator} + def assign({parameters}): + value.field = 42 + pass + value = Value() + Owner.assign(value) + print(value.field) + try: + print(Owner.field) + except AttributeError: + print("absent") + """), + ) + before = source.read() + assert execute(project, source) == "42\nabsent\n" + owner = project.get_pymodule(source)["Owner"].get_object() + owner.get_scope().get_defined_names() + assert owner["assign"].get_object().get_kind() == "staticmethod" + assert "field" not in owner.get_attributes() + offset = before.index("print(Owner.field)") + len("print(Owner.") + 1 + with pytest.raises(exceptions.RefactoringError): + inline.create_inline(project, source, offset).get_changes() + assert source.read() == before + assert execute(project, source) == "42\nabsent\n" + + +@pytest.mark.parametrize( + "declaration,decorator,receiver,call,kind", + [ + ("", "", "self", "owner.assign()", "method"), + ("", "@classmethod", "cls", "Owner.assign()", "classmethod"), + ( + "class_method = classmethod", + "@class_method", + "cls", + "Owner.assign()", + "classmethod", + ), + ( + "import builtins as namespace", + "@namespace.classmethod", + "cls", + "Owner.assign()", + "classmethod", + ), + ( + "def staticmethod(func): return func", + "@staticmethod", + "self", + "owner.assign()", + "method", + ), + ], +) +def test_receiver_attributes_survive_kind_resolution_and_early_scope_access( + project, declaration, decorator, receiver, call, kind +): + inspected = "Owner" if receiver == "cls" else "owner" + source = module( + project, + dedent(f"""\ + class Owner: + {declaration} + {decorator} + def assign({receiver}, /): + {receiver}.field = 42 + owner = Owner() + {call} + print({inspected}.field) + """), + ) + assert execute(project, source) == "42\n" + owner = project.get_pymodule(source)["Owner"].get_object() + owner.get_scope().get_defined_names() + assert owner["assign"].get_object().get_kind() == kind + assert "field" in owner.get_attributes() + assert ( + owner["field"].get_object().get_type() is builtins.builtins["int"].get_object() + ) + assert execute(project, source) == "42\n" + + +def test_explicit_unknown_argument_does_not_use_default(project): + source = module( + project, + dedent("""\ + def supply(): + return eval("str") + def func(*, value=int): + return value + result = func(value=supply()) + print(result is str, func() is int) + """), + ) + assert execute(project, source) == "True True\n" + pymodule = project.get_pymodule(source) + function = pymodule["func"].get_object() + call = pymodule.get_ast().body[-2].value + actual = arguments.create_arguments(None, function, call, pymodule.get_scope()) + value = actual.get_arguments(function.get_param_names(False))[0] + assert value is None or value is not builtins.builtins["int"].get_object() + assert pymodule["result"].get_object() is not builtins.builtins["int"].get_object() + + +def test_variadic_method_with_no_ordinary_slots_keeps_receiver_in_items(project): + source = module( + project, + dedent("""\ + class Owner: + def func(*items): + return items[0] + owner = Owner() + result = owner.func() + print(result is owner) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + function = pymodule["Owner"].get_object()["func"].get_object() + actual = arguments.create_arguments( + pymodule["owner"], + function, + pymodule.get_ast().body[-2].value, + pymodule.get_scope(), + ) + assert function.get_param_names(False) == [] + assert actual.get_arguments(function.get_param_names(False)) == [] + assert isinstance( + function.get_parameters()["items"].get_object().get_type(), builtins.List + ) + assert execute(project, source) == "True\n" + + +@pytest.mark.parametrize("record_size", [1, 3]) +def test_persisted_wrong_size_calls_do_not_poison_readers_or_fallback( + tmp_path, record_size +): + folder = str(tmp_path) + project = Project( + folder, save_objectdb=True, save_history=False, automatic_soa=False + ) + source = module( + project, + dedent("""\ + class Pos: + pass + class Key: + pass + class Wrong: + pass + def func(pos, /, *, key): + return key + print(func(Pos, key=Key) is Key) + """), + ) + assert execute(project, source) == "True\n" + pymodule = project.get_pymodule(source) + function = pymodule["func"].get_object() + manager = project.pycore.object_info + path, key = manager._get_call_scope(function) + wrong = manager.to_textual(pymodule["Wrong"].get_object()) + manager.objectdb.add_callinfo(path, key, (wrong,) * record_size, wrong) + project.close() + reopened = Project( + folder, save_objectdb=True, save_history=False, automatic_soa=False + ) + try: + source = reopened.get_file("source.py") + pymodule = reopened.get_pymodule(source) + function = pymodule["func"].get_object() + manager = reopened.pycore.object_info + call = pymodule.get_ast().body[-1].value.args[0].left + actual = arguments.create_arguments(None, function, call, pymodule.get_scope()) + assert manager.get_parameter_objects(function) is None + assert manager.get_passed_objects(function, 0) == [] + assert manager.get_passed_objects(function, 1) == [] + assert manager.get_returned(function, None) is None + assert manager.get_exact_returned(function, actual) is None + pos, expected = pymodule["Pos"].get_object(), pymodule["Key"].get_object() + manager.function_called(function, [pos], expected) + assert manager.get_returned(function, None) is None + manager.function_called(function, [pos, expected], expected) + assert manager.get_parameter_objects(function) == [pos, expected] + assert manager.get_passed_objects(function, 1) == [expected] + assert manager.get_passed_objects(function, 2) == [] + assert manager.get_returned(function, None) is expected + assert manager.get_exact_returned(function, actual) is expected + assert execute(reopened, source) == "True\n" + finally: + reopened.close() + + +def test_calltip_without_positional_receiver_keeps_keyword_only_parameter(project): + source = module( + project, + dedent("""\ + class Owner: + def func(*, value: int = 2): + return value + print(Owner.func(value=2)) + """), + ) + code = source.read() + assert execute(project, source) == "2\n" + offset = code.rindex(".func(") + 1 + assert get_calltip(project, code, offset, resource=source, remove_self=True) == ( + "source.Owner.func(*, value: int = 2)" + ) + + +def test_direct_inline_parameter_on_assignment_preserves_runtime(project): + source = module(project, "target = int\nprint(target is int)\n") + before = source.read() + assert execute(project, source) == "True\n" + with pytest.raises(exceptions.RefactoringError): + inline.InlineParameter(project, source, before.index("target") + 1) + assert source.read() == before + assert execute(project, source) == "True\n"