"""L1: tht.execute._inject_limit — the AST-based LIMIT injection. The gate runs reviewer SQL through _inject_limit before execution so previews never return unbounded rows (and the +1 lets run_controlled detect truncation). These tests pin: LIMIT added when absent, respected when present, never added to non-query statements, and the truncation-detection contract (limit+1 rows). Pure logic, no DB. """ import sqlglot from tht.execute import _inject_limit def test_limit_injected_when_absent(): out, injected = _inject_limit("SELECT * FROM t", 10) assert injected is True # parse it back and confirm the LIMIT is 11 (10 + 1, for truncation detection) ast = sqlglot.parse_one(out, read="postgres") assert ast.args.get("limit") is not None # the limit expression should evaluate to 11 limit_expr = ast.args["limit"].expression assert int(limit_expr.to_py()) == 11 def test_existing_limit_respected_not_overwritten(): out, injected = _inject_limit("SELECT * FROM t LIMIT 5", 10) assert injected is False ast = sqlglot.parse_one(out, read="postgres") assert int(ast.args["limit"].expression.to_py()) == 5 # unchanged def test_no_limit_injected_on_non_query(): # a non-query statement (DDL): _inject_limit must leave it untouched (the # READ-ONLY transaction downstream rejects it, not the injector). out, injected = _inject_limit("INSERT INTO t VALUES (1)", 10) assert injected is False assert out == "INSERT INTO t VALUES (1)" def test_union_query_accepts_limit(): out, injected = _inject_limit("SELECT 1 UNION SELECT 2", 10) assert injected is True ast = sqlglot.parse_one(out, read="postgres") assert int(ast.args["limit"].expression.to_py()) == 11 def test_with_cte_query_accepts_limit(): sql = "WITH cte AS (SELECT 1) SELECT * FROM cte" _out, injected = _inject_limit(sql, 10) assert injected is True def test_limit_one_plus_n_for_truncation_detection(): # the whole point of +1: run_controlled fetches limit+1 rows, if it gets > # limit it knows truncation happened. Verify the arithmetic for several limits. for n in (1, 5, 100, 1000): out, _ = _inject_limit("SELECT * FROM t", n) ast = sqlglot.parse_one(out, read="postgres") assert int(ast.args["limit"].expression.to_py()) == n + 1