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")