WinslowFan Claude Opus 4.8 commited on
Commit
e5bf473
·
1 Parent(s): 6d84fd1

Migrate to new MuSProt.db schema (aliased columns, new stats)

Browse files

Alias renamed edge columns to existing API names in SQL, switch node path to
binding_status/cath_id, use COLLATE NOCASE for the new idx_pdb_chain index, and
refresh the summary stat cards and schema docs for the new DB.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

backend/app/protein/api/routes/protein.py CHANGED
@@ -1,267 +1,267 @@
1
- """Protein API routes."""
2
- import logging
3
- import math
4
- from typing import Dict, List
5
-
6
- from fastapi import APIRouter, HTTPException, Query
7
-
8
- from app.protein.chain_resolver import (
9
- batch_resolve_chains_async,
10
- get_cache,
11
- resolve_chain_async,
12
- )
13
- from app.protein.data_loader import DataManager
14
- from app.protein import tsv_loader
15
- from app.protein.models import (
16
- ChainMetadata,
17
- DataRecord,
18
- DataResponse,
19
- FilterOptions,
20
- FilterParams,
21
- SummaryStats,
22
- )
23
-
24
- logger = logging.getLogger(__name__)
25
- router = APIRouter()
26
-
27
- data_manager: DataManager | None = None
28
-
29
-
30
- def set_data_manager(manager: DataManager | None) -> None:
31
- """Set the active data manager for route handlers."""
32
- global data_manager
33
- data_manager = manager
34
-
35
-
36
- def _require_data_manager() -> DataManager:
37
- """Return data manager or raise if not initialized."""
38
- if data_manager is None:
39
- raise HTTPException(
40
- status_code=503,
41
- detail="Protein data is not initialized yet"
42
- )
43
- return data_manager
44
-
45
-
46
- @router.get("/")
47
- async def root():
48
- """Root endpoint."""
49
- return {
50
- "message": "Protein Visualisation API",
51
- "version": "1.0.0",
52
- "endpoints": {
53
- "/filters": "Get available filter options",
54
- "/data": "Get filtered protein data",
55
- "/summary": "Get aggregate statistics"
56
- }
57
- }
58
-
59
-
60
- @router.get("/filters", response_model=FilterOptions)
61
- async def get_filters():
62
- """Get available filter options and ranges."""
63
- try:
64
- manager = _require_data_manager()
65
- chains = sorted(manager.get_chain_sample())
66
- fd = manager.get_filters_data()
67
-
68
- return FilterOptions(
69
- chain_ids=chains[:100],
70
- rmsd_range={"min": float(fd["rmsd_min"]), "max": float(fd["rmsd_max"])},
71
- tm_score_range={"min": float(fd["tm_min"]), "max": float(fd["tm_max"])},
72
- length_range={"min": int(fd["length_min"]), "max": int(fd["length_max"])},
73
- clusters=[
74
- "cluster_ultra_high", "cluster_very_high", "cluster_high",
75
- "cluster_medium", "cluster_low",
76
- ],
77
- total_records=fd["total_records"],
78
- )
79
- except HTTPException:
80
- raise
81
- except Exception as e:
82
- logger.error(f"Error getting filters: {e}")
83
- raise HTTPException(status_code=500, detail=str(e))
84
-
85
-
86
- @router.post("/data", response_model=DataResponse)
87
- async def get_data(filters: FilterParams):
88
- """Get filtered protein data. Searches only by pdb_id_A and auth_asym_id_A."""
89
- try:
90
- manager = _require_data_manager()
91
- limit = filters.limit or 1000
92
- fetch_limit = limit * 10 if filters.cluster_id else limit
93
-
94
- df = manager.query_edge(
95
- pdb_id_a=filters.pdb_id,
96
- auth_asym_id_a=filters.auth_asym_id,
97
- chain_ids=filters.chain_ids,
98
- rmsd_min=filters.rmsd_min,
99
- rmsd_max=filters.rmsd_max,
100
- tm_min=filters.tm_score_min,
101
- tm_max=filters.tm_score_max,
102
- limit=fetch_limit,
103
- )
104
-
105
- if filters.cluster_id:
106
- df = df[df["cluster_id"] == filters.cluster_id]
107
-
108
- filtered_total = len(df)
109
- df = df.head(limit)
110
-
111
- def _safe_float(val):
112
- try:
113
- v = float(val)
114
- return None if math.isnan(v) else round(v, 6)
115
- except (TypeError, ValueError):
116
- return None
117
-
118
- records = []
119
- for row in df.itertuples(index=False):
120
- raw_ph = getattr(row, "pH", None)
121
- raw_temp = getattr(row, "temp_K", None)
122
- records.append(DataRecord(
123
- pdb_id_a=row.pdb_id_A,
124
- auth_asym_id_a=row.auth_asym_id_A,
125
- pdb_id_b=row.pdb_id_B,
126
- auth_asym_id_b=row.auth_asym_id_B,
127
- tm_score=float(row.TM1) if row.TM1 is not None else 0.0,
128
- rmsd=float(row.RMSD) if row.RMSD is not None else 0.0,
129
- structure_sim=_safe_float(row.structure_sim),
130
- length_a=int(row.length_a) if row.length_a is not None else None,
131
- length_b=int(row.length_b) if row.length_b is not None else None,
132
- cluster_id=row.cluster_id,
133
- exptl_method=getattr(row, "experimental_method", None) or None,
134
- pH=_safe_float(raw_ph) if raw_ph else None,
135
- temp=_safe_float(raw_temp) if raw_temp else None,
136
- binding_status=getattr(row, "base_label", None) or None,
137
- chain_composition=getattr(row, "chain_composition", None) or None,
138
- delta_rosetta=_safe_float(getattr(row, "delta_Rosetta", None)),
139
- delta_foldx=_safe_float(getattr(row, "delta_FoldX", None)),
140
- delta_evoef2=_safe_float(getattr(row, "delta_EvoEF2", None)),
141
- delta_rw=_safe_float(getattr(row, "delta_RW", None)),
142
- delta_rw_plus=_safe_float(getattr(row, "delta_RW_plus", None)),
143
- state_id_b=getattr(row, "state_id_B", None) or None,
144
- state_fidelity=getattr(row, "state_fidelity", None) or None,
145
- avg_sim=getattr(row, "avg_sim", None) or None,
146
- observation_fidelity=getattr(row, "observation_fidelity", None) or None,
147
- ))
148
-
149
- return DataResponse(
150
- data=records,
151
- total=manager.get_total(),
152
- filtered=filtered_total,
153
- )
154
-
155
- except HTTPException:
156
- raise
157
- except Exception as e:
158
- logger.error(f"Error filtering data: {e}")
159
- raise HTTPException(status_code=500, detail=str(e))
160
-
161
-
162
- @router.get("/summary", response_model=SummaryStats)
163
- async def get_summary(
164
- chain_ids: List[str] = Query(default=None),
165
- rmsd_min: float = Query(default=None),
166
- rmsd_max: float = Query(default=None),
167
- tm_score_min: float = Query(default=None),
168
- tm_score_max: float = Query(default=None),
169
- ):
170
- """Get aggregate statistics for the dataset or filtered subset."""
171
- try:
172
- manager = _require_data_manager()
173
- stats = manager.get_summary_stats(
174
- chain_ids=chain_ids,
175
- rmsd_min=rmsd_min,
176
- rmsd_max=rmsd_max,
177
- tm_min=tm_score_min,
178
- tm_max=tm_score_max,
179
- )
180
- return SummaryStats(
181
- total_records=stats["total_records"],
182
- unique_chains=stats["unique_chains"],
183
- avg_rmsd=float(stats["avg_rmsd"]),
184
- avg_tm_score=float(stats["avg_tm_score"]),
185
- avg_sequence_length=float(stats["avg_sequence_length"]),
186
- rmsd_distribution={str(k): int(v) for k, v in stats["rmsd_distribution"].items()},
187
- tm_score_distribution={str(k): int(v) for k, v in stats["tm_score_distribution"].items()},
188
- length_distribution={str(k): int(v) for k, v in stats["length_distribution"].items()},
189
- )
190
- except HTTPException:
191
- raise
192
- except Exception as e:
193
- logger.error(f"Error computing summary: {e}")
194
- raise HTTPException(status_code=500, detail=str(e))
195
-
196
-
197
- @router.get("/chain/resolve", response_model=ChainMetadata)
198
- async def resolve_chain_endpoint(
199
- pdb_id: str = Query(..., description="4-character PDB ID"),
200
- auth_asym_id: str = Query(..., description="Chain identifier (e.g., 'A')"),
201
- use_cache: bool = Query(True, description="Use cached data if available")
202
- ):
203
- """Resolve chain metadata from RCSB web API."""
204
- metadata = await resolve_chain_async(pdb_id, auth_asym_id, use_cache=use_cache)
205
- return metadata.model_copy(update={
206
- "binding_status": tsv_loader.get_binding_status(pdb_id, auth_asym_id),
207
- "chain_composition": tsv_loader.get_chain_composition(pdb_id, auth_asym_id),
208
- "cath_id": tsv_loader.get_cath_id(pdb_id, auth_asym_id),
209
- "cath_superfamily": tsv_loader.get_cath_superfamily(pdb_id, auth_asym_id),
210
- "state_id": tsv_loader.get_state_id(pdb_id, auth_asym_id),
211
- })
212
-
213
-
214
- @router.post("/chain/batch-resolve")
215
- async def batch_resolve_chains_endpoint(
216
- chains: List[Dict[str, str]],
217
- use_cache: bool = Query(True, description="Use cached data if available")
218
- ):
219
- """Batch resolve multiple chains at once."""
220
- chain_tuples = []
221
- for chain in chains:
222
- pdb_id = chain.get("pdb_id")
223
- auth_asym_id = chain.get("auth_asym_id")
224
- if pdb_id and auth_asym_id:
225
- chain_tuples.append((pdb_id, auth_asym_id))
226
-
227
- result = await batch_resolve_chains_async(chain_tuples, use_cache=use_cache)
228
- return {
229
- key: metadata.model_copy(update={
230
- "binding_status": tsv_loader.get_binding_status(*key.split("|")),
231
- "chain_composition": tsv_loader.get_chain_composition(*key.split("|")),
232
- "cath_id": tsv_loader.get_cath_id(*key.split("|")),
233
- "cath_superfamily": tsv_loader.get_cath_superfamily(*key.split("|")),
234
- "state_id": tsv_loader.get_state_id(*key.split("|")),
235
- })
236
- for key, metadata in result.items()
237
- }
238
-
239
-
240
- @router.get("/chain/cache/stats")
241
- async def cache_stats():
242
- """Get cache statistics."""
243
- cache = get_cache()
244
- return {
245
- "cached_chains": len(cache.cache),
246
- "ttl_hours": cache.ttl.total_seconds() / 3600
247
- }
248
-
249
-
250
- @router.get("/chain/functions")
251
- async def get_chain_functions(
252
- pdb_id: str = Query(..., description="4-character PDB ID"),
253
- auth_asym_id: str = Query(..., description="Chain identifier (e.g., 'A')"),
254
- ):
255
- """Return ranked functional annotations for a chain from node table."""
256
- functions = tsv_loader.get_functions(pdb_id, auth_asym_id)
257
- if functions is None:
258
- raise HTTPException(status_code=404, detail=f"No annotation found for {pdb_id}:{auth_asym_id}")
259
- return {"pdb_id": pdb_id.lower(), "chain_id": auth_asym_id.upper(), "functions": functions}
260
-
261
-
262
- @router.post("/chain/cache/clear")
263
- async def clear_cache():
264
- """Clear the chain metadata cache."""
265
- cache = get_cache()
266
- cache.clear()
267
- return {"message": "Cache cleared successfully"}
 
