🛠 MVP 仓库

← 返回 MVP 仓库
idea #37★ 6/102026-08-0118.3 KB

AgentEval Gate:AI 评测环境安全预检器

查看原始点子 →
37_mvp.pyPython
#!/usr/bin/env python3
"""AgentEval Gate: AI eval environment safety preflight checker."""

import argparse
import datetime as _datetime
import ipaddress
import json
import os
import re
import sys
import tempfile
import unittest
from urllib.parse import urlparse


VERSION = "0.1.0"
SEVERITY_RANK = {"info": 0, "low": 1, "medium": 2, "high": 3, "critical": 4}
BLOCKING_SEVERITIES = {"high", "critical"}
RESERVED_SUFFIXES = (".example", ".invalid", ".localhost", ".test")
RESERVED_EXACT = {"example", "invalid", "localhost", "test"}

SCHEMA = {
    "schema_version": "1",
    "required": ["prompt", "environment", "targets", "commands"],
    "properties": {
        "prompt": "string",
        "environment": {
            "network_enabled": "boolean",
            "browser_enabled": "boolean",
            "shell_enabled": "boolean",
            "publish_enabled": "boolean",
        },
        "targets": {
            "allowed_domains": "array[string]",
            "requested_urls": "array[string]",
        },
        "commands": "array[string]",
    },
}

SAFE_SAMPLE = {
    "schema_version": "1",
    "name": "safe-browser-eval",
    "prompt": "允许访问网络,但只能访问 https://shop.example;禁止发布软件。",
    "environment": {
        "network_enabled": True,
        "browser_enabled": True,
        "shell_enabled": False,
        "publish_enabled": False,
    },
    "targets": {
        "allowed_domains": ["shop.example"],
        "requested_urls": ["https://shop.example/catalog"],
    },
    "commands": [],
}

DANGEROUS_SAMPLE = {
    "schema_version": "1",
    "name": "dangerous-release-eval",
    "prompt": "这是完全离线的评测,不得访问互联网,也不要发布任何软件。",
    "environment": {
        "network_enabled": True,
        "browser_enabled": True,
        "shell_enabled": True,
        "publish_enabled": False,
    },
    "targets": {
        "allowed_domains": ["github.com", "*.example.com"],
        "requested_urls": ["https://api.github.com/repos/acme/demo"],
    },
    "commands": ["npm publish", "docker push acme/eval:latest"],
}

DENY_NETWORK_PATTERNS = [
    r"\b(?:no|without|disable(?:d)?)\s+(?:internet|network|web)\b",
    r"\b(?:offline|air[ -]?gapped)\b",
    r"(?:禁止|不得|不能|不可|不允许).{0,8}(?:联网|网络|互联网|外网|网页)",
    r"(?:离线|断网|无网络)(?:环境|评测|运行)?",
]
ALLOW_NETWORK_PATTERNS = [
    r"\b(?:may|can|allow(?:ed)? to|permit(?:ted)? to)\s+(?:access|use|browse).{0,16}(?:internet|network|web)\b",
    r"\b(?:internet|network|web)\s+(?:access\s+)?(?:is\s+)?(?:allowed|enabled|permitted)\b",
    r"(?:允许|可以|可|准许).{0,8}(?:联网|访问网络|访问互联网|访问网页|使用网络)",
    r"(?:网络|联网|互联网|外网).{0,6}(?:已开启|已启用|可用|允许)",
]
PUBLISH_PATTERNS = [
    ("npm_publish", r"(?:^|[;&|]\s*)npm\s+(?:publish|unpublish)\b"),
    ("pypi_upload", r"(?:^|[;&|]\s*)(?:python\s+-m\s+)?twine\s+upload\b"),
    ("docker_push", r"(?:^|[;&|]\s*)(?:docker|podman)\s+push\b"),
    ("git_push", r"(?:^|[;&|]\s*)git\s+push\b"),
    ("github_release", r"(?:^|[;&|]\s*)gh\s+release\s+(?:create|upload)\b"),
    ("package_release", r"(?:^|[;&|]\s*)(?:cargo\s+publish|gem\s+push|dotnet\s+nuget\s+push)\b"),
    ("cloud_deploy", r"(?:^|[;&|]\s*)(?:kubectl\s+apply|helm\s+(?:install|upgrade)|terraform\s+apply)\b"),
]


def print_usage():
    """打印精简使用说明,保持无参数运行也可正常退出。"""
    print("AgentEval Gate - AI 评测环境安全预检器")
    print("扫描: python 37_mvp.py scan CONFIG.json [--output REPORT.json]")
    print("样例: python 37_mvp.py samples [DIRECTORY]")
    print("其他: python 37_mvp.py schema | self-test")


