"""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( """ Starting FiftyOne

Starting your FiftyOne session…

""" ) 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