ref: Fix some ruff rules for private access (SLF) (#4139)

* Add ruff rules for private access (SLF)

* Changes following review

* Rename Vertex._data to Vertex.full_data
This commit is contained in:
Christophe Bornet 2024-10-23 01:04:12 +02:00 • committed by GitHub
commit f44aca5b41
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 199 additions and 199 deletions

View file

@ -143,12 +143,12 @@ def format_elapsed_time(elapsed_time: float) -> str:
async def build_graph_from_data(flow_id: str, payload: dict, **kwargs): async def build_graph_from_data(flow_id: str, payload: dict, **kwargs):
"""Build and cache the graph.""" """Build and cache the graph."""
graph = Graph.from_payload(payload, flow_id, **kwargs) graph = Graph.from_payload(payload, flow_id, **kwargs)
for vertex_id in graph._has_session_id_vertices: for vertex_id in graph.has_session_id_vertices:
vertex = graph.get_vertex(vertex_id) vertex = graph.get_vertex(vertex_id)
if vertex is None: if vertex is None:
msg = f"Vertex {vertex_id} not found" msg = f"Vertex {vertex_id} not found"
raise ValueError(msg) raise ValueError(msg)
if not vertex._raw_params.get("session_id"): if not vertex.raw_params.get("session_id"):
vertex.update_raw_params({"session_id": flow_id}, overwrite=True) vertex.update_raw_params({"session_id": flow_id}, overwrite=True)
run_id = uuid.uuid4() run_id = uuid.uuid4()

View file

@ -219,7 +219,7 @@ async def build_flow(
try: try:
vertex = graph.get_vertex(vertex_id) vertex = graph.get_vertex(vertex_id)
try: try:
lock = chat_service._async_cache_locks[flow_id_str] lock = chat_service.async_cache_locks[flow_id_str]
vertex_build_result = await graph.build_vertex( vertex_build_result = await graph.build_vertex(
vertex_id=vertex_id, vertex_id=vertex_id,
user_id=str(current_user.id), user_id=str(current_user.id),
@ -505,7 +505,7 @@ async def build_vertex(
vertex = graph.get_vertex(vertex_id) vertex = graph.get_vertex(vertex_id)
try: try:
lock = chat_service._async_cache_locks[flow_id_str] lock = chat_service.async_cache_locks[flow_id_str]
vertex_build_result = await graph.build_vertex( vertex_build_result = await graph.build_vertex(
vertex_id=vertex_id, vertex_id=vertex_id,
user_id=str(current_user.id), user_id=str(current_user.id),
@ -643,7 +643,7 @@ async def _stream_vertex(flow_id: str, vertex_id: str, chat_service: ChatService
yield str(StreamData(event="error", data={"error": msg})) yield str(StreamData(event="error", data={"error": msg}))
return return
if isinstance(vertex._built_result, str) and vertex._built_result: if isinstance(vertex.built_result, str) and vertex.built_result:
stream_data = StreamData( stream_data = StreamData(
event="message", event="message",
data={"message": f"Streaming vertex {vertex_id}"}, data={"message": f"Streaming vertex {vertex_id}"},
@ -651,11 +651,11 @@ async def _stream_vertex(flow_id: str, vertex_id: str, chat_service: ChatService
yield str(stream_data) yield str(stream_data)
stream_data = StreamData( stream_data = StreamData(
event="message", event="message",
data={"chunk": vertex._built_result}, data={"chunk": vertex.built_result},
) )
yield str(stream_data) yield str(stream_data)
elif not vertex.frozen or not vertex._built: elif not vertex.frozen or not vertex.built:
logger.debug(f"Streaming vertex {vertex_id}") logger.debug(f"Streaming vertex {vertex_id}")
stream_data = StreamData( stream_data = StreamData(
event="message", event="message",
@ -678,7 +678,7 @@ async def _stream_vertex(flow_id: str, vertex_id: str, chat_service: ChatService
elif vertex.result is not None: elif vertex.result is not None:
stream_data = StreamData( stream_data = StreamData(
event="message", event="message",
data={"chunk": vertex._built_result}, data={"chunk": vertex.built_result},
) )
yield str(stream_data) yield str(stream_data)
else: else:

View file

@ -44,7 +44,7 @@ class ChatComponent(Component):
def _update_stored_message(self, message_id: str, complete_message: str) -> Message: def _update_stored_message(self, message_id: str, complete_message: str) -> Message:
message_table = update_message(message_id=message_id, message={"text": complete_message}) message_table = update_message(message_id=message_id, message={"text": complete_message})
updated_message = Message(**message_table.model_dump()) updated_message = Message(**message_table.model_dump())
self.vertex._added_message = updated_message self.vertex.added_message = updated_message
return updated_message return updated_message
def _process_chunk(self, chunk: str, complete_message: str, message: Message, message_id: str) -> str: def _process_chunk(self, chunk: str, complete_message: str, message: Message, message_id: str) -> str:

View file

@ -32,7 +32,7 @@ def check_cached_vector_store(f):
self._cached_vector_store = result self._cached_vector_store = result
return result return result
check_cached._is_cached_vector_store_checked = True check_cached.is_cached_vector_store_checked = True
return check_cached return check_cached
@ -45,7 +45,7 @@ class LCVectorStoreComponent(Component):
super().__init_subclass__(**kwargs) super().__init_subclass__(**kwargs)
if hasattr(cls, "build_vector_store"): if hasattr(cls, "build_vector_store"):
method = cls.build_vector_store method = cls.build_vector_store
if not hasattr(method, "_is_cached_vector_store_checked"): if not hasattr(method, "is_cached_vector_store_checked"):
msg = ( msg = (
f"The method 'build_vector_store' in class {cls.__name__} " f"The method 'build_vector_store' in class {cls.__name__} "
"must be decorated with @check_cached_vector_store" "must be decorated with @check_cached_vector_store"

View file

@ -30,8 +30,8 @@ class Edge:
except Exception as e: except Exception as e:
if "inputTypes" in self._target_handle and self._target_handle["inputTypes"] is None: if "inputTypes" in self._target_handle and self._target_handle["inputTypes"] is None:
# Check if self._target_handle['fieldName'] # Check if self._target_handle['fieldName']
if hasattr(target, "_custom_component"): if hasattr(target, "custom_component"):
display_name = getattr(target._custom_component, "display_name", "") display_name = getattr(target.custom_component, "display_name", "")
msg = ( msg = (
f"Component {display_name} field '{self._target_handle['fieldName']}' " f"Component {display_name} field '{self._target_handle['fieldName']}' "
"might not be a valid input." "might not be a valid input."
@ -214,8 +214,8 @@ class CycleEdge(Edge):
self.is_fulfilled = False # Whether the contract has been fulfilled. self.is_fulfilled = False # Whether the contract has been fulfilled.
self.result: Any = None self.result: Any = None
self.is_cycle = True self.is_cycle = True
source._has_cycle_edges = True source.has_cycle_edges = True
target._has_cycle_edges = True target.has_cycle_edges = True
async def honor(self, source: Vertex, target: Vertex) -> None: async def honor(self, source: Vertex, target: Vertex) -> None:
"""Fulfills the contract by setting the result of the source vertex to the target vertex's parameter. """Fulfills the contract by setting the result of the source vertex to the target vertex's parameter.
@ -228,16 +228,16 @@ class CycleEdge(Edge):
if self.is_fulfilled: if self.is_fulfilled:
return return
if not source._built: if not source.built:
# The system should be read-only, so we should not be building vertices # The system should be read-only, so we should not be building vertices
# that are not already built. # that are not already built.
msg = f"Source vertex {source.id} is not built." msg = f"Source vertex {source.id} is not built."
raise ValueError(msg) raise ValueError(msg)
if self.matched_type == "Text": if self.matched_type == "Text":
self.result = source._built_result self.result = source.built_result
else: else:
self.result = source._built_object self.result = source.built_object
target.params[self.target_param] = self.result target.params[self.target_param] = self.result
self.is_fulfilled = True self.is_fulfilled = True

View file

@ -91,7 +91,7 @@ class Graph:
self._is_input_vertices: list[str] = [] self._is_input_vertices: list[str] = []
self._is_output_vertices: list[str] = [] self._is_output_vertices: list[str] = []
self._is_state_vertices: list[str] = [] self._is_state_vertices: list[str] = []
self._has_session_id_vertices: list[str] = [] self.has_session_id_vertices: list[str] = []
self._sorted_vertices_layers: list[list[str]] = [] self._sorted_vertices_layers: list[list[str]] = []
self._run_id = "" self._run_id = ""
self._start_time = datetime.now(timezone.utc) self._start_time = datetime.now(timezone.utc)
@ -258,10 +258,10 @@ class Graph:
msg = f"Target vertex {target_id} is not a component vertex." msg = f"Target vertex {target_id} is not a component vertex."
raise TypeError(msg) raise TypeError(msg)
output_name, input_name = output_input_tuple output_name, input_name = output_input_tuple
if source_vertex._custom_component is None: if source_vertex.custom_component is None:
msg = f"Source vertex {source_id} does not have a custom component." msg = f"Source vertex {source_id} does not have a custom component."
raise ValueError(msg) raise ValueError(msg)
if target_vertex._custom_component is None: if target_vertex.custom_component is None:
msg = f"Target vertex {target_id} does not have a custom component." msg = f"Target vertex {target_id} does not have a custom component."
raise ValueError(msg) raise ValueError(msg)
@ -282,8 +282,8 @@ class Graph:
"target": target_id, "target": target_id,
"data": { "data": {
"sourceHandle": { "sourceHandle": {
"dataType": source_vertex._custom_component.name "dataType": source_vertex.custom_component.name
or source_vertex._custom_component.__class__.__name__, or source_vertex.custom_component.__class__.__name__,
"id": source_vertex.id, "id": source_vertex.id,
"name": output_name, "name": output_name,
"output_types": source_vertex.get_output(output_name).types, "output_types": source_vertex.get_output(output_name).types,
@ -339,9 +339,9 @@ class Graph:
def __apply_config(self, config: StartConfigDict) -> None: def __apply_config(self, config: StartConfigDict) -> None:
for vertex in self.vertices: for vertex in self.vertices:
if vertex._custom_component is None: if vertex.custom_component is None:
continue continue
for output in vertex._custom_component._outputs_map.values(): for output in vertex.custom_component._outputs_map.values():
for key, value in config["output"].items(): for key, value in config["output"].items():
setattr(output, key, value) setattr(output, key, value)
@ -439,8 +439,8 @@ class Graph:
if vertex_id == caller or vertex.display_name == caller_vertex.display_name: if vertex_id == caller or vertex.display_name == caller_vertex.display_name:
continue continue
if ( if (
isinstance(vertex._raw_params["name"], str) isinstance(vertex.raw_params["name"], str)
and name in vertex._raw_params["name"] and name in vertex.raw_params["name"]
and vertex_id != caller and vertex_id != caller
and isinstance(vertex, StateVertex) and isinstance(vertex, StateVertex)
): ):
@ -604,11 +604,15 @@ class Graph:
def define_vertices_lists(self) -> None: def define_vertices_lists(self) -> None:
"""Defines the lists of vertices that are inputs, outputs, and have session_id.""" """Defines the lists of vertices that are inputs, outputs, and have session_id."""
attributes = ["is_input", "is_output", "has_session_id", "is_state"]
for vertex in self.vertices: for vertex in self.vertices:
for attribute in attributes: if vertex.is_input:
if getattr(vertex, attribute): self._is_input_vertices.append(vertex.id)
getattr(self, f"_{attribute}_vertices").append(vertex.id) if vertex.is_output:
self._is_output_vertices.append(vertex.id)
if vertex.has_session_id:
self.has_session_id_vertices.append(vertex.id)
if vertex.is_state:
self._is_state_vertices.append(vertex.id)
def _set_inputs(self, input_components: list[str], inputs: dict[str, str], input_type: InputType | None) -> None: def _set_inputs(self, input_components: list[str], inputs: dict[str, str], input_type: InputType | None) -> None:
for vertex_id in self._is_input_vertices: for vertex_id in self._is_input_vertices:
@ -662,7 +666,7 @@ class Graph:
if inputs: if inputs:
self._set_inputs(input_components, inputs, input_type) self._set_inputs(input_components, inputs, input_type)
# Update all the vertices with the session_id # Update all the vertices with the session_id
for vertex_id in self._has_session_id_vertices: for vertex_id in self.has_session_id_vertices:
vertex = self.get_vertex(vertex_id) vertex = self.get_vertex(vertex_id)
if vertex is None: if vertex is None:
msg = f"Vertex {vertex_id} not found" msg = f"Vertex {vertex_id} not found"
@ -690,7 +694,7 @@ class Graph:
# Get the outputs # Get the outputs
vertex_outputs = [] vertex_outputs = []
for vertex in self.vertices: for vertex in self.vertices:
if not vertex._built: if not vertex.built:
continue continue
if vertex is None: if vertex is None:
msg = f"Vertex {vertex_id} not found" msg = f"Vertex {vertex_id} not found"
@ -781,7 +785,7 @@ class Graph:
List[RunOutputs]: The outputs of the graph. List[RunOutputs]: The outputs of the graph.
""" """
# inputs is {"message": "Hello, world!"} # inputs is {"message": "Hello, world!"}
# we need to go through self.inputs and update the self._raw_params # we need to go through self.inputs and update the self.raw_params
# of the vertices that are inputs # of the vertices that are inputs
# if the value is a list, we need to run multiple times # if the value is a list, we need to run multiple times
vertex_outputs = [] vertex_outputs = []
@ -948,7 +952,7 @@ class Graph:
"_edges": self._edges, "_edges": self._edges,
"_is_input_vertices": self._is_input_vertices, "_is_input_vertices": self._is_input_vertices,
"_is_output_vertices": self._is_output_vertices, "_is_output_vertices": self._is_output_vertices,
"_has_session_id_vertices": self._has_session_id_vertices, "has_session_id_vertices": self.has_session_id_vertices,
"_sorted_vertices_layers": self._sorted_vertices_layers, "_sorted_vertices_layers": self._sorted_vertices_layers,
} }
@ -1123,17 +1127,17 @@ class Graph:
vertex (Vertex): The vertex to be updated. vertex (Vertex): The vertex to be updated.
other_vertex (Vertex): The vertex to update from. other_vertex (Vertex): The vertex to update from.
""" """
vertex._data = other_vertex._data vertex.full_data = other_vertex.full_data
vertex._parse_data() vertex.parse_data()
# Now we update the edges of the vertex # Now we update the edges of the vertex
self.update_edges_from_vertex(other_vertex) self.update_edges_from_vertex(other_vertex)
vertex.params = {} vertex.params = {}
vertex._build_params() vertex.build_params()
vertex.graph = self vertex.graph = self
# If the vertex is frozen, we don't want # If the vertex is frozen, we don't want
# to reset the results nor the _built attribute # to reset the results nor the built attribute
if not vertex.frozen: if not vertex.frozen:
vertex._built = False vertex.built = False
vertex.result = None vertex.result = None
vertex.artifacts = {} vertex.artifacts = {}
vertex.set_top_level(self.top_level_vertices) vertex.set_top_level(self.top_level_vertices)
@ -1146,7 +1150,7 @@ class Graph:
if vid in self.vertex_map: if vid in self.vertex_map:
_vertex = self.vertex_map[vid] _vertex = self.vertex_map[vid]
if not _vertex.frozen: if not _vertex.frozen:
_vertex._build_params() _vertex.build_params()
def _add_vertex(self, vertex: Vertex) -> None: def _add_vertex(self, vertex: Vertex) -> None:
"""Adds a vertex to the graph.""" """Adds a vertex to the graph."""
@ -1206,7 +1210,7 @@ class Graph:
def _build_vertex_params(self) -> None: def _build_vertex_params(self) -> None:
"""Identifies and handles the LLM vertex within the graph.""" """Identifies and handles the LLM vertex within the graph."""
for vertex in self.vertices: for vertex in self.vertices:
vertex._build_params() vertex.build_params()
def _validate_vertex(self, vertex: Vertex) -> bool: def _validate_vertex(self, vertex: Vertex) -> bool:
"""Validates a vertex.""" """Validates a vertex."""
@ -1359,14 +1363,14 @@ class Graph:
try: try:
cached_vertex_dict = cached_result["result"] cached_vertex_dict = cached_result["result"]
# Now set update the vertex with the cached vertex # Now set update the vertex with the cached vertex
vertex._built = cached_vertex_dict["_built"] vertex.built = cached_vertex_dict["built"]
vertex.artifacts = cached_vertex_dict["artifacts"] vertex.artifacts = cached_vertex_dict["artifacts"]
vertex._built_object = cached_vertex_dict["_built_object"] vertex.built_object = cached_vertex_dict["built_object"]
vertex._built_result = cached_vertex_dict["_built_result"] vertex.built_result = cached_vertex_dict["built_result"]
vertex._data = cached_vertex_dict["_data"] vertex.full_data = cached_vertex_dict["full_data"]
vertex.results = cached_vertex_dict["results"] vertex.results = cached_vertex_dict["results"]
try: try:
vertex._finalize_build() vertex.finalize_build()
if vertex.result is not None: if vertex.result is not None:
vertex.result.used_frozen_result = True vertex.result.used_frozen_result = True
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
@ -1385,12 +1389,12 @@ class Graph:
) )
if set_cache is not None: if set_cache is not None:
vertex_dict = { vertex_dict = {
"_built": vertex._built, "built": vertex.built,
"results": vertex.results, "results": vertex.results,
"artifacts": vertex.artifacts, "artifacts": vertex.artifacts,
"_built_object": vertex._built_object, "built_object": vertex.built_object,
"_built_result": vertex._built_result, "built_result": vertex.built_result,
"_data": vertex._data, "full_data": vertex.full_data,
} }
await set_cache(key=vertex.id, data=vertex_dict) await set_cache(key=vertex.id, data=vertex_dict)
@ -1401,7 +1405,7 @@ class Graph:
raise raise
if vertex.result is not None: if vertex.result is not None:
params = f"{vertex._built_object_repr()}{params}" params = f"{vertex.built_object_repr()}{params}"
valid = True valid = True
result_dict = vertex.result result_dict = vertex.result
artifacts = vertex.artifacts artifacts = vertex.artifacts
@ -1452,7 +1456,7 @@ class Graph:
self.set_run_id(run_id) self.set_run_id(run_id)
self.set_run_name() self.set_run_name()
await self.initialize_run() await self.initialize_run()
lock = chat_service._async_cache_locks[self.run_id] lock = chat_service.async_cache_locks[self.run_id]
while to_process: while to_process:
current_batch = list(to_process) # Copy current deque items to a list current_batch = list(to_process) # Copy current deque items to a list
to_process.clear() # Clear the deque for new items to_process.clear() # Clear the deque for new items
@ -1540,7 +1544,7 @@ class Graph:
# This could usually happen with input vertices like ChatInput # This could usually happen with input vertices like ChatInput
self.run_manager.remove_vertex_from_runnables(v.id) self.run_manager.remove_vertex_from_runnables(v.id)
logger.debug(f"Vertex {v.id}, result: {v._built_result}, object: {v._built_object}") logger.debug(f"Vertex {v.id}, result: {v.built_result}, object: {v.built_object}")
for v in vertices: for v in vertices:
next_runnable_vertices = await self.get_next_runnable_vertices(lock, vertex=v, cache=False) next_runnable_vertices = await self.get_next_runnable_vertices(lock, vertex=v, cache=False)

View file

@ -28,18 +28,18 @@ def create_state_model_from_graph(graph: BaseModel) -> type[BaseModel]:
Raises: Raises:
ValueError: If any vertex in the graph does not have a properly initialized ValueError: If any vertex in the graph does not have a properly initialized
component instance (i.e., if vertex._custom_component is None). component instance (i.e., if vertex.custom_component is None).
Notes: Notes:
- Each vertex in the graph must have a '_custom_component' attribute. - Each vertex in the graph must have a 'custom_component' attribute.
- The '_custom_component' must have a 'get_state_model_instance_getter' method. - The 'custom_component' must have a 'get_state_model_instance_getter' method.
- Vertex IDs are converted from camel case to snake case for the resulting model's field names. - Vertex IDs are converted from camel case to snake case for the resulting model's field names.
- The resulting model uses the 'create_state_model' function with validation disabled. - The resulting model uses the 'create_state_model' function with validation disabled.
Example: Example:
>>> class Vertex(BaseModel): >>> class Vertex(BaseModel):
... id: str ... id: str
... _custom_component: Any ... custom_component: Any
>>> class Graph(BaseModel): >>> class Graph(BaseModel):
... vertices: List[Vertex] ... vertices: List[Vertex]
>>> # Assume proper setup of vertices and components >>> # Assume proper setup of vertices and components
@ -50,14 +50,14 @@ def create_state_model_from_graph(graph: BaseModel) -> type[BaseModel]:
>>> print(graph_state.some_component_name) >>> print(graph_state.some_component_name)
""" """
for vertex in graph.vertices: for vertex in graph.vertices:
if hasattr(vertex, "_custom_component") and vertex._custom_component is None: if hasattr(vertex, "custom_component") and vertex.custom_component is None:
msg = f"Vertex {vertex.id} does not have a component instance." msg = f"Vertex {vertex.id} does not have a component instance."
raise ValueError(msg) raise ValueError(msg)
state_model_getters = [ state_model_getters = [
vertex._custom_component.get_state_model_instance_getter() vertex.custom_component.get_state_model_instance_getter()
for vertex in graph.vertices for vertex in graph.vertices
if hasattr(vertex, "_custom_component") and hasattr(vertex._custom_component, "get_state_model_instance_getter") if hasattr(vertex, "custom_component") and hasattr(vertex.custom_component, "get_state_model_instance_getter")
] ]
fields = { fields = {
camel_to_snake(vertex.id): state_model_getter camel_to_snake(vertex.id): state_model_getter

View file

@ -67,17 +67,17 @@ class Vertex:
self.is_input = any(input_component_name in self.id for input_component_name in INPUT_COMPONENTS) self.is_input = any(input_component_name in self.id for input_component_name in INPUT_COMPONENTS)
self.is_output = any(output_component_name in self.id for output_component_name in OUTPUT_COMPONENTS) self.is_output = any(output_component_name in self.id for output_component_name in OUTPUT_COMPONENTS)
self.has_session_id = None self.has_session_id = None
self._custom_component = None self.custom_component = None
self.has_external_input = False self.has_external_input = False
self.has_external_output = False self.has_external_output = False
self.graph = graph self.graph = graph
self._data = data.copy() self.full_data = data.copy()
self.base_type: str | None = base_type self.base_type: str | None = base_type
self.outputs: list[dict] = [] self.outputs: list[dict] = []
self._parse_data() self.parse_data()
self._built_object: Any = UnbuiltObject() self.built_object: Any = UnbuiltObject()
self._built_result: Any = None self.built_result: Any = None
self._built = False self.built = False
self._successors_ids: list[str] | None = None self._successors_ids: list[str] | None = None
self.artifacts: dict[str, Any] = {} self.artifacts: dict[str, Any] = {}
self.artifacts_raw: dict[str, Any] = {} self.artifacts_raw: dict[str, Any] = {}
@ -87,7 +87,7 @@ class Vertex:
self.task_id: str | None = None self.task_id: str | None = None
self.is_task = is_task self.is_task = is_task
self.params = params or {} self.params = params or {}
self.parent_node_id: str | None = self._data.get("parent_node_id") self.parent_node_id: str | None = self.full_data.get("parent_node_id")
self.load_from_db_fields: list[str] = [] self.load_from_db_fields: list[str] = []
self.parent_is_top_level = False self.parent_is_top_level = False
self.layer = None self.layer = None
@ -95,7 +95,7 @@ class Vertex:
self.results: dict[str, Any] = {} self.results: dict[str, Any] = {}
self.outputs_logs: dict[str, OutputValue] = {} self.outputs_logs: dict[str, OutputValue] = {}
self.logs: dict[str, Log] = {} self.logs: dict[str, Log] = {}
self._has_cycle_edges = False self.has_cycle_edges = False
try: try:
self.is_interface_component = self.vertex_type in InterfaceComponentTypes self.is_interface_component = self.vertex_type in InterfaceComponentTypes
except ValueError: except ValueError:
@ -107,17 +107,17 @@ class Vertex:
self.log_transaction_tasks: set[asyncio.Task] = set() self.log_transaction_tasks: set[asyncio.Task] = set()
def set_input_value(self, name: str, value: Any) -> None: def set_input_value(self, name: str, value: Any) -> None:
if self._custom_component is None: if self.custom_component is None:
msg = f"Vertex {self.id} does not have a component instance." msg = f"Vertex {self.id} does not have a component instance."
raise ValueError(msg) raise ValueError(msg)
self._custom_component._set_input_value(name, value) self.custom_component._set_input_value(name, value)
def to_data(self): def to_data(self):
return self._data return self.full_data
def add_component_instance(self, component_instance: Component) -> None: def add_component_instance(self, component_instance: Component) -> None:
component_instance.set_vertex(self) component_instance.set_vertex(self)
self._custom_component = component_instance self.custom_component = component_instance
def add_result(self, name: str, result: Any) -> None: def add_result(self, name: str, result: Any) -> None:
self.results[name] = result self.results[name] = result
@ -154,19 +154,19 @@ class Vertex:
# If the Vertex.type is a power component # If the Vertex.type is a power component
# then we need to return the built object # then we need to return the built object
# instead of the result dict # instead of the result dict
if self.is_interface_component and not isinstance(self._built_object, UnbuiltObject): if self.is_interface_component and not isinstance(self.built_object, UnbuiltObject):
result = self._built_object result = self.built_object
# if it is not a dict or a string and hasattr model_dump then # if it is not a dict or a string and hasattr model_dump then
# return the model_dump # return the model_dump
if not isinstance(result, dict | str) and hasattr(result, "content"): if not isinstance(result, dict | str) and hasattr(result, "content"):
return result.content return result.content
return result return result
if isinstance(self._built_object, str): if isinstance(self.built_object, str):
self._built_result = self._built_object self.built_result = self.built_object
if isinstance(self._built_result, UnbuiltResult): if isinstance(self.built_result, UnbuiltResult):
return {} return {}
return self._built_result if isinstance(self._built_result, dict) else {"result": self._built_result} return self.built_result if isinstance(self.built_result, dict) else {"result": self.built_result}
def set_artifacts(self) -> None: def set_artifacts(self) -> None:
pass pass
@ -202,21 +202,21 @@ class Vertex:
def __getstate__(self): def __getstate__(self):
state = self.__dict__.copy() state = self.__dict__.copy()
state["_lock"] = None # Locks are not serializable state["_lock"] = None # Locks are not serializable
state["_built_object"] = None if isinstance(self._built_object, UnbuiltObject) else self._built_object state["built_object"] = None if isinstance(self.built_object, UnbuiltObject) else self.built_object
state["_built_result"] = None if isinstance(self._built_result, UnbuiltResult) else self._built_result state["built_result"] = None if isinstance(self.built_result, UnbuiltResult) else self.built_result
return state return state
def __setstate__(self, state): def __setstate__(self, state):
self.__dict__.update(state) self.__dict__.update(state)
self._lock = asyncio.Lock() # Reinitialize the lock self._lock = asyncio.Lock() # Reinitialize the lock
self._built_object = state.get("_built_object") or UnbuiltObject() self.built_object = state.get("built_object") or UnbuiltObject()
self._built_result = state.get("_built_result") or UnbuiltResult() self.built_result = state.get("built_result") or UnbuiltResult()
def set_top_level(self, top_level_vertices: list[str]) -> None: def set_top_level(self, top_level_vertices: list[str]) -> None:
self.parent_is_top_level = self.parent_node_id in top_level_vertices self.parent_is_top_level = self.parent_node_id in top_level_vertices
def _parse_data(self) -> None: def parse_data(self) -> None:
self.data = self._data["data"] self.data = self.full_data["data"]
if self.data["node"]["template"]["_type"] == "Component": if self.data["node"]["template"]["_type"] == "Component":
if "outputs" not in self.data["node"]: if "outputs" not in self.data["node"]:
msg = f"Outputs not found for {self.display_name}" msg = f"Outputs not found for {self.display_name}"
@ -300,7 +300,7 @@ class Vertex:
params[param_key] = self.graph.get_vertex(edge.source_id) params[param_key] = self.graph.get_vertex(edge.source_id)
return params return params
def _build_params(self) -> None: def build_params(self) -> None:
# sourcery skip: merge-list-append, remove-redundant-if # sourcery skip: merge-list-append, remove-redundant-if
# Some params are required, some are optional # Some params are required, some are optional
# but most importantly, some params are python base classes # but most importantly, some params are python base classes
@ -436,7 +436,7 @@ class Vertex:
# Add _type to params # Add _type to params
self.params = params self.params = params
self.load_from_db_fields = load_from_db_fields self.load_from_db_fields = load_from_db_fields
self._raw_params = params.copy() self.raw_params = params.copy()
def update_raw_params(self, new_params: Mapping[str, str | list[str]], *, overwrite: bool = False) -> None: def update_raw_params(self, new_params: Mapping[str, str | list[str]], *, overwrite: bool = False) -> None:
"""Update the raw parameters of the vertex with the given new parameters. """Update the raw parameters of the vertex with the given new parameters.
@ -447,28 +447,24 @@ class Vertex:
Defaults to False. Defaults to False.
Raises: Raises:
ValueError: If any key in new_params is not found in self._raw_params. ValueError: If any key in new_params is not found in self.raw_params.
""" """
# First check if the input_value in _raw_params is not a vertex # First check if the input_value in raw_params is not a vertex
if not new_params: if not new_params:
return return
if any(isinstance(self._raw_params.get(key), Vertex) for key in new_params): if any(isinstance(self.raw_params.get(key), Vertex) for key in new_params):
return return
if not overwrite: if not overwrite:
for key in new_params.copy(): # type: ignore[attr-defined] for key in new_params.copy(): # type: ignore[attr-defined]
if key not in self._raw_params: if key not in self.raw_params:
new_params.pop(key) # type: ignore[attr-defined] new_params.pop(key) # type: ignore[attr-defined]
self._raw_params.update(new_params) self.raw_params.update(new_params)
self.params = self._raw_params.copy() self.params = self.raw_params.copy()
self.updated_raw_params = True self.updated_raw_params = True
def has_cycle_edges(self):
"""Checks if the vertex has any cycle edges."""
return self._has_cycle_edges
async def instantiate_component(self, user_id=None) -> None: async def instantiate_component(self, user_id=None) -> None:
if not self._custom_component: if not self.custom_component:
self._custom_component, _ = await initialize.loading.instantiate_class( self.custom_component, _ = await initialize.loading.instantiate_class(
user_id=user_id, user_id=user_id,
vertex=self, vertex=self,
) )
@ -487,13 +483,13 @@ class Vertex:
msg = f"Base type for vertex {self.display_name} not found" msg = f"Base type for vertex {self.display_name} not found"
raise ValueError(msg) raise ValueError(msg)
if not self._custom_component: if not self.custom_component:
custom_component, custom_params = await initialize.loading.instantiate_class( custom_component, custom_params = await initialize.loading.instantiate_class(
user_id=user_id, vertex=self, event_manager=event_manager user_id=user_id, vertex=self, event_manager=event_manager
) )
else: else:
custom_component = self._custom_component custom_component = self.custom_component
self._custom_component.set_event_manager(event_manager) self.custom_component.set_event_manager(event_manager)
custom_params = initialize.loading.get_params(self.params) custom_params = initialize.loading.get_params(self.params)
await self._build_results( await self._build_results(
@ -505,7 +501,7 @@ class Vertex:
self._validate_built_object() self._validate_built_object()
self._built = True self.built = True
def extract_messages_from_artifacts(self, artifacts: dict[str, Any]) -> list[dict]: def extract_messages_from_artifacts(self, artifacts: dict[str, Any]) -> list[dict]:
"""Extracts messages from the artifacts. """Extracts messages from the artifacts.
@ -546,7 +542,7 @@ class Vertex:
return messages return messages
def _finalize_build(self) -> None: def finalize_build(self) -> None:
result_dict = self.get_built_result() result_dict = self.get_built_result()
# We need to set the artifacts to pass information # We need to set the artifacts to pass information
# to the frontend # to the frontend
@ -566,7 +562,7 @@ class Vertex:
async def _build_each_vertex_in_params_dict(self) -> None: async def _build_each_vertex_in_params_dict(self) -> None:
"""Iterates over each vertex in the params dictionary and builds it.""" """Iterates over each vertex in the params dictionary and builds it."""
for key, value in self._raw_params.items(): for key, value in self.raw_params.items():
if self._is_vertex(value): if self._is_vertex(value):
if value == self: if value == self:
del self.params[key] del self.params[key]
@ -637,13 +633,13 @@ class Vertex:
The built result if use_result is True, else the built object. The built result if use_result is True, else the built object.
""" """
flow_id = self.graph.flow_id flow_id = self.graph.flow_id
if not self._built: if not self.built:
if flow_id: if flow_id:
self._log_transaction_async(str(flow_id), source=self, target=requester, status="error") self._log_transaction_async(str(flow_id), source=self, target=requester, status="error")
msg = f"Component {self.display_name} has not been built yet" msg = f"Component {self.display_name} has not been built yet"
raise ValueError(msg) raise ValueError(msg)
result = self._built_result if self.use_result else self._built_object result = self.built_result if self.use_result else self.built_object
if flow_id: if flow_id:
self._log_transaction_async(str(flow_id), source=self, target=requester, status="success") self._log_transaction_async(str(flow_id), source=self, target=requester, status="success")
return result return result
@ -729,50 +725,50 @@ class Vertex:
"""Updates the built object and its artifacts.""" """Updates the built object and its artifacts."""
if isinstance(result, tuple): if isinstance(result, tuple):
if len(result) == 2: # noqa: PLR2004 if len(result) == 2: # noqa: PLR2004
self._built_object, self.artifacts = result self.built_object, self.artifacts = result
elif len(result) == 3: # noqa: PLR2004 elif len(result) == 3: # noqa: PLR2004
self._custom_component, self._built_object, self.artifacts = result self.custom_component, self.built_object, self.artifacts = result
self.logs = self._custom_component._output_logs self.logs = self.custom_component._output_logs
self.artifacts_raw = self.artifacts.get("raw", None) self.artifacts_raw = self.artifacts.get("raw", None)
self.artifacts_type = { self.artifacts_type = {
self.outputs[0]["name"]: self.artifacts.get("type", None) or ArtifactType.UNKNOWN.value self.outputs[0]["name"]: self.artifacts.get("type", None) or ArtifactType.UNKNOWN.value
} }
self.artifacts = {self.outputs[0]["name"]: self.artifacts} self.artifacts = {self.outputs[0]["name"]: self.artifacts}
else: else:
self._built_object = result self.built_object = result
def _validate_built_object(self) -> None: def _validate_built_object(self) -> None:
"""Checks if the built object is None and raises a ValueError if so.""" """Checks if the built object is None and raises a ValueError if so."""
if isinstance(self._built_object, UnbuiltObject): if isinstance(self.built_object, UnbuiltObject):
msg = f"{self.display_name}: {self._built_object_repr()}" msg = f"{self.display_name}: {self.built_object_repr()}"
raise TypeError(msg) raise TypeError(msg)
if self._built_object is None: if self.built_object is None:
message = f"{self.display_name} returned None." message = f"{self.display_name} returned None."
if self.base_type == "custom_components": if self.base_type == "custom_components":
message += " Make sure your build method returns a component." message += " Make sure your build method returns a component."
logger.warning(message) logger.warning(message)
elif isinstance(self._built_object, Iterator | AsyncIterator): elif isinstance(self.built_object, Iterator | AsyncIterator):
if self.display_name == "Text Output": if self.display_name == "Text Output":
msg = f"You are trying to stream to a {self.display_name}. Try using a Chat Output instead." msg = f"You are trying to stream to a {self.display_name}. Try using a Chat Output instead."
raise ValueError(msg) raise ValueError(msg)
def _reset(self) -> None: def _reset(self) -> None:
self._built = False self.built = False
self._built_object = UnbuiltObject() self.built_object = UnbuiltObject()
self._built_result = UnbuiltResult() self.built_result = UnbuiltResult()
self.artifacts = {} self.artifacts = {}
self.steps_ran = [] self.steps_ran = []
self._build_params() self.build_params()
def _is_chat_input(self) -> bool: def _is_chat_input(self) -> bool:
return False return False
def build_inactive(self) -> None: def build_inactive(self) -> None:
# Just set the results to None # Just set the results to None
self._built = True self.built = True
self._built_object = None self.built_object = None
self._built_result = None self.built_result = None
async def build( async def build(
self, self,
@ -789,9 +785,9 @@ class Vertex:
self.build_inactive() self.build_inactive()
return None return None
if self.frozen and self._built: if self.frozen and self.built:
return await self.get_requester_result(requester) return await self.get_requester_result(requester)
if self._built and requester is not None: if self.built and requester is not None:
# This means that the vertex has already been built # This means that the vertex has already been built
# and we are just getting the result for the requester # and we are just getting the result for the requester
return await self.get_requester_result(requester) return await self.get_requester_result(requester)
@ -824,7 +820,7 @@ class Vertex:
step(user_id=user_id, event_manager=event_manager, **kwargs) step(user_id=user_id, event_manager=event_manager, **kwargs)
self.steps_ran.append(step) self.steps_ran.append(step)
self._finalize_build() self.finalize_build()
return await self.get_requester_result(requester) return await self.get_requester_result(requester)
@ -832,7 +828,7 @@ class Vertex:
# If the requester is None, this means that # If the requester is None, this means that
# the Vertex is the root of the graph # the Vertex is the root of the graph
if requester is None: if requester is None:
return self._built_object return self.built_object
# Get the requester edge # Get the requester edge
requester_edge = next((edge for edge in self.edges if edge.target_id == requester.id), None) requester_edge = next((edge for edge in self.edges if edge.target_id == requester.id), None)
@ -857,7 +853,7 @@ class Vertex:
# We should create a more robust comparison # We should create a more robust comparison
# for the Vertex class # for the Vertex class
ids_are_equal = self.id == __o.id ids_are_equal = self.id == __o.id
# self._data is a dict and we need to compare them # self.data is a dict and we need to compare them
# to check if they are equal # to check if they are equal
data_are_equal = self.data == __o.data data_are_equal = self.data == __o.data
except AttributeError: except AttributeError:
@ -868,13 +864,13 @@ class Vertex:
def __hash__(self) -> int: def __hash__(self) -> int:
return id(self) return id(self)
def _built_object_repr(self) -> str: def built_object_repr(self) -> str:
# Add a message with an emoji, stars for sucess, # Add a message with an emoji, stars for success,
return "Built successfully ✨" if self._built_object is not None else "Failed to build 😵‍💫" return "Built successfully ✨" if self.built_object is not None else "Failed to build 😵‍💫"
def apply_on_outputs(self, func: Callable[[Any], Any]) -> None: def apply_on_outputs(self, func: Callable[[Any], Any]) -> None:
"""Applies a function to the outputs of the vertex.""" """Applies a function to the outputs of the vertex."""
if not self._custom_component or not self._custom_component.outputs: if not self.custom_component or not self.custom_component.outputs:
return return
# Apply the function to each output # Apply the function to each output
[func(output) for output in self._custom_component.outputs] [func(output) for output in self.custom_component.outputs]

View file

@ -31,9 +31,9 @@ class CustomComponentVertex(Vertex):
def __init__(self, data: NodeData, graph): def __init__(self, data: NodeData, graph):
super().__init__(data, graph=graph, base_type="custom_components") super().__init__(data, graph=graph, base_type="custom_components")
def _built_object_repr(self): def built_object_repr(self):
if self.artifacts and "repr" in self.artifacts: if self.artifacts and "repr" in self.artifacts:
return self.artifacts["repr"] or super()._built_object_repr() return self.artifacts["repr"] or super().built_object_repr()
return None return None
@ -42,36 +42,36 @@ class ComponentVertex(Vertex):
super().__init__(data, graph=graph, base_type="component") super().__init__(data, graph=graph, base_type="component")
def get_input(self, name: str) -> InputTypes: def get_input(self, name: str) -> InputTypes:
if self._custom_component is None: if self.custom_component is None:
msg = f"Vertex {self.id} does not have a component instance." msg = f"Vertex {self.id} does not have a component instance."
raise ValueError(msg) raise ValueError(msg)
return self._custom_component.get_input(name) return self.custom_component.get_input(name)
def get_output(self, name: str) -> Output: def get_output(self, name: str) -> Output:
if self._custom_component is None: if self.custom_component is None:
raise NoComponentInstanceError(self.id) raise NoComponentInstanceError(self.id)
return self._custom_component.get_output(name) return self.custom_component.get_output(name)
def _built_object_repr(self): def built_object_repr(self):
if self.artifacts and "repr" in self.artifacts: if self.artifacts and "repr" in self.artifacts:
return self.artifacts["repr"] or super()._built_object_repr() return self.artifacts["repr"] or super().built_object_repr()
return None return None
def _update_built_object_and_artifacts(self, result) -> None: def _update_built_object_and_artifacts(self, result) -> None:
"""Updates the built object and its artifacts.""" """Updates the built object and its artifacts."""
if isinstance(result, tuple): if isinstance(result, tuple):
if len(result) == 2: # noqa: PLR2004 if len(result) == 2: # noqa: PLR2004
self._built_object, self.artifacts = result self.built_object, self.artifacts = result
elif len(result) == 3: # noqa: PLR2004 elif len(result) == 3: # noqa: PLR2004
self._custom_component, self._built_object, self.artifacts = result self.custom_component, self.built_object, self.artifacts = result
self.logs = self._custom_component._output_logs self.logs = self.custom_component._output_logs
for key in self.artifacts: for key in self.artifacts:
self.artifacts_raw[key] = self.artifacts[key].get("raw", None) self.artifacts_raw[key] = self.artifacts[key].get("raw", None)
self.artifacts_type[key] = self.artifacts[key].get("type", None) or ArtifactType.UNKNOWN.value self.artifacts_type[key] = self.artifacts[key].get("type", None) or ArtifactType.UNKNOWN.value
else: else:
self._built_object = result self.built_object = result
for key, value in self._built_object.items(): for key, value in self.built_object.items():
self.add_result(key, value) self.add_result(key, value)
def get_edge_with_target(self, target_id: str) -> Generator[CycleEdge, None, None]: def get_edge_with_target(self, target_id: str) -> Generator[CycleEdge, None, None]:
@ -96,7 +96,7 @@ class ComponentVertex(Vertex):
The built result if use_result is True, else the built object. The built result if use_result is True, else the built object.
""" """
flow_id = self.graph.flow_id flow_id = self.graph.flow_id
if not self._built: if not self.built:
default_value = UNDEFINED default_value = UNDEFINED
for edge in self.get_edge_with_target(requester.id): for edge in self.get_edge_with_target(requester.id):
# We need to check if the edge is a normal edge # We need to check if the edge is a normal edge
@ -181,7 +181,7 @@ class ComponentVertex(Vertex):
) )
return messages return messages
def _finalize_build(self) -> None: def finalize_build(self) -> None:
result_dict = self.get_built_result() result_dict = self.get_built_result()
# We need to set the artifacts to pass information # We need to set the artifacts to pass information
# to the frontend # to the frontend
@ -201,14 +201,14 @@ class ComponentVertex(Vertex):
class InterfaceVertex(ComponentVertex): class InterfaceVertex(ComponentVertex):
def __init__(self, data: NodeData, graph): def __init__(self, data: NodeData, graph):
super().__init__(data, graph=graph) super().__init__(data, graph=graph)
self._added_message = None self.added_message = None
self.steps = [self._build, self._run] self.steps = [self._build, self._run]
self.is_interface_component = True self.is_interface_component = True
def build_stream_url(self) -> str: def build_stream_url(self) -> str:
return f"/api/v1/build/{self.graph.flow_id}/{self.id}/stream" return f"/api/v1/build/{self.graph.flow_id}/{self.id}/stream"
def _built_object_repr(self): def built_object_repr(self):
if self.task_id and self.is_task: if self.task_id and self.is_task:
if task := self.get_task(): if task := self.get_task():
return str(task.info) return str(task.info)
@ -227,22 +227,22 @@ class InterfaceVertex(ComponentVertex):
_artifact = {k.title().replace("_", " "): v for k, v in artifact.items() if v is not None} _artifact = {k.title().replace("_", " "): v for k, v in artifact.items() if v is not None}
artifacts.append(_artifact) artifacts.append(_artifact)
return yaml.dump(artifacts, default_flow_style=False, allow_unicode=True) return yaml.dump(artifacts, default_flow_style=False, allow_unicode=True)
return super()._built_object_repr() return super().built_object_repr()
def _process_chat_component(self): def _process_chat_component(self):
"""Process the chat component and return the message. """Process the chat component and return the message.
This method processes the chat component by extracting the necessary parameters This method processes the chat component by extracting the necessary parameters
such as sender, sender_name, and message from the `params` dictionary. It then such as sender, sender_name, and message from the `params` dictionary. It then
performs additional operations based on the type of the `_built_object` attribute. performs additional operations based on the type of the `built_object` attribute.
If `_built_object` is an instance of `AIMessage`, it creates a `ChatOutputResponse` If `built_object` is an instance of `AIMessage`, it creates a `ChatOutputResponse`
object using the `from_message` method. If `_built_object` is not an instance of object using the `from_message` method. If `built_object` is not an instance of
`UnbuiltObject`, it checks the type of `_built_object` and performs specific `UnbuiltObject`, it checks the type of `built_object` and performs specific
operations accordingly. If `_built_object` is a dictionary, it converts it into a operations accordingly. If `built_object` is a dictionary, it converts it into a
code block. If `_built_object` is an instance of `Data`, it assigns the `text` code block. If `built_object` is an instance of `Data`, it assigns the `text`
attribute to the `message` variable. If `message` is an instance of `AsyncIterator` attribute to the `message` variable. If `message` is an instance of `AsyncIterator`
or `Iterator`, it builds a stream URL and sets `message` to an empty string. If or `Iterator`, it builds a stream URL and sets `message` to an empty string. If
`_built_object` is not a string, it converts it to a string. If `message` is a `built_object` is not a string, it converts it to a string. If `message` is a
generator or iterator, it assigns it to the `message` variable. Finally, it creates generator or iterator, it assigns it to the `message` variable. Finally, it creates
a `ChatOutputResponse` object using the extracted parameters and assigns it to the a `ChatOutputResponse` object using the extracted parameters and assigns it to the
`artifacts` attribute. If `artifacts` is not None, it calls the `model_dump` method `artifacts` attribute. If `artifacts` is not None, it calls the `model_dump` method
@ -288,7 +288,7 @@ class InterfaceVertex(ComponentVertex):
message = "" message = ""
self.results["text"] = message self.results["text"] = message
self.results["message"].text = message self.results["message"].text = message
self._built_object = self.results self.built_object = self.results
elif not isinstance(text_output, str): elif not isinstance(text_output, str):
message = str(text_output) message = str(text_output)
# if the message is a generator or iterator # if the message is a generator or iterator
@ -335,12 +335,12 @@ class InterfaceVertex(ComponentVertex):
ValueError: If an element in the list is not an instance of `Data` and ValueError: If an element in the list is not an instance of `Data` and
`ignore_errors` is set to `False`. `ignore_errors` is set to `False`.
""" """
if isinstance(self._built_object, Data): if isinstance(self.built_object, Data):
artifacts = [self._built_object.data] artifacts = [self.built_object.data]
elif isinstance(self._built_object, list): elif isinstance(self.built_object, list):
artifacts = [] artifacts = []
ignore_errors = self.params.get("ignore_errors", False) ignore_errors = self.params.get("ignore_errors", False)
for value in self._built_object: for value in self.built_object:
if isinstance(value, Data): if isinstance(value, Data):
artifacts.append(value.data) artifacts.append(value.data)
elif ignore_errors: elif ignore_errors:
@ -349,19 +349,19 @@ class InterfaceVertex(ComponentVertex):
msg = f"Data expected, but got {value} of type {type(value)}" msg = f"Data expected, but got {value} of type {type(value)}"
raise ValueError(msg) raise ValueError(msg)
self.artifacts = DataOutputResponse(data=artifacts) self.artifacts = DataOutputResponse(data=artifacts)
return self._built_object return self.built_object
async def _run(self, *args, **kwargs) -> None: # noqa: ARG002 async def _run(self, *args, **kwargs) -> None: # noqa: ARG002
if self.vertex_type in CHAT_COMPONENTS: if self.vertex_type in CHAT_COMPONENTS:
message = self._process_chat_component() message = self._process_chat_component()
elif self.vertex_type in RECORDS_COMPONENTS: elif self.vertex_type in RECORDS_COMPONENTS:
message = self._process_data_component() message = self._process_data_component()
if isinstance(self._built_object, AsyncIterator | Iterator): if isinstance(self.built_object, AsyncIterator | Iterator):
if self.params.get("return_data", False): if self.params.get("return_data", False):
self._built_object = Data(text=message, data=self.artifacts) self.built_object = Data(text=message, data=self.artifacts)
else: else:
self._built_object = message self.built_object = message
self._built_result = self._built_object self.built_result = self.built_object
async def stream(self): async def stream(self):
iterator = self.params.get(INPUT_FIELD_NAME, None) iterator = self.params.get(INPUT_FIELD_NAME, None)
@ -410,17 +410,17 @@ class InterfaceVertex(ComponentVertex):
session_id=self.params.get("session_id", ""), session_id=self.params.get("session_id", ""),
) )
self.params[INPUT_FIELD_NAME] = complete_message self.params[INPUT_FIELD_NAME] = complete_message
if isinstance(self._built_object, dict): if isinstance(self.built_object, dict):
for key, value in self._built_object.items(): for key, value in self.built_object.items():
if hasattr(value, "text") and (isinstance(value.text, AsyncIterator | Iterator) or value.text == ""): if hasattr(value, "text") and (isinstance(value.text, AsyncIterator | Iterator) or value.text == ""):
self._built_object[key] = message self.built_object[key] = message
else: else:
self._built_object = message self.built_object = message
self.artifacts_type = ArtifactType.MESSAGE self.artifacts_type = ArtifactType.MESSAGE
# Update artifacts with the message # Update artifacts with the message
# and remove the stream_url # and remove the stream_url
self._finalize_build() self.finalize_build()
logger.debug(f"Streamed message: {complete_message}") logger.debug(f"Streamed message: {complete_message}")
# Set the result in the vertex of origin # Set the result in the vertex of origin
edges = self.get_edge_with_target(self.id) edges = self.get_edge_with_target(self.id)
@ -430,22 +430,22 @@ class InterfaceVertex(ComponentVertex):
if isinstance(value, AsyncIterator | Iterator): if isinstance(value, AsyncIterator | Iterator):
origin_vertex.results[key] = complete_message origin_vertex.results[key] = complete_message
if ( if (
self._custom_component self.custom_component
and hasattr(self._custom_component, "should_store_message") and hasattr(self.custom_component, "should_store_message")
and hasattr(self._custom_component, "store_message") and hasattr(self.custom_component, "store_message")
): ):
self._custom_component.store_message(message) self.custom_component.store_message(message)
log_vertex_build( log_vertex_build(
flow_id=self.graph.flow_id, flow_id=self.graph.flow_id,
vertex_id=self.id, vertex_id=self.id,
valid=True, valid=True,
params=self._built_object_repr(), params=self.built_object_repr(),
data=self.result, data=self.result,
artifacts=self.artifacts, artifacts=self.artifacts,
) )
self._validate_built_object() self._validate_built_object()
self._built = True self.built = True
async def consume_async_generator(self) -> None: async def consume_async_generator(self) -> None:
async for _ in self.stream(): async for _ in self.stream():
@ -468,9 +468,9 @@ class StateVertex(ComponentVertex):
return super().successors_ids return super().successors_ids
return self._successors_ids return self._successors_ids
def _built_object_repr(self): def built_object_repr(self):
if self.artifacts and "repr" in self.artifacts: if self.artifacts and "repr" in self.artifacts:
return self.artifacts["repr"] or super()._built_object_repr() return self.artifacts["repr"] or super().built_object_repr()
return None return None

View file

@ -14,7 +14,7 @@ class ChatService(Service):
name = "chat_service" name = "chat_service"
def __init__(self) -> None: def __init__(self) -> None:
self._async_cache_locks: dict[str, asyncio.Lock] = defaultdict(asyncio.Lock) self.async_cache_locks: dict[str, asyncio.Lock] = defaultdict(asyncio.Lock)
self._sync_cache_locks: dict[str, RLock] = defaultdict(RLock) self._sync_cache_locks: dict[str, RLock] = defaultdict(RLock)
self.cache_service = get_cache_service() self.cache_service = get_cache_service()
@ -28,7 +28,7 @@ class ChatService(Service):
threading.Lock or asyncio.Lock: The lock associated with the given key. threading.Lock or asyncio.Lock: The lock associated with the given key.
""" """
if isinstance(self.cache_service, AsyncBaseCacheService): if isinstance(self.cache_service, AsyncBaseCacheService):
return self._async_cache_locks[key] return self.async_cache_locks[key]
return self._sync_cache_locks[key] return self._sync_cache_locks[key]
async def _perform_cache_operation( async def _perform_cache_operation(

View file

@ -63,9 +63,9 @@ async def build_vertex(
return return
start_time = time.perf_counter() start_time = time.perf_counter()
try: try:
if isinstance(vertex, Vertex) or not vertex._built: if isinstance(vertex, Vertex) or not vertex.built:
await vertex.build(user_id=None, session_id=sid) await vertex.build(user_id=None, session_id=sid)
params = vertex._built_object_repr() params = vertex.built_object_repr()
valid = True valid = True
result_dict = vertex.get_built_result() result_dict = vertex.get_built_result()
# We need to set the artifacts to pass information # We need to set the artifacts to pass information

View file

@ -37,7 +37,7 @@ async def user_data_context(store_service: StoreService, api_key: str | None = N
# Fetch and set user data to the context variable # Fetch and set user data to the context variable
if api_key: if api_key:
try: try:
user_data, _ = await store_service._get( user_data, _ = await store_service.get(
f"{store_service.base_url}/users/me", api_key, params={"fields": "id"} f"{store_service.base_url}/users/me", api_key, params={"fields": "id"}
) )
user_data_var.set(user_data[0]) user_data_var.set(user_data[0])
@ -112,7 +112,7 @@ class StoreService(Service):
# If it is, return True # If it is, return True
# If it is not, return False # If it is not, return False
try: try:
user_data, _ = await self._get(f"{self.base_url}/users/me", api_key, params={"fields": "id"}) user_data, _ = await self.get(f"{self.base_url}/users/me", api_key, params={"fields": "id"})
return "id" in user_data[0] return "id" in user_data[0]
except HTTPStatusError as exc: except HTTPStatusError as exc:
@ -124,7 +124,7 @@ class StoreService(Service):
msg = f"Unexpected error: {exc}" msg = f"Unexpected error: {exc}"
raise ValueError(msg) from exc raise ValueError(msg) from exc
async def _get( async def get(
self, url: str, api_key: str | None = None, params: dict[str, Any] | None = None self, url: str, api_key: str | None = None, params: dict[str, Any] | None = None
) -> tuple[list[dict[str, Any]], dict[str, Any]]: ) -> tuple[list[dict[str, Any]], dict[str, Any]]:
"""Utility method to perform GET requests.""" """Utility method to perform GET requests."""
@ -183,7 +183,7 @@ class StoreService(Service):
api_key = api_key if use_api_key else None api_key = api_key if use_api_key else None
results, _ = await self._get(self.components_url, api_key, params) results, _ = await self.get(self.components_url, api_key, params)
return int(results[0].get("count", 0)) return int(results[0].get("count", 0))
@staticmethod @staticmethod
@ -289,7 +289,7 @@ class StoreService(Service):
# so we don't need to risk passing an invalid api_key # so we don't need to risk passing an invalid api_key
# and getting 401 # and getting 401
api_key = api_key if use_api_key else None api_key = api_key if use_api_key else None
results, metadata = await self._get(self.components_url, api_key, params) results, metadata = await self.get(self.components_url, api_key, params)
if isinstance(results, dict): if isinstance(results, dict):
results = [results] results = [results]
@ -316,7 +316,7 @@ class StoreService(Service):
} }
), ),
} }
results, _ = await self._get(self.components_url, api_key, params) results, _ = await self.get(self.components_url, api_key, params)
return [result["id"] for result in results] return [result["id"] for result in results]
# Which of the components is parent of the user's components # Which of the components is parent of the user's components
@ -336,7 +336,7 @@ class StoreService(Service):
} }
), ),
} }
results, _ = await self._get(self.components_url, api_key, params) results, _ = await self.get(self.components_url, api_key, params)
return [result["id"] for result in results] return [result["id"] for result in results]
async def download(self, api_key: str, component_id: UUID) -> DownloadComponentResponse: async def download(self, api_key: str, component_id: UUID) -> DownloadComponentResponse:
@ -345,7 +345,7 @@ class StoreService(Service):
if not self.download_webhook_url: if not self.download_webhook_url:
msg = "DOWNLOAD_WEBHOOK_URL is not set" msg = "DOWNLOAD_WEBHOOK_URL is not set"
raise ValueError(msg) raise ValueError(msg)
component, _ = await self._get(url, api_key, params) component, _ = await self.get(url, api_key, params)
await self.call_webhook(api_key, self.download_webhook_url, component_id) await self.call_webhook(api_key, self.download_webhook_url, component_id)
if len(component) > 1: if len(component) > 1:
msg = "Something went wrong while downloading the component" msg = "Something went wrong while downloading the component"
@ -437,7 +437,7 @@ class StoreService(Service):
async def get_tags(self) -> list[dict[str, Any]]: async def get_tags(self) -> list[dict[str, Any]]:
url = f"{self.base_url}/items/tags" url = f"{self.base_url}/items/tags"
params = {"fields": "id,name"} params = {"fields": "id,name"}
tags, _ = await self._get(url, api_key=None, params=params) tags, _ = await self.get(url, api_key=None, params=params)
return tags return tags
async def get_user_likes(self, api_key: str) -> list[dict[str, Any]]: async def get_user_likes(self, api_key: str) -> list[dict[str, Any]]:
@ -445,7 +445,7 @@ class StoreService(Service):
params = { params = {
"fields": "id,likes", "fields": "id,likes",
} }
likes, _ = await self._get(url, api_key, params) likes, _ = await self.get(url, api_key, params)
return likes return likes
async def get_component_likes_count(self, component_id: str, api_key: str | None = None) -> int: async def get_component_likes_count(self, component_id: str, api_key: str | None = None) -> int:
@ -454,7 +454,7 @@ class StoreService(Service):
params = { params = {
"fields": "id,count(liked_by)", "fields": "id,count(liked_by)",
} }
result, _ = await self._get(url, api_key=api_key, params=params) result, _ = await self.get(url, api_key=api_key, params=params)
if len(result) == 0: if len(result) == 0:
msg = "Component not found" msg = "Component not found"
raise ValueError(msg) raise ValueError(msg)

View file

@ -59,7 +59,7 @@ ignore = [
# Rules that are TODOs # Rules that are TODOs
"ANN", # Missing type annotations "ANN", # Missing type annotations
"D1", # Missing docstrings "D1", # Missing docstrings
"SLF", "SLF001", # Using private attributes outside of class
] ]
[tool.ruff.lint.per-file-ignores] [tool.ruff.lint.per-file-ignores]

View file

@ -169,4 +169,4 @@ async def run_single_component(
_, _ = await run_graph_internal( _, _ = await run_graph_internal(
graph, flow_id, session_id=session_id, inputs=graph_run_inputs, outputs=[component_id] graph, flow_id, session_id=session_id, inputs=graph_run_inputs, outputs=[component_id]
) )
return graph.get_vertex(component_id)._built_object return graph.get_vertex(component_id).built_object

View file

@ -126,13 +126,13 @@ def test_that_outputs_cache_is_set_to_false_in_cycle():
graph = Graph(chat_input, chat_output) graph = Graph(chat_input, chat_output)
cycle_vertices = find_cycle_vertices(graph._get_edges_as_list_of_tuples()) cycle_vertices = find_cycle_vertices(graph._get_edges_as_list_of_tuples())
cycle_outputs_lists = [graph.vertex_map[vertex_id]._custom_component.outputs for vertex_id in cycle_vertices] cycle_outputs_lists = [graph.vertex_map[vertex_id].custom_component.outputs for vertex_id in cycle_vertices]
cycle_outputs = [output for outputs in cycle_outputs_lists for output in outputs] cycle_outputs = [output for outputs in cycle_outputs_lists for output in outputs]
for output in cycle_outputs: for output in cycle_outputs:
assert output.cache is False assert output.cache is False
non_cycle_outputs_lists = [ non_cycle_outputs_lists = [
vertex._custom_component.outputs for vertex in graph.vertices if vertex.id not in cycle_vertices vertex.custom_component.outputs for vertex in graph.vertices if vertex.id not in cycle_vertices
] ]
non_cycle_outputs = [output for outputs in non_cycle_outputs_lists for output in outputs] non_cycle_outputs = [output for outputs in non_cycle_outputs_lists for output in outputs]
for output in non_cycle_outputs: for output in non_cycle_outputs: