Skip to content

Commit d2c09a9

Browse files
flyersworderclaude
andcommitted
fix: reject unknown fields on SemanticRule, remove dead CheckResult.severity
- Add extra="forbid" to SemanticRule so old YAML with filter_column fails loudly instead of silently losing enforcement - Remove CheckResult.severity field (was set to "block" by every checker but never read by the Validator — enforcement comes from the rule config) - Document session cost timing design choice in run_query Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
1 parent ce176d4 commit d2c09a9

4 files changed

Lines changed: 28 additions & 25 deletions

File tree

src/agentic_data_contracts/core/schema.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,8 @@ class ResultCheck(BaseModel):
4444

4545

4646
class SemanticRule(BaseModel):
47+
model_config = {"extra": "forbid"}
48+
4749
name: str
4850
description: str
4951
enforcement: Enforcement

src/agentic_data_contracts/tools/factory.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -269,7 +269,10 @@ async def run_query(args: dict[str, Any]) -> dict[str, Any]:
269269
)
270270
return _text_response(msg)
271271

272-
# Record estimated cost from EXPLAIN
272+
# Record estimated cost from EXPLAIN — charged before execution because
273+
# the cost budget tracks database resource consumption, not successful
274+
# operations. Even if result checks later block the output, the database
275+
# work was performed.
273276
if vresult.estimated_cost_usd is not None:
274277
session.record_cost(vresult.estimated_cost_usd)
275278

src/agentic_data_contracts/validation/checkers.py

Lines changed: 9 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,6 @@
1313
@dataclass
1414
class CheckResult:
1515
passed: bool
16-
severity: str # "block" | "warn" | "log"
1716
message: str
1817

1918

@@ -45,10 +44,9 @@ def check_ast(self, ast: exp.Expression, contract: DataContract) -> CheckResult:
4544
if disallowed:
4645
return CheckResult(
4746
passed=False,
48-
severity="block",
4947
message=f"Tables not in allowlist: {', '.join(sorted(disallowed))}",
5048
)
51-
return CheckResult(passed=True, severity="block", message="")
49+
return CheckResult(passed=True, message="")
5250

5351

5452
class OperationBlocklistChecker:
@@ -68,7 +66,6 @@ def check_ast(self, ast: exp.Expression, contract: DataContract) -> CheckResult:
6866
if isinstance(ast, expr_type) and op_name in forbidden:
6967
return CheckResult(
7068
passed=False,
71-
severity="block",
7269
message=f"Forbidden operation: {op_name}",
7370
)
7471

@@ -82,11 +79,10 @@ def check_ast(self, ast: exp.Expression, contract: DataContract) -> CheckResult:
8279
):
8380
return CheckResult(
8481
passed=False,
85-
severity="block",
8682
message="Forbidden operation: TRUNCATE",
8783
)
8884

89-
return CheckResult(passed=True, severity="block", message="")
85+
return CheckResult(passed=True, message="")
9086

9187

9288
class NoSelectStarChecker:
@@ -96,10 +92,9 @@ def check_ast(self, ast: exp.Expression) -> CheckResult:
9692
if any(ast.find_all(exp.Star)):
9793
return CheckResult(
9894
passed=False,
99-
severity="block",
10095
message="SELECT * is not allowed — specify explicit columns",
10196
)
102-
return CheckResult(passed=True, severity="block", message="")
97+
return CheckResult(passed=True, message="")
10398

10499

105100
class RequiredFilterChecker:
@@ -117,10 +112,9 @@ def check_ast(self, ast: exp.Expression) -> CheckResult:
117112
if self.column.lower() not in where_columns:
118113
return CheckResult(
119114
passed=False,
120-
severity="block",
121115
message=f"Missing required filter: {self.column}",
122116
)
123-
return CheckResult(passed=True, severity="block", message="")
117+
return CheckResult(passed=True, message="")
124118

125119

