mohnishi commited on
Commit
048e5cf
Β·
1 Parent(s): 3518200

brush up app.py

Browse files
Files changed (1) hide show
  1. app.py +584 -423
app.py CHANGED
@@ -1,5 +1,6 @@
1
  import json
2
  import os
 
3
  import urllib.parse
4
  from datetime import datetime, timezone
5
  from typing import Any
@@ -41,7 +42,7 @@ DATASET_FILES = {
41
  "non_ladder": "corresponding_non_ladder_polymers.csv",
42
  }
43
 
44
- # Columns that are known numeric property columns (for filtering/sorting hints)
45
  NUMERIC_PROPERTY_COLUMNS = [
46
  "density", "thermal_conductivity", "thermal_diffusivity",
47
  "tg", "refractive_index", "static_dielectric_const", "dielectric_const_dc",
@@ -54,9 +55,17 @@ NUMERIC_PROPERTY_COLUMNS = [
54
  "nematic_order_parameter",
55
  ]
56
 
 
 
 
 
 
 
 
 
57
 
58
  # ---------------------------------------------------------------------
59
- # Data
60
  # ---------------------------------------------------------------------
61
 
62
  _datasets: dict[str, pd.DataFrame] = {}
@@ -86,7 +95,7 @@ def preload_data() -> None:
86
  df = _load_csv(fname)
87
  if df is not None:
88
  df["_source"] = key
89
- # Coerce known numeric columns to numeric dtype for reliable filtering
90
  for col in NUMERIC_PROPERTY_COLUMNS:
91
  if col in df.columns:
92
  df[col] = pd.to_numeric(df[col], errors="coerce")
@@ -121,56 +130,28 @@ def get_chi_df(solvent: str) -> pd.DataFrame | None:
121
 
122
 
123
  # ---------------------------------------------------------------------
124
- # Utility
125
  # ---------------------------------------------------------------------
126
 
127
- def safe_jsonrpc_result(req_id: Any, result: Any) -> JSONResponse:
128
- return JSONResponse({
129
- "jsonrpc": "2.0",
130
- "id": req_id,
131
- "result": result,
132
- })
133
-
134
-
135
- def safe_jsonrpc_error(req_id: Any, code: int, message: str) -> JSONResponse:
136
- return JSONResponse({
137
- "jsonrpc": "2.0",
138
- "id": req_id,
139
- "error": {
140
- "code": code,
141
- "message": message,
142
- },
143
- })
144
-
145
-
146
  def _safe_value(v: Any) -> Any:
147
- """Convert non-JSON-serialisable values (NaN, inf, numpy types) to Python-native."""
148
  if isinstance(v, float) and (np.isnan(v) or np.isinf(v)):
149
  return None
150
- if isinstance(v, (np.integer,)):
151
  return int(v)
152
- if isinstance(v, (np.floating,)):
153
  return float(v)
154
  return v
155
 
156
 
157
  def records_to_json(df: pd.DataFrame, fields: list[str] | None = None) -> list[dict[str, Any]]:
158
- """Convert DataFrame rows to clean JSON-serialisable dicts.
159
-
160
- Parameters
161
- ----------
162
- df:
163
- Source DataFrame (already sliced to the desired rows).
164
- fields:
165
- If provided, only include these columns in the output.
166
- """
167
  if df.empty:
168
  return []
169
  if fields:
170
  existing = [f for f in fields if f in df.columns]
171
  df = df[existing]
172
- raw = df.to_dict(orient="records")
173
- return [{k: _safe_value(v) for k, v in row.items()} for row in raw]
174
 
175
 
176
  def stringify_records(records: list[dict[str, Any]]) -> str:
@@ -179,25 +160,30 @@ def stringify_records(records: list[dict[str, Any]]) -> str:
179
  return json.dumps(records, ensure_ascii=False, indent=2, default=str)
180
 
181
 
 
 
 
 
 
 
 
 
182
  # ---------------------------------------------------------------------
183
- # Numeric filter parser
184
  #
185
- # Accepts a list of filter strings such as:
186
- # ["density>1.0", "density<=1.5", "tg>=300", "thermal_conductivity!="]
187
- #
188
- # Supported operators: >, >=, <, <=, ==, !=
189
- # Special case: "col!=" means "column is not null/empty"
190
  # ---------------------------------------------------------------------
191
 
192
  _OPERATORS = [">=", "<=", "!=", ">", "<", "=="]
193
 
194
 
195
- def _apply_numeric_filters(df: pd.DataFrame, filters: list[str]) -> tuple[pd.DataFrame, list[str]]:
196
- """Apply numeric range filters to a DataFrame.
197
-
198
- Returns the filtered DataFrame and a list of warning messages for
199
- unrecognised or unapplicable filters.
200
- """
201
  warnings: list[str] = []
202
  for f in filters:
203
  f = f.strip()
@@ -214,13 +200,11 @@ def _apply_numeric_filters(df: pd.DataFrame, filters: list[str]) -> tuple[pd.Dat
214
  parsed = True
215
  break
216
 
217
- # Coerce column to numeric if not already
218
  if df[col].dtype == object:
219
  df = df.copy()
220
  df[col] = pd.to_numeric(df[col], errors="coerce")
221
 
222
  if op == "!=" and val_str == "":
223
- # Special: "col!=" means "has a value"
224
  df = df[df[col].notna()]
225
  parsed = True
226
  break
@@ -232,19 +216,15 @@ def _apply_numeric_filters(df: pd.DataFrame, filters: list[str]) -> tuple[pd.Dat
232
  parsed = True
233
  break
234
 
235
- if op == ">":
236
- df = df[df[col] > val]
237
- elif op == ">=":
238
- df = df[df[col] >= val]
239
- elif op == "<":
240
- df = df[df[col] < val]
241
- elif op == "<=":
242
- df = df[df[col] <= val]
243
- elif op == "==":
244
- df = df[df[col] == val]
245
- elif op == "!=":
246
- df = df[df[col] != val]
247
-
248
  parsed = True
249
  break
250
 
@@ -255,21 +235,26 @@ def _apply_numeric_filters(df: pd.DataFrame, filters: list[str]) -> tuple[pd.Dat
255
 
256
 
257
  # ---------------------------------------------------------------------
258
- # Tool implementations
259
  # ---------------------------------------------------------------------
260
 
261
  def tool_list_datasets() -> dict[str, Any]:
262
- return {
263
- "datasets": {
264
- name: {
265
- "rows": int(len(df)),
266
- "columns": list(df.columns),
 
267
  }
268
- for name, df in _datasets.items()
269
  }
270
- }
 
271
 
272
 
 
 
 
 
273
  def tool_search_polymers(
274
  query: str = "",
275
  dataset: str = "general",
@@ -282,243 +267,346 @@ def tool_search_polymers(
282
  ) -> dict[str, Any]:
283
  """Search and filter polymers.
284
 
285
- Parameters
286
- ----------
287
- query:
288
- Full-text search string across all text columns. Empty string returns all rows
289
- (subject to other filters).
290
- dataset:
291
- Target dataset name.
292
- limit:
293
- Maximum number of records to return (default 10, max 200).
294
- filters:
295
- List of numeric filter expressions, e.g. ["density>1.0", "tg>=300"].
296
- Supported operators: >, >=, <, <=, ==, !=
297
- Use "col!=" to require a non-null value.
298
- sort_by:
299
- Column name to sort results by.
300
- sort_ascending:
301
- Sort direction (default True = ascending).
302
- fields:
303
- List of column names to include in the output. If None, a sensible
304
- default set of columns is returned to keep responses compact.
305
- require_numeric:
306
- Convenience shorthand: list of column names that must be non-null numeric.
307
- Equivalent to adding "col!=" entries to filters.
308
  """
309
- if dataset not in _datasets:
310
- return {"error": f"Unknown dataset: {dataset}"}
311
-
312
- limit = min(int(limit), 200)
313
- df = _datasets[dataset].copy()
314
-
315
- # --- full-text search ---
316
- q = (query or "").strip()
317
- if q:
318
- q_lower = q.lower()
319
- mask = pd.Series(False, index=df.index)
320
- for col in df.columns:
321
- try:
322
- if df[col].dtype == object:
323
- mask = mask | df[col].astype(str).str.lower().str.contains(q_lower, na=False)
324
- except Exception:
325
- continue
326
- df = df[mask]
327
 
328
- # --- require_numeric shorthand ---
329
- if require_numeric:
330
- extra = [f"{col}!=" for col in require_numeric]
331
- filters = list(filters or []) + extra
332
 
333
- # --- numeric filters ---
334
- warnings: list[str] = []
335
- if filters:
336
- df, warnings = _apply_numeric_filters(df, filters)
337
-
338
- total_after_filter = len(df)
339
-
340
- # --- sort ---
341
- if sort_by and sort_by in df.columns:
342
- df = df.sort_values(by=sort_by, ascending=sort_ascending, na_position="last")
343
-
344
- # --- field selection ---
345
- # Default: return a compact set of useful columns
346
- if fields is None:
347
- default_fields = [
348
- "UUID", "smiles_list", "polymer_class", "_source",
349
- "density", "thermal_conductivity", "thermal_diffusivity",
350
- "tg", "refractive_index", "static_dielectric_const",
351
- "bulk_modulus", "sp_total", "abbe_number_sos",
352
- ]
353
- fields = [f for f in default_fields if f in df.columns]
354
 
355
- records = records_to_json(df.head(limit), fields=fields)
 
 
356
 
357
- return {
358
- "dataset": dataset,
359
- "query": query,
360
- "filters": filters or [],
361
- "total_matched": total_after_filter,
362
- "returned": len(records),
363
- "sort_by": sort_by,
364
- "warnings": warnings,
365
- "records": records,
366
- }
367
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
368
 
369
  def tool_get_dataset_columns(dataset: str = "general") -> dict[str, Any]:
370
- if dataset not in _datasets:
371
- return {"error": f"Unknown dataset: {dataset}"}
372
- df = _datasets[dataset]
373
- numeric_cols = [c for c in df.columns if pd.api.types.is_numeric_dtype(df[c])]
374
- return {
375
- "dataset": dataset,
376
- "num_columns": int(len(df.columns)),
377
- "columns": list(df.columns),
378
- "numeric_columns": numeric_cols,
379
- "known_property_columns": [c for c in NUMERIC_PROPERTY_COLUMNS if c in df.columns],
380
- }
 
 
 
 
381
 
 
 
 
382
 
383
  def tool_get_statistics(
384
  dataset: str = "general",
385
  columns: list[str] | None = None,
386
  filters: list[str] | None = None,
387
  ) -> dict[str, Any]:
388
- """Return descriptive statistics (count, mean, std, min, max, median) for
389
- numeric property columns in a dataset, optionally after applying filters.
390
-
391
- Parameters
392
- ----------
393
- dataset:
394
- Target dataset.
395
- columns:
396
- Specific columns to summarise. Defaults to all known property columns
397
- that are present and non-empty in the dataset.
398
- filters:
399
- Numeric filter expressions applied before computing statistics.
400
  """
401
- if dataset not in _datasets:
402
- return {"error": f"Unknown dataset: {dataset}"}
 
403
 
404
- df = _datasets[dataset].copy()
405
 
406
- if filters:
407
- df, warnings = _apply_numeric_filters(df, filters)
408
- else:
409
- warnings = []
410
 
411
- if columns is None:
412
- columns = [c for c in NUMERIC_PROPERTY_COLUMNS if c in df.columns]
413
 
414
- stats: dict[str, Any] = {}
415
- for col in columns:
416
- if col not in df.columns:
417
- stats[col] = {"error": "column not found"}
418
- continue
419
- s = pd.to_numeric(df[col], errors="coerce").dropna()
420
- if s.empty:
421
- stats[col] = {"count": 0, "note": "no numeric data"}
422
- continue
423
- stats[col] = {
424
- "count": int(s.count()),
425
- "mean": round(float(s.mean()), 6),
426
- "std": round(float(s.std()), 6),
427
- "min": round(float(s.min()), 6),
428
- "p25": round(float(s.quantile(0.25)), 6),
429
- "median": round(float(s.median()), 6),
430
- "p75": round(float(s.quantile(0.75)), 6),
431
- "max": round(float(s.max()), 6),
 
 
 
 
 
 
 
 
 
432
  }
 
 
433
 
434
- return {
435
- "dataset": dataset,
436
- "filters": filters or [],
437
- "warnings": warnings,
438
- "rows_after_filter": len(df),
439
- "statistics": stats,
440
- }
441
 
 
 
 
442
 
443
  def tool_get_correlation(
444
  dataset: str = "general",
445
  col_x: str = "density",
446
  col_y: str = "thermal_conductivity",
447
  filters: list[str] | None = None,
448
- limit: int = 200,
449
  ) -> dict[str, Any]:
450
- """Return Pearson and Spearman correlation between two numeric columns,
451
- plus sample data points for scatter plotting.
452
-
453
- Parameters
454
- ----------
455
- dataset:
456
- Target dataset.
457
- col_x:
458
- First numeric column.
459
- col_y:
460
- Second numeric column.
461
- filters:
462
- Numeric filter expressions applied before computing the correlation.
463
- limit:
464
- Maximum number of data points to return for plotting (default 200).
465
  """
466
- if dataset not in _datasets:
467
- return {"error": f"Unknown dataset: {dataset}"}
 
468
 
469
- df = _datasets[dataset].copy()
470
 
471
- if filters:
472
- df, warnings = _apply_numeric_filters(df, filters)
473
- else:
474
- warnings = []
 
 
 
 
 
 
 
475
 
476
- for col in [col_x, col_y]:
477
- if col not in df.columns:
478
- return {"error": f"Column '{col}' not found in dataset '{dataset}'."}
479
- df[col] = pd.to_numeric(df[col], errors="coerce")
 
480
 
481
- pair = df[[col_x, col_y]].dropna()
482
- n = len(pair)
 
 
 
 
 
 
 
483
 
484
- if n < 2:
485
  return {
486
  "dataset": dataset,
487
  "col_x": col_x,
488
  "col_y": col_y,
489
- "n": n,
 
 
 
490
  "warnings": warnings,
491
- "error": "Not enough data points after filtering.",
 
492
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
493
 
494
- pearson = float(pair[col_x].corr(pair[col_y], method="pearson"))
495
- spearman = float(pair[col_x].corr(pair[col_y], method="spearman"))
496
 
497
- sample = pair.sample(n=min(limit, n), random_state=42).sort_values(col_x)
498
- points = [
499
- {col_x: _safe_value(row[col_x]), col_y: _safe_value(row[col_y])}
500
- for _, row in sample.iterrows()
501
- ]
502
 
503
- return {
504
- "dataset": dataset,
505
- "col_x": col_x,
506
- "col_y": col_y,
507
- "n": n,
508
- "pearson_r": round(pearson, 4),
509
- "spearman_r": round(spearman, 4),
510
- "filters": filters or [],
511
- "warnings": warnings,
512
- "sample_points": points,
513
- }
514
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
515
 
516
  def tool_get_chi_solvents() -> dict[str, Any]:
517
- loaded = [s for s in SOLVENTS if get_chi_df(s) is not None]
518
- return {
519
- "count": len(loaded),
520
- "solvents": loaded,
521
- }
522
 
523
 
524
  def tool_search_chi(
@@ -527,58 +615,58 @@ def tool_search_chi(
527
  limit: int = 10,
528
  fields: list[str] | None = None,
529
  ) -> dict[str, Any]:
530
- df = get_chi_df(solvent)
531
- if df is None:
532
- return {"error": f"Chi dataset not found for solvent: {solvent}"}
533
-
534
- if not query:
535
- result_df = df
536
- else:
537
- q_lower = query.lower()
538
- mask = pd.Series(False, index=df.index)
539
- for col in df.columns:
540
- try:
541
- if df[col].dtype == object:
542
- mask = mask | df[col].astype(str).str.lower().str.contains(q_lower, na=False)
543
- except Exception:
544
- continue
545
- result_df = df[mask]
 
546
 
547
- records = records_to_json(result_df.head(limit), fields=fields)
548
- return {
549
- "solvent": solvent,
550
- "query": query,
551
- "total_matched": len(result_df),
552
- "returned": len(records),
553
- "records": records,
554
- }
 
 
555
 
556
 
557
  # ---------------------------------------------------------------------
558
- # MCP
559
  # ---------------------------------------------------------------------
560
 
561
  SERVER_INFO = {
562
  "name": "polyomics-mcp-server",
563
- "version": "0.2.0",
564
  }
565
 
566
  TOOLS = [
567
  {
568
  "name": "list_datasets",
569
- "description": "List available polymer datasets and basic metadata.",
570
- "inputSchema": {
571
- "type": "object",
572
- "properties": {},
573
- "additionalProperties": False,
574
- },
575
  },
576
  {
577
  "name": "search_polymers",
578
  "description": (
579
- "Search and filter polymers in a dataset. Supports full-text search across "
580
- "text columns and numeric range filters (e.g. density>1.0). "
581
- "Results can be sorted by a numeric column and limited to specific fields."
 
582
  ),
583
  "inputSchema": {
584
  "type": "object",
@@ -589,13 +677,17 @@ TOOLS = [
589
  "default": "",
590
  },
591
  "dataset": {"type": "string", "default": "general"},
592
- "limit": {"type": "integer", "default": 10, "description": "Max rows to return (max 200)."},
 
 
 
 
593
  "filters": {
594
  "type": "array",
595
  "items": {"type": "string"},
596
  "description": (
597
- "Numeric filter expressions, e.g. [\"density>1.0\", \"tg>=300\", "
598
- "\"thermal_conductivity!=\"]. "
599
  "Operators: >, >=, <, <=, ==, !=. "
600
  "Use 'col!=' to require a non-null value."
601
  ),
@@ -606,12 +698,12 @@ TOOLS = [
606
  "fields": {
607
  "type": "array",
608
  "items": {"type": "string"},
609
- "description": "Column names to include in output. Defaults to a compact property set.",
610
  },
611
  "require_numeric": {
612
  "type": "array",
613
  "items": {"type": "string"},
614
- "description": "Columns that must have a non-null numeric value.",
615
  },
616
  },
617
  "additionalProperties": False,
@@ -620,22 +712,22 @@ TOOLS = [
620
  {
621
  "name": "get_dataset_columns",
622
  "description": (
623
- "Return all columns of a dataset, plus lists of numeric columns "
624
  "and known property columns useful for filtering."
625
  ),
626
  "inputSchema": {
627
  "type": "object",
628
- "properties": {
629
- "dataset": {"type": "string", "default": "general"},
630
- },
631
  "additionalProperties": False,
632
  },
633
  },
634
  {
635
  "name": "get_statistics",
636
  "description": (
637
- "Return descriptive statistics (count, mean, std, min, quartiles, max) "
638
- "for numeric property columns, optionally after applying numeric filters."
 
 
639
  ),
640
  "inputSchema": {
641
  "type": "object",
@@ -649,7 +741,7 @@ TOOLS = [
649
  "filters": {
650
  "type": "array",
651
  "items": {"type": "string"},
652
- "description": "Numeric filter expressions applied before computing statistics.",
653
  "default": [],
654
  },
655
  },
@@ -660,7 +752,8 @@ TOOLS = [
660
  "name": "get_correlation",
661
  "description": (
662
  "Compute Pearson and Spearman correlation between two numeric columns "
663
- "and return sample data points for scatter-plot visualisation."
 
664
  ),
665
  "inputSchema": {
666
  "type": "object",
@@ -674,10 +767,10 @@ TOOLS = [
674
  "description": "Numeric filters applied before computing the correlation.",
675
  "default": [],
676
  },
677
- "limit": {
678
  "type": "integer",
679
- "default": 200,
680
- "description": "Max sample points to return for plotting.",
681
  },
682
  },
683
  "required": ["col_x", "col_y"],
@@ -685,14 +778,78 @@ TOOLS = [
685
  },
686
  },
687
  {
688
- "name": "get_chi_solvents",
689
- "description": "List available solvents for chi parameter tables.",
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
690
  "inputSchema": {
691
  "type": "object",
692
- "properties": {},
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
693
  "additionalProperties": False,
694
  },
695
  },
 
 
 
 
 
696
  {
697
  "name": "search_chi",
698
  "description": "Search a chi parameter table for a given solvent.",
@@ -715,55 +872,81 @@ TOOLS = [
715
  ]
716
 
717
 
718
- def handle_tool_call(name: str, arguments: dict[str, Any]) -> dict[str, Any]:
719
- if name == "list_datasets":
720
- return tool_list_datasets()
721
-
722
- if name == "search_polymers":
723
- return tool_search_polymers(
724
- query=arguments.get("query", ""),
725
- dataset=arguments.get("dataset", "general"),
726
- limit=int(arguments.get("limit", 10)),
727
- filters=arguments.get("filters") or [],
728
- sort_by=arguments.get("sort_by"),
729
- sort_ascending=bool(arguments.get("sort_ascending", True)),
730
- fields=arguments.get("fields"),
731
- require_numeric=arguments.get("require_numeric"),
732
- )
733
-
734
- if name == "get_dataset_columns":
735
- return tool_get_dataset_columns(
736
- dataset=arguments.get("dataset", "general")
737
- )
738
-
739
- if name == "get_statistics":
740
- return tool_get_statistics(
741
- dataset=arguments.get("dataset", "general"),
742
- columns=arguments.get("columns"),
743
- filters=arguments.get("filters") or [],
744
- )
745
-
746
- if name == "get_correlation":
747
- return tool_get_correlation(
748
- dataset=arguments.get("dataset", "general"),
749
- col_x=arguments.get("col_x", "density"),
750
- col_y=arguments.get("col_y", "thermal_conductivity"),
751
- filters=arguments.get("filters") or [],
752
- limit=int(arguments.get("limit", 200)),
753
- )
754
-
755
- if name == "get_chi_solvents":
756
- return tool_get_chi_solvents()
757
 
758
- if name == "search_chi":
759
- return tool_search_chi(
760
- solvent=arguments.get("solvent", ""),
761
- query=arguments.get("query", ""),
762
- limit=int(arguments.get("limit", 10)),
763
- fields=arguments.get("fields"),
764
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
765
 
766
- return {"error": f"Unknown tool: {name}"}
 
 
 
 
767
 
768
 
769
  # ---------------------------------------------------------------------
@@ -772,9 +955,8 @@ def handle_tool_call(name: str, arguments: dict[str, Any]) -> dict[str, Any]:
772
 
773
  async def mcp_sse_get(request: Request) -> Response:
774
  print(">>> GET /mcp/sse")
775
- body = "event: ping\ndata: alive\n\n"
776
  return PlainTextResponse(
777
- body,
778
  media_type="text/event-stream",
779
  headers={
780
  "Cache-Control": "no-cache",
@@ -789,21 +971,20 @@ async def mcp_sse_post(request: Request) -> JSONResponse:
789
  payload = await request.json()
790
  except Exception:
791
  raw = await request.body()
792
- print(">>> POST /mcp/sse invalid json")
793
- print(raw.decode("utf-8", errors="replace")[:1000])
794
  return safe_jsonrpc_error(None, -32700, "Parse error")
795
 
796
  print(">>> POST /mcp/sse")
797
- print(json.dumps(payload, ensure_ascii=False)[:2000])
798
 
799
  req_id = payload.get("id")
800
- method = payload.get("method")
801
- params = payload.get("params", {}) or {}
802
 
803
  if method == "initialize":
804
- protocol_version = params.get("protocolVersion", "2025-03-26")
805
  return safe_jsonrpc_result(req_id, {
806
- "protocolVersion": protocol_version,
807
  "capabilities": {"tools": {}},
808
  "serverInfo": SERVER_INFO,
809
  })
@@ -818,25 +999,20 @@ async def mcp_sse_post(request: Request) -> JSONResponse:
818
  return safe_jsonrpc_result(req_id, {"tools": TOOLS})
819
 
820
  if method == "tools/call":
821
- name = params.get("name")
822
- arguments = params.get("arguments", {}) or {}
823
- result = handle_tool_call(name, arguments)
824
-
825
  return safe_jsonrpc_result(req_id, {
826
- "content": [
827
- {
828
- "type": "text",
829
- "text": stringify_records([result]) if isinstance(result, dict) else str(result),
830
- }
831
- ],
832
- "isError": bool(isinstance(result, dict) and "error" in result),
833
  })
834
 
835
  return safe_jsonrpc_error(req_id, -32601, f"Method not found: {method}")
836
 
837
 
838
  # ---------------------------------------------------------------------
839
- # Auxiliary endpoints
840
  # ---------------------------------------------------------------------
841
 
842
  async def health(request: Request) -> JSONResponse:
@@ -851,19 +1027,16 @@ async def health(request: Request) -> JSONResponse:
851
 
852
 
853
  async def oauth_protected_resource(request: Request) -> Response:
854
- print(f">>> {request.method} {request.url.path}")
855
  return Response(status_code=204)
856
 
857
 
858
  async def oauth_authorization_server(request: Request) -> Response:
859
- print(f">>> {request.method} {request.url.path}")
860
  return Response(status_code=204)
861
 
862
 
863
  async def register(request: Request) -> Response:
864
  raw = await request.body()
865
- print(">>> POST /register")
866
- print(raw.decode("utf-8", errors="replace")[:1000])
867
  return Response(status_code=204)
868
 
869
 
@@ -879,7 +1052,7 @@ async def options_handler(request: Request) -> Response:
879
 
880
 
881
  # ---------------------------------------------------------------------
882
- # UI
883
  # ---------------------------------------------------------------------
884
 
885
  def build_gradio_app() -> GradioApp:
@@ -888,31 +1061,32 @@ def build_gradio_app() -> GradioApp:
888
  "# PolyOmics MCP Server\n\n"
889
  f"**Dataset repo:** `{DATASET_REPO}`\n\n"
890
  "**Claude connector endpoint:** `https://mohnishi-polyomics-mcp-server.hf.space/mcp/sse`\n\n"
891
- "**Available MCP tools (v0.2.0):**\n"
892
  "- `list_datasets` β€” list datasets and row counts\n"
893
- "- `search_polymers` β€” full-text + numeric filter + sort + field selection\n"
894
  "- `get_dataset_columns` β€” column names, numeric columns, known property columns\n"
895
- "- `get_statistics` β€” descriptive stats per property column\n"
896
- "- `get_correlation` β€” Pearson/Spearman correlation + scatter data\n"
 
 
897
  "- `get_chi_solvents` β€” list loaded chi-parameter solvents\n"
898
  "- `search_chi` β€” search chi parameter tables\n"
899
  )
900
 
901
  df = get_main_df()
902
- summary = {
903
- "dataset_repo": DATASET_REPO,
904
  "server_version": SERVER_INFO["version"],
905
- "main_rows": int(len(df)),
906
- "main_columns": int(len(df.columns)) if not df.empty else 0,
907
- "datasets": list(_datasets.keys()),
908
- }
909
- gr.JSON(summary)
910
 
911
  return GradioApp.create_app(demo, app_kwargs={"docs_url": "/docs"})
912
 
913
 
914
  # ---------------------------------------------------------------------
915
- # App
916
  # ---------------------------------------------------------------------
917
 
918
  def build_app() -> Starlette:
@@ -921,46 +1095,33 @@ def build_app() -> Starlette:
921
  app = Starlette(
922
  routes=[
923
  Route("/health", endpoint=health, methods=["GET"]),
924
- Route("/mcp/sse", endpoint=mcp_sse_get, methods=["GET"]),
925
  Route("/mcp/sse", endpoint=mcp_sse_post, methods=["POST"]),
926
  Route("/mcp/sse", endpoint=options_handler, methods=["OPTIONS"]),
927
- Route(
928
- "/.well-known/oauth-protected-resource",
929
- endpoint=oauth_protected_resource,
930
- methods=["GET"],
931
- ),
932
- Route(
933
- "/.well-known/oauth-protected-resource/mcp/sse",
934
- endpoint=oauth_protected_resource,
935
- methods=["GET"],
936
- ),
937
- Route(
938
- "/.well-known/oauth-authorization-server",
939
- endpoint=oauth_authorization_server,
940
- methods=["GET"],
941
- ),
942
  Route("/register", endpoint=register, methods=["POST"]),
943
  Mount("/", app=gradio_app),
944
  ]
945
  )
946
 
947
- print("MCP JSON-RPC server ready at /mcp/sse")
948
  for r in app.routes:
949
- print(type(r).__name__, getattr(r, "path", None))
950
 
951
  return app
952
 
953
 
954
  # ---------------------------------------------------------------------
955
- # Main
956
  # ---------------------------------------------------------------------
957
 
958
  if __name__ == "__main__":
959
- now = datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
960
  print(f"===== Application Startup at {now} =====")
961
  preload_data()
962
  app = build_app()
963
  port = int(os.environ.get("PORT", 7860))
964
- print(f"Starting uvicorn on port {port} ...")
965
- uvicorn.run(app, host="0.0.0.0", port=port)
966
-
 
1
  import json
2
  import os
3
+ import traceback
4
  import urllib.parse
5
  from datetime import datetime, timezone
6
  from typing import Any
 
42
  "non_ladder": "corresponding_non_ladder_polymers.csv",
43
  }
44
 
45
+ # Known numeric property columns β€” used as defaults for statistics / filtering
46
  NUMERIC_PROPERTY_COLUMNS = [
47
  "density", "thermal_conductivity", "thermal_diffusivity",
48
  "tg", "refractive_index", "static_dielectric_const", "dielectric_const_dc",
 
55
  "nematic_order_parameter",
56
  ]
57
 
58
+ # Default compact field set returned by search_polymers
59
+ _SEARCH_DEFAULT_FIELDS = [
60
+ "UUID", "smiles_list", "polymer_class", "_source",
61
+ "density", "thermal_conductivity", "thermal_diffusivity",
62
+ "tg", "refractive_index", "static_dielectric_const",
63
+ "bulk_modulus", "sp_total", "abbe_number_sos",
64
+ ]
65
+
66
 
67
  # ---------------------------------------------------------------------
68
+ # Data layer
69
  # ---------------------------------------------------------------------
70
 
71
  _datasets: dict[str, pd.DataFrame] = {}
 
95
  df = _load_csv(fname)
96
  if df is not None:
97
  df["_source"] = key
98
+ # Coerce known numeric columns at load time for reliable downstream ops
99
  for col in NUMERIC_PROPERTY_COLUMNS:
100
  if col in df.columns:
101
  df[col] = pd.to_numeric(df[col], errors="coerce")
 
130
 
131
 
132
  # ---------------------------------------------------------------------
133
+ # JSON / serialisation helpers
134
  # ---------------------------------------------------------------------
135
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
136
  def _safe_value(v: Any) -> Any:
137
+ """Convert non-JSON-serialisable values to Python-native types."""
138
  if isinstance(v, float) and (np.isnan(v) or np.isinf(v)):
139
  return None
140
+ if isinstance(v, np.integer):
141
  return int(v)
142
+ if isinstance(v, np.floating):
143
  return float(v)
144
  return v
145
 
146
 
147
  def records_to_json(df: pd.DataFrame, fields: list[str] | None = None) -> list[dict[str, Any]]:
148
+ """Convert a DataFrame to a clean list of JSON-serialisable dicts."""
 
 
 
 
 
 
 
 
149
  if df.empty:
150
  return []
151
  if fields:
152
  existing = [f for f in fields if f in df.columns]
153
  df = df[existing]
154
+ return [{k: _safe_value(v) for k, v in row.items()} for row in df.to_dict(orient="records")]
 
155
 
156
 
157
  def stringify_records(records: list[dict[str, Any]]) -> str:
 
160
  return json.dumps(records, ensure_ascii=False, indent=2, default=str)
161
 
162
 
163
+ def safe_jsonrpc_result(req_id: Any, result: Any) -> JSONResponse:
164
+ return JSONResponse({"jsonrpc": "2.0", "id": req_id, "result": result})
165
+
166
+
167
+ def safe_jsonrpc_error(req_id: Any, code: int, message: str) -> JSONResponse:
168
+ return JSONResponse({"jsonrpc": "2.0", "id": req_id, "error": {"code": code, "message": message}})
169
+
170
+
171
  # ---------------------------------------------------------------------
172
+ # Numeric filter engine
173
  #
174
+ # Syntax:
175
+ # "density>1.0" strict greater-than
176
+ # "density>=1.0" greater-or-equal
177
+ # "thermal_conductivity!=" column must be non-null
 
178
  # ---------------------------------------------------------------------
179
 
180
  _OPERATORS = [">=", "<=", "!=", ">", "<", "=="]
181
 
182
 
183
+ def _apply_numeric_filters(
184
+ df: pd.DataFrame, filters: list[str]
185
+ ) -> tuple[pd.DataFrame, list[str]]:
186
+ """Apply numeric range filters; return (filtered_df, warning_list)."""
 
 
187
  warnings: list[str] = []
188
  for f in filters:
189
  f = f.strip()
 
200
  parsed = True
201
  break
202
 
 
203
  if df[col].dtype == object:
204
  df = df.copy()
205
  df[col] = pd.to_numeric(df[col], errors="coerce")
206
 
207
  if op == "!=" and val_str == "":
 
208
  df = df[df[col].notna()]
209
  parsed = True
210
  break
 
216
  parsed = True
217
  break
218
 
219
+ ops_map = {
220
+ ">": df[col] > val,
221
+ ">=": df[col] >= val,
222
+ "<": df[col] < val,
223
+ "<=": df[col] <= val,
224
+ "==": df[col] == val,
225
+ "!=": df[col] != val,
226
+ }
227
+ df = df[ops_map[op]]
 
 
 
 
228
  parsed = True
229
  break
230
 
 
235
 
236
 
237
  # ---------------------------------------------------------------------
238
+ # Tool: list_datasets
239
  # ---------------------------------------------------------------------
240
 
241
  def tool_list_datasets() -> dict[str, Any]:
242
+ """List all loaded datasets with row counts and column names."""
243
+ try:
244
+ return {
245
+ "datasets": {
246
+ name: {"rows": int(len(df)), "columns": list(df.columns)}
247
+ for name, df in _datasets.items()
248
  }
 
249
  }
250
+ except Exception as e:
251
+ return {"error": str(e)}
252
 
253
 
254
+ # ---------------------------------------------------------------------
255
+ # Tool: search_polymers
256
+ # ---------------------------------------------------------------------
257
+
258
  def tool_search_polymers(
259
  query: str = "",
260
  dataset: str = "general",
 
267
  ) -> dict[str, Any]:
268
  """Search and filter polymers.
269
 
270
+ - limit is capped at 1000; use the aggregation tools for whole-dataset analysis.
271
+ - Empty query with filters returns all rows passing those filters.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
272
  """
273
+ try:
274
+ if dataset not in _datasets:
275
+ return {"error": f"Unknown dataset: '{dataset}'. Available: {list(_datasets.keys())}"}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
276
 
277
+ limit = min(int(limit), 1000)
278
+ df = _datasets[dataset].copy()
 
 
279
 
280
+ q = (query or "").strip()
281
+ if q:
282
+ q_lower = q.lower()
283
+ mask = pd.Series(False, index=df.index)
284
+ for col in df.columns:
285
+ try:
286
+ if df[col].dtype == object:
287
+ mask = mask | df[col].astype(str).str.lower().str.contains(q_lower, na=False)
288
+ except Exception:
289
+ continue
290
+ df = df[mask]
291
+
292
+ if require_numeric:
293
+ filters = list(filters or []) + [f"{col}!=" for col in require_numeric]
 
 
 
 
 
 
 
294
 
295
+ warnings: list[str] = []
296
+ if filters:
297
+ df, warnings = _apply_numeric_filters(df, filters)
298
 
299
+ total_after_filter = len(df)
 
 
 
 
 
 
 
 
 
300
 
301
+ if sort_by and sort_by in df.columns:
302
+ df = df.sort_values(by=sort_by, ascending=sort_ascending, na_position="last")
303
+
304
+ if fields is None:
305
+ fields = [f for f in _SEARCH_DEFAULT_FIELDS if f in df.columns]
306
+
307
+ records = records_to_json(df.head(limit), fields=fields)
308
+
309
+ return {
310
+ "dataset": dataset,
311
+ "query": query,
312
+ "filters": filters or [],
313
+ "total_matched": total_after_filter,
314
+ "returned": len(records),
315
+ "sort_by": sort_by,
316
+ "warnings": warnings,
317
+ "records": records,
318
+ }
319
+ except Exception as e:
320
+ return {"error": str(e), "traceback": traceback.format_exc()}
321
+
322
+
323
+ # ---------------------------------------------------------------------
324
+ # Tool: get_dataset_columns
325
+ # ---------------------------------------------------------------------
326
 
327
  def tool_get_dataset_columns(dataset: str = "general") -> dict[str, Any]:
328
+ try:
329
+ if dataset not in _datasets:
330
+ return {"error": f"Unknown dataset: '{dataset}'"}
331
+ df = _datasets[dataset]
332
+ numeric_cols = [c for c in df.columns if pd.api.types.is_numeric_dtype(df[c])]
333
+ return {
334
+ "dataset": dataset,
335
+ "num_columns": int(len(df.columns)),
336
+ "columns": list(df.columns),
337
+ "numeric_columns": numeric_cols,
338
+ "known_property_columns": [c for c in NUMERIC_PROPERTY_COLUMNS if c in df.columns],
339
+ }
340
+ except Exception as e:
341
+ return {"error": str(e)}
342
+
343
 
344
+ # ---------------------------------------------------------------------
345
+ # Tool: get_statistics β€” whole-dataset aggregation
346
+ # ---------------------------------------------------------------------
347
 
348
  def tool_get_statistics(
349
  dataset: str = "general",
350
  columns: list[str] | None = None,
351
  filters: list[str] | None = None,
352
  ) -> dict[str, Any]:
353
+ """Descriptive statistics over ALL rows (or a filtered subset).
354
+
355
+ No per-row limit β€” aggregation runs server-side.
 
 
 
 
 
 
 
 
 
356
  """
357
+ try:
358
+ if dataset not in _datasets:
359
+ return {"error": f"Unknown dataset: '{dataset}'"}
360
 
361
+ df = _datasets[dataset].copy()
362
 
363
+ warnings: list[str] = []
364
+ if filters:
365
+ df, warnings = _apply_numeric_filters(df, filters)
 
366
 
367
+ if columns is None:
368
+ columns = [c for c in NUMERIC_PROPERTY_COLUMNS if c in df.columns]
369
 
370
+ stats: dict[str, Any] = {}
371
+ for col in columns:
372
+ if col not in df.columns:
373
+ stats[col] = {"error": "column not found"}
374
+ continue
375
+ s = pd.to_numeric(df[col], errors="coerce").dropna()
376
+ if s.empty:
377
+ stats[col] = {"count": 0, "note": "no numeric data"}
378
+ continue
379
+ stats[col] = {
380
+ "count": int(s.count()),
381
+ "mean": round(float(s.mean()), 6),
382
+ "std": round(float(s.std()), 6),
383
+ "min": round(float(s.min()), 6),
384
+ "p25": round(float(s.quantile(0.25)), 6),
385
+ "median": round(float(s.median()), 6),
386
+ "p75": round(float(s.quantile(0.75)), 6),
387
+ "max": round(float(s.max()), 6),
388
+ }
389
+
390
+ return {
391
+ "dataset": dataset,
392
+ "total_rows": len(_datasets[dataset]),
393
+ "rows_after_filter": len(df),
394
+ "filters": filters or [],
395
+ "warnings": warnings,
396
+ "statistics": stats,
397
  }
398
+ except Exception as e:
399
+ return {"error": str(e), "traceback": traceback.format_exc()}
400
 
 
 
 
 
 
 
 
401
 
402
+ # ---------------------------------------------------------------------
403
+ # Tool: get_correlation β€” whole-dataset correlation
404
+ # ---------------------------------------------------------------------
405
 
406
  def tool_get_correlation(
407
  dataset: str = "general",
408
  col_x: str = "density",
409
  col_y: str = "thermal_conductivity",
410
  filters: list[str] | None = None,
411
+ sample_limit: int = 500,
412
  ) -> dict[str, Any]:
413
+ """Pearson + Spearman correlation computed on ALL matching rows.
414
+
415
+ Returns a random sample of up to `sample_limit` points for plotting.
 
 
 
 
 
 
 
 
 
 
 
 
416
  """
417
+ try:
418
+ if dataset not in _datasets:
419
+ return {"error": f"Unknown dataset: '{dataset}'"}
420
 
421
+ df = _datasets[dataset].copy()
422
 
423
+ warnings: list[str] = []
424
+ if filters:
425
+ df, warnings = _apply_numeric_filters(df, filters)
426
+
427
+ for col in [col_x, col_y]:
428
+ if col not in df.columns:
429
+ return {"error": f"Column '{col}' not found in dataset '{dataset}'."}
430
+ df[col] = pd.to_numeric(df[col], errors="coerce")
431
+
432
+ pair = df[[col_x, col_y]].dropna()
433
+ n = len(pair)
434
 
435
+ if n < 2:
436
+ return {
437
+ "dataset": dataset, "col_x": col_x, "col_y": col_y, "n_total": n,
438
+ "warnings": warnings, "error": "Not enough data points after filtering.",
439
+ }
440
 
441
+ pearson_r = float(pair[col_x].corr(pair[col_y], method="pearson"))
442
+ spearman_r = float(pair[col_x].corr(pair[col_y], method="spearman"))
443
+
444
+ sample_n = min(sample_limit, n)
445
+ sample = pair.sample(n=sample_n, random_state=42).sort_values(col_x)
446
+ sample_points = [
447
+ {col_x: _safe_value(row[col_x]), col_y: _safe_value(row[col_y])}
448
+ for _, row in sample.iterrows()
449
+ ]
450
 
 
451
  return {
452
  "dataset": dataset,
453
  "col_x": col_x,
454
  "col_y": col_y,
455
+ "n_total": n,
456
+ "pearson_r": round(pearson_r, 4),
457
+ "spearman_r": round(spearman_r, 4),
458
+ "filters": filters or [],
459
  "warnings": warnings,
460
+ "sample_n": sample_n,
461
+ "sample_points": sample_points,
462
  }
463
+ except Exception as e:
464
+ return {"error": str(e), "traceback": traceback.format_exc()}
465
+
466
+
467
+ # ---------------------------------------------------------------------
468
+ # Tool: get_distribution β€” whole-dataset histogram
469
+ # ---------------------------------------------------------------------
470
+
471
+ def tool_get_distribution(
472
+ dataset: str = "general",
473
+ column: str = "density",
474
+ bins: int = 20,
475
+ filters: list[str] | None = None,
476
+ ) -> dict[str, Any]:
477
+ """Histogram (bin counts + frequencies) for a numeric column over ALL rows.
478
+
479
+ No per-row limit β€” computed server-side.
480
+ """
481
+ try:
482
+ if dataset not in _datasets:
483
+ return {"error": f"Unknown dataset: '{dataset}'"}
484
 
485
+ df = _datasets[dataset].copy()
 
486
 
487
+ warnings: list[str] = []
488
+ if filters:
489
+ df, warnings = _apply_numeric_filters(df, filters)
 
 
490
 
491
+ if column not in df.columns:
492
+ return {"error": f"Column '{column}' not found in dataset '{dataset}'."}
 
 
 
 
 
 
 
 
 
493
 
494
+ s = pd.to_numeric(df[column], errors="coerce").dropna()
495
+ if s.empty:
496
+ return {"error": f"No numeric data in column '{column}' after filtering."}
497
+
498
+ bins = max(1, min(int(bins), 200))
499
+ counts, edges = np.histogram(s.values, bins=bins)
500
+
501
+ histogram = [
502
+ {
503
+ "bin_start": round(float(edges[i]), 6),
504
+ "bin_end": round(float(edges[i + 1]), 6),
505
+ "bin_mid": round(float((edges[i] + edges[i + 1]) / 2), 6),
506
+ "count": int(counts[i]),
507
+ "frequency": round(float(counts[i] / len(s)), 6),
508
+ }
509
+ for i in range(len(counts))
510
+ ]
511
+
512
+ return {
513
+ "dataset": dataset,
514
+ "column": column,
515
+ "total_rows": len(_datasets[dataset]),
516
+ "rows_used": int(len(s)),
517
+ "bins": bins,
518
+ "filters": filters or [],
519
+ "warnings": warnings,
520
+ "min": round(float(s.min()), 6),
521
+ "max": round(float(s.max()), 6),
522
+ "mean": round(float(s.mean()), 6),
523
+ "median": round(float(s.median()), 6),
524
+ "std": round(float(s.std()), 6),
525
+ "histogram": histogram,
526
+ }
527
+ except Exception as e:
528
+ return {"error": str(e), "traceback": traceback.format_exc()}
529
+
530
+
531
+ # ---------------------------------------------------------------------
532
+ # Tool: get_group_stats β€” GROUP BY aggregation
533
+ # ---------------------------------------------------------------------
534
+
535
+ def tool_get_group_stats(
536
+ dataset: str = "general",
537
+ group_by: str = "polymer_class",
538
+ value_column: str = "thermal_conductivity",
539
+ filters: list[str] | None = None,
540
+ min_group_size: int = 5,
541
+ ) -> dict[str, Any]:
542
+ """Aggregate a numeric column by a categorical column over ALL rows.
543
+
544
+ Returns count, mean, std, min, p25, median, p75, max per group,
545
+ sorted by descending count.
546
+ """
547
+ try:
548
+ if dataset not in _datasets:
549
+ return {"error": f"Unknown dataset: '{dataset}'"}
550
+
551
+ df = _datasets[dataset].copy()
552
+
553
+ warnings: list[str] = []
554
+ if filters:
555
+ df, warnings = _apply_numeric_filters(df, filters)
556
+
557
+ for col in [group_by, value_column]:
558
+ if col not in df.columns:
559
+ return {"error": f"Column '{col}' not found in dataset '{dataset}'."}
560
+
561
+ df[value_column] = pd.to_numeric(df[value_column], errors="coerce")
562
+ df_clean = df[[group_by, value_column]].dropna(subset=[value_column])
563
+
564
+ groups: list[dict[str, Any]] = []
565
+ for name, grp in df_clean.groupby(group_by, sort=False):
566
+ s = grp[value_column]
567
+ if len(s) < min_group_size:
568
+ continue
569
+ groups.append({
570
+ "group": _safe_value(name),
571
+ "count": int(len(s)),
572
+ "mean": round(float(s.mean()), 6),
573
+ "std": round(float(s.std()), 6),
574
+ "min": round(float(s.min()), 6),
575
+ "p25": round(float(s.quantile(0.25)), 6),
576
+ "median": round(float(s.median()), 6),
577
+ "p75": round(float(s.quantile(0.75)), 6),
578
+ "max": round(float(s.max()), 6),
579
+ })
580
+
581
+ groups.sort(key=lambda g: g["count"], reverse=True)
582
+
583
+ return {
584
+ "dataset": dataset,
585
+ "group_by": group_by,
586
+ "value_column": value_column,
587
+ "total_rows": len(_datasets[dataset]),
588
+ "rows_after_filter": len(df),
589
+ "rows_with_value": int(len(df_clean)),
590
+ "num_groups": len(groups),
591
+ "min_group_size": min_group_size,
592
+ "filters": filters or [],
593
+ "warnings": warnings,
594
+ "groups": groups,
595
+ }
596
+ except Exception as e:
597
+ return {"error": str(e), "traceback": traceback.format_exc()}
598
+
599
+
600
+ # ---------------------------------------------------------------------
601
+ # Tool: get_chi_solvents / search_chi
602
+ # ---------------------------------------------------------------------
603
 
604
  def tool_get_chi_solvents() -> dict[str, Any]:
605
+ try:
606
+ loaded = [s for s in SOLVENTS if get_chi_df(s) is not None]
607
+ return {"count": len(loaded), "solvents": loaded}
608
+ except Exception as e:
609
+ return {"error": str(e)}
610
 
611
 
612
  def tool_search_chi(
 
615
  limit: int = 10,
616
  fields: list[str] | None = None,
617
  ) -> dict[str, Any]:
618
+ try:
619
+ df = get_chi_df(solvent)
620
+ if df is None:
621
+ return {"error": f"Chi dataset not found for solvent: '{solvent}'"}
622
+
623
+ if query:
624
+ q_lower = query.lower()
625
+ mask = pd.Series(False, index=df.index)
626
+ for col in df.columns:
627
+ try:
628
+ if df[col].dtype == object:
629
+ mask = mask | df[col].astype(str).str.lower().str.contains(q_lower, na=False)
630
+ except Exception:
631
+ continue
632
+ result_df = df[mask]
633
+ else:
634
+ result_df = df
635
 
636
+ records = records_to_json(result_df.head(limit), fields=fields)
637
+ return {
638
+ "solvent": solvent,
639
+ "query": query,
640
+ "total_matched": len(result_df),
641
+ "returned": len(records),
642
+ "records": records,
643
+ }
644
+ except Exception as e:
645
+ return {"error": str(e)}
646
 
647
 
648
  # ---------------------------------------------------------------------
649
+ # MCP metadata
650
  # ---------------------------------------------------------------------
651
 
652
  SERVER_INFO = {
653
  "name": "polyomics-mcp-server",
654
+ "version": "0.3.0",
655
  }
656
 
657
  TOOLS = [
658
  {
659
  "name": "list_datasets",
660
+ "description": "List available polymer datasets with row counts and column names.",
661
+ "inputSchema": {"type": "object", "properties": {}, "additionalProperties": False},
 
 
 
 
662
  },
663
  {
664
  "name": "search_polymers",
665
  "description": (
666
+ "Search and filter polymers in a dataset. Supports full-text search, "
667
+ "numeric range filters (e.g. 'density>1.0'), sorting, and field selection. "
668
+ "Returns up to 1000 rows. For whole-dataset aggregations use "
669
+ "get_statistics, get_correlation, get_distribution, or get_group_stats."
670
  ),
671
  "inputSchema": {
672
  "type": "object",
 
677
  "default": "",
678
  },
679
  "dataset": {"type": "string", "default": "general"},
680
+ "limit": {
681
+ "type": "integer",
682
+ "default": 10,
683
+ "description": "Max rows to return (max 1000).",
684
+ },
685
  "filters": {
686
  "type": "array",
687
  "items": {"type": "string"},
688
  "description": (
689
+ "Numeric filter expressions, e.g. ['density>1.0', 'tg>=300', "
690
+ "'thermal_conductivity!=']. "
691
  "Operators: >, >=, <, <=, ==, !=. "
692
  "Use 'col!=' to require a non-null value."
693
  ),
 
698
  "fields": {
699
  "type": "array",
700
  "items": {"type": "string"},
701
+ "description": "Columns to include in output. Defaults to a compact property set.",
702
  },
703
  "require_numeric": {
704
  "type": "array",
705
  "items": {"type": "string"},
706
+ "description": "Convenience: columns that must have a non-null numeric value.",
707
  },
708
  },
709
  "additionalProperties": False,
 
712
  {
713
  "name": "get_dataset_columns",
714
  "description": (
715
+ "Return all column names for a dataset, plus lists of numeric columns "
716
  "and known property columns useful for filtering."
717
  ),
718
  "inputSchema": {
719
  "type": "object",
720
+ "properties": {"dataset": {"type": "string", "default": "general"}},
 
 
721
  "additionalProperties": False,
722
  },
723
  },
724
  {
725
  "name": "get_statistics",
726
  "description": (
727
+ "Compute descriptive statistics (count, mean, std, min, p25, median, p75, max) "
728
+ "for numeric property columns using ALL rows in the dataset (server-side aggregation). "
729
+ "Optional filters are applied before aggregation. "
730
+ "Use this instead of search_polymers for whole-dataset summaries."
731
  ),
732
  "inputSchema": {
733
  "type": "object",
 
741
  "filters": {
742
  "type": "array",
743
  "items": {"type": "string"},
744
+ "description": "Numeric filters applied before computing statistics.",
745
  "default": [],
746
  },
747
  },
 
752
  "name": "get_correlation",
753
  "description": (
754
  "Compute Pearson and Spearman correlation between two numeric columns "
755
+ "using ALL matching rows (server-side β€” no row limit). "
756
+ "Returns correlation coefficients plus a random scatter-plot sample."
757
  ),
758
  "inputSchema": {
759
  "type": "object",
 
767
  "description": "Numeric filters applied before computing the correlation.",
768
  "default": [],
769
  },
770
+ "sample_limit": {
771
  "type": "integer",
772
+ "default": 500,
773
+ "description": "Max data points to return for scatter-plot visualisation.",
774
  },
775
  },
776
  "required": ["col_x", "col_y"],
 
778
  },
779
  },
780
  {
781
+ "name": "get_distribution",
782
+ "description": (
783
+ "Compute a histogram (bin counts and frequencies) for a numeric column "
784
+ "over ALL matching rows (server-side β€” no row limit). "
785
+ "Ideal for visualising the distribution of density, TC, Tg, etc."
786
+ ),
787
+ "inputSchema": {
788
+ "type": "object",
789
+ "properties": {
790
+ "dataset": {"type": "string", "default": "general"},
791
+ "column": {
792
+ "type": "string",
793
+ "default": "density",
794
+ "description": "Numeric column to compute the histogram for.",
795
+ },
796
+ "bins": {
797
+ "type": "integer",
798
+ "default": 20,
799
+ "description": "Number of histogram bins (1–200).",
800
+ },
801
+ "filters": {
802
+ "type": "array",
803
+ "items": {"type": "string"},
804
+ "description": "Numeric filters applied before computing the histogram.",
805
+ "default": [],
806
+ },
807
+ },
808
+ "additionalProperties": False,
809
+ },
810
+ },
811
+ {
812
+ "name": "get_group_stats",
813
+ "description": (
814
+ "Aggregate a numeric column grouped by a categorical column (e.g. polymer_class). "
815
+ "Runs on ALL matching rows server-side. "
816
+ "Returns count, mean, std, min, p25, median, p75, max per group. "
817
+ "Useful for comparing thermal_conductivity or density across polymer classes."
818
+ ),
819
  "inputSchema": {
820
  "type": "object",
821
+ "properties": {
822
+ "dataset": {"type": "string", "default": "general"},
823
+ "group_by": {
824
+ "type": "string",
825
+ "default": "polymer_class",
826
+ "description": "Categorical column to group on (e.g. 'polymer_class', '_source').",
827
+ },
828
+ "value_column": {
829
+ "type": "string",
830
+ "default": "thermal_conductivity",
831
+ "description": "Numeric column to aggregate.",
832
+ },
833
+ "filters": {
834
+ "type": "array",
835
+ "items": {"type": "string"},
836
+ "description": "Numeric filters applied before grouping.",
837
+ "default": [],
838
+ },
839
+ "min_group_size": {
840
+ "type": "integer",
841
+ "default": 5,
842
+ "description": "Groups with fewer rows than this are omitted.",
843
+ },
844
+ },
845
  "additionalProperties": False,
846
  },
847
  },
848
+ {
849
+ "name": "get_chi_solvents",
850
+ "description": "List available solvents for chi parameter tables.",
851
+ "inputSchema": {"type": "object", "properties": {}, "additionalProperties": False},
852
+ },
853
  {
854
  "name": "search_chi",
855
  "description": "Search a chi parameter table for a given solvent.",
 
872
  ]
873
 
874
 
875
+ # ---------------------------------------------------------------------
876
+ # Tool dispatcher
877
+ # ---------------------------------------------------------------------
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
878
 
879
+ def handle_tool_call(name: str, arguments: dict[str, Any]) -> dict[str, Any]:
880
+ try:
881
+ if name == "list_datasets":
882
+ return tool_list_datasets()
883
+
884
+ if name == "search_polymers":
885
+ return tool_search_polymers(
886
+ query=arguments.get("query", ""),
887
+ dataset=arguments.get("dataset", "general"),
888
+ limit=int(arguments.get("limit", 10)),
889
+ filters=arguments.get("filters") or [],
890
+ sort_by=arguments.get("sort_by"),
891
+ sort_ascending=bool(arguments.get("sort_ascending", True)),
892
+ fields=arguments.get("fields"),
893
+ require_numeric=arguments.get("require_numeric"),
894
+ )
895
+
896
+ if name == "get_dataset_columns":
897
+ return tool_get_dataset_columns(dataset=arguments.get("dataset", "general"))
898
+
899
+ if name == "get_statistics":
900
+ return tool_get_statistics(
901
+ dataset=arguments.get("dataset", "general"),
902
+ columns=arguments.get("columns"),
903
+ filters=arguments.get("filters") or [],
904
+ )
905
+
906
+ if name == "get_correlation":
907
+ return tool_get_correlation(
908
+ dataset=arguments.get("dataset", "general"),
909
+ col_x=arguments.get("col_x", "density"),
910
+ col_y=arguments.get("col_y", "thermal_conductivity"),
911
+ filters=arguments.get("filters") or [],
912
+ sample_limit=int(arguments.get("sample_limit", 500)),
913
+ )
914
+
915
+ if name == "get_distribution":
916
+ return tool_get_distribution(
917
+ dataset=arguments.get("dataset", "general"),
918
+ column=arguments.get("column", "density"),
919
+ bins=int(arguments.get("bins", 20)),
920
+ filters=arguments.get("filters") or [],
921
+ )
922
+
923
+ if name == "get_group_stats":
924
+ return tool_get_group_stats(
925
+ dataset=arguments.get("dataset", "general"),
926
+ group_by=arguments.get("group_by", "polymer_class"),
927
+ value_column=arguments.get("value_column", "thermal_conductivity"),
928
+ filters=arguments.get("filters") or [],
929
+ min_group_size=int(arguments.get("min_group_size", 5)),
930
+ )
931
+
932
+ if name == "get_chi_solvents":
933
+ return tool_get_chi_solvents()
934
+
935
+ if name == "search_chi":
936
+ return tool_search_chi(
937
+ solvent=arguments.get("solvent", ""),
938
+ query=arguments.get("query", ""),
939
+ limit=int(arguments.get("limit", 10)),
940
+ fields=arguments.get("fields"),
941
+ )
942
+
943
+ return {"error": f"Unknown tool: '{name}'"}
944
 
945
+ except Exception as e:
946
+ return {
947
+ "error": f"Unhandled exception in tool '{name}': {e}",
948
+ "traceback": traceback.format_exc(),
949
+ }
950
 
951
 
952
  # ---------------------------------------------------------------------
 
955
 
956
  async def mcp_sse_get(request: Request) -> Response:
957
  print(">>> GET /mcp/sse")
 
958
  return PlainTextResponse(
959
+ "event: ping\ndata: alive\n\n",
960
  media_type="text/event-stream",
961
  headers={
962
  "Cache-Control": "no-cache",
 
971
  payload = await request.json()
972
  except Exception:
973
  raw = await request.body()
974
+ print(">>> POST /mcp/sse β€” invalid JSON")
975
+ print(raw.decode("utf-8", errors="replace")[:500])
976
  return safe_jsonrpc_error(None, -32700, "Parse error")
977
 
978
  print(">>> POST /mcp/sse")
979
+ print(json.dumps(payload, ensure_ascii=False)[:1000])
980
 
981
  req_id = payload.get("id")
982
+ method = payload.get("method")
983
+ params = payload.get("params") or {}
984
 
985
  if method == "initialize":
 
986
  return safe_jsonrpc_result(req_id, {
987
+ "protocolVersion": params.get("protocolVersion", "2025-03-26"),
988
  "capabilities": {"tools": {}},
989
  "serverInfo": SERVER_INFO,
990
  })
 
999
  return safe_jsonrpc_result(req_id, {"tools": TOOLS})
1000
 
1001
  if method == "tools/call":
1002
+ name = params.get("name")
1003
+ arguments = params.get("arguments") or {}
1004
+ result = handle_tool_call(name, arguments)
1005
+ is_error = isinstance(result, dict) and "error" in result
1006
  return safe_jsonrpc_result(req_id, {
1007
+ "content": [{"type": "text", "text": stringify_records([result])}],
1008
+ "isError": is_error,
 
 
 
 
 
1009
  })
1010
 
1011
  return safe_jsonrpc_error(req_id, -32601, f"Method not found: {method}")
1012
 
1013
 
1014
  # ---------------------------------------------------------------------
1015
+ # Auxiliary HTTP endpoints
1016
  # ---------------------------------------------------------------------
1017
 
1018
  async def health(request: Request) -> JSONResponse:
 
1027
 
1028
 
1029
  async def oauth_protected_resource(request: Request) -> Response:
 
1030
  return Response(status_code=204)
1031
 
1032
 
1033
  async def oauth_authorization_server(request: Request) -> Response:
 
1034
  return Response(status_code=204)
1035
 
1036
 
1037
  async def register(request: Request) -> Response:
1038
  raw = await request.body()
1039
+ print(">>> POST /register:", raw.decode("utf-8", errors="replace")[:200])
 
1040
  return Response(status_code=204)
1041
 
1042
 
 
1052
 
1053
 
1054
  # ---------------------------------------------------------------------
1055
+ # Gradio UI
1056
  # ---------------------------------------------------------------------
1057
 
1058
  def build_gradio_app() -> GradioApp:
 
1061
  "# PolyOmics MCP Server\n\n"
1062
  f"**Dataset repo:** `{DATASET_REPO}`\n\n"
1063
  "**Claude connector endpoint:** `https://mohnishi-polyomics-mcp-server.hf.space/mcp/sse`\n\n"
1064
+ "**Available MCP tools (v0.3.0):**\n"
1065
  "- `list_datasets` β€” list datasets and row counts\n"
1066
+ "- `search_polymers` β€” full-text + numeric filter + sort + field selection (up to 1000 rows)\n"
1067
  "- `get_dataset_columns` β€” column names, numeric columns, known property columns\n"
1068
+ "- `get_statistics` β€” **whole-dataset** descriptive stats per property column\n"
1069
+ "- `get_correlation` β€” **whole-dataset** Pearson/Spearman + scatter sample\n"
1070
+ "- `get_distribution` β€” **whole-dataset** histogram for any numeric column\n"
1071
+ "- `get_group_stats` β€” **whole-dataset** GROUP BY aggregation (e.g. TC by polymer_class)\n"
1072
  "- `get_chi_solvents` β€” list loaded chi-parameter solvents\n"
1073
  "- `search_chi` β€” search chi parameter tables\n"
1074
  )
1075
 
1076
  df = get_main_df()
1077
+ gr.JSON({
1078
+ "dataset_repo": DATASET_REPO,
1079
  "server_version": SERVER_INFO["version"],
1080
+ "main_rows": int(len(df)),
1081
+ "main_columns": int(len(df.columns)) if not df.empty else 0,
1082
+ "datasets": list(_datasets.keys()),
1083
+ })
 
1084
 
1085
  return GradioApp.create_app(demo, app_kwargs={"docs_url": "/docs"})
1086
 
1087
 
1088
  # ---------------------------------------------------------------------
1089
+ # Starlette app assembly
1090
  # ---------------------------------------------------------------------
1091
 
1092
  def build_app() -> Starlette:
 
1095
  app = Starlette(
1096
  routes=[
1097
  Route("/health", endpoint=health, methods=["GET"]),
1098
+ Route("/mcp/sse", endpoint=mcp_sse_get, methods=["GET"]),
1099
  Route("/mcp/sse", endpoint=mcp_sse_post, methods=["POST"]),
1100
  Route("/mcp/sse", endpoint=options_handler, methods=["OPTIONS"]),
1101
+ Route("/.well-known/oauth-protected-resource", endpoint=oauth_protected_resource, methods=["GET"]),
1102
+ Route("/.well-known/oauth-protected-resource/mcp/sse", endpoint=oauth_protected_resource, methods=["GET"]),
1103
+ Route("/.well-known/oauth-authorization-server", endpoint=oauth_authorization_server, methods=["GET"]),
 
 
 
 
 
 
 
 
 
 
 
 
1104
  Route("/register", endpoint=register, methods=["POST"]),
1105
  Mount("/", app=gradio_app),
1106
  ]
1107
  )
1108
 
1109
+ print(f"MCP JSON-RPC server v{SERVER_INFO['version']} ready at /mcp/sse")
1110
  for r in app.routes:
1111
+ print(f" {type(r).__name__}: {getattr(r, 'path', None)}")
1112
 
1113
  return app
1114
 
1115
 
1116
  # ---------------------------------------------------------------------
1117
+ # Entry point
1118
  # ---------------------------------------------------------------------
1119
 
1120
  if __name__ == "__main__":
1121
+ now = datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S UTC")
1122
  print(f"===== Application Startup at {now} =====")
1123
  preload_data()
1124
  app = build_app()
1125
  port = int(os.environ.get("PORT", 7860))
1126
+ print(f"Starting uvicorn on 0.0.0.0:{port} ...")
1127
+ uvicorn.run(app, host="0.0.0.0", port=port)