stasaking commited on
Commit
b740423
·
verified ·
1 Parent(s): c253486

Sync main.py from blood-brain-omics repo

Browse files
Files changed (1) hide show
  1. main.py +132 -17
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 = db.execute(f"""
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 = db.execute("""
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 = db.execute(f"""
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 = db.execute(f"""
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 = db.execute(f"""
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 = db.execute(f"""
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 = db.execute(f"""
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 = db.execute(f"""
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 = db.execute("""
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 = db.execute("""
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 = db.execute(f"""
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 = db.execute(f"""
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 = db.execute(f"""
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 = db.execute("SELECT COUNT(*) FROM results").fetchone()[0]
742
- n_features = db.execute("SELECT COUNT(*) FROM features").fetchone()[0]
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
- db.execute("SELECT 1").fetchone()
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(