Spaces:
Sleeping
Sleeping
Sync main.py from blood-brain-omics repo
Browse files
main.py
CHANGED
|
@@ -27,6 +27,7 @@ Endpoints:
|
|
| 27 |
GET /api/v1/ready - Readiness probe (SELECT 1)
|
| 28 |
"""
|
| 29 |
|
|
|
|
| 30 |
import json
|
| 31 |
import logging
|
| 32 |
import os
|
|
@@ -38,7 +39,7 @@ from typing import Optional
|
|
| 38 |
import duckdb
|
| 39 |
from fastapi import FastAPI, HTTPException, Query, Request
|
| 40 |
from fastapi.middleware.cors import CORSMiddleware
|
| 41 |
-
from fastapi.responses import JSONResponse
|
| 42 |
|
| 43 |
logging.basicConfig(
|
| 44 |
level=logging.INFO,
|
|
@@ -68,6 +69,14 @@ DUCKDB_MEMORY_LIMIT = os.environ.get("DUCKDB_MEMORY_LIMIT", "512MB")
|
|
| 68 |
DUCKDB_THREADS = int(os.environ.get("DUCKDB_THREADS", "1"))
|
| 69 |
DUCKDB_TEMP_DIR = os.environ.get("DUCKDB_TEMP_DIR", "/tmp/duckdb_tmp")
|
| 70 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 71 |
app = FastAPI(
|
| 72 |
title="Blood-Brain Omics Benchmark API",
|
| 73 |
version="1.0.0",
|
|
@@ -82,6 +91,99 @@ app.add_middleware(
|
|
| 82 |
)
|
| 83 |
|
| 84 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 85 |
# =============================================================================
|
| 86 |
# Middleware: request logging + unhandled-exception handler
|
| 87 |
# =============================================================================
|
|
@@ -212,6 +314,19 @@ def shutdown():
|
|
| 212 |
VALID_METRICS = {"r2", "pearson", "mse", "mae"}
|
| 213 |
|
| 214 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 215 |
def validate_metric(metric: str) -> str:
|
| 216 |
if metric not in VALID_METRICS:
|
| 217 |
raise HTTPException(400, f"Invalid metric: {metric}. "
|
|
@@ -253,7 +368,7 @@ def maxn_heatmap(
|
|
| 253 |
validate_metric(metric)
|
| 254 |
phase = covariate_phase("maxn", covariates)
|
| 255 |
|
| 256 |
-
rows =
|
| 257 |
SELECT blood_platform, brain_target,
|
| 258 |
MEDIAN({metric}) as value,
|
| 259 |
MEDIAN(n_samples) as n_samples,
|
|
@@ -305,7 +420,7 @@ def maxn_detail(
|
|
| 305 |
"""Per-module results for one blood x brain combination."""
|
| 306 |
phase = covariate_phase("maxn", covariates)
|
| 307 |
|
| 308 |
-
rows =
|
| 309 |
SELECT target, n_samples, r2, pearson, mse, mae
|
| 310 |
FROM results
|
| 311 |
WHERE phase = ? AND blood_platform = ? AND brain_target = ?
|
|
@@ -338,7 +453,7 @@ def h2h_blood(
|
|
| 338 |
validate_metric(metric)
|
| 339 |
phase = "h2h_blood_withcov" if covariates == "included" else "h2h_blood"
|
| 340 |
|
| 341 |
-
rows =
|
| 342 |
SELECT target, blood_platform, {metric}
|
| 343 |
FROM results
|
| 344 |
WHERE phase = ? AND h2h_pair = ? AND brain_target = ? AND model = ?
|
|
@@ -375,7 +490,7 @@ def h2h_blood_summary(
|
|
| 375 |
validate_metric(metric)
|
| 376 |
phase = "h2h_blood_withcov" if covariates == "included" else "h2h_blood"
|
| 377 |
|
| 378 |
-
rows =
|
| 379 |
SELECT h2h_pair, brain_target, target, blood_platform, {metric}
|
| 380 |
FROM results
|
| 381 |
WHERE phase = ? AND model = ? AND h2h_pair IS NOT NULL
|
|
@@ -434,7 +549,7 @@ def h2h_brain(
|
|
| 434 |
validate_metric(metric)
|
| 435 |
phase = "h2h_brain_withcov" if covariates == "included" else "h2h_brain"
|
| 436 |
|
| 437 |
-
rows =
|
| 438 |
SELECT h2h_pair, brain_target, target, {metric}
|
| 439 |
FROM results
|
| 440 |
WHERE phase = ? AND blood_platform = ? AND model = ?
|
|
@@ -468,7 +583,7 @@ def h2h_models(
|
|
| 468 |
validate_metric(metric)
|
| 469 |
phase = covariate_phase("maxn", covariates)
|
| 470 |
|
| 471 |
-
rows =
|
| 472 |
SELECT target, model, {metric}
|
| 473 |
FROM results
|
| 474 |
WHERE phase = ? AND blood_platform = ? AND brain_target = ?
|
|
@@ -506,14 +621,14 @@ def temporal(
|
|
| 506 |
phase = "temporal_withcov" if covariates == "included" else "temporal"
|
| 507 |
|
| 508 |
if target:
|
| 509 |
-
rows =
|
| 510 |
SELECT temporal_bin, brain_target, target, {metric}, n_samples
|
| 511 |
FROM results
|
| 512 |
WHERE phase = ? AND model = ? AND target = ?
|
| 513 |
ORDER BY temporal_bin
|
| 514 |
""", [phase, model, target]).fetchall()
|
| 515 |
else:
|
| 516 |
-
rows =
|
| 517 |
SELECT temporal_bin, brain_target, target, {metric}, n_samples
|
| 518 |
FROM results
|
| 519 |
WHERE phase = ? AND model = ?
|
|
@@ -544,7 +659,7 @@ def feature_importance(
|
|
| 544 |
limit: int = Query(30, ge=1, le=100),
|
| 545 |
):
|
| 546 |
"""Top features for a specific module/target."""
|
| 547 |
-
rows =
|
| 548 |
SELECT feature_name, importance, rank
|
| 549 |
FROM features
|
| 550 |
WHERE phase = ? AND blood_platform = ? AND brain_target = ?
|
|
@@ -574,7 +689,7 @@ def feature_cross_target(
|
|
| 574 |
limit: int = Query(30, ge=1, le=100),
|
| 575 |
):
|
| 576 |
"""Aggregated feature importance across all modules in a brain target."""
|
| 577 |
-
rows =
|
| 578 |
SELECT feature_name,
|
| 579 |
AVG(importance) as mean_importance,
|
| 580 |
COUNT(DISTINCT target) as n_targets,
|
|
@@ -613,7 +728,7 @@ def maxn_panel_summary(
|
|
| 613 |
validate_metric(metric)
|
| 614 |
phase = covariate_phase("maxn", covariates)
|
| 615 |
|
| 616 |
-
rows =
|
| 617 |
SELECT brain_target,
|
| 618 |
AVG({metric}) as mean_val,
|
| 619 |
MIN({metric}) as min_val,
|
|
@@ -659,7 +774,7 @@ def maxn_outcome_summary(
|
|
| 659 |
validate_metric(metric)
|
| 660 |
phase = covariate_phase("maxn", covariates)
|
| 661 |
|
| 662 |
-
rows =
|
| 663 |
SELECT blood_platform,
|
| 664 |
AVG({metric}) as mean_val,
|
| 665 |
MIN({metric}) as min_val,
|
|
@@ -706,7 +821,7 @@ def maxn_target_comparison(
|
|
| 706 |
validate_metric(metric)
|
| 707 |
phase = covariate_phase("maxn", covariates)
|
| 708 |
|
| 709 |
-
rows =
|
| 710 |
SELECT blood_platform, {metric}, n_samples
|
| 711 |
FROM results
|
| 712 |
WHERE phase = ? AND brain_target = ? AND target = ? AND model = ?
|
|
@@ -738,8 +853,8 @@ def health():
|
|
| 738 |
content={"status": "degraded", "detail": "Database not initialized"},
|
| 739 |
)
|
| 740 |
try:
|
| 741 |
-
n_results =
|
| 742 |
-
n_features =
|
| 743 |
except Exception as exc:
|
| 744 |
logger.error("Health check query failed: %s", exc)
|
| 745 |
return JSONResponse(
|
|
@@ -758,7 +873,7 @@ def ready():
|
|
| 758 |
content={"ready": False, "detail": "Database not initialized"},
|
| 759 |
)
|
| 760 |
try:
|
| 761 |
-
|
| 762 |
except Exception as exc:
|
| 763 |
logger.error("Ready check query failed: %s", exc)
|
| 764 |
return JSONResponse(
|
|
|
|
| 27 |
GET /api/v1/ready - Readiness probe (SELECT 1)
|
| 28 |
"""
|
| 29 |
|
| 30 |
+
import asyncio
|
| 31 |
import json
|
| 32 |
import logging
|
| 33 |
import os
|
|
|
|
| 39 |
import duckdb
|
| 40 |
from fastapi import FastAPI, HTTPException, Query, Request
|
| 41 |
from fastapi.middleware.cors import CORSMiddleware
|
| 42 |
+
from fastapi.responses import JSONResponse, Response
|
| 43 |
|
| 44 |
logging.basicConfig(
|
| 45 |
level=logging.INFO,
|
|
|
|
| 69 |
DUCKDB_THREADS = int(os.environ.get("DUCKDB_THREADS", "1"))
|
| 70 |
DUCKDB_TEMP_DIR = os.environ.get("DUCKDB_TEMP_DIR", "/tmp/duckdb_tmp")
|
| 71 |
|
| 72 |
+
# Concurrency cap — bound in-flight queries to avoid OOM/process crash on
|
| 73 |
+
# CPU Basic. Excess requests get 429 and the frontend can retry.
|
| 74 |
+
MAX_INFLIGHT = int(os.environ.get("API_MAX_INFLIGHT", "8"))
|
| 75 |
+
|
| 76 |
+
# Response cache TTL (seconds) — absorbs repeat queries from rapid clicks.
|
| 77 |
+
CACHE_TTL = float(os.environ.get("API_CACHE_TTL", "60"))
|
| 78 |
+
CACHE_MAX_ENTRIES = int(os.environ.get("API_CACHE_MAX_ENTRIES", "256"))
|
| 79 |
+
|
| 80 |
app = FastAPI(
|
| 81 |
title="Blood-Brain Omics Benchmark API",
|
| 82 |
version="1.0.0",
|
|
|
|
| 91 |
)
|
| 92 |
|
| 93 |
|
| 94 |
+
# =============================================================================
|
| 95 |
+
# Concurrency cap + response cache
|
| 96 |
+
# =============================================================================
|
| 97 |
+
|
| 98 |
+
# Bound the number of in-flight requests so that bursts cannot OOM/crash the
|
| 99 |
+
# DuckDB process on CPU Basic. Excess requests fast-fail with 429 instead of
|
| 100 |
+
# wedging the worker.
|
| 101 |
+
_inflight_sem = asyncio.Semaphore(MAX_INFLIGHT)
|
| 102 |
+
|
| 103 |
+
# Tiny TTL cache keyed on (path, query_string). Frontend rapid clicks
|
| 104 |
+
# (e.g. switching panels) repeatedly hit the same URLs — caching them
|
| 105 |
+
# absorbs the load before it ever reaches DuckDB.
|
| 106 |
+
_cache: dict[str, tuple[float, int, bytes, str]] = {}
|
| 107 |
+
_CACHEABLE_PATH_PREFIXES = (
|
| 108 |
+
"/api/v1/registry",
|
| 109 |
+
"/api/v1/maxn/",
|
| 110 |
+
"/api/v1/h2h/",
|
| 111 |
+
"/api/v1/temporal",
|
| 112 |
+
"/api/v1/features",
|
| 113 |
+
)
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
def _cache_key(request: Request) -> Optional[str]:
|
| 117 |
+
if request.method != "GET":
|
| 118 |
+
return None
|
| 119 |
+
path = request.url.path
|
| 120 |
+
if not any(path.startswith(p) for p in _CACHEABLE_PATH_PREFIXES):
|
| 121 |
+
return None
|
| 122 |
+
return f"{path}?{request.url.query}"
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
def _cache_get(key: str) -> Optional[tuple[int, bytes, str]]:
|
| 126 |
+
entry = _cache.get(key)
|
| 127 |
+
if entry is None:
|
| 128 |
+
return None
|
| 129 |
+
expires, status, body, ctype = entry
|
| 130 |
+
if expires < time.monotonic():
|
| 131 |
+
_cache.pop(key, None)
|
| 132 |
+
return None
|
| 133 |
+
return status, body, ctype
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
def _cache_put(key: str, status: int, body: bytes, ctype: str) -> None:
|
| 137 |
+
if status != 200:
|
| 138 |
+
return
|
| 139 |
+
if len(_cache) >= CACHE_MAX_ENTRIES:
|
| 140 |
+
# Cheap eviction: drop the oldest-expiring entry.
|
| 141 |
+
oldest = min(_cache.items(), key=lambda kv: kv[1][0])[0]
|
| 142 |
+
_cache.pop(oldest, None)
|
| 143 |
+
_cache[key] = (time.monotonic() + CACHE_TTL, status, body, ctype)
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
@app.middleware("http")
|
| 147 |
+
async def cache_and_throttle(request: Request, call_next):
|
| 148 |
+
key = _cache_key(request)
|
| 149 |
+
if key is not None:
|
| 150 |
+
hit = _cache_get(key)
|
| 151 |
+
if hit is not None:
|
| 152 |
+
status, body, ctype = hit
|
| 153 |
+
return Response(
|
| 154 |
+
content=body, status_code=status, media_type=ctype,
|
| 155 |
+
headers={"x-cache": "hit"},
|
| 156 |
+
)
|
| 157 |
+
|
| 158 |
+
# Bound concurrency. If saturated, fail fast with 429 rather than queue
|
| 159 |
+
# indefinitely (queueing under sustained load is what crashes the worker).
|
| 160 |
+
try:
|
| 161 |
+
await asyncio.wait_for(_inflight_sem.acquire(), timeout=0.05)
|
| 162 |
+
except asyncio.TimeoutError:
|
| 163 |
+
return JSONResponse(
|
| 164 |
+
status_code=429,
|
| 165 |
+
content={"error": "busy", "detail": "Server busy, please retry."},
|
| 166 |
+
headers={"retry-after": "1"},
|
| 167 |
+
)
|
| 168 |
+
|
| 169 |
+
try:
|
| 170 |
+
response = await call_next(request)
|
| 171 |
+
finally:
|
| 172 |
+
_inflight_sem.release()
|
| 173 |
+
|
| 174 |
+
if key is not None and response.status_code == 200:
|
| 175 |
+
# Buffer the streaming body so we can both cache and return it.
|
| 176 |
+
body_chunks = [chunk async for chunk in response.body_iterator]
|
| 177 |
+
body = b"".join(body_chunks)
|
| 178 |
+
ctype = response.headers.get("content-type", "application/json")
|
| 179 |
+
_cache_put(key, response.status_code, body, ctype)
|
| 180 |
+
return Response(
|
| 181 |
+
content=body, status_code=response.status_code, media_type=ctype,
|
| 182 |
+
headers={"x-cache": "miss"},
|
| 183 |
+
)
|
| 184 |
+
return response
|
| 185 |
+
|
| 186 |
+
|
| 187 |
# =============================================================================
|
| 188 |
# Middleware: request logging + unhandled-exception handler
|
| 189 |
# =============================================================================
|
|
|
|
| 314 |
VALID_METRICS = {"r2", "pearson", "mse", "mae"}
|
| 315 |
|
| 316 |
|
| 317 |
+
def q(sql: str, params=None):
|
| 318 |
+
"""Run a query on an isolated DuckDB cursor.
|
| 319 |
+
|
| 320 |
+
DuckDB Connection objects are NOT safe for concurrent use across
|
| 321 |
+
requests — sharing the global ``db`` connection between overlapping
|
| 322 |
+
FastAPI requests causes corrupted result state and hung responses.
|
| 323 |
+
Cursors created via ``db.cursor()`` are cheap, isolated, and safe
|
| 324 |
+
for concurrent queries against the same in-memory database.
|
| 325 |
+
"""
|
| 326 |
+
cur = db.cursor()
|
| 327 |
+
return cur.execute(sql, params) if params is not None else cur.execute(sql)
|
| 328 |
+
|
| 329 |
+
|
| 330 |
def validate_metric(metric: str) -> str:
|
| 331 |
if metric not in VALID_METRICS:
|
| 332 |
raise HTTPException(400, f"Invalid metric: {metric}. "
|
|
|
|
| 368 |
validate_metric(metric)
|
| 369 |
phase = covariate_phase("maxn", covariates)
|
| 370 |
|
| 371 |
+
rows = q(f"""
|
| 372 |
SELECT blood_platform, brain_target,
|
| 373 |
MEDIAN({metric}) as value,
|
| 374 |
MEDIAN(n_samples) as n_samples,
|
|
|
|
| 420 |
"""Per-module results for one blood x brain combination."""
|
| 421 |
phase = covariate_phase("maxn", covariates)
|
| 422 |
|
| 423 |
+
rows = q("""
|
| 424 |
SELECT target, n_samples, r2, pearson, mse, mae
|
| 425 |
FROM results
|
| 426 |
WHERE phase = ? AND blood_platform = ? AND brain_target = ?
|
|
|
|
| 453 |
validate_metric(metric)
|
| 454 |
phase = "h2h_blood_withcov" if covariates == "included" else "h2h_blood"
|
| 455 |
|
| 456 |
+
rows = q(f"""
|
| 457 |
SELECT target, blood_platform, {metric}
|
| 458 |
FROM results
|
| 459 |
WHERE phase = ? AND h2h_pair = ? AND brain_target = ? AND model = ?
|
|
|
|
| 490 |
validate_metric(metric)
|
| 491 |
phase = "h2h_blood_withcov" if covariates == "included" else "h2h_blood"
|
| 492 |
|
| 493 |
+
rows = q(f"""
|
| 494 |
SELECT h2h_pair, brain_target, target, blood_platform, {metric}
|
| 495 |
FROM results
|
| 496 |
WHERE phase = ? AND model = ? AND h2h_pair IS NOT NULL
|
|
|
|
| 549 |
validate_metric(metric)
|
| 550 |
phase = "h2h_brain_withcov" if covariates == "included" else "h2h_brain"
|
| 551 |
|
| 552 |
+
rows = q(f"""
|
| 553 |
SELECT h2h_pair, brain_target, target, {metric}
|
| 554 |
FROM results
|
| 555 |
WHERE phase = ? AND blood_platform = ? AND model = ?
|
|
|
|
| 583 |
validate_metric(metric)
|
| 584 |
phase = covariate_phase("maxn", covariates)
|
| 585 |
|
| 586 |
+
rows = q(f"""
|
| 587 |
SELECT target, model, {metric}
|
| 588 |
FROM results
|
| 589 |
WHERE phase = ? AND blood_platform = ? AND brain_target = ?
|
|
|
|
| 621 |
phase = "temporal_withcov" if covariates == "included" else "temporal"
|
| 622 |
|
| 623 |
if target:
|
| 624 |
+
rows = q(f"""
|
| 625 |
SELECT temporal_bin, brain_target, target, {metric}, n_samples
|
| 626 |
FROM results
|
| 627 |
WHERE phase = ? AND model = ? AND target = ?
|
| 628 |
ORDER BY temporal_bin
|
| 629 |
""", [phase, model, target]).fetchall()
|
| 630 |
else:
|
| 631 |
+
rows = q(f"""
|
| 632 |
SELECT temporal_bin, brain_target, target, {metric}, n_samples
|
| 633 |
FROM results
|
| 634 |
WHERE phase = ? AND model = ?
|
|
|
|
| 659 |
limit: int = Query(30, ge=1, le=100),
|
| 660 |
):
|
| 661 |
"""Top features for a specific module/target."""
|
| 662 |
+
rows = q("""
|
| 663 |
SELECT feature_name, importance, rank
|
| 664 |
FROM features
|
| 665 |
WHERE phase = ? AND blood_platform = ? AND brain_target = ?
|
|
|
|
| 689 |
limit: int = Query(30, ge=1, le=100),
|
| 690 |
):
|
| 691 |
"""Aggregated feature importance across all modules in a brain target."""
|
| 692 |
+
rows = q("""
|
| 693 |
SELECT feature_name,
|
| 694 |
AVG(importance) as mean_importance,
|
| 695 |
COUNT(DISTINCT target) as n_targets,
|
|
|
|
| 728 |
validate_metric(metric)
|
| 729 |
phase = covariate_phase("maxn", covariates)
|
| 730 |
|
| 731 |
+
rows = q(f"""
|
| 732 |
SELECT brain_target,
|
| 733 |
AVG({metric}) as mean_val,
|
| 734 |
MIN({metric}) as min_val,
|
|
|
|
| 774 |
validate_metric(metric)
|
| 775 |
phase = covariate_phase("maxn", covariates)
|
| 776 |
|
| 777 |
+
rows = q(f"""
|
| 778 |
SELECT blood_platform,
|
| 779 |
AVG({metric}) as mean_val,
|
| 780 |
MIN({metric}) as min_val,
|
|
|
|
| 821 |
validate_metric(metric)
|
| 822 |
phase = covariate_phase("maxn", covariates)
|
| 823 |
|
| 824 |
+
rows = q(f"""
|
| 825 |
SELECT blood_platform, {metric}, n_samples
|
| 826 |
FROM results
|
| 827 |
WHERE phase = ? AND brain_target = ? AND target = ? AND model = ?
|
|
|
|
| 853 |
content={"status": "degraded", "detail": "Database not initialized"},
|
| 854 |
)
|
| 855 |
try:
|
| 856 |
+
n_results = q("SELECT COUNT(*) FROM results").fetchone()[0]
|
| 857 |
+
n_features = q("SELECT COUNT(*) FROM features").fetchone()[0]
|
| 858 |
except Exception as exc:
|
| 859 |
logger.error("Health check query failed: %s", exc)
|
| 860 |
return JSONResponse(
|
|
|
|
| 873 |
content={"ready": False, "detail": "Database not initialized"},
|
| 874 |
)
|
| 875 |
try:
|
| 876 |
+
q("SELECT 1").fetchone()
|
| 877 |
except Exception as exc:
|
| 878 |
logger.error("Ready check query failed: %s", exc)
|
| 879 |
return JSONResponse(
|