126120
class BlockedColumnsChecker:
@@ -138,7 +132,6 @@ def check_ast(self, ast: exp.Expression) -> CheckResult:
138132
if any(ast.find_all(exp.Star)):
139133
return CheckResult(
140134
passed=False,
141-
severity="block",
142135
message=(
143136
"SELECT * may expose blocked columns: "
144137
f"{', '.join(sorted(self.blocked))}"
@@ -155,10 +148,9 @@ def check_ast(self, ast: exp.Expression) -> CheckResult:
155148
if found:
156149
return CheckResult(
157150
passed=False,
158-
severity="block",
159151
message=f"Blocked columns in SELECT: {', '.join(sorted(found))}",
160152
)
161-
return CheckResult(passed=True, severity="block", message="")
153+
return CheckResult(passed=True, message="")
162154

163155

164156
class RequireLimitChecker:
@@ -168,10 +160,9 @@ def check_ast(self, ast: exp.Expression) -> CheckResult:
168160
if not list(ast.find_all(exp.Limit)):
169161
return CheckResult(
170162
passed=False,
171-
severity="block",
172163
message="Query must include a LIMIT clause",
173164
)
174-
return CheckResult(passed=True, severity="block", message="")
165+
return CheckResult(passed=True, message="")
175166

176167

177168
class MaxJoinsChecker:
@@ -185,12 +176,11 @@ def check_ast(self, ast: exp.Expression) -> CheckResult:
185176
if join_count > self.max_joins:
186177
return CheckResult(
187178
passed=False,
188-
severity="block",
189179
message=(
190180
f"Query has {join_count} JOINs, exceeds maximum of {self.max_joins}"
191181
),
192182
)
193-
return CheckResult(passed=True, severity="block", message="")
183+
return CheckResult(passed=True, message="")
194184

195185

196186
class ResultCheckRunner:
@@ -219,7 +209,6 @@ def check_results(self, columns: list[str], rows: list[tuple]) -> CheckResult:
219209
if self.min_rows is not None and row_count < self.min_rows:
220210
return CheckResult(
221211
passed=False,
222-
severity="block",
223212
message=(
224213
f"Rule '{self.rule_name}': query returned {row_count} rows, "
225214
f"minimum is {self.min_rows}"
@@ -228,7 +217,6 @@ def check_results(self, columns: list[str], rows: list[tuple]) -> CheckResult:
228217
if self.max_rows is not None and row_count > self.max_rows:
229218
return CheckResult(
230219
passed=False,
231-
severity="block",
232220
message=(
233221
f"Rule '{self.rule_name}': query returned {row_count} rows, "
234222
f"maximum is {self.max_rows}"
@@ -239,15 +227,14 @@ def check_results(self, columns: list[str], rows: list[tuple]) -> CheckResult:
239227
col_lower = {c.lower(): i for i, c in enumerate(columns)}
240228
idx = col_lower.get(self.column.lower())
241229
if idx is None:
242-
return CheckResult(passed=True, severity="block", message="")
230+
return CheckResult(passed=True, message="")
243231

244232
values = [row[idx] for row in rows]
245233

246234
if self.not_null and any(v is None for v in values):
247235
null_count = sum(1 for v in values if v is None)
248236
return CheckResult(
249237
passed=False,
250-
severity="block",
251238
message=(
252239
f"Rule '{self.rule_name}': column '{self.column}' "
253240
f"contains {null_count} null values"
@@ -265,7 +252,6 @@ def check_results(self, columns: list[str], rows: list[tuple]) -> CheckResult:
265252
if actual_min < self.min_value:
266253
return CheckResult(
267254
passed=False,
268-
severity="block",
269255
message=(
270256
f"Rule '{self.rule_name}': column '{self.column}' "
271257
f"min value {actual_min} "
@@ -277,11 +263,10 @@ def check_results(self, columns: list[str], rows: list[tuple]) -> CheckResult:
277263
if actual_max > self.max_value:
278264
return CheckResult(
279265
passed=False,
280-
severity="block",
281266
message=(
282267
f"Rule '{self.rule_name}': column '{self.column}' "
283268
f"max value {actual_max} exceeds limit {self.max_value}"
284269
),
285270
)
286271

287-
return CheckResult(passed=True, severity="block", message="")
272+
return CheckResult(passed=True, message="")

tests/test_core/test_schema.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -118,6 +118,19 @@ def test_rule_rejects_both_checks() -> None:
118118
)
119119

120120

121+
def test_old_filter_column_rejected() -> None:
122+
"""Old YAML with filter_column should fail loudly, not silently lose enforcement."""
123+
with pytest.raises(ValueError, match="extra"):
124+
SemanticRule.model_validate(
125+
{
126+
"name": "tenant_filter",
127+
"description": "Must filter by tenant_id",
128+
"enforcement": "block",
129+
"filter_column": "tenant_id",
130+
}
131+
)
132+
133+
121134
def test_advisory_rule_no_checks() -> None:
122135
"""Rules with neither check are advisory — shown in prompt only."""
123136
rule = SemanticRule(

0 commit comments

Comments
 (0)