Init
This commit is contained in:
@@ -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)]
|
||||
@@ -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))
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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()
|
||||
@@ -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"]
|
||||
@@ -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}
|
||||
Reference in New Issue
Block a user