Files
info-privacy/tests/test_classifier.py
T
qiurui cbfe4a23dc feat: info-privacy PII detection service with frame analysis support
- Pipeline: analyze(), redact(), analyze_image() methods
- API: /analyze, /redact, /analyze/frame, /analyze/frame/base64 endpoints
- Detectors: regex, NER, face (RKNN NPU)
- Privacy frame route added for KVM-Privacy Hub integration
2026-02-28 17:33:11 +08:00

86 lines
3.6 KiB
Python

# tests/test_classifier.py
import pytest
from info_privacy.models import TextBlock, Classification, Entity, EntityType
from info_privacy.classifier.doc_classifier import DocClassifier
def test_classified_by_keyword():
classifier = DocClassifier("configs/pii_rules.yaml")
blocks = [TextBlock(text="【机密】本文件仅供内部使用", bbox=[0,0,200,20], page=1, layer="text")]
result = classifier.classify(blocks)
assert result.classification == Classification.CLASSIFIED
assert result.blocked is True
assert result.block_reason is not None
def test_sensitive_partial_many_entities():
from info_privacy.models import Entity, EntityType
classifier = DocClassifier("configs/pii_rules.yaml")
entities = [
Entity(id=str(i), type=EntityType.ID_CARD, value="x", page=1,
bbox=[0,0,1,1], layer="text", security_level="high")
for i in range(6)
]
result = classifier.classify([], entities=entities)
assert result.classification == Classification.SENSITIVE_PARTIAL
assert result.warning is not None
def test_normal_document():
classifier = DocClassifier("configs/pii_rules.yaml")
blocks = [TextBlock(text="本季度销售报告", bbox=[0,0,200,20], page=1, layer="text")]
result = classifier.classify(blocks)
assert result.classification == Classification.NORMAL
assert result.blocked is False
_ALL_KEYWORDS = [
"机密", "绝密", "保密", "内部文件", "内部资料",
"CONFIDENTIAL", "SECRET", "TOP SECRET", "RESTRICTED", "FOR INTERNAL USE ONLY",
]
@pytest.mark.parametrize("keyword", _ALL_KEYWORDS)
def test_all_confidential_keywords_trigger_blocked(keyword):
"""pii_rules.yaml 中每个保密关键词都应触发 classified 拦截。"""
classifier = DocClassifier("configs/pii_rules.yaml")
blocks = [TextBlock(text=f"标题:{keyword}", bbox=[0,0,300,20], page=1, layer="text")]
result = classifier.classify(blocks)
assert result.classification == Classification.CLASSIFIED, f"关键词 '{keyword}' 未触发拦截"
assert result.blocked is True
assert result.block_reason is not None
def _make_high_risk_entities(n: int) -> list[Entity]:
return [
Entity(id=str(i), type=EntityType.PHONE, value="13800000000",
page=1, bbox=[0, 0, 1, 1], layer="text", security_level="high")
for i in range(n)
]
def test_exactly_4_high_risk_no_warning():
"""4 个高风险实体:低于阈值,不触发 sensitive_partial 警告。"""
classifier = DocClassifier("configs/pii_rules.yaml")
result = classifier.classify([], entities=_make_high_risk_entities(4))
assert result.classification == Classification.NORMAL
assert result.warning is None
def test_exactly_5_high_risk_triggers_warning():
"""5 个高风险实体:恰好达到阈值,应触发 sensitive_partial 警告。"""
classifier = DocClassifier("configs/pii_rules.yaml")
result = classifier.classify([], entities=_make_high_risk_entities(5))
assert result.classification == Classification.SENSITIVE_PARTIAL
assert result.warning is not None
def test_medium_entities_alone_do_not_trigger_warning():
"""10 个中风险实体(email)不应触发 sensitive_partial 警告。"""
classifier = DocClassifier("configs/pii_rules.yaml")
entities = [
Entity(id=str(i), type=EntityType.EMAIL, value="x@x.com",
page=1, bbox=[0, 0, 1, 1], layer="text", security_level="medium")
for i in range(10)
]
result = classifier.classify([], entities=entities)
assert result.classification == Classification.NORMAL
assert result.warning is None