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