harpreetsahota's picture
Upload folder using huggingface_hub
087ddd7 verified
Raw
History Blame Contribute Delete
15 kB
"""Provides browser-local FiftyOne dataset clones behind a small gateway."""
import asyncio
from contextlib import asynccontextmanager, suppress
from dataclasses import dataclass
import logging
from pathlib import Path
import re
import secrets
import time
from urllib.parse import urlparse
import httpx
from starlette.applications import Starlette
from starlette.background import BackgroundTask
from starlette.responses import (
HTMLResponse,
JSONResponse,
RedirectResponse,
Response,
StreamingResponse,
)
from starlette.routing import Route
import fiftyone as fo
COOKIE_NAME = "fiftyone_demo_session"
SESSION_HEADER = "x-fiftyone-session"
INTERNAL_APP_URL = "http://127.0.0.1:5151"
MAX_ACTIVE_SESSIONS = 20
DATASET_ROUTE_PATTERN = re.compile(r"^/datasets/([^/]+)/?$")
SESSION_TTL_SECONDS = 30 * 60
CLEANUP_INTERVAL_SECONDS = 5 * 60
logger = logging.getLogger(__name__)
@dataclass
class SessionRecord:
"""Tracks one browser's cloned dataset."""
token: str
dataset_names: tuple[str, ...]
default_dataset_name: str
last_seen: float
class SessionManager:
"""Creates and expires browser-local dataset clones."""
def __init__(self, base_datasets, default_dataset_name):
self._base_datasets = base_datasets
self._default_dataset_name = default_dataset_name
self._lock = asyncio.Lock()
self._sessions = {}
self._datasets = {}
@property
def active_count(self):
"""Returns the number of active browser sessions."""
return len(self._sessions)
async def get_or_create(self, token=None):
"""Returns an existing session or creates a new dataset clone."""
async with self._lock:
now = time.monotonic()
record = self._sessions.get(token)
if record is not None:
record.last_seen = now
return record
await self._delete_expired_locked(now)
if len(self._sessions) >= MAX_ACTIVE_SESSIONS:
return None
token = secrets.token_hex(16)
dataset_names = tuple(
base_name + "-session-" + token
for base_name in self._base_datasets
)
logger.info(
"Creating browser dataset clones: datasets=%s active_before=%d",
dataset_names,
len(self._sessions),
)
started_at = time.monotonic()
created_names = []
try:
for base_name, base_dataset in self._base_datasets.items():
dataset_name = base_name + "-session-" + token
await asyncio.to_thread(
base_dataset.clone,
dataset_name,
True,
)
created_names.append(dataset_name)
except Exception:
logger.exception(
"Failed to create all browser dataset clones: datasets=%s",
dataset_names,
)
for dataset_name in created_names:
await asyncio.to_thread(
_delete_dataset_if_exists, dataset_name
)
raise
record = SessionRecord(
token=token,
dataset_names=dataset_names,
default_dataset_name=(
self._default_dataset_name + "-session-" + token
),
last_seen=time.monotonic(),
)
self._sessions[token] = record
for dataset_name in dataset_names:
self._datasets[dataset_name] = record
logger.info(
"Browser dataset clones ready in %.2fs: datasets=%s active=%d",
time.monotonic() - started_at,
dataset_names,
len(self._sessions),
)
return record
async def resolve(self, request):
"""Resolves and touches the session represented by a request."""
token = _get_session_token(request)
record = self._sessions.get(token)
if record is None:
dataset_name = request.query_params.get("dataset")
if dataset_name not in self._datasets:
dataset_name = _extract_dataset_name(request.url.path)
if dataset_name is None:
dataset_name = _extract_dataset_name(
request.headers.get("referer", "")
)
record = self._datasets.get(dataset_name)
if record is not None:
record.last_seen = time.monotonic()
return record
async def cleanup_loop(self):
"""Periodically deletes clones whose browsers stopped heartbeating."""
while True:
await asyncio.sleep(CLEANUP_INTERVAL_SECONDS)
try:
await self.delete_expired()
except Exception:
logger.exception("Browser clone cleanup failed")
async def delete_expired(self):
"""Deletes all expired browser dataset clones."""
async with self._lock:
await self._delete_expired_locked(time.monotonic())
async def close(self):
"""Deletes all active clones during a graceful shutdown."""
async with self._lock:
records = list(self._sessions.values())
for record in records:
await self._delete_locked(record, reason="shutdown")
async def _delete_expired_locked(self, now):
expired = [
record
for record in self._sessions.values()
if now - record.last_seen >= SESSION_TTL_SECONDS
]
for record in expired:
await self._delete_locked(record, reason="inactive")
async def _delete_locked(self, record, reason):
logger.info(
"Deleting browser dataset clones: datasets=%s reason=%s",
record.dataset_names,
reason,
)
failed = False
for dataset_name in record.dataset_names:
try:
await asyncio.to_thread(
_delete_dataset_if_exists, dataset_name
)
except Exception:
failed = True
logger.exception(
"Failed to delete browser dataset clone '%s'",
dataset_name,
)
if failed:
return
self._sessions.pop(record.token, None)
for dataset_name in record.dataset_names:
self._datasets.pop(dataset_name, None)
logger.info(
"Browser dataset clones deleted: datasets=%s active=%d",
record.dataset_names,
len(self._sessions),
)
def create_gateway(
base_datasets,
default_dataset_name,
shared_media_roots,
):
"""Creates the session gateway ASGI application."""
manager = SessionManager(base_datasets, default_dataset_name)
resolved_media_roots = tuple(
Path(path).resolve() for path in shared_media_roots
)
@asynccontextmanager
async def lifespan(app):
limits = httpx.Limits(
max_connections=100,
max_keepalive_connections=20,
)
app.state.client = httpx.AsyncClient(
timeout=None,
limits=limits,
follow_redirects=False,
)
app.state.manager = manager
app.state.shared_media_roots = resolved_media_roots
cleanup_task = asyncio.create_task(manager.cleanup_loop())
logger.info(
"Session gateway ready: max_sessions=%d ttl_minutes=%d "
"cleanup_minutes=%d",
MAX_ACTIVE_SESSIONS,
SESSION_TTL_SECONDS // 60,
CLEANUP_INTERVAL_SECONDS // 60,
)
try:
yield
finally:
cleanup_task.cancel()
with suppress(asyncio.CancelledError):
await cleanup_task
await app.state.client.aclose()
await manager.close()
routes = [
Route("/", _landing_page, methods=["GET"]),
Route("/__health", _health, methods=["GET"]),
Route("/__session/start", _start_session, methods=["POST"]),
Route(
"/__session/heartbeat",
_heartbeat,
methods=["POST"],
),
Route(
"/{path:path}",
_proxy,
methods=[
"GET",
"POST",
"PUT",
"PATCH",
"DELETE",
"OPTIONS",
"HEAD",
],
),
]
return Starlette(routes=routes, lifespan=lifespan)
async def _landing_page(request):
"""Serves a probe-safe page that creates sessions only in browsers."""
return HTMLResponse(
"""<!doctype html>
<html>
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width,initial-scale=1">
<title>Starting FiftyOne</title>
<style>
body { background:#111; color:#eee; font:16px system-ui; display:grid;
min-height:100vh; margin:0; place-items:center; }
main { text-align:center; max-width:32rem; padding:2rem; }
</style>
</head>
<body>
<main><h1>Starting your FiftyOne session…</h1><p id="status"></p></main>
<script>
const sessionKey = "fiftyone_demo_session";
const savedSession = localStorage.getItem(sessionKey);
const headers = savedSession
? {"X-FiftyOne-Session": savedSession}
: {};
fetch("/__session/start", {
method:"POST",
credentials:"same-origin",
headers,
})
.then(async response => {
const data = await response.json();
if (!response.ok) throw new Error(data.error || "Session unavailable");
localStorage.setItem(sessionKey, data.token);
location.replace(data.url);
})
.catch(error => {
document.getElementById("status").textContent = error.message;
});
</script>
</body>
</html>"""
)
async def _health(request):
"""Returns gateway health without allocating a browser session."""
return JSONResponse(
{
"status": "ok",
"active_sessions": request.app.state.manager.active_count,
"max_sessions": MAX_ACTIVE_SESSIONS,
}
)
async def _start_session(request):
"""Allocates a clone and returns its dataset URL."""
manager = request.app.state.manager
record = await manager.get_or_create(_get_session_token(request))
if record is None:
return JSONResponse(
{
"error": (
"This demo currently has the maximum number of active "
"sessions. Please try again later."
)
},
status_code=503,
)
response = JSONResponse(
{
"url": f"/datasets/{record.default_dataset_name}",
"token": record.token,
}
)
response.set_cookie(
COOKIE_NAME,
record.token,
max_age=SESSION_TTL_SECONDS,
httponly=True,
secure=True,
samesite="none",
)
return response
async def _heartbeat(request):
"""Keeps a browser clone alive while its App tab remains open."""
record = await request.app.state.manager.resolve(request)
if record is None:
return Response(status_code=401)
return Response(status_code=204)
async def _proxy(request):
"""Streams an authorized browser request to the internal FiftyOne App."""
is_public_request = _is_public_request(request)
if not is_public_request:
record = await request.app.state.manager.resolve(request)
if record is None:
return RedirectResponse("/", status_code=303)
requested_dataset = _extract_dataset_name(request.url.path)
if (
request.url.path == "/datasets"
or (
DATASET_ROUTE_PATTERN.fullmatch(request.url.path)
and requested_dataset not in record.dataset_names
)
):
return RedirectResponse(
f"/datasets/{record.default_dataset_name}",
status_code=303,
)
client = request.app.state.client
upstream_url = INTERNAL_APP_URL + request.url.path
if request.url.query:
upstream_url += "?" + request.url.query
headers = {
key: value
for key, value in request.headers.items()
if key.lower() not in {"host", "content-length", "connection"}
}
return await _forward_request(
request,
request.app.state.client,
upstream_url,
headers,
)
async def _forward_request(request, client, upstream_url, headers):
"""Streams a request and response through an HTTP client."""
body = await request.body()
upstream_request = client.build_request(
request.method,
upstream_url,
headers=headers,
content=body,
)
upstream = await client.send(upstream_request, stream=True)
response_headers = {
key: value
for key, value in upstream.headers.items()
if key.lower()
not in {"connection", "keep-alive", "transfer-encoding"}
}
return StreamingResponse(
upstream.aiter_raw(),
status_code=upstream.status_code,
headers=response_headers,
background=BackgroundTask(upstream.aclose),
)
def _get_session_token(request):
"""Gets the browser session token from a request."""
return request.headers.get(SESSION_HEADER) or request.cookies.get(
COOKIE_NAME
)
def _delete_dataset_if_exists(dataset_name):
"""Deletes a dataset if it still exists."""
if fo.dataset_exists(dataset_name):
fo.delete_dataset(dataset_name)
def _extract_dataset_name(value):
"""Extracts a dataset name from a dataset path or referrer."""
if not value:
return None
path = urlparse(value).path
match = DATASET_ROUTE_PATTERN.fullmatch(path)
return match.group(1) if match else None
def _is_public_request(request):
"""Checks whether a read-only request can omit browser identity."""
if request.method not in {"GET", "HEAD"}:
return False
path = request.url.path
if path.startswith(("/datasets/assets/", "/assets/")):
return True
if path == "/favicon.ico":
return True
if path != "/media":
return False
filepath = request.query_params.get("filepath")
if not filepath:
return False
try:
resolved = Path(filepath).resolve()
is_shared_media = any(
resolved.is_relative_to(root)
for root in request.app.state.shared_media_roots
)
except OSError:
is_shared_media = False
if not is_shared_media:
logger.warning("Rejected media path outside local roots: %s", filepath)
return False
return True