Skip to content

Commit a47df36

Browse files
authored
Add compatibility with diagram schema version 2 (#2)
* Support new diagram schema. * Add NumMols evaluator.
1 parent f03f53b commit a47df36

5 files changed

Lines changed: 514 additions & 364 deletions

File tree

api.py

Lines changed: 97 additions & 47 deletions
Original file line numberDiff line numberDiff line change
@@ -4,11 +4,14 @@
44
from dotenv import load_dotenv
55
from rdkit import Chem
66
from rdkit import RDLogger
7+
from rdkit.Chem import AllChem, rdmolops, Descriptors
78

89

910
RDLogger.DisableLog('rdApp.*')
1011

1112

13+
# === Reading MolView diagram JSON ===
14+
1215
bond_type = {
1316
'single': Chem.rdchem.BondType.SINGLE,
1417
'double': Chem.rdchem.BondType.DOUBLE,
@@ -24,22 +27,32 @@
2427
}
2528

2629

27-
def check_authorization(provided_key: str) -> bool:
28-
allowed_key = environ.get('API_KEY')
29-
if not allowed_key:
30-
return False
31-
return provided_key == allowed_key
32-
33-
34-
def read_vec2(data) -> tuple[float, float]:
30+
def vec2(data) -> tuple[float, float]:
3531
if isinstance(data, list):
3632
return (data[0] / 100, data[1] / 100)
3733
else:
3834
return (data['x'] / 100, data['y'] / 100)
3935

4036

37+
def lone_pairs(atom: dict):
38+
if 'non_bonded_ve' in atom:
39+
return atom['non_bonded_ve'] // 2
40+
else:
41+
return atom.get('lone_pairs', 0)
42+
43+
44+
def unpaired_electrons(atom: dict):
45+
if 'non_bonded_ve' in atom:
46+
return atom['non_bonded_ve'] % 2
47+
else:
48+
return atom.get('unpaired_electrons', 0)
49+
50+
4151
def bond_edge(bond):
42-
return bond['atoms'] if 'atoms' in bond else [bond['from'], bond['to']]
52+
if 'atoms' in bond:
53+
return bond['atoms']
54+
else:
55+
return [bond['from'], bond['to']]
4356

4457

4558
def attach_arrow_endpoint(g: nx.Graph, endpoint, anchor, bonds):
@@ -56,12 +69,13 @@ def json_to_graph(diagram, add_arrows=False) -> nx.Graph:
5669
g = nx.Graph()
5770

5871
for i, atom in enumerate(diagram['atoms']):
59-
g.add_node(i, **{
72+
g.add_node(atom.get('id') or i, **{
6073
'type': 'atom',
61-
'symbol': atom['symbol'],
62-
'position': read_vec2(atom['position']),
74+
'label': atom.get('label') or atom['symbol'],
75+
'position': vec2(atom['position']),
6376
'formal_charge': atom.get('formal_charge', 0),
64-
'non_bonded_ve': atom.get('non_bonded_ve', 0)
77+
'lone_pairs': lone_pairs(atom),
78+
'unpaired_electrons': unpaired_electrons(atom)
6579
})
6680

6781
for bond in diagram['bonds']:
@@ -92,7 +106,7 @@ def graph_to_mol(g: nx.Graph) -> Chem.rdchem.RWMol:
92106

93107
for i, attr in g.nodes(data=True):
94108
if attr['type'] == 'atom':
95-
a = Chem.rdchem.Atom(attr['symbol'])
109+
a = Chem.rdchem.Atom(attr['label'])
96110
a.SetFormalCharge(attr['formal_charge'])
97111
atom_index = mol.AddAtom(a)
98112
atoms[i] = atom_index
@@ -131,6 +145,9 @@ def json_to_mol(data) -> Chem.rdchem.RWMol:
131145
return graph_to_mol(json_to_graph(data))
132146

133147

148+
# === Utilities based on NetworkX ===
149+
150+
134151
def match_props(x, y, props):
135152
return all(map(lambda p: x.get(p[0], p[1]) == y.get(p[0], p[1]), props))
136153

@@ -140,10 +157,10 @@ def match_nodes(v1, v2):
140157
if v1['type'] == 'atom':
141158
return match_props(v1, v2, [
142159
['type', ''],
143-
['symbol', ''],
160+
['label', ''],
144161
['formal_charge', 0],
145-
['non_bonded_ve', 0],
146-
['mark', -1]
162+
['lone_pairs', 0],
163+
['unpaired_electrons', 0]
147164
])
148165
elif v1['type'] == 'arrow':
149166
return match_props(v1, v2, [
@@ -165,18 +182,13 @@ def match_edges(e1, e2):
165182
return True
166183

167184

168-
def graph_to_smiles(g: nx.Graph, isomeric = True) -> str:
169-
mol = Chem.RemoveHs(graph_to_mol(g))
170-
return Chem.rdmolfiles.MolToSmiles(mol, isomericSmiles=isomeric)
171-
172-
173185
def graph_basic_hydrogens(g: nx.Graph) -> list:
174186
# Select hdrogen atoms _where all other attributes are empty_.
175187
return [v for v, attr in g.nodes(data=True) if
176188
attr['type'] == 'atom' and
177-
attr['symbol'] == 'H' and
189+
attr['label'] == 'H' and
178190
attr['formal_charge'] == 0 and
179-
attr['non_bonded_ve'] == 0]
191+
attr['unpaired_electrons'] == 0]
180192

181193

182194
def compare_component(g1: nx.Graph, g2: nx.Graph, match_stereo: bool) -> bool:
@@ -185,20 +197,46 @@ def compare_component(g1: nx.Graph, g2: nx.Graph, match_stereo: bool) -> bool:
185197
g1_no_h.remove_nodes_from(graph_basic_hydrogens(g1_no_h))
186198
g2_no_h.remove_nodes_from(graph_basic_hydrogens(g2_no_h))
187199

188-
# Don't compare edge types to avoid filtering resonance structures.
200+
# Don't compare bond types to avoid filtering resonance structures.
189201
if nx.is_isomorphic(g1_no_h, g2_no_h, match_nodes, match_edges):
190202
# This can fail because of non-standard valences, like in ClF3, which is
191203
# not supported by RDKit. Therefore, in the case of an error, we default
192204
# to accepting the bare isomorphism.
193205
try:
194-
smi1 = graph_to_smiles(g1, match_stereo)
195-
smi2 = graph_to_smiles(g2, match_stereo)
206+
smi1 = get_smiles(graph_to_mol(g1), match_stereo)
207+
smi2 = get_smiles(graph_to_mol(g2), match_stereo)
196208
return smi1 == smi2
197209
except ValueError:
198210
return True
199211

212+
else:
213+
return False
214+
215+
216+
# === Utilities based on RDKit ===
217+
218+
219+
def get_smiles(mol: Chem.Mol, isomeric=True) -> str:
220+
mol = Chem.RemoveHs(mol)
221+
return Chem.MolToSmiles(mol, isomericSmiles=isomeric)
200222

201-
def compare(diagram1, diagram2, match_stereo = True) -> bool:
223+
224+
def get_fragments(mol):
225+
return rdmolops.GetMolFrags(mol, asMols=True)
226+
227+
228+
# === Main operations ===
229+
230+
231+
def validate(diagram) -> bool:
232+
try:
233+
json_to_mol(diagram)
234+
return True
235+
except Exception:
236+
return False
237+
238+
239+
def compare(diagram1, diagram2, match_stereo=True) -> bool:
202240
g1 = json_to_graph(diagram1, True)
203241
g2 = json_to_graph(diagram2, True)
204242
cs1 = list(map(lambda c: g1.subgraph(c), nx.connected_components(g1)))
@@ -218,27 +256,23 @@ def compare(diagram1, diagram2, match_stereo = True) -> bool:
218256
return len(cs2) == 0
219257

220258

221-
def validate(diagram) -> bool:
222-
try:
223-
json_to_mol(diagram)
224-
return True
225-
except:
226-
return False
259+
def evaluate(diagram, evaluator, params):
260+
mol = json_to_mol(diagram)
261+
match evaluator:
262+
case 'NumMols':
263+
return len(set(map(get_smiles, get_fragments(mol))))
264+
case _:
265+
return {'err': 'Unknown evaluator'}
227266

228267

229-
def handle_validate(body):
230-
return {
231-
'valid': validate(body['diagram'])
232-
}
268+
# === AWS Lambda handler ===
233269

234270

235-
def handle_compare(body):
236-
return {
237-
'equal': compare(
238-
body['reference_diagram'],
239-
body['student_diagram'],
240-
body['match_stereo'])
241-
}
271+
def check_authorization(provided_key: str) -> bool:
272+
allowed_key = environ.get('API_KEY')
273+
if not allowed_key:
274+
return False
275+
return provided_key == allowed_key
242276

243277

244278
def handler(event, context):
@@ -268,11 +302,27 @@ def handler(event, context):
268302
request_path_raw: str = event.get('requestContext', {}).get('path', '')
269303
request_path = request_path_raw.replace('/api/v1/', '')
270304
request_body = json.loads(event['body'])
305+
271306
match request_path:
272307
case 'compare':
273-
response_body = handle_compare(request_body)
308+
response_body = {
309+
'equal': compare(
310+
request_body['reference_diagram'],
311+
request_body['student_diagram'],
312+
request_body['match_stereo'])
313+
}
274314
case 'validate':
275-
response_body = handle_validate(request_body)
315+
response_body = {
316+
'valid': validate(
317+
request_body['diagram'])
318+
}
319+
case 'evaluate':
320+
response_body = {
321+
'result': evaluate(
322+
request_body['diagram'],
323+
request_body['evaluator'],
324+
request_body['params'])
325+
}
276326
case _:
277327
response_body = {'err': 'Invalid request'}
278328

0 commit comments

Comments
 (0)