mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(ci): resolve class-local constants when detecting extra="allow"
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
e1dcb740a0
commit
4084f38128
2 changed files with 52 additions and 17 deletions
|
|
@ -119,7 +119,7 @@ def _statements(node: ast.AST) -> tuple[ast.stmt, ...]:
|
|||
return tuple(statement for statement, _ in _scoped_statements(node))
|
||||
|
||||
|
||||
def _module_bindings(node: ast.AST) -> tuple[Binding, ...]:
|
||||
def _bindings(node: ast.AST) -> tuple[Binding, ...]:
|
||||
return tuple(
|
||||
Binding(statement.lineno, name, value, branch)
|
||||
for statement, branch in _scoped_statements(node)
|
||||
|
|
@ -166,26 +166,39 @@ def _assigns_extra_allow(statement: ast.stmt, target_names: Sequence[str], scope
|
|||
)
|
||||
|
||||
|
||||
def _legacy_config_sets_extra_allow(class_def: ast.ClassDef, scope: Scope) -> bool:
|
||||
def _body_scope(node: ast.AST, enclosing: tuple[Binding, ...]) -> tuple[Binding, ...]:
|
||||
"""The bindings a statement in ``node``'s body reads, so a class-local constant is resolved
|
||||
and shadows a module-level name of its own, the way the class body itself would read it."""
|
||||
local: Final = _bindings(node)
|
||||
shadowed: Final = frozenset(binding.name for binding in local)
|
||||
return tuple(binding for binding in enclosing if binding.name not in shadowed) + local
|
||||
|
||||
|
||||
def _legacy_config_sets_extra_allow(config: ast.ClassDef, enclosing: tuple[Binding, ...]) -> bool:
|
||||
bindings: Final = _body_scope(config, enclosing)
|
||||
return any(
|
||||
isinstance(statement, ast.ClassDef)
|
||||
and statement.name == "Config"
|
||||
and any(
|
||||
isinstance(inner, ast.Assign)
|
||||
and any(isinstance(target, ast.Name) and target.id == "extra" for target in inner.targets)
|
||||
and _is_allow_literal(inner.value, scope)
|
||||
for inner in _statements(statement)
|
||||
)
|
||||
for statement in _statements(class_def)
|
||||
isinstance(inner, ast.Assign)
|
||||
and any(isinstance(target, ast.Name) and target.id == "extra" for target in inner.targets)
|
||||
and _is_allow_literal(inner.value, Scope(bindings, inner.lineno))
|
||||
for inner in _statements(config)
|
||||
)
|
||||
|
||||
|
||||
def _class_allows_extra(class_def: ast.ClassDef, scope: Scope) -> bool:
|
||||
if any(_is_extra_allow_keyword(keyword, scope) for keyword in class_def.keywords):
|
||||
def _class_allows_extra(class_def: ast.ClassDef, enclosing: tuple[Binding, ...]) -> bool:
|
||||
if any(_is_extra_allow_keyword(keyword, Scope(enclosing, class_def.lineno)) for keyword in class_def.keywords):
|
||||
return True
|
||||
if any(_assigns_extra_allow(statement, ["model_config"], scope) for statement in _statements(class_def)):
|
||||
bindings: Final = _body_scope(class_def, enclosing)
|
||||
statements: Final = _statements(class_def)
|
||||
if any(
|
||||
_assigns_extra_allow(statement, ["model_config"], Scope(bindings, statement.lineno)) for statement in statements
|
||||
):
|
||||
return True
|
||||
return _legacy_config_sets_extra_allow(class_def, scope)
|
||||
return any(
|
||||
isinstance(statement, ast.ClassDef)
|
||||
and statement.name == "Config"
|
||||
and _legacy_config_sets_extra_allow(statement, bindings)
|
||||
for statement in statements
|
||||
)
|
||||
|
||||
|
||||
def _iter_classes(node: ast.AST, prefix: str = "") -> Iterator[tuple[str, ast.ClassDef]]:
|
||||
|
|
@ -203,11 +216,11 @@ def find_violations_in_source(source: str, relative_path: str) -> tuple[Violatio
|
|||
if "allow" not in source:
|
||||
return ()
|
||||
tree: Final = ast.parse(source, filename=relative_path)
|
||||
bindings: Final = _module_bindings(tree)
|
||||
bindings: Final = _bindings(tree)
|
||||
return tuple(
|
||||
Violation(file=relative_path, line=class_def.lineno, model=qualified)
|
||||
for qualified, class_def in _iter_classes(tree)
|
||||
if _class_allows_extra(class_def, Scope(bindings, class_def.lineno))
|
||||
if _class_allows_extra(class_def, bindings)
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -124,6 +124,19 @@ _ratchet_spec.loader.exec_module(ratchet)
|
|||
"class Foo(BaseModel):\n model_config = CONFIG\n",
|
||||
id="one_branch_of_a_conditional_binding_allows",
|
||||
),
|
||||
pytest.param(
|
||||
'class Foo(BaseModel):\n _CONFIG = ConfigDict(extra="allow")\n model_config = _CONFIG\n',
|
||||
id="class_local_constant",
|
||||
),
|
||||
pytest.param(
|
||||
'FORBID = ConfigDict(extra="forbid")\n\n\nclass Foo(BaseModel):\n'
|
||||
' FORBID = ConfigDict(extra="allow")\n model_config = FORBID\n',
|
||||
id="class_local_constant_shadowing_a_harmless_module_one",
|
||||
),
|
||||
pytest.param(
|
||||
'class Foo(BaseModel):\n class Config:\n _EXTRA = "allow"\n extra = _EXTRA\n',
|
||||
id="legacy_inner_config_reading_its_own_constant",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_detects_extra_allow(source):
|
||||
|
|
@ -181,6 +194,15 @@ def test_detects_extra_allow(source):
|
|||
"class Foo(BaseModel):\n model_config = CONFIG\n",
|
||||
id="conditional_allow_rebound_unconditionally_before_the_class",
|
||||
),
|
||||
pytest.param(
|
||||
'ALLOW = ConfigDict(extra="allow")\n\n\nclass Foo(BaseModel):\n'
|
||||
' ALLOW = ConfigDict(extra="forbid")\n model_config = ALLOW\n',
|
||||
id="class_local_constant_shadowing_a_permissive_module_one",
|
||||
),
|
||||
pytest.param(
|
||||
'class Foo(BaseModel):\n model_config = LATER\n LATER = ConfigDict(extra="allow")\n',
|
||||
id="class_local_constant_bound_below_the_line_reading_it",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_ignores_non_violations(source):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue