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)