diff --git a/src/modelfuzz/rules.py b/src/modelfuzz/rules.py index e2b7aef..c4706b0 100644 --- a/src/modelfuzz/rules.py +++ b/src/modelfuzz/rules.py @@ -154,11 +154,17 @@ def _check_recursive(self, data: object) -> Violation | None: reason=f"String contains sensitive keyword: '{keyword}'", ) elif isinstance(data, dict): - for value in data.values(): + for key, value in data.items(): + violation = self._check_recursive(key) + if violation: + return violation + violation = self._check_recursive(value) if violation: return violation - elif isinstance(data, (list, tuple)): + elif isinstance(data, (bytes, bytearray)): + return self._check_recursive(data.decode("utf-8", errors="ignore")) + elif isinstance(data, (list, tuple, set, frozenset)): for item in data: violation = self._check_recursive(item) if violation: diff --git a/tests/test_rules.py b/tests/test_rules.py index 389dcd4..d267e2f 100644 --- a/tests/test_rules.py +++ b/tests/test_rules.py @@ -246,3 +246,33 @@ def test_allows_clean_data(self, filter: SensitiveDataFilter): data = {"user": "alice", "action": "login"} violation = filter(data) assert violation is None + + def test_blocks_sensitive_dict_key(self, filter: SensitiveDataFilter): + """Ensure it blocks sensitive keywords in dictionary keys.""" + violation = filter({"api_key": "abc123"}) + assert violation is not None + assert "api_key" in violation.reason + + def test_blocks_sensitive_bytes(self, filter: SensitiveDataFilter): + """Ensure bytes are inspected.""" + violation = filter(b"contains password") + assert violation is not None + assert "password" in violation.reason + + def test_blocks_sensitive_bytearray(self, filter: SensitiveDataFilter): + """Ensure bytearrays are inspected.""" + violation = filter(bytearray(b"contains api_key")) + assert violation is not None + assert "api_key" in violation.reason + + def test_blocks_sensitive_set(self, filter: SensitiveDataFilter): + """Ensure sets are inspected.""" + violation = filter({"contains secret"}) + assert violation is not None + assert "secret" in violation.reason + + def test_blocks_sensitive_frozenset(self, filter: SensitiveDataFilter): + """Ensure frozensets are inspected.""" + violation = filter(frozenset({"contains password"})) + assert violation is not None + assert "password" in violation.reason