1
+ """Protein API routes."""
2
+ import logging
3
+ import math
4
+ from typing import Dict, List
5
+
6
+ from fastapi import APIRouter, HTTPException, Query
7
+
8
+ from app.protein.chain_resolver import (
9
+ batch_resolve_chains_async,
10
+ get_cache,
11
+ resolve_chain_async,
12
+ )
13
+ from app.protein.data_loader import DataManager
14
+ from app.protein import tsv_loader
15
+ from app.protein.models import (
16
+ ChainMetadata,
17
+ DataRecord,
18
+ DataResponse,
19
+ FilterOptions,
20
+ FilterParams,
21
+ SummaryStats,
22
+ )
23
+
24
+ logger = logging.getLogger(__name__)
25
+ router = APIRouter()
26
+
27
+ data_manager: DataManager | None = None
28
+
29
+
30
+ def set_data_manager(manager: DataManager | None) -> None:
31
+ """Set the active data manager for route handlers."""
32
+ global data_manager
33
+ data_manager = manager
34
+
35
+
36
+ def _require_data_manager() -> DataManager:
37
+ """Return data manager or raise if not initialized."""
38
+ if data_manager is None:
39
+ raise HTTPException(
40
+ status_code=503,
41
+ detail="Protein data is not initialized yet"
42
+ )
43
+ return data_manager
44
+
45
+
46
+ @router.get("/")
47
+ async def root():
48
+ """Root endpoint."""
49
+ return {
50
+ "message": "Protein Visualisation API",
51
+ "version": "1.0.0",
52
+ "endpoints": {
53
+ "/filters": "Get available filter options",
54
+ "/data": "Get filtered protein data",
55
+ "/summary": "Get aggregate statistics"
56
+ }
57
+ }
58
+
59
+
60
+ @router.get("/filters", response_model=FilterOptions)
61
+ async def get_filters():
62
+ """Get available filter options and ranges."""
63
+ try:
64
+ manager = _require_data_manager()
65
+ chains = sorted(manager.get_chain_sample())
66
+ fd = manager.get_filters_data()
67
+
68
+ return FilterOptions(
69
+ chain_ids=chains[:100],
70
+ rmsd_range={"min": float(fd["rmsd_min"]), "max": float(fd["rmsd_max"])},
71
+ tm_score_range={"min": float(fd["tm_min"]), "max": float(fd["tm_max"])},
72
+ length_range={"min": int(fd["length_min"]), "max": int(fd["length_max"])},
73
+ clusters=[
74
+ "cluster_ultra_high", "cluster_very_high", "cluster_high",
75
+ "cluster_medium", "cluster_low",
76
+ ],
77
+ total_records=fd["total_records"],
78
+ )
79
+ except HTTPException:
80
+ raise
81
+ except Exception as e:
82
+ logger.error(f"Error getting filters: {e}")
83
+ raise HTTPException(status_code=500, detail=str(e))
84
+
85
+
86
+ @router.post("/data", response_model=DataResponse)
87
+ async def get_data(filters: FilterParams):
88
+ """Get filtered protein data. Searches only by pdb_id_A and auth_asym_id_A."""
89
+ try:
90
+ manager = _require_data_manager()
91
+ limit = filters.limit or 1000
92
+ fetch_limit = limit * 10 if filters.cluster_id else limit
93
+
94
+ df = manager.query_edge(
95
+ pdb_id_a=filters.pdb_id,
96
+ auth_asym_id_a=filters.auth_asym_id,
97
+ chain_ids=filters.chain_ids,
98
+ rmsd_min=filters.rmsd_min,
99
+ rmsd_max=filters.rmsd_max,
100
+ tm_min=filters.tm_score_min,
101
+ tm_max=filters.tm_score_max,
102
+ limit=fetch_limit,
103
+ )
104
+
105
+ if filters.cluster_id:
106
+ df = df[df["cluster_id"] == filters.cluster_id]
107
+
108
+ filtered_total = len(df)
109
+ df = df.head(limit)
110
+
111
+ def _safe_float(val):
112
+ try:
113
+ v = float(val)
114
+ return None if math.isnan(v) else round(v, 6)
115
+ except (TypeError, ValueError):
116
+ return None
117
+
118
+ records = []
119
+ for row in df.itertuples(index=False):
120
+ raw_ph = getattr(row, "pH", None)
121
+ raw_temp = getattr(row, "temp_K", None)
122
+ records.append(DataRecord(
123
+ pdb_id_a=row.pdb_id_A,
124
+ auth_asym_id_a=row.auth_asym_id_A,
125
+ pdb_id_b=row.pdb_id_B,
126
+ auth_asym_id_b=row.auth_asym_id_B,
127
+ tm_score=float(row.TM1) if row.TM1 is not None else 0.0,
128
+ rmsd=float(row.RMSD) if row.RMSD is not None else 0.0,
129
+ structure_sim=_safe_float(row.structure_sim),
130
+ length_a=int(row.length_a) if row.length_a is not None else None,
131
+ length_b=int(row.length_b) if row.length_b is not None else None,
132
+ cluster_id=row.cluster_id,
133
+ exptl_method=getattr(row, "experimental_method", None) or None,
134
+ pH=_safe_float(raw_ph) if raw_ph else None,
135
+ temp=_safe_float(raw_temp) if raw_temp else None,
136
+ binding_status=getattr(row, "binding_status", None) or None,
137
+ chain_composition=getattr(row, "chain_composition", None) or None,
138
+ delta_rosetta=_safe_float(getattr(row, "delta_Rosetta", None)),
139
+ delta_foldx=_safe_float(getattr(row, "delta_FoldX", None)),
140
+ delta_evoef2=_safe_float(getattr(row, "delta_EvoEF2", None)),
141
+ delta_rw=_safe_float(getattr(row, "delta_RW", None)),
142
+ delta_rw_plus=_safe_float(getattr(row, "delta_RW_plus", None)),
143
+ state_id_b=getattr(row, "state_id_B", None) or None,
144
+ state_fidelity=getattr(row, "state_fidelity", None) or None,
145
+ avg_sim=getattr(row, "avg_sim", None) or None,
146
+ observation_fidelity=getattr(row, "observation_fidelity", None) or None,
147
+ ))
148
+
149
+ return DataResponse(
150
+ data=records,
151
+ total=manager.get_total(),
152
+ filtered=filtered_total,
153
+ )
154
+
155
+ except HTTPException:
156
+ raise
157
+ except Exception as e:
158
+ logger.error(f"Error filtering data: {e}")
159
+ raise HTTPException(status_code=500, detail=str(e))
160
+
161
+
162
+ @router.get("/summary", response_model=SummaryStats)
163
+ async def get_summary(
164
+ chain_ids: List[str] = Query(default=None),
165
+ rmsd_min: float = Query(default=None),
166
+ rmsd_max: float = Query(default=None),
167
+ tm_score_min: float = Query(default=None),
168
+ tm_score_max: float = Query(default=None),
169
+ ):
170
+ """Get aggregate statistics for the dataset or filtered subset."""
171
+ try:
172
+ manager = _require_data_manager()
173
+ stats = manager.get_summary_stats(
174
+ chain_ids=chain_ids,
175
+ rmsd_min=rmsd_min,
176
+ rmsd_max=rmsd_max,
177
+ tm_min=tm_score_min,
178
+ tm_max=tm_score_max,
179
+ )
180
+ return SummaryStats(
181
+ total_records=stats["total_records"],
182
+ unique_chains=stats["unique_chains"],
183
+ avg_rmsd=float(stats["avg_rmsd"]),
184
+ avg_tm_score=float(stats["avg_tm_score"]),
185
+ avg_sequence_length=float(stats["avg_sequence_length"]),
186
+ rmsd_distribution={str(k): int(v) for k, v in stats["rmsd_distribution"].items()},
187
+ tm_score_distribution={str(k): int(v) for k, v in stats["tm_score_distribution"].items()},
188
+ length_distribution={str(k): int(v) for k, v in stats["length_distribution"].items()},
189
+ )
190
+ except HTTPException:
191
+ raise
192
+ except Exception as e:
193
+ logger.error(f"Error computing summary: {e}")
194
+ raise HTTPException(status_code=500, detail=str(e))
195
+
196
+
197
+ @router.get("/chain/resolve", response_model=ChainMetadata)
198
+ async def resolve_chain_endpoint(
199
+ pdb_id: str = Query(..., description="4-character PDB ID"),
200
+ auth_asym_id: str = Query(..., description="Chain identifier (e.g., 'A')"),
201
+ use_cache: bool = Query(True, description="Use cached data if available")
202
+ ):
203
+ """Resolve chain metadata from RCSB web API."""
204
+ metadata = await resolve_chain_async(pdb_id, auth_asym_id, use_cache=use_cache)
205
+ return metadata.model_copy(update={
206
+ "binding_status": tsv_loader.get_binding_status(pdb_id, auth_asym_id),
207
+ "chain_composition": tsv_loader.get_chain_composition(pdb_id, auth_asym_id),
208
+ "cath_id": tsv_loader.get_cath_id(pdb_id, auth_asym_id),
209
+ "cath_superfamily": tsv_loader.get_cath_superfamily(pdb_id, auth_asym_id),
210
+ "state_id": tsv_loader.get_state_id(pdb_id, auth_asym_id),
211
+ })
212
+
213
+
214
+ @router.post("/chain/batch-resolve")
215
+ async def batch_resolve_chains_endpoint(
216
+ chains: List[Dict[str, str]],
217
+ use_cache: bool = Query(True, description="Use cached data if available")
218
+ ):
219
+ """Batch resolve multiple chains at once."""
220
+ chain_tuples = []
221
+ for chain in chains:
222
+ pdb_id = chain.get("pdb_id")
223
+ auth_asym_id = chain.get("auth_asym_id")
224
+ if pdb_id and auth_asym_id:
225
+ chain_tuples.append((pdb_id, auth_asym_id))
226
+
227
+ result = await batch_resolve_chains_async(chain_tuples, use_cache=use_cache)
228
+ return {
229
+ key: metadata.model_copy(update={
230
+ "binding_status": tsv_loader.get_binding_status(*key.split("|")),
231
+ "chain_composition": tsv_loader.get_chain_composition(*key.split("|")),
232
+ "cath_id": tsv_loader.get_cath_id(*key.split("|")),
233
+ "cath_superfamily": tsv_loader.get_cath_superfamily(*key.split("|")),
234
+ "state_id": tsv_loader.get_state_id(*key.split("|")),
235
+ })
236
+ for key, metadata in result.items()
237
+ }
238
+
239
+
240
+ @router.get("/chain/cache/stats")
241
+ async def cache_stats():
242
+ """Get cache statistics."""
243
+ cache = get_cache()
244
+ return {
245
+ "cached_chains": len(cache.cache),
246
+ "ttl_hours": cache.ttl.total_seconds() / 3600
247
+ }
248
+
249
+
250
+ @router.get("/chain/functions")
251
+ async def get_chain_functions(
252
+ pdb_id: str = Query(..., description="4-character PDB ID"),
253
+ auth_asym_id: str = Query(..., description="Chain identifier (e.g., 'A')"),
254
+ ):
255
+ """Return ranked functional annotations for a chain from node table."""
256
+ functions = tsv_loader.get_functions(pdb_id, auth_asym_id)
257
+ if functions is None:
258
+ raise HTTPException(status_code=404, detail=f"No annotation found for {pdb_id}:{auth_asym_id}")
259
+ return {"pdb_id": pdb_id.lower(), "chain_id": auth_asym_id.upper(), "functions": functions}
260
+
261
+
262
+ @router.post("/chain/cache/clear")
263
+ async def clear_cache():
264
+ """Clear the chain metadata cache."""
265
+ cache = get_cache()
266
+ cache.clear()
267
+ return {"message": "Cache cleared successfully"}
backend/app/protein/data_loader.py CHANGED
@@ -1,337 +1,338 @@
1
- """
2
- SQLite-based data loader for protein structure data.
3
- Reads from MuSProt.db: edge table (pairwise comparisons from CSV)
4
- and node table (chain annotations from TSV).
5
- """
6
- import json
7
- import pandas as pd
8
- from pathlib import Path
9
- from typing import Any, Dict, List, Optional, Set
10
- import logging
11
-
12
- from app.protein.config import get_database_path, get_summary_path
13
- from app.protein.database import connect_readonly
14
-
15
- logger = logging.getLogger(__name__)
16
-
17
-
18
- def _assign_cluster(tm: float) -> str:
19
- if tm >= 0.99:
20
- return "cluster_ultra_high"
21
- if tm >= 0.95:
22
- return "cluster_very_high"
23
- if tm >= 0.90:
24
- return "cluster_high"
25
- if tm >= 0.80:
26
- return "cluster_medium"
27
- return "cluster_low"
28
-
29
-
30
- class DataManager:
31
- """Manages SQL queries against MuSProt.db edge and node tables."""
32
-
33
- def __init__(self, db_path: Optional[Path] = None, **_kwargs):
34
- self.db_path = db_path or get_database_path()
35
- self._loaded = False
36
- self._total_records = 0
37
- self._summary: Dict[str, Any] = {}
38
-
39
- def load_data(self) -> None:
40
- """Verify DB is accessible and load optional precomputed metadata."""
41
- if self._loaded:
42
- return
43
- summary_path = get_summary_path()
44
- if summary_path:
45
- self._summary = json.loads(summary_path.read_text(encoding="utf-8"))
46
- self._total_records = int(self._summary.get("total_records", 0))
47
-
48
- conn = connect_readonly(self.db_path)
49
- try:
50
- cur = conn.cursor()
51
- cur.execute("SELECT 1 FROM edge LIMIT 1")
52
- if not self._total_records:
53
- logger.warning(
54
- "No musprot_summary.json found; counting edge rows during startup. "
55
- "Publish the summary sidecar to avoid this scan."
56
- )
57
- cur.execute("SELECT COUNT(*) FROM edge")
58
- self._total_records = cur.fetchone()[0]
59
- finally:
60
- conn.close()
61
- self._loaded = True
62
- logger.info(f"DataManager ready: {self._total_records:,} edge records in {self.db_path}")
63
-
64
- def _connect(self):
65
- if not self._loaded:
66
- raise RuntimeError("Call load_data() first.")
67
- return connect_readonly(self.db_path)
68
-
69
- def get_chain_sample(self, limit: int = 500) -> Set[str]:
70
- """Return a small chain sample without loading the full node index."""
71
- conn = self._connect()
72
- try:
73
- cur = conn.cursor()
74
- cur.execute(
75
- "SELECT pdb_id, auth_asym_id FROM node "
76
- "WHERE pdb_id IS NOT NULL AND auth_asym_id IS NOT NULL LIMIT ?",
77
- (limit,),
78
- )
79
- return {
80
- f"{str(pdb).upper()}_{str(chain).upper()}"
81
- for pdb, chain in cur.fetchall()
82
- }
83
- finally:
84
- conn.close()
85
-
86
- def get_filters_data(self) -> Dict[str, Any]:
87
- """Return aggregated ranges for filter UI."""
88
- if self._summary:
89
- return {
90
- "tm_min": self._summary["tm_score_range"]["min"],
91
- "tm_max": self._summary["tm_score_range"]["max"],
92
- "rmsd_min": self._summary["rmsd_range"]["min"],
93
- "rmsd_max": self._summary["rmsd_range"]["max"],
94
- "length_min": self._summary["length_range"]["min"],
95
- "length_max": self._summary["length_range"]["max"],
96
- "total_records": self._total_records,
97
- }
98
-
99
- conn = self._connect()
100
- try:
101
- cur = conn.cursor()
102
- cur.execute(
103
- "SELECT MIN(CAST(TM1 AS REAL)), MAX(CAST(TM1 AS REAL)),"
104
- " MIN(CAST(RMSD AS REAL)), MAX(CAST(RMSD AS REAL))"
105
- " FROM edge"
106
- )
107
- tm_min, tm_max, rmsd_min, rmsd_max = cur.fetchone()
108
- finally:
109
- conn.close()
110
-
111
- from app.protein import tsv_loader
112
- lengths = []
113
- for row in tsv_loader._get_index().values():
114
- try:
115
- lengths.append(int(row["sequence_length"]))
116
- except (ValueError, TypeError):
117
- pass
118
-
119
- return {
120
- "tm_min": tm_min or 0.0,
121
- "tm_max": tm_max or 1.0,
122
- "rmsd_min": rmsd_min or 0.0,
123
- "rmsd_max": rmsd_max or 10.0,
124
- "length_min": min(lengths) if lengths else 0,
125
- "length_max": max(lengths) if lengths else 0,
126
- "total_records": self._total_records,
127
- }
128
-
129
- def query_edge(
130
- self,
131
- pdb_id_a: Optional[str] = None,
132
- auth_asym_id_a: Optional[str] = None,
133
- chain_ids: Optional[List[str]] = None,
134
- rmsd_min: Optional[float] = None,
135
- rmsd_max: Optional[float] = None,
136
- tm_min: Optional[float] = None,
137
- tm_max: Optional[float] = None,
138
- limit: int = 1000,
139
- ) -> pd.DataFrame:
140
- """Query edge table with filters. Returns DataFrame with cluster_id and lengths."""
141
- where_clauses: List[str] = []
142
- params: List[Any] = []
143
-
144
- if pdb_id_a and auth_asym_id_a:
145
- where_clauses.append("LOWER(pdb_id_A) = LOWER(?) AND LOWER(auth_asym_id_A) = LOWER(?)")
146
- params.extend([pdb_id_a, auth_asym_id_a])
147
- elif chain_ids:
148
- sub: List[str] = []
149
- for cid in chain_ids:
150
- parts = cid.split("_", 1)
151
- if len(parts) == 2:
152
- sub.append("(LOWER(pdb_id_A) = LOWER(?) AND LOWER(auth_asym_id_A) = LOWER(?))")
153
- params.extend(parts)
154
- if sub:
155
- where_clauses.append("(" + " OR ".join(sub) + ")")
156
-
157
- if rmsd_min is not None:
158
- where_clauses.append("CAST(RMSD AS REAL) >= ?")
159
- params.append(rmsd_min)
160
- if rmsd_max is not None:
161
- where_clauses.append("CAST(RMSD AS REAL) <= ?")
162
- params.append(rmsd_max)
163
- if tm_min is not None:
164
- where_clauses.append("CAST(TM1 AS REAL) >= ?")
165
- params.append(tm_min)
166
- if tm_max is not None:
167
- where_clauses.append("CAST(TM1 AS REAL) <= ?")
168
- params.append(tm_max)
169
-
170
- where_sql = ("WHERE " + " AND ".join(where_clauses)) if where_clauses else ""
171
- sql = f"""
172
- SELECT e.pdb_id_A, e.auth_asym_id_A, e.pdb_id_B, e.auth_asym_id_B,
173
- CAST(e.TM1 AS REAL) AS TM1,
174
- CAST(e.RMSD AS REAL) AS RMSD,
175
- CAST(e.structure_sim AS REAL) AS structure_sim,
176
- e."delta_Rosetta", e."delta_FoldX", e."delta_EvoEF2",
177
- e."delta_RW", e."delta_RW+",
178
- e.state_id_B, e.state_fidelity, e.avg_sim, e.observation_fidelity
179
- FROM edge e
180
- {where_sql}
181
- LIMIT ?
182
- """
183
- params.append(limit)
184
-
185
- conn = self._connect()
186
- try:
187
- df = pd.read_sql_query(sql, conn, params=params)
188
- finally:
189
- conn.close()
190
-
191
- # Rename columns with special characters so itertuples works cleanly
192
- df = df.rename(columns={"delta_RW+": "delta_RW_plus"})
193
-
194
- df["cluster_id"] = df["TM1"].apply(
195
- lambda t: _assign_cluster(t) if pd.notna(t) else "cluster_low"
196
- )
197
-
198
- from app.protein import tsv_loader
199
- df["length_a"] = [
200
- tsv_loader.get_sequence(r.pdb_id_A, r.auth_asym_id_A)[1]
201
- for r in df.itertuples(index=False)
202
- ]
203
- df["length_b"] = [
204
- tsv_loader.get_sequence(r.pdb_id_B, r.auth_asym_id_B)[1]
205
- for r in df.itertuples(index=False)
206
- ]
207
- node_rows = [
208
- tsv_loader.get_node_row(r.pdb_id_B, r.auth_asym_id_B)
209
- for r in df.itertuples(index=False)
210
- ]
211
- df["experimental_method"] = [row.get("experimental_method") if row else None for row in node_rows]
212
- df["pH"] = [row.get("pH") if row else None for row in node_rows]
213
- df["temp_K"] = [row.get("temp_K") if row else None for row in node_rows]
214
- df["base_label"] = [row.get("base_label") if row else None for row in node_rows]
215
- df["chain_composition"] = [row.get("chain_composition") if row else None for row in node_rows]
216
-
217
- return df
218
-
219
- def get_summary_stats(
220
- self,
221
- chain_ids: Optional[List[str]] = None,
222
- rmsd_min: Optional[float] = None,
223
- rmsd_max: Optional[float] = None,
224
- tm_min: Optional[float] = None,
225
- tm_max: Optional[float] = None,
226
- ) -> Dict[str, Any]:
227
- """Compute aggregate statistics via SQL, with optional filters."""
228
- if not any([chain_ids, rmsd_min, rmsd_max, tm_min, tm_max]) and self._summary:
229
- return self._summary
230
-
231
- where_clauses: List[str] = []
232
- params: List[Any] = []
233
-
234
- if chain_ids:
235
- sub: List[str] = []
236
- for cid in chain_ids:
237
- parts = cid.split("_", 1)
238
- if len(parts) == 2:
239
- sub.append("(LOWER(pdb_id_A) = LOWER(?) AND LOWER(auth_asym_id_A) = LOWER(?))")
240
- params.extend(parts)
241
- if sub:
242
- where_clauses.append("(" + " OR ".join(sub) + ")")
243
-
244
- if rmsd_min is not None:
245
- where_clauses.append("CAST(RMSD AS REAL) >= ?")
246
- params.append(rmsd_min)
247
- if rmsd_max is not None:
248
- where_clauses.append("CAST(RMSD AS REAL) <= ?")
249
- params.append(rmsd_max)
250
- if tm_min is not None:
251
- where_clauses.append("CAST(TM1 AS REAL) >= ?")
252
- params.append(tm_min)
253
- if tm_max is not None:
254
- where_clauses.append("CAST(TM1 AS REAL) <= ?")
255
- params.append(tm_max)
256
-
257
- where_sql = ("WHERE " + " AND ".join(where_clauses)) if where_clauses else ""
258
-
259
- conn = self._connect()
260
- try:
261
- cur = conn.cursor()
262
- cur.execute(
263
- f"SELECT COUNT(*), AVG(CAST(RMSD AS REAL)), AVG(CAST(TM1 AS REAL))"
264
- f" FROM edge {where_sql}",
265
- params,
266
- )
267
- total, avg_rmsd, avg_tm = cur.fetchone()
268
-
269
- cur.execute(
270
- f"SELECT COUNT(DISTINCT LOWER(pdb_id_A) || '_' || LOWER(auth_asym_id_A))"
271
- f" FROM edge {where_sql}",
272
- params,
273
- )
274
- unique_chains = cur.fetchone()[0]
275
-
276
- def _dist(col, bins, labels):
277
- cases = " ".join(
278
- f"WHEN CAST({col} AS REAL) >= {lo} AND CAST({col} AS REAL) < {hi} THEN '{lbl}'"
279
- for (lo, hi), lbl in zip(zip(bins, bins[1:]), labels)
280
- )
281
- sql = (
282
- f"SELECT CASE {cases} ELSE '{labels[-1]}' END AS bucket, COUNT(*)"
283
- f" FROM edge {where_sql} GROUP BY bucket"
284
- )
285
- cur.execute(sql, params)
286
- return {r[0]: r[1] for r in cur.fetchall()}
287
-
288
- rmsd_dist = _dist(
289
- "RMSD",
290
- [0, 0.5, 1.0, 1.5, 2.0, 5.0, 1e9],
291
- ["0-0.5", "0.5-1.0", "1.0-1.5", "1.5-2.0", "2.0-5.0", ">5.0"],
292
- )
293
- tm_dist = _dist(
294
- "TM1",
295
- [0, 0.5, 0.7, 0.85, 0.95, 1.01],
296
- ["0-0.5", "0.5-0.7", "0.7-0.85", "0.85-0.95", "0.95-1.0"],
297
- )
298
- finally:
299
- conn.close()
300
-
301
- from app.protein import tsv_loader
302
- lengths = []
303
- for row in tsv_loader._get_index().values():
304
- try:
305
- lengths.append(int(row["sequence_length"]))
306
- except (ValueError, TypeError):
307
- pass
308
- avg_len = sum(lengths) / len(lengths) if lengths else 0.0
309
- len_dist = {}
310
- for length in lengths:
311
- if length <= 100:
312
- bucket = "0-100"
313
- elif length <= 200:
314
- bucket = "100-200"
315
- elif length <= 300:
316
- bucket = "200-300"
317
- elif length <= 500:
318
- bucket = "300-500"
319
- elif length <= 1000:
320
- bucket = "500-1000"
321
- else:
322
- bucket = ">1000"
323
- len_dist[bucket] = len_dist.get(bucket, 0) + 1
324
-
325
- return {
326
- "total_records": total or 0,
327
- "unique_chains": unique_chains or 0,
328
- "avg_rmsd": avg_rmsd or 0.0,
329
- "avg_tm_score": avg_tm or 0.0,
330
- "avg_sequence_length": avg_len,
331
- "rmsd_distribution": rmsd_dist,
332
- "tm_score_distribution": tm_dist,
333
- "length_distribution": len_dist,
334
- }
335
-
336
- def get_total(self) -> int:
337
- return self._total_records
 
 
1
+ """
2
+ SQLite-based data loader for protein structure data.
3
+ Reads from MuSProt.db: edge table (pairwise comparisons from CSV)
4
+ and node table (chain annotations from TSV).
5
+ """
6
+ import json
7
+ import pandas as pd
8
+ from pathlib import Path
9
+ from typing import Any, Dict, List, Optional, Set
10
+ import logging
11
+
12
+ from app.protein.config import get_database_path, get_summary_path
13
+ from app.protein.database import connect_readonly
14
+
15
+ logger = logging.getLogger(__name__)
16
+
17
+
18
+ def _assign_cluster(tm: float) -> str:
19
+ if tm >= 0.99:
20
+ return "cluster_ultra_high"
21
+ if tm >= 0.95:
22
+ return "cluster_very_high"
23
+ if tm >= 0.90:
24
+ return "cluster_high"
25
+ if tm >= 0.80:
26
+ return "cluster_medium"
27
+ return "cluster_low"
28
+
29
+
30
+ class DataManager:
31
+ """Manages SQL queries against MuSProt.db edge and node tables."""
32
+
33
+ def __init__(self, db_path: Optional[Path] = None, **_kwargs):
34
+ self.db_path = db_path or get_database_path()
35
+ self._loaded = False
36
+ self._total_records = 0
37
+ self._summary: Dict[str, Any] = {}
38
+
39
+ def load_data(self) -> None:
40
+ """Verify DB is accessible and load optional precomputed metadata."""
41
+ if self._loaded:
42
+ return
43
+ summary_path = get_summary_path()
44
+ if summary_path:
45
+ self._summary = json.loads(summary_path.read_text(encoding="utf-8"))
46
+ self._total_records = int(self._summary.get("total_records", 0))
47
+
48
+ conn = connect_readonly(self.db_path)
49
+ try:
50
+ cur = conn.cursor()
51
+ cur.execute("SELECT 1 FROM edge LIMIT 1")
52
+ if not self._total_records:
53
+ logger.warning(
54
+ "No musprot_summary.json found; counting edge rows during startup. "
55
+ "Publish the summary sidecar to avoid this scan."
56
+ )
57
+ cur.execute("SELECT COUNT(*) FROM edge")
58
+ self._total_records = cur.fetchone()[0]
59
+ finally:
60
+ conn.close()
61
+ self._loaded = True
62
+ logger.info(f"DataManager ready: {self._total_records:,} edge records in {self.db_path}")
63
+
64
+ def _connect(self):
65
+ if not self._loaded:
66
+ raise RuntimeError("Call load_data() first.")
67
+ return connect_readonly(self.db_path)
68
+
69
+ def get_chain_sample(self, limit: int = 500) -> Set[str]:
70
+ """Return a small chain sample without loading the full node index."""
71
+ conn = self._connect()
72
+ try:
73
+ cur = conn.cursor()
74
+ cur.execute(
75
+ "SELECT pdb_id, auth_asym_id FROM node "
76
+ "WHERE pdb_id IS NOT NULL AND auth_asym_id IS NOT NULL LIMIT ?",
77
+ (limit,),
78
+ )
79
+ return {
80
+ f"{str(pdb).upper()}_{str(chain).upper()}"
81
+ for pdb, chain in cur.fetchall()
82
+ }
83
+ finally:
84
+ conn.close()
85
+
86
+ def get_filters_data(self) -> Dict[str, Any]:
87
+ """Return aggregated ranges for filter UI."""
88
+ if self._summary:
89
+ return {
90
+ "tm_min": self._summary["tm_score_range"]["min"],
91
+ "tm_max": self._summary["tm_score_range"]["max"],
92
+ "rmsd_min": self._summary["rmsd_range"]["min"],
93
+ "rmsd_max": self._summary["rmsd_range"]["max"],
94
+ "length_min": self._summary["length_range"]["min"],
95
+ "length_max": self._summary["length_range"]["max"],
96
+ "total_records": self._total_records,
97
+ }
98
+
99
+ conn = self._connect()
100
+ try:
101
+ cur = conn.cursor()
102
+ cur.execute(
103
+ "SELECT MIN(CAST(TM1 AS REAL)), MAX(CAST(TM1 AS REAL)),"
104
+ " MIN(CAST(RMSD AS REAL)), MAX(CAST(RMSD AS REAL))"
105
+ " FROM edge"
106
+ )
107
+ tm_min, tm_max, rmsd_min, rmsd_max = cur.fetchone()
108
+ finally:
109
+ conn.close()
110
+
111
+ from app.protein import tsv_loader
112
+ lengths = []
113
+ for row in tsv_loader._get_index().values():
114
+ try:
115
+ lengths.append(int(row["sequence_length"]))
116
+ except (ValueError, TypeError):
117
+ pass
118
+
119
+ return {
120
+ "tm_min": tm_min or 0.0,
121
+ "tm_max": tm_max or 1.0,
122
+ "rmsd_min": rmsd_min or 0.0,
123
+ "rmsd_max": rmsd_max or 10.0,
124
+ "length_min": min(lengths) if lengths else 0,
125
+ "length_max": max(lengths) if lengths else 0,
126
+ "total_records": self._total_records,
127
+ }
128
+
129
+ def query_edge(
130
+ self,
131
+ pdb_id_a: Optional[str] = None,
132
+ auth_asym_id_a: Optional[str] = None,
133
+ chain_ids: Optional[List[str]] = None,
134
+ rmsd_min: Optional[float] = None,
135
+ rmsd_max: Optional[float] = None,
136
+ tm_min: Optional[float] = None,
137
+ tm_max: Optional[float] = None,
138
+ limit: int = 1000,
139
+ ) -> pd.DataFrame:
140
+ """Query edge table with filters. Returns DataFrame with cluster_id and lengths."""
141
+ where_clauses: List[str] = []
142
+ params: List[Any] = []
143
+
144
+ if pdb_id_a and auth_asym_id_a:
145
+ where_clauses.append("pdb_id_A = ? COLLATE NOCASE AND auth_asym_id_A = ? COLLATE NOCASE")
146
+ params.extend([pdb_id_a, auth_asym_id_a])
147
+ elif chain_ids:
148
+ sub: List[str] = []
149
+ for cid in chain_ids:
150
+ parts = cid.split("_", 1)
151
+ if len(parts) == 2:
152
+ sub.append("(pdb_id_A = ? COLLATE NOCASE AND auth_asym_id_A = ? COLLATE NOCASE)")
153
+ params.extend(parts)
154
+ if sub:
155
+ where_clauses.append("(" + " OR ".join(sub) + ")")
156
+
157
+ if rmsd_min is not None:
158
+ where_clauses.append("CAST(RMSD AS REAL) >= ?")
159
+ params.append(rmsd_min)
160
+ if rmsd_max is not None:
161
+ where_clauses.append("CAST(RMSD AS REAL) <= ?")
162
+ params.append(rmsd_max)
163
+ if tm_min is not None:
164
+ where_clauses.append("CAST(TM1 AS REAL) >= ?")
165
+ params.append(tm_min)
166
+ if tm_max is not None:
167
+ where_clauses.append("CAST(TM1 AS REAL) <= ?")
168
+ params.append(tm_max)
169
+
170
+ where_sql = ("WHERE " + " AND ".join(where_clauses)) if where_clauses else ""
171
+ sql = f"""
172
+ SELECT e.pdb_id_A, e.auth_asym_id_A, e.pdb_id_B, e.auth_asym_id_B,
173
+ CAST(e.TM1 AS REAL) AS TM1,
174
+ CAST(e.RMSD AS REAL) AS RMSD,
175
+ CAST(e.pair_similarity AS REAL) AS structure_sim,
176
+ e."delta_Rosetta", e."delta_FoldX", e."delta_EvoEF2",
177
+ e."delta_RW", e."delta_RW+",
178
+ e.state_id_B, e.state_fidelity,
179
+ e.state_similarity AS avg_sim, e.pair_fidelity AS observation_fidelity
180
+ FROM edge e
181
+ {where_sql}
182
+ LIMIT ?
183
+ """
184
+ params.append(limit)
185
+
186
+ conn = self._connect()
187
+ try:
188
+ df = pd.read_sql_query(sql, conn, params=params)
189
+ finally:
190
+ conn.close()
191
+
192
+ # Rename columns with special characters so itertuples works cleanly
193
+ df = df.rename(columns={"delta_RW+": "delta_RW_plus"})
194
+
195
+ df["cluster_id"] = df["TM1"].apply(
196
+ lambda t: _assign_cluster(t) if pd.notna(t) else "cluster_low"
197
+ )
198
+
199
+ from app.protein import tsv_loader
200
+ df["length_a"] = [
201
+ tsv_loader.get_sequence(r.pdb_id_A, r.auth_asym_id_A)[1]
202
+ for r in df.itertuples(index=False)
203
+ ]
204
+ df["length_b"] = [
205
+ tsv_loader.get_sequence(r.pdb_id_B, r.auth_asym_id_B)[1]
206
+ for r in df.itertuples(index=False)
207
+ ]
208
+ node_rows = [
209
+ tsv_loader.get_node_row(r.pdb_id_B, r.auth_asym_id_B)
210
+ for r in df.itertuples(index=False)
211
+ ]
212
+ df["experimental_method"] = [row.get("experimental_method") if row else None for row in node_rows]
213
+ df["pH"] = [row.get("pH") if row else None for row in node_rows]
214
+ df["temp_K"] = [row.get("temp_K") if row else None for row in node_rows]
215
+ df["binding_status"] = [row.get("binding_status") if row else None for row in node_rows]
216
+ df["chain_composition"] = [row.get("chain_composition") if row else None for row in node_rows]
217
+
218
+ return df
219
+
220
+ def get_summary_stats(
221
+ self,
222
+ chain_ids: Optional[List[str]] = None,
223
+ rmsd_min: Optional[float] = None,
224
+ rmsd_max: Optional[float] = None,
225
+ tm_min: Optional[float] = None,
226
+ tm_max: Optional[float] = None,
227
+ ) -> Dict[str, Any]:
228
+ """Compute aggregate statistics via SQL, with optional filters."""
229
+ if not any([chain_ids, rmsd_min, rmsd_max, tm_min, tm_max]) and self._summary:
230
+ return self._summary
231
+
232
+ where_clauses: List[str] = []
233
+ params: List[Any] = []
234
+
235
+ if chain_ids:
236
+ sub: List[str] = []
237
+ for cid in chain_ids:
238
+ parts = cid.split("_", 1)
239
+ if len(parts) == 2:
240
+ sub.append("(pdb_id_A = ? COLLATE NOCASE AND auth_asym_id_A = ? COLLATE NOCASE)")
241
+ params.extend(parts)
242
+ if sub:
243
+ where_clauses.append("(" + " OR ".join(sub) + ")")
244
+
245
+ if rmsd_min is not None:
246
+ where_clauses.append("CAST(RMSD AS REAL) >= ?")
247
+ params.append(rmsd_min)
248
+ if rmsd_max is not None:
249
+ where_clauses.append("CAST(RMSD AS REAL) <= ?")
250
+ params.append(rmsd_max)
251
+ if tm_min is not None:
252
+ where_clauses.append("CAST(TM1 AS REAL) >= ?")
253
+ params.append(tm_min)
254
+ if tm_max is not None:
255
+ where_clauses.append("CAST(TM1 AS REAL) <= ?")
256
+ params.append(tm_max)
257
+
258
+ where_sql = ("WHERE " + " AND ".join(where_clauses)) if where_clauses else ""
259
+
260
+ conn = self._connect()
261
+ try:
262
+ cur = conn.cursor()
263
+ cur.execute(
264
+ f"SELECT COUNT(*), AVG(CAST(RMSD AS REAL)), AVG(CAST(TM1 AS REAL))"
265
+ f" FROM edge {where_sql}",
266
+ params,
267
+ )
268
+ total, avg_rmsd, avg_tm = cur.fetchone()
269
+
270
+ cur.execute(
271
+ f"SELECT COUNT(DISTINCT LOWER(pdb_id_A) || '_' || LOWER(auth_asym_id_A))"
272
+ f" FROM edge {where_sql}",
273
+ params,
274
+ )
275
+ unique_chains = cur.fetchone()[0]
276
+
277
+ def _dist(col, bins, labels):
278
+ cases = " ".join(
279
+ f"WHEN CAST({col} AS REAL) >= {lo} AND CAST({col} AS REAL) < {hi} THEN '{lbl}'"
280
+ for (lo, hi), lbl in zip(zip(bins, bins[1:]), labels)
281
+ )
282
+ sql = (
283
+ f"SELECT CASE {cases} ELSE '{labels[-1]}' END AS bucket, COUNT(*)"
284
+ f" FROM edge {where_sql} GROUP BY bucket"
285
+ )
286
+ cur.execute(sql, params)
287
+ return {r[0]: r[1] for r in cur.fetchall()}
288
+
289
+ rmsd_dist = _dist(
290
+ "RMSD",
291
+ [0, 0.5, 1.0, 1.5, 2.0, 5.0, 1e9],
292
+ ["0-0.5", "0.5-1.0", "1.0-1.5", "1.5-2.0", "2.0-5.0", ">5.0"],
293
+ )
294
+ tm_dist = _dist(
295
+ "TM1",
296
+ [0, 0.5, 0.7, 0.85, 0.95, 1.01],
297
+ ["0-0.5", "0.5-0.7", "0.7-0.85", "0.85-0.95", "0.95-1.0"],
298
+ )
299
+ finally:
300
+ conn.close()
301
+
302
+ from app.protein import tsv_loader
303
+ lengths = []
304
+ for row in tsv_loader._get_index().values():
305
+ try:
306
+ lengths.append(int(row["sequence_length"]))
307
+ except (ValueError, TypeError):
308
+ pass
309
+ avg_len = sum(lengths) / len(lengths) if lengths else 0.0
310
+ len_dist = {}
311
+ for length in lengths:
312
+ if length <= 100:
313
+ bucket = "0-100"
314
+ elif length <= 200:
315
+ bucket = "100-200"
316
+ elif length <= 300:
317
+ bucket = "200-300"
318
+ elif length <= 500:
319
+ bucket = "300-500"
320
+ elif length <= 1000:
321
+ bucket = "500-1000"
322
+ else:
323
+ bucket = ">1000"
324
+ len_dist[bucket] = len_dist.get(bucket, 0) + 1
325
+
326
+ return {
327
+ "total_records": total or 0,
328
+ "unique_chains": unique_chains or 0,
329
+ "avg_rmsd": avg_rmsd or 0.0,
330
+ "avg_tm_score": avg_tm or 0.0,
331
+ "avg_sequence_length": avg_len,
332
+ "rmsd_distribution": rmsd_dist,
333
+ "tm_score_distribution": tm_dist,
334
+ "length_distribution": len_dist,
335
+ }
336
+
337
+ def get_total(self) -> int:
338
+ return self._total_records
backend/app/protein/tsv_loader.py CHANGED
@@ -1,210 +1,210 @@
1
- """
2
- Node table lookup for chain annotations.
3
-
4
- Loads the node table from MuSProt.db once on first use and indexes rows
5
- by (pdb_id_lower, auth_asym_id_upper). When multiple rows share the same
6
- key the first occurrence is kept.
7
- """
8
- import ast
9
- import sqlite3
10
- import logging
11
- import zlib
12
- from functools import lru_cache
13
- from typing import Dict, List, Optional, Tuple, Any
14
-
15
- from app.protein.config import get_database_path, get_node_database_path
16
- from app.protein.database import connect_readonly
17
-
18
- logger = logging.getLogger(__name__)
19
-
20
- _Row = Dict[str, str]
21
-
22
- # (pdb_id_lower, auth_asym_id_upper) -> relevant fields
23
- _index: Optional[Dict[Tuple[str, str], _Row]] = None
24
-
25
-
26
- def _load() -> Dict[Tuple[str, str], _Row]:
27
- index: Dict[Tuple[str, str], _Row] = {}
28
- db_path = get_database_path()
29
- conn = connect_readonly(db_path)
30
- try:
31
- conn.row_factory = sqlite3.Row
32
- cur = conn.cursor()
33
- cur.execute(
34
- 'SELECT pdb_id, auth_asym_id, base_label, sequence, sequence_length,'
35
- ' CATH_ID, cath_superfamily, Rosetta, FoldX, EvoEF2, RW, "RW+", ranked_functions,'
36
- ' state_id, experimental_method, pH, temp_K, chain_composition'
37
- ' FROM node'
38
- )
39
- for row in cur:
40
- pdb_id = row["pdb_id"] or ""
41
- auth_asym_id = row["auth_asym_id"] or ""
42
- key = (pdb_id.lower(), auth_asym_id.upper())
43
- if key not in index:
44
- index[key] = {
45
- "base_label": row["base_label"] or "",
46
- "sequence": row["sequence"] or "",
47
- "sequence_length": row["sequence_length"] or "",
48
- "CATH_ID": row["CATH_ID"] or "",
49
- "cath_superfamily": row["cath_superfamily"] or "",
50
- "Rosetta": row["Rosetta"] or "",
51
- "FoldX": row["FoldX"] or "",
52
- "EvoEF2": row["EvoEF2"] or "",
53
- "RW": row["RW"] or "",
54
- "RW+": row["RW+"] or "",
55
- "ranked_functions": row["ranked_functions"] or "",
56
- "state_id": row["state_id"] or "",
57
- "experimental_method": row["experimental_method"] or "",
58
- "pH": row["pH"] or "",
59
- "temp_K": row["temp_K"] or "",
60
- "chain_composition": row["chain_composition"] or "",
61
- }
62
- finally:
63
- conn.close()
64
-
65
- logger.info(f"Node index loaded: {len(index):,} entries from {db_path}")
66
- return index
67
-
68
-
69
- def _get_index() -> Dict[Tuple[str, str], _Row]:
70
- global _index
71
- if _index is None:
72
- _index = _load()
73
- return _index
74
-
75
-
76
- def reset_index() -> None:
77
- """Force reload of node index."""
78
- global _index
79
- _index = None
80
- _lookup_node_row.cache_clear()
81
-
82
-
83
- @lru_cache(maxsize=4096)
84
- def _lookup_node_row(pdb_id: str, auth_asym_id: str) -> Optional[_Row]:
85
- """Read one node row without blocking on a full-table in-memory index."""
86
- conn = connect_readonly(get_node_database_path())
87
- try:
88
- conn.row_factory = sqlite3.Row
89
- row = conn.execute(
90
- 'SELECT base_label, sequence, sequence_length, CATH_ID, cath_superfamily,'
91
- ' Rosetta, FoldX, EvoEF2, RW, "RW+", ranked_functions,'
92
- ' state_id, experimental_method, pH, temp_K, chain_composition'
93
- ' FROM node WHERE LOWER(pdb_id) = ? AND UPPER(auth_asym_id) = ? LIMIT 1',
94
- (pdb_id.lower(), auth_asym_id.upper()),
95
- ).fetchone()
96
- if row is None:
97
- return None
98
- return {
99
- key: zlib.decompress(value).decode("utf-8") if isinstance(value, bytes) else value or ""
100
- for key, value in ((key, row[key]) for key in row.keys())
101
- }
102
- finally:
103
- conn.close()
104
-
105
-
106
- def _parse_functions(raw: str) -> List[str]:
107
- raw = raw.strip()
108
- if not raw:
109
- return []
110
- try:
111
- parsed = ast.literal_eval(raw)
112
- if isinstance(parsed, list):
113
- return [str(f).strip() for f in parsed if str(f).strip()]
114
- except (ValueError, SyntaxError):
115
- pass
116
- return [f.strip() for f in raw.split(";") if f.strip()]
117
-
118
-
119
- def _parse_float(val: str) -> Optional[float]:
120
- val = str(val).strip()
121
- if not val:
122
- return None
123
- try:
124
- return float(val)
125
- except ValueError:
126
- return None
127
-
128
-
129
- def get_energy_scores(pdb_id: str, auth_asym_id: str) -> Optional[Dict[str, Any]]:
130
- """Return energy scores (Rosetta, FoldX, EvoEF2, RW, RW+) or None if not found."""
131
- row = get_node_row(pdb_id, auth_asym_id)
132
- if row is None:
133
- return None
134
- return {
135
- "Rosetta": _parse_float(row.get("Rosetta", "")),
136
- "FoldX": _parse_float(row.get("FoldX", "")),
137
- "EvoEF2": _parse_float(row.get("EvoEF2", "")),
138
- "RW": _parse_float(row.get("RW", "")),
139
- "RW+": _parse_float(row.get("RW+", "")),
140
- }
141
-
142
-
143
- def get_node_row(pdb_id: str, auth_asym_id: str) -> Optional[_Row]:
144
- """Return the cached node row for a chain."""
145
- key = (pdb_id.lower(), auth_asym_id.upper())
146
- return _index.get(key) if _index is not None else _lookup_node_row(*key)
147
-
148
-
149
- def get_binding_status(pdb_id: str, auth_asym_id: str) -> Optional[str]:
150
- """Return binding status (base_label: apo/holo) or None if not found."""
151
- row = get_node_row(pdb_id, auth_asym_id)
152
- if row is None:
153
- return None
154
- return row.get("base_label") or None
155
-
156
-
157
- def get_chain_composition(pdb_id: str, auth_asym_id: str) -> Optional[str]:
158
- """Return chain composition (monomeric/homomeric/heteromeric) or None if not found."""
159
- row = get_node_row(pdb_id, auth_asym_id)
160
- if row is None:
161
- return None
162
- return row.get("chain_composition") or None
163
-
164
-
165
- def get_cath_id(pdb_id: str, auth_asym_id: str) -> Optional[str]:
166
- """Return CATH_ID or 'uncategorized' if not found."""
167
- row = get_node_row(pdb_id, auth_asym_id)
168
- if row is None:
169
- return "uncategorized"
170
- val = row.get("CATH_ID", "").strip()
171
- return val if val else "uncategorized"
172
-
173
-
174
- def get_cath_superfamily(pdb_id: str, auth_asym_id: str) -> Optional[str]:
175
- """Return CATH superfamily code (e.g. '3.40.190.10') or None if not found."""
176
- row = get_node_row(pdb_id, auth_asym_id)
177
- if row is None:
178
- return None
179
- val = row.get("cath_superfamily", "").strip()
180
- return val if val else None
181
-
182
-
183
- def get_sequence(pdb_id: str, auth_asym_id: str) -> Tuple[Optional[str], Optional[int]]:
184
- """Return (seq_can, sequence_length) from node table, or (None, None) if not found."""
185
- row = get_node_row(pdb_id, auth_asym_id)
186
- if row is None:
187
- return None, None
188
- seq = row["sequence"] or None
189
- try:
190
- length = int(row["sequence_length"]) if row["sequence_length"].strip() else None
191
- except ValueError:
192
- length = None
193
- return seq, length
194
-
195
-
196
- def get_functions(pdb_id: str, auth_asym_id: str) -> Optional[List[str]]:
197
- """Return ranked function list for a chain, or None if not in node table."""
198
- row = get_node_row(pdb_id, auth_asym_id)
199
- if row is None:
200
- return None
201
- return [f for f in _parse_functions(row["ranked_functions"]) if len(f) > 10]
202
-
203
-
204
- def get_state_id(pdb_id: str, auth_asym_id: str) -> Optional[str]:
205
- """Return state_id for a chain, or None if not found."""
206
- row = get_node_row(pdb_id, auth_asym_id)
207
- if row is None:
208
- return None
209
- val = row.get("state_id", "").strip()
210
- return val if val else None
 
