Skip to content

Commit 21c4306

Browse files
committed
chore(retrosynthesis): modification tools
1 parent a8626b5 commit 21c4306

4 files changed

Lines changed: 48 additions & 26 deletions

File tree

ChemCoScientist/agents/agents.py

Lines changed: 25 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -251,10 +251,23 @@ def chemist_node(state: dict, config: dict) -> Command:
251251
config["configurable"]["state"] = state
252252
agent_response = chem_agent.invoke({"messages": [("user", task_formatted)]})
253253

254+
def _parse_tool_content(content):
255+
try:
256+
return json.loads(content)
257+
except Exception:
258+
try:
259+
return ast.literal_eval(content)
260+
except Exception:
261+
return None
262+
254263
updated_metadata = state.get("metadata", {}).copy()
255264
for message in agent_response["messages"]:
265+
if isinstance(message, ToolMessage):
266+
print(f"TOOL MESSAGE: {message.name}")
256267
if isinstance(message, ToolMessage) and message.name in ["detect_molecules", "detect_reactions"]:
257-
result = ast.literal_eval(message.content)
268+
result = _parse_tool_content(message.content)
269+
if result is None:
270+
continue
258271
ocr_metadata = {"chem_ocr": result.get("metadata", None)}
259272
if ocr_metadata["chem_ocr"]:
260273
if "chem_ocr" in updated_metadata.keys():
@@ -263,7 +276,9 @@ def chemist_node(state: dict, config: dict) -> Command:
263276
updated_metadata.update(ocr_metadata)
264277

265278
elif isinstance(message, ToolMessage) and message.name in ["calculate_docking_score"]:
266-
result = ast.literal_eval(message.content)
279+
result = _parse_tool_content(message.content)
280+
if result is None:
281+
continue
267282
docking_metadata = {"docking": result.get("metadata", None)}
268283
if docking_metadata["docking"]:
269284
if "docking" in updated_metadata.keys():
@@ -272,23 +287,22 @@ def chemist_node(state: dict, config: dict) -> Command:
272287
updated_metadata.update(docking_metadata)
273288

274289
elif isinstance(message, ToolMessage) and message.name in ["retrosynthesis_tree_search"]:
275-
try:
276-
result = ast.literal_eval(message.content)
277-
except Exception:
290+
result = _parse_tool_content(message.content)
291+
if result is None:
278292
continue
293+
if isinstance(result, dict):
294+
print(f"RETRO RESULT KEYS: {list(result.keys())[:10]}")
279295
updated_metadata.update({"retrosynthesis": result})
280296

281297
elif isinstance(message, ToolMessage) and message.name in ["classify_reaction"]:
282-
try:
283-
result = ast.literal_eval(message.content)
284-
except Exception:
298+
result = _parse_tool_content(message.content)
299+
if result is None:
285300
continue
286301
updated_metadata.update({"reaction_classification": result})
287302

288303
elif isinstance(message, ToolMessage) and message.name in ["forward_predict"]:
289-
try:
290-
result = ast.literal_eval(message.content)
291-
except Exception:
304+
result = _parse_tool_content(message.content)
305+
if result is None:
292306
continue
293307
updated_metadata.update({"forward_prediction": result})
294308

ChemCoScientist/agents/agents_prompts.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -127,4 +127,4 @@
127127
You must detect and output every plausible chemical structure present in the image, even if the image is low-quality,
128128
sketchy, partial, or ambiguous. When uncertain, infer the most likely structure based on visible atoms, bonds, and geometry.
129129
Never return ‘no molecules detected’—instead describe all candidate structures with confidence scores.
130-
"""
130+
"""

