ci: catch a config constant that a later rebinding overwrites

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-08-08 03:49:40 +00:00
parent 49a38d1da2
commit 520d6793fa
2 changed files with 53 additions and 22 deletions

View file

@ -93,52 +93,68 @@ def _assigned_names(statement: ast.stmt) -> tuple[tuple[str, ast.expr], ...]:
return tuple((target.id, value) for target in targets if isinstance(target, ast.Name))
def _module_constants(body: Sequence[ast.stmt]) -> Mapping[str, ast.expr]:
"""Module-level ``NAME = value`` bindings, so a config held in a constant is still seen."""
return MappingProxyType(dict(binding for statement in body for binding in _assigned_names(statement)))
def _module_constants(body: Sequence[ast.stmt]) -> Mapping[str, tuple[ast.expr, ...]]:
"""Every module-level ``NAME = value`` binding, keyed by name, so a config held in a
constant is still seen and rebinding the name later can't hide an earlier one.
"""
bindings: Final = tuple(binding for statement in body for binding in _assigned_names(statement))
return MappingProxyType({name: tuple(value for bound, value in bindings if bound == name) for name, _ in bindings})
def _resolve(node: ast.expr, constants: Mapping[str, ast.expr], seen: frozenset[str] = frozenset()) -> ast.expr:
def _resolve(
node: ast.expr, constants: Mapping[str, tuple[ast.expr, ...]], seen: frozenset[str] = frozenset()
) -> tuple[ast.expr, ...]:
if isinstance(node, ast.Name) and node.id in constants and node.id not in seen:
return _resolve(constants[node.id], constants, seen | {node.id})
return node
return tuple(
candidate for value in constants[node.id] for candidate in _resolve(value, constants, seen | {node.id})
)
return (node,)
def _is_allow_literal(node: ast.expr, constants: Mapping[str, ast.expr]) -> bool:
resolved = _resolve(node, constants)
if isinstance(resolved, ast.Constant):
return resolved.value == "allow"
return isinstance(resolved, ast.Attribute) and resolved.attr == "allow"
def _is_allow_literal(node: ast.expr, constants: Mapping[str, tuple[ast.expr, ...]]) -> bool:
return any(
resolved.value == "allow" if isinstance(resolved, ast.Constant) else _is_allow_attribute(resolved)
for resolved in _resolve(node, constants)
)
def _is_extra_allow_keyword(keyword: ast.keyword, constants: Mapping[str, ast.expr]) -> bool:
def _is_allow_attribute(node: ast.expr) -> bool:
return isinstance(node, ast.Attribute) and node.attr == "allow"
def _is_extra_allow_keyword(keyword: ast.keyword, constants: Mapping[str, tuple[ast.expr, ...]]) -> bool:
return keyword.arg == "extra" and _is_allow_literal(keyword.value, constants)
def _mapping_sets_extra_allow(node: ast.Dict, constants: Mapping[str, ast.expr]) -> bool:
def _mapping_sets_extra_allow(node: ast.Dict, constants: Mapping[str, tuple[ast.expr, ...]]) -> bool:
return any(
isinstance(key, ast.Constant) and key.value == "extra" and _is_allow_literal(value, constants)
for key, value in zip(node.keys, node.values)
)
def _is_extra_allow_value(node: ast.expr, constants: Mapping[str, ast.expr]) -> bool:
resolved = _resolve(node, constants)
if isinstance(resolved, ast.Call):
return any(_is_extra_allow_keyword(keyword, constants) for keyword in resolved.keywords)
if isinstance(resolved, ast.Dict):
return _mapping_sets_extra_allow(resolved, constants)
def _config_sets_extra_allow(node: ast.expr, constants: Mapping[str, tuple[ast.expr, ...]]) -> bool:
if isinstance(node, ast.Call):
return any(_is_extra_allow_keyword(keyword, constants) for keyword in node.keywords)
if isinstance(node, ast.Dict):
return _mapping_sets_extra_allow(node, constants)
return False
def _assigns_extra_allow(statement: ast.stmt, target_names: Sequence[str], constants: Mapping[str, ast.expr]) -> bool:
def _is_extra_allow_value(node: ast.expr, constants: Mapping[str, tuple[ast.expr, ...]]) -> bool:
return any(_config_sets_extra_allow(resolved, constants) for resolved in _resolve(node, constants))
def _assigns_extra_allow(
statement: ast.stmt, target_names: Sequence[str], constants: Mapping[str, tuple[ast.expr, ...]]
) -> bool:
return any(
name in set(target_names) and _is_extra_allow_value(value, constants)
for name, value in _assigned_names(statement)
)
def _legacy_config_sets_extra_allow(class_def: ast.ClassDef, constants: Mapping[str, ast.expr]) -> bool:
def _legacy_config_sets_extra_allow(class_def: ast.ClassDef, constants: Mapping[str, tuple[ast.expr, ...]]) -> bool:
return any(
isinstance(statement, ast.ClassDef)
and statement.name == "Config"
@ -152,7 +168,7 @@ def _legacy_config_sets_extra_allow(class_def: ast.ClassDef, constants: Mapping[
)
def _class_allows_extra(class_def: ast.ClassDef, constants: Mapping[str, ast.expr]) -> bool:
def _class_allows_extra(class_def: ast.ClassDef, constants: Mapping[str, tuple[ast.expr, ...]]) -> bool:
if any(_is_extra_allow_keyword(keyword, constants) for keyword in class_def.keywords):
return True
if any(_assigns_extra_allow(statement, ["model_config"], constants) for statement in class_def.body):

View file

@ -60,6 +60,21 @@ _REPO_ROOT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", ".."
'EXTRA = "allow"\n\n\nclass Foo(BaseModel):\n model_config = ConfigDict(extra=EXTRA)\n',
id="module_constant_extra_value",
),
pytest.param(
'ALLOW = ConfigDict(extra="allow")\n\n\nclass Foo(BaseModel):\n'
' model_config = ALLOW\n\n\nALLOW = ConfigDict(extra="forbid")\n',
id="constant_rebound_to_forbid_after_the_class",
),
pytest.param(
'ALLOW = ConfigDict(extra="forbid")\nALLOW = ConfigDict(extra="allow")\n\n\n'
"class Foo(BaseModel):\n model_config = ALLOW\n",
id="constant_rebound_to_allow_before_the_class",
),
pytest.param(
'EXTRA = "allow"\n\n\nclass Foo(BaseModel):\n'
' model_config = ConfigDict(extra=EXTRA)\n\n\nEXTRA = "forbid"\n',
id="extra_value_constant_rebound_after_the_class",
),
pytest.param(
'EXTRA = "allow"\n\n\nclass Foo(BaseModel, extra=EXTRA):\n pass\n',
id="module_constant_class_keyword",