test_sql_guard.py 9.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259
  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_unqualified_partition_columns_are_valid_inside_single_source_ctes() -> None:
  43. sql = """
  44. WITH dau AS (
  45. SELECT COUNT(*) FROM loghubods.useractive_log_per5min
  46. WHERE dt LIKE '20260812%'
  47. ), video AS (
  48. SELECT COUNT(*) FROM loghubods.video_action_log_flow
  49. WHERE year='2026' AND month='08' AND dt='20260812' AND hh='16'
  50. ), shares AS (
  51. SELECT COUNT(*) FROM loghubods.user_share_log_per5min
  52. WHERE dt LIKE '20260812%'
  53. )
  54. SELECT * FROM dau CROSS JOIN video CROSS JOIN shares
  55. """
  56. SQLGuard.validate_partition_predicates(
  57. sql,
  58. {
  59. "loghubods.useractive_log_per5min": ["dt"],
  60. "loghubods.video_action_log_flow": ["year", "month", "dt", "hh"],
  61. "loghubods.user_share_log_per5min": ["dt"],
  62. },
  63. )
  64. def test_every_repeated_partitioned_source_must_have_its_own_filter() -> None:
  65. sql = """
  66. WITH source AS (
  67. SELECT shareid FROM loghubods.user_share_log_per5min
  68. WHERE dt LIKE '20260812%' AND topic='share'
  69. ), click AS (
  70. SELECT shareid FROM loghubods.user_share_log_per5min
  71. WHERE topic='click'
  72. )
  73. SELECT * FROM source JOIN click USING (shareid)
  74. """
  75. with pytest.raises(SQLValidationError, match="user_share_log_per5min"):
  76. SQLGuard.validate_partition_predicates(
  77. sql, {"loghubods.user_share_log_per5min": ["dt"]}
  78. )
  79. def test_video_action_applet_partition_requires_dt_but_not_business() -> None:
  80. partitions = {"loghubods.video_action_log_applet": ["dt", "business"]}
  81. SQLGuard.validate_partition_predicates(
  82. "SELECT mid FROM loghubods.video_action_log_applet v "
  83. "WHERE v.dt='20260810' AND v.businesstype='videoView'",
  84. partitions,
  85. )
  86. with pytest.raises(SQLValidationError, match=r"video_action_log_applet\(dt\)"):
  87. SQLGuard.validate_partition_predicates(
  88. "SELECT mid FROM loghubods.video_action_log_applet v WHERE v.businesstype='videoView'",
  89. partitions,
  90. )
  91. def test_offline_product_efficiency_requires_extparams_root_session() -> None:
  92. bad = """
  93. SELECT SUBSTR(u.rootsessionid, LENGTH(u.rootsessionid) - 2, 1) bucket,
  94. COUNT(DISTINCT u.machinecode) dau
  95. FROM loghubods.useractive_log u
  96. WHERE u.dt='20260808'
  97. GROUP BY SUBSTR(u.rootsessionid, LENGTH(u.rootsessionid) - 2, 1)
  98. """
  99. with pytest.raises(SQLValidationError, match="extparams.*rootSessionId"):
  100. SQLGuard.validate_product_efficiency_contract(bad, "offline")
  101. good = """
  102. SELECT SUBSTR(root_session_id, LENGTH(root_session_id) - 2, 1) bucket,
  103. COUNT(DISTINCT machinecode) dau
  104. FROM (
  105. SELECT machinecode, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  106. FROM loghubods.useractive_log
  107. WHERE dt='20260808'
  108. ) u
  109. GROUP BY SUBSTR(root_session_id, LENGTH(root_session_id) - 2, 1)
  110. """
  111. SQLGuard.validate_product_efficiency_contract(good, "offline")
  112. def test_offline_total_only_product_efficiency_does_not_require_bucketing() -> None:
  113. sql = """
  114. SELECT COUNT(DISTINCT machinecode) AS dau
  115. FROM loghubods.useractive_log
  116. WHERE dt='20260810' AND apptype='0' AND businesstype='path'
  117. """
  118. SQLGuard.validate_product_efficiency_contract(sql, "offline", bucketed=False)
  119. def test_realtime_product_efficiency_does_not_use_offline_contract() -> None:
  120. SQLGuard.validate_product_efficiency_contract("SELECT 1", "realtime")
  121. def test_realtime_product_efficiency_per5min_sources_require_day_prefix() -> None:
  122. invalid = """
  123. WITH dau AS (
  124. SELECT COUNT(*) FROM loghubods.useractive_log_per5min
  125. WHERE dt='20260812'
  126. ), shares AS (
  127. SELECT COUNT(*) FROM loghubods.user_share_log_per5min
  128. WHERE dt='20260812'
  129. )
  130. SELECT * FROM dau CROSS JOIN shares
  131. """
  132. with pytest.raises(SQLValidationError, match="dt LIKE"):
  133. SQLGuard.validate_product_efficiency_contract(invalid, "realtime")
  134. valid = invalid.replace("dt='20260812'", "dt LIKE '20260812%'")
  135. SQLGuard.validate_product_efficiency_contract(valid, "realtime")
  136. def test_realtime_all_version_product_efficiency_prefers_video_per5min() -> None:
  137. sql = """
  138. SELECT COUNT(*)
  139. FROM loghubods.video_action_log_flow v
  140. WHERE v.year='2026' AND v.month='08' AND v.dt='12'
  141. AND v.apptype='0' AND v.businesstype='videoView'
  142. """
  143. with pytest.raises(SQLValidationError, match="全版本.*video_action_log_per5min"):
  144. SQLGuard.validate_product_efficiency_contract(
  145. sql, "realtime", bucketed=False, version="all"
  146. )
  147. def test_realtime_video_per5min_requires_day_prefix_partition() -> None:
  148. invalid = """
  149. SELECT COUNT(*)
  150. FROM loghubods.video_action_log_per5min v
  151. WHERE v.dt='20260812' AND v.apptype='0' AND v.businesstype='videoView'
  152. """
  153. with pytest.raises(SQLValidationError, match="dt LIKE"):
  154. SQLGuard.validate_product_efficiency_contract(
  155. invalid, "realtime", bucketed=False, version="all"
  156. )
  157. valid = invalid.replace("v.dt='20260812'", "v.dt LIKE '20260812%'")
  158. SQLGuard.validate_product_efficiency_contract(
  159. valid, "realtime", bucketed=False, version="all"
  160. )
  161. def test_realtime_video_flow_uses_day_not_full_date_partition() -> None:
  162. invalid = """
  163. SELECT COUNT(*)
  164. FROM loghubods.video_action_log_flow v
  165. WHERE v.year='2026' AND v.month='08' AND v.dt='20260812' AND v.hh='16'
  166. AND v.apptype='0' AND v.versioncode='1578' AND v.businesstype='videoView'
  167. """
  168. with pytest.raises(SQLValidationError, match="dt='DD'"):
  169. SQLGuard.validate_product_efficiency_contract(
  170. invalid, "realtime", bucketed=False, version="1578"
  171. )
  172. valid = invalid.replace("v.dt='20260812'", "v.dt='12'")
  173. SQLGuard.validate_product_efficiency_contract(
  174. valid, "realtime", bucketed=False, version="1578"
  175. )
  176. def test_realtime_product_efficiency_rejects_video_business_filter() -> None:
  177. sql = """
  178. SELECT mid, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  179. FROM loghubods.video_action_log_flow v
  180. WHERE v.year='2026' AND v.month='08' AND v.dt='12'
  181. AND v.business='applet'
  182. AND v.apptype='4' AND v.businesstype='videoView'
  183. """
  184. with pytest.raises(SQLValidationError, match="禁止使用 business"):
  185. SQLGuard.validate_product_efficiency_contract(sql, "realtime")
  186. def test_realtime_video_per5min_rejects_business_filter() -> None:
  187. sql = """
  188. SELECT COUNT(*)
  189. FROM loghubods.video_action_log_per5min v
  190. WHERE v.dt LIKE '20260812%'
  191. AND v.business='applet'
  192. AND v.apptype='0' AND v.businesstype='videoView'
  193. """
  194. with pytest.raises(SQLValidationError, match="禁止使用 business"):
  195. SQLGuard.validate_product_efficiency_contract(
  196. sql, "realtime", bucketed=False, version="all"
  197. )
  198. @pytest.mark.parametrize(
  199. "business_filter",
  200. ["v.business='applet'", "v.business IN ('videoView', 'videoPlay', 'videoShareFriend')"],
  201. )
  202. def test_offline_product_efficiency_rejects_video_business_filter(business_filter: str) -> None:
  203. sql = f"""
  204. WITH dau AS (
  205. SELECT machinecode, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  206. FROM loghubods.useractive_log
  207. WHERE dt='20260810'
  208. ), video AS (
  209. SELECT mid, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  210. FROM loghubods.video_action_log_applet v
  211. WHERE v.dt='20260810'
  212. AND {business_filter}
  213. AND v.businesstype IN ('videoView', 'videoPlay', 'videoShareFriend')
  214. )
  215. SELECT COUNT(*) FROM dau JOIN video ON dau.root_session_id=video.root_session_id
  216. """
  217. with pytest.raises(SQLValidationError, match="禁止使用 business"):
  218. SQLGuard.validate_product_efficiency_contract(sql, "offline")