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