Refactor vertex lookup using a dictionary
This commit is contained in:
parent
b5c27724a6
commit
86ce36426a
2 changed files with 5 additions and 4 deletions
|
|
@ -70,7 +70,7 @@ class Graph:
|
||||||
def _build_graph(self) -> None:
|
def _build_graph(self) -> None:
|
||||||
"""Builds the graph from the vertices and edges."""
|
"""Builds the graph from the vertices and edges."""
|
||||||
self.vertices = self._build_vertices()
|
self.vertices = self._build_vertices()
|
||||||
self.vertex_ids = [vertex.id for vertex in self.vertices]
|
self.vertex_map = {vertex.id: vertex for vertex in self.vertices}
|
||||||
self.edges = self._build_edges()
|
self.edges = self._build_edges()
|
||||||
|
|
||||||
# This is a hack to make sure that the LLM vertex is sent to
|
# This is a hack to make sure that the LLM vertex is sent to
|
||||||
|
|
@ -107,7 +107,7 @@ class Graph:
|
||||||
|
|
||||||
def get_vertex(self, vertex_id: str) -> Union[None, Vertex]:
|
def get_vertex(self, vertex_id: str) -> Union[None, Vertex]:
|
||||||
"""Returns a vertex by id."""
|
"""Returns a vertex by id."""
|
||||||
return next((vertex for vertex in self.vertices if vertex.id == vertex_id), None)
|
return self.vertex_map.get(vertex_id)
|
||||||
|
|
||||||
def get_vertex_edges(self, vertex_id: str) -> List[Edge]:
|
def get_vertex_edges(self, vertex_id: str) -> List[Edge]:
|
||||||
"""Returns a list of edges for a given vertex."""
|
"""Returns a list of edges for a given vertex."""
|
||||||
|
|
@ -249,3 +249,4 @@ class Graph:
|
||||||
vertex_ids = [vertex.id for vertex in self.vertices]
|
vertex_ids = [vertex.id for vertex in self.vertices]
|
||||||
edges_repr = "\n".join([f"{edge.source_id} --> {edge.target_id}" for edge in self.edges])
|
edges_repr = "\n".join([f"{edge.source_id} --> {edge.target_id}" for edge in self.edges])
|
||||||
return f"Graph:\nNodes: {vertex_ids}\nConnections:\n{edges_repr}"
|
return f"Graph:\nNodes: {vertex_ids}\nConnections:\n{edges_repr}"
|
||||||
|
return f"Graph:\nNodes: {vertex_ids}\nConnections:\n{edges_repr}"
|
||||||
|
|
|
||||||
|
|
@ -81,8 +81,8 @@ def test_graph_structure(basic_graph):
|
||||||
assert isinstance(node, Vertex)
|
assert isinstance(node, Vertex)
|
||||||
for edge in basic_graph.edges:
|
for edge in basic_graph.edges:
|
||||||
assert isinstance(edge, Edge)
|
assert isinstance(edge, Edge)
|
||||||
assert edge.source_id in basic_graph.vertex_ids
|
assert edge.source_id in basic_graph.vertex_map.keys()
|
||||||
assert edge.target_id in basic_graph.vertex_ids
|
assert edge.target_id in basic_graph.vertex_map.keys()
|
||||||
|
|
||||||
|
|
||||||
def test_circular_dependencies(basic_graph):
|
def test_circular_dependencies(basic_graph):
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue