From 520d6793fa87c2c72d51862f08ce348d877b61a9 Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 03:49:40 +0000 Subject: [PATCH] ci: catch a config constant that a later rebinding overwrites Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../ban_pydantic_extra_allow.py | 60 ++++++++++++------- .../test_ban_pydantic_extra_allow.py | 15 +++++ 2 files changed, 53 insertions(+), 22 deletions(-) diff --git a/tests/code_coverage_tests/ban_pydantic_extra_allow.py b/tests/code_coverage_tests/ban_pydantic_extra_allow.py index e1671720cfb..e5917378cf4 100644 --- a/tests/code_coverage_tests/ban_pydantic_extra_allow.py +++ b/tests/code_coverage_tests/ban_pydantic_extra_allow.py @@ -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): diff --git a/tests/test_litellm/test_ban_pydantic_extra_allow.py b/tests/test_litellm/test_ban_pydantic_extra_allow.py index b6f050e1ce9..b9cd43736e0 100644 --- a/tests/test_litellm/test_ban_pydantic_extra_allow.py +++ b/tests/test_litellm/test_ban_pydantic_extra_allow.py @@ -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",