Enterprise-AI-Copilot / tests /unit /test_sql_agent.py
Punit1's picture
Initial commit
939c0c0
Raw
History Blame Contribute Delete
2.2 kB
"""
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"