study-buddy / app /services /graph_algorithm_compiler.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
12.2 kB
from __future__ import annotations
import heapq
import math
import re
from collections import deque
from app.schemas.visual_lesson import (
CompiledGraphAlgorithm,
GraphAlgorithmSpec,
GraphEdgeSpec,
GraphNodeSpec,
GraphSearchBranch,
GraphSearchStep,
)
ALGORITHMS = ("bfs", "dfs", "dijkstra", "astar")
class GraphAlgorithmCompilationError(ValueError):
pass
class GraphAlgorithmCompiler:
@staticmethod
def _default_graph() -> tuple[list[GraphNodeSpec], list[GraphEdgeSpec]]:
positions = {
"A": (90, 290), "B": (270, 120), "C": (270, 450), "D": (500, 80),
"E": (520, 300), "F": (500, 520), "G": (750, 180), "H": (900, 350),
}
edge_rows = [
("A", "B", 4), ("A", "C", 2), ("B", "C", 1), ("B", "D", 5),
("B", "E", 10), ("C", "E", 8), ("C", "F", 4), ("D", "G", 3),
("E", "G", 1), ("E", "H", 5), ("F", "E", 2), ("F", "H", 7), ("G", "H", 2),
]
nodes = [GraphNodeSpec(node_id=node_id, label=node_id, x=x, y=y) for node_id, (x, y) in positions.items()]
edges = [GraphEdgeSpec(edge_id=f"e{index}", source=source, target=target, weight=weight) for index, (source, target, weight) in enumerate(edge_rows)]
return nodes, edges
@staticmethod
def _custom_graph(prompt: str) -> tuple[list[GraphNodeSpec], list[GraphEdgeSpec], bool] | None:
matches = re.findall(r"\b([A-Za-z][A-Za-z0-9_]{0,7})\s*(->|-)\s*([A-Za-z][A-Za-z0-9_]{0,7})\s*(?::|=)\s*(\d+(?:\.\d+)?)", prompt)
if len(matches) < 2:
return None
node_ids = sorted({source for source, _, _, _ in matches} | {target for _, _, target, _ in matches})
if len(node_ids) > 10:
raise GraphAlgorithmCompilationError("Graph visualizations support at most 10 nodes")
directed = "directed" in prompt.lower() or any(arrow == "->" for _, arrow, _, _ in matches)
nodes = []
for index, node_id in enumerate(node_ids):
angle = 2 * math.pi * index / len(node_ids) - math.pi / 2
nodes.append(GraphNodeSpec(node_id=node_id, label=node_id, x=500 + 380 * math.cos(angle), y=300 + 235 * math.sin(angle)))
edges = []
seen: set[tuple[str, str]] = set()
for source, _, target, weight_text in matches:
key = (source, target) if directed else tuple(sorted((source, target)))
if key in seen:
continue
seen.add(key)
weight = float(weight_text)
if not 0 < weight <= 1e6:
raise GraphAlgorithmCompilationError("Graph weights must be positive and at most 1,000,000")
edges.append(GraphEdgeSpec(edge_id=f"e{len(edges)}", source=source, target=target, weight=weight))
return nodes, edges, directed
def parse_prompt(self, prompt: str) -> tuple[GraphAlgorithmSpec, list[str]]:
text = prompt.lower()
custom = self._custom_graph(prompt)
if custom:
nodes, edges, directed = custom
illustrative = False
warnings: list[str] = []
else:
nodes, edges = self._default_graph()
directed = "directed" in text
illustrative = True
warnings = []
if "breadth" in text or re.search(r"\bbfs\b", text):
algorithm = "bfs"
elif "depth" in text or re.search(r"\bdfs\b", text):
algorithm = "dfs"
elif "a*" in text or "a-star" in text or "astar" in text:
algorithm = "astar"
else:
algorithm = "dijkstra"
node_ids = {node.node_id for node in nodes}
start, target = nodes[0].node_id, nodes[-1].node_id
endpoint = re.search(r"\bfrom\s+([A-Za-z][A-Za-z0-9_]{0,7})\s+to\s+([A-Za-z][A-Za-z0-9_]{0,7})\b", prompt, flags=re.IGNORECASE)
if endpoint:
proposed_start, proposed_target = endpoint.group(1), endpoint.group(2)
lookup = {node_id.lower(): node_id for node_id in node_ids}
if proposed_start.lower() not in lookup or proposed_target.lower() not in lookup:
raise GraphAlgorithmCompilationError("The requested start or target node is not present in the graph")
start, target = lookup[proposed_start.lower()], lookup[proposed_target.lower()]
if start == target:
target = nodes[-1].node_id if start != nodes[-1].node_id else nodes[0].node_id
return GraphAlgorithmSpec(
project_id="",
prompt=prompt,
directed=directed,
nodes=nodes,
edges=edges,
primary_algorithm=algorithm,
initial_start_node=start,
initial_target_node=target,
graph_is_illustrative=illustrative,
assumptions=[
"Edge weights are non-negative and remain fixed during a search.",
"Ties are resolved deterministically by node label.",
"A* uses a graph-derived admissible Euclidean heuristic; the other algorithms do not use node positions.",
],
created_at=0.0,
), warnings
@staticmethod
def _adjacency(spec: GraphAlgorithmSpec) -> dict[str, list[tuple[str, float, str]]]:
adjacency = {node.node_id: [] for node in spec.nodes}
for edge in spec.edges:
adjacency[edge.source].append((edge.target, edge.weight, edge.edge_id))
if not spec.directed:
adjacency[edge.target].append((edge.source, edge.weight, edge.edge_id))
for rows in adjacency.values():
rows.sort(key=lambda row: row[0])
return adjacency
@staticmethod
def _path(target: str, parents: dict[str, tuple[str, str]], start: str) -> tuple[list[str], list[str]]:
if target != start and target not in parents:
return [], []
nodes = [target]
edges: list[str] = []
while nodes[-1] != start:
parent, edge_id = parents[nodes[-1]]
nodes.append(parent)
edges.append(edge_id)
return list(reversed(nodes)), list(reversed(edges))
@staticmethod
def _heuristic_scale(spec: GraphAlgorithmSpec) -> float:
positions = {node.node_id: (node.x, node.y) for node in spec.nodes}
ratios = []
for edge in spec.edges:
ax, ay = positions[edge.source]
bx, by = positions[edge.target]
distance = math.hypot(ax - bx, ay - by)
if distance > 0:
ratios.append(edge.weight / distance)
return min(ratios) if ratios else 0.0
def _compile_branch(self, spec: GraphAlgorithmSpec, algorithm: str, start: str, target: str) -> GraphSearchBranch:
adjacency = self._adjacency(spec)
parents: dict[str, tuple[str, str]] = {}
distances: dict[str, float] = {start: 0.0}
visited: list[str] = []
steps = [GraphSearchStep(step_index=0, frontier=[start], distances={start: 0.0}, description=f"Start at {start}.")]
edge_weights = {edge.edge_id: edge.weight for edge in spec.edges}
found = False
if algorithm in {"bfs", "dfs"}:
frontier = deque([start]) if algorithm == "bfs" else [start]
discovered = {start}
while frontier:
current = frontier.popleft() if algorithm == "bfs" else frontier.pop()
if current in visited:
continue
visited.append(current)
active_edges: list[str] = []
if current == target:
found = True
else:
neighbours = adjacency[current] if algorithm == "bfs" else list(reversed(adjacency[current]))
for neighbour, weight, edge_id in neighbours:
if neighbour in discovered:
continue
discovered.add(neighbour)
parents[neighbour] = (current, edge_id)
distances[neighbour] = distances[current] + weight
active_edges.append(edge_id)
frontier.append(neighbour)
snapshot = list(frontier)
parent_edges = [edge_id for _, edge_id in parents.values()]
steps.append(GraphSearchStep(step_index=len(steps), current_node=current, frontier=snapshot, visited=list(visited), distances=dict(distances), parent_edge_ids=parent_edges, active_edge_ids=active_edges, description=f"Visit {current}; add undiscovered neighbours to the {'queue' if algorithm == 'bfs' else 'stack'}."))
if found:
break
else:
positions = {node.node_id: (node.x, node.y) for node in spec.nodes}
scale = self._heuristic_scale(spec) if algorithm == "astar" else 0.0
def heuristic(node_id: str) -> float:
ax, ay = positions[node_id]
bx, by = positions[target]
return math.hypot(ax - bx, ay - by) * scale
heap: list[tuple[float, str]] = [(heuristic(start), start)]
settled: set[str] = set()
while heap:
_, current = heapq.heappop(heap)
if current in settled:
continue
settled.add(current)
visited.append(current)
active_edges = []
if current == target:
found = True
else:
for neighbour, weight, edge_id in adjacency[current]:
candidate = distances[current] + weight
if candidate + 1e-12 < distances.get(neighbour, math.inf):
distances[neighbour] = candidate
parents[neighbour] = (current, edge_id)
heapq.heappush(heap, (candidate + heuristic(neighbour), neighbour))
active_edges.append(edge_id)
frontier_nodes = sorted({node_id for _, node_id in heap if node_id not in settled}, key=lambda node_id: (distances.get(node_id, math.inf) + heuristic(node_id), node_id))
steps.append(GraphSearchStep(step_index=len(steps), current_node=current, frontier=frontier_nodes, visited=list(visited), distances=dict(distances), parent_edge_ids=[edge_id for _, edge_id in parents.values()], active_edge_ids=active_edges, description=f"Settle {current}; relax outgoing edges with lower tentative cost."))
if found:
break
path, path_edges = self._path(target, parents, start)
if start == target:
found, path = True, [start]
cost = sum(edge_weights[edge_id] for edge_id in path_edges)
steps.append(GraphSearchStep(step_index=len(steps), current_node=target if found else "", frontier=[], visited=list(visited), distances=dict(distances), parent_edge_ids=[edge_id for _, edge_id in parents.values()], active_edge_ids=path_edges, final_path=path, description=f"Path found with cost {cost:g}." if found else "No path connects the selected nodes."))
return GraphSearchBranch(algorithm=algorithm, start_node=start, target_node=target, found=found, path=path, path_cost=cost, steps=steps)
def compile_spec(self, spec: GraphAlgorithmSpec) -> CompiledGraphAlgorithm:
node_ids = [node.node_id for node in spec.nodes]
branches = [self._compile_branch(spec, algorithm, start, target) for algorithm in ALGORITHMS for start in node_ids for target in node_ids]
index = {(branch.algorithm, branch.start_node, branch.target_node): branch for branch in branches}
for start in node_ids:
for target in node_ids:
dijkstra = index[("dijkstra", start, target)]
astar = index[("astar", start, target)]
if dijkstra.found != astar.found or abs(dijkstra.path_cost - astar.path_cost) > 1e-8:
raise GraphAlgorithmCompilationError("A* and Dijkstra disagree on the optimal path cost")
return CompiledGraphAlgorithm(branches=branches, assertions_passed=True)