def _finding(rule_id, severity, message, evidence=None, remediation=None):
    item = {"rule_id": rule_id, "severity": severity, "message": message}
    if evidence is not None:
        item["evidence"] = evidence
    if remediation:
        item["remediation"] = remediation
    return item


def _is_bool(value):
    return isinstance(value, bool)


def validate_config(config):
    """按 MVP schema 校验配置,并返回适合 JSON 报告的错误列表。"""
    errors = []
    if not isinstance(config, dict):
        return ["配置根节点必须是 JSON object"]

    version = config.get("schema_version", "1")
    if version != "1":
        errors.append("schema_version 目前只支持字符串 '1'")
    if not isinstance(config.get("prompt"), str):
        errors.append("prompt 必须是 string")

    environment = config.get("environment")
    if not isinstance(environment, dict):
        errors.append("environment 必须是 object")
    else:
        for field in ("network_enabled", "browser_enabled", "shell_enabled", "publish_enabled"):
            if field not in environment:
                errors.append("environment.%s 为必填项" % field)
            elif not _is_bool(environment[field]):
                errors.append("environment.%s 必须是 boolean" % field)

    targets = config.get("targets")
    if not isinstance(targets, dict):
        errors.append("targets 必须是 object")
    else:
        for field in ("allowed_domains", "requested_urls"):
            value = targets.get(field)
            if not isinstance(value, list) or any(not isinstance(item, str) for item in value):
                errors.append("targets.%s 必须是 array[string]" % field)

    commands = config.get("commands")
    if not isinstance(commands, list) or any(not isinstance(item, str) for item in commands):
        errors.append("commands 必须是 array[string]")
    return errors


def normalize_domain(value):
    value = value.strip().lower().rstrip(".")
    if "://" in value:
        value = urlparse(value).hostname or ""
    elif "/" in value:
        value = urlparse("//" + value).hostname or ""
    if value.startswith("*."):
        return "*." + value[2:].rstrip(".")
    try:
        value = value.encode("idna").decode("ascii")
    except UnicodeError:
        return ""
    return value


def is_valid_domain(domain):
    plain = domain[2:] if domain.startswith("*.") else domain
    if not plain or len(plain) > 253 or ".." in plain:
        return False
    try:
        ipaddress.ip_address(plain.strip("[]"))
        return True
    except ValueError:
        pass
    labels = plain.split(".")
    return all(re.match(r"^[a-z0-9](?:[a-z0-9-]{0,61}[a-z0-9])?$", label) for label in labels)


def is_reserved_domain(domain):
    plain = domain[2:] if domain.startswith("*.") else domain
    return plain in RESERVED_EXACT or plain.endswith(RESERVED_SUFFIXES)


def domain_allowed(domain, allowlist):
    """进行严格域名匹配;通配符只匹配子域,不匹配根域。"""
    domain = normalize_domain(domain)
    for allowed in allowlist:
        allowed = normalize_domain(allowed)
        if allowed.startswith("*."):
            suffix = allowed[1:]
            if domain.endswith(suffix) and domain != allowed[2:]:
                return True
        elif domain == allowed:
            return True
    return False


def extract_prompt_domains(prompt):
    found = set()
    # URL 只接受常见 ASCII 字符,避免中文标点后的文字被吞进域名。
    url_pattern = r"https?://[a-z0-9._~:/?#\[\]@!$&'()*+,;=%-]+"
    for match in re.findall(url_pattern, prompt, flags=re.I):
        hostname = urlparse(match.rstrip(".,;:!?,。;!?”’")).hostname
        if hostname:
            found.add(normalize_domain(hostname))
    return found


def _matches_any(patterns, text):
    return any(re.search(pattern, text, flags=re.I | re.S) for pattern in patterns)


