rosdiff / api /storage.py
Chandra Kiran
Refuse new jobs while the R2 bucket is near its free tier
115459a unverified
Raw History Blame Contribute Delete
7.12 kB
"""Cloudflare R2 (S3-compatible) access: uploads from the worker, presigned links for the report page."""
from __future__ import annotations
import os
from dataclasses import dataclass
from pathlib import Path
from typing import Protocol
CONTENT_TYPES = {
".mcap": "application/octet-stream",
".mp4": "video/mp4",
".json": "application/json",
".xml": "application/xml",
".webm": "video/webm",
".npz": "application/octet-stream",
".log": "text/plain; charset=utf-8",
}
class Storage(Protocol):
def upload(self, local: Path, key: str) -> None: ...
def url(self, key: str, expires_s: int = 7 * 24 * 3600) -> str: ...
def read_text(self, key: str) -> str | None: ...
def total_bytes(self) -> int: ...
@dataclass(frozen=True)
class R2Config:
account_id: str
access_key_id: str
secret_access_key: str
bucket: str
@classmethod
def from_env(cls) -> R2Config | None:
values = [
os.environ.get(k, "") for k in ("R2_ACCOUNT_ID", "R2_ACCESS_KEY_ID", "R2_SECRET_ACCESS_KEY", "R2_BUCKET")
]
return cls(*values) if all(values) else None
class R2Storage:
def __init__(self, config: R2Config):
import boto3
from botocore.config import Config
self.bucket = config.bucket
self._s3 = boto3.client(
"s3",
endpoint_url=f"https://{config.account_id}.r2.cloudflarestorage.com",
aws_access_key_id=config.access_key_id,
aws_secret_access_key=config.secret_access_key,
region_name="auto",
config=Config(signature_version="s3v4", retries={"max_attempts": 5}),
)
def upload(self, local: Path, key: str) -> None:
extra = {"ContentType": CONTENT_TYPES.get(Path(key).suffix, "application/octet-stream")}
self._s3.upload_file(str(local), self.bucket, key, ExtraArgs=extra)
def url(self, key: str, expires_s: int = 7 * 24 * 3600) -> str:
# S3 SigV4 presigned URLs are valid for at most 7 days.
return self._s3.generate_presigned_url(
"get_object", Params={"Bucket": self.bucket, "Key": key}, ExpiresIn=min(expires_s, 7 * 24 * 3600)
)
def read_text(self, key: str) -> str | None:
try:
return self._s3.get_object(Bucket=self.bucket, Key=key)["Body"].read().decode()
except self._s3.exceptions.NoSuchKey:
return None
def total_bytes(self) -> int:
"""Everything stored in the bucket (one listing request per 1000 files)."""
total = 0
for page in self._s3.get_paginator("list_objects_v2").paginate(Bucket=self.bucket):
total += sum(obj["Size"] for obj in page.get("Contents", []))
return total
def set_expiry(self, days: int, prefixes: tuple[str, ...] = ("runs/", "jobs/")) -> None:
"""Delete run and job files automatically after `days` (keeps the bucket inside R2's free 10 GB)."""
self._s3.put_bucket_lifecycle_configuration(
Bucket=self.bucket,
LifecycleConfiguration={
"Rules": [
{
"ID": f"expire-{p.strip('/')}",
"Status": "Enabled",
"Filter": {"Prefix": p},
"Expiration": {"Days": days},
}
for p in prefixes
]
},
)
def set_cors(self, origins: list[str]) -> None:
"""Let the Foxglove web app (and the report page) fetch run files directly from the bucket."""
self._s3.put_bucket_cors(
Bucket=self.bucket,
CORSConfiguration={
"CORSRules": [
{
"AllowedOrigins": origins,
"AllowedMethods": ["GET", "HEAD"],
"AllowedHeaders": ["*"],
"ExposeHeaders": ["Content-Length", "Content-Range", "Accept-Ranges", "ETag"],
"MaxAgeSeconds": 3600,
}
]
},
)
class LocalStorage:
"""Directory-backed stand-in for R2 (tests and offline dry runs)."""
def __init__(self, root: Path, base_url: str = "file://"):
self.root = Path(root)
self.base_url = base_url
def upload(self, local: Path, key: str) -> None:
dest = self.root / key
dest.parent.mkdir(parents=True, exist_ok=True)
dest.write_bytes(Path(local).read_bytes())
def url(self, key: str, expires_s: int = 0) -> str:
return (
f"{self.base_url}{(self.root / key).resolve()}" if self.base_url == "file://" else f"{self.base_url}/{key}"
)
def read_text(self, key: str) -> str | None:
p = self.root / key
return p.read_text() if p.is_file() else None
def total_bytes(self) -> int:
return sum(p.stat().st_size for p in self.root.rglob("*") if p.is_file()) if self.root.exists() else 0
def run_prefix(run_id: str) -> str:
return f"runs/{run_id}"
def main(argv: list[str] | None = None) -> int:
"""python -m api.storage cors | check"""
import argparse
import sys
import tempfile
parser = argparse.ArgumentParser(description="Cloudflare R2 setup for simulation runs")
sub = parser.add_subparsers(dest="cmd", required=True)
p = sub.add_parser("cors", help="allow the Foxglove web app (and extra origins) to read run files")
p.add_argument("--origin", action="append", default=[], help="extra allowed origin, e.g. your API's URL")
sub.add_parser("check", help="upload a test object and print a presigned link")
sub.add_parser("usage", help="how much the bucket holds")
p = sub.add_parser("expire", help="delete run/job files automatically after N days")
p.add_argument("--days", type=int, default=30)
args = parser.parse_args(argv)
config = R2Config.from_env()
if not config:
print("error: set R2_ACCOUNT_ID, R2_ACCESS_KEY_ID, R2_SECRET_ACCESS_KEY and R2_BUCKET", file=sys.stderr)
return 1
r2 = R2Storage(config)
if args.cmd == "cors":
# embed.foxglove.dev: the viewer embedded in the report page; app.foxglove.dev: the full web app.
origins = ["https://embed.foxglove.dev", "https://app.foxglove.dev", *args.origin]
r2.set_cors(origins)
print(f"CORS on bucket {config.bucket}: GET/HEAD from {', '.join(origins)}")
return 0
if args.cmd == "usage":
print(f"bucket {config.bucket}: {r2.total_bytes() / 1e9:.2f} GB")
return 0
if args.cmd == "expire":
r2.set_expiry(args.days)
print(f"bucket {config.bucket}: files under runs/ and jobs/ are deleted after {args.days} days")
return 0
with tempfile.NamedTemporaryFile("w", suffix=".json", delete=False) as f:
f.write('{"ok": true}\n')
r2.upload(Path(f.name), "runs/_check/check.json")
print(f"uploaded runs/_check/check.json; presigned link (7 days):\n{r2.url('runs/_check/check.json')}")
return 0
if __name__ == "__main__":
raise SystemExit(main())