|
7 | 7 |
|
8 | 8 | from __future__ import annotations |
9 | 9 |
|
| 10 | +import dataclasses |
10 | 11 | from typing import Annotated, Any |
11 | 12 |
|
12 | 13 | from pydantic import Field |
|
16 | 17 | AnalysisRun, |
17 | 18 | BranchView, |
18 | 19 | BuildTreeResult, |
| 20 | + EditOp, |
19 | 21 | EvpiResult, |
20 | 22 | KeyValue, |
21 | 23 | NodeDiff, |
@@ -149,6 +151,122 @@ def build_tree( |
149 | 151 | ) |
150 | 152 |
|
151 | 153 |
|
| 154 | +def _apply_edits(tree: DecisionTree, edits: list[EditOp]) -> DecisionTree: |
| 155 | + """Return a new tree with the edits applied. Frozen dataclasses are |
| 156 | + rebuilt via dataclasses.replace; unknown ops or missing targets raise |
| 157 | + TreeParseError.""" |
| 158 | + nodes = dict(tree.nodes) |
| 159 | + maximize = tree.maximize |
| 160 | + model_name = tree.model_name |
| 161 | + |
| 162 | + for e in edits: |
| 163 | + op = e.op.lower() |
| 164 | + if op == "set_objective": |
| 165 | + if e.maximize is None: |
| 166 | + raise TreeParseError("set_objective needs `maximize`.") |
| 167 | + maximize = e.maximize |
| 168 | + continue |
| 169 | + |
| 170 | + if not e.node_id or e.node_id not in nodes: |
| 171 | + raise TreeParseError(f"edit {op!r}: no node {e.node_id!r}.") |
| 172 | + node = nodes[e.node_id] |
| 173 | + |
| 174 | + if op == "rename_node": |
| 175 | + if not e.name: |
| 176 | + raise TreeParseError("rename_node needs `name`.") |
| 177 | + nodes[e.node_id] = dataclasses.replace(node, name=e.name) |
| 178 | + elif op == "set_terminal_value": |
| 179 | + if node.kind != "terminal": |
| 180 | + raise TreeParseError(f"{e.node_id!r} is not a terminal node.") |
| 181 | + if e.value is None: |
| 182 | + raise TreeParseError("set_terminal_value needs `value`.") |
| 183 | + nodes[e.node_id] = dataclasses.replace(node, value=e.value) |
| 184 | + elif op in ("set_probability", "set_branch_value", "rename_branch"): |
| 185 | + found = False |
| 186 | + new_branches = [] |
| 187 | + for b in node.branches: |
| 188 | + if b.name == e.branch_name: |
| 189 | + found = True |
| 190 | + if op == "set_probability": |
| 191 | + b = dataclasses.replace(b, probability=e.value) |
| 192 | + elif op == "set_branch_value": |
| 193 | + b = dataclasses.replace(b, value=e.value or 0.0) |
| 194 | + else: # rename_branch |
| 195 | + if not e.name: |
| 196 | + raise TreeParseError("rename_branch needs `name`.") |
| 197 | + b = dataclasses.replace(b, name=e.name) |
| 198 | + new_branches.append(b) |
| 199 | + if not found: |
| 200 | + raise TreeParseError( |
| 201 | + f"node {e.node_id!r} has no branch named {e.branch_name!r}." |
| 202 | + ) |
| 203 | + nodes[e.node_id] = dataclasses.replace(node, branches=new_branches) |
| 204 | + else: |
| 205 | + raise TreeParseError(f"unknown edit op {e.op!r}.") |
| 206 | + |
| 207 | + return dataclasses.replace( |
| 208 | + tree, nodes=nodes, maximize=maximize, model_name=model_name |
| 209 | + ) |
| 210 | + |
| 211 | + |
| 212 | +@mcp.tool( |
| 213 | + description=( |
| 214 | + "ModelChoice: Edit an existing decision tree in place — change " |
| 215 | + "probabilities, branch cash flows, terminal payoffs, node/branch " |
| 216 | + "labels, or the maximize/minimize objective — then re-roll it. Pass " |
| 217 | + "`edits` as a list of operations ('set_probability', 'set_branch_value', " |
| 218 | + "'set_terminal_value', 'rename_node', 'rename_branch', 'set_objective'). " |
| 219 | + "Reads the named tree, applies the edits, validates by rolling back, and " |
| 220 | + "returns the new EV + optimal policy. dry_run=True (default) previews; " |
| 221 | + "dry_run=False writes and re-renders. The 'tweak it by talking' path." |
| 222 | + ) |
| 223 | +) |
| 224 | +def edit_tree( |
| 225 | + edits: list[EditOp], |
| 226 | + tree_name: Annotated[ |
| 227 | + str | None, Field(description="Tree sheet name. Omit for the first tree.") |
| 228 | + ] = None, |
| 229 | + dry_run: bool = True, |
| 230 | + workbook_name: str | None = None, |
| 231 | +) -> BuildTreeResult: |
| 232 | + bridge = get_bridge() |
| 233 | + trees = bridge.list_trees(workbook_name) |
| 234 | + if tree_name is None: |
| 235 | + tree_name = next(iter(trees)) |
| 236 | + if tree_name not in trees: |
| 237 | + raise TreeParseError(f"no tree {tree_name!r}; available: {', '.join(trees)}.") |
| 238 | + |
| 239 | + edited = _apply_edits(parse_model(trees[tree_name]), edits) |
| 240 | + model_json = to_model_json(edited) |
| 241 | + parsed = parse_model(model_json) |
| 242 | + r = rollup(parsed) |
| 243 | + |
| 244 | + direction = "maximize" if edited.maximize else "minimize" |
| 245 | + if r.optimal_path: |
| 246 | + rec = ( |
| 247 | + f"After edits, optimal decision ({direction} EV): take " |
| 248 | + f"{' → '.join(r.optimal_path)}. Expected value {r.expected_value:,.2f}." |
| 249 | + ) |
| 250 | + else: |
| 251 | + rec = f"After edits, expected value {r.expected_value:,.2f} ({direction})." |
| 252 | + |
| 253 | + written = False |
| 254 | + if not dry_run: |
| 255 | + bridge.write_tree(model_json, sheet_name=tree_name, workbook=workbook_name) |
| 256 | + bridge.render_tree(tree_name, workbook_name) |
| 257 | + written = True |
| 258 | + |
| 259 | + return BuildTreeResult( |
| 260 | + written=written, |
| 261 | + sheet=tree_name if written else None, |
| 262 | + node_count=len(parsed.nodes), |
| 263 | + expected_value=r.expected_value, |
| 264 | + optimal_path=r.optimal_path, |
| 265 | + recommendation=rec, |
| 266 | + model_json=model_json, |
| 267 | + ) |
| 268 | + |
| 269 | + |
152 | 270 | def _counts(tree: DecisionTree) -> tuple[int, int, int]: |
153 | 271 | d = sum(1 for n in tree.nodes.values() if n.kind == "decision") |
154 | 272 | c = sum(1 for n in tree.nodes.values() if n.kind == "chance") |
@@ -559,6 +677,7 @@ def read_sheet( |
559 | 677 |
|
560 | 678 | __all__ = [ |
561 | 679 | "build_tree", |
| 680 | + "edit_tree", |
562 | 681 | "get_bridge", |
563 | 682 | "get_tree", |
564 | 683 | "list_trees", |
|
0 commit comments