Skip to content

Commit 6d2a256

Browse files
committed
fix: assembly_substituent_tree() now accepts precomputed trace_segments, so callers do not recompute assembly_trace_segments(parts) when both trace and tree are requested, component_namer.name_component() and recursive name_subgraph() now compute trace segments once and reuse them for both return value and tree. Grouped same-name substituents no longer silently overwrite substituent_tree. If multiple instances have different trees, they are preserved
1 parent 4aa9452 commit 6d2a256

4 files changed

Lines changed: 74 additions & 14 deletions

File tree

src/bluenamer/component_namer.py

Lines changed: 13 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -593,17 +593,21 @@ def name_component_again(next_mol: Molecule, next_atoms: set[int], is_substituen
593593
"name_rewrite_history": parts.name_rewrite_history,
594594
},
595595
)
596-
tree = assembly_substituent_tree(
597-
parts,
598-
name=name,
599-
atom_ids=state.component_atoms,
600-
bond_ids=bond_ids_within(mol, state.component_atoms),
601-
)
602-
tree["kind"] = "component"
596+
trace_segments = assembly_trace_segments(parts) if return_trace or return_tree else []
597+
tree = None
598+
if return_tree:
599+
tree = assembly_substituent_tree(
600+
parts,
601+
name=name,
602+
atom_ids=state.component_atoms,
603+
bond_ids=bond_ids_within(mol, state.component_atoms),
604+
trace_segments=trace_segments,
605+
)
606+
tree["kind"] = "component"
603607
if return_trace and return_tree:
604-
return name, assembly_trace_segments(parts), tree
608+
return name, trace_segments, tree
605609
if return_trace:
606-
return name, assembly_trace_segments(parts)
610+
return name, trace_segments
607611
if return_tree:
608612
return name, tree
609613
return name

src/bluenamer/namer.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1236,6 +1236,7 @@ def name_subgraph(
12361236
)
12371237

12381238
name = _assemble_parent_name(mol, parts, numbered_path, get_loc, finalize_subgraph=True)
1239+
trace_segments = _assembly_trace_segments(parts) if return_trace or return_tree else []
12391240
trace_decision(
12401241
decision_trace,
12411242
TracePhase.ASSEMBLY,
@@ -1245,11 +1246,10 @@ def name_subgraph(
12451246
bonds=_bond_ids_within(mol, component),
12461247
data={
12471248
"name": name,
1248-
"trace_segment_count": len(_assembly_trace_segments(parts)) if return_trace else 0,
1249+
"trace_segment_count": len(trace_segments),
12491250
},
12501251
)
12511252
if return_trace:
1252-
trace_segments = _assembly_trace_segments(parts)
12531253
if return_tree:
12541254
return (
12551255
name,
@@ -1260,6 +1260,7 @@ def name_subgraph(
12601260
atom_ids=component,
12611261
bond_ids=_bond_ids_within(mol, component),
12621262
decisions=decision_trace_data(decision_trace),
1263+
trace_segments=trace_segments,
12631264
),
12641265
)
12651266
return name, trace_segments
@@ -1272,6 +1273,7 @@ def name_subgraph(
12721273
atom_ids=component,
12731274
bond_ids=_bond_ids_within(mol, component),
12741275
decisions=decision_trace_data(decision_trace),
1276+
trace_segments=trace_segments,
12751277
),
12761278
)
12771279
return name

src/bluenamer/tests/test_analysis.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -290,6 +290,29 @@ def test_add_substituent_trace_preserves_nested_decisions_into_segments():
290290
assert assembly_trace_segments(parts)[0]["nested_decisions"] == nested_decisions
291291

292292

