tmp-files / app.py
jiayi.xie
feat(mirror): allow relay gateway hosts
c48ce72
Raw
History Blame Contribute Delete
5.93 kB
import os
import pathlib
import uuid
from urllib.parse import urlparse
import httpx
from fastapi import FastAPI, File, HTTPException, Request, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse, StreamingResponse
DATA_DIR = pathlib.Path(os.environ.get("DATA_DIR", "/tmp/uploads"))
DATA_DIR.mkdir(parents=True, exist_ok=True)
MAX_BYTES = int(os.environ.get("MAX_BYTES", str(30 * 1024 * 1024)))
# /mirror streams somebody else's object storage back with CORS headers. It stores nothing.
MIRROR_MAX_BYTES = int(os.environ.get("MIRROR_MAX_BYTES", str(300 * 1024 * 1024)))
DEFAULT_MIRROR_HOSTS = (
# object storage the generated videos actually live on
"aliyuncs.com,r2.cloudflarestorage.com,amazonaws.com,myqcloud.com,bcebos.com,"
"volces.com,byteimg.com,volccdn.com,googleapis.com,cloudfront.net,oaiusercontent.com,"
# relay gateways whose own /v1/videos/.../content links may need mirroring
"diwdiw.cn,hjmie.cc.cd"
)
MIRROR_ALLOW_HOSTS = tuple(host.strip().lower() for host in os.environ.get("MIRROR_ALLOW_HOSTS", DEFAULT_MIRROR_HOSTS).split(",") if host.strip())
# Optional shared upload token. If UPLOAD_TOKEN is unset, uploads are open (temp use).
UPLOAD_TOKEN = os.environ.get("UPLOAD_TOKEN", "").strip()
# Optional public base override, e.g. https://looknicemm1-tmp-files.hf.space
PUBLIC_BASE = os.environ.get("PUBLIC_BASE", "").strip().rstrip("/")
EXT_BY_TYPE = {
"image/jpeg": "jpg",
"image/jpg": "jpg",
"image/png": "png",
"image/webp": "webp",
"image/gif": "gif",
"video/mp4": "mp4",
"video/quicktime": "mov",
"video/webm": "webm",
"audio/mpeg": "mp3",
"audio/mp4": "m4a",
"audio/wav": "wav",
"audio/x-wav": "wav",
"audio/aac": "aac",
"audio/ogg": "ogg",
}
app = FastAPI(title="tmp-file-service")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
def public_base(request: Request) -> str:
if PUBLIC_BASE:
return PUBLIC_BASE
host = request.headers.get("x-forwarded-host") or request.headers.get("host") or "localhost"
proto = request.headers.get("x-forwarded-proto", "https")
return f"{proto}://{host}"
def check_auth(request: Request) -> None:
if not UPLOAD_TOKEN:
return
header = request.headers.get("authorization", "")
token = header[7:].strip() if header.lower().startswith("bearer ") else request.query_params.get("token", "")
if token != UPLOAD_TOKEN:
raise HTTPException(401, "invalid upload token")
def pick_ext(file: UploadFile) -> str:
ext = EXT_BY_TYPE.get((file.content_type or "").lower())
if ext:
return ext
suffix = pathlib.Path(file.filename or "").suffix.lstrip(".").lower()
return suffix or "bin"
def check_mirror_url(raw: str) -> str:
url = (raw or "").strip()
parsed = urlparse(url)
if parsed.scheme not in ("http", "https"):
raise HTTPException(400, "url must be http(s)")
host = (parsed.hostname or "").lower()
if not host:
raise HTTPException(400, "url has no host")
if not any(host == allowed or host.endswith(f".{allowed}") for allowed in MIRROR_ALLOW_HOSTS):
raise HTTPException(403, f"host not allowed: {host}")
return url
@app.get("/")
def health():
return {"ok": True, "service": "tmp-file-service", "auth": bool(UPLOAD_TOKEN), "mirror_hosts": list(MIRROR_ALLOW_HOSTS)}
@app.get("/mirror")
async def mirror(request: Request, url: str = ""):
"""Fetch a signed object-storage URL server-side and stream it back, so browsers can read
bytes that the origin serves without CORS headers. Nothing is written to disk."""
target = check_mirror_url(url)
client = httpx.AsyncClient(follow_redirects=True, timeout=httpx.Timeout(connect=15.0, read=120.0, write=60.0, pool=15.0))
try:
upstream = await client.send(client.build_request("GET", target), stream=True)
except httpx.HTTPError as error:
await client.aclose()
raise HTTPException(502, f"fetch failed: {error}") from error
if upstream.status_code >= 400:
await upstream.aclose()
await client.aclose()
raise HTTPException(502, f"upstream returned {upstream.status_code}")
declared = upstream.headers.get("content-length")
if declared and declared.isdigit() and int(declared) > MIRROR_MAX_BYTES:
await upstream.aclose()
await client.aclose()
raise HTTPException(413, f"object too large (> {MIRROR_MAX_BYTES} bytes)")
async def body():
sent = 0
try:
async for chunk in upstream.aiter_bytes(chunk_size=256 * 1024):
sent += len(chunk)
if sent > MIRROR_MAX_BYTES:
break
yield chunk
finally:
await upstream.aclose()
await client.aclose()
headers = {"Cache-Control": "no-store"}
if declared:
headers["Content-Length"] = declared
return StreamingResponse(body(), media_type=upstream.headers.get("content-type", "application/octet-stream"), headers=headers)
@app.post("/upload")
async def upload(request: Request, file: UploadFile = File(...)):
check_auth(request)
data = await file.read()
if not data:
raise HTTPException(400, "empty file")
if len(data) > MAX_BYTES:
raise HTTPException(413, f"file too large (> {MAX_BYTES} bytes)")
name = f"{uuid.uuid4().hex}.{pick_ext(file)}"
(DATA_DIR / name).write_bytes(data)
url = f"{public_base(request)}/files/{name}"
return {"success": True, "url": url, "filename": name, "size": len(data)}
@app.get("/files/{name}")
def get_file(name: str):
if "/" in name or ".." in name:
raise HTTPException(400, "bad name")
path = DATA_DIR / name
if not path.is_file():
raise HTTPException(404, "not found")
return FileResponse(path)