Files
han-app/codebase/services/message-safety/tests/test_api_contract.py
T

191 lines
6.7 KiB
Python

from __future__ import annotations
from uuid import uuid4
import httpx
import pytest
from app.api import create_app
from app.db import TaskStatus
from app.repository import ConflictError
from app.service import SafetyService
from app.settings import EmergencyMode
class FakeRepository:
def __init__(self) -> None:
self.requests = {}
self.text = {}
self.tasks = {}
self.audits = []
async def get_request(self, message_id):
return self.requests.get(message_id)
async def reserve_request(self, row, task=None, **kwargs):
existing = self.requests.get(row.message_id)
if existing:
if existing.request_fingerprint != row.request_fingerprint:
raise ConflictError
return existing, False
self.requests[row.message_id] = row
if task:
self.tasks[task.id] = task
return row, True
async def text_cache(self, digest, version):
return self.text.get((digest, version))
async def put_text_cache(self, row):
self.text[(row.analysis_sha256, row.rules_version)] = row
async def file_cache(self, digest, config, signatures_version):
return None
async def task(self, task_id):
return self.tasks.get(task_id)
async def audit(self, row):
self.audits.append(row)
class ForbiddenResolver:
async def resolve(self, hostname):
raise AssertionError("MOCK must not call DNS")
def body(kind: str, message_id=None) -> dict:
value = {
"message_id": str(message_id or uuid4()),
"content_kind": kind,
"text": "hello" if kind == "text" else "",
"attachment": None,
}
if kind == "file":
value["attachment"] = {
"attachment_id": str(uuid4()),
"quarantine_object_key": (
"quarantine/users/00000000-0000-4000-8000-000000000001/"
"dialogs/00000000-0000-4000-8000-000000000002/"
"00000000-0000-4000-8000-000000000003"
),
"quarantine_version_id": "v1",
"quarantine_etag": '"e"',
"mime_type": "application/pdf",
"size_bytes": 10,
"checksum": "sha256:" + "0" * 64,
}
return value
@pytest.mark.parametrize(
"text_free,file_free,kind,status",
[
(True, True, "text", 200),
(True, True, "file", 200),
(True, False, "text", 200),
(True, False, "file", 403),
(False, True, "text", 403),
(False, True, "file", 200),
(False, False, "text", 403),
(False, False, "file", 403),
],
)
async def test_mock_2x2_is_sync(active_config, text_free, file_free, kind, status) -> None:
repo = FakeRepository()
service = SafetyService(
repo,
active_config,
EmergencyMode(mock=True, text_free=text_free, file_free=file_free),
ForbiddenResolver(),
)
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=create_app(service, "secret")),
base_url="http://test",
headers={"X-Service-Token": "secret"},
) as client:
response = await client.post("/internal/safety/v2/messages/check", json=body(kind))
assert response.status_code == status
assert response.json()["processing_mode"] == "mock"
assert response.json()["verdict"] in {"allow", "deny"}
assert len(repo.audits) == 1
async def test_auth_strict_dto_idempotency_and_conflict(active_config) -> None:
repo = FakeRepository()
service = SafetyService(repo, active_config, EmergencyMode(), ForbiddenResolver())
app = create_app(service, "secret")
message_id = uuid4()
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app), base_url="http://test"
) as client:
assert (
await client.post("/internal/safety/v2/messages/check", json=body("text"))
).status_code == 401
invalid = body("text")
invalid["unknown"] = True
assert (
await client.post(
"/internal/safety/v2/messages/check",
json=invalid,
headers={"X-Service-Token": "secret"},
)
).status_code == 400
headers = {"X-Service-Token": "secret"}
first = await client.post(
"/internal/safety/v2/messages/check", json=body("text", message_id), headers=headers
)
replay = await client.post(
"/internal/safety/v2/messages/check", json=body("text", message_id), headers=headers
)
changed = body("text", message_id)
changed["text"] = "different"
conflict = await client.post(
"/internal/safety/v2/messages/check", json=changed, headers=headers
)
assert first.status_code == replay.status_code == 200
assert first.json() == replay.json()
assert conflict.status_code == 409
async def test_standard_text_deny_and_file_pending(active_config) -> None:
repo = FakeRepository()
service = SafetyService(repo, active_config, EmergencyMode(), ForbiddenResolver())
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=create_app(service, "secret")),
base_url="http://test",
headers={"X-Service-Token": "secret"},
) as client:
denied = body("text")
denied["text"] = "<script>alert(1)</script>"
deny_response = await client.post("/internal/safety/v2/messages/check", json=denied)
pending_response = await client.post(
"/internal/safety/v2/messages/check", json=body("file")
)
task_response = await client.get(pending_response.headers["Location"])
assert deny_response.status_code == 403
assert deny_response.json()["rule_id"] == "text.active_script"
assert pending_response.status_code == task_response.status_code == 202
assert pending_response.json() == task_response.json()
assert pending_response.headers["Retry-After"] == "2"
async def test_final_task_response_keeps_task_config_snapshot(active_config) -> None:
repo = FakeRepository()
service = SafetyService(repo, active_config, EmergencyMode(), ForbiddenResolver())
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=create_app(service, "secret")),
base_url="http://test",
headers={"X-Service-Token": "secret"},
) as client:
pending = await client.post("/internal/safety/v2/messages/check", json=body("file"))
task = repo.tasks[next(iter(repo.tasks))]
task.status = TaskStatus.allowed
task.verdict = "allow"
task.rule_id = "safety.all_checks_passed"
task.config_version = 99
final = await client.get(pending.headers["Location"])
assert final.status_code == 200
assert final.json()["config_version"] == 99