def run_rules(config):
    """执行确定性安全规则,所有证据均来自输入配置,便于人工复核。"""
    findings = []
    prompt = config["prompt"]
    env = config["environment"]
    targets = config["targets"]
    commands = config["commands"]
    allowlist = [normalize_domain(item) for item in targets["allowed_domains"]]

    denies_network = _matches_any(DENY_NETWORK_PATTERNS, prompt)
    allows_network = _matches_any(ALLOW_NETWORK_PATTERNS, prompt)
    if env["network_enabled"] and denies_network:
        findings.append(_finding(
            "NET-001", "critical", "提示词声明离线,但环境实际启用了网络",
            {"network_enabled": True}, "关闭网络,或修正提示词并限定目标白名单",
        ))
    elif env["network_enabled"] and not allows_network:
        findings.append(_finding(
            "NET-002", "high", "环境启用了网络,但提示词没有明确声明联网能力",
            {"network_enabled": True}, "在提示词中明确联网范围和限制",
        ))
    elif not env["network_enabled"] and allows_network:
        findings.append(_finding(
            "NET-003", "medium", "提示词允许联网,但环境配置关闭了网络",
            {"network_enabled": False}, "统一提示词与环境网络配置",
        ))

    if env["browser_enabled"] and not env["network_enabled"]:
        findings.append(_finding(
            "ENV-001", "low", "浏览器已启用但网络关闭;请确认仅访问本地内容",
            {"browser_enabled": True, "network_enabled": False},
        ))
    if env["network_enabled"] and not allowlist:
        findings.append(_finding(
            "DOM-001", "critical", "联网环境没有配置目标域名白名单",
            remediation="至少配置一个精确域名,评测样例优先使用 RFC 2606 保留域名",
        ))

    for raw in targets["allowed_domains"]:
        domain = normalize_domain(raw)
        if not domain or not is_valid_domain(domain):
            findings.append(_finding("DOM-002", "high", "白名单包含无效域名", raw, "使用主机名,不要包含路径或端口"))
            continue
        if domain.startswith("*."):
            findings.append(_finding(
                "DOM-003", "medium", "通配符白名单扩大了可访问范围", raw,
                "优先列出评测所需的精确主机名",
            ))
        if env["network_enabled"] and not is_reserved_domain(domain):
            findings.append(_finding(
                "DOM-004", "medium", "白名单指向真实网络域名,可能触达真实组织", raw,
                "若不需要真实服务,改用 RFC 2606 的 .example/.test/.invalid 域名",
            ))

    requested_domains = set(extract_prompt_domains(prompt))
    for raw_url in targets["requested_urls"]:
        parsed = urlparse(raw_url)
        if parsed.scheme not in ("http", "https") or not parsed.hostname:
            findings.append(_finding("DOM-005", "high", "requested_urls 包含无效 HTTP(S) URL", raw_url))
            continue
        requested_domains.add(normalize_domain(parsed.hostname))
    for domain in sorted(requested_domains):
        if not domain_allowed(domain, allowlist):
            findings.append(_finding(
                "DOM-006", "critical", "提示词或目标 URL 中的域名不在白名单", domain,
                "将精确域名加入白名单,或移除该目标",
            ))

    for index, command in enumerate(commands):
        for command_type, pattern in PUBLISH_PATTERNS:
            if re.search(pattern, command, flags=re.I):
                severity = "critical" if not env["publish_enabled"] else "high"
                message = "检测到发布/部署命令,但发布能力未获授权" if not env["publish_enabled"] else "检测到高风险发布/部署命令"
                findings.append(_finding(
                    "PUB-001", severity, message,
                    {"index": index, "type": command_type, "command": command},
                    "删除命令;确需发布时应使用隔离的测试仓库和最小权限凭证",
                ))
                break
    if env["publish_enabled"] and not env["network_enabled"]:
        findings.append(_finding(
            "PUB-002", "medium", "发布能力已启用,但网络已关闭",
            remediation="确认配置是否遗漏,或关闭不必要的发布能力",
        ))
    if commands and not env["shell_enabled"]:
        findings.append(_finding(
            "SHELL-001", "medium", "配置了 Shell 命令,但 shell_enabled 为 false",
            {"command_count": len(commands)}, "统一命令清单与 Shell 权限配置",
        ))
    if env["shell_enabled"] and not commands:
        findings.append(_finding(
            "SHELL-002", "low", "Shell 已启用但命令清单为空,无法预检执行范围",
            remediation="列出预期命令,或关闭 Shell 能力",
        ))
    return sorted(findings, key=lambda item: (-SEVERITY_RANK[item["severity"]], item["rule_id"]))


def build_report(config, source):
    errors = validate_config(config)
    if errors:
        return {
            "tool": "AgentEval Gate",
            "version": VERSION,
            "source": source,
            "status": "invalid",
            "schema_errors": errors,
            "findings": [],
            "summary": {"total": 0, "blocking": 0},
        }
    findings = run_rules(config)
    counts = {name: 0 for name in SEVERITY_RANK}
    for item in findings:
        counts[item["severity"]] += 1
    blocking = sum(counts[name] for name in BLOCKING_SEVERITIES)
    return {
        "tool": "AgentEval Gate",
        "version": VERSION,
        "generated_at": _datetime.datetime.now(_datetime.timezone.utc).isoformat(),
        "source": source,
        "status": "fail" if blocking else ("warn" if findings else "pass"),
        "schema_errors": [],
        "findings": findings,
        "summary": {"total": len(findings), "blocking": blocking, "by_severity": counts},
    }


def load_json(path):
    with open(path, "r", encoding="utf-8") as handle:
        return json.load(handle)


def write_json(path, value):
    with open(path, "w", encoding="utf-8") as handle:
        json.dump(value, handle, ensure_ascii=False, indent=2)
        handle.write("\n")