1
+ """
2
+ Node table lookup for chain annotations.
3
+
4
+ Loads the node table from MuSProt.db once on first use and indexes rows
5
+ by (pdb_id_lower, auth_asym_id_upper). When multiple rows share the same
6
+ key the first occurrence is kept.
7
+ """
8
+ import ast
9
+ import sqlite3
10
+ import logging
11
+ import zlib
12
+ from functools import lru_cache
13
+ from typing import Dict, List, Optional, Tuple, Any
14
+
15
+ from app.protein.config import get_database_path, get_node_database_path
16
+ from app.protein.database import connect_readonly
17
+
18
+ logger = logging.getLogger(__name__)
19
+
20
+ _Row = Dict[str, str]
21
+
22
+ # (pdb_id_lower, auth_asym_id_upper) -> relevant fields
23
+ _index: Optional[Dict[Tuple[str, str], _Row]] = None
24
+
25
+
26
+ def _load() -> Dict[Tuple[str, str], _Row]:
27
+ index: Dict[Tuple[str, str], _Row] = {}
28
+ db_path = get_database_path()
29
+ conn = connect_readonly(db_path)
30
+ try:
31
+ conn.row_factory = sqlite3.Row
32
+ cur = conn.cursor()
33
+ cur.execute(
34
+ 'SELECT pdb_id, auth_asym_id, binding_status, sequence, sequence_length,'
35
+ ' cath_id, cath_superfamily, Rosetta, FoldX, EvoEF2, RW, "RW+", ranked_functions,'
36
+ ' state_id, experimental_method, pH, temp_K, chain_composition'
37
+ ' FROM node'
38
+ )
39
+ for row in cur:
40
+ pdb_id = row["pdb_id"] or ""
41
+ auth_asym_id = row["auth_asym_id"] or ""
42
+ key = (pdb_id.lower(), auth_asym_id.upper())
43
+ if key not in index:
44
+ index[key] = {
45
+ "binding_status": row["binding_status"] or "",
46
+ "sequence": row["sequence"] or "",
47
+ "sequence_length": row["sequence_length"] or "",
48
+ "cath_id": row["cath_id"] or "",
49
+ "cath_superfamily": row["cath_superfamily"] or "",
50
+ "Rosetta": row["Rosetta"] or "",
51
+ "FoldX": row["FoldX"] or "",
52
+ "EvoEF2": row["EvoEF2"] or "",
53
+ "RW": row["RW"] or "",
54
+ "RW+": row["RW+"] or "",
55
+ "ranked_functions": row["ranked_functions"] or "",
56
+ "state_id": row["state_id"] or "",
57
+ "experimental_method": row["experimental_method"] or "",
58
+ "pH": row["pH"] or "",
59
+ "temp_K": row["temp_K"] or "",
60
+ "chain_composition": row["chain_composition"] or "",
61
+ }
62
+ finally:
63
+ conn.close()
64
+
65
+ logger.info(f"Node index loaded: {len(index):,} entries from {db_path}")
66
+ return index
67
+
68
+
69
+ def _get_index() -> Dict[Tuple[str, str], _Row]:
70
+ global _index
71
+ if _index is None:
72
+ _index = _load()
73
+ return _index
74
+
75
+
76
+ def reset_index() -> None:
77
+ """Force reload of node index."""
78
+ global _index
79
+ _index = None
80
+ _lookup_node_row.cache_clear()
81
+
82
+
83
+ @lru_cache(maxsize=4096)
84
+ def _lookup_node_row(pdb_id: str, auth_asym_id: str) -> Optional[_Row]:
85
+ """Read one node row without blocking on a full-table in-memory index."""
86
+ conn = connect_readonly(get_node_database_path())
87
+ try:
88
+ conn.row_factory = sqlite3.Row
89
+ row = conn.execute(
90
+ 'SELECT binding_status, sequence, sequence_length, cath_id, cath_superfamily,'
91
+ ' Rosetta, FoldX, EvoEF2, RW, "RW+", ranked_functions,'
92
+ ' state_id, experimental_method, pH, temp_K, chain_composition'
93
+ ' FROM node WHERE LOWER(pdb_id) = ? AND UPPER(auth_asym_id) = ? LIMIT 1',
94
+ (pdb_id.lower(), auth_asym_id.upper()),
95
+ ).fetchone()
96
+ if row is None:
97
+ return None
98
+ return {
99
+ key: zlib.decompress(value).decode("utf-8") if isinstance(value, bytes) else value or ""
100
+ for key, value in ((key, row[key]) for key in row.keys())
101
+ }
102
+ finally:
103
+ conn.close()
104
+
105
+
106
+ def _parse_functions(raw: str) -> List[str]:
107
+ raw = raw.strip()
108
+ if not raw:
109
+ return []
110
+ try:
111
+ parsed = ast.literal_eval(raw)
112
+ if isinstance(parsed, list):
113
+ return [str(f).strip() for f in parsed if str(f).strip()]
114
+ except (ValueError, SyntaxError):
115
+ pass
116
+ return [f.strip() for f in raw.split(";") if f.strip()]
117
+
118
+
119
+ def _parse_float(val: str) -> Optional[float]:
120
+ val = str(val).strip()
121
+ if not val:
122
+ return None
123
+ try:
124
+ return float(val)
125
+ except ValueError:
126
+ return None
127
+
128
+
129
+ def get_energy_scores(pdb_id: str, auth_asym_id: str) -> Optional[Dict[str, Any]]:
130
+ """Return energy scores (Rosetta, FoldX, EvoEF2, RW, RW+) or None if not found."""
131
+ row = get_node_row(pdb_id, auth_asym_id)
132
+ if row is None:
133
+ return None
134
+ return {
135
+ "Rosetta": _parse_float(row.get("Rosetta", "")),
136
+ "FoldX": _parse_float(row.get("FoldX", "")),
137
+ "EvoEF2": _parse_float(row.get("EvoEF2", "")),
138
+ "RW": _parse_float(row.get("RW", "")),
139
+ "RW+": _parse_float(row.get("RW+", "")),
140
+ }
141
+
142
+
143
+ def get_node_row(pdb_id: str, auth_asym_id: str) -> Optional[_Row]:
144
+ """Return the cached node row for a chain."""
145
+ key = (pdb_id.lower(), auth_asym_id.upper())
146
+ return _index.get(key) if _index is not None else _lookup_node_row(*key)
147
+
148
+
149
+ def get_binding_status(pdb_id: str, auth_asym_id: str) -> Optional[str]:
150
+ """Return binding status (apo/holo) or None if not found."""
151
+ row = get_node_row(pdb_id, auth_asym_id)
152
+ if row is None:
153
+ return None
154
+ return row.get("binding_status") or None
155
+
156
+
157
+ def get_chain_composition(pdb_id: str, auth_asym_id: str) -> Optional[str]:
158
+ """Return chain composition (monomeric/homomeric/heteromeric) or None if not found."""
159
+ row = get_node_row(pdb_id, auth_asym_id)
160
+ if row is None:
161
+ return None
162
+ return row.get("chain_composition") or None
163
+
164
+
165
+ def get_cath_id(pdb_id: str, auth_asym_id: str) -> Optional[str]:
166
+ """Return cath_id or 'uncategorized' if not found."""
167
+ row = get_node_row(pdb_id, auth_asym_id)
168
+ if row is None:
169
+ return "uncategorized"
170
+ val = row.get("cath_id", "").strip()
171
+ return val if val else "uncategorized"
172
+
173
+
174
+ def get_cath_superfamily(pdb_id: str, auth_asym_id: str) -> Optional[str]:
175
+ """Return CATH superfamily code (e.g. '3.40.190.10') or None if not found."""
176
+ row = get_node_row(pdb_id, auth_asym_id)
177
+ if row is None:
178
+ return None
179
+ val = row.get("cath_superfamily", "").strip()
180
+ return val if val else None
181
+
182
+
183
+ def get_sequence(pdb_id: str, auth_asym_id: str) -> Tuple[Optional[str], Optional[int]]:
184
+ """Return (seq_can, sequence_length) from node table, or (None, None) if not found."""
185
+ row = get_node_row(pdb_id, auth_asym_id)
186
+ if row is None:
187
+ return None, None
188
+ seq = row["sequence"] or None
189
+ try:
190
+ length = int(row["sequence_length"]) if row["sequence_length"].strip() else None
191
+ except ValueError:
192
+ length = None
193
+ return seq, length
194
+
195
+
196
+ def get_functions(pdb_id: str, auth_asym_id: str) -> Optional[List[str]]:
197
+ """Return ranked function list for a chain, or None if not in node table."""
198
+ row = get_node_row(pdb_id, auth_asym_id)
199
+ if row is None:
200
+ return None
201
+ return [f for f in _parse_functions(row["ranked_functions"]) if len(f) > 10]
202
+
203
+
204
+ def get_state_id(pdb_id: str, auth_asym_id: str) -> Optional[str]:
205
+ """Return state_id for a chain, or None if not found."""
206
+ row = get_node_row(pdb_id, auth_asym_id)
207
+ if row is None:
208
+ return None
209
+ val = row.get("state_id", "").strip()
210
+ return val if val else None
backend/assets/MuSProt_documentation.md CHANGED
@@ -1,152 +1,151 @@
1
- # MuSProt Dataset Documentation
2
-
3
- **MuSProt** (Multistate Protein Database) is a million-scale multimodal database for multistate proteins, designed to support programmable protein design and AI model development. It links experimentally observed conformational states of identical protein sequences from the PDB, organizes them into state clusters and transition relationships, and enriches each record with structural similarity, experimental context, state-specific function rankings and transition fidelity labels. Users can search, browse and download MuSProt records to study conformational diversity, state-dependent functions and feasible protein state transitions.
4
-
5
- ---
6
-
7
- ## Download
8
-
9
- The full dataset is distributed as a single SQLite database file (`MuSProt.db`, ~6.3 GB). Download it from the dataset page using the **Download DB** button.
10
-
11
- For detailed protein structures and atom coordinates, users are expected to download [Protein Data Bank (PDB)](https://rcsb.org) and fetch from `mmCIF` or `.pdb` files based on the PDB entry and chain ID.
12
-
13
- ---
14
-
15
- ## Database Schema
16
-
17
- The database contains two tables: **`node`** and **`edge`**.
18
-
19
- ---
20
-
21
- ### Table: `node`
22
-
23
- Each row represents a single protein chain instance (one PDB entry + chain).
24
-
25
- | Column | Type | Description |
26
- |---|---|---|
27
- | `uniprot_id` | TEXT | UniProt accession number |
28
- | `pdb_id` | TEXT | 4-character PDB entry ID (lowercase) |
29
- | `auth_asym_id` | TEXT | Author chain identifier (e.g. `A`) |
30
- | `base_label` | TEXT | Canonical label combining UniProt ID and chain state |
31
- | `sequence` | TEXT | SEQRES amino acid sequence |
32
- | `sequence_length` | INT | Number of residues in the chain |
33
- | `original_metals` | TEXT | Metal ions present in the structure |
34
- | `original_ligands` | TEXT | Small-molecule ligands present in the structure |
35
- | `sequence_id` | TEXT | Internal sequence cluster identifier |
36
- | `state_id` | TEXT | Conformational state cluster this chain is assigned to (`0`, `1`, `2`, …) |
37
- | `CATH_ID` | TEXT | CATH domain assignment |
38
- | `cath_class` | TEXT | CATH class (e.g. `1` = Mainly Alpha) |
39
- | `cath_arch` | TEXT | CATH architecture |
40
- | `cath_topo` | TEXT | CATH topology |
41
- | `cath_homology` | TEXT | CATH homology superfamily |
42
- | `cath_superfamily` | TEXT | Full CATH superfamily code (e.g. `1.10.10.10`) |
43
- | `domain_length` | INT | Length of the matched CATH domain |
44
- | `experimental_method` | TEXT | Structure determination method (e.g. `X-RAY DIFFRACTION`, `ELECTRON MICROSCOPY`, `SOLUTION NMR`) |
45
- | `pH` | FLOAT | pH of the experimental / crystallization condition |
46
- | `temp_K` | FLOAT | Temperature of the experiment, in Kelvin |
47
- | `experimental_details` | TEXT | Free-text crystallization / sample-preparation details |
48
- | `resolution` | FLOAT | Experimental resolution in Å (lower is sharper; empty for methods without a resolution) |
49
- | `Rosetta` | FLOAT | Rosetta total energy score |
50
- | `FoldX` | FLOAT | FoldX total energy score |
51
- | `EvoEF2` | FLOAT | EvoEF2 total energy score |
52
- | `RW` | FLOAT | Random-Walk (RW) energy score |
53
- | `RW+` | FLOAT | Random-Walk+ (RWplus) energy score |
54
- | `ranked_functions` | TEXT | JSON-encoded list of ranked GO/functional annotations |
55
- | `chain_composition` | TEXT | Quaternary composition of the deposited assembly: `monomeric`, `homomeric`, or `heteromeric` |
56
- | `non_protein_polymer_binding` | TEXT | Non-protein polymer bound in the structure (`DNA`, `RNA`, or `DNA/RNA`); empty when none is present |
57
-
58
-
59
- ---
60
-
61
- ### Table: `edge`
62
-
63
- Each row represents a pairwise structural comparison between two chain instances (A and B) that share the same UniProt identity.
64
-
65
- | Column | Type | Description |
66
- |---|---|---|
67
- | `pdb_id_A` | TEXT | PDB ID of chain A |
68
- | `auth_asym_id_A` | TEXT | Chain identifier of chain A |
69
- | `pdb_id_B` | TEXT | PDB ID of chain B |
70
- | `auth_asym_id_B` | TEXT | Chain identifier of chain B |
71
- | `TM1` | FLOAT | TM-score of the alignment (chain A as reference) |
72
- | `RMSD` | FLOAT | Root-mean-square deviation of atoms ) |
73
- | `structure_sim` | FLOAT | Composite structural similarity score |
74
- | `delta_Rosetta` | FLOAT | Rosetta energy difference (B − A) |
75
- | `delta_FoldX` | FLOAT | FoldX energy difference (B − A) |
76
- | `delta_EvoEF2` | FLOAT | EvoEF2 energy difference (B − A) |
77
- | `delta_RW` | FLOAT | Random-Walk (RW) energy difference (B A) |
78
- | `delta_RW+` | FLOAT | Random-Walk+ (RWplus) energy difference (B − A) |
79
- | `state_id_A` | TEXT | Conformational state cluster of chain A |
80
- | `state_id_B` | TEXT | Conformational state cluster of chain B |
81
- | `avg_sim` | TEXT | Average structural similarity within the state cluster (`>0.95` or a numeric value) |
82
- | `state_fidelity` | TEXT | State-level transition fidelity label: `identical`, `high`, `medium`, or `low` |
83
- | `observation_fidelity` | TEXT | Observation-level transition fidelity label: `identical`, `high`, `medium`, or `low` |
84
-
85
- ---
86
-
87
- ## Usage Examples
88
-
89
- ### Python (sqlite3)
90
-
91
- ```python
92
- import sqlite3
93
- import pandas as pd
94
-
95
- conn = sqlite3.connect("MuSProt.db")
96
-
97
- # Load all chains for a UniProt entry
98
- df_nodes = pd.read_sql(
99
- "SELECT * FROM node WHERE uniprot_id = 'P00533'",
100
- conn
101
- )
102
-
103
- # Find all structural neighbours of a given chain
104
- df_edges = pd.read_sql(
105
- "SELECT * FROM edge WHERE pdb_id_A = '1ivo' AND auth_asym_id_A = 'A'",
106
- conn
107
- )
108
-
109
- conn.close()
110
- ```
111
-
112
- ### Filter by structural similarity
113
-
114
- ```python
115
- # Retrieve pairs with high TM-score and low RMSD
116
- df = pd.read_sql("""
117
- SELECT *
118
- FROM edge
119
- WHERE CAST(TM1 AS REAL) > 0.8
120
- AND CAST(RMSD AS REAL) < 2.0
121
- LIMIT 1000
122
- """, conn)
123
- ```
124
-
125
- ### Join nodes and edges
126
-
127
- ```python
128
- # Get full info for both chains in each pair
129
- df = pd.read_sql("""
130
- SELECT
131
- e.pdb_id_A, e.auth_asym_id_A,
132
- e.pdb_id_B, e.auth_asym_id_B,
133
- e.TM1, e.RMSD,
134
- nA.sequence_length AS len_A,
135
- nB.sequence_length AS len_B,
136
- nA.cath_superfamily
137
- FROM edge e
138
- JOIN node nA ON e.pdb_id_A = nA.pdb_id AND e.auth_asym_id_A = nA.auth_asym_id
139
- JOIN node nB ON e.pdb_id_B = nB.pdb_id AND e.auth_asym_id_B = nB.auth_asym_id
140
- WHERE nA.uniprot_id = 'P00533'
141
- LIMIT 500
142
- """, conn)
143
- ```
144
-
145
- ---
146
-
147
- ## Notes
148
-
149
- - All numeric fields (energies, TM-score, RMSD, lengths) are stored as `TEXT`; cast them with `CAST(col AS REAL)` or `CAST(col AS INTEGER)` as needed.
150
- - `ranked_functions` in the `node` table is a JSON string. Parse it with `json.loads()`.
151
- - The `edge` table is directional: (A→B) and (B→A) are separate rows and may differ slightly in TM-score.
152
- - Energy delta values represent B − A; a negative delta means chain B is lower energy than chain A.
 
