Проект разделен на два репозитория
This commit is contained in:
@@ -0,0 +1,190 @@
|
||||
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
|
||||
Reference in New Issue
Block a user