| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284 |
- import pytest
- from data_query_agent.sql_guard import SQLGuard, SQLValidationError
- def test_allows_one_select_and_tracks_cte_sources() -> None:
- guard = SQLGuard(frozenset({"loghubods"}))
- refs = guard.validate(
- "WITH daily AS (SELECT uid FROM loghubods.events WHERE ds='20260810') SELECT count(*) FROM daily"
- )
- assert [(ref.project, ref.name) for ref in refs] == [("loghubods", "events")]
- @pytest.mark.parametrize(
- "sql",
- [
- "DELETE FROM loghubods.events WHERE ds='20260810'",
- "SELECT * FROM loghubods.events; SELECT * FROM loghubods.users",
- "CREATE TABLE x AS SELECT 1",
- ],
- )
- def test_rejects_non_read_only_or_multiple_statements(sql: str) -> None:
- with pytest.raises(SQLValidationError):
- SQLGuard(frozenset({"loghubods"})).validate(sql)
- def test_rejects_cross_project() -> None:
- with pytest.raises(SQLValidationError, match="跨项目"):
- SQLGuard(frozenset({"loghubods"})).validate("SELECT * FROM other_project.events")
- def test_output_labels_require_chinese_business_meanings() -> None:
- invalid = """
- SELECT stat_date, dau, risk_scene, risk_user_uv
- FROM loghubods.risk_result
- WHERE dt='20260812'
- """
- with pytest.raises(SQLValidationError, match="中文业务含义"):
- SQLGuard.validate_output_labels(invalid)
- valid = """
- SELECT stat_date AS `日期`, dau AS DAU,
- risk_scene AS `风险场景`, risk_user_uv AS `风险用户UV`
- FROM loghubods.risk_result
- WHERE dt='20260812'
- """
- SQLGuard.validate_output_labels(valid)
- def test_output_labels_allow_common_time_and_metric_abbreviations() -> None:
- SQLGuard.validate_output_labels(
- "SELECT date, time, COUNT(*) AS DAU, pv, uv "
- "FROM loghubods.metrics WHERE dt='20260812' GROUP BY date, time, pv, uv"
- )
- def test_partition_predicate_is_required() -> None:
- with pytest.raises(SQLValidationError, match="分区"):
- SQLGuard.validate_partition_predicates(
- "SELECT uid FROM events WHERE uid > 0", {"events": ["ds"]}
- )
- SQLGuard.validate_partition_predicates(
- "SELECT uid FROM events WHERE ds BETWEEN '20260801' AND '20260810'", {"events": ["ds"]}
- )
- def test_each_partitioned_join_source_must_be_filtered() -> None:
- partitions = {"events": ["ds"], "users": ["ds"]}
- with pytest.raises(SQLValidationError, match="users"):
- SQLGuard.validate_partition_predicates(
- "SELECT a.uid FROM events a JOIN users b ON a.uid=b.uid WHERE a.ds='20260810'",
- partitions,
- )
- SQLGuard.validate_partition_predicates(
- "SELECT a.uid FROM events a JOIN users b ON a.uid=b.uid AND b.ds='20260810' WHERE a.ds='20260810'",
- partitions,
- )
- def test_unqualified_partition_columns_are_valid_inside_single_source_ctes() -> None:
- sql = """
- WITH dau AS (
- SELECT COUNT(*) FROM loghubods.useractive_log_per5min
- WHERE dt LIKE '20260812%'
- ), video AS (
- SELECT COUNT(*) FROM loghubods.video_action_log_flow
- WHERE year='2026' AND month='08' AND dt='20260812' AND hh='16'
- ), shares AS (
- SELECT COUNT(*) FROM loghubods.user_share_log_per5min
- WHERE dt LIKE '20260812%'
- )
- SELECT * FROM dau CROSS JOIN video CROSS JOIN shares
- """
- SQLGuard.validate_partition_predicates(
- sql,
- {
- "loghubods.useractive_log_per5min": ["dt"],
- "loghubods.video_action_log_flow": ["year", "month", "dt", "hh"],
- "loghubods.user_share_log_per5min": ["dt"],
- },
- )
- def test_every_repeated_partitioned_source_must_have_its_own_filter() -> None:
- sql = """
- WITH source AS (
- SELECT shareid FROM loghubods.user_share_log_per5min
- WHERE dt LIKE '20260812%' AND topic='share'
- ), click AS (
- SELECT shareid FROM loghubods.user_share_log_per5min
- WHERE topic='click'
- )
- SELECT * FROM source JOIN click USING (shareid)
- """
- with pytest.raises(SQLValidationError, match="user_share_log_per5min"):
- SQLGuard.validate_partition_predicates(
- sql, {"loghubods.user_share_log_per5min": ["dt"]}
- )
- def test_video_action_applet_partition_requires_dt_but_not_business() -> None:
- partitions = {"loghubods.video_action_log_applet": ["dt", "business"]}
- SQLGuard.validate_partition_predicates(
- "SELECT mid FROM loghubods.video_action_log_applet v "
- "WHERE v.dt='20260810' AND v.businesstype='videoView'",
- partitions,
- )
- with pytest.raises(SQLValidationError, match=r"video_action_log_applet\(dt\)"):
- SQLGuard.validate_partition_predicates(
- "SELECT mid FROM loghubods.video_action_log_applet v WHERE v.businesstype='videoView'",
- partitions,
- )
- def test_offline_product_efficiency_requires_extparams_root_session() -> None:
- bad = """
- SELECT SUBSTR(u.rootsessionid, LENGTH(u.rootsessionid) - 2, 1) bucket,
- COUNT(DISTINCT u.machinecode) dau
- FROM loghubods.useractive_log u
- WHERE u.dt='20260808'
- GROUP BY SUBSTR(u.rootsessionid, LENGTH(u.rootsessionid) - 2, 1)
- """
- with pytest.raises(SQLValidationError, match="extparams.*rootSessionId"):
- SQLGuard.validate_product_efficiency_contract(bad, "offline")
- good = """
- SELECT SUBSTR(root_session_id, LENGTH(root_session_id) - 2, 1) bucket,
- COUNT(DISTINCT machinecode) dau
- FROM (
- SELECT machinecode, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
- FROM loghubods.useractive_log
- WHERE dt='20260808'
- ) u
- GROUP BY SUBSTR(root_session_id, LENGTH(root_session_id) - 2, 1)
- """
- SQLGuard.validate_product_efficiency_contract(good, "offline")
- def test_offline_total_only_product_efficiency_does_not_require_bucketing() -> None:
- sql = """
- SELECT COUNT(DISTINCT machinecode) AS dau
- FROM loghubods.useractive_log
- WHERE dt='20260810' AND apptype='0' AND businesstype='path'
- """
- SQLGuard.validate_product_efficiency_contract(sql, "offline", bucketed=False)
- def test_realtime_product_efficiency_does_not_use_offline_contract() -> None:
- SQLGuard.validate_product_efficiency_contract("SELECT 1", "realtime")
- def test_realtime_product_efficiency_per5min_sources_require_day_prefix() -> None:
- invalid = """
- WITH dau AS (
- SELECT COUNT(*) FROM loghubods.useractive_log_per5min
- WHERE dt='20260812'
- ), shares AS (
- SELECT COUNT(*) FROM loghubods.user_share_log_per5min
- WHERE dt='20260812'
- )
- SELECT * FROM dau CROSS JOIN shares
- """
- with pytest.raises(SQLValidationError, match="dt LIKE"):
- SQLGuard.validate_product_efficiency_contract(invalid, "realtime")
- valid = invalid.replace("dt='20260812'", "dt LIKE '20260812%'")
- SQLGuard.validate_product_efficiency_contract(valid, "realtime")
- def test_realtime_all_version_product_efficiency_prefers_video_per5min() -> None:
- sql = """
- SELECT COUNT(*)
- FROM loghubods.video_action_log_flow v
- WHERE v.year='2026' AND v.month='08' AND v.dt='12'
- AND v.apptype='0' AND v.businesstype='videoView'
- """
- with pytest.raises(SQLValidationError, match="全版本.*video_action_log_per5min"):
- SQLGuard.validate_product_efficiency_contract(
- sql, "realtime", bucketed=False, version="all"
- )
- def test_realtime_video_per5min_requires_day_prefix_partition() -> None:
- invalid = """
- SELECT COUNT(*)
- FROM loghubods.video_action_log_per5min v
- WHERE v.dt='20260812' AND v.apptype='0' AND v.businesstype='videoView'
- """
- with pytest.raises(SQLValidationError, match="dt LIKE"):
- SQLGuard.validate_product_efficiency_contract(
- invalid, "realtime", bucketed=False, version="all"
- )
- valid = invalid.replace("v.dt='20260812'", "v.dt LIKE '20260812%'")
- SQLGuard.validate_product_efficiency_contract(
- valid, "realtime", bucketed=False, version="all"
- )
- def test_realtime_video_flow_uses_day_not_full_date_partition() -> None:
- invalid = """
- SELECT COUNT(*)
- FROM loghubods.video_action_log_flow v
- WHERE v.year='2026' AND v.month='08' AND v.dt='20260812' AND v.hh='16'
- AND v.apptype='0' AND v.versioncode='1578' AND v.businesstype='videoView'
- """
- with pytest.raises(SQLValidationError, match="dt='DD'"):
- SQLGuard.validate_product_efficiency_contract(
- invalid, "realtime", bucketed=False, version="1578"
- )
- valid = invalid.replace("v.dt='20260812'", "v.dt='12'")
- SQLGuard.validate_product_efficiency_contract(
- valid, "realtime", bucketed=False, version="1578"
- )
- def test_realtime_product_efficiency_rejects_video_business_filter() -> None:
- sql = """
- SELECT mid, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
- FROM loghubods.video_action_log_flow v
- WHERE v.year='2026' AND v.month='08' AND v.dt='12'
- AND v.business='applet'
- AND v.apptype='4' AND v.businesstype='videoView'
- """
- with pytest.raises(SQLValidationError, match="禁止使用 business"):
- SQLGuard.validate_product_efficiency_contract(sql, "realtime")
- def test_realtime_video_per5min_rejects_business_filter() -> None:
- sql = """
- SELECT COUNT(*)
- FROM loghubods.video_action_log_per5min v
- WHERE v.dt LIKE '20260812%'
- AND v.business='applet'
- AND v.apptype='0' AND v.businesstype='videoView'
- """
- with pytest.raises(SQLValidationError, match="禁止使用 business"):
- SQLGuard.validate_product_efficiency_contract(
- sql, "realtime", bucketed=False, version="all"
- )
- @pytest.mark.parametrize(
- "business_filter",
- ["v.business='applet'", "v.business IN ('videoView', 'videoPlay', 'videoShareFriend')"],
- )
- def test_offline_product_efficiency_rejects_video_business_filter(business_filter: str) -> None:
- sql = f"""
- WITH dau AS (
- SELECT machinecode, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
- FROM loghubods.useractive_log
- WHERE dt='20260810'
- ), video AS (
- SELECT mid, GET_JSON_OBJECT(extparams, '$.rootSessionId') root_session_id
- FROM loghubods.video_action_log_applet v
- WHERE v.dt='20260810'
- AND {business_filter}
- AND v.businesstype IN ('videoView', 'videoPlay', 'videoShareFriend')
- )
- SELECT COUNT(*) FROM dau JOIN video ON dau.root_session_id=video.root_session_id
- """
- with pytest.raises(SQLValidationError, match="禁止使用 business"):
- SQLGuard.validate_product_efficiency_contract(sql, "offline")
|