191 lines
6.7 KiB
Python
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
|