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"] = "" 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