1
+ # MuSProt Dataset Documentation
2
+
3
+ **MuSProt** (Multistate Protein Database) is a million-scale multimodal database for multistate proteins, designed to support programmable protein design and AI model development. It links experimentally observed conformational states of identical protein sequences from the PDB, organizes them into state clusters and transition relationships, and enriches each record with structural similarity, experimental context, state-specific function rankings and transition fidelity labels. Users can search, browse and download MuSProt records to study conformational diversity, state-dependent functions and feasible protein state transitions.
4
+
5
+ ---
6
+
7
+ ## Download
8
+
9
+ The full dataset is distributed as a single SQLite database file (`MuSProt.db`, ~31 GB). Download it from the dataset page using the **Download DB** button.
10
+
11
+ For detailed protein structures and atom coordinates, users are expected to download [Protein Data Bank (PDB)](https://rcsb.org) and fetch from `mmCIF` or `.pdb` files based on the PDB entry and chain ID.
12
+
13
+ ---
14
+
15
+ ## Database Schema
16
+
17
+ The database contains two tables: **`node`** and **`edge`**.
18
+
19
+ ---
20
+
21
+ ### Table: `node`
22
+
23
+ Each row represents a single protein chain instance (one PDB entry + chain).
24
+
25
+ | Column | Type | Description |
26
+ |---|---|---|
27
+ | `sequence_id` | TEXT | Internal sequence cluster identifier (e.g. `SEQ087470`); shared by chains with the same sequence |
28
+ | `uniprot_id` | TEXT | UniProt accession number |
29
+ | `pdb_id` | TEXT | 4-character PDB entry ID (lowercase) |
30
+ | `auth_asym_id` | TEXT | Author chain identifier (e.g. `A`) |
31
+ | `sequence` | TEXT | SEQRES amino acid sequence |
32
+ | `sequence_length` | INT | Number of residues in the SEQRES sequence |
33
+ | `modeled_sequence` | TEXT | Amino acid sequence actually resolved (modeled) in the structure |
34
+ | `modeled_sequence_length` | INT | Number of modeled residues |
35
+ | `binders` | TEXT | Bound ligand / small-molecule component IDs present in the structure; empty when none |
36
+ | `binding_status` | TEXT | Ligand-binding state of the chain: `apo` or `holo` |
37
+ | `experimental_method` | TEXT | Structure determination method (e.g. `X-RAY DIFFRACTION`, `ELECTRON MICROSCOPY`, `SOLUTION NMR`) |
38
+ | `resolution` | FLOAT | Experimental resolution in Å (lower is sharper; empty for methods without a resolution) |
39
+ | `pH` | FLOAT | pH of the experimental / crystallization condition |
40
+ | `temp_K` | FLOAT | Temperature of the experiment, in Kelvin |
41
+ | `experimental_details` | TEXT | Free-text crystallization / sample-preparation details |
42
+ | `chain_composition` | TEXT | Quaternary composition of the deposited assembly: `monomeric`, `homomeric`, or `heteromeric` |
43
+ | `non_protein_polymer_binding` | TEXT | Non-protein polymer bound in the structure (`DNA`, `RNA`, or `DNA/RNA`); `None` when absent |
44
+ | `initial_release_date` | TEXT | PDB entry initial release date (`YYYY-MM-DD`) |
45
+ | `cath_id` | TEXT | CATH domain assignment; empty when unassigned |
46
+ | `cath_superfamily` | TEXT | Full CATH superfamily code (e.g. `1.10.10.10`) |
47
+ | `state_id` | TEXT | Conformational state cluster this chain is assigned to (`0`, `1`, `2`, …) |
48
+ | `Rosetta` | FLOAT | Rosetta total energy score |
49
+ | `FoldX` | FLOAT | FoldX total energy score |
50
+ | `EvoEF2` | FLOAT | EvoEF2 total energy score |
51
+ | `RW` | FLOAT | Random-Walk (RW) energy score |
52
+ | `RW+` | FLOAT | Random-Walk+ (RWplus) energy score |
53
+ | `ranked_functions` | TEXT | JSON-encoded list of ranked GO/functional annotations |
54
+ | `matching_scores` | TEXT | JSON-encoded list of function-matching scores, aligned with `ranked_functions` |
55
+
56
+
57
+ ---
58
+
59
+ ### Table: `edge`
60
+
61
+ Each row represents a pairwise structural comparison between two chain instances (A and B) that share the same UniProt identity.
62
+
63
+ | Column | Type | Description |
64
+ |---|---|---|
65
+ | `pdb_id_A` | TEXT | PDB ID of chain A |
66
+ | `auth_asym_id_A` | TEXT | Chain identifier of chain A |
67
+ | `pdb_id_B` | TEXT | PDB ID of chain B |
68
+ | `auth_asym_id_B` | TEXT | Chain identifier of chain B |
69
+ | `TM1` | FLOAT | TM-score of the alignment (chain A as reference) |
70
+ | `RMSD` | FLOAT | Root-mean-square deviation of atoms (Å) |
71
+ | `pair_similarity` | FLOAT | Composite structural similarity score for the A–B pair |
72
+ | `delta_Rosetta` | FLOAT | Rosetta energy difference (B A) |
73
+ | `delta_FoldX` | FLOAT | FoldX energy difference (B − A) |
74
+ | `delta_EvoEF2` | FLOAT | EvoEF2 energy difference (B − A) |
75
+ | `delta_RW` | FLOAT | Random-Walk (RW) energy difference (B − A) |
76
+ | `delta_RW+` | FLOAT | Random-Walk+ (RWplus) energy difference (B − A) |
77
+ | `state_id_A` | TEXT | Conformational state cluster of chain A |
78
+ | `state_id_B` | TEXT | Conformational state cluster of chain B |
79
+ | `state_similarity` | TEXT | Average structural similarity within the state cluster |
80
+ | `state_fidelity` | TEXT | State-level transition fidelity label: `identical`, `high`, `medium`, or `low` |
81
+ | `pair_fidelity` | TEXT | Pairwise (observation-level) transition fidelity label: `identical`, `high`, `medium`, or `low` |
82
+ | `sequence_id` | TEXT | Internal sequence cluster identifier shared by chains A and B |
83
+
84
+ ---
85
+
86
+ ## Usage Examples
87
+
88
+ ### Python (sqlite3)
89
+
90
+ ```python
91
+ import sqlite3
92
+ import pandas as pd
93
+
94
+ conn = sqlite3.connect("MuSProt.db")
95
+
96
+ # Load all chains for a UniProt entry
97
+ df_nodes = pd.read_sql(
98
+ "SELECT * FROM node WHERE uniprot_id = 'P00533'",
99
+ conn
100
+ )
101
+
102
+ # Find all structural neighbours of a given chain
103
+ df_edges = pd.read_sql(
104
+ "SELECT * FROM edge WHERE pdb_id_A = '1ivo' AND auth_asym_id_A = 'A'",
105
+ conn
106
+ )
107
+
108
+ conn.close()
109
+ ```
110
+
111
+ ### Filter by structural similarity
112
+
113
+ ```python
114
+ # Retrieve pairs with high TM-score and low RMSD
115
+ df = pd.read_sql("""
116
+ SELECT *
117
+ FROM edge
118
+ WHERE CAST(TM1 AS REAL) > 0.8
119
+ AND CAST(RMSD AS REAL) < 2.0
120
+ LIMIT 1000
121
+ """, conn)
122
+ ```
123
+
124
+ ### Join nodes and edges
125
+
126
+ ```python
127
+ # Get full info for both chains in each pair
128
+ df = pd.read_sql("""
129
+ SELECT
130
+ e.pdb_id_A, e.auth_asym_id_A,
131
+ e.pdb_id_B, e.auth_asym_id_B,
132
+ e.TM1, e.RMSD,
133
+ nA.sequence_length AS len_A,
134
+ nB.sequence_length AS len_B,
135
+ nA.cath_superfamily
136
+ FROM edge e
137
+ JOIN node nA ON e.pdb_id_A = nA.pdb_id AND e.auth_asym_id_A = nA.auth_asym_id
138
+ JOIN node nB ON e.pdb_id_B = nB.pdb_id AND e.auth_asym_id_B = nB.auth_asym_id
139
+ WHERE nA.uniprot_id = 'P00533'
140
+ LIMIT 500
141
+ """, conn)
142
+ ```
143
+
144
+ ---
145
+
146
+ ## Notes
147
+
148
+ - All numeric fields (energies, TM-score, RMSD, lengths) are stored as `TEXT`; cast them with `CAST(col AS REAL)` or `CAST(col AS INTEGER)` as needed.
149
+ - `ranked_functions` in the `node` table is a JSON string. Parse it with `json.loads()`.
150
+ - The `edge` table is directional: (A→B) and (B→A) are separate rows and may differ slightly in TM-score.
151
+ - Energy delta values represent B A; a negative delta means chain B is lower energy than chain A.
 
backend/scripts/generate_node_lookup.py CHANGED
@@ -1,68 +1,68 @@
1
- """Generate an indexed read-only node lookup sidecar from MuSProt.db."""
2
- from __future__ import annotations
3
-
4
- import argparse
5
- import sqlite3
6
- import zlib
7
- from pathlib import Path
8
-
9
-
10
- NODE_COLUMNS = (
11
- "pdb_id", "auth_asym_id", "base_label", "sequence", "sequence_length",
12
- "CATH_ID", "cath_superfamily", "Rosetta", "FoldX", "EvoEF2", "RW", "RW+",
13
- "ranked_functions", "state_id", "experimental_method", "pH", "temp_K",
14
- "chain_composition",
15
- )
16
-
17
-
18
- def main() -> None:
19
- parser = argparse.ArgumentParser()
20
- parser.add_argument("source", type=Path)
21
- parser.add_argument("-o", "--output", type=Path, default=Path("MuSProt-node-lookup.db"))
22
- args = parser.parse_args()
23
-
24
- output = args.output.resolve()
25
- output.unlink(missing_ok=True)
26
-
27
- conn = sqlite3.connect(output)
28
- try:
29
- conn.execute("PRAGMA journal_mode = OFF")
30
- conn.execute("PRAGMA synchronous = OFF")
31
- conn.execute(
32
- 'CREATE TABLE node (pdb_id, auth_asym_id, base_label, sequence BLOB,'
33
- ' sequence_length, CATH_ID, cath_superfamily, Rosetta, FoldX, EvoEF2,'
34
- ' RW, "RW+", ranked_functions BLOB, state_id, experimental_method, pH, temp_K,'
35
- ' chain_composition)'
36
- )
37
- source = sqlite3.connect(f"file:{args.source.resolve()}?mode=ro", uri=True)
38
- try:
39
- select_columns = ", ".join(f'"{column}"' for column in NODE_COLUMNS)
40
- rows = source.execute(f"SELECT {select_columns} FROM node")
41
- placeholders = ", ".join("?" for _ in NODE_COLUMNS)
42
- batch = []
43
- for row in rows:
44
- row = list(row)
45
- for index in (3, 12):
46
- row[index] = zlib.compress((row[index] or "").encode("utf-8"), 1)
47
- batch.append(row)
48
- if len(batch) == 1000:
49
- conn.executemany(f"INSERT INTO node VALUES ({placeholders})", batch)
50
- batch.clear()
51
- if batch:
52
- conn.executemany(f"INSERT INTO node VALUES ({placeholders})", batch)
53
- finally:
54
- source.close()
55
- conn.execute(
56
- "CREATE INDEX idx_node_chain "
57
- "ON node(LOWER(pdb_id), UPPER(auth_asym_id))"
58
- )
59
- conn.execute("ANALYZE")
60
- conn.commit()
61
- finally:
62
- conn.close()
63
-
64
- print(f"Created {output} ({output.stat().st_size:,} bytes)")
65
-
66
-
67
- if __name__ == "__main__":
68
- main()
 
1
+ """Generate an indexed read-only node lookup sidecar from MuSProt.db."""
2
+ from __future__ import annotations
3
+
4
+ import argparse
5
+ import sqlite3
6
+ import zlib
7
+ from pathlib import Path
8
+
9
+
10
+ NODE_COLUMNS = (
11
+ "pdb_id", "auth_asym_id", "binding_status", "sequence", "sequence_length",
12
+ "cath_id", "cath_superfamily", "Rosetta", "FoldX", "EvoEF2", "RW", "RW+",
13
+ "ranked_functions", "state_id", "experimental_method", "pH", "temp_K",
14
+ "chain_composition",
15
+ )
16
+
17
+
18
+ def main() -> None:
19
+ parser = argparse.ArgumentParser()
20
+ parser.add_argument("source", type=Path)
21
+ parser.add_argument("-o", "--output", type=Path, default=Path("MuSProt-node-lookup.db"))
22
+ args = parser.parse_args()
23
+
24
+ output = args.output.resolve()
25
+ output.unlink(missing_ok=True)
26
+
27
+ conn = sqlite3.connect(output)
28
+ try:
29
+ conn.execute("PRAGMA journal_mode = OFF")
30
+ conn.execute("PRAGMA synchronous = OFF")
31
+ conn.execute(
32
+ 'CREATE TABLE node (pdb_id, auth_asym_id, binding_status, sequence BLOB,'
33
+ ' sequence_length, cath_id, cath_superfamily, Rosetta, FoldX, EvoEF2,'
34
+ ' RW, "RW+", ranked_functions BLOB, state_id, experimental_method, pH, temp_K,'
35
+ ' chain_composition)'
36
+ )
37
+ source = sqlite3.connect(f"file:{args.source.resolve()}?mode=ro", uri=True)
38
+ try:
39
+ select_columns = ", ".join(f'"{column}"' for column in NODE_COLUMNS)
40
+ rows = source.execute(f"SELECT {select_columns} FROM node")
41
+ placeholders = ", ".join("?" for _ in NODE_COLUMNS)
42
+ batch = []
43
+ for row in rows:
44
+ row = list(row)
45
+ for index in (3, 12):
46
+ row[index] = zlib.compress((row[index] or "").encode("utf-8"), 1)
47
+ batch.append(row)
48
+ if len(batch) == 1000:
49
+ conn.executemany(f"INSERT INTO node VALUES ({placeholders})", batch)
50
+ batch.clear()
51
+ if batch:
52
+ conn.executemany(f"INSERT INTO node VALUES ({placeholders})", batch)
53
+ finally:
54
+ source.close()
55
+ conn.execute(
56
+ "CREATE INDEX idx_node_chain "
57
+ "ON node(LOWER(pdb_id), UPPER(auth_asym_id))"
58
+ )
59
+ conn.execute("ANALYZE")
60
+ conn.commit()
61
+ finally:
62
+ conn.close()
63
+
64
+ print(f"Created {output} ({output.stat().st_size:,} bytes)")
65
+
66
+
67
+ if __name__ == "__main__":
68
+ main()
frontend/src/components/datasets/protein/StatsPanel.tsx CHANGED
@@ -5,32 +5,32 @@ export const StatsPanel: React.FC = () => {
5
  return (
6
  <div className="stats-panel">
7
  <div className="stat-card">
8
- <div className="stat-value">64,812</div>
9
  <div className="stat-label">Sequences with multiple conformational observations</div>
10
  </div>
11
 
12
  <div className="stat-card">
13
- <div className="stat-value">11,946</div>
14
  <div className="stat-label">Sequences with multiple state clusters</div>
15
  </div>
16
 
17
  <div className="stat-card">
18
- <div className="stat-value">355,731</div>
19
  <div className="stat-label">Conformational observations</div>
20
  </div>
21
 
22
  <div className="stat-card">
23
- <div className="stat-value">90,745</div>
24
  <div className="stat-label">Unique state clusters</div>
25
  </div>
26
 
27
  <div className="stat-card">
28
- <div className="stat-value">27.6 million</div>
29
  <div className="stat-label">Transition triplets between conformational observations</div>
30
  </div>
31
 
32
  <div className="stat-card">
33
- <div className="stat-value">5.7 million</div>
34
  <div className="stat-label">Transition triplets between distinct states</div>
35
  </div>
36
  </div>
 
5
  return (
6
  <div className="stats-panel">
7
  <div className="stat-card">
8
+ <div className="stat-value">102,790</div>
9
  <div className="stat-label">Sequences with multiple conformational observations</div>
10
  </div>
11
 
12
  <div className="stat-card">
13
+ <div className="stat-value">24,796</div>
14
  <div className="stat-label">Sequences with multiple state clusters</div>
15
  </div>
16
 
17
  <div className="stat-card">
18
+ <div className="stat-value">863,643</div>
19
  <div className="stat-label">Conformational observations</div>
20
  </div>
21
 
22
  <div className="stat-card">
23
+ <div className="stat-value">163,869</div>
24
  <div className="stat-label">Unique state clusters</div>
25
  </div>
26
 
27
  <div className="stat-card">
28
+ <div className="stat-value">113.6 million</div>
29
  <div className="stat-label">Transition triplets between conformational observations</div>
30
  </div>
31
 
32
  <div className="stat-card">
33
+ <div className="stat-value">24.95 million</div>
34
  <div className="stat-label">Transition triplets between distinct states</div>
35
  </div>
36
  </div>