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