Merge remote-tracking branch 'origin/zustand/io/migration' into globalVariables

This commit is contained in:
Lucas Oliveira 2024-03-21 15:35:52 +01:00
commit 99561d0d4c
26 changed files with 1381 additions and 887 deletions

View file

@ -13,6 +13,7 @@ from langflow.services.store.schema import StoreComponentCreate
from langflow.services.store.utils import get_lf_version_from_pypi
if TYPE_CHECKING:
from langflow.graph.vertex.base import Vertex
from langflow.services.database.models.flow.model import Flow
@ -238,3 +239,62 @@ def format_exception_message(exc: Exception) -> str:
if isinstance(causing_exception, SyntaxError):
return format_syntax_error_message(causing_exception)
return str(exc)
async def get_next_runnable_vertices(
graph: Graph,
vertex: "Vertex",
vertex_id: str,
chat_service: ChatService,
flow_id: str,
):
"""
Retrieves the next runnable vertices in the graph for a given vertex.
Args:
graph (Graph): The graph object representing the flow.
vertex (Vertex): The current vertex.
vertex_id (str): The ID of the current vertex.
chat_service (ChatService): The chat service object.
flow_id (str): The ID of the flow.
Returns:
list: A list of IDs of the next runnable vertices.
"""
async with chat_service._cache_locks[flow_id] as lock:
graph.remove_from_predecessors(vertex_id)
direct_successors_ready = [v for v in vertex.successors_ids if graph.is_vertex_runnable(v)]
if not direct_successors_ready:
# No direct successors ready, look for runnable predecessors of successors
next_runnable_vertices = graph.find_runnable_predecessors_for_successors(vertex_id)
else:
next_runnable_vertices = direct_successors_ready
for v_id in set(next_runnable_vertices): # Use set to avoid duplicates
graph.vertices_to_run.remove(v_id)
graph.remove_from_predecessors(v_id)
await chat_service.set_cache(flow_id=flow_id, data=graph, lock=lock)
return next_runnable_vertices
def get_top_level_vertices(graph, vertices_ids):
"""
Retrieves the top-level vertices from the given graph based on the provided vertex IDs.
Args:
graph (Graph): The graph object containing the vertices.
vertices_ids (list): A list of vertex IDs.
Returns:
list: A list of top-level vertex IDs.
"""
top_level_vertices = []
for vertex_id in vertices_ids:
vertex = graph.get_vertex(vertex_id)
if vertex.parent_is_top_level:
top_level_vertices.append(vertex.parent_node_id)
else:
top_level_vertices.append(vertex_id)
return top_level_vertices

View file

