Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -499,7 +499,8 @@ it can't see through):
- **`sql-fstring-injection`** / **`sql-concat-injection`** /
**`sql-percent-format-injection`** / **`sql-format-injection`** — SQL built
with an f-string, `+` concatenation, the `%` operator, or `str.format()`
instead of a psycopg t-string (3.14+), `psycopg.sql`, or bound params.
instead of a psycopg t-string (3.14+), `psycopg.sql`, or bound params. For an f-string the fix is
usually just `f"..."` -> `t"..."` (`{value}` is bound, `{name:i}` quotes an identifier).
- **`sql-positional-param`** — positional `%s` instead of named `%(name)s`.
- **`sql-forbidden-join`** — `RIGHT JOIN`/`LATERAL JOIN`/`CROSS APPLY` (same
patterns the `prek` skill's `check_files.py` forbids in `.sql` files).
Expand Down
3 changes: 2 additions & 1 deletion bmsdna/devtools/find_injection.py
Original file line number Diff line number Diff line change
Expand Up @@ -265,7 +265,8 @@ def check_csp(root: Path, scanned: list[Path], scope: _ScanFilter, only: set[Pat
continue
if has_csp(text):
found = True
weakened.extend(find_csp_weakening(candidate, text))
if not _is_test_file(candidate, root):
weakened.extend(find_csp_weakening(candidate, text))
if found:
return [f for f in weakened if only is None or f.path.resolve() in only]
return [
Expand Down
202 changes: 193 additions & 9 deletions bmsdna/devtools/injection_python.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,10 +44,195 @@ def _dotted(func: ast.expr, aliases: dict[str, str]) -> str:
return ".".join(reversed(parts))


def _is_literal(node: ast.expr) -> bool:
class _Scope:
"""Every binding of each name in one function/module scope. `None` marks a binding whose value
isn't statically known (parameter, loop over a variable, `+=`, import, ...)."""

def __init__(self, parent: _Scope | None, is_class: bool = False) -> None:
self.parent = parent
self.is_class = is_class
self.bindings: dict[str, list[ast.expr | None]] = {}
self.global_names: set[str] = set()


def _target_names(target: ast.AST) -> list[str]:
return [n.id for n in ast.walk(target) if isinstance(n, ast.Name)]


class _ScopeBuilder(ast.NodeVisitor):
def __init__(self, tree: ast.Module) -> None:
self.module = _Scope(None)
self.scope = self.module
self.calls: list[tuple[ast.Call, _Scope]] = []
self._body(tree)

def _body(self, node: ast.AST) -> None:
for child in ast.iter_child_nodes(node):
self.visit(child)

def _bind(self, name: str, value: ast.expr | None) -> None:
scope = self.module if name in self.scope.global_names else self.scope
scope.bindings.setdefault(name, []).append(value)

def _unknown(self, target: ast.AST) -> None:
for name in _target_names(target):
self._bind(name, None)

def _function(self, node: ast.FunctionDef | ast.AsyncFunctionDef | ast.Lambda) -> None:
if not isinstance(node, ast.Lambda):
self._bind(node.name, None)
for deco in node.decorator_list:
self.visit(deco)
for default in [*node.args.defaults, *(d for d in node.args.kw_defaults if d is not None)]:
self.visit(default)
outer = self.scope
parent = outer
while parent.is_class and parent.parent is not None:
parent = parent.parent
self.scope = _Scope(parent)
args = node.args
for a in [*args.posonlyargs, *args.args, *args.kwonlyargs, *(x for x in (args.vararg, args.kwarg) if x)]:
self._bind(a.arg, None)
for stmt in node.body if isinstance(node.body, list) else [node.body]:
self.visit(stmt)
self.scope = outer

visit_FunctionDef = visit_AsyncFunctionDef = visit_Lambda = _function

def visit_ClassDef(self, node: ast.ClassDef) -> None:
self._bind(node.name, None)
for expr in [*node.bases, *(k.value for k in node.keywords), *node.decorator_list]:
self.visit(expr)
outer = self.scope
self.scope = _Scope(outer, is_class=True)
for stmt in node.body:
self.visit(stmt)
self.scope = outer

def visit_Global(self, node: ast.Global) -> None:
self.scope.global_names.update(node.names)
for name in node.names:
self.module.bindings.setdefault(name, [])

def visit_Nonlocal(self, node: ast.Nonlocal) -> None:
scope = self.scope.parent
while scope is not None and scope is not self.module:
for name in node.names:
scope.bindings.setdefault(name, []).append(None)
scope = scope.parent

def visit_Assign(self, node: ast.Assign) -> None:
self.visit(node.value)
for target in node.targets:
if isinstance(target, ast.Name):
self._bind(target.id, node.value)
else:
self._unknown(target)

def visit_AnnAssign(self, node: ast.AnnAssign) -> None:
if node.value is not None:
self.visit(node.value)
if isinstance(node.target, ast.Name):
self._bind(node.target.id, node.value)
else:
self._unknown(node.target)

def visit_AugAssign(self, node: ast.AugAssign) -> None:
self.visit(node.value)
self._unknown(node.target)

def visit_NamedExpr(self, node: ast.NamedExpr) -> None:
self.visit(node.value)
self._bind(node.target.id, None)

def visit_For(self, node: ast.For | ast.AsyncFor) -> None:
self.visit(node.iter)
if isinstance(node.target, ast.Name) and isinstance(node.iter, (ast.List, ast.Tuple, ast.Set)):
for element in node.iter.elts:
self._bind(node.target.id, element)
else:
self._unknown(node.target)
for stmt in [*node.body, *node.orelse]:
self.visit(stmt)

visit_AsyncFor = visit_For

def visit_comprehension(self, node: ast.comprehension) -> None:
self._unknown(node.target)
self._body(node)

def visit_withitem(self, node: ast.withitem) -> None:
if node.optional_vars is not None:
self._unknown(node.optional_vars)
self._body(node)

def visit_ExceptHandler(self, node: ast.ExceptHandler) -> None:
if node.name:
self._bind(node.name, None)
self._body(node)

def visit_Import(self, node: ast.Import) -> None:
for a in node.names:
self._bind((a.asname or a.name).split(".")[0], None)

def visit_ImportFrom(self, node: ast.ImportFrom) -> None:
for a in node.names:
self._bind(a.asname or a.name, None)

def visit_Delete(self, node: ast.Delete) -> None:
for target in node.targets:
self._unknown(target)

def visit_MatchAs(self, node: ast.MatchAs) -> None:
if node.name:
self._bind(node.name, None)
self._body(node)

def visit_MatchStar(self, node: ast.MatchStar) -> None:
if node.name:
self._bind(node.name, None)

def visit_MatchMapping(self, node: ast.MatchMapping) -> None:
if node.rest:
self._bind(node.rest, None)
self._body(node)

def visit_Call(self, node: ast.Call) -> None:
self.calls.append((node, self.scope))
self._body(node)


def _lookup(scope: _Scope, name: str) -> list[ast.expr | None] | None:
current: _Scope | None = scope
while current is not None:
if name in current.bindings:
return current.bindings[name]
current = current.parent
return None


def _is_const(node: ast.expr, scope: _Scope, seen: frozenset[str] = frozenset()) -> bool:
"""True when every value `node` can take is a string built only from literals -- including names
whose every assignment in scope is such a constant."""
if isinstance(node, ast.Constant):
return isinstance(node.value, str)
return isinstance(node, ast.BinOp) and isinstance(node.op, ast.Add) and _is_literal(node.left) and _is_literal(node.right)
if isinstance(node, ast.BinOp):
return isinstance(node.op, ast.Add) and _is_const(node.left, scope, seen) and _is_const(node.right, scope, seen)
if isinstance(node, ast.JoinedStr):
return all(_is_const(v.value if isinstance(v, ast.FormattedValue) else v, scope, seen) for v in node.values)
if isinstance(node, ast.IfExp):
return _is_const(node.body, scope, seen) and _is_const(node.orelse, scope, seen)
if isinstance(node, ast.Name):
if node.id in seen:
return False
bound = _lookup(scope, node.id)
return bool(bound) and all(b is not None and _is_const(b, scope, seen | {node.id}) for b in bound)
return False


def _is_arg_list(node: ast.expr | None, scope: _Scope) -> bool:
"""A list/tuple command whose program is a literal: extra (even dynamic) items are arguments, not shell code."""
return isinstance(node, (ast.List, ast.Tuple)) and bool(node.elts) and _is_const(node.elts[0], scope)


def _shell_true(call: ast.Call) -> bool:
Expand All @@ -58,11 +243,11 @@ def _finding(path: Path, node: ast.Call, rule: str, message: str, severity: str)
return Finding(path, node.lineno, rule, message, severity=severity)


def _check_call(path: Path, node: ast.Call, aliases: dict[str, str]) -> list[Finding]:
def _check_call(path: Path, node: ast.Call, aliases: dict[str, str], scope: _Scope) -> list[Finding]:
name = _dotted(node.func, aliases)
last = name.rsplit(".", 1)[-1]
first_arg = node.args[0] if node.args else next((kw.value for kw in node.keywords if kw.arg in _FIRST_ARG_KEYWORDS), None)
dynamic = first_arg is not None and not _is_literal(first_arg)
dynamic = first_arg is not None and not _is_const(first_arg, scope)

if name in ("eval", "exec") and dynamic:
return [
Expand All @@ -84,8 +269,8 @@ def _check_call(path: Path, node: ast.Call, aliases: dict[str, str]) -> list[Fin
"error",
)
]
if name.startswith("subprocess.") and _shell_true(node):
if dynamic or isinstance(first_arg, (ast.JoinedStr, ast.BinOp)):
if name.startswith("subprocess.") and _shell_true(node) and not _is_arg_list(first_arg, scope):
if dynamic:
return [
_finding(
path,
Expand Down Expand Up @@ -183,7 +368,6 @@ def check_python_sinks(path: Path, source: str | None = None) -> list[Finding]:
return []
aliases = _import_aliases(tree)
findings: list[Finding] = []
for node in ast.walk(tree):
if isinstance(node, ast.Call):
findings.extend(_check_call(path, node, aliases))
for node, scope in _ScopeBuilder(tree).calls:
findings.extend(_check_call(path, node, aliases, scope))
return findings
21 changes: 11 additions & 10 deletions bmsdna/devtools/lint_sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,7 +109,8 @@ def _is_sql_composed_call(node: ast.AST) -> bool:
if not isinstance(node, ast.Call):
return False
func = node.func
return isinstance(func, ast.Attribute) and func.attr in ("SQL", "Identifier", "Composed")
name = func.attr if isinstance(func, ast.Attribute) else func.id if isinstance(func, ast.Name) else ""
return name in ("SQL", "Identifier", "Composed")


_UNWRAP_METHODS = frozenset({"strip", "lstrip", "rstrip"})
Expand Down Expand Up @@ -265,18 +266,18 @@ def visit(operand: ast.expr) -> None:
return "".join(parts) if saw_literal else None


_FSTRING_FIX = (
"Change the `f` prefix to `t` (psycopg t-string, Python 3.14+): `{value}` becomes a bound parameter and "
"`{name:i}` quotes a table/column identifier. Or use load_sql()/a .sql file for static SQL."
)
_GENERIC_FIX = "Use a psycopg t-string (Python 3.14+), psycopg.sql for dynamic SQL, or load_sql()/a .sql file for static SQL."


def _injection_finding(text: str, path: Path, lineno: int, rule: str, how: str) -> list[Finding]:
if parse_sql_text(text) is None:
return []
return [
Finding(
path,
lineno,
rule,
f"SQL built with {how} -- injection risk. Use a psycopg t-string (Python 3.14+), "
"psycopg.sql for dynamic SQL, or load_sql()/a .sql file for static SQL.",
)
]
fix = _FSTRING_FIX if rule == "sql-fstring-injection" else _GENERIC_FIX
return [Finding(path, lineno, rule, f"SQL built with {how} -- injection risk. {fix}")]


def _literal_findings(text: str, path: Path, lineno: int) -> list[Finding]:
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ packages = ["bmsdna"]

[project]
name = "bmsdna-devtools"
version = "0.27.0"
version = "0.28.0"
description = "Shared Azure DevOps / GitHub / git / Azure Monitor developer tooling for BMS projects"
readme = "README.md"
requires-python = ">=3.14" # pgdevkit>=0.7.1 requires 3.14; was >=3.11 before adding it as a dependency
Expand Down
94 changes: 94 additions & 0 deletions tests/test_find_injection.py
Original file line number Diff line number Diff line change
Expand Up @@ -298,3 +298,97 @@ def run(conn, t):
"""
)
assert find_injection.run([], root=tmp_path).findings == []


def test_subprocess_arg_list_with_shell_true_and_bare_sql(tmp_path: Path) -> None:
(tmp_path / "ok.py").write_text(
"""import subprocess
from psycopg.sql import SQL


def run(conn, target):
subprocess.check_call(["bun", "run", "build"], shell=True)
subprocess.check_call(["bun", "run", target], cwd=target, shell=True)
conn.execute(SQL("UPDATE t SET a = 1 WHERE id = %s"), (1,))
"""
)
(tmp_path / "bad.py").write_text(
"""import subprocess


def run(cmd, target):
subprocess.check_call([cmd, "x"], shell=True)
subprocess.check_call(f"bun run {target}", shell=True)
"""
)
result = find_injection.run([], root=tmp_path)
assert [(f.path.name, f.line, f.rule) for f in result.findings] == [("bad.py", 5, "py-shell-command"), ("bad.py", 6, "py-shell-command")]


def test_csp_weakening_in_test_files_is_ignored(tmp_path: Path) -> None:
(tmp_path / "app.js").write_text("run();\n")
(tmp_path / "web.config").write_text("Content-Security-Policy: default-src 'self'")
(tmp_path / "test_headers.py").write_text("H = \"Content-Security-Policy: script-src 'unsafe-eval'\"\n")
assert find_injection.run([], root=tmp_path).findings == []


def _lines(tmp_path: Path, src: str) -> list[int]:
path = tmp_path / "m.py"
path.write_text(src)
return [f.line for f in check_python_sinks(path)]


def test_names_assigned_only_constants_are_constant(tmp_path: Path) -> None:
src = """import os, subprocess
GREETING = "echo hi"
CMD = GREETING + " there"


def ok(flag):
os.system(CMD)
cmd = "ls" if flag else "pwd"
os.system(cmd)
for c in ("a", "b"):
os.system(c)
os.system(f"{GREETING} now")
prog = "bun"
subprocess.check_call([prog, "x"], shell=True)


def bad(arg, flag):
cmd = "ls"
cmd += arg
os.system(cmd)
other = "ls"
if flag:
other = arg
os.system(other)
os.system(arg)
for c in arg:
os.system(c)
os.system(undefined_name)
"""
assert _lines(tmp_path, src) == [20, 24, 25, 27, 28]


def test_global_rebinding_and_class_scope_defeat_constness(tmp_path: Path) -> None:
src = """import os
GREETING = "echo hi"


def rebind():
global GREETING
GREETING = input()


def use():
os.system(GREETING)


class K:
x = "ls"

def m(self):
os.system(x)
"""
assert _lines(tmp_path, src) == [11, 18]
Loading
Loading