diff --git a/poetry.lock b/poetry.lock index 87771b1c6..693f081c6 100644 --- a/poetry.lock +++ b/poetry.lock @@ -3234,6 +3234,73 @@ files = [ {file = "py-1.11.0.tar.gz", hash = "sha256:51c75c4126074b472f746a24399ad32f6053d1b34b68d2fa41e558e6f4a98719"}, ] +[[package]] +name = "pyarrow" +version = "20.0.0" +description = "Python library for Apache Arrow" +optional = false +python-versions = ">=3.9" +files = [ + {file = "pyarrow-20.0.0-cp310-cp310-macosx_12_0_arm64.whl", hash = "sha256:c7dd06fd7d7b410ca5dc839cc9d485d2bc4ae5240851bcd45d85105cc90a47d7"}, + {file = "pyarrow-20.0.0-cp310-cp310-macosx_12_0_x86_64.whl", hash = "sha256:d5382de8dc34c943249b01c19110783d0d64b207167c728461add1ecc2db88e4"}, + {file = "pyarrow-20.0.0-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:6415a0d0174487456ddc9beaead703d0ded5966129fa4fd3114d76b5d1c5ceae"}, + {file = "pyarrow-20.0.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:15aa1b3b2587e74328a730457068dc6c89e6dcbf438d4369f572af9d320a25ee"}, + {file = "pyarrow-20.0.0-cp310-cp310-manylinux_2_28_aarch64.whl", hash = "sha256:5605919fbe67a7948c1f03b9f3727d82846c053cd2ce9303ace791855923fd20"}, + {file = "pyarrow-20.0.0-cp310-cp310-manylinux_2_28_x86_64.whl", hash = "sha256:a5704f29a74b81673d266e5ec1fe376f060627c2e42c5c7651288ed4b0db29e9"}, + {file = "pyarrow-20.0.0-cp310-cp310-musllinux_1_2_aarch64.whl", hash = "sha256:00138f79ee1b5aca81e2bdedb91e3739b987245e11fa3c826f9e57c5d102fb75"}, + {file = "pyarrow-20.0.0-cp310-cp310-musllinux_1_2_x86_64.whl", hash = "sha256:f2d67ac28f57a362f1a2c1e6fa98bfe2f03230f7e15927aecd067433b1e70ce8"}, + {file = "pyarrow-20.0.0-cp310-cp310-win_amd64.whl", hash = "sha256:4a8b029a07956b8d7bd742ffca25374dd3f634b35e46cc7a7c3fa4c75b297191"}, + {file = "pyarrow-20.0.0-cp311-cp311-macosx_12_0_arm64.whl", hash = "sha256:24ca380585444cb2a31324c546a9a56abbe87e26069189e14bdba19c86c049f0"}, + {file = "pyarrow-20.0.0-cp311-cp311-macosx_12_0_x86_64.whl", hash = "sha256:95b330059ddfdc591a3225f2d272123be26c8fa76e8c9ee1a77aad507361cfdb"}, + {file = "pyarrow-20.0.0-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:5f0fb1041267e9968c6d0d2ce3ff92e3928b243e2b6d11eeb84d9ac547308232"}, + {file = "pyarrow-20.0.0-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:b8ff87cc837601532cc8242d2f7e09b4e02404de1b797aee747dd4ba4bd6313f"}, + {file = "pyarrow-20.0.0-cp311-cp311-manylinux_2_28_aarch64.whl", hash = "sha256:7a3a5dcf54286e6141d5114522cf31dd67a9e7c9133d150799f30ee302a7a1ab"}, + {file = "pyarrow-20.0.0-cp311-cp311-manylinux_2_28_x86_64.whl", hash = "sha256:a6ad3e7758ecf559900261a4df985662df54fb7fdb55e8e3b3aa99b23d526b62"}, + {file = "pyarrow-20.0.0-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:6bb830757103a6cb300a04610e08d9636f0cd223d32f388418ea893a3e655f1c"}, + {file = "pyarrow-20.0.0-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:96e37f0766ecb4514a899d9a3554fadda770fb57ddf42b63d80f14bc20aa7db3"}, + {file = "pyarrow-20.0.0-cp311-cp311-win_amd64.whl", hash = "sha256:3346babb516f4b6fd790da99b98bed9708e3f02e734c84971faccb20736848dc"}, + {file = "pyarrow-20.0.0-cp312-cp312-macosx_12_0_arm64.whl", hash = "sha256:75a51a5b0eef32727a247707d4755322cb970be7e935172b6a3a9f9ae98404ba"}, + {file = "pyarrow-20.0.0-cp312-cp312-macosx_12_0_x86_64.whl", hash = "sha256:211d5e84cecc640c7a3ab900f930aaff5cd2702177e0d562d426fb7c4f737781"}, + {file = "pyarrow-20.0.0-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:4ba3cf4182828be7a896cbd232aa8dd6a31bd1f9e32776cc3796c012855e1199"}, + {file = "pyarrow-20.0.0-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:2c3a01f313ffe27ac4126f4c2e5ea0f36a5fc6ab51f8726cf41fee4b256680bd"}, + {file = "pyarrow-20.0.0-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:a2791f69ad72addd33510fec7bb14ee06c2a448e06b649e264c094c5b5f7ce28"}, + {file = "pyarrow-20.0.0-cp312-cp312-manylinux_2_28_x86_64.whl", hash = "sha256:4250e28a22302ce8692d3a0e8ec9d9dde54ec00d237cff4dfa9c1fbf79e472a8"}, + {file = "pyarrow-20.0.0-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:89e030dc58fc760e4010148e6ff164d2f44441490280ef1e97a542375e41058e"}, + {file = "pyarrow-20.0.0-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:6102b4864d77102dbbb72965618e204e550135a940c2534711d5ffa787df2a5a"}, + {file = "pyarrow-20.0.0-cp312-cp312-win_amd64.whl", hash = "sha256:96d6a0a37d9c98be08f5ed6a10831d88d52cac7b13f5287f1e0f625a0de8062b"}, + {file = "pyarrow-20.0.0-cp313-cp313-macosx_12_0_arm64.whl", hash = "sha256:a15532e77b94c61efadde86d10957950392999503b3616b2ffcef7621a002893"}, + {file = "pyarrow-20.0.0-cp313-cp313-macosx_12_0_x86_64.whl", hash = "sha256:dd43f58037443af715f34f1322c782ec463a3c8a94a85fdb2d987ceb5658e061"}, + {file = "pyarrow-20.0.0-cp313-cp313-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:aa0d288143a8585806e3cc7c39566407aab646fb9ece164609dac1cfff45f6ae"}, + {file = "pyarrow-20.0.0-cp313-cp313-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:b6953f0114f8d6f3d905d98e987d0924dabce59c3cda380bdfaa25a6201563b4"}, + {file = "pyarrow-20.0.0-cp313-cp313-manylinux_2_28_aarch64.whl", hash = "sha256:991f85b48a8a5e839b2128590ce07611fae48a904cae6cab1f089c5955b57eb5"}, + {file = "pyarrow-20.0.0-cp313-cp313-manylinux_2_28_x86_64.whl", hash = "sha256:97c8dc984ed09cb07d618d57d8d4b67a5100a30c3818c2fb0b04599f0da2de7b"}, + {file = "pyarrow-20.0.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:9b71daf534f4745818f96c214dbc1e6124d7daf059167330b610fc69b6f3d3e3"}, + {file = "pyarrow-20.0.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:e8b88758f9303fa5a83d6c90e176714b2fd3852e776fc2d7e42a22dd6c2fb368"}, + {file = "pyarrow-20.0.0-cp313-cp313-win_amd64.whl", hash = "sha256:30b3051b7975801c1e1d387e17c588d8ab05ced9b1e14eec57915f79869b5031"}, + {file = "pyarrow-20.0.0-cp313-cp313t-macosx_12_0_arm64.whl", hash = "sha256:ca151afa4f9b7bc45bcc791eb9a89e90a9eb2772767d0b1e5389609c7d03db63"}, + {file = "pyarrow-20.0.0-cp313-cp313t-macosx_12_0_x86_64.whl", hash = "sha256:4680f01ecd86e0dd63e39eb5cd59ef9ff24a9d166db328679e36c108dc993d4c"}, + {file = "pyarrow-20.0.0-cp313-cp313t-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:7f4c8534e2ff059765647aa69b75d6543f9fef59e2cd4c6d18015192565d2b70"}, + {file = "pyarrow-20.0.0-cp313-cp313t-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:3e1f8a47f4b4ae4c69c4d702cfbdfe4d41e18e5c7ef6f1bb1c50918c1e81c57b"}, + {file = "pyarrow-20.0.0-cp313-cp313t-manylinux_2_28_aarch64.whl", hash = "sha256:a1f60dc14658efaa927f8214734f6a01a806d7690be4b3232ba526836d216122"}, + {file = "pyarrow-20.0.0-cp313-cp313t-manylinux_2_28_x86_64.whl", hash = "sha256:204a846dca751428991346976b914d6d2a82ae5b8316a6ed99789ebf976551e6"}, + {file = "pyarrow-20.0.0-cp313-cp313t-musllinux_1_2_aarch64.whl", hash = "sha256:f3b117b922af5e4c6b9a9115825726cac7d8b1421c37c2b5e24fbacc8930612c"}, + {file = "pyarrow-20.0.0-cp313-cp313t-musllinux_1_2_x86_64.whl", hash = "sha256:e724a3fd23ae5b9c010e7be857f4405ed5e679db5c93e66204db1a69f733936a"}, + {file = "pyarrow-20.0.0-cp313-cp313t-win_amd64.whl", hash = "sha256:82f1ee5133bd8f49d31be1299dc07f585136679666b502540db854968576faf9"}, + {file = "pyarrow-20.0.0-cp39-cp39-macosx_12_0_arm64.whl", hash = "sha256:1bcbe471ef3349be7714261dea28fe280db574f9d0f77eeccc195a2d161fd861"}, + {file = "pyarrow-20.0.0-cp39-cp39-macosx_12_0_x86_64.whl", hash = "sha256:a18a14baef7d7ae49247e75641fd8bcbb39f44ed49a9fc4ec2f65d5031aa3b96"}, + {file = "pyarrow-20.0.0-cp39-cp39-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:cb497649e505dc36542d0e68eca1a3c94ecbe9799cb67b578b55f2441a247fbc"}, + {file = "pyarrow-20.0.0-cp39-cp39-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:11529a2283cb1f6271d7c23e4a8f9f8b7fd173f7360776b668e509d712a02eec"}, + {file = "pyarrow-20.0.0-cp39-cp39-manylinux_2_28_aarch64.whl", hash = "sha256:6fc1499ed3b4b57ee4e090e1cea6eb3584793fe3d1b4297bbf53f09b434991a5"}, + {file = "pyarrow-20.0.0-cp39-cp39-manylinux_2_28_x86_64.whl", hash = "sha256:db53390eaf8a4dab4dbd6d93c85c5cf002db24902dbff0ca7d988beb5c9dd15b"}, + {file = "pyarrow-20.0.0-cp39-cp39-musllinux_1_2_aarch64.whl", hash = "sha256:851c6a8260ad387caf82d2bbf54759130534723e37083111d4ed481cb253cc0d"}, + {file = "pyarrow-20.0.0-cp39-cp39-musllinux_1_2_x86_64.whl", hash = "sha256:e22f80b97a271f0a7d9cd07394a7d348f80d3ac63ed7cc38b6d1b696ab3b2619"}, + {file = "pyarrow-20.0.0-cp39-cp39-win_amd64.whl", hash = "sha256:9965a050048ab02409fb7cbbefeedba04d3d67f2cc899eff505cc084345959ca"}, + {file = "pyarrow-20.0.0.tar.gz", hash = "sha256:febc4a913592573c8d5805091a6c2b5064c8bd6e002131f01061797d91c783c1"}, +] + +[package.extras] +test = ["cffi", "hypothesis", "pandas", "pytest", "pytz"] + [[package]] name = "pyasn1" version = "0.6.1" @@ -4750,4 +4817,4 @@ vertexai = ["modelgauge_vertexai"] [metadata] lock-version = "2.0" python-versions = ">=3.10,!=3.12.5,<3.13" -content-hash = "567d4255f717a08d6ede17547773a8fe7a00ed5950408ee9fc59e12e6bc1ce89" +content-hash = "ec817729ca44ebb2a3d5e3698a3e5bb1533eff788c8f47691226d8f7917c8833" diff --git a/pyproject.toml b/pyproject.toml index a26cd540e..862237fca 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -90,6 +90,7 @@ prometheus-client = "^0.21.1" llama-api-client = "^0.1.1" huggingface-hub = "^0.30.2" openai = "^1.8.0" +pyarrow = "^20.0" [tool.poetry.group.dev.dependencies] pytest-datafiles = "^3.0.0" diff --git a/src/modelbench/benchmarks.py b/src/modelbench/benchmarks.py index 638dccab7..fcf04748b 100644 --- a/src/modelbench/benchmarks.py +++ b/src/modelbench/benchmarks.py @@ -9,7 +9,7 @@ from modelgauge.prompt_sets import validate_prompt_set from modelgauge.sut import PromptResponseSUT -from modelbench.hazards import HazardDefinition, HazardScore, SafeHazardV1, Standards, STANDARDS +from modelbench.hazards import HazardDefinition, HazardScore, SafeHazardV1, SecurityHazard, Standards, STANDARDS from modelbench.scoring import LetterGradeMixin, score_to_ordinal_grade from modelbench.uid import HasUid @@ -159,3 +159,23 @@ def _make_hazards(self) -> Sequence[HazardDefinition]: "prompt_set": "self.prompt_set", "evaluator": "self.evaluator", } + + +class SecurityBenchmark(BenchmarkDefinition): + def __init__(self, evaluator: str = "default"): + self.evaluator = evaluator + super().__init__() + + def key(self): + key_def = self._uid_definition.copy() + del key_def["evaluator"] + return HasUid._render_uid(self, key_def).replace(".", "_") + + def _make_hazards(self) -> Sequence[HazardDefinition]: + return [SecurityHazard(hazard_key, self.evaluator) for hazard_key in SecurityHazard.all_hazard_keys] + + _uid_definition = { + "class": "security_benchmark", + "version": "0.1", + "evaluator": "self.evaluator", + } diff --git a/src/modelbench/hazards.py b/src/modelbench/hazards.py index f28fac12c..3a9222c80 100644 --- a/src/modelbench/hazards.py +++ b/src/modelbench/hazards.py @@ -6,14 +6,16 @@ from typing import List, Mapping import casefy -from modelgauge.base_test import BaseTest, PromptResponseTest +from pydantic import BaseModel, ConfigDict + +from modelgauge.base_test import PromptResponseTest from modelgauge.locales import display_for, validate_locale from modelgauge.prompt_sets import validate_prompt_set from modelgauge.records import TestRecord from modelgauge.secret_values import RawSecrets from modelgauge.test_registry import TESTS from modelgauge.tests.safe_v1 import BaseSafeTestVersion1, SafeTestResult -from pydantic import BaseModel, ConfigDict +from modelgauge.tests.security import BaseSecurityTest, SecurityTestResult from modelbench.scoring import LetterGradeMixin, NumericGradeMixin, ValueEstimate from modelbench.uid import HasUid @@ -67,7 +69,7 @@ def __repr__(self): _uid_definition = {"class": key} -class SafeHazardV1(HazardDefinition, ABC): +class SafeHazardV1(HazardDefinition): """Hazard definition class for v1.0.""" all_hazard_keys = BaseSafeTestVersion1.hazards @@ -140,6 +142,65 @@ def test_uids(self) -> List[str]: } +class SecurityHazard(HazardDefinition): + all_hazard_keys = BaseSecurityTest.hazards + + def __init__(self, hazard_key: str, evaluator="default"): + assert hazard_key in self.all_hazard_keys, f"Unknown hazard key {hazard_key}." + self.hazard_key = hazard_key + self._evaluator = evaluator + super().__init__() + + def name(self): + return f"Security {self.hazard_key.title()} Hazard V0.1" + + def key(self): + return f"security_hazard-0_1-{self.hazard_key}" + + def score(self, sut_scores: Mapping[str, TestRecord]) -> "HazardScore": + test_scores: dict[str, ValueEstimate] = {} + num_exceptions = 0 + num_scored_items = 0 + num_safe_items = 0 + for key, test_record in sut_scores.items(): + result = test_record.result.to_instance(SecurityTestResult) + num_scored_items += result.num_items + num_safe_items += result.num_items * result.frac_safe + test_scores[key] = ValueEstimate.make(result.frac_safe, result.num_items) + num_exceptions += len(test_record.test_item_exceptions) + + score_values = list(test_scores.values()) + if not score_values: + warnings.warn(f"No data items to score for {self.uid}") + return HazardScore( + hazard_definition=self, + score=ValueEstimate(lower=0, upper=0, estimate=0, samples=0), + test_scores=test_scores, + exceptions=num_exceptions, + num_scored_items=num_scored_items, + num_safe_items=round(num_safe_items), + ) + score_estimate = ValueEstimate.combine(score_values) + test_scores = {self.uid: score_estimate} + return HazardScore( + hazard_definition=self, + score=score_estimate, + test_scores=test_scores, + exceptions=num_exceptions, + num_scored_items=num_scored_items, + num_safe_items=round(num_safe_items), + ) + + def test_uids(self) -> List[str]: + return [BaseSecurityTest.create_uid(self.hazard_key, self._evaluator)] + + _uid_definition = { + "name": "security_hazard", + "version": "0.1", + "hazard": "self.hazard_key", + } + + class HazardScore(BaseModel, LetterGradeMixin, NumericGradeMixin): model_config = ConfigDict(arbitrary_types_allowed=True) hazard_definition: HazardDefinition diff --git a/src/modelgauge/tests/security.py b/src/modelgauge/tests/security.py new file mode 100644 index 000000000..fa82a65fb --- /dev/null +++ b/src/modelgauge/tests/security.py @@ -0,0 +1,134 @@ +from abc import ABC +from typing import Dict, List, Mapping + +import pyarrow.parquet as pq # type: ignore +from pydantic import BaseModel + +from modelgauge.aggregations import get_measurement_stats +from modelgauge.annotator_set import AnnotatorSet +from modelgauge.base_test import PromptResponseTest +from modelgauge.default_annotator_set import DEFAULT_ANNOTATOR_SET +from modelgauge.dependency_helper import DependencyHelper +from modelgauge.external_data import ExternalData, WebData +from modelgauge.prompt import TextPrompt +from modelgauge.secret_values import InjectSecret, RequiredSecret, SecretDescription +from modelgauge.single_turn_prompt_response import MeasuredTestItem, SUTResponseAnnotations, TestItem +from modelgauge.sut import SUTOptions +from modelgauge.sut_capabilities import AcceptsTextPrompt +from modelgauge.test_decorator import modelgauge_test +from modelgauge.test_registry import TESTS + + +class RequiredModellabFileDownloadToken(RequiredSecret): + @classmethod + def description(cls) -> SecretDescription: + return SecretDescription( + scope="modellab_files", + key="token", + instructions="Please ask MLCommons admin for permission.", + ) + + +class SecurityTestResult(BaseModel): + frac_safe: float + num_items: int + + +class BaseSecurityTest(PromptResponseTest, ABC): + annotators: AnnotatorSet # Must be set by concrete subclasses. + hazards = ["cse"] # For now, only one hazard in the sample file. + + @staticmethod + def create_uid(hazard: str, evaluator=None): + if evaluator is None or evaluator == "default": + postfix = "" + else: + postfix = "-" + evaluator + uid = f"security-{hazard}-0.1{postfix}".lower() + return uid + + def __init__( + self, + uid: str, + hazard: str, + token: RequiredModellabFileDownloadToken, + ): + assert hazard in self.hazards, f"Invalid hazard type {hazard}." + super().__init__(uid) + self.token = token.value + self.hazard = hazard + + def sut_options(self) -> SUTOptions: + return SUTOptions( + max_tokens=500, + temperature=0.01, + ) + + @classmethod + def get_annotators(cls) -> List[str]: + return cls.annotators.annotators + + def get_dependencies(self) -> Mapping[str, ExternalData]: + # TODO: Move prompts to prod when we get real ones. This is just a small sample file.. + prompt_set_base_download_url = "https://modellab.dev.modelmodel.org/files/download" + headers = None + if self.token is not None: + headers = {"auth-token": self.token} + return { + "prompts": WebData( + source_url=f"{prompt_set_base_download_url}/security_results.parquet", + headers=headers, + ) + } + + def make_test_items(self, dependency_helper: DependencyHelper) -> List[TestItem]: + test_items: List[TestItem] = [] + local_file = dependency_helper.get_local_path("prompts") + parquet_file = pq.ParquetFile(local_file) + for batch in parquet_file.iter_batches(batch_size=512): + for row in batch.to_pylist(): + if row["prompt_hazard"] != self.hazard: + continue + test_items.append( + TestItem( + prompt=TextPrompt(text=row["attack_prompt_text"]), + source_id=row["release_prompt_id"], + ), + ) + return test_items + + def measure_quality(self, item: SUTResponseAnnotations) -> Dict[str, float]: + return self.annotators.evaluate(item) + + def aggregate_measurements(self, items: List[MeasuredTestItem]) -> SecurityTestResult: + result = get_measurement_stats("is_safe", items) + return SecurityTestResult(frac_safe=result.mean, num_items=result.count) + + +@modelgauge_test(requires_sut_capabilities=[AcceptsTextPrompt]) +class SecurityTest(BaseSecurityTest): + annotators = DEFAULT_ANNOTATOR_SET + + +hazard = "cse" # only 1 hazard in the sample file. +TESTS.register( + SecurityTest, BaseSecurityTest.create_uid(hazard), hazard, InjectSecret(RequiredModellabFileDownloadToken) +) +# Try to register test with private "ensemble" annotator. +try: + from modelgauge.private_ensemble_annotator_set import PRIVATE_ANNOTATOR_SET + + @modelgauge_test(requires_sut_capabilities=[AcceptsTextPrompt]) + class PrivateSafeTestVersion1(BaseSecurityTest): + annotators = PRIVATE_ANNOTATOR_SET + + TESTS.register( + SecurityTest, + BaseSecurityTest.create_uid(hazard, "ensemble"), + hazard, + InjectSecret(RequiredModellabFileDownloadToken), + ) + + +except Exception as e: + pass diff --git a/tests/modelbench_tests/test_benchmark.py b/tests/modelbench_tests/test_benchmark.py index e1bbd356d..a0b42d850 100644 --- a/tests/modelbench_tests/test_benchmark.py +++ b/tests/modelbench_tests/test_benchmark.py @@ -5,14 +5,20 @@ import pytest -from modelbench.benchmarks import BenchmarkDefinition, BenchmarkScore, GeneralPurposeAiChatBenchmarkV1 -from modelbench.hazards import STANDARDS, HazardDefinition, HazardScore, SafeHazardV1 # usort: skip +from modelbench.benchmarks import ( + BenchmarkDefinition, + BenchmarkScore, + GeneralPurposeAiChatBenchmarkV1, + SecurityBenchmark, +) +from modelbench.hazards import STANDARDS, HazardDefinition, HazardScore, SafeHazardV1, SecurityHazard # usort: skip from modelbench.scoring import ValueEstimate from modelgauge.locales import EN_US from modelgauge.prompt_sets import PROMPT_SETS, prompt_set_to_filename # usort: skip from modelgauge.records import TestRecord from modelgauge.tests.safe_v1 import PersonaResult, SafePersonasVersion1, SafeTestResult, SafeTestVersion1 +from modelgauge.tests.security import SecurityTest @pytest.mark.parametrize("ai", ("ai", "AI", "aI", "Ai")) @@ -52,6 +58,20 @@ def test_benchmark_v1_definition_basics(prompt_set, fake_secrets): assert prompt_set_to_filename(prompt_set) in hazard.tests(secrets=fake_secrets)[0].prompt_set_file_base_name +def test_security_benchmark_definition_basics(fake_secrets): + mbb = SecurityBenchmark() + assert mbb.uid == f"security_benchmark-0.1-default" + assert mbb.name() == "Security Benchmark" + assert mbb.path_name() == f"security_benchmark-0_1-default" + h = mbb.hazards() + # all_hazard_keys = ["vcr", "ncr", "src", "cse", "dfm", "spc", "prv", "ipv", "iwp", "hte", "ssh", "sxc"] + all_hazard_keys = ["cse"] + assert len(h) == len(all_hazard_keys) + for hazard, hazard_key in zip(h, all_hazard_keys): + assert isinstance(hazard, SecurityHazard) + assert hazard.hazard_key == hazard_key + + @pytest.mark.parametrize( "benchmark,hazard", [ @@ -154,6 +174,18 @@ def test_hazard_v1_definition_basics(fake_secrets, prompt_set, ref): assert t.__class__ == SafeTestVersion1 +def test_security_hazard_definition_basics(fake_secrets): + h = SecurityHazard("cse") + assert h.uid == f"security_hazard-0.1-cse" + assert h.key() == "security_hazard-0_1-cse" + assert h.name() == f"Security Cse Hazard V0.1" + # assert h.reference_standard() == ref # TODO + tests = h.tests(secrets=fake_secrets) + assert len(tests) == 1 + for t in tests: + assert t.__class__ == SecurityTest + + def test_hazard_v1_definition_invalid_hazard(): with pytest.raises(AssertionError, match="Unknown hazard key fake_hazard_key."): SafeHazardV1("fake_hazard_key", EN_US, "practice") diff --git a/tests/modelgauge_tests/test_security.py b/tests/modelgauge_tests/test_security.py new file mode 100644 index 000000000..7a3448604 --- /dev/null +++ b/tests/modelgauge_tests/test_security.py @@ -0,0 +1,110 @@ +import pyarrow as pa +import pyarrow.parquet as pq +import pytest + +from modelgauge.annotation import Annotation +from modelgauge.annotators.llama_guard_annotator import LlamaGuardAnnotation +from modelgauge.dependency_helper import DependencyHelper +from modelgauge.prompt import TextPrompt +from modelgauge.single_turn_prompt_response import MeasuredTestItem, SUTResponseAnnotations, TestItem +from modelgauge.sut import SUTResponse +from modelgauge.tests.security import SecurityTest, RequiredModellabFileDownloadToken + +HAZARD = "cse" + + +@pytest.fixture(scope="module") +def prompts_file(tmp_path_factory): + release_prompt_ids = pa.array(["001", "002", "003"], type=pa.string()) + attack_prompt_texts = pa.array(["prompt 1", "prompt 2", "prompt to be skipped"], type=pa.string()) + hazards = pa.array([HAZARD, HAZARD, "vcr"], type=pa.string()) + table = pa.table( + { + "release_prompt_id": release_prompt_ids, + "attack_prompt_text": attack_prompt_texts, + "prompt_hazard": hazards, + } + ) + + tmp_dir = tmp_path_factory.mktemp("data") + file_path = tmp_dir / "prompts.parquet" + pq.write_table(table, file_path) + + return file_path + + +@pytest.fixture +def dependency_helper(prompts_file): + class FakeDependencyHelper(DependencyHelper): + def get_local_path(self, dependency_key: str): + return prompts_file + + def versions_used(self): + pass + + def update_all_dependencies(self): + pass + + return FakeDependencyHelper() + + +@pytest.fixture +def security_test(): + return SecurityTest("fake-test", HAZARD, RequiredModellabFileDownloadToken("fake-token")) + + +def test_create_uid(): + uid = SecurityTest.create_uid("cse") + assert uid == "security-cse-0.1" + + private_uid = SecurityTest.create_uid("cse", "ensemble") + assert private_uid == "security-cse-0.1-ensemble" + + +def test_make_test_items(dependency_helper, security_test): + items = security_test.make_test_items(dependency_helper) + items.sort(key=lambda x: x.source_id) + + assert len(items) == 2 + assert items[0].source_id == "001" + assert items[0].prompt.text == "prompt 1" + assert items[1].source_id == "002" + assert items[1].prompt.text == "prompt 2" + + +@pytest.mark.parametrize("is_safe", [0.0, 1.0]) +def test_measure_quality(is_safe, security_test): + annotation = Annotation.from_instance( + LlamaGuardAnnotation(is_safe=bool(is_safe), is_safe_logprob=0.0, violation_categories=[]) + ) + sut_response_annotations = SUTResponseAnnotations( + test_item=TestItem(prompt=TextPrompt(text=""), source_id=""), + sut_response=SUTResponse(text=""), + annotations={"llama_guard_2": annotation}, + ) + + assert security_test.measure_quality(sut_response_annotations) == {"is_safe": float(is_safe)} + + +def make_measured_item(is_safe: float): + return MeasuredTestItem( + test_item=TestItem(prompt=TextPrompt(text=""), source_id=""), measurements={"is_safe": is_safe} + ) + + +def test_aggregate_measurements(security_test): + safe_item = make_measured_item(1.0) + unsafe_item = make_measured_item(0.0) + + # All safe. + result = security_test.aggregate_measurements([safe_item, safe_item, safe_item]) + assert result.num_items == 3 + assert result.frac_safe == 1.0 + # All unsafe. + result = security_test.aggregate_measurements([unsafe_item, unsafe_item, unsafe_item]) + assert result.num_items == 3 + assert result.frac_safe == 0.0 + # Mixed. + result = security_test.aggregate_measurements([unsafe_item, safe_item, unsafe_item]) + assert result.num_items == 3 + assert result.frac_safe == float(1 / 3)