@ -10,6 +10,8 @@ from langflow.api.utils import (
build_and_cache_graph,
format_elapsed_time,
format_exception_message,
get_next_runnable_vertices,
get_top_level_vertices,
)
from langflow.api.v1.schemas import (
InputValueRequest,
@ -95,7 +97,8 @@ async def build_vertex(
"""Build a vertex instead of the entire graph."""
start_time = time.perf_counter()
next_vertices_ids = []
next_runnable_vertices = []
top_level_vertices = []
try:
start_time = time.perf_counter()
cache = await chat_service.get_cache(flow_id)
@ -121,12 +124,9 @@ async def build_vertex(
artifacts = vertex.artifacts
else:
raise ValueError(f"No result found for vertex {vertex_id}")
async with chat_service._cache_locks[flow_id] as lock:
graph.remove_from_predecessors(vertex_id)
next_vertices_ids = vertex.successors_ids
next_vertices_ids = [v for v in next_vertices_ids if graph.should_run_vertex(v)]
await chat_service.set_cache(flow_id=flow_id, data=graph, lock=lock)
next_runnable_vertices = await get_next_runnable_vertices(graph, vertex, vertex_id, chat_service, flow_id)
top_level_vertices = get_top_level_vertices(graph, next_runnable_vertices)
result_data_response = ResultDataResponse(**result_dict.model_dump())
except Exception as exc:
@ -166,12 +166,13 @@ async def build_vertex(
# to stop the build of the graph at a certain vertex
# if it is in next_vertices_ids, we need to remove other
# vertices from next_vertices_ids
if graph.stop_vertex and graph.stop_vertex in next_vertices_ids:
next_vertices_ids = [graph.stop_vertex]
if graph.stop_vertex and graph.stop_vertex in next_runnable_vertices:
next_runnable_vertices = [graph.stop_vertex]
build_response = VertexBuildResponse(
inactivated_vertices=inactivated_vertices,
next_vertices_ids=next_vertices_ids,
next_vertices_ids=next_runnable_vertices,
top_level_vertices=top_level_vertices,
valid=valid,
params=params,
id=vertex.id,
@ -201,7 +202,7 @@ async def build_vertex_stream(
async def stream_vertex():
try:
if not session_id:
cache = chat_service.get_cache(flow_id)
cache = await chat_service.get_cache(flow_id)
if not cache:
# If there's no cache
raise ValueError(f"No cache found for {flow_id}.")
@ -251,7 +252,7 @@ async def build_vertex_stream(
raise ValueError(f"No result found for vertex {vertex_id}")
except Exception as exc:
logger.error(f"Error building vertex: {exc}")
logger.exception(f"Error building vertex: {exc}")
yield str(StreamData(event="error", data={"error": str(exc)}))
finally:
logger.debug("Closing stream")

View file

@ -247,6 +247,7 @@ class VertexBuildResponse(BaseModel):
id: Optional[str] = None
inactivated_vertices: Optional[List[str]] = None
next_vertices_ids: Optional[List[str]] = None
top_level_vertices: Optional[List[str]] = None
valid: bool
params: Optional[Any] = Field(default_factory=dict)
"""JSON string of the params."""

View file

@ -44,13 +44,13 @@ def post_validate_prompt(prompt_request: ValidatePromptRequest):
input_variables=input_variables,
frontend_node=None,
)
if not prompt_request.custom_fields:
prompt_request.custom_fields = defaultdict(list)
old_custom_fields = get_old_custom_fields(prompt_request.custom_fields, prompt_request.name)
if not prompt_request.frontend_node.custom_fields:
prompt_request.frontend_node.custom_fields = defaultdict(list)
old_custom_fields = get_old_custom_fields(prompt_request.frontend_node.custom_fields, prompt_request.name)
add_new_variables_to_template(
input_variables,
prompt_request.custom_fields,
prompt_request.frontend_node.custom_fields,
prompt_request.frontend_node.template,
prompt_request.name,
)
@ -58,13 +58,25 @@ def post_validate_prompt(prompt_request: ValidatePromptRequest):
remove_old_variables_from_template(
old_custom_fields,
input_variables,
prompt_request.custom_fields,
prompt_request.frontend_node.custom_fields,
prompt_request.frontend_node.template,
prompt_request.name,
)
update_input_variables_field(input_variables, prompt_request.frontend_node.template)
# If frontend_node.template contains only one field that is type == 'prompt', then we can remove all fields that are not
# 'code', and not in the input_variables list.
prompt_fields = [
key
for key, field in prompt_request.frontend_node.template.items()
if isinstance(field, dict) and field["type"] == "prompt"
]
if len(prompt_fields) == 1:
for key, field in prompt_request.frontend_node.template.copy().items():
if isinstance(field, dict) and field["type"] != "code" and key not in input_variables + prompt_fields:
del prompt_request.frontend_node.template[key]
return PromptValidationResponse(
input_variables=input_variables,
frontend_node=prompt_request.frontend_node,

View file

@ -1,6 +1,6 @@
from typing import Optional
from langchain_community.chat_models.anthropic import ChatAnthropic
from langchain_anthropic.chat_models import ChatAnthropic
from pydantic.v1 import SecretStr
from langflow.components.models.base.model import LCModelComponent

View file

@ -35,6 +35,14 @@ class Graph:
edges: List[Dict[str, str]],
flow_id: Optional[str] = None,
) -> None:
"""
Initializes a new instance of the Graph class.
Args:
nodes (List[Dict]): A list of dictionaries representing the vertices of the graph.
edges (List[Dict[str, str]]): A list of dictionaries representing the edges of the graph.
flow_id (Optional[str], optional): The ID of the flow. Defaults to None.
"""
self._vertices = nodes
self._edges = edges
self.raw_graph_data = {"nodes": nodes, "edges": edges}
@ -71,11 +79,26 @@ class Graph:
self.state_manager = GraphStateManager()
def get_state(self, name: str) -> Optional[Record]:
"""Returns the state of the graph."""
"""
Returns the state of the graph with the given name.
Args:
name (str): The name of the state.
Returns:
Optional[Record]: The state record, or None if the state does not exist.
"""
return self.state_manager.get_state(name, run_id=self._run_id)
def update_state(self, name: str, record: Union[str, Record], caller: Optional[str] = None) -> None:
"""Updates the state of the graph."""
"""
Updates the state of the graph with the given name.
Args:
name (str): The name of the state.
record (Union[str, Record]): The new state record.
caller (Optional[str], optional): The ID of the vertex that is updating the state. Defaults to None.
"""
if caller:
# If there is a caller which is a vertex_id, I want to activate
# all StateVertex in self.vertices that are not the caller
@ -86,6 +109,13 @@ class Graph:
self.state_manager.update_state(name, record, run_id=self._run_id)
def activate_state_vertices(self, name: str, caller: str):
"""
Activates the state vertices in the graph with the given name and caller.
Args:
name (str): The name of the state.
caller (str): The ID of the vertex that is updating the state.
"""
vertices_ids = []
for vertex_id in self._is_state_vertices:
if vertex_id == caller:
@ -104,10 +134,20 @@ class Graph:
self.vertices_to_run.update(vertices_ids)
def reset_activated_vertices(self):
"""
Resets the activated vertices in the graph.
"""
self.activated_vertices = []
def append_state(self, name: str, record: Union[str, Record], caller: Optional[str] = None) -> None:
"""Appends the state of the graph."""
"""
Appends the state of the graph with the given name.
Args:
name (str): The name of the state.
record (Union[str, Record]): The state record to append.
caller (Optional[str], optional): The ID of the vertex that is updating the state. Defaults to None.
"""
if caller:
self.activate_state_vertices(name, caller)
@ -115,17 +155,38 @@ class Graph:
@property
def run_id(self):
"""
The ID of the current run.
Returns:
str: The run ID.
Raises:
ValueError: If the run ID is not set.
"""
if not self._run_id:
raise ValueError("Run ID not set")
return self._run_id
def set_run_id(self, run_id: str):
"""
Sets the ID of the current run.
Args:
run_id (str): The run ID.
"""
for vertex in self.vertices:
self.state_manager.subscribe(run_id, vertex.update_graph_state)
self._run_id = run_id
@property
def sorted_vertices_layers(self) -> List[List[str]]:
"""
The sorted layers of vertices in the graph.
Returns:
List[List[str]]: The sorted layers of vertices.
"""
if not self._sorted_vertices_layers:
self.sort_vertices()
return self._sorted_vertices_layers
@ -148,7 +209,19 @@ class Graph:
stream: bool,
session_id: str,
) -> List[Optional["ResultData"]]:
"""Runs the graph with the given inputs."""
"""
Runs the graph with the given inputs.
Args:
inputs (Dict[str, str]): The input values for the graph.
input_components (list[str]): The components to run for the inputs.
outputs (list[str]): The outputs to retrieve from the graph.
stream (bool): Whether to stream the results or not.
session_id (str): The session ID for the graph.
Returns:
List[Optional["ResultData"]]: The outputs of the graph.
"""
for vertex_id in self._is_input_vertices:
vertex = self.get_vertex(vertex_id)
if input_components and (vertex_id not in input_components or vertex.display_name not in input_components):
@ -190,7 +263,19 @@ class Graph:
session_id: Optional[str] = None,
stream: bool = False,
) -> List[RunOutputs]:
"""Runs the graph with the given inputs."""
"""
Runs the graph with the given inputs.
Args:
inputs (list[Dict[str, str]]): The input values for the graph.
inputs_components (Optional[list[list[str]]], optional): The components to run for the inputs. Defaults to None.
outputs (Optional[list[str]], optional): The outputs to retrieve from the graph. Defaults to None.
session_id (Optional[str], optional): The session ID for the graph. Defaults to None.
stream (bool, optional): Whether to stream the results or not. Defaults to False.
Returns:
List[RunOutputs]: The outputs of the graph.
"""
# inputs is {"message": "Hello, world!"}
# we need to go through self.inputs and update the self._raw_params
# of the vertices that are inputs
@ -218,16 +303,23 @@ class Graph:
vertex_outputs.append(run_output_object)
return vertex_outputs
# vertices_layers is a list of lists ordered by the order the vertices
# should be built.
# We need to create a new method that will take the vertices_layers
# and return the next vertex to be built.
def next_vertex_to_build(self):
"""Returns the next vertex to be built."""
"""
Returns the next vertex to be built.
Yields:
str: The ID of the next vertex to be built.
"""
yield from chain.from_iterable(self.vertices_layers)
@property
def metadata(self):
"""
The metadata of the graph.
Returns:
dict: The metadata of the graph.
"""
return {
"runs": self._runs,
"updates": self._updates,
@ -235,12 +327,19 @@ class Graph:
}
def build_graph_maps(self):
"""
Builds the adjacency maps for the graph.
"""
self.predecessor_map, self.successor_map = self.build_adjacency_maps()
self.in_degree_map = self.build_in_degree()
self.parent_child_map = self.build_parent_child_map()
def reset_inactivated_vertices(self):
"""
Resets the inactivated vertices in the graph.
"""
self.inactivated_vertices = []
self.inactivated_vertices = set()
def mark_all_vertices(self, state: str):
@ -377,7 +476,10 @@ class Graph:
# Remove vertices that are not in the other graph
for vertex_id in removed_vertex_ids:
self.remove_vertex(vertex_id)
try:
self.remove_vertex(vertex_id)
except ValueError:
pass
# The order here matters because adding the vertex is required
# if any of them have edges that point to any of the new vertices
@ -741,8 +843,11 @@ class Graph:
vertex_data = vertex["data"]
vertex_type: str = vertex_data["type"] # type: ignore
vertex_base_type: str = vertex_data["node"]["template"]["_type"] # type: ignore
if "id" not in vertex_data:
raise ValueError(f"Vertex data for {vertex_data['display_name']} does not contain an id")
VertexClass = self._get_vertex_class(vertex_type, vertex_base_type, vertex_data["id"])
vertex_instance = VertexClass(vertex, graph=self)
vertex_instance.set_top_level(self.top_level_vertices)
vertices.append(vertex_instance)
@ -953,22 +1058,26 @@ class Graph:
# Return just the first layer
return first_layer
def vertex_has_no_more_predecessors(self, vertex_id: str) -> bool:
"""Returns whether a vertex has no more predecessors."""
return not self.run_predecessors.get(vertex_id)
def is_vertex_runnable(self, vertex_id: str) -> bool:
"""Returns whether a vertex is runnable."""
return vertex_id in self.vertices_to_run and not self.run_predecessors.get(vertex_id)
def should_run_vertex(self, vertex_id: str) -> bool:
"""Returns whether a component should be run."""
# the self.run_map is a map of vertex_id to a list of predecessors
# each time a vertex is run, we remove it from the list of predecessors
# if a vertex has no more predecessors, it should be run
should_run = vertex_id in self.vertices_to_run and self.vertex_has_no_more_predecessors(vertex_id)
def find_runnable_predecessors_for_successors(self, vertex_id: str) -> List[str]:
"""
For each successor of the current vertex, find runnable predecessors if any.
This checks the direct predecessors of each successor to identify any that are
immediately runnable, expanding the search to ensure progress can be made.
"""
runnable_vertices = []
visited = set()
if should_run:
self.vertices_to_run.remove(vertex_id)
# remove the vertex from the run_map
self.remove_from_predecessors(vertex_id)
return should_run
for successor_id in self.run_map.get(vertex_id, []):
for predecessor_id in self.run_predecessors.get(successor_id, []):
if predecessor_id not in visited and self.is_vertex_runnable(predecessor_id):
runnable_vertices.append(predecessor_id)
visited.add(predecessor_id)
return runnable_vertices
def remove_from_predecessors(self, vertex_id: str):
predecessors = self.run_map.get(vertex_id, [])