ChemCoScientist/chemical_utils/retrosynthesis.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -129,8 +129,8 @@ def classify_reaction_smiles(smiles: List[str], num_results: int = 10) -> Dict[s
129129

130130
def forward_predict_products(
131131
smiles: List[str],
132-
backend: str,
133-
model_name: str = "wldn5",
132+
backend: str = "wldn5",
133+
model_name: str = "pistachio",
134134
reagents: str = "",
135135
solvent: str = "",
136136
) -> Dict[str, Any]:
@@ -178,6 +178,7 @@ def forward_predict_products(
178178
error_msg = "Forward Prediction API returned None JSON response"
179179
logger.error(error_msg)
180180
raise ValueError(error_msg)
181+
logger.info(f"FORWARD PREDICTION JSON RESPONSE: {json_response}")
181182
return json_response
182183
except requests.exceptions.RequestException as e:
183184
error_msg = (

ChemCoScientist/frontend/chat.py

Lines changed: 19 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,8 @@
55
import os
66
import streamlit as st
77
import threading
8+
import re
9+
from typing import Optional, List
810
from typing import Optional
911

1012
from io import BytesIO
@@ -356,6 +358,8 @@ def message_handler(user_query: str, placeholder: st.delta_generator.DeltaGenera
356358
#os.remove(file)
357359

358360
# Store metadata in the message for later display
361+
362+
logger.info(f"RESULT METADATA: {result}")
359363
if "paper_analysis" in result["metadata"].keys():
360364
st.session_state.messages[-1]["paper_analysis"] = result["metadata"]["paper_analysis"]
361365
# Display the metadata immediately after storing it
@@ -452,10 +456,13 @@ def _reaction_smiles_to_image(reaction_smiles: str):
452456

453457

454458
def _render_reaction_smiles(reaction_smiles: str, caption: Optional[str] = None):
459+
logger.info(f"RENDERING REACTION SMILES: {reaction_smiles}")
455460
normalized = reaction_smiles
456461
if isinstance(normalized, str):
457462
normalized = normalized.replace(" -> ", ">>").replace(" → ", ">>")
458463
normalized = normalized.replace(" + ", ".")
464+
normalized = re.sub(r">{3,}", ">>", normalized)
465+
normalized = re.sub(r">>\s*>>", ">>", normalized)
459466
img = _reaction_smiles_to_image(normalized)
460467
if img is not None:
461468
st.image(img, caption=caption)
@@ -464,6 +471,7 @@ def _render_reaction_smiles(reaction_smiles: str, caption: Optional[str] = None)
464471

465472

466473
def display_retrosynthesis_metadata(message):
474+
logger.info(f"DISPLAYING RETROSYNTHESIS METADATA: {message}")
467475
data = message.get("retrosynthesis") or {}
468476
if not data:
469477
return
@@ -522,10 +530,13 @@ def _route_sort_key(r):
522530
continue
523531
for step_idx, step in enumerate(steps, start=1):
524532
reaction_smiles = step.get("reaction_smiles") or step.get("mapped_smiles")
533+
logger.info(f"REACTION SMILES: {reaction_smiles}")
534+
logger.info(f"STEP: {step}")
525535
caption = f"Step {step_idx}"
526536
if step.get("plausibility") is not None:
527537
caption += f" | plausibility={step.get('plausibility')}"
528538
if reaction_smiles:
539+
logger.info(f"RENDERING REACTION SMILES: {reaction_smiles}")
529540
_render_reaction_smiles(reaction_smiles, caption=caption)
530541
else:
531542
st.markdown(f"**Step {step_idx}**")
@@ -548,10 +559,6 @@ def display_forward_prediction_metadata(message):
548559
if backend or model_name:
549560
st.markdown(f"**Backend:** `{backend}` **Model:** `{model_name}`")
550561
inputs = data.get("inputs") or []
551-
if inputs:
552-
st.markdown("**Inputs:**")
553-
for item in inputs:
554-
st.code(item)
555562
predictions = data.get("predictions") or []
556563
if predictions and isinstance(predictions, list):
557564
predictions = sorted(
@@ -564,16 +571,16 @@ def display_forward_prediction_metadata(message):
564571
return
565572
if len(inputs) == 1:
566573
base = inputs[0]
567-
for idx, pred in enumerate(predictions, start=1):
574+
pred = predictions[0] if predictions else None
575+
if pred:
568576
prod = pred.get("smiles")
569577
score = pred.get("score")
570-
if not prod:
571-
continue
572-
reaction_smiles = f"{base}>>{prod}"
573-
caption = f"Prediction {idx}"
574-
if score is not None:
575-
caption += f" | score={score}"
576-
_render_reaction_smiles(reaction_smiles, caption=caption)
578+
if prod:
579+
reaction_smiles = f"{base}>>{prod}"
580+
caption = "Best prediction"
581+
if score is not None:
582+
caption += f" | score={score}"
583+
_render_reaction_smiles(reaction_smiles, caption=caption)
577584
else:
578585
rows = [{"smiles": p.get("smiles"), "score": p.get("score")} for p in predictions]
579586
st.dataframe(rows)

0 commit comments

Comments
 (0)