test_sql_guard.py 5.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129
  1. import pytest
  2. from data_query_agent.sql_guard import SQLGuard, SQLValidationError
  3. def test_allows_one_select_and_tracks_cte_sources() -> None:
  4. guard = SQLGuard(frozenset({"loghubods"}))
  5. refs = guard.validate(
  6. "WITH daily AS (SELECT uid FROM loghubods.events WHERE ds='20260810') SELECT count(*) FROM daily"
  7. )
  8. assert [(ref.project, ref.name) for ref in refs] == [("loghubods", "events")]
  9. @pytest.mark.parametrize(
  10. "sql",
  11. [
  12. "DELETE FROM loghubods.events WHERE ds='20260810'",
  13. "SELECT * FROM loghubods.events; SELECT * FROM loghubods.users",
  14. "CREATE TABLE x AS SELECT 1",
  15. ],
  16. )
  17. def test_rejects_non_read_only_or_multiple_statements(sql: str) -> None:
  18. with pytest.raises(SQLValidationError):
  19. SQLGuard(frozenset({"loghubods"})).validate(sql)
  20. def test_rejects_cross_project() -> None:
  21. with pytest.raises(SQLValidationError, match="跨项目"):
  22. SQLGuard(frozenset({"loghubods"})).validate("SELECT * FROM other_project.events")
  23. def test_partition_predicate_is_required() -> None:
  24. with pytest.raises(SQLValidationError, match="分区"):
  25. SQLGuard.validate_partition_predicates(
  26. "SELECT uid FROM events WHERE uid > 0", {"events": ["ds"]}
  27. )
  28. SQLGuard.validate_partition_predicates(
  29. "SELECT uid FROM events WHERE ds BETWEEN '20260801' AND '20260810'", {"events": ["ds"]}
  30. )
  31. def test_each_partitioned_join_source_must_be_filtered() -> None:
  32. partitions = {"events": ["ds"], "users": ["ds"]}
  33. with pytest.raises(SQLValidationError, match="users"):
  34. SQLGuard.validate_partition_predicates(
  35. "SELECT a.uid FROM events a JOIN users b ON a.uid=b.uid WHERE a.ds='20260810'",
  36. partitions,
  37. )
  38. SQLGuard.validate_partition_predicates(
  39. "SELECT a.uid FROM events a JOIN users b ON a.uid=b.uid AND b.ds='20260810' WHERE a.ds='20260810'",
  40. partitions,
  41. )
  42. def test_video_action_applet_partition_requires_dt_but_not_business() -> None:
  43. partitions = {"loghubods.video_action_log_applet": ["dt", "business"]}
  44. SQLGuard.validate_partition_predicates(
  45. "SELECT mid FROM loghubods.video_action_log_applet v "
  46. "WHERE v.dt='20260810' AND v.businesstype='videoView'",
  47. partitions,
  48. )
  49. with pytest.raises(SQLValidationError, match=r"video_action_log_applet\(dt\)"):
  50. SQLGuard.validate_partition_predicates(
  51. "SELECT mid FROM loghubods.video_action_log_applet v WHERE v.businesstype='videoView'",
  52. partitions,
  53. )
  54. def test_offline_product_efficiency_requires_extparams_root_session() -> None:
  55. bad = """
  56. SELECT SUBSTR(u.rootsessionid, LENGTH(u.rootsessionid) - 2, 1) bucket,
  57. COUNT(DISTINCT u.machinecode) dau
  58. FROM loghubods.useractive_log u
  59. WHERE u.dt='20260808'
  60. GROUP BY SUBSTR(u.rootsessionid, LENGTH(u.rootsessionid) - 2, 1)
  61. """
  62. with pytest.raises(SQLValidationError, match="extparams.*rootSessionId"):
  63. SQLGuard.validate_product_efficiency_contract(bad, "offline")
  64. good = """
  65. SELECT SUBSTR(root_session_id, LENGTH(root_session_id) - 2, 1) bucket,
  66. COUNT(DISTINCT machinecode) dau
  67. FROM (
  68. SELECT machinecode, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  69. FROM loghubods.useractive_log
  70. WHERE dt='20260808'
  71. ) u
  72. GROUP BY SUBSTR(root_session_id, LENGTH(root_session_id) - 2, 1)
  73. """
  74. SQLGuard.validate_product_efficiency_contract(good, "offline")
  75. def test_realtime_product_efficiency_does_not_use_offline_contract() -> None:
  76. SQLGuard.validate_product_efficiency_contract("SELECT 1", "realtime")
  77. def test_realtime_product_efficiency_rejects_video_business_filter() -> None:
  78. sql = """
  79. SELECT mid, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  80. FROM loghubods.video_action_log_flow v
  81. WHERE v.year='2026' AND v.month='08' AND v.dt='12'
  82. AND v.business='applet'
  83. AND v.apptype='4' AND v.businesstype='videoView'
  84. """
  85. with pytest.raises(SQLValidationError, match="禁止使用 business"):
  86. SQLGuard.validate_product_efficiency_contract(sql, "realtime")
  87. @pytest.mark.parametrize(
  88. "business_filter",
  89. ["v.business='applet'", "v.business IN ('videoView', 'videoPlay', 'videoShareFriend')"],
  90. )
  91. def test_offline_product_efficiency_rejects_video_business_filter(business_filter: str) -> None:
  92. sql = f"""
  93. WITH dau AS (
  94. SELECT machinecode, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  95. FROM loghubods.useractive_log
  96. WHERE dt='20260810'
  97. ), video AS (
  98. SELECT mid, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  99. FROM loghubods.video_action_log_applet v
  100. WHERE v.dt='20260810'
  101. AND {business_filter}
  102. AND v.businesstype IN ('videoView', 'videoPlay', 'videoShareFriend')
  103. )
  104. SELECT COUNT(*) FROM dau JOIN video ON dau.root_session_id=video.root_session_id
  105. """
  106. with pytest.raises(SQLValidationError, match="禁止使用 business"):
  107. SQLGuard.validate_product_efficiency_contract(sql, "offline")