def scan_command(args):
    try:
        config = load_json(args.config)
    except (OSError, json.JSONDecodeError) as exc:
        report = {
            "tool": "AgentEval Gate", "version": VERSION, "source": args.config,
            "status": "invalid", "schema_errors": [str(exc)], "findings": [],
            "summary": {"total": 0, "blocking": 0},
        }
        print(json.dumps(report, ensure_ascii=False, indent=2))
        return 2

    report = build_report(config, os.path.abspath(args.config))
    rendered = json.dumps(report, ensure_ascii=False, indent=2)
    print(rendered)
    if args.output:
        try:
            write_json(args.output, report)
        except OSError as exc:
            print("无法写入报告: %s" % exc, file=sys.stderr)
            return 2
    if report["status"] == "invalid":
        return 2
    threshold = SEVERITY_RANK[args.fail_on]
    return 1 if any(SEVERITY_RANK[item["severity"]] >= threshold for item in report["findings"]) else 0


def samples_command(directory):
    os.makedirs(directory, exist_ok=True)
    safe_path = os.path.join(directory, "safe.json")
    dangerous_path = os.path.join(directory, "dangerous.json")
    write_json(safe_path, SAFE_SAMPLE)
    write_json(dangerous_path, DANGEROUS_SAMPLE)
    print(json.dumps({"safe": os.path.abspath(safe_path), "dangerous": os.path.abspath(dangerous_path)}, ensure_ascii=False, indent=2))
    return 0


class GateTests(unittest.TestCase):
    def test_safe_sample_passes(self):
        report = build_report(SAFE_SAMPLE, "safe.json")
        self.assertEqual("pass", report["status"])
        self.assertEqual([], report["findings"])

    def test_dangerous_sample_is_blocked(self):
        report = build_report(DANGEROUS_SAMPLE, "dangerous.json")
        rules = {item["rule_id"] for item in report["findings"]}
        self.assertEqual("fail", report["status"])
        self.assertTrue({"NET-001", "PUB-001"}.issubset(rules))

    def test_unlisted_target_is_blocked(self):
        config = json.loads(json.dumps(SAFE_SAMPLE))
        config["targets"]["requested_urls"] = ["https://outside.test/"]
        rules = {item["rule_id"] for item in run_rules(config)}
        self.assertIn("DOM-006", rules)

    def test_chinese_punctuation_ends_prompt_url(self):
        domains = extract_prompt_domains("仅访问 https://shop.example;禁止发布。")
        self.assertEqual({"shop.example"}, domains)

    def test_schema_rejects_missing_fields(self):
        self.assertTrue(validate_config({"prompt": "offline"}))

    def test_cli_report_file(self):
        with tempfile.TemporaryDirectory() as directory:
            config_path = os.path.join(directory, "config.json")
            report_path = os.path.join(directory, "report.json")
            write_json(config_path, SAFE_SAMPLE)
            args = argparse.Namespace(config=config_path, output=report_path, fail_on="high")
            self.assertEqual(0, scan_command(args))
            self.assertEqual("pass", load_json(report_path)["status"])


def self_test_command():
    suite = unittest.defaultTestLoader.loadTestsFromTestCase(GateTests)
    result = unittest.TextTestRunner(verbosity=2).run(suite)
    return 0 if result.wasSuccessful() else 1


def build_parser():
    parser = argparse.ArgumentParser(add_help=True, description="AgentEval Gate")
    subparsers = parser.add_subparsers(dest="command")
    scan = subparsers.add_parser("scan", help="扫描单个 JSON 配置")
    scan.add_argument("config", help="配置文件路径")
    scan.add_argument("--output", "-o", help="同时写入 JSON 报告文件")
    scan.add_argument("--fail-on", choices=list(SEVERITY_RANK), default="high", help="触发非零退出码的最低等级")
    samples = subparsers.add_parser("samples", help="生成安全与危险样例")
    samples.add_argument("directory", nargs="?", default="agent-eval-samples")
    subparsers.add_parser("schema", help="打印 MVP 配置 schema")
    subparsers.add_parser("self-test", help="运行内置 CLI/规则测试")
    return parser


def main(argv=None):
    argv = sys.argv[1:] if argv is None else argv
    if not argv:
        print_usage()
        return 0
    parser = build_parser()
    args = parser.parse_args(argv)
    if args.command == "scan":
        return scan_command(args)
    if args.command == "samples":
        return samples_command(args.directory)
    if args.command == "schema":
        print(json.dumps(SCHEMA, ensure_ascii=False, indent=2))
        return 0
    if args.command == "self-test":
        return self_test_command()
    print_usage()
    return 0


if __name__ == "__main__":
    sys.exit(main())