From c94cb3ea4ebd14248b6610482a30107866d069ac Mon Sep 17 00:00:00 2001 From: Neil Watts <1678092075@qq.com> Date: Wed, 2 Sep 2026 19:09:30 +0800 Subject: [PATCH] B008: resolve imported immutable calls --- README.rst | 2 + bugbear.py | 279 +++++++++++++++++++- tests/eval_files/b008_extended.py | 211 ++++++++++++++- tests/eval_files/b008_extended_shadowing.py | 16 ++ tests/eval_files/b039_extended.py | 5 + 5 files changed, 506 insertions(+), 7 deletions(-) create mode 100644 tests/eval_files/b008_extended_shadowing.py create mode 100644 tests/eval_files/b039_extended.py diff --git a/README.rst b/README.rst index dee86d3..2d8bcbb 100644 --- a/README.rst +++ b/README.rst @@ -517,6 +517,8 @@ Change Log UNRELEASED ~~~~~~~~~~ +* B008: resolve direct module-level imports and aliases when matching + ``extend-immutable-calls`` (#252) * B015: allow comparisons inside pytest/unittest exception and warning assertion context managers, where an overloaded comparison may intentionally raise (#462). * B020: stop flagging names bound inside nested destructuring patterns, like diff --git a/bugbear.py b/bugbear.py index 260bf1f..a42a004 100644 --- a/bugbear.py +++ b/bugbear.py @@ -43,6 +43,7 @@ ast.GeneratorExp, ) FUNCTION_NODES = (ast.AsyncFunctionDef, ast.FunctionDef, ast.Lambda) +EAGER_COMPREHENSION_NODES = (ast.ListComp, ast.SetComp, ast.DictComp) FUNCTIONS_WITHOUT_SIDE_EFFECTS = ( "all", "any", @@ -463,6 +464,10 @@ class BugBearVisitor(ast.NodeVisitor): _b023_seen: set[ast.Name] = attr.ib(factory=set, init=False) _b023_scopes: dict[int, tuple] = attr.ib(factory=dict, init=False) _b005_imports: set[str] = attr.ib(factory=set, init=False) + # None marks an imported name that has since been rebound at module scope. + _b008_imports: dict[str, str | None] = attr.ib(factory=dict, init=False) + _b008_class_imports: list[dict[str, str | None]] = attr.ib(factory=list, init=False) + _b008_class_globals: list[set[str]] = attr.ib(factory=list, init=False) # set to "*" when inside a try/except*, for correctly printing errors in_trystar: str = attr.ib(default="") @@ -485,6 +490,99 @@ def node_stack(self) -> list[ast.AST]: context, stack = self.contexts[-1] return stack + def _b008_shadow_imports( + self, + names: Iterable[str], + imports: dict[str, str | None] | None = None, + ) -> None: + if imports is None: + imports = self._b008_imports + for name in names: + if imports.get(name) is not None: + imports[name] = None + + def _b008_shadow_bindings( + self, + names: Iterable[str], + imports: dict[str, str | None] | None = None, + ) -> None: + names = tuple(names) + self._b008_shadow_imports(names, imports) + if ( + imports is not None + and self._b008_class_imports + and imports is self._b008_class_imports[-1] + and self._b008_class_globals + ): + global_names = self._b008_class_globals[-1].intersection(names) + self._b008_shadow_imports(global_names) + + def _b008_shadow_named_expr_targets( + self, + nodes: Iterable[ast.AST], + imports: dict[str, str | None] | None = None, + ) -> None: + finder = B008NamedExprFinder() + finder.visit(list(nodes)) + self._b008_shadow_bindings(finder.names, imports) + + def _b008_in_module_scope(self) -> bool: + return len(self.contexts) == 1 and isinstance(self.contexts[0].node, ast.Module) + + def _b008_is_direct_module_statement(self) -> bool: + return self._b008_in_module_scope() and len(self.node_stack) == 2 + + def _b008_in_module_child(self) -> bool: + return len(self.contexts) == 2 and isinstance(self.contexts[0].node, ast.Module) + + def _b008_in_direct_module_child(self) -> bool: + return self._b008_in_module_child() and len(self.contexts[0].stack) == 1 + + def _b008_in_direct_module_class_scope(self) -> bool: + return self._b008_in_direct_module_child() and isinstance( + self.contexts[-1].node, ast.ClassDef + ) + + def _b008_in_direct_module_class_child(self) -> bool: + return ( + len(self.contexts) == 3 + and isinstance(self.contexts[0].node, ast.Module) + and isinstance(self.contexts[1].node, ast.ClassDef) + and len(self.contexts[1].stack) == 1 + ) + + def _b008_binding_imports(self) -> dict[str, str | None] | None: + if self._b008_in_module_scope(): + return self._b008_imports + if self._b008_in_direct_module_class_scope() and self._b008_class_imports: + return self._b008_class_imports[-1] + return None + + def _b008_function_imports(self) -> dict[str, str | None] | None: + if self._b008_in_direct_module_child(): + return self._b008_imports + if self._b008_in_direct_module_class_child() and self._b008_class_imports: + return self._b008_class_imports[-1] + return None + + @staticmethod + def _b008_function_annotations( + node: ast.FunctionDef | ast.AsyncFunctionDef, + ) -> list[ast.expr]: + arguments = [ + *node.args.posonlyargs, + *node.args.args, + *node.args.kwonlyargs, + ] + if node.args.vararg is not None: + arguments.append(node.args.vararg) + if node.args.kwarg is not None: + arguments.append(node.args.kwarg) + annotations = [argument.annotation for argument in arguments] + if node.returns is not None: + annotations.append(node.returns) + return [annotation for annotation in annotations if annotation is not None] + def in_class_init(self) -> bool: return ( len(self.contexts) >= 2 @@ -555,6 +653,9 @@ def visit_ExceptHandler(self, node: ast.ExceptHandler) -> None: ): self.add_error("B040", node) self.b040_caught_exception = old_b040_caught_exception + imports = self._b008_binding_imports() + if imports is not None and node.name is not None: + self._b008_shadow_bindings((node.name,), imports) def visit_UAdd(self, node: ast.UAdd) -> None: trailing_nodes = list(map(type, self.node_window[-4:])) @@ -631,6 +732,22 @@ def visit_Call(self, node: ast.Call) -> None: def visit_Module(self, node: ast.Module) -> None: self.generic_visit(node) + def visit_Name( # noqa: B906 # names don't contain other names + self, node: ast.Name + ) -> None: + imports = self._b008_binding_imports() + if ( + imports is not None + and isinstance(node.ctx, (ast.Store, ast.Del)) + and not ( + isinstance(node.ctx, ast.Store) + and len(self.node_stack) >= 2 + and isinstance(self.node_stack[-2], ast.AnnAssign) + and self.node_stack[-2].value is None + ) + ): + self._b008_shadow_bindings((node.id,), imports) + def visit_Assign(self, node: ast.Assign) -> None: self.check_for_b040_usage(node.value) if len(node.targets) == 1: @@ -682,25 +799,63 @@ def visit_Assert(self, node: ast.Assert) -> None: def visit_AsyncFunctionDef(self, node: ast.AsyncFunctionDef) -> None: self.check_for_b902(node) + imports = self._b008_function_imports() + if imports is not None: + self._b008_shadow_named_expr_targets(node.decorator_list, imports) self.check_for_b006_and_b008(node) + if imports is not None: + self._b008_shadow_named_expr_targets( + self._b008_function_annotations(node), imports + ) self.check_for_b019(node) self.generic_visit(node) + if self._b008_in_module_child(): + self._b008_shadow_imports((node.name,)) + elif self._b008_in_direct_module_class_child() and imports is not None: + self._b008_shadow_bindings((node.name,), imports) def visit_FunctionDef(self, node: ast.FunctionDef) -> None: self.check_for_b901(node) self.check_for_b902(node) + imports = self._b008_function_imports() + if imports is not None: + self._b008_shadow_named_expr_targets(node.decorator_list, imports) self.check_for_b006_and_b008(node) + if imports is not None: + self._b008_shadow_named_expr_targets( + self._b008_function_annotations(node), imports + ) self.check_for_b019(node) self.check_for_b021(node) self.check_for_b906(node) self.generic_visit(node) + if self._b008_in_module_child(): + self._b008_shadow_imports((node.name,)) + elif self._b008_in_direct_module_class_child() and imports is not None: + self._b008_shadow_bindings((node.name,), imports) def visit_ClassDef(self, node: ast.ClassDef) -> None: self.check_for_b903(node) + is_module_class = self._b008_in_direct_module_child() + parent_class_imports = None + if self._b008_in_direct_module_class_child() and self._b008_class_imports: + parent_class_imports = self._b008_class_imports[-1] + if is_module_class: + definition_nodes = [*node.decorator_list, *node.bases, *node.keywords] + self._b008_shadow_named_expr_targets(definition_nodes) + self._b008_class_imports.append(self._b008_imports.copy()) + self._b008_class_globals.append(set()) self.check_for_b021(node) self.check_for_b024_and_b027(node) self.check_for_b042(node) self.generic_visit(node) + if is_module_class: + self._b008_class_imports.pop() + self._b008_class_globals.pop() + if self._b008_in_module_child(): + self._b008_shadow_imports((node.name,)) + elif parent_class_imports is not None: + self._b008_shadow_bindings((node.name,), parent_class_imports) def visit_Try(self, node: ast.Try | ast.TryStar) -> None: self.check_for_b012(node) @@ -733,6 +888,28 @@ def visit_With(self, node: ast.With) -> None: self.check_for_b908(node) self.generic_visit(node) + def visit_MatchAs(self, node: ast.MatchAs) -> None: + imports = self._b008_binding_imports() + if imports is not None and node.name is not None: + self._b008_shadow_bindings((node.name,), imports) + self.generic_visit(node) + + def visit_MatchMapping(self, node: ast.MatchMapping) -> None: + imports = self._b008_binding_imports() + if imports is not None and node.rest is not None: + self._b008_shadow_bindings((node.rest,), imports) + self.generic_visit(node) + + def visit_MatchStar(self, node: ast.MatchStar) -> None: + imports = self._b008_binding_imports() + if imports is not None and node.name is not None: + self._b008_shadow_bindings((node.name,), imports) + self.generic_visit(node) + + def visit_Global(self, node: ast.Global) -> None: + if self._b008_in_direct_module_class_scope() and self._b008_class_globals: + self._b008_class_globals[-1].update(node.names) + def visit_JoinedStr(self, node: ast.JoinedStr) -> None: self.check_for_b907(node) self.generic_visit(node) @@ -742,12 +919,61 @@ def visit_AnnAssign(self, node: ast.AnnAssign) -> None: self.check_for_b040_usage(node.value) self.generic_visit(node) + def visit_NamedExpr(self, node: ast.NamedExpr) -> None: + if isinstance(node.target, ast.Name) and ( + self._b008_in_module_scope() + or ( + len(self.contexts) >= 2 + and isinstance(self.contexts[0].node, ast.Module) + and all( + isinstance(context.node, EAGER_COMPREHENSION_NODES) + for context in self.contexts[1:] + ) + ) + ): + self._b008_shadow_imports((node.target.id,)) + self.generic_visit(node) + def visit_Import(self, node: ast.Import) -> None: self.check_for_b005(node) + if self.b008_b039_extend_immutable_calls and self._b008_in_module_scope(): + for name in node.names: + bound_name = name.asname or name.name.partition(".")[0] + qualified_name = name.name if name.asname else bound_name + if self._b008_is_direct_module_statement(): + self._b008_imports[bound_name] = qualified_name + else: + self._b008_shadow_imports((bound_name,)) + elif (imports := self._b008_binding_imports()) is not None: + self._b008_shadow_bindings( + (name.asname or name.name.partition(".")[0] for name in node.names), + imports, + ) self.generic_visit(node) def visit_ImportFrom(self, node: ast.ImportFrom) -> None: self.check_for_b005(node) + if self.b008_b039_extend_immutable_calls and self._b008_in_module_scope(): + for name in node.names: + if name.name == "*": + self._b008_shadow_imports(self._b008_imports) + elif ( + self._b008_is_direct_module_statement() + and node.level == 0 + and node.module is not None + ): + self._b008_imports[name.asname or name.name] = ( + f"{node.module}.{name.name}" + ) + else: + self._b008_shadow_imports((name.asname or name.name,)) + elif (imports := self._b008_binding_imports()) is not None: + if any(name.name == "*" for name in node.names): + self._b008_shadow_bindings(imports, imports) + else: + self._b008_shadow_bindings( + (name.asname or name.name for name in node.names), imports + ) self.generic_visit(node) def visit_Set(self, node: ast.Set) -> None: @@ -820,12 +1046,18 @@ def check_for_b005(self, node: ast.Import | ast.ImportFrom | ast.Call) -> None: def check_for_b006_and_b008( self, node: ast.FunctionDef | ast.AsyncFunctionDef ) -> None: + imported_names = self._b008_function_imports() + if not self.b008_b039_extend_immutable_calls: + imported_names = None visitor = FunctionDefDefaultsVisitor( error_codes["B006"], error_codes["B008"], self.b008_b039_extend_immutable_calls, + imported_names, ) - visitor.visit(node.args.defaults + node.args.kw_defaults) + for default in node.args.defaults + node.args.kw_defaults: + if default is not None: + visitor.visit(default) self.errors.extend(visitor.errors) def check_for_b039(self, node: ast.Call) -> None: @@ -2596,6 +2828,17 @@ def visit(self, node): return node +class B008NamedExprFinder(NamedExprFinder): + def visit_GeneratorExp(self, node: ast.GeneratorExp) -> None: + self.visit(node.generators[0].iter) + + def visit_Lambda(self, node: ast.Lambda) -> None: # noqa: B906 + self.visit(node.args.defaults) + self.visit( + [default for default in node.args.kw_defaults if default is not None] + ) + + class FunctionDefDefaultsVisitor(ast.NodeVisitor): """Used by B006, B008, and B039. B039 is essentially B006+B008 but for ContextVar.""" @@ -2604,12 +2847,14 @@ def __init__( error_code_calls: "Error", # B006 or B039 error_code_literals: "Error", # B008 or B039 b008_b039_extend_immutable_calls: set[str] | None = None, + imported_names: dict[str, str | None] | None = None, ) -> None: self.b008_b039_extend_immutable_calls = ( b008_b039_extend_immutable_calls or set() ) self.error_code_calls = error_code_calls self.error_code_literals = error_code_literals + self.imported_names = imported_names or {} for node in B006_MUTABLE_LITERALS + B006_MUTABLE_COMPREHENSIONS: setattr(self, f"visit_{node}", self.visit_mutable_literal_or_comprehension) self.errors: list[error] = [] @@ -2639,7 +2884,22 @@ def visit_Call(self, node: ast.Call) -> None: self.generic_visit(node) return - if call_path in B008_IMMUTABLE_CALLS | self.b008_b039_extend_immutable_calls: + if call_path in B008_IMMUTABLE_CALLS: + self.generic_visit(node) + return + + head, separator, tail = call_path.partition(".") + if head not in self.imported_names: + extended_call_paths = {call_path} + elif (qualified_name := self.imported_names[head]) is None: + extended_call_paths = set() + else: + resolved_call_path = qualified_name + if separator: + resolved_call_path = f"{resolved_call_path}.{tail}" + extended_call_paths = {call_path, resolved_call_path} + + if extended_call_paths & self.b008_b039_extend_immutable_calls: self.generic_visit(node) return @@ -2660,10 +2920,19 @@ def visit_Call(self, node: ast.Call) -> None: # Check for nested functions. self.generic_visit(node) + def visit_NamedExpr(self, node: ast.NamedExpr) -> None: + if isinstance(node.target, ast.Name) and node.target.id in self.imported_names: + self.imported_names[node.target.id] = None + self.generic_visit(node) + def visit_Lambda(self, node) -> None: # noqa: B906 - # Don't recurse into lambda expressions - # as they are evaluated at call time. - pass + self.visit(node.args.defaults) + self.visit( + [default for default in node.args.kw_defaults if default is not None] + ) + + def visit_GeneratorExp(self, node: ast.GeneratorExp) -> None: + self.visit(node.generators[0].iter) def visit(self, node) -> None: """Like super-visit but supports iteration over lists.""" diff --git a/tests/eval_files/b008_extended.py b/tests/eval_files/b008_extended.py index 4f2c179..979ea15 100644 --- a/tests/eval_files/b008_extended.py +++ b/tests/eval_files/b008_extended.py @@ -2,7 +2,15 @@ from typing import List import fastapi +import fastapi as fastapi_alias +from fastapi import Depends +from fastapi import Depends as DependsAlias +from fastapi import Depends as loop_depends from fastapi import Query +from other import Depends as OtherDepends + +if condition: + from fastapi import Depends as conditional_depends def this_is_okay_extended(db=fastapi.Depends(get_db)): ... @@ -11,5 +19,204 @@ def this_is_okay_extended(db=fastapi.Depends(get_db)): ... def this_is_okay_extended_second(data: List[str] = fastapi.Query(None)): ... -# not okay, relative import not listed -def not_okay(data: List[str] = Query(None)): ... # B008: 31 +def this_is_okay_imported(db=Depends(get_db)): ... + + +def this_is_okay_imported_alias(db=DependsAlias(get_db)): ... + + +def this_is_okay_module_alias(db=fastapi_alias.Depends(get_db)): ... + + +class API: + def this_is_okay_method(self, db=Depends(get_db)): ... + + +class APIWithRebinding: + Depends = other + + def not_okay_method(self, db=Depends(get_db)): ... # B008: 33 + + +class APIWithImportRebinding: + from other import Depends + + def not_okay_method(self, db=Depends(get_db)): ... # B008: 33 + + +def not_okay_other_import(db=OtherDepends(get_db)): ... # B008: 29 + + +Depends: Callable + + +def this_is_okay_after_annotation_only(db=Depends(get_db)): ... + + +def Depends(): ... + + +def not_okay_redefined(db=Depends(get_db)): ... # B008: 26 + + +Query = lambda value: value + + +def not_okay_reassigned(data: List[str] = Query(None)): ... # B008: 42 + + +def not_okay_conditional_import( + db=conditional_depends(get_db), # B008: 7 +): ... + + +for loop_depends in providers: + _ = loop_depends + + +def not_okay_loop_rebound(db=loop_depends(get_db)): ... # B008: 29 + + +from fastapi import Depends as walrus_depends + + +def not_okay_walrus_default( + first=(walrus_depends := other), + db=walrus_depends(get_db), # B008: 7 +): ... + + +from fastapi import Depends as decorator_depends + + +@decorate((decorator_depends := other)) +def not_okay_walrus_decorator(db=decorator_depends(get_db)): ... # B008: 33 + + +from fastapi import Depends as class_base_depends + + +class RebindsInBase((class_base_depends := OtherBase)): ... + + +def not_okay_walrus_class_base(db=class_base_depends(get_db)): ... # B008: 34 + + +from fastapi import Depends as class_keyword_depends + + +class RebindsInKeyword(metaclass=(class_keyword_depends := Meta)): ... + + +def not_okay_walrus_class_keyword( + db=class_keyword_depends(get_db), # B008: 7 +): ... + + +from fastapi import Depends as nested_comprehension_depends + +_ = [ + [(nested_comprehension_depends := other) for inner in values] + for outer in values +] + + +def not_okay_walrus_nested_comprehension( + db=nested_comprehension_depends(get_db), # B008: 7 +): ... + + +from fastapi import Depends as lambda_default_depends + + +@decorate(lambda value=(lambda_default_depends := other): value) +def not_okay_walrus_lambda_default( + db=lambda_default_depends(get_db), # B008: 7 +): ... + + +from fastapi import Depends as function_lambda_depends + + +def not_okay_walrus_function_lambda_default( + first=(lambda value=(function_lambda_depends := other): value), + db=function_lambda_depends(get_db), # B008: 7 +): ... + + +from fastapi import Depends as annotation_depends + + +def annotation_rebinds_import(value: (annotation_depends := other)): ... + + +def not_okay_walrus_annotation( + db=annotation_depends(get_db), # B008: 7 +): ... + + +from fastapi import Depends as generator_depends + +_ = ((generator_depends := other) for value in values) + + +def this_is_okay_deferred_generator(db=generator_depends(get_db)): ... + + +from fastapi import Depends as default_generator_depends + + +def this_is_okay_deferred_default_generator( + generator=((default_generator_depends := other) for value in values), + db=default_generator_depends(get_db), +): ... + + +from fastapi import Depends as decorator_generator_depends + + +@decorate((decorator_generator_depends := other) for value in values) +def this_is_okay_deferred_decorator_generator( + db=decorator_generator_depends(get_db), +): ... + + +from fastapi import Depends as global_depends + + +class GlobalRebinding: + global global_depends + global_depends = other + + +def not_okay_class_global_rebinding( + db=global_depends(get_db), # B008: 7 +): ... + + +from fastapi import Depends as global_def_depends + + +class GlobalDefRebinding: + global global_def_depends + + @staticmethod + def global_def_depends(): ... + + +def not_okay_class_global_def( + db=global_def_depends(get_db), # B008: 7 +): ... + + +from fastapi import Depends as global_import_depends + + +class GlobalImportRebinding: + global global_import_depends + from other import Depends as global_import_depends + + +def not_okay_class_global_import( + db=global_import_depends(get_db), # B008: 7 +): ... diff --git a/tests/eval_files/b008_extended_shadowing.py b/tests/eval_files/b008_extended_shadowing.py new file mode 100644 index 0000000..554d0c1 --- /dev/null +++ b/tests/eval_files/b008_extended_shadowing.py @@ -0,0 +1,16 @@ +# OPTIONS: extend_immutable_calls=["Depends", "factory"] +from fastapi import Depends + + +def this_is_okay_imported(db=Depends(get_db)): ... + + +def Depends(): ... + + +def not_okay_redefined(db=Depends(get_db)): ... # B008: 26 + + +def source_path_stays_configured( + first=(factory := other), second=factory() +): ... diff --git a/tests/eval_files/b039_extended.py b/tests/eval_files/b039_extended.py new file mode 100644 index 0000000..def2ce4 --- /dev/null +++ b/tests/eval_files/b039_extended.py @@ -0,0 +1,5 @@ +# OPTIONS: extend_immutable_calls=["factory"] +from contextvars import ContextVar + +ContextVar("configured", default=((factory := other), factory())[1]) +ContextVar("not_configured", default=other()) # B039: 37