File size: 2,195 Bytes
939c0c0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
"""
Unit tests for SQL Agent, AST/Regex security guardrails, and read-only validator.
"""
from __future__ import annotations

import pytest
from agents.sql_agent import is_safe_sql, sql_node
from agents.state import CopilotState


def test_is_safe_sql_valid_selects():
    """Verify valid SELECT queries pass security check."""
    q1 = "SELECT * FROM users WHERE role = 'admin'"
    safe, msg = is_safe_sql(q1)
    assert safe is True

    q2 = "WITH total_sales AS (SELECT SUM(amount) as s FROM orders) SELECT * FROM total_sales"
    safe, msg = is_safe_sql(q2)
    assert safe is True


def test_is_safe_sql_blocks_dml_ddl():
    """Verify mutating queries (DROP, DELETE, UPDATE, INSERT) are blocked."""
    forbidden_queries = [
        "DROP TABLE users;",
        "DELETE FROM audit_logs;",
        "UPDATE users SET role = 'admin';",
        "INSERT INTO tenants (name) VALUES ('Hacked');",
        "ALTER TABLE users ADD COLUMN secret text;",
        "TRUNCATE TABLE documents;"
    ]

    for q in forbidden_queries:
        safe, msg = is_safe_sql(q)
        assert safe is False, f"Failed to block: {q}"
        assert "Security Violation" in msg


from unittest.mock import patch, AsyncMock
from app.services.llm_gateway import LLMResponse


@pytest.mark.asyncio
async def test_sql_node_execution():
    """Test SQL Node execution flow with safe query."""
    state: CopilotState = {
        "query": "How many users are in the system?",
        "user_role": "admin",
        "tenant_id": "test-tenant"
    }

    mock_resp = LLMResponse(
        content="SELECT * FROM users",
        provider="groq",
        model="llama3-8b-8192",
        tokens_used=10,
        latency_ms=50.0
    )

    with patch("agents.sql_agent.llm_gateway.generate", new_callable=AsyncMock, return_value=mock_resp), \
         patch("agents.sql_agent.get_database_schema", new_callable=AsyncMock, return_value="Table: users(id, email)"), \
         patch("agents.sql_agent.execute_sql_query", new_callable=AsyncMock, return_value=(["count"], [{"count": 10}])):
        res_state = await sql_node(state)
        assert "retrieved_chunks" in res_state
        assert res_state.get("active_agent") == "sql"