File size: 1,413 Bytes
cf796c5
 
 
8d73fdc
cf796c5
 
f9e8eed
cf796c5
 
 
 
f9e8eed
cf796c5
 
 
f9e8eed
 
 
 
 
cf796c5
 
 
 
f9e8eed
 
 
 
 
cf796c5
 
 
 
f9e8eed
 
 
 
 
cf796c5
f9e8eed
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
import re
from src.state import AgentState

BLOCKED = r'\b(DROP|DELETE|INSERT|UPDATE|ALTER|TRUNCATE|CREATE|EXEC|GRANT|REVOKE|UNION)\b'
SYSTEM_TABLES = ["sqlite_master", "sqlite_sequence", "information_schema"]

def check_sql_safety(state: AgentState) -> dict:
    """Validate SQL query for safety before execution"""
    query = state.get("sql_query", "")

    if not query:
        return {"sql_result": {"error": "no_query"}}

    # Check for dangerous keywords
    if re.search(BLOCKED, query, re.IGNORECASE):
        return {
            "sql_query": None,
            "sql_result": {"error": "unsafe_query"},
            "error_log": state.get("error_log", []) + ["[sql_safety] Blocked dangerous SQL keyword"]
        }

    # Check for system table access
    for sys_table in SYSTEM_TABLES:
        if sys_table in query.lower():
            return {
                "sql_query": None,
                "sql_result": {"error": "system_table_access"},
                "error_log": state.get("error_log", []) + [f"[sql_safety] Blocked system table access: {sys_table}"]
            }

    # Ensure SELECT only
    stripped = query.strip().upper()
    if not stripped.startswith("SELECT"):
        return {
            "sql_query": None,
            "sql_result": {"error": "non_select"},
            "error_log": state.get("error_log", []) + ["[sql_safety] Non-SELECT query blocked"]
        }

    return {}