Fix styleUtils import and remove unnecessary lines

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-02-22 17:07:36 -03:00
commit 9cfb03fbc9
47 changed files with 219 additions and 574 deletions

View file

@ -3,8 +3,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional
from uuid import UUID from uuid import UUID
from langchain.schema import AgentAction, AgentFinish from langchain.schema import AgentAction, AgentFinish
from langchain_core.callbacks.base import (AsyncCallbackHandler, from langchain_core.callbacks.base import AsyncCallbackHandler, BaseCallbackHandler
BaseCallbackHandler)
from langflow.api.v1.schemas import ChatResponse, PromptResponse from langflow.api.v1.schemas import ChatResponse, PromptResponse
from langflow.services.deps import get_chat_service from langflow.services.deps import get_chat_service
from langflow.utils.util import remove_ansi_escape_codes from langflow.utils.util import remove_ansi_escape_codes

View file

@ -43,13 +43,9 @@ async def chat(
user = await get_current_user_for_websocket(websocket, db) user = await get_current_user_for_websocket(websocket, db)
await websocket.accept() await websocket.accept()
if not user: if not user:
await websocket.close( await websocket.close(code=status.WS_1008_POLICY_VIOLATION, reason="Unauthorized")
code=status.WS_1008_POLICY_VIOLATION, reason="Unauthorized"
)
elif not user.is_active: elif not user.is_active:
await websocket.close( await websocket.close(code=status.WS_1008_POLICY_VIOLATION, reason="Unauthorized")
code=status.WS_1008_POLICY_VIOLATION, reason="Unauthorized"
)
if client_id in chat_service.cache_service: if client_id in chat_service.cache_service:
await chat_service.handle_websocket(client_id, websocket) await chat_service.handle_websocket(client_id, websocket)
@ -65,9 +61,7 @@ async def chat(
logger.error(f"Error in chat websocket: {exc}") logger.error(f"Error in chat websocket: {exc}")
messsage = exc.detail if isinstance(exc, HTTPException) else str(exc) messsage = exc.detail if isinstance(exc, HTTPException) else str(exc)
if "Could not validate credentials" in str(exc): if "Could not validate credentials" in str(exc):
await websocket.close( await websocket.close(code=status.WS_1008_POLICY_VIOLATION, reason="Unauthorized")
code=status.WS_1008_POLICY_VIOLATION, reason="Unauthorized"
)
else: else:
await websocket.close(code=status.WS_1011_INTERNAL_ERROR, reason=messsage) await websocket.close(code=status.WS_1011_INTERNAL_ERROR, reason=messsage)
@ -137,12 +131,8 @@ async def build_vertex(
cache = chat_service.get_cache(flow_id) cache = chat_service.get_cache(flow_id)
if not cache: if not cache:
# If there's no cache # If there's no cache
logger.warning( logger.warning(f"No cache found for {flow_id}. Building graph starting at {vertex_id}")
f"No cache found for {flow_id}. Building graph starting at {vertex_id}" graph = build_and_cache_graph(flow_id=flow_id, session=next(get_session()), chat_service=chat_service)
)
graph = build_and_cache_graph(
flow_id=flow_id, session=next(get_session()), chat_service=chat_service
)
else: else:
graph = cache.get("result") graph = cache.get("result")
result_dict = {} result_dict = {}

View file

@ -66,8 +66,6 @@ async def get_transactions(
monitor_service: MonitorService = Depends(get_monitor_service), monitor_service: MonitorService = Depends(get_monitor_service),
): ):
try: try:
return monitor_service.get_transactions( return monitor_service.get_transactions(source=source, target=target, status=status, order_by=order_by)
source=source, target=target, status=status, order_by=order_by
)
except Exception as e: except Exception as e:
raise HTTPException(status_code=500, detail=str(e)) raise HTTPException(status_code=500, detail=str(e))

View file

@ -161,9 +161,7 @@ class StreamData(BaseModel):
data: dict data: dict
def __str__(self) -> str: def __str__(self) -> str:
return ( return f"event: {self.event}\ndata: {orjson_dumps(self.data, indent_2=False)}\n\n"
f"event: {self.event}\ndata: {orjson_dumps(self.data, indent_2=False)}\n\n"
)
class CustomComponentCode(BaseModel): class CustomComponentCode(BaseModel):

View file

@ -40,9 +40,7 @@ def post_validate_prompt(prompt_request: ValidatePromptRequest):
add_new_variables_to_template(input_variables, prompt_request) add_new_variables_to_template(input_variables, prompt_request)
remove_old_variables_from_template( remove_old_variables_from_template(old_custom_fields, input_variables, prompt_request)
old_custom_fields, input_variables, prompt_request
)
update_input_variables_field(input_variables, prompt_request) update_input_variables_field(input_variables, prompt_request)
@ -57,19 +55,12 @@ def post_validate_prompt(prompt_request: ValidatePromptRequest):
def get_old_custom_fields(prompt_request): def get_old_custom_fields(prompt_request):
try: try:
if ( if len(prompt_request.frontend_node.custom_fields) == 1 and prompt_request.name == "":
len(prompt_request.frontend_node.custom_fields) == 1
and prompt_request.name == ""
):
# If there is only one custom field and the name is empty string # If there is only one custom field and the name is empty string
# then we are dealing with the first prompt request after the node was created # then we are dealing with the first prompt request after the node was created
prompt_request.name = list( prompt_request.name = list(prompt_request.frontend_node.custom_fields.keys())[0]
prompt_request.frontend_node.custom_fields.keys()
)[0]
old_custom_fields = prompt_request.frontend_node.custom_fields[ old_custom_fields = prompt_request.frontend_node.custom_fields[prompt_request.name]
prompt_request.name
]
if old_custom_fields is None: if old_custom_fields is None:
old_custom_fields = [] old_custom_fields = []
@ -95,40 +86,26 @@ def add_new_variables_to_template(input_variables, prompt_request):
) )
if variable in prompt_request.frontend_node.template: if variable in prompt_request.frontend_node.template:
# Set the new field with the old value # Set the new field with the old value
template_field.value = prompt_request.frontend_node.template[variable][ template_field.value = prompt_request.frontend_node.template[variable]["value"]
"value"
]
prompt_request.frontend_node.template[variable] = template_field.to_dict() prompt_request.frontend_node.template[variable] = template_field.to_dict()
# Check if variable is not already in the list before appending # Check if variable is not already in the list before appending
if ( if variable not in prompt_request.frontend_node.custom_fields[prompt_request.name]:
variable prompt_request.frontend_node.custom_fields[prompt_request.name].append(variable)
not in prompt_request.frontend_node.custom_fields[prompt_request.name]
):
prompt_request.frontend_node.custom_fields[prompt_request.name].append(
variable
)
except Exception as exc: except Exception as exc:
logger.exception(exc) logger.exception(exc)
raise HTTPException(status_code=500, detail=str(exc)) from exc raise HTTPException(status_code=500, detail=str(exc)) from exc
def remove_old_variables_from_template( def remove_old_variables_from_template(old_custom_fields, input_variables, prompt_request):
old_custom_fields, input_variables, prompt_request
):
for variable in old_custom_fields: for variable in old_custom_fields:
if variable not in input_variables: if variable not in input_variables:
try: try:
# Remove the variable from custom_fields associated with the given name # Remove the variable from custom_fields associated with the given name
if ( if variable in prompt_request.frontend_node.custom_fields[prompt_request.name]:
variable prompt_request.frontend_node.custom_fields[prompt_request.name].remove(variable)
in prompt_request.frontend_node.custom_fields[prompt_request.name]
):
prompt_request.frontend_node.custom_fields[
prompt_request.name
].remove(variable)
# Remove the variable from the template # Remove the variable from the template
prompt_request.frontend_node.template.pop(variable, None) prompt_request.frontend_node.template.pop(variable, None)
@ -140,6 +117,4 @@ def remove_old_variables_from_template(
def update_input_variables_field(input_variables, prompt_request): def update_input_variables_field(input_variables, prompt_request):
if "input_variables" in prompt_request.frontend_node.template: if "input_variables" in prompt_request.frontend_node.template:
prompt_request.frontend_node.template["input_variables"][ prompt_request.frontend_node.template["input_variables"]["value"] = input_variables
"value"
] = input_variables

View file

@ -52,9 +52,7 @@ class StoreMessages(CustomComponent):
if not records: if not records:
records = [] records = []
if not session_id or not sender or not sender_name: if not session_id or not sender or not sender_name:
raise ValueError( raise ValueError("If passing texts, session_id, sender, and sender_name must be provided.")
"If passing texts, session_id, sender, and sender_name must be provided."
)
for text in texts: for text in texts:
record = Record( record = Record(
text=text, text=text,

View file

@ -126,8 +126,7 @@ class ChatLiteLLMComponent(CustomComponent):
litellm.set_verbose = verbose litellm.set_verbose = verbose
except ImportError: except ImportError:
raise ChatLiteLLMException( raise ChatLiteLLMException(
"Could not import litellm python package. " "Could not import litellm python package. " "Please install it with `pip install litellm`"
"Please install it with `pip install litellm`"
) )
provider_map = { provider_map = {
"OpenAI": "openai_api_key", "OpenAI": "openai_api_key",

View file

@ -54,7 +54,7 @@ class AnthropicLLM(CustomComponent):
def build( def build(
self, self,
model: str, model: str,
inputs:str, inputs: str,
anthropic_api_key: Optional[str] = None, anthropic_api_key: Optional[str] = None,
max_tokens: Optional[int] = None, max_tokens: Optional[int] = None,
temperature: Optional[float] = None, temperature: Optional[float] = None,
@ -78,4 +78,3 @@ class AnthropicLLM(CustomComponent):
result = message.content if hasattr(message, "content") else message result = message.content if hasattr(message, "content") else message
self.status = result self.status = result
return result return result

View file

@ -31,7 +31,7 @@ class CTransformersComponent(CustomComponent):
"inputs": {"display_name": "Input"}, "inputs": {"display_name": "Input"},
} }
def build(self, model: str, model_file: str,inputs:str, model_type: str, config: Optional[Dict] = None) -> Text: def build(self, model: str, model_file: str, inputs: str, model_type: str, config: Optional[Dict] = None) -> Text:
output = CTransformers(model=model, model_file=model_file, model_type=model_type, config=config) output = CTransformers(model=model, model_file=model_file, model_type=model_type, config=config)
message = output.invoke(inputs) message = output.invoke(inputs)
result = message.content if hasattr(message, "content") else message result = message.content if hasattr(message, "content") else message

View file

@ -14,41 +14,41 @@ class GoogleGenerativeAIComponent(CustomComponent):
def build_config(self): def build_config(self):
return { return {
"google_api_key": "google_api_key": {
{ "display_name":"Google API Key", "display_name": "Google API Key",
"info":"The Google API Key to use for the Google Generative AI.", "info": "The Google API Key to use for the Google Generative AI.",
} , },
"max_output_tokens":{ "max_output_tokens": {
"display_name":"Max Output Tokens", "display_name": "Max Output Tokens",
"info":"The maximum number of tokens to generate.", "info": "The maximum number of tokens to generate.",
}, },
"temperature": { "temperature": {
"display_name":"Temperature", "display_name": "Temperature",
"info":"Run inference with this temperature. Must by in the closed interval [0.0, 1.0].", "info": "Run inference with this temperature. Must by in the closed interval [0.0, 1.0].",
}, },
"top_k": { "top_k": {
"display_name":"Top K", "display_name": "Top K",
"info":"Decode using top-k sampling: consider the set of top_k most probable tokens. Must be positive.", "info": "Decode using top-k sampling: consider the set of top_k most probable tokens. Must be positive.",
"range_spec":RangeSpec(min=0, max=2, step=0.1), "range_spec": RangeSpec(min=0, max=2, step=0.1),
"advanced":True, "advanced": True,
}, },
"top_p": { "top_p": {
"display_name":"Top P", "display_name": "Top P",
"info":"The maximum cumulative probability of tokens to consider when sampling.", "info": "The maximum cumulative probability of tokens to consider when sampling.",
"advanced":True, "advanced": True,
}, },
"n": { "n": {
"display_name":"N", "display_name": "N",
"info":"Number of chat completions to generate for each prompt. Note that the API may not return the full n completions if duplicates are generated.", "info": "Number of chat completions to generate for each prompt. Note that the API may not return the full n completions if duplicates are generated.",
"advanced":True, "advanced": True,
}, },
"model": { "model": {
"display_name":"Model", "display_name": "Model",
"info":"The name of the model to use. Supported examples: gemini-pro", "info": "The name of the model to use. Supported examples: gemini-pro",
"options":["gemini-pro", "gemini-pro-vision"], "options": ["gemini-pro", "gemini-pro-vision"],
}, },
"code": { "code": {
"advanced":True, "advanced": True,
}, },
"inputs": {"display_name": "Input"}, "inputs": {"display_name": "Input"},
} }
@ -57,7 +57,7 @@ class GoogleGenerativeAIComponent(CustomComponent):
self, self,
google_api_key: str, google_api_key: str,
model: str, model: str,
inputs:str, inputs: str,
max_output_tokens: Optional[int] = None, max_output_tokens: Optional[int] = None,
temperature: float = 0.1, temperature: float = 0.1,
top_k: Optional[int] = None, top_k: Optional[int] = None,

View file

@ -47,4 +47,3 @@ class HuggingFaceEndpointsComponent(CustomComponent):
result = message.content if hasattr(message, "content") else message result = message.content if hasattr(message, "content") else message
self.status = result self.status = result
return result return result

View file

@ -57,7 +57,7 @@ class LlamaCppComponent(CustomComponent):
def build( def build(
self, self,
model_path: str, model_path: str,
inputs:str, inputs: str,
grammar: Optional[str] = None, grammar: Optional[str] = None,
cache: Optional[bool] = None, cache: Optional[bool] = None,
client: Optional[Any] = None, client: Optional[Any] = None,

View file

@ -171,7 +171,7 @@ class ChatOllamaComponent(CustomComponent):
self, self,
base_url: Optional[str], base_url: Optional[str],
model: str, model: str,
inputs:str, inputs: str,
mirostat: Optional[str], mirostat: Optional[str],
mirostat_eta: Optional[float] = None, mirostat_eta: Optional[float] = None,
mirostat_tau: Optional[float] = None, mirostat_tau: Optional[float] = None,

View file

@ -19,7 +19,6 @@ class PromptComponent(CustomComponent):
template: Prompt, template: Prompt,
**kwargs, **kwargs,
) -> Text: ) -> Text:
prompt_template = PromptTemplate.from_template(template) prompt_template = PromptTemplate.from_template(template)
attributes_to_check = ["text", "page_content"] attributes_to_check = ["text", "page_content"]

View file

@ -27,9 +27,7 @@ class RecordsAsTextComponent(CustomComponent):
if isinstance(records, Record): if isinstance(records, Record):
records = [records] records = [records]
formated_records = [ formated_records = [template.format(text=record.text, **record.data) for record in records]
template.format(text=record.text, **record.data) for record in records
]
result_string = "\n".join(formated_records) result_string = "\n".join(formated_records)
self.status = result_string self.status = result_string
return result_string return result_string

View file

@ -84,8 +84,7 @@ class ChromaComponent(CustomComponent):
if chroma_server_host is not None: if chroma_server_host is not None:
chroma_settings = chromadb.config.Settings( chroma_settings = chromadb.config.Settings(
chroma_server_cors_allow_origins=chroma_server_cors_allow_origins chroma_server_cors_allow_origins=chroma_server_cors_allow_origins or None,
or None,
chroma_server_host=chroma_server_host, chroma_server_host=chroma_server_host,
chroma_server_port=chroma_server_port or None, chroma_server_port=chroma_server_port or None,
chroma_server_grpc_port=chroma_server_grpc_port or None, chroma_server_grpc_port=chroma_server_grpc_port or None,
@ -100,9 +99,7 @@ class ChromaComponent(CustomComponent):
if documents is not None and embedding is not None: if documents is not None and embedding is not None:
if len(documents) == 0: if len(documents) == 0:
raise ValueError( raise ValueError("If documents are provided, there must be at least one document.")
"If documents are provided, there must be at least one document."
)
chroma = Chroma.from_documents( chroma = Chroma.from_documents(
documents=documents, # type: ignore documents=documents, # type: ignore
persist_directory=index_directory, persist_directory=index_directory,
@ -111,7 +108,5 @@ class ChromaComponent(CustomComponent):
client_settings=chroma_settings, client_settings=chroma_settings,
) )
else: else:
chroma = Chroma( chroma = Chroma(persist_directory=index_directory, client_settings=chroma_settings)
persist_directory=index_directory, client_settings=chroma_settings
)
return chroma return chroma

View file

@ -92,8 +92,7 @@ class ChromaSearchComponent(CustomComponent):
if chroma_server_host is not None: if chroma_server_host is not None:
chroma_settings = chromadb.config.Settings( chroma_settings = chromadb.config.Settings(
chroma_server_cors_allow_origins=chroma_server_cors_allow_origins chroma_server_cors_allow_origins=chroma_server_cors_allow_origins or None,
or None,
chroma_server_host=chroma_server_host, chroma_server_host=chroma_server_host,
chroma_server_port=chroma_server_port or None, chroma_server_port=chroma_server_port or None,
chroma_server_grpc_port=chroma_server_grpc_port or None, chroma_server_grpc_port=chroma_server_grpc_port or None,

View file

@ -11,9 +11,7 @@ if TYPE_CHECKING:
class SourceHandle(BaseModel): class SourceHandle(BaseModel):
baseClasses: List[str] = Field( baseClasses: List[str] = Field(..., description="List of base classes for the source handle.")
..., description="List of base classes for the source handle."
)
dataType: str = Field(..., description="Data type for the source handle.") dataType: str = Field(..., description="Data type for the source handle.")
id: str = Field(..., description="Unique identifier for the source handle.") id: str = Field(..., description="Unique identifier for the source handle.")
@ -21,9 +19,7 @@ class SourceHandle(BaseModel):
class TargetHandle(BaseModel): class TargetHandle(BaseModel):
fieldName: str = Field(..., description="Field name for the target handle.") fieldName: str = Field(..., description="Field name for the target handle.")
id: str = Field(..., description="Unique identifier for the target handle.") id: str = Field(..., description="Unique identifier for the target handle.")
inputTypes: Optional[List[str]] = Field( inputTypes: Optional[List[str]] = Field(None, description="List of input types for the target handle.")
None, description="List of input types for the target handle."
)
type: str = Field(..., description="Type of the target handle.") type: str = Field(..., description="Type of the target handle.")
@ -52,24 +48,16 @@ class Edge:
def validate_handles(self, source, target) -> None: def validate_handles(self, source, target) -> None:
if self.target_handle.inputTypes is None: if self.target_handle.inputTypes is None:
self.valid_handles = ( self.valid_handles = self.target_handle.type in self.source_handle.baseClasses
self.target_handle.type in self.source_handle.baseClasses
)
else: else:
self.valid_handles = ( self.valid_handles = (
any( any(baseClass in self.target_handle.inputTypes for baseClass in self.source_handle.baseClasses)
baseClass in self.target_handle.inputTypes
for baseClass in self.source_handle.baseClasses
)
or self.target_handle.type in self.source_handle.baseClasses or self.target_handle.type in self.source_handle.baseClasses
) )
if not self.valid_handles: if not self.valid_handles:
logger.debug(self.source_handle) logger.debug(self.source_handle)
logger.debug(self.target_handle) logger.debug(self.target_handle)
raise ValueError( raise ValueError(f"Edge between {source.vertex_type} and {target.vertex_type} " f"has invalid handles")
f"Edge between {source.vertex_type} and {target.vertex_type} "
f"has invalid handles"
)
def __setstate__(self, state): def __setstate__(self, state):
self.source_id = state["source_id"] self.source_id = state["source_id"]
@ -86,11 +74,7 @@ class Edge:
# Both lists contain strings and sometimes a string contains the value we are # Both lists contain strings and sometimes a string contains the value we are
# looking for e.g. comgin_out=["Chain"] and target_reqs=["LLMChain"] # looking for e.g. comgin_out=["Chain"] and target_reqs=["LLMChain"]
# so we need to check if any of the strings in source_types is in target_reqs # so we need to check if any of the strings in source_types is in target_reqs
self.valid = any( self.valid = any(output in target_req for output in self.source_types for target_req in self.target_reqs)
output in target_req
for output in self.source_types
for target_req in self.target_reqs
)
# Get what type of input the target node is expecting # Get what type of input the target node is expecting
self.matched_type = next( self.matched_type = next(
@ -101,10 +85,7 @@ class Edge:
if no_matched_type: if no_matched_type:
logger.debug(self.source_types) logger.debug(self.source_types)
logger.debug(self.target_reqs) logger.debug(self.target_reqs)
raise ValueError( raise ValueError(f"Edge between {source.vertex_type} and {target.vertex_type} " f"has no matched type")
f"Edge between {source.vertex_type} and {target.vertex_type} "
f"has no matched type"
)
def __repr__(self) -> str: def __repr__(self) -> str:
return ( return (
@ -116,11 +97,7 @@ class Edge:
return hash(self.__repr__()) return hash(self.__repr__())
def __eq__(self, __value: object) -> bool: def __eq__(self, __value: object) -> bool:
return ( return self.__repr__() == __value.__repr__() if isinstance(__value, Edge) else False
self.__repr__() == __value.__repr__()
if isinstance(__value, Edge)
else False
)
class ContractEdge(Edge): class ContractEdge(Edge):
@ -176,9 +153,7 @@ class ContractEdge(Edge):
return f"{self.source_id} -[{self.target_param}]-> {self.target_id}" return f"{self.source_id} -[{self.target_param}]-> {self.target_id}"
def log_transaction( def log_transaction(edge: ContractEdge, source: "Vertex", target: "Vertex", status, error=None):
edge: ContractEdge, source: "Vertex", target: "Vertex", status, error=None
):
try: try:
monitor_service = get_monitor_service() monitor_service = get_monitor_service()
clean_params = build_clean_params(target) clean_params = build_clean_params(target)

View file

@ -225,11 +225,7 @@ class Graph:
return return
self.vertices.remove(vertex) self.vertices.remove(vertex)
self.vertex_map.pop(vertex_id) self.vertex_map.pop(vertex_id)
self.edges = [ self.edges = [edge for edge in self.edges if edge.source_id != vertex_id and edge.target_id != vertex_id]
edge
for edge in self.edges
if edge.source_id != vertex_id and edge.target_id != vertex_id
]
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."""
@ -250,9 +246,7 @@ class Graph:
return return
for vertex in self.vertices: for vertex in self.vertices:
if not self._validate_vertex(vertex): if not self._validate_vertex(vertex):
raise ValueError( raise ValueError(f"{vertex.vertex_type} is not connected to any other components")
f"{vertex.vertex_type} is not connected to any other components"
)
def _validate_vertex(self, vertex: Vertex) -> bool: def _validate_vertex(self, vertex: Vertex) -> bool:
"""Validates a vertex.""" """Validates a vertex."""
@ -268,11 +262,7 @@ class Graph:
def get_vertex_edges(self, vertex_id: str) -> List[ContractEdge]: def get_vertex_edges(self, vertex_id: str) -> List[ContractEdge]:
"""Returns a list of edges for a given vertex.""" """Returns a list of edges for a given vertex."""
return [ return [edge for edge in self.edges if edge.source_id == vertex_id or edge.target_id == vertex_id]
edge
for edge in self.edges
if edge.source_id == vertex_id or edge.target_id == vertex_id
]
def get_vertices_with_target(self, vertex_id: str) -> List[Vertex]: def get_vertices_with_target(self, vertex_id: str) -> List[Vertex]:
"""Returns the vertices connected to a vertex.""" """Returns the vertices connected to a vertex."""
@ -310,9 +300,7 @@ class Graph:
def dfs(vertex): def dfs(vertex):
if state[vertex] == 1: if state[vertex] == 1:
# We have a cycle # We have a cycle
raise ValueError( raise ValueError("Graph contains a cycle, cannot perform topological sort")
"Graph contains a cycle, cannot perform topological sort"
)
if state[vertex] == 0: if state[vertex] == 0:
state[vertex] = 1 state[vertex] = 1
for edge in vertex.edges: for edge in vertex.edges:
@ -336,17 +324,11 @@ class Graph:
def get_predecessors(self, vertex): def get_predecessors(self, vertex):
"""Returns the predecessors of a vertex.""" """Returns the predecessors of a vertex."""
return [ return [self.get_vertex(source_id) for source_id in self.predecessor_map.get(vertex.id, [])]
self.get_vertex(source_id)
for source_id in self.predecessor_map.get(vertex.id, [])
]
def get_successors(self, vertex): def get_successors(self, vertex):
"""Returns the successors of a vertex.""" """Returns the successors of a vertex."""
return [ return [self.get_vertex(target_id) for target_id in self.successor_map.get(vertex.id, [])]
self.get_vertex(target_id)
for target_id in self.successor_map.get(vertex.id, [])
]
def get_vertex_neighbors(self, vertex: Vertex) -> Dict[Vertex, int]: def get_vertex_neighbors(self, vertex: Vertex) -> Dict[Vertex, int]:
"""Returns the neighbors of a vertex.""" """Returns the neighbors of a vertex."""
@ -385,9 +367,7 @@ class Graph:
edges.append(ContractEdge(source, target, edge)) edges.append(ContractEdge(source, target, edge))
return edges return edges
def _get_vertex_class( def _get_vertex_class(self, node_type: str, node_base_type: str, node_id: str) -> Type[Vertex]:
self, node_type: str, node_base_type: str, node_id: str
) -> Type[Vertex]:
"""Returns the node class based on the node type.""" """Returns the node class based on the node type."""
# First we check for the node_base_type # First we check for the node_base_type
node_name = node_id.split("-")[0] node_name = node_id.split("-")[0]
@ -417,18 +397,14 @@ class Graph:
vertex_type: str = vertex_data["type"] # type: ignore vertex_type: str = vertex_data["type"] # type: ignore
vertex_base_type: str = vertex_data["node"]["template"]["_type"] # type: ignore vertex_base_type: str = vertex_data["node"]["template"]["_type"] # type: ignore
VertexClass = self._get_vertex_class( VertexClass = self._get_vertex_class(vertex_type, vertex_base_type, vertex_data["id"])
vertex_type, vertex_base_type, vertex_data["id"]
)
vertex_instance = VertexClass(vertex, graph=self) vertex_instance = VertexClass(vertex, graph=self)
vertex_instance.set_top_level(self.top_level_vertices) vertex_instance.set_top_level(self.top_level_vertices)
vertices.append(vertex_instance) vertices.append(vertex_instance)
return vertices return vertices
def get_children_by_vertex_type( def get_children_by_vertex_type(self, vertex: Vertex, vertex_type: str) -> List[Vertex]:
self, vertex: Vertex, vertex_type: str
) -> List[Vertex]:
"""Returns the children of a vertex based on the vertex type.""" """Returns the children of a vertex based on the vertex type."""
children = [] children = []
vertex_types = [vertex.data["type"]] vertex_types = [vertex.data["type"]]
@ -440,9 +416,7 @@ class Graph:
def __repr__(self): def __repr__(self):
vertex_ids = [vertex.id for vertex in self.vertices] vertex_ids = [vertex.id for vertex in self.vertices]
edges_repr = "\n".join( edges_repr = "\n".join([f"{edge.source_id} --> {edge.target_id}" for edge in self.edges])
[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}"
def sort_up_to_vertex(self, vertex_id: str) -> "Graph": def sort_up_to_vertex(self, vertex_id: str) -> "Graph":
@ -473,9 +447,7 @@ class Graph:
"""Performs a layered topological sort of the vertices in the graph.""" """Performs a layered topological sort of the vertices in the graph."""
# Queue for vertices with no incoming edges # Queue for vertices with no incoming edges
queue = deque( queue = deque(vertex.id for vertex in vertices if self.in_degree_map[vertex.id] == 0)
vertex.id for vertex in vertices if self.in_degree_map[vertex.id] == 0
)
layers = [] layers = []
current_layer = 0 current_layer = 0
@ -531,10 +503,7 @@ class Graph:
return refined_layers return refined_layers
def sort_chat_inputs_first( def sort_chat_inputs_first(self, vertices_layers: List[List[str]]) -> List[List[str]]:
self, vertices_layers: List[List[str]]
) -> List[List[str]]:
chat_inputs_first = [] chat_inputs_first = []
for layer in vertices_layers: for layer in vertices_layers:
for vertex_id in layer: for vertex_id in layer:
@ -561,15 +530,11 @@ class Graph:
self.increment_run_count() self.increment_run_count()
return vertices_layers return vertices_layers
def sort_interface_components_first( def sort_interface_components_first(self, vertices_layers: List[List[str]]) -> List[List[str]]:
self, vertices_layers: List[List[str]]
) -> List[List[str]]:
"""Sorts the vertices in the graph so that vertices containing ChatInput or ChatOutput come first.""" """Sorts the vertices in the graph so that vertices containing ChatInput or ChatOutput come first."""
def contains_interface_component(vertex): def contains_interface_component(vertex):
return any( return any(component.value in vertex for component in InterfaceComponentTypes)
component.value in vertex for component in InterfaceComponentTypes
)
# Sort each inner list so that vertices containing ChatInput or ChatOutput come first # Sort each inner list so that vertices containing ChatInput or ChatOutput come first
sorted_vertices = [ sorted_vertices = [
@ -588,13 +553,9 @@ class Graph:
"""Sorts the vertices in the graph so that vertices with the lowest average build time come first.""" """Sorts the vertices in the graph so that vertices with the lowest average build time come first."""
if len(vertices_ids) == 1: if len(vertices_ids) == 1:
return vertices_ids return vertices_ids
vertices_ids.sort( vertices_ids.sort(key=lambda vertex_id: self.get_vertex(vertex_id).avg_build_time)
key=lambda vertex_id: self.get_vertex(vertex_id).avg_build_time
)
return vertices_ids return vertices_ids
sorted_vertices = [ sorted_vertices = [sort_layer_by_avg_build_time(layer) for layer in vertices_layers]
sort_layer_by_avg_build_time(layer) for layer in vertices_layers
]
return sorted_vertices return sorted_vertices

View file

@ -84,9 +84,7 @@ class Vertex:
): ):
if edge.target_id not in edge_results: if edge.target_id not in edge_results:
edge_results[edge.target_id] = {} edge_results[edge.target_id] = {}
edge_results[edge.target_id][edge.target_param] = await edge.get_result( edge_results[edge.target_id][edge.target_param] = await edge.get_result(source=self, target=target)
source=self, target=target
)
return edge_results return edge_results
def set_result(self, result: "ResultData") -> None: def set_result(self, result: "ResultData") -> None:
@ -96,9 +94,7 @@ 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( if self.is_interface_component and not isinstance(self._built_object, UnbuiltObject):
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
@ -108,11 +104,7 @@ class Vertex:
if isinstance(self._built_result, UnbuiltResult): if isinstance(self._built_result, UnbuiltResult):
return {} return {}
return ( return self._built_result if isinstance(self._built_result, dict) else {"result": self._built_result}
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
@ -174,29 +166,17 @@ class Vertex:
self.data = self._data["data"] self.data = self._data["data"]
self.output = self.data["node"]["base_classes"] self.output = self.data["node"]["base_classes"]
self.pinned = self.data["node"].get("pinned", False) self.pinned = self.data["node"].get("pinned", False)
template_dicts = { template_dicts = {key: value for key, value in self.data["node"]["template"].items() if isinstance(value, dict)}
key: value
for key, value in self.data["node"]["template"].items()
if isinstance(value, dict)
}
self.required_inputs = [ self.required_inputs = [
template_dicts[key]["type"] template_dicts[key]["type"] for key, value in template_dicts.items() if value["required"]
for key, value in template_dicts.items()
if value["required"]
] ]
self.optional_inputs = [ self.optional_inputs = [
template_dicts[key]["type"] template_dicts[key]["type"] for key, value in template_dicts.items() if not value["required"]
for key, value in template_dicts.items()
if not value["required"]
] ]
# Add the template_dicts[key]["input_types"] to the optional_inputs # Add the template_dicts[key]["input_types"] to the optional_inputs
self.optional_inputs.extend( self.optional_inputs.extend(
[ [input_type for value in template_dicts.values() for input_type in value.get("input_types", [])]
input_type
for value in template_dicts.values()
for input_type in value.get("input_types", [])
]
) )
template_dict = self.data["node"]["template"] template_dict = self.data["node"]["template"]
@ -239,11 +219,7 @@ class Vertex:
if self.graph is None: if self.graph is None:
raise ValueError("Graph not found") raise ValueError("Graph not found")
template_dict = { template_dict = {key: value for key, value in self.data["node"]["template"].items() if isinstance(value, dict)}
key: value
for key, value in self.data["node"]["template"].items()
if isinstance(value, dict)
}
params = {} params = {}
for edge in self.edges: for edge in self.edges:
@ -294,11 +270,7 @@ class Vertex:
# list of dicts, so we need to convert it to a dict # list of dicts, so we need to convert it to a dict
# before passing it to the build method # before passing it to the build method
if isinstance(val, list): if isinstance(val, list):
params[key] = { params[key] = {k: v for item in value.get("value", []) for k, v in item.items()}
k: v
for item in value.get("value", [])
for k, v in item.items()
}
elif isinstance(val, dict): elif isinstance(val, dict):
params[key] = val params[key] = val
elif value.get("type") == "int" and val is not None: elif value.get("type") == "int" and val is not None:
@ -358,9 +330,7 @@ class Vertex:
if isinstance(self._built_object, str): if isinstance(self._built_object, str):
self._built_result = self._built_object self._built_result = self._built_object
result = await generate_result( result = await generate_result(self._built_object, inputs, self.has_external_output, session_id)
self._built_object, inputs, self.has_external_output, session_id
)
self._built_result = result self._built_result = result
async def _build_each_node_in_params_dict(self, user_id=None): async def _build_each_node_in_params_dict(self, user_id=None):
@ -388,9 +358,7 @@ class Vertex:
""" """
return all(self._is_node(node) for node in value) return all(self._is_node(node) for node in value)
async def get_result( async def get_result(self, requester: Optional["Vertex"] = None, user_id=None, timeout=None) -> Any:
self, requester: Optional["Vertex"] = None, user_id=None, timeout=None
) -> Any:
# PLEASE REVIEW THIS IF STATEMENT # PLEASE REVIEW THIS IF STATEMENT
# Check if the Vertex was built already # Check if the Vertex was built already
if self._built: if self._built:
@ -424,9 +392,7 @@ class Vertex:
self._extend_params_list_with_result(key, result) self._extend_params_list_with_result(key, result)
self.params[key] = result self.params[key] = result
async def _build_list_of_nodes_and_update_params( async def _build_list_of_nodes_and_update_params(self, key, nodes: List["Vertex"], user_id=None):
self, key, nodes: List["Vertex"], user_id=None
):
""" """
Iterates over a list of nodes, builds each and updates the params dictionary. Iterates over a list of nodes, builds each and updates the params dictionary.
""" """
@ -478,9 +444,7 @@ class Vertex:
self._update_built_object_and_artifacts(result) self._update_built_object_and_artifacts(result)
except Exception as exc: except Exception as exc:
logger.exception(exc) logger.exception(exc)
raise ValueError( raise ValueError(f"Error building node {self.vertex_type}(ID:{self.id}): {str(exc)}") from exc
f"Error building node {self.vertex_type}(ID:{self.id}): {str(exc)}"
) from exc
def _update_built_object_and_artifacts(self, result): def _update_built_object_and_artifacts(self, result):
""" """
@ -539,15 +503,9 @@ class Vertex:
return self._built_object return self._built_object
# Get the requester edge # Get the requester edge
requester_edge = next( requester_edge = next((edge for edge in self.edges if edge.target_id == requester.id), None)
(edge for edge in self.edges if edge.target_id == requester.id), None
)
# Return the result of the requester edge # Return the result of the requester edge
return ( return None if requester_edge is None else await requester_edge.get_result(source=self, target=requester)
None
if requester_edge is None
else await requester_edge.get_result(source=self, target=requester)
)
def add_edge(self, edge: "ContractEdge") -> None: def add_edge(self, edge: "ContractEdge") -> None:
if edge not in self.edges: if edge not in self.edges:
@ -567,11 +525,7 @@ class Vertex:
def _built_object_repr(self): def _built_object_repr(self):
# Add a message with an emoji, stars for sucess, # Add a message with an emoji, stars for sucess,
return ( return "Built sucessfully ✨" if self._built_object is not None else "Failed to build 😵‍💫"
"Built sucessfully ✨"
if self._built_object is not None
else "Failed to build 😵‍💫"
)
class StatefulVertex(Vertex): class StatefulVertex(Vertex):

View file

@ -119,11 +119,9 @@ class DocumentLoaderVertex(StatefulVertex):
# show how many documents are in the list? # show how many documents are in the list?
if not isinstance(self._built_object, UnbuiltObject): if not isinstance(self._built_object, UnbuiltObject):
avg_length = sum( avg_length = sum(len(doc.page_content) for doc in self._built_object if hasattr(doc, "page_content")) / len(
len(doc.page_content) self._built_object
for doc in self._built_object )
if hasattr(doc, "page_content")
) / len(self._built_object)
return f"""{self.vertex_type}({len(self._built_object)} documents) return f"""{self.vertex_type}({len(self._built_object)} documents)
\nAvg. Document Length (characters): {int(avg_length)} \nAvg. Document Length (characters): {int(avg_length)}
Documents: {self._built_object[:3]}...""" Documents: {self._built_object[:3]}..."""
@ -196,9 +194,7 @@ class TextSplitterVertex(StatefulVertex):
# show how many documents are in the list? # show how many documents are in the list?
if not isinstance(self._built_object, UnbuiltObject): if not isinstance(self._built_object, UnbuiltObject):
avg_length = sum(len(doc.page_content) for doc in self._built_object) / len( avg_length = sum(len(doc.page_content) for doc in self._built_object) / len(self._built_object)
self._built_object
)
return f"""{self.vertex_type}({len(self._built_object)} documents) return f"""{self.vertex_type}({len(self._built_object)} documents)
\nAvg. Document Length (characters): {int(avg_length)} \nAvg. Document Length (characters): {int(avg_length)}
\nDocuments: {self._built_object[:3]}...""" \nDocuments: {self._built_object[:3]}..."""
@ -245,27 +241,18 @@ class PromptVertex(StatelessVertex):
user_id = kwargs.get("user_id", None) user_id = kwargs.get("user_id", None)
tools = kwargs.get("tools", []) tools = kwargs.get("tools", [])
if not self._built or force: if not self._built or force:
if ( if "input_variables" not in self.params or self.params["input_variables"] is None:
"input_variables" not in self.params
or self.params["input_variables"] is None
):
self.params["input_variables"] = [] self.params["input_variables"] = []
# Check if it is a ZeroShotPrompt and needs a tool # Check if it is a ZeroShotPrompt and needs a tool
if "ShotPrompt" in self.vertex_type: if "ShotPrompt" in self.vertex_type:
tools = ( tools = [tool_node.build(user_id=user_id) for tool_node in tools] if tools is not None else []
[tool_node.build(user_id=user_id) for tool_node in tools]
if tools is not None
else []
)
# flatten the list of tools if it is a list of lists # flatten the list of tools if it is a list of lists
# first check if it is a list # first check if it is a list
if tools and isinstance(tools, list) and isinstance(tools[0], list): if tools and isinstance(tools, list) and isinstance(tools[0], list):
tools = flatten_list(tools) tools = flatten_list(tools)
self.params["tools"] = tools self.params["tools"] = tools
prompt_params = [ prompt_params = [
key key for key, value in self.params.items() if isinstance(value, str) and key != "format_instructions"
for key, value in self.params.items()
if isinstance(value, str) and key != "format_instructions"
] ]
else: else:
prompt_params = ["template"] prompt_params = ["template"]
@ -275,20 +262,14 @@ class PromptVertex(StatelessVertex):
prompt_text = self.params[param] prompt_text = self.params[param]
variables = extract_input_variables_from_prompt(prompt_text) variables = extract_input_variables_from_prompt(prompt_text)
self.params["input_variables"].extend(variables) self.params["input_variables"].extend(variables)
self.params["input_variables"] = list( self.params["input_variables"] = list(set(self.params["input_variables"]))
set(self.params["input_variables"])
)
elif isinstance(self.params, dict): elif isinstance(self.params, dict):
self.params.pop("input_variables", None) self.params.pop("input_variables", None)
await self._build(user_id=user_id) await self._build(user_id=user_id)
def _built_object_repr(self): def _built_object_repr(self):
if ( if not self.artifacts or self._built_object is None or not hasattr(self._built_object, "format"):
not self.artifacts
or self._built_object is None
or not hasattr(self._built_object, "format")
):
return super()._built_object_repr() return super()._built_object_repr()
elif isinstance(self._built_object, UnbuiltObject): elif isinstance(self._built_object, UnbuiltObject):
return super()._built_object_repr() return super()._built_object_repr()
@ -300,9 +281,7 @@ class PromptVertex(StatelessVertex):
# so the prompt format doesn't break # so the prompt format doesn't break
artifacts.pop("handle_keys", None) artifacts.pop("handle_keys", None)
try: try:
if not hasattr(self._built_object, "template") and hasattr( if not hasattr(self._built_object, "template") and hasattr(self._built_object, "prompt"):
self._built_object, "prompt"
):
template = self._built_object.prompt.template template = self._built_object.prompt.template
else: else:
template = self._built_object.template template = self._built_object.template
@ -310,11 +289,7 @@ class PromptVertex(StatelessVertex):
if value: if value:
replace_key = "{" + key + "}" replace_key = "{" + key + "}"
template = template.replace(replace_key, value) template = template.replace(replace_key, value)
return ( return template if isinstance(template, str) else f"{self.vertex_type}({template})"
template
if isinstance(template, str)
else f"{self.vertex_type}({template})"
)
except KeyError: except KeyError:
return str(self._built_object) return str(self._built_object)
@ -422,15 +397,11 @@ class RoutingVertex(StatelessVertex):
else: else:
target_vertex.should_run = False target_vertex.should_run = False
else: else:
raise ValueError( raise ValueError(f"RoutingVertex {self.id} must have a condition in the _built_object")
f"RoutingVertex {self.id} must have a condition in the _built_object"
)
self._built_result = result self._built_result = result
else: else:
raise ValueError( raise ValueError(f"RoutingVertex {self.id} must have a _built_object with a condition and a result")
f"RoutingVertex {self.id} must have a _built_object with a condition and a result"
)
def dict_to_codeblock(d: dict) -> str: def dict_to_codeblock(d: dict) -> str:

View file

@ -21,9 +21,7 @@ class ComponentFunctionEntrypointNameNullError(HTTPException):
class Component: class Component:
ERROR_CODE_NULL: ClassVar[str] = "Python code must be provided." ERROR_CODE_NULL: ClassVar[str] = "Python code must be provided."
ERROR_FUNCTION_ENTRYPOINT_NAME_NULL: ClassVar[str] = ( ERROR_FUNCTION_ENTRYPOINT_NAME_NULL: ClassVar[str] = "The name of the entrypoint function must be provided."
"The name of the entrypoint function must be provided."
)
code: Optional[str] = None code: Optional[str] = None
_function_entrypoint_name: str = "build" _function_entrypoint_name: str = "build"

View file

@ -100,8 +100,7 @@ class CustomComponent(Component):
detail={ detail={
"error": "Type hint Error", "error": "Type hint Error",
"traceback": ( "traceback": (
"Prompt type is not supported in the build method." "Prompt type is not supported in the build method." " Try using PromptTemplate instead."
" Try using PromptTemplate instead."
), ),
}, },
) )
@ -115,20 +114,14 @@ class CustomComponent(Component):
if not self.code: if not self.code:
return {} return {}
component_classes = [ component_classes = [cls for cls in self.tree["classes"] if self.code_class_base_inheritance in cls["bases"]]
cls
for cls in self.tree["classes"]
if self.code_class_base_inheritance in cls["bases"]
]
if not component_classes: if not component_classes:
return {} return {}
# Assume the first Component class is the one we're interested in # Assume the first Component class is the one we're interested in
component_class = component_classes[0] component_class = component_classes[0]
build_methods = [ build_methods = [
method method for method in component_class["methods"] if method["name"] == self.function_entrypoint_name
for method in component_class["methods"]
if method["name"] == self.function_entrypoint_name
] ]
return build_methods[0] if build_methods else {} return build_methods[0] if build_methods else {}
@ -185,9 +178,7 @@ class CustomComponent(Component):
# Retrieve and decrypt the credential by name for the current user # Retrieve and decrypt the credential by name for the current user
db_service = get_db_service() db_service = get_db_service()
with session_getter(db_service) as session: with session_getter(db_service) as session:
return credential_service.get_credential( return credential_service.get_credential(user_id=self._user_id or "", name=name, session=session)
user_id=self._user_id or "", name=name, session=session
)
return get_credential return get_credential
@ -197,9 +188,7 @@ class CustomComponent(Component):
credential_service = get_credential_service() credential_service = get_credential_service()
db_service = get_db_service() db_service = get_db_service()
with session_getter(db_service) as session: with session_getter(db_service) as session:
return credential_service.list_credentials( return credential_service.list_credentials(user_id=self._user_id, session=session)
user_id=self._user_id, session=session
)
def index(self, value: int = 0): def index(self, value: int = 0):
"""Returns a function that returns the value at the given index in the iterable.""" """Returns a function that returns the value at the given index in the iterable."""
@ -250,11 +239,7 @@ class CustomComponent(Component):
if flow_id: if flow_id:
flow = session.query(Flow).get(flow_id) flow = session.query(Flow).get(flow_id)
elif flow_name: elif flow_name:
flow = ( flow = (session.query(Flow).filter(Flow.name == flow_name).filter(Flow.user_id == self.user_id)).first()
session.query(Flow)
.filter(Flow.name == flow_name)
.filter(Flow.user_id == self.user_id)
).first()
else: else:
raise ValueError("Either flow_name or flow_id must be provided") raise ValueError("Either flow_name or flow_id must be provided")

View file

@ -79,13 +79,9 @@ class DirectoryReader:
except Exception as e: except Exception as e:
logger.error(f"Error while loading component: {e}") logger.error(f"Error while loading component: {e}")
continue continue
items.append( items.append({"name": menu["name"], "path": menu["path"], "components": components})
{"name": menu["name"], "path": menu["path"], "components": components}
)
filtered = [menu for menu in items if menu["components"]] filtered = [menu for menu in items if menu["components"]]
logger.debug( logger.debug(f'Filtered components {"with errors" if with_errors else ""}: {len(filtered)}')
f'Filtered components {"with errors" if with_errors else ""}: {len(filtered)}'
)
return {"menu": filtered} return {"menu": filtered}
def validate_code(self, file_content): def validate_code(self, file_content):
@ -118,9 +114,7 @@ class DirectoryReader:
Walk through the directory path and return a list of all .py files. Walk through the directory path and return a list of all .py files.
""" """
if not (safe_path := self.get_safe_path()): if not (safe_path := self.get_safe_path()):
raise CustomComponentPathValueError( raise CustomComponentPathValueError(f"The path needs to start with '{self.base_path}'.")
f"The path needs to start with '{self.base_path}'."
)
file_list = [] file_list = []
for root, _, files in os.walk(safe_path): for root, _, files in os.walk(safe_path):
@ -165,9 +159,7 @@ class DirectoryReader:
for node in ast.walk(module): for node in ast.walk(module):
if isinstance(node, ast.FunctionDef): if isinstance(node, ast.FunctionDef):
for arg in node.args.args: for arg in node.args.args:
if self._is_type_hint_in_arg_annotation( if self._is_type_hint_in_arg_annotation(arg.annotation, type_hint_name):
arg.annotation, type_hint_name
):
return True return True
except SyntaxError: except SyntaxError:
# Returns False if the code is not valid Python # Returns False if the code is not valid Python
@ -185,16 +177,14 @@ class DirectoryReader:
and annotation.value.id == type_hint_name and annotation.value.id == type_hint_name
) )
def is_type_hint_used_but_not_imported( def is_type_hint_used_but_not_imported(self, type_hint_name: str, code: str) -> bool:
self, type_hint_name: str, code: str
) -> bool:
""" """
Check if a type hint is used but not imported in the given code. Check if a type hint is used but not imported in the given code.
""" """
try: try:
return self._is_type_hint_used_in_args( return self._is_type_hint_used_in_args(type_hint_name, code) and not self._is_type_hint_imported(
type_hint_name, code type_hint_name, code
) and not self._is_type_hint_imported(type_hint_name, code) )
except SyntaxError: except SyntaxError:
# Returns True if there's something wrong with the code # Returns True if there's something wrong with the code
# TODO : Find a better way to handle this # TODO : Find a better way to handle this
@ -215,9 +205,9 @@ class DirectoryReader:
return False, "Syntax error" return False, "Syntax error"
elif not self.validate_build(file_content): elif not self.validate_build(file_content):
return False, "Missing build function" return False, "Missing build function"
elif self._is_type_hint_used_in_args( elif self._is_type_hint_used_in_args("Optional", file_content) and not self._is_type_hint_imported(
"Optional", file_content "Optional", file_content
) and not self._is_type_hint_imported("Optional", file_content): ):
return ( return (
False, False,
"Type hint 'Optional' is used but not imported in the code.", "Type hint 'Optional' is used but not imported in the code.",
@ -233,9 +223,7 @@ class DirectoryReader:
from the .py files in the directory. from the .py files in the directory.
""" """
response = {"menu": []} response = {"menu": []}
logger.debug( logger.debug("-------------------- Building component menu list --------------------")
"-------------------- Building component menu list --------------------"
)
for file_path in file_paths: for file_path in file_paths:
menu_name = os.path.basename(os.path.dirname(file_path)) menu_name = os.path.basename(os.path.dirname(file_path))
@ -255,9 +243,7 @@ class DirectoryReader:
# first check if it's already CamelCase # first check if it's already CamelCase
if "_" in component_name: if "_" in component_name:
component_name_camelcase = " ".join( component_name_camelcase = " ".join(word.title() for word in component_name.split("_"))
word.title() for word in component_name.split("_")
)
else: else:
component_name_camelcase = component_name component_name_camelcase = component_name
@ -265,9 +251,7 @@ class DirectoryReader:
try: try:
output_types = self.get_output_types_from_code(result_content) output_types = self.get_output_types_from_code(result_content)
except Exception as exc: except Exception as exc:
logger.exception( logger.exception(f"Error while getting output types from code: {str(exc)}")
f"Error while getting output types from code: {str(exc)}"
)
output_types = [component_name_camelcase] output_types = [component_name_camelcase]
else: else:
output_types = [component_name_camelcase] output_types = [component_name_camelcase]
@ -283,9 +267,7 @@ class DirectoryReader:
if menu_result not in response["menu"]: if menu_result not in response["menu"]:
response["menu"].append(menu_result) response["menu"].append(menu_result)
logger.debug( logger.debug("-------------------- Component menu list built --------------------")
"-------------------- Component menu list built --------------------"
)
return response return response
@staticmethod @staticmethod

View file

@ -27,18 +27,14 @@ from langflow.utils import validate
from langflow.utils.util import get_base_classes from langflow.utils.util import get_base_classes
def add_output_types( def add_output_types(frontend_node: CustomComponentFrontendNode, return_types: List[str]):
frontend_node: CustomComponentFrontendNode, return_types: List[str]
):
"""Add output types to the frontend node""" """Add output types to the frontend node"""
for return_type in return_types: for return_type in return_types:
if return_type is None: if return_type is None:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": ( "error": ("Invalid return type. Please check your code and try again."),
"Invalid return type. Please check your code and try again."
),
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) )
@ -67,18 +63,14 @@ def reorder_fields(frontend_node: CustomComponentFrontendNode, field_order: List
frontend_node.template.fields = reordered_fields frontend_node.template.fields = reordered_fields
def add_base_classes( def add_base_classes(frontend_node: CustomComponentFrontendNode, return_types: List[str]):
frontend_node: CustomComponentFrontendNode, return_types: List[str]
):
"""Add base classes to the frontend node""" """Add base classes to the frontend node"""
for return_type_instance in return_types: for return_type_instance in return_types:
if return_type_instance is None: if return_type_instance is None:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": ( "error": ("Invalid return type. Please check your code and try again."),
"Invalid return type. Please check your code and try again."
),
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) )
@ -153,14 +145,10 @@ def add_new_custom_field(
# If options is a list, then it's a dropdown # If options is a list, then it's a dropdown
# If options is None, then it's a list of strings # If options is None, then it's a list of strings
is_list = isinstance(field_config.get("options"), list) is_list = isinstance(field_config.get("options"), list)
field_config["is_list"] = ( field_config["is_list"] = is_list or field_config.get("is_list", False) or field_contains_list
is_list or field_config.get("is_list", False) or field_contains_list
)
if "name" in field_config: if "name" in field_config:
warnings.warn( warnings.warn("The 'name' key in field_config is used to build the object and can't be changed.")
"The 'name' key in field_config is used to build the object and can't be changed."
)
required = field_config.pop("required", field_required) required = field_config.pop("required", field_required)
placeholder = field_config.pop("placeholder", "") placeholder = field_config.pop("placeholder", "")
@ -191,9 +179,7 @@ def add_extra_fields(frontend_node, field_config, function_args):
if "name" not in extra_field or extra_field["name"] == "self": if "name" not in extra_field or extra_field["name"] == "self":
continue continue
field_name, field_type, field_value, field_required = get_field_properties( field_name, field_type, field_value, field_required = get_field_properties(extra_field)
extra_field
)
config = field_config.get(field_name, {}) config = field_config.get(field_name, {})
frontend_node = add_new_custom_field( frontend_node = add_new_custom_field(
frontend_node, frontend_node,
@ -231,9 +217,7 @@ def run_build_config(
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": ( "error": ("Invalid type convertion. Please check your code and try again."),
"Invalid type convertion. Please check your code and try again."
),
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) from exc ) from exc
@ -261,9 +245,7 @@ def run_build_config(
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": ( "error": ("Invalid type convertion. Please check your code and try again."),
"Invalid type convertion. Please check your code and try again."
),
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) from exc ) from exc
@ -318,24 +300,16 @@ def build_custom_component_template(
frontend_node = build_frontend_node(custom_component.template_config) frontend_node = build_frontend_node(custom_component.template_config)
logger.debug("Updated attributes") logger.debug("Updated attributes")
field_config, custom_instance = run_build_config( field_config, custom_instance = run_build_config(custom_component, user_id=user_id, update_field=update_field)
custom_component, user_id=user_id, update_field=update_field
)
logger.debug("Built field config") logger.debug("Built field config")
entrypoint_args = custom_component.get_function_entrypoint_args entrypoint_args = custom_component.get_function_entrypoint_args
add_extra_fields(frontend_node, field_config, entrypoint_args) add_extra_fields(frontend_node, field_config, entrypoint_args)
frontend_node = add_code_field( frontend_node = add_code_field(frontend_node, custom_component.code, field_config.get("code", {}))
frontend_node, custom_component.code, field_config.get("code", {})
)
add_base_classes( add_base_classes(frontend_node, custom_component.get_function_entrypoint_return_type)
frontend_node, custom_component.get_function_entrypoint_return_type add_output_types(frontend_node, custom_component.get_function_entrypoint_return_type)
)
add_output_types(
frontend_node, custom_component.get_function_entrypoint_return_type
)
logger.debug("Added base classes") logger.debug("Added base classes")
reorder_fields(frontend_node, custom_instance._get_field_order()) reorder_fields(frontend_node, custom_instance._get_field_order())
@ -347,9 +321,7 @@ def build_custom_component_template(
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": ( "error": ("Invalid type convertion. Please check your code and try again."),
"Invalid type convertion. Please check your code and try again."
),
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) from exc ) from exc
@ -373,9 +345,7 @@ def build_custom_components(settings_service):
if not settings_service.settings.COMPONENTS_PATH: if not settings_service.settings.COMPONENTS_PATH:
return {} return {}
logger.info( logger.info(f"Building custom components from {settings_service.settings.COMPONENTS_PATH}")
f"Building custom components from {settings_service.settings.COMPONENTS_PATH}"
)
custom_components_from_file = {} custom_components_from_file = {}
processed_paths = set() processed_paths = set()
for path in settings_service.settings.COMPONENTS_PATH: for path in settings_service.settings.COMPONENTS_PATH:
@ -386,9 +356,7 @@ def build_custom_components(settings_service):
custom_component_dict = build_custom_component_list_from_path(path_str) custom_component_dict = build_custom_component_list_from_path(path_str)
if custom_component_dict: if custom_component_dict:
category = next(iter(custom_component_dict)) category = next(iter(custom_component_dict))
logger.info( logger.info(f"Loading {len(custom_component_dict[category])} component(s) from category {category}")
f"Loading {len(custom_component_dict[category])} component(s) from category {category}"
)
custom_components_from_file = merge_nested_dicts_with_renaming( custom_components_from_file = merge_nested_dicts_with_renaming(
custom_components_from_file, custom_component_dict custom_components_from_file, custom_component_dict
) )

View file

@ -146,9 +146,7 @@ async def process_runnable(runnable: Runnable, inputs: Union[dict, List[dict]]):
elif isinstance(inputs, dict) and hasattr(runnable, "ainvoke"): elif isinstance(inputs, dict) and hasattr(runnable, "ainvoke"):
result = await runnable.ainvoke(inputs) result = await runnable.ainvoke(inputs)
else: else:
raise ValueError( raise ValueError(f"Runnable {runnable} does not support inputs of type {type(inputs)}")
f"Runnable {runnable} does not support inputs of type {type(inputs)}"
)
# Check if the result is a list of AIMessages # Check if the result is a list of AIMessages
if isinstance(result, list) and all(isinstance(r, AIMessage) for r in result): if isinstance(result, list) and all(isinstance(r, AIMessage) for r in result):
result = [r.content for r in result] result = [r.content for r in result]
@ -157,9 +155,7 @@ async def process_runnable(runnable: Runnable, inputs: Union[dict, List[dict]]):
return result return result
async def process_inputs_dict( async def process_inputs_dict(built_object: Union[Chain, VectorStore, Runnable], inputs: dict):
built_object: Union[Chain, VectorStore, Runnable], inputs: dict
):
if isinstance(built_object, Chain): if isinstance(built_object, Chain):
if inputs is None: if inputs is None:
raise ValueError("Inputs must be provided for a Chain") raise ValueError("Inputs must be provided for a Chain")
@ -194,9 +190,7 @@ async def process_inputs_list(built_object: Runnable, inputs: List[dict]):
return await process_runnable(built_object, inputs) return await process_runnable(built_object, inputs)
async def generate_result( async def generate_result(built_object: Union[Chain, VectorStore, Runnable], inputs: Union[dict, List[dict]]):
built_object: Union[Chain, VectorStore, Runnable], inputs: Union[dict, List[dict]]
):
if isinstance(inputs, dict): if isinstance(inputs, dict):
result = await process_inputs_dict(built_object, inputs) result = await process_inputs_dict(built_object, inputs)
elif isinstance(inputs, List) and isinstance(built_object, Runnable): elif isinstance(inputs, List) and isinstance(built_object, Runnable):
@ -228,9 +222,7 @@ async def process_graph_cached(
if clear_cache: if clear_cache:
session_service.clear_session(session_id) session_service.clear_session(session_id)
if session_id is None: if session_id is None:
session_id = session_service.generate_key( session_id = session_service.generate_key(session_id=session_id, data_graph=data_graph)
session_id=session_id, data_graph=data_graph
)
# Load the graph using SessionService # Load the graph using SessionService
session = await session_service.load_session(session_id, data_graph) session = await session_service.load_session(session_id, data_graph)
graph, artifacts = session if session else (None, None) graph, artifacts = session if session else (None, None)
@ -266,18 +258,14 @@ async def build_graph_and_generate_result(
return Result(result=result, session_id=session_id) return Result(result=result, session_id=session_id)
def validate_input( def validate_input(graph_data: Dict[str, Any], tweaks: Dict[str, Dict[str, Any]]) -> List[Dict[str, Any]]:
graph_data: Dict[str, Any], tweaks: Dict[str, Dict[str, Any]]
) -> List[Dict[str, Any]]:
if not isinstance(graph_data, dict) or not isinstance(tweaks, dict): if not isinstance(graph_data, dict) or not isinstance(tweaks, dict):
raise ValueError("graph_data and tweaks should be dictionaries") raise ValueError("graph_data and tweaks should be dictionaries")
nodes = graph_data.get("data", {}).get("nodes") or graph_data.get("nodes") nodes = graph_data.get("data", {}).get("nodes") or graph_data.get("nodes")
if not isinstance(nodes, list): if not isinstance(nodes, list):
raise ValueError( raise ValueError("graph_data should contain a list of nodes under 'data' key or directly under 'nodes' key")
"graph_data should contain a list of nodes under 'data' key or directly under 'nodes' key"
)
return nodes return nodes
@ -286,9 +274,7 @@ def apply_tweaks(node: Dict[str, Any], node_tweaks: Dict[str, Any]) -> None:
template_data = node.get("data", {}).get("node", {}).get("template") template_data = node.get("data", {}).get("node", {}).get("template")
if not isinstance(template_data, dict): if not isinstance(template_data, dict):
logger.warning( logger.warning(f"Template data for node {node.get('id')} should be a dictionary")
f"Template data for node {node.get('id')} should be a dictionary"
)
return return
for tweak_name, tweak_value in node_tweaks.items(): for tweak_name, tweak_value in node_tweaks.items():
@ -303,9 +289,7 @@ def apply_tweaks_on_vertex(vertex: Vertex, node_tweaks: Dict[str, Any]) -> None:
vertex.params[tweak_name] = tweak_value vertex.params[tweak_name] = tweak_value
def process_tweaks( def process_tweaks(graph_data: Dict[str, Any], tweaks: Dict[str, Dict[str, Any]]) -> Dict[str, Any]:
graph_data: Dict[str, Any], tweaks: Dict[str, Dict[str, Any]]
) -> Dict[str, Any]:
""" """
This function is used to tweak the graph data using the node id and the tweaks dict. This function is used to tweak the graph data using the node id and the tweaks dict.
@ -326,9 +310,7 @@ def process_tweaks(
if node_tweaks := tweaks.get(node_id): if node_tweaks := tweaks.get(node_id):
apply_tweaks(node, node_tweaks) apply_tweaks(node, node_tweaks)
else: else:
logger.warning( logger.warning("Each node should be a dictionary with an 'id' key of type str")
"Each node should be a dictionary with an 'id' key of type str"
)
return graph_data return graph_data
@ -340,8 +322,6 @@ def process_tweaks_on_graph(graph: Graph, tweaks: Dict[str, Dict[str, Any]]):
if node_tweaks := tweaks.get(node_id): if node_tweaks := tweaks.get(node_id):
apply_tweaks_on_vertex(vertex, node_tweaks) apply_tweaks_on_vertex(vertex, node_tweaks)
else: else:
logger.warning( logger.warning("Each node should be a Vertex with an 'id' attribute of type str")
"Each node should be a Vertex with an 'id' attribute of type str"
)
return graph return graph

View file

@ -10,9 +10,7 @@ if TYPE_CHECKING:
class TransactionModel(BaseModel): class TransactionModel(BaseModel):
id: Optional[int] = Field(default=None, alias="id") id: Optional[int] = Field(default=None, alias="id")
timestamp: Optional[datetime] = Field( timestamp: Optional[datetime] = Field(default_factory=datetime.now, alias="timestamp")
default_factory=datetime.now, alias="timestamp"
)
source: str source: str
target: str target: str
target_args: dict target_args: dict
@ -54,9 +52,7 @@ class MessageModel(BaseModel):
def from_record(cls, record: "Record"): def from_record(cls, record: "Record"):
# first check if the record has all the required fields # first check if the record has all the required fields
if "sender" not in record.data and "sender_name" not in record.data: if "sender" not in record.data and "sender_name" not in record.data:
raise ValueError( raise ValueError("The record does not have the required fields 'sender' and 'sender_name' in the data.")
"The record does not have the required fields 'sender' and 'sender_name' in the data."
)
return cls( return cls(
sender=record.data["sender"], sender=record.data["sender"],
sender_name=record.data["sender_name"], sender_name=record.data["sender_name"],
@ -110,7 +106,6 @@ class VertexBuildModel(BaseModel):
class VertexBuildResponseModel(VertexBuildModel): class VertexBuildResponseModel(VertexBuildModel):
@field_serializer("data", "artifacts") @field_serializer("data", "artifacts")
def serialize_dict(v): def serialize_dict(v):
return v return v

View file

@ -43,9 +43,7 @@ class MonitorService(Service):
def ensure_tables_exist(self): def ensure_tables_exist(self):
for table_name, model in self.table_map.items(): for table_name, model in self.table_map.items():
drop_and_create_table_if_schema_mismatch( drop_and_create_table_if_schema_mismatch(str(self.db_path), table_name, model)
str(self.db_path), table_name, model
)
def add_row( def add_row(
self, self,

View file

@ -45,9 +45,7 @@ def model_to_sql_column_definitions(model: Type[BaseModel]) -> dict:
return columns return columns
def drop_and_create_table_if_schema_mismatch( def drop_and_create_table_if_schema_mismatch(db_path: str, table_name: str, model: Type[BaseModel]):
db_path: str, table_name: str, model: Type[BaseModel]
):
with duckdb.connect(db_path) as conn: with duckdb.connect(db_path) as conn:
# Get the current schema from the database # Get the current schema from the database
try: try:
@ -68,12 +66,8 @@ def drop_and_create_table_if_schema_mismatch(
conn.execute(f"CREATE SEQUENCE seq_{table_name} START 1;") conn.execute(f"CREATE SEQUENCE seq_{table_name} START 1;")
except duckdb.CatalogException: except duckdb.CatalogException:
pass pass
desired_schema[INDEX_KEY] = ( desired_schema[INDEX_KEY] = f"INTEGER PRIMARY KEY DEFAULT NEXTVAL('seq_{table_name}')"
f"INTEGER PRIMARY KEY DEFAULT NEXTVAL('seq_{table_name}')" columns_sql = ", ".join(f"{name} {data_type}" for name, data_type in desired_schema.items())
)
columns_sql = ", ".join(
f"{name} {data_type}" for name, data_type in desired_schema.items()
)
create_table_sql = f"CREATE TABLE {table_name} ({columns_sql})" create_table_sql = f"CREATE TABLE {table_name} ({columns_sql})"
conn.execute(create_table_sql) conn.execute(create_table_sql)

View file

@ -31,9 +31,7 @@ class SettingsService(Service):
for key in settings_dict: for key in settings_dict:
if key not in Settings.model_fields.keys(): if key not in Settings.model_fields.keys():
raise KeyError(f"Key {key} not found in settings") raise KeyError(f"Key {key} not found in settings")
logger.debug( logger.debug(f"Loading {len(settings_dict[key])} {key} from {file_path}")
f"Loading {len(settings_dict[key])} {key} from {file_path}"
)
settings = Settings(**settings_dict) settings = Settings(**settings_dict)
if not settings.CONFIG_DIR: if not settings.CONFIG_DIR:

View file

@ -95,9 +95,7 @@ async def build_vertex(
) )
# Emit the vertex build response # Emit the vertex build response
response = VertexBuildResponse( response = VertexBuildResponse(valid=valid, params=params, id=vertex.id, data=result_dict)
valid=valid, params=params, id=vertex.id, data=result_dict
)
await sio.emit("vertex_build", data=response.model_dump(), to=sid) await sio.emit("vertex_build", data=response.model_dump(), to=sid)
except Exception as exc: except Exception as exc:

View file

@ -88,9 +88,7 @@ class LocalStorageService(StorageService):
file_path.unlink() file_path.unlink()
logger.info(f"File {file_name} deleted successfully from flow {flow_id}.") logger.info(f"File {file_name} deleted successfully from flow {flow_id}.")
else: else:
logger.warning( logger.warning(f"Attempted to delete non-existent file {file_name} in flow {flow_id}.")
f"Attempted to delete non-existent file {file_name} in flow {flow_id}."
)
def teardown(self): def teardown(self):
"""Perform any cleanup operations when the service is being torn down.""" """Perform any cleanup operations when the service is being torn down."""

View file

@ -60,9 +60,7 @@ class TemplateField(BaseModel):
refresh: Optional[bool] = None refresh: Optional[bool] = None
"""Specifies if the field should be refreshed. Defaults to False.""" """Specifies if the field should be refreshed. Defaults to False."""
range_spec: Optional[RangeSpec] = Field( range_spec: Optional[RangeSpec] = Field(default=None, serialization_alias="rangeSpec")
default=None, serialization_alias="rangeSpec"
)
"""Range specification for the field. Defaults to None.""" """Range specification for the field. Defaults to None."""
title_case: bool = True title_case: bool = True

View file

@ -88,11 +88,7 @@ class FrontendNode(BaseModel):
def process_base_classes(self, base_classes: List[str]) -> List[str]: def process_base_classes(self, base_classes: List[str]) -> List[str]:
"""Removes unwanted base classes from the list of base classes.""" """Removes unwanted base classes from the list of base classes."""
return [ return [base_class for base_class in base_classes if base_class not in CLASSES_TO_REMOVE]
base_class
for base_class in base_classes
if base_class not in CLASSES_TO_REMOVE
]
@field_serializer("display_name") @field_serializer("display_name")
def process_display_name(self, display_name: str) -> str: def process_display_name(self, display_name: str) -> str:
@ -172,9 +168,7 @@ class FrontendNode(BaseModel):
return _type return _type
@staticmethod @staticmethod
def handle_special_field( def handle_special_field(field, key: str, _type: str, SPECIAL_FIELD_HANDLERS) -> str:
field, key: str, _type: str, SPECIAL_FIELD_HANDLERS
) -> str:
"""Handles special field by using the respective handler if present.""" """Handles special field by using the respective handler if present."""
handler = SPECIAL_FIELD_HANDLERS.get(key) handler = SPECIAL_FIELD_HANDLERS.get(key)
return handler(field) if handler else _type return handler(field) if handler else _type
@ -185,11 +179,7 @@ class FrontendNode(BaseModel):
if "dict" in _type.lower() and field.name == "dict_": if "dict" in _type.lower() and field.name == "dict_":
field.field_type = "file" field.field_type = "file"
field.file_types = [".json", ".yaml", ".yml"] field.file_types = [".json", ".yaml", ".yml"]
elif ( elif _type.startswith("Dict") or _type.startswith("Mapping") or _type.startswith("dict"):
_type.startswith("Dict")
or _type.startswith("Mapping")
or _type.startswith("dict")
):
field.field_type = "dict" field.field_type = "dict"
return _type return _type
@ -200,9 +190,7 @@ class FrontendNode(BaseModel):
field.value = value["default"] field.value = value["default"]
@staticmethod @staticmethod
def handle_specific_field_values( def handle_specific_field_values(field: TemplateField, key: str, name: Optional[str] = None) -> None:
field: TemplateField, key: str, name: Optional[str] = None
) -> None:
"""Handles specific field values for certain fields.""" """Handles specific field values for certain fields."""
if key == "headers": if key == "headers":
field.value = """{"Authorization": "Bearer <token>"}""" field.value = """{"Authorization": "Bearer <token>"}"""
@ -210,9 +198,7 @@ class FrontendNode(BaseModel):
FrontendNode._handle_api_key_specific_field_values(field, key, name) FrontendNode._handle_api_key_specific_field_values(field, key, name)
@staticmethod @staticmethod
def _handle_model_specific_field_values( def _handle_model_specific_field_values(field: TemplateField, key: str, name: Optional[str] = None) -> None:
field: TemplateField, key: str, name: Optional[str] = None
) -> None:
"""Handles specific field values related to models.""" """Handles specific field values related to models."""
model_dict = { model_dict = {
"OpenAI": constants.OPENAI_MODELS, "OpenAI": constants.OPENAI_MODELS,
@ -225,9 +211,7 @@ class FrontendNode(BaseModel):
field.is_list = True field.is_list = True
@staticmethod @staticmethod
def _handle_api_key_specific_field_values( def _handle_api_key_specific_field_values(field: TemplateField, key: str, name: Optional[str] = None) -> None:
field: TemplateField, key: str, name: Optional[str] = None
) -> None:
"""Handles specific field values related to API keys.""" """Handles specific field values related to API keys."""
if "api_key" in key and "OpenAI" in str(name): if "api_key" in key and "OpenAI" in str(name):
field.display_name = "OpenAI API Key" field.display_name = "OpenAI API Key"
@ -267,10 +251,7 @@ class FrontendNode(BaseModel):
@staticmethod @staticmethod
def should_be_password(key: str, show: bool) -> bool: def should_be_password(key: str, show: bool) -> bool:
"""Determines whether the field should be a password field.""" """Determines whether the field should be a password field."""
return ( return any(text in key.lower() for text in {"password", "token", "api", "key"}) and show
any(text in key.lower() for text in {"password", "token", "api", "key"})
and show
)
@staticmethod @staticmethod
def should_be_multiline(key: str) -> bool: def should_be_multiline(key: str) -> bool:

View file

@ -15,7 +15,6 @@ from langflow.template.template.base import Template
class MemoryFrontendNode(FrontendNode): class MemoryFrontendNode(FrontendNode):
pinned: bool = True pinned: bool = True
def add_extra_fields(self) -> None: def add_extra_fields(self) -> None:
@ -81,9 +80,7 @@ class MemoryFrontendNode(FrontendNode):
field.show = True field.show = True
field.advanced = False field.advanced = False
field.value = "" field.value = ""
field.info = ( field.info = INPUT_KEY_INFO if field.name == "input_key" else OUTPUT_KEY_INFO
INPUT_KEY_INFO if field.name == "input_key" else OUTPUT_KEY_INFO
)
if field.name == "memory_key": if field.name == "memory_key":
field.value = "chat_history" field.value = "chat_history"

View file

@ -46,9 +46,7 @@ def validate_code(code):
# Evaluate the function definition # Evaluate the function definition
for node in tree.body: for node in tree.body:
if isinstance(node, ast.FunctionDef): if isinstance(node, ast.FunctionDef):
code_obj = compile( code_obj = compile(ast.Module(body=[node], type_ignores=[]), "<string>", "exec")
ast.Module(body=[node], type_ignores=[]), "<string>", "exec"
)
try: try:
exec(code_obj) exec(code_obj)
except Exception as e: except Exception as e:
@ -92,23 +90,15 @@ def execute_function(code, function_name, *args, **kwargs):
exec_globals, exec_globals,
locals(), locals(),
) )
exec_globals[alias.asname or alias.name] = importlib.import_module( exec_globals[alias.asname or alias.name] = importlib.import_module(alias.name)
alias.name
)
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
raise ModuleNotFoundError( raise ModuleNotFoundError(f"Module {alias.name} not found. Please install it and try again.") from e
f"Module {alias.name} not found. Please install it and try again."
) from e
function_code = next( function_code = next(
node node for node in module.body if isinstance(node, ast.FunctionDef) and node.name == function_name
for node in module.body
if isinstance(node, ast.FunctionDef) and node.name == function_name
) )
function_code.parent = None function_code.parent = None
code_obj = compile( code_obj = compile(ast.Module(body=[function_code], type_ignores=[]), "<string>", "exec")
ast.Module(body=[function_code], type_ignores=[]), "<string>", "exec"
)
try: try:
exec(code_obj, exec_globals, locals()) exec(code_obj, exec_globals, locals())
except Exception as exc: except Exception as exc:
@ -135,23 +125,15 @@ def create_function(code, function_name):
if isinstance(node, ast.Import): if isinstance(node, ast.Import):
for alias in node.names: for alias in node.names:
try: try:
exec_globals[alias.asname or alias.name] = importlib.import_module( exec_globals[alias.asname or alias.name] = importlib.import_module(alias.name)
alias.name
)
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
raise ModuleNotFoundError( raise ModuleNotFoundError(f"Module {alias.name} not found. Please install it and try again.") from e
f"Module {alias.name} not found. Please install it and try again."
) from e
function_code = next( function_code = next(
node node for node in module.body if isinstance(node, ast.FunctionDef) and node.name == function_name
for node in module.body
if isinstance(node, ast.FunctionDef) and node.name == function_name
) )
function_code.parent = None function_code.parent = None
code_obj = compile( code_obj = compile(ast.Module(body=[function_code], type_ignores=[]), "<string>", "exec")
ast.Module(body=[function_code], type_ignores=[]), "<string>", "exec"
)
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
exec(code_obj, exec_globals, locals()) exec(code_obj, exec_globals, locals())
exec_globals[function_name] = locals()[function_name] exec_globals[function_name] = locals()[function_name]
@ -213,22 +195,16 @@ def prepare_global_scope(code, module):
if isinstance(node, ast.Import): if isinstance(node, ast.Import):
for alias in node.names: for alias in node.names:
try: try:
exec_globals[alias.asname or alias.name] = importlib.import_module( exec_globals[alias.asname or alias.name] = importlib.import_module(alias.name)
alias.name
)
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
raise ModuleNotFoundError( raise ModuleNotFoundError(f"Module {alias.name} not found. Please install it and try again.") from e
f"Module {alias.name} not found. Please install it and try again."
) from e
elif isinstance(node, ast.ImportFrom) and node.module is not None: elif isinstance(node, ast.ImportFrom) and node.module is not None:
try: try:
imported_module = importlib.import_module(node.module) imported_module = importlib.import_module(node.module)
for alias in node.names: for alias in node.names:
exec_globals[alias.name] = getattr(imported_module, alias.name) exec_globals[alias.name] = getattr(imported_module, alias.name)
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
raise ModuleNotFoundError( raise ModuleNotFoundError(f"Module {node.module} not found. Please install it and try again.") from e
f"Module {node.module} not found. Please install it and try again."
) from e
return exec_globals return exec_globals
@ -240,11 +216,7 @@ def extract_class_code(module, class_name):
:param class_name: Name of the class to extract :param class_name: Name of the class to extract
:return: AST node of the specified class :return: AST node of the specified class
""" """
class_code = next( class_code = next(node for node in module.body if isinstance(node, ast.ClassDef) and node.name == class_name)
node
for node in module.body
if isinstance(node, ast.ClassDef) and node.name == class_name
)
class_code.parent = None class_code.parent = None
return class_code return class_code
@ -257,9 +229,7 @@ def compile_class_code(class_code):
:param class_code: AST node of the class :param class_code: AST node of the class
:return: Compiled code object of the class :return: Compiled code object of the class
""" """
code_obj = compile( code_obj = compile(ast.Module(body=[class_code], type_ignores=[]), "<string>", "exec")
ast.Module(body=[class_code], type_ignores=[]), "<string>", "exec"
)
return code_obj return code_obj
@ -303,9 +273,7 @@ def get_default_imports(code_string):
langflow_imports = list(CUSTOM_COMPONENT_SUPPORTED_TYPES.keys()) langflow_imports = list(CUSTOM_COMPONENT_SUPPORTED_TYPES.keys())
necessary_imports = find_names_in_code(code_string, langflow_imports) necessary_imports = find_names_in_code(code_string, langflow_imports)
langflow_module = importlib.import_module("langflow.field_typing") langflow_module = importlib.import_module("langflow.field_typing")
default_imports.update( default_imports.update({name: getattr(langflow_module, name) for name in necessary_imports})
{name: getattr(langflow_module, name) for name in necessary_imports}
)
return default_imports return default_imports

View file

@ -5,6 +5,7 @@ import Tooltip from "../../components/TooltipComponent";
import IconComponent from "../../components/genericIconComponent"; import IconComponent from "../../components/genericIconComponent";
import InputComponent from "../../components/inputComponent"; import InputComponent from "../../components/inputComponent";
import { Button } from "../../components/ui/button"; import { Button } from "../../components/ui/button";
import Loading from "../../components/ui/loading";
import { Textarea } from "../../components/ui/textarea"; import { Textarea } from "../../components/ui/textarea";
import { priorityFields } from "../../constants/constants"; import { priorityFields } from "../../constants/constants";
import { BuildStatus } from "../../constants/enums"; import { BuildStatus } from "../../constants/enums";
@ -18,7 +19,6 @@ import { handleKeyDown, scapedJSONStringfy } from "../../utils/reactflowUtils";
import { nodeColors, nodeIconsLucide } from "../../utils/styleUtils"; import { nodeColors, nodeIconsLucide } from "../../utils/styleUtils";
import { classNames, cn, getFieldTitle } from "../../utils/utils"; import { classNames, cn, getFieldTitle } from "../../utils/utils";
import ParameterComponent from "./components/parameterComponent"; import ParameterComponent from "./components/parameterComponent";
import Loading from "../../components/ui/loading";
export default function GenericNode({ export default function GenericNode({
data, data,
@ -166,7 +166,7 @@ export default function GenericNode({
); );
const getStatusClassName = ( const getStatusClassName = (
validationStatus: validationStatusType | null, validationStatus: validationStatusType | null
) => { ) => {
if (validationStatus && validationStatus.valid) { if (validationStatus && validationStatus.valid) {
return "green-status"; return "green-status";
@ -181,10 +181,10 @@ export default function GenericNode({
const renderIconPlayOrPauseComponents = ( const renderIconPlayOrPauseComponents = (
buildStatus: BuildStatus | undefined, buildStatus: BuildStatus | undefined,
validationStatus: validationStatusType | null, validationStatus: validationStatusType | null
) => { ) => {
if (buildStatus === BuildStatus.BUILDING) { if (buildStatus === BuildStatus.BUILDING) {
return <Loading/> return <Loading />;
} else { } else {
const className = getStatusClassName(validationStatus); const className = getStatusClassName(validationStatus);
return <>{getIconPlayOrPauseComponent("Play", className)}</>; return <>{getIconPlayOrPauseComponent("Play", className)}</>;
@ -446,7 +446,9 @@ export default function GenericNode({
})); }));
}} }}
> >
<Tooltip title={<span>{pinned ? "Pin Output" : "Unpin Output"}</span>}> <Tooltip
title={<span>{pinned ? "Pin Output" : "Unpin Output"}</span>}
>
<div className="generic-node-status-position flex items-center"> <div className="generic-node-status-position flex items-center">
<IconComponent <IconComponent
name={"Pin"} name={"Pin"}
@ -461,12 +463,12 @@ export default function GenericNode({
)} )}
{showNode && ( {showNode && (
<Button <Button
variant="outline" variant="outline"
className={"h-9 px-1.5"} className={"h-9 px-1.5"}
onClick={() => { onClick={() => {
if(data?.build_status === BuildStatus.BUILDING || isBuilding) return; if (data?.build_status === BuildStatus.BUILDING || isBuilding)
buildFlow(data.id) return;
buildFlow(data.id);
}} }}
> >
<div> <div>
@ -499,7 +501,8 @@ export default function GenericNode({
<div className="generic-node-status-position flex items-center justify-center"> <div className="generic-node-status-position flex items-center justify-center">
{renderIconPlayOrPauseComponents( {renderIconPlayOrPauseComponents(
data?.build_status, data?.build_status,
validationStatus)} validationStatus
)}
</div> </div>
</Tooltip> </Tooltip>
</div> </div>

View file

@ -3,9 +3,6 @@ import IconComponent from "../../../components/genericIconComponent";
import { Textarea } from "../../../components/ui/textarea"; import { Textarea } from "../../../components/ui/textarea";
import { chatInputType } from "../../../types/components"; import { chatInputType } from "../../../types/components";
import { classNames } from "../../../utils/utils"; import { classNames } from "../../../utils/utils";
import { Button } from "../../ui/button";
import { Input } from "../../ui/input";
import { Popover, PopoverContent, PopoverTrigger } from "../../ui/popover";
export default function ChatInput({ export default function ChatInput({
lockChat, lockChat,
@ -113,7 +110,8 @@ export default function ChatInput({
)} )}
</button> </button>
</div> </div>
</div>{/* </div>
{/*
<Popover> <Popover>
<PopoverTrigger asChild> <PopoverTrigger asChild>
<Button variant="primary" className="h-13 px-4"> <Button variant="primary" className="h-13 px-4">

View file

@ -106,7 +106,7 @@ export default function newChatView(): JSX.Element {
}, []); }, []);
async function sendMessage(count = 1): Promise<void> { async function sendMessage(count = 1): Promise<void> {
if(isBuilding) return; if (isBuilding) return;
const { nodes, edges } = getFlow(); const { nodes, edges } = getFlow();
let nodeValidationErrors = validateNodes(nodes, edges); let nodeValidationErrors = validateNodes(nodes, edges);
if (nodeValidationErrors.length === 0) { if (nodeValidationErrors.length === 0) {

View file

@ -80,7 +80,6 @@ export default function NodeToolbarComponent({
window.open(url, "_blank", "noreferrer"); window.open(url, "_blank", "noreferrer");
}; };
useEffect(() => { useEffect(() => {
if (!showModalAdvanced) { if (!showModalAdvanced) {
onCloseAdvancedModal!(false); onCloseAdvancedModal!(false);

View file

@ -2,7 +2,6 @@ import { cloneDeep } from "lodash";
import { import {
Edge, Edge,
EdgeChange, EdgeChange,
MarkerType,
Node, Node,
NodeChange, NodeChange,
addEdge, addEdge,
@ -10,7 +9,6 @@ import {
applyNodeChanges, applyNodeChanges,
} from "reactflow"; } from "reactflow";
import { create } from "zustand"; import { create } from "zustand";
import { INPUT_TYPES, OUTPUT_TYPES } from "../constants/constants";
import { BuildStatus } from "../constants/enums"; import { BuildStatus } from "../constants/enums";
import { getFlowPool, updateFlowInDatabase } from "../controllers/API"; import { getFlowPool, updateFlowInDatabase } from "../controllers/API";
import { import {
@ -123,7 +121,7 @@ const useFlowStore = create<FlowStoreType>((set, get) => ({
}); });
const flowsManager = useFlowsManagerStore.getState(); const flowsManager = useFlowsManagerStore.getState();
if(!(get().isBuilding)){ if (!get().isBuilding) {
flowsManager.autoSaveCurrentFlow( flowsManager.autoSaveCurrentFlow(
newChange, newChange,
newEdges, newEdges,
@ -139,7 +137,7 @@ const useFlowStore = create<FlowStoreType>((set, get) => ({
}); });
const flowsManager = useFlowsManagerStore.getState(); const flowsManager = useFlowsManagerStore.getState();
if(!(get().isBuilding)){ if (!get().isBuilding) {
flowsManager.autoSaveCurrentFlow( flowsManager.autoSaveCurrentFlow(
get().nodes, get().nodes,
newChange, newChange,
@ -345,7 +343,7 @@ const useFlowStore = create<FlowStoreType>((set, get) => ({
sourceHandle: scapeJSONParse(connection.sourceHandle!), sourceHandle: scapeJSONParse(connection.sourceHandle!),
}, },
style: { stroke: "#555" }, style: { stroke: "#555" },
className:"stroke-foreground stroke-connection", className: "stroke-foreground stroke-connection",
}, },
oldEdges oldEdges
); );

View file

@ -1,5 +1,5 @@
import { create } from "zustand"; import { create } from "zustand";
import { checkHasApiKey, checkHasStore } from "../controllers/API"; import { checkHasStore } from "../controllers/API";
import { StoreStoreType } from "../types/zustand/store"; import { StoreStoreType } from "../types/zustand/store";
export const useStoreStore = create<StoreStoreType>((set) => ({ export const useStoreStore = create<StoreStoreType>((set) => ({

View file

@ -63,13 +63,15 @@ export async function buildVertices({
onBuildError, onBuildError,
verticesIds, verticesIds,
buildResults, buildResults,
stopBuild:()=>{stop=true} stopBuild: () => {
stop = true;
},
}); });
if(stop){ if (stop) {
break; break;
} }
} }
if(stop){ if (stop) {
break; break;
} }
} }
@ -96,7 +98,7 @@ async function buildVertex({
onBuildError?: (title, list, idList: string[]) => void; onBuildError?: (title, list, idList: string[]) => void;
verticesIds: string[]; verticesIds: string[];
buildResults: boolean[]; buildResults: boolean[];
stopBuild:()=>void; stopBuild: () => void;
}) { }) {
try { try {
const buildRes = await postBuildVertex(flowId, id); const buildRes = await postBuildVertex(flowId, id);

View file

@ -4,8 +4,8 @@ import {
Bell, Bell,
BookMarked, BookMarked,
BookmarkPlus, BookmarkPlus,
Boxes,
Bot, Bot,
Boxes,
Cable, Cable,
Check, Check,
CheckCircle2, CheckCircle2,