Download tests/unit/test_sql_agent.py from Punit1/Enterprise-AI-Copilot: direct link, hf CLI and curl.
- Browser
- Download file 2.2 kB
-
https://huggingface.co/spaces/Punit1/Enterprise-AI-Copilot/resolve/main/tests/unit/test_sql_agent.py
- Command line
-
hf download hf://spaces/Punit1/Enterprise-AI-Copilot/tests/unit/test_sql_agent.py
-
curl -L -o test_sql_agent.py https://huggingface.co/spaces/Punit1/Enterprise-AI-Copilot/resolve/main/tests/unit/test_sql_agent.py
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 | |
| 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" | |