This commit is contained in:
2026-07-17 15:57:36 +03:00
commit 8f2d798e1a
72 changed files with 6227 additions and 0 deletions
View File
+93
View File
@@ -0,0 +1,93 @@
from __future__ import annotations
import secrets
from typing import Annotated
from fastapi import Depends, HTTPException, Request
from fastapi.security import HTTPBasic, HTTPBasicCredentials
from itsdangerous import BadSignature, SignatureExpired, URLSafeTimedSerializer
from starlette.responses import Response
from app.config import get_settings
http_basic = HTTPBasic(auto_error=False)
def _serializer() -> URLSafeTimedSerializer:
settings = get_settings()
return URLSafeTimedSerializer(settings.app_secret_key, salt="wrapped-admin")
def create_session_token(username: str) -> str:
return _serializer().dumps({"u": username})
def read_session_token(token: str, max_age: int) -> str | None:
try:
data = _serializer().loads(token, max_age=max_age)
return data.get("u")
except (BadSignature, SignatureExpired, Exception):
return None
def set_admin_cookie(response: Response, username: str) -> None:
settings = get_settings()
token = create_session_token(username)
response.set_cookie(
settings.session_cookie_name,
token,
httponly=True,
samesite="lax",
max_age=settings.session_max_age,
secure=settings.app_env == "production",
path="/",
)
def clear_admin_cookie(response: Response) -> None:
settings = get_settings()
response.delete_cookie(settings.session_cookie_name, path="/")
def get_admin_user(request: Request) -> str | None:
settings = get_settings()
token = request.cookies.get(settings.session_cookie_name)
if not token:
return None
return read_session_token(token, settings.session_max_age)
def _basic_ok(credentials: HTTPBasicCredentials) -> bool:
settings = get_settings()
user_ok = secrets.compare_digest(credentials.username, settings.admin_username)
pass_ok = secrets.compare_digest(credentials.password, settings.admin_password)
return user_ok and pass_ok
async def require_admin_web(request: Request) -> str:
"""Cookie-only auth for HTML admin pages (redirect via exception handler)."""
cookie_user = get_admin_user(request)
if cookie_user:
return cookie_user
raise HTTPException(status_code=401, detail="unauthorized")
async def require_admin_api(
request: Request,
credentials: Annotated[HTTPBasicCredentials | None, Depends(http_basic)] = None,
) -> str:
"""Cookie or HTTP Basic — for JSON admin API / Swagger."""
cookie_user = get_admin_user(request)
if cookie_user:
return cookie_user
if credentials and _basic_ok(credentials):
return credentials.username
raise HTTPException(
status_code=401,
detail="unauthorized",
headers={"WWW-Authenticate": "Basic"},
)
AdminWebAuth = Annotated[str, Depends(require_admin_web)]
AdminAuth = Annotated[str, Depends(require_admin_api)]
+97
View File
@@ -0,0 +1,97 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import delete, func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models import AuditEvent
KNOWN_AUDIT_EVENT_TYPES = [
"wrap.create",
"wrap.unwrap",
"admin.login",
"admin.logout",
"admin.settings_update",
"admin.purge",
]
async def write_audit(
db: AsyncSession,
*,
event_type: str,
success: bool,
wrap_id: str | None = None,
ip: str | None = None,
user_agent: str | None = None,
accept_language: str | None = None,
forwarded_for: str | None = None,
details: dict[str, Any] | None = None,
) -> None:
event = AuditEvent(
event_type=event_type,
success=success,
wrap_id=wrap_id,
ip=ip,
user_agent=(user_agent or "")[:512] or None,
accept_language=(accept_language or "")[:128] or None,
forwarded_for=(forwarded_for or "")[:256] or None,
details=details or {},
)
db.add(event)
await db.commit()
async def purge_old_audit(db: AsyncSession, retention_days: int) -> int:
cutoff = datetime.now(timezone.utc) - timedelta(days=retention_days)
result = await db.execute(delete(AuditEvent).where(AuditEvent.created_at < cutoff))
await db.commit()
return result.rowcount or 0
def _audit_filters(event_type: str | None, wrap_id: str | None):
filters = []
if event_type:
filters.append(AuditEvent.event_type == event_type)
if wrap_id:
filters.append(AuditEvent.wrap_id == wrap_id)
return filters
async def list_audit(
db: AsyncSession,
*,
limit: int = 100,
offset: int = 0,
event_type: str | None = None,
wrap_id: str | None = None,
) -> list[AuditEvent]:
stmt = select(AuditEvent).order_by(AuditEvent.created_at.desc()).offset(offset).limit(limit)
for f in _audit_filters(event_type, wrap_id):
stmt = stmt.where(f)
result = await db.execute(stmt)
return list(result.scalars().all())
async def count_audit(
db: AsyncSession,
*,
event_type: str | None = None,
wrap_id: str | None = None,
) -> int:
stmt = select(func.count()).select_from(AuditEvent)
for f in _audit_filters(event_type, wrap_id):
stmt = stmt.where(f)
result = await db.execute(stmt)
return int(result.scalar_one() or 0)
async def list_audit_event_types(db: AsyncSession) -> list[str]:
"""Known types plus distinct values already stored in DB (dynamic dropdown)."""
result = await db.execute(
select(AuditEvent.event_type).distinct().order_by(AuditEvent.event_type)
)
from_db = [row[0] for row in result.all() if row[0]]
return sorted(set(KNOWN_AUDIT_EVENT_TYPES) | set(from_db))
+59
View File
@@ -0,0 +1,59 @@
from __future__ import annotations
import httpx
from app.config import get_settings
from app.models import CaptchaProvider
async def verify_captcha(
provider: CaptchaProvider,
token: str | None,
remote_ip: str | None,
*,
turnstile_site_configured: bool,
hcaptcha_site_configured: bool,
) -> tuple[bool, str]:
if provider == CaptchaProvider.off:
return True, "off"
if not token:
return False, "missing_token"
settings = get_settings()
if provider == CaptchaProvider.turnstile:
if not turnstile_site_configured:
return False, "turnstile_not_configured"
secret = settings.turnstile_secret_key
if not secret:
return False, "turnstile_secret_missing"
return await _post_verify(
"https://challenges.cloudflare.com/turnstile/v0/siteverify",
{"secret": secret, "response": token, "remoteip": remote_ip or ""},
)
if provider == CaptchaProvider.hcaptcha:
if not hcaptcha_site_configured:
return False, "hcaptcha_not_configured"
secret = settings.hcaptcha_secret_key
if not secret:
return False, "hcaptcha_secret_missing"
return await _post_verify(
"https://hcaptcha.com/siteverify",
{"secret": secret, "response": token, "remoteip": remote_ip or ""},
)
return False, "unknown_provider"
async def _post_verify(url: str, data: dict) -> tuple[bool, str]:
try:
async with httpx.AsyncClient(timeout=10.0) as client:
resp = await client.post(url, data=data)
payload = resp.json()
if payload.get("success"):
return True, "ok"
return False, "verify_failed"
except Exception:
return False, "verify_error"
+47
View File
@@ -0,0 +1,47 @@
from __future__ import annotations
from datetime import datetime, timezone
from sqlalchemy import and_, or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models import Wrap, WrapStatus
from app.services.storage import storage
async def purge_expired_wraps(db: AsyncSession) -> int:
now = datetime.now(timezone.utc)
result = await db.execute(
select(Wrap).where(
and_(
Wrap.status == WrapStatus.pending,
Wrap.expires_at < now,
)
)
)
wraps = list(result.scalars().all())
count = 0
for wrap in wraps:
try:
await storage.delete(wrap.object_key)
except Exception:
pass
wrap.status = WrapStatus.expired
count += 1
if count:
await db.commit()
return count
async def cleanup_consumed_orphans(db: AsyncSession) -> int:
"""Best-effort: mark very old pending with missing logic already handled elsewhere."""
result = await db.execute(
select(Wrap).where(
or_(
Wrap.status == WrapStatus.consumed,
Wrap.status == WrapStatus.expired,
)
).limit(0)
)
_ = result
return 0
+139
View File
@@ -0,0 +1,139 @@
from __future__ import annotations
import ipaddress
import re
from fastapi import Request
from app.config import get_settings
_FORWARDED_FOR_RE = re.compile(r"for=(?:\"?\[?)([^;\"\\]\S*)", re.IGNORECASE)
def _parse_ip(value: str | None) -> str | None:
if not value:
return None
raw = value.strip().strip('"').strip()
if not raw or raw.lower() == "unknown":
return None
# [IPv6]:port or IPv4:port
if raw.startswith("["):
end = raw.find("]")
if end != -1:
raw = raw[1:end]
elif raw.count(":") == 1 and "." in raw:
raw = raw.rsplit(":", 1)[0]
try:
addr = ipaddress.ip_address(raw)
except ValueError:
return None
# Normalize IPv4-mapped IPv6 (::ffff:1.2.3.4 → 1.2.3.4)
if isinstance(addr, ipaddress.IPv6Address) and addr.ipv4_mapped:
return str(addr.ipv4_mapped)
return str(addr)
def _peer_ip(request: Request) -> str | None:
if request.client and request.client.host:
return _parse_ip(request.client.host) or request.client.host
return None
def _trusted_proxy_nets() -> list[ipaddress.IPv4Network | ipaddress.IPv6Network] | None:
"""None means trust any peer (TRUSTED_PROXIES=*)."""
raw = (get_settings().trusted_proxies or "").strip()
if not raw or raw == "*":
return None
nets: list[ipaddress.IPv4Network | ipaddress.IPv6Network] = []
for part in raw.split(","):
part = part.strip()
if not part:
continue
try:
if "/" in part:
nets.append(ipaddress.ip_network(part, strict=False))
else:
ip = ipaddress.ip_address(part)
nets.append(ipaddress.ip_network(f"{ip}/{ip.max_prefixlen}"))
except ValueError:
continue
return nets
def _is_trusted_peer(peer: str | None) -> bool:
if not peer:
return False
nets = _trusted_proxy_nets()
if nets is None:
return True
try:
addr = ipaddress.ip_address(peer)
except ValueError:
return False
return any(addr in net for net in nets)
def _ips_from_forwarded_for(header: str | None) -> list[str]:
if not header:
return []
out: list[str] = []
for part in header.split(","):
ip = _parse_ip(part)
if ip:
out.append(ip)
return out
def _ips_from_forwarded(header: str | None) -> list[str]:
if not header:
return []
out: list[str] = []
for match in _FORWARDED_FOR_RE.finditer(header):
ip = _parse_ip(match.group(1))
if ip:
out.append(ip)
return out
def client_ip(request: Request) -> str | None:
"""
Resolve the original client IP.
When the immediate peer is a trusted proxy, prefer (in order):
CF-Connecting-IP, True-Client-IP, X-Real-IP, left-most X-Forwarded-For,
RFC 7239 Forwarded. Otherwise use the TCP peer address.
"""
peer = _peer_ip(request)
if _is_trusted_peer(peer):
for header in ("cf-connecting-ip", "true-client-ip", "x-real-ip"):
ip = _parse_ip(request.headers.get(header))
if ip:
return ip
chain = _ips_from_forwarded_for(request.headers.get("x-forwarded-for"))
if chain:
return chain[0]
chain = _ips_from_forwarded(request.headers.get("forwarded"))
if chain:
return chain[0]
return peer
def request_meta(request: Request) -> dict[str, str | None]:
forwarded = request.headers.get("x-forwarded-for")
real_ip = request.headers.get("x-real-ip")
cf_ip = request.headers.get("cf-connecting-ip")
forwarded_store = forwarded or real_ip or cf_ip
return {
"ip": client_ip(request),
"user_agent": request.headers.get("user-agent"),
"accept_language": request.headers.get("accept-language"),
"forwarded_for": (forwarded_store or "")[:256] or None,
}
def uvicorn_forwarded_allow_ips() -> str:
"""Value for uvicorn --forwarded-allow-ips (CIDRs not supported → * or IP list)."""
raw = (get_settings().trusted_proxies or "").strip()
if not raw or raw == "*" or "/" in raw:
return "*"
return raw
+19
View File
@@ -0,0 +1,19 @@
from __future__ import annotations
from argon2 import PasswordHasher
from argon2.exceptions import VerifyMismatchError
_ph = PasswordHasher()
def hash_password(password: str) -> str:
return _ph.hash(password)
def verify_password(password_hash: str, password: str) -> bool:
try:
return _ph.verify(password_hash, password)
except VerifyMismatchError:
return False
except Exception:
return False
+29
View File
@@ -0,0 +1,29 @@
from __future__ import annotations
import time
from collections import defaultdict, deque
from threading import Lock
class RateLimiter:
"""In-memory sliding window limiter (per process). Good enough for single-replica / dev."""
def __init__(self) -> None:
self._hits: dict[str, deque[float]] = defaultdict(deque)
self._lock = Lock()
def allow(self, key: str, limit: int, window_seconds: int = 60) -> bool:
if limit <= 0:
return True
now = time.monotonic()
with self._lock:
q = self._hits[key]
while q and now - q[0] > window_seconds:
q.popleft()
if len(q) >= limit:
return False
q.append(now)
return True
rate_limiter = RateLimiter()
+72
View File
@@ -0,0 +1,72 @@
from __future__ import annotations
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models import (
DEFAULT_MIME_ALLOWLIST,
AppSettings,
CaptchaProvider,
PasswordMode,
)
PASSWORD_MODE_HELP = {
"client_only": (
"Password is verified only in the browser after ciphertext download. "
"Maximum zero-knowledge: the server never checks the password. "
"Best when you trust link secrecy and want strongest ZK."
),
"server_gate": (
"Server stores an Argon2id hash and releases ciphertext only after a correct password. "
"Slightly weaker ZK metadata (server knows password success/fail) but stops offline "
"bruteforce if someone steals the share link without the password."
),
}
async def get_or_create_settings(db: AsyncSession) -> AppSettings:
result = await db.execute(select(AppSettings).where(AppSettings.id == 1))
row = result.scalar_one_or_none()
if row:
return row
row = AppSettings(
id=1,
mime_allowlist=list(DEFAULT_MIME_ALLOWLIST),
captcha_provider=CaptchaProvider.off,
password_mode=PasswordMode.client_only,
)
db.add(row)
await db.commit()
await db.refresh(row)
return row
def mime_allowed(allowlist: list[str], content_type: str) -> bool:
ct = (content_type or "").split(";")[0].strip().lower()
if not ct:
return False
for pattern in allowlist:
p = pattern.strip().lower()
if not p:
continue
if p.endswith("/*"):
if ct.startswith(p[:-1]):
return True
elif ct == p:
return True
return False
def public_settings_payload(row: AppSettings) -> dict:
return {
"max_upload_bytes": row.max_upload_bytes,
"default_ttl_seconds": row.default_ttl_seconds,
"max_ttl_seconds": row.max_ttl_seconds,
"mime_allowlist": list(row.mime_allowlist or []),
"captcha_provider": row.captcha_provider.value,
"turnstile_site_key": row.turnstile_site_key or "",
"hcaptcha_site_key": row.hcaptcha_site_key or "",
"password_mode": row.password_mode.value,
"password_mode_description": PASSWORD_MODE_HELP,
}
+89
View File
@@ -0,0 +1,89 @@
from __future__ import annotations
from contextlib import asynccontextmanager
from typing import Any, AsyncIterator
from aiobotocore.session import get_session
from app.config import get_settings
class ObjectStorage:
def __init__(self) -> None:
self.settings = get_settings()
def _client_kwargs(self) -> dict[str, Any]:
return {
"endpoint_url": self.settings.s3_endpoint_url,
"aws_access_key_id": self.settings.s3_access_key,
"aws_secret_access_key": self.settings.s3_secret_key,
"region_name": self.settings.s3_region,
"use_ssl": self.settings.s3_use_ssl,
}
@asynccontextmanager
async def client(self) -> AsyncIterator[Any]:
session = get_session()
async with session.create_client("s3", **self._client_kwargs()) as client:
yield client
async def ensure_bucket(self) -> None:
if not self.settings.s3_create_bucket:
return
async with self.client() as client:
try:
await client.head_bucket(Bucket=self.settings.s3_bucket)
except Exception:
await client.create_bucket(Bucket=self.settings.s3_bucket)
async def put_bytes(self, key: str, data: bytes, content_type: str = "application/octet-stream") -> None:
async with self.client() as client:
await client.put_object(
Bucket=self.settings.s3_bucket,
Key=key,
Body=data,
ContentType=content_type,
)
async def get_bytes(self, key: str) -> bytes:
async with self.client() as client:
resp = await client.get_object(Bucket=self.settings.s3_bucket, Key=key)
async with resp["Body"] as stream:
return await stream.read()
async def delete(self, key: str) -> None:
async with self.client() as client:
await client.delete_object(Bucket=self.settings.s3_bucket, Key=key)
async def delete_prefix(self, prefix: str = "wraps/") -> int:
"""Delete all objects under prefix. Returns number of deleted keys."""
deleted = 0
async with self.client() as client:
token: str | None = None
while True:
kwargs: dict[str, Any] = {
"Bucket": self.settings.s3_bucket,
"Prefix": prefix,
"MaxKeys": 1000,
}
if token:
kwargs["ContinuationToken"] = token
resp = await client.list_objects_v2(**kwargs)
contents = resp.get("Contents") or []
if contents:
# delete_objects accepts up to 1000 keys
await client.delete_objects(
Bucket=self.settings.s3_bucket,
Delete={
"Objects": [{"Key": obj["Key"]} for obj in contents],
"Quiet": True,
},
)
deleted += len(contents)
if not resp.get("IsTruncated"):
break
token = resp.get("NextContinuationToken")
return deleted
storage = ObjectStorage()
+17
View File
@@ -0,0 +1,17 @@
from __future__ import annotations
import secrets
from datetime import datetime, timedelta, timezone
from app.services.client_ip import client_ip, request_meta
def new_wrap_id() -> str:
return secrets.token_urlsafe(18).replace("-", "").replace("_", "")[:24]
def expires_at_from_ttl(ttl_seconds: int) -> datetime:
return datetime.now(timezone.utc) + timedelta(seconds=ttl_seconds)
__all__ = ["client_ip", "request_meta", "new_wrap_id", "expires_at_from_ttl"]
+75
View File
@@ -0,0 +1,75 @@
from __future__ import annotations
import secrets
from datetime import datetime, timezone
from sqlalchemy import delete, select, update
from sqlalchemy.ext.asyncio import AsyncSession
from app.models import Wrap, WrapStatus
from app.services.client_ip import client_ip, request_meta
from app.services.storage import storage
def new_wrap_id() -> str:
return secrets.token_urlsafe(18).replace("-", "").replace("_", "")[:24]
async def expire_due_wraps(db: AsyncSession) -> int:
now = datetime.now(timezone.utc)
result = await db.execute(
select(Wrap).where(Wrap.status == WrapStatus.pending, Wrap.expires_at <= now)
)
rows = list(result.scalars().all())
for wrap in rows:
try:
await storage.delete(wrap.object_key)
except Exception:
pass
wrap.status = WrapStatus.expired
if rows:
await db.commit()
return len(rows)
async def delete_wrap_object(db: AsyncSession, wrap: Wrap) -> None:
try:
await storage.delete(wrap.object_key)
except Exception:
pass
async def cleanup_expired_meta(db: AsyncSession) -> int:
"""Remove long-expired wrap rows (objects already deleted)."""
now = datetime.now(timezone.utc)
result = await db.execute(
delete(Wrap).where(
Wrap.status.in_([WrapStatus.expired, WrapStatus.consumed]),
Wrap.expires_at < now,
)
)
await db.commit()
return result.rowcount or 0
async def mark_consumed(db: AsyncSession, wrap_id: str) -> Wrap | None:
now = datetime.now(timezone.utc)
result = await db.execute(
update(Wrap)
.where(Wrap.id == wrap_id, Wrap.status == WrapStatus.pending, Wrap.expires_at > now)
.values(status=WrapStatus.consumed, consumed_at=now)
.returning(Wrap)
)
row = result.scalar_one_or_none()
await db.commit()
return row
async def purge_all_wraps(db: AsyncSession) -> dict[str, int]:
"""Force-delete all wrap DB rows and all objects under wraps/ in MinIO."""
count_result = await db.execute(select(Wrap.id))
db_count = len(list(count_result.scalars().all()))
await db.execute(delete(Wrap))
await db.commit()
objects_deleted = await storage.delete_prefix("wraps/")
return {"wraps_deleted": db_count, "objects_deleted": objects_deleted}