| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129 |
- 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_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_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_realtime_product_efficiency_does_not_use_offline_contract() -> None:
- SQLGuard.validate_product_efficiency_contract("SELECT 1", "realtime")
- 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")
- @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")
|