Проект разделен на два репозитория
This commit is contained in:
@@ -0,0 +1,334 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import random
|
||||
import unicodedata
|
||||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Annotated, Any, Literal, Protocol
|
||||
|
||||
import redis.asyncio as redis
|
||||
import uvicorn
|
||||
from fastapi import Depends, FastAPI, Header, HTTPException, Request
|
||||
from fastapi.exceptions import RequestValidationError
|
||||
from fastapi.responses import JSONResponse
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = SettingsConfigDict(extra="ignore")
|
||||
app_env: str = "production-like"
|
||||
message_safety_redis_url: str = "redis://redis:6379/2"
|
||||
message_safety_service_token: str = Field(min_length=16)
|
||||
message_safety_rules_version: str = "2026-01-01"
|
||||
message_safety_task_ttl_sec: int = Field(default=900, ge=330)
|
||||
message_safety_poll_after_ms: int = Field(default=2000, ge=100, le=30000)
|
||||
safety_stub_worker_mode: Literal["emulated_on_poll"] = "emulated_on_poll"
|
||||
safety_stub_rng_seed: int | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def forbid_seed_outside_tests(self) -> Settings:
|
||||
if self.safety_stub_rng_seed is not None and self.app_env != "test":
|
||||
raise ValueError("SAFETY_STUB_RNG_SEED is allowed only when APP_ENV=test")
|
||||
return self
|
||||
|
||||
|
||||
class Attachment(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
attachment_id: uuid.UUID
|
||||
quarantine_object_key: str = Field(min_length=1, max_length=1024)
|
||||
mime_type: str = Field(min_length=1, max_length=255)
|
||||
size_bytes: int = Field(ge=0, le=10 * 1024 * 1024)
|
||||
checksum: str = Field(pattern=r"^sha256:[0-9a-fA-F]{64}$")
|
||||
|
||||
|
||||
class CheckRequest(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
message_id: uuid.UUID
|
||||
content_kind: Literal["text", "file"]
|
||||
text: str = Field(default="", max_length=10000)
|
||||
attachment: Attachment | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_kind(self) -> CheckRequest:
|
||||
if self.content_kind == "file" and self.attachment is None:
|
||||
raise ValueError("attachment is required for file content")
|
||||
if self.content_kind == "text" and self.attachment is not None:
|
||||
raise ValueError("attachment is forbidden for text content")
|
||||
return self
|
||||
|
||||
|
||||
class TaskStore(Protocol):
|
||||
async def reserve(self, message_id: str, fingerprint: str, task_id: str, ttl: int) -> str: ...
|
||||
async def poll(self, task_id: str) -> int | None: ...
|
||||
async def ready(self) -> bool: ...
|
||||
async def close(self) -> None: ...
|
||||
|
||||
|
||||
class RedisTaskStore:
|
||||
_reserve_lua = """
|
||||
local prior = redis.call('GET', KEYS[1])
|
||||
if prior then
|
||||
local sep = string.find(prior, '|', 1, true)
|
||||
local old_fp = string.sub(prior, 1, sep - 1)
|
||||
if old_fp ~= ARGV[1] then return {'conflict'} end
|
||||
return {'existing', string.sub(prior, sep + 1)}
|
||||
end
|
||||
redis.call('HSET', KEYS[2], 'schema_version', '1', 'message_id', ARGV[2],
|
||||
'created_at_ms', ARGV[3], 'poll_count', '0', 'rules_version', ARGV[5])
|
||||
redis.call('EXPIRE', KEYS[2], ARGV[4])
|
||||
redis.call('SET', KEYS[1], ARGV[1] .. '|' .. ARGV[6], 'EX', ARGV[4])
|
||||
return {'created', ARGV[6]}
|
||||
"""
|
||||
_poll_lua = """
|
||||
if redis.call('EXISTS', KEYS[1]) == 0 then return nil end
|
||||
return redis.call('HINCRBY', KEYS[1], 'poll_count', 1)
|
||||
"""
|
||||
|
||||
def __init__(self, client: redis.Redis, rules_version: str) -> None:
|
||||
self.client = client
|
||||
self.rules_version = rules_version
|
||||
|
||||
async def reserve(self, message_id: str, fingerprint: str, task_id: str, ttl: int) -> str:
|
||||
result = await self.client.eval(
|
||||
self._reserve_lua,
|
||||
2,
|
||||
f"han:safety:task-by-message:{message_id}",
|
||||
f"han:safety:task:{task_id}",
|
||||
fingerprint,
|
||||
message_id,
|
||||
str(int(datetime.now(UTC).timestamp() * 1000)),
|
||||
str(ttl),
|
||||
self.rules_version,
|
||||
task_id,
|
||||
)
|
||||
status = _decode(result[0])
|
||||
if status == "conflict":
|
||||
raise ValueError("conflict")
|
||||
return _decode(result[1])
|
||||
|
||||
async def poll(self, task_id: str) -> int | None:
|
||||
key = f"han:safety:task:{task_id}"
|
||||
count = await self.client.eval(self._poll_lua, 1, key)
|
||||
return int(count) if count is not None else None
|
||||
|
||||
async def ready(self) -> bool:
|
||||
key = f"han:safety:ready:{uuid.uuid4()}"
|
||||
try:
|
||||
await self.client.set(key, "1", ex=5)
|
||||
return await self.client.get(key) == b"1"
|
||||
finally:
|
||||
await self.client.delete(key)
|
||||
|
||||
async def close(self) -> None:
|
||||
await self.client.aclose()
|
||||
|
||||
|
||||
def _decode(value: Any) -> str:
|
||||
return value.decode() if isinstance(value, bytes) else str(value)
|
||||
|
||||
|
||||
def normalize(text: str) -> str:
|
||||
return unicodedata.normalize("NFKC", text.replace("\r\n", "\n").replace("\r", "\n")).lstrip()
|
||||
|
||||
|
||||
def fingerprint(dto: CheckRequest) -> str:
|
||||
body = dto.model_dump(mode="json")
|
||||
body["text"] = normalize(dto.text)
|
||||
encoded = json.dumps(body, ensure_ascii=False, sort_keys=True, separators=(",", ":")).encode()
|
||||
return hashlib.sha256(encoded).hexdigest()
|
||||
|
||||
|
||||
def error(code: str, message: str, request_id: str, details: dict[str, Any] | None = None) -> dict:
|
||||
return {
|
||||
"error": {
|
||||
"code": code,
|
||||
"message": message,
|
||||
"request_id": request_id,
|
||||
"details": details or {},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def create_app(
|
||||
settings: Settings | None = None,
|
||||
store: TaskStore | None = None,
|
||||
rng: random.Random | None = None,
|
||||
) -> FastAPI:
|
||||
cfg = settings or Settings()
|
||||
verdict_rng = rng or random.Random(cfg.safety_stub_rng_seed)
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
app.state.store = store or RedisTaskStore(
|
||||
redis.from_url(cfg.message_safety_redis_url, decode_responses=False),
|
||||
cfg.message_safety_rules_version,
|
||||
)
|
||||
yield
|
||||
await app.state.store.close()
|
||||
|
||||
app = FastAPI(
|
||||
title="HAN Message Safety",
|
||||
version="1.0.0",
|
||||
lifespan=lifespan,
|
||||
docs_url=None if cfg.app_env != "test" else "/docs",
|
||||
)
|
||||
app.state.settings = cfg
|
||||
|
||||
@app.middleware("http")
|
||||
async def request_id_middleware(request: Request, call_next):
|
||||
request.state.request_id = request.headers.get("X-Request-ID") or str(uuid.uuid4())
|
||||
response = await call_next(request)
|
||||
response.headers["X-Request-ID"] = request.state.request_id
|
||||
return response
|
||||
|
||||
def authorize(
|
||||
request: Request,
|
||||
token: Annotated[str | None, Header(alias="X-Service-Token")] = None,
|
||||
) -> None:
|
||||
if token is None or not hmac.compare_digest(token, cfg.message_safety_service_token):
|
||||
raise HTTPException(
|
||||
401,
|
||||
error(
|
||||
"service_unauthorized",
|
||||
"Service authentication failed",
|
||||
request.state.request_id,
|
||||
),
|
||||
)
|
||||
|
||||
@app.exception_handler(HTTPException)
|
||||
async def http_error(_: Request, exc: HTTPException):
|
||||
return JSONResponse(status_code=exc.status_code, content=exc.detail)
|
||||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
async def validation_error(request: Request, _: RequestValidationError):
|
||||
return JSONResponse(
|
||||
error("validation_error", "Request is invalid", request.state.request_id),
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
@app.get("/health/live")
|
||||
async def live() -> dict[str, str]:
|
||||
return {"status": "live"}
|
||||
|
||||
@app.get("/health/ready")
|
||||
async def ready(request: Request):
|
||||
try:
|
||||
ok = await request.app.state.store.ready()
|
||||
except Exception:
|
||||
ok = False
|
||||
body = {
|
||||
"status": "ready" if ok else "not_ready",
|
||||
"components": {"redis": "ok" if ok else "down"},
|
||||
}
|
||||
return JSONResponse(body, status_code=200 if ok else 503)
|
||||
|
||||
@app.post("/internal/safety/v1/messages/check", dependencies=[Depends(authorize)])
|
||||
async def check(dto: CheckRequest, request: Request):
|
||||
text = normalize(dto.text)
|
||||
common = {"rules_version": cfg.message_safety_rules_version}
|
||||
if text and text[0] in {"ф", "Ф"}:
|
||||
return JSONResponse(
|
||||
{
|
||||
"verdict": "deny",
|
||||
"rule_id": "stub.starts_with_cyrillic_ef",
|
||||
"reason_code": "stub_blocked",
|
||||
**common,
|
||||
},
|
||||
status_code=403,
|
||||
)
|
||||
if text and unicodedata.category(text[0]) == "Nd":
|
||||
task_id = str(uuid.uuid4())
|
||||
try:
|
||||
task_id = await request.app.state.store.reserve(
|
||||
str(dto.message_id), fingerprint(dto), task_id, cfg.message_safety_task_ttl_sec
|
||||
)
|
||||
except ValueError:
|
||||
return JSONResponse(
|
||||
error(
|
||||
"safety_request_conflict", "message_id was reused", request.state.request_id
|
||||
),
|
||||
status_code=409,
|
||||
)
|
||||
except Exception:
|
||||
return JSONResponse(
|
||||
error(
|
||||
"redis_unavailable", "Task storage is unavailable", request.state.request_id
|
||||
),
|
||||
status_code=503,
|
||||
)
|
||||
return JSONResponse(
|
||||
{
|
||||
"verdict": "pending",
|
||||
"task_id": task_id,
|
||||
"poll_after_ms": cfg.message_safety_poll_after_ms,
|
||||
"expires_at": (
|
||||
datetime.now(UTC) + timedelta(seconds=cfg.message_safety_task_ttl_sec)
|
||||
)
|
||||
.isoformat()
|
||||
.replace("+00:00", "Z"),
|
||||
**common,
|
||||
},
|
||||
status_code=203,
|
||||
)
|
||||
return {"verdict": "allow", "rule_id": "stub.default_allow", **common}
|
||||
|
||||
@app.get("/internal/safety/v1/messages/tasks/{task_id}", dependencies=[Depends(authorize)])
|
||||
async def task(task_id: str, request: Request):
|
||||
try:
|
||||
parsed = str(uuid.UUID(task_id))
|
||||
except ValueError:
|
||||
return JSONResponse(
|
||||
error("validation_error", "Request is invalid", request.state.request_id),
|
||||
status_code=400,
|
||||
)
|
||||
try:
|
||||
count = await request.app.state.store.poll(parsed)
|
||||
except Exception:
|
||||
return JSONResponse(
|
||||
error("redis_unavailable", "Task storage is unavailable", request.state.request_id),
|
||||
status_code=503,
|
||||
)
|
||||
if count is None:
|
||||
return JSONResponse(
|
||||
error("task_not_found", "Task was not found", request.state.request_id),
|
||||
status_code=404,
|
||||
)
|
||||
outcome = verdict_rng.choice(("pending", "allow", "final_error"))
|
||||
if outcome == "pending":
|
||||
return JSONResponse(
|
||||
{
|
||||
"verdict": "pending",
|
||||
"task_id": parsed,
|
||||
"poll_after_ms": cfg.message_safety_poll_after_ms,
|
||||
},
|
||||
status_code=203,
|
||||
)
|
||||
if outcome == "allow":
|
||||
return {"verdict": "allow", "task_id": parsed, "rule_id": "stub.random_allow"}
|
||||
return JSONResponse(
|
||||
{
|
||||
"verdict": "deny",
|
||||
"task_id": parsed,
|
||||
**error(
|
||||
"stub_final_error",
|
||||
"Stub task returned a final negative verdict",
|
||||
request.state.request_id,
|
||||
{"terminal": True},
|
||||
),
|
||||
},
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
return app
|
||||
|
||||
|
||||
app = create_app()
|
||||
|
||||
|
||||
def run() -> None:
|
||||
uvicorn.run("app.main:app", host="0.0.0.0", port=8080)
|
||||
Reference in New Issue
Block a user