293+
def test_add_substituent_trace_preserves_grouped_tree_instances():
294+
parts = AssemblyParts(parent_length=1, parent_atom_ids={0})
295+
296+
add_substituent_trace(
297+
parts,
298+
"ethyl",
299+
"1",
300+
atom_ids={1, 2},
301+
substituent_tree={"kind": "substituent", "name": "ethyl", "atoms": [1, 2]},
302+
)
303+
add_substituent_trace(
304+
parts,
305+
"ethyl",
306+
"2",
307+
atom_ids={3, 4},
308+
substituent_tree={"kind": "substituent", "name": "ethyl", "atoms": [3, 4]},
309+
)
310+
311+
tree = parts.substituents[0].substituent_tree
312+
assert tree["kind"] == "grouped_substituent_instances"
313+
assert [instance["atoms"] for instance in tree["instances"]] == [[1, 2], [3, 4]]
314+
315+
293316
def test_principal_suffix_tokens_are_emitted_from_functional_group_renderer():
294317
parts = AssemblyParts(
295318
parent_length=3,

src/bluenamer/trace_helpers.py

Lines changed: 34 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -113,7 +113,11 @@ def add_substituent_trace(
113113
existing.trace_segments.extend(trace_segments)
114114
existing.nested_decisions.extend(nested_decisions)
115115
if substituent_tree:
116-
existing.substituent_tree = substituent_tree
116+
existing.substituent_tree = _merge_substituent_tree_instances(
117+
existing.substituent_tree,
118+
substituent_tree,
119+
existing.name,
120+
)
117121
existing.emitted_tokens = existing.emitted_tokens + tuple(emitted_tokens)
118122
else:
119123
parts.substituents.append(
@@ -186,7 +190,11 @@ def assembly_trace_segments(parts: AssemblyParts) -> list[dict]:
186190
target.trace_segments.extend(item.trace_segments)
187191
target.nested_decisions.extend(item.nested_decisions)
188192
if item.substituent_tree:
189-
target.substituent_tree = item.substituent_tree
193+
target.substituent_tree = _merge_substituent_tree_instances(
194+
target.substituent_tree,
195+
item.substituent_tree,
196+
target.name,
197+
)
190198
for item in grouped.values():
191199
if item.trace_segments and strip_outer_parentheses(item.name) != "methyl":
192200
for segment in item.trace_segments:
@@ -305,9 +313,12 @@ def assembly_substituent_tree(
305313
atom_ids=None,
306314
bond_ids=None,
307315
decisions=None,
316+
trace_segments=None,
308317
) -> dict:
309318
"""Return a nested substituent tree from the graph-bound assembly parts."""
310319

320+
if trace_segments is None:
321+
trace_segments = assembly_trace_segments(parts)
311322
component_atoms = set(atom_ids or parts.parent_atom_ids)
312323
component_bonds = set(bond_ids or parts.parent_bond_ids)
313324
parent_node = {
@@ -352,12 +363,32 @@ def assembly_substituent_tree(
352363
}
353364
for item in parts.parent_charges
354365
],
355-
"trace_segments": assembly_trace_segments(parts),
366+
"trace_segments": list(trace_segments),
356367
"nested_decisions": list(decisions or ()),
357368
}
358369
return parent_node
359370

360371

372+
def _merge_substituent_tree_instances(existing: dict | None, new: dict, name: str) -> dict:
373+
"""Preserve all tree instances when same-name substituents are grouped."""
374+
375+
if existing is None:
376+
return new
377+
if existing == new:
378+
merged = dict(existing)
379+
merged["instance_count"] = int(merged.get("instance_count", 1)) + 1
380+
return merged
381+
if existing.get("kind") == "grouped_substituent_instances":
382+
merged = dict(existing)
383+
merged["instances"] = [*existing.get("instances", ()), new]
384+
return merged
385+
return {
386+
"kind": "grouped_substituent_instances",
387+
"name": name,
388+
"instances": [existing, new],
389+
}
390+
391+
361392
def _parent_tree_node(parts: AssemblyParts) -> dict:
362393
return {
363394
"kind": "parent",

0 commit comments

Comments
 (0)