Files
group_xinghuo_jinrong/tests/test_sql_guard.py
T

84 lines
3.1 KiB
Python
Raw Normal View History

"""sql_guard 单元测试(unittest,无需外部服务)。"""
import unittest
from app.service.sql_guard import (
SqlGuardError,
extract_tables,
inject_ownership,
validate,
)
class TestSqlGuard(unittest.TestCase):
def test_select_allowed(self):
r = validate("SELECT risk_code, COUNT(*) AS c FROM core_customer GROUP BY risk_code", "full")
self.assertTrue(r.allowed)
self.assertIn("core_customer", r.tables)
def test_with_select_allowed(self):
r = validate("WITH t AS (SELECT * FROM core_holding) SELECT * FROM t", "full")
self.assertTrue(r.allowed)
def test_insert_rejected(self):
with self.assertRaises(SqlGuardError) as cm:
validate("INSERT INTO core_customer VALUES (1)", "full")
self.assertEqual(cm.exception.error_code, "SQL_NOT_SELECT")
def test_multi_statement_rejected(self):
with self.assertRaises(SqlGuardError) as cm:
validate("SELECT 1; DROP TABLE core_customer;", "full")
self.assertEqual(cm.exception.error_code, "SQL_MULTI_STATEMENT")
def test_drop_rejected(self):
with self.assertRaises(SqlGuardError):
validate("SELECT * FROM core_customer; DROP TABLE core_customer", "full")
def test_unknown_table_rejected(self):
with self.assertRaises(SqlGuardError) as cm:
validate("SELECT * FROM mysql.user", "full")
self.assertEqual(cm.exception.error_code, "SQL_TABLE_NOT_ALLOWED")
def test_advisor_out_of_scope_rejected(self):
with self.assertRaises(SqlGuardError) as cm:
validate(
"SELECT * FROM core_holding WHERE customer_id = 'CUST-1004'",
"assigned", ["CUST-1001"],
)
self.assertEqual(cm.exception.error_code, "AUTH_403_NOT_ASSIGNED")
def test_advisor_in_scope_allowed(self):
r = validate(
"SELECT * FROM core_holding WHERE customer_id = 'CUST-1001'",
"assigned", ["CUST-1001", "CUST-1002"],
)
self.assertTrue(r.allowed)
self.assertTrue(r.has_customer_detail)
def test_ops_customer_detail_rejected(self):
with self.assertRaises(SqlGuardError) as cm:
validate("SELECT customer_id, market_value FROM core_holding", "aggregate")
self.assertEqual(cm.exception.error_code, "AUTH_403_SCOPE")
def test_ops_aggregate_allowed(self):
r = validate("SELECT COUNT(DISTINCT customer_id) AS cnt FROM core_holding", "aggregate")
self.assertTrue(r.allowed)
def test_ops_specific_customer_rejected(self):
with self.assertRaises(SqlGuardError):
validate("SELECT * FROM core_holding WHERE customer_id='CUST-1001'", "aggregate")
def test_inject_ownership(self):
out = inject_ownership("SELECT * FROM core_holding", ["CUST-1", "CUST-2"])
self.assertIn("CUST-1", out)
self.assertIn("WHERE customer_id IN", out)
def test_extract_tables(self):
self.assertEqual(
extract_tables("SELECT * FROM core_customer c JOIN core_holding h ON h.customer_id=c.customer_id"),
["core_customer", "core_holding"],
)
if __name__ == "__main__":
unittest.main()