test_sql_guard.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284
  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_output_labels_require_chinese_business_meanings() -> None:
  24. invalid = """
  25. SELECT stat_date, dau, risk_scene, risk_user_uv
  26. FROM loghubods.risk_result
  27. WHERE dt='20260812'
  28. """
  29. with pytest.raises(SQLValidationError, match="中文业务含义"):
  30. SQLGuard.validate_output_labels(invalid)
  31. valid = """
  32. SELECT stat_date AS `日期`, dau AS DAU,
  33. risk_scene AS `风险场景`, risk_user_uv AS `风险用户UV`
  34. FROM loghubods.risk_result
  35. WHERE dt='20260812'
  36. """
  37. SQLGuard.validate_output_labels(valid)
  38. def test_output_labels_allow_common_time_and_metric_abbreviations() -> None:
  39. SQLGuard.validate_output_labels(
  40. "SELECT date, time, COUNT(*) AS DAU, pv, uv "
  41. "FROM loghubods.metrics WHERE dt='20260812' GROUP BY date, time, pv, uv"
  42. )
  43. def test_partition_predicate_is_required() -> None:
  44. with pytest.raises(SQLValidationError, match="分区"):
  45. SQLGuard.validate_partition_predicates(
  46. "SELECT uid FROM events WHERE uid > 0", {"events": ["ds"]}
  47. )
  48. SQLGuard.validate_partition_predicates(
  49. "SELECT uid FROM events WHERE ds BETWEEN '20260801' AND '20260810'", {"events": ["ds"]}
  50. )
  51. def test_each_partitioned_join_source_must_be_filtered() -> None:
  52. partitions = {"events": ["ds"], "users": ["ds"]}
  53. with pytest.raises(SQLValidationError, match="users"):
  54. SQLGuard.validate_partition_predicates(
  55. "SELECT a.uid FROM events a JOIN users b ON a.uid=b.uid WHERE a.ds='20260810'",
  56. partitions,
  57. )
  58. SQLGuard.validate_partition_predicates(
  59. "SELECT a.uid FROM events a JOIN users b ON a.uid=b.uid AND b.ds='20260810' WHERE a.ds='20260810'",
  60. partitions,
  61. )
  62. def test_unqualified_partition_columns_are_valid_inside_single_source_ctes() -> None:
  63. sql = """
  64. WITH dau AS (
  65. SELECT COUNT(*) FROM loghubods.useractive_log_per5min
  66. WHERE dt LIKE '20260812%'
  67. ), video AS (
  68. SELECT COUNT(*) FROM loghubods.video_action_log_flow
  69. WHERE year='2026' AND month='08' AND dt='20260812' AND hh='16'
  70. ), shares AS (
  71. SELECT COUNT(*) FROM loghubods.user_share_log_per5min
  72. WHERE dt LIKE '20260812%'
  73. )
  74. SELECT * FROM dau CROSS JOIN video CROSS JOIN shares
  75. """
  76. SQLGuard.validate_partition_predicates(
  77. sql,
  78. {
  79. "loghubods.useractive_log_per5min": ["dt"],
  80. "loghubods.video_action_log_flow": ["year", "month", "dt", "hh"],
  81. "loghubods.user_share_log_per5min": ["dt"],
  82. },
  83. )
  84. def test_every_repeated_partitioned_source_must_have_its_own_filter() -> None:
  85. sql = """
  86. WITH source AS (
  87. SELECT shareid FROM loghubods.user_share_log_per5min
  88. WHERE dt LIKE '20260812%' AND topic='share'
  89. ), click AS (
  90. SELECT shareid FROM loghubods.user_share_log_per5min
  91. WHERE topic='click'
  92. )
  93. SELECT * FROM source JOIN click USING (shareid)
  94. """
  95. with pytest.raises(SQLValidationError, match="user_share_log_per5min"):
  96. SQLGuard.validate_partition_predicates(
  97. sql, {"loghubods.user_share_log_per5min": ["dt"]}
  98. )
  99. def test_video_action_applet_partition_requires_dt_but_not_business() -> None:
  100. partitions = {"loghubods.video_action_log_applet": ["dt", "business"]}
  101. SQLGuard.validate_partition_predicates(
  102. "SELECT mid FROM loghubods.video_action_log_applet v "
  103. "WHERE v.dt='20260810' AND v.businesstype='videoView'",
  104. partitions,
  105. )
  106. with pytest.raises(SQLValidationError, match=r"video_action_log_applet\(dt\)"):
  107. SQLGuard.validate_partition_predicates(
  108. "SELECT mid FROM loghubods.video_action_log_applet v WHERE v.businesstype='videoView'",
  109. partitions,
  110. )
  111. def test_offline_product_efficiency_requires_extparams_root_session() -> None:
  112. bad = """
  113. SELECT SUBSTR(u.rootsessionid, LENGTH(u.rootsessionid) - 2, 1) bucket,
  114. COUNT(DISTINCT u.machinecode) dau
  115. FROM loghubods.useractive_log u
  116. WHERE u.dt='20260808'
  117. GROUP BY SUBSTR(u.rootsessionid, LENGTH(u.rootsessionid) - 2, 1)
  118. """
  119. with pytest.raises(SQLValidationError, match="extparams.*rootSessionId"):
  120. SQLGuard.validate_product_efficiency_contract(bad, "offline")
  121. good = """
  122. SELECT SUBSTR(root_session_id, LENGTH(root_session_id) - 2, 1) bucket,
  123. COUNT(DISTINCT machinecode) dau
  124. FROM (
  125. SELECT machinecode, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  126. FROM loghubods.useractive_log
  127. WHERE dt='20260808'
  128. ) u
  129. GROUP BY SUBSTR(root_session_id, LENGTH(root_session_id) - 2, 1)
  130. """
  131. SQLGuard.validate_product_efficiency_contract(good, "offline")
  132. def test_offline_total_only_product_efficiency_does_not_require_bucketing() -> None:
  133. sql = """
  134. SELECT COUNT(DISTINCT machinecode) AS dau
  135. FROM loghubods.useractive_log
  136. WHERE dt='20260810' AND apptype='0' AND businesstype='path'
  137. """
  138. SQLGuard.validate_product_efficiency_contract(sql, "offline", bucketed=False)
  139. def test_realtime_product_efficiency_does_not_use_offline_contract() -> None:
  140. SQLGuard.validate_product_efficiency_contract("SELECT 1", "realtime")
  141. def test_realtime_product_efficiency_per5min_sources_require_day_prefix() -> None:
  142. invalid = """
  143. WITH dau AS (
  144. SELECT COUNT(*) FROM loghubods.useractive_log_per5min
  145. WHERE dt='20260812'
  146. ), shares AS (
  147. SELECT COUNT(*) FROM loghubods.user_share_log_per5min
  148. WHERE dt='20260812'
  149. )
  150. SELECT * FROM dau CROSS JOIN shares
  151. """
  152. with pytest.raises(SQLValidationError, match="dt LIKE"):
  153. SQLGuard.validate_product_efficiency_contract(invalid, "realtime")
  154. valid = invalid.replace("dt='20260812'", "dt LIKE '20260812%'")
  155. SQLGuard.validate_product_efficiency_contract(valid, "realtime")
  156. def test_realtime_all_version_product_efficiency_prefers_video_per5min() -> None:
  157. sql = """
  158. SELECT COUNT(*)
  159. FROM loghubods.video_action_log_flow v
  160. WHERE v.year='2026' AND v.month='08' AND v.dt='12'
  161. AND v.apptype='0' AND v.businesstype='videoView'
  162. """
  163. with pytest.raises(SQLValidationError, match="全版本.*video_action_log_per5min"):
  164. SQLGuard.validate_product_efficiency_contract(
  165. sql, "realtime", bucketed=False, version="all"
  166. )
  167. def test_realtime_video_per5min_requires_day_prefix_partition() -> None:
  168. invalid = """
  169. SELECT COUNT(*)
  170. FROM loghubods.video_action_log_per5min v
  171. WHERE v.dt='20260812' AND v.apptype='0' AND v.businesstype='videoView'
  172. """
  173. with pytest.raises(SQLValidationError, match="dt LIKE"):
  174. SQLGuard.validate_product_efficiency_contract(
  175. invalid, "realtime", bucketed=False, version="all"
  176. )
  177. valid = invalid.replace("v.dt='20260812'", "v.dt LIKE '20260812%'")
  178. SQLGuard.validate_product_efficiency_contract(
  179. valid, "realtime", bucketed=False, version="all"
  180. )
  181. def test_realtime_video_flow_uses_day_not_full_date_partition() -> None:
  182. invalid = """
  183. SELECT COUNT(*)
  184. FROM loghubods.video_action_log_flow v
  185. WHERE v.year='2026' AND v.month='08' AND v.dt='20260812' AND v.hh='16'
  186. AND v.apptype='0' AND v.versioncode='1578' AND v.businesstype='videoView'
  187. """
  188. with pytest.raises(SQLValidationError, match="dt='DD'"):
  189. SQLGuard.validate_product_efficiency_contract(
  190. invalid, "realtime", bucketed=False, version="1578"
  191. )
  192. valid = invalid.replace("v.dt='20260812'", "v.dt='12'")
  193. SQLGuard.validate_product_efficiency_contract(
  194. valid, "realtime", bucketed=False, version="1578"
  195. )
  196. def test_realtime_product_efficiency_rejects_video_business_filter() -> None:
  197. sql = """
  198. SELECT mid, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  199. FROM loghubods.video_action_log_flow v
  200. WHERE v.year='2026' AND v.month='08' AND v.dt='12'
  201. AND v.business='applet'
  202. AND v.apptype='4' AND v.businesstype='videoView'
  203. """
  204. with pytest.raises(SQLValidationError, match="禁止使用 business"):
  205. SQLGuard.validate_product_efficiency_contract(sql, "realtime")
  206. def test_realtime_video_per5min_rejects_business_filter() -> None:
  207. sql = """
  208. SELECT COUNT(*)
  209. FROM loghubods.video_action_log_per5min v
  210. WHERE v.dt LIKE '20260812%'
  211. AND v.business='applet'
  212. AND v.apptype='0' AND v.businesstype='videoView'
  213. """
  214. with pytest.raises(SQLValidationError, match="禁止使用 business"):
  215. SQLGuard.validate_product_efficiency_contract(
  216. sql, "realtime", bucketed=False, version="all"
  217. )
  218. @pytest.mark.parametrize(
  219. "business_filter",
  220. ["v.business='applet'", "v.business IN ('videoView', 'videoPlay', 'videoShareFriend')"],
  221. )
  222. def test_offline_product_efficiency_rejects_video_business_filter(business_filter: str) -> None:
  223. sql = f"""
  224. WITH dau AS (
  225. SELECT machinecode, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  226. FROM loghubods.useractive_log
  227. WHERE dt='20260810'
  228. ), video AS (
  229. SELECT mid, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
  230. FROM loghubods.video_action_log_applet v
  231. WHERE v.dt='20260810'
  232. AND {business_filter}
  233. AND v.businesstype IN ('videoView', 'videoPlay', 'videoShareFriend')
  234. )
  235. SELECT COUNT(*) FROM dau JOIN video ON dau.root_session_id=video.root_session_id
  236. """
  237. with pytest.raises(SQLValidationError, match="禁止使用 business"):
  238. SQLGuard.validate_product_efficiency_contract(sql, "offline")