Fix formatting and icon naming conventions

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-02-28 21:50:10 -03:00
commit f8fe58bd08
72 changed files with 274 additions and 888 deletions

View file

@ -20,9 +20,7 @@ API_WORDS = ["api", "key", "token"]
def has_api_terms(word: str): def has_api_terms(word: str):
return "api" in word and ( return "api" in word and ("key" in word or ("token" in word and "tokens" not in word))
"key" in word or ("token" in word and "tokens" not in word)
)
def remove_api_keys(flow: dict): def remove_api_keys(flow: dict):
@ -32,11 +30,7 @@ def remove_api_keys(flow: dict):
node_data = node.get("data").get("node") node_data = node.get("data").get("node")
template = node_data.get("template") template = node_data.get("template")
for value in template.values(): for value in template.values():
if ( if isinstance(value, dict) and has_api_terms(value["name"]) and value.get("password"):
isinstance(value, dict)
and has_api_terms(value["name"])
and value.get("password")
):
value["value"] = None value["value"] = None
return flow return flow
@ -57,9 +51,7 @@ def build_input_keys_response(langchain_object, artifacts):
input_keys_response["input_keys"][key] = value input_keys_response["input_keys"][key] = value
# If the object has memory, that memory will have a memory_variables attribute # If the object has memory, that memory will have a memory_variables attribute
# memory variables should be removed from the input keys # memory variables should be removed from the input keys
if hasattr(langchain_object, "memory") and hasattr( if hasattr(langchain_object, "memory") and hasattr(langchain_object.memory, "memory_variables"):
langchain_object.memory, "memory_variables"
):
# Remove memory variables from input keys # Remove memory variables from input keys
input_keys_response["input_keys"] = { input_keys_response["input_keys"] = {
key: value key: value
@ -69,9 +61,7 @@ def build_input_keys_response(langchain_object, artifacts):
# Add memory variables to memory_keys # Add memory variables to memory_keys
input_keys_response["memory_keys"] = langchain_object.memory.memory_variables input_keys_response["memory_keys"] = langchain_object.memory.memory_variables
if hasattr(langchain_object, "prompt") and hasattr( if hasattr(langchain_object, "prompt") and hasattr(langchain_object.prompt, "template"):
langchain_object.prompt, "template"
):
input_keys_response["template"] = langchain_object.prompt.template input_keys_response["template"] = langchain_object.prompt.template
return input_keys_response return input_keys_response
@ -106,11 +96,7 @@ def raw_frontend_data_is_valid(raw_frontend_data):
def is_valid_data(frontend_node, raw_frontend_data): def is_valid_data(frontend_node, raw_frontend_data):
"""Check if the data is valid for processing.""" """Check if the data is valid for processing."""
return ( return frontend_node and "template" in frontend_node and raw_frontend_data_is_valid(raw_frontend_data)
frontend_node
and "template" in frontend_node
and raw_frontend_data_is_valid(raw_frontend_data)
)
def update_template_values(frontend_template, raw_template): def update_template_values(frontend_template, raw_template):
@ -150,9 +136,7 @@ def get_file_path_value(file_path):
# If the path is not in the cache dir, return empty string # If the path is not in the cache dir, return empty string
# This is to prevent access to files outside the cache dir # This is to prevent access to files outside the cache dir
# If the path is not a file, return empty string # If the path is not a file, return empty string
if not path.exists() or not str(path).startswith( if not path.exists() or not str(path).startswith(user_cache_dir("langflow", "langflow")):
user_cache_dir("langflow", "langflow")
):
return "" return ""
return file_path return file_path
@ -183,9 +167,7 @@ async def check_langflow_version(component: StoreComponentCreate):
langflow_version = get_lf_version_from_pypi() langflow_version = get_lf_version_from_pypi()
if langflow_version is None: if langflow_version is None:
raise HTTPException( raise HTTPException(status_code=500, detail="Unable to verify the latest version of Langflow")
status_code=500, detail="Unable to verify the latest version of Langflow"
)
elif langflow_version != component.last_tested_version: elif langflow_version != component.last_tested_version:
warnings.warn( warnings.warn(
f"Your version of Langflow ({component.last_tested_version}) is outdated. " f"Your version of Langflow ({component.last_tested_version}) is outdated. "

View file

@ -28,9 +28,7 @@ class AsyncStreamingLLMCallbackHandleSIO(AsyncCallbackHandler):
resp = ChatResponse(message=token, type="stream", intermediate_steps="") resp = ChatResponse(message=token, type="stream", intermediate_steps="")
await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump()) await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump())
async def on_tool_start( async def on_tool_start(self, serialized: Dict[str, Any], input_str: str, **kwargs: Any) -> Any:
self, serialized: Dict[str, Any], input_str: str, **kwargs: Any
) -> Any:
"""Run when tool starts running.""" """Run when tool starts running."""
resp = ChatResponse( resp = ChatResponse(
message="", message="",
@ -68,9 +66,7 @@ class AsyncStreamingLLMCallbackHandleSIO(AsyncCallbackHandler):
try: try:
# This is to emulate the stream of tokens # This is to emulate the stream of tokens
for resp in resps: for resp in resps:
await self.socketio_service.emit_token( await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump())
to=self.sid, data=resp.model_dump()
)
except Exception as exc: except Exception as exc:
logger.error(f"Error sending response: {exc}") logger.error(f"Error sending response: {exc}")
@ -96,9 +92,7 @@ class AsyncStreamingLLMCallbackHandleSIO(AsyncCallbackHandler):
resp = PromptResponse( resp = PromptResponse(
prompt=text, prompt=text,
) )
await self.socketio_service.emit_message( await self.socketio_service.emit_message(to=self.sid, data=resp.model_dump())
to=self.sid, data=resp.model_dump()
)
async def on_agent_action(self, action: AgentAction, **kwargs: Any): async def on_agent_action(self, action: AgentAction, **kwargs: Any):
log = f"Thought: {action.log}" log = f"Thought: {action.log}"
@ -108,9 +102,7 @@ class AsyncStreamingLLMCallbackHandleSIO(AsyncCallbackHandler):
logs = log.split("\n") logs = log.split("\n")
for log in logs: for log in logs:
resp = ChatResponse(message="", type="stream", intermediate_steps=log) resp = ChatResponse(message="", type="stream", intermediate_steps=log)
await self.socketio_service.emit_token( await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump())
to=self.sid, data=resp.model_dump()
)
else: else:
resp = ChatResponse(message="", type="stream", intermediate_steps=log) resp = ChatResponse(message="", type="stream", intermediate_steps=log)
await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump()) await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump())

View file

@ -98,12 +98,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_data_response = ResultDataResponse(results={}) result_data_response = ResultDataResponse(results={})
@ -195,9 +191,7 @@ async def build_vertex_stream(
else: else:
graph = cache.get("result") graph = cache.get("result")
else: else:
session_data = await session_service.load_session( session_data = await session_service.load_session(session_id, flow_id=flow_id)
session_id, flow_id=flow_id
)
graph, artifacts = session_data if session_data else (None, None) graph, artifacts = session_data if session_data else (None, None)
if not graph: if not graph:
raise ValueError(f"No graph found for {flow_id}.") raise ValueError(f"No graph found for {flow_id}.")

View file

@ -49,9 +49,7 @@ def get_all(
raise HTTPException(status_code=500, detail=str(exc)) from exc raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.post( @router.post("/run/{flow_id}", response_model=RunResponse, response_model_exclude_none=True)
"/run/{flow_id}", response_model=RunResponse, response_model_exclude_none=True
)
async def run_flow_with_caching( async def run_flow_with_caching(
session: Annotated[Session, Depends(get_session)], session: Annotated[Session, Depends(get_session)],
flow_id: str, flow_id: str,
@ -69,9 +67,7 @@ async def run_flow_with_caching(
input_values_dict = {} input_values_dict = {}
if session_id: if session_id:
session_data = await session_service.load_session( session_data = await session_service.load_session(session_id, flow_id=flow_id)
session_id, flow_id=flow_id
)
graph, artifacts = session_data if session_data else (None, None) graph, artifacts = session_data if session_data else (None, None)
task_result: Any = None task_result: Any = None
if not graph: if not graph:
@ -89,11 +85,7 @@ async def run_flow_with_caching(
else: else:
# Get the flow that matches the flow_id and belongs to the user # Get the flow that matches the flow_id and belongs to the user
# flow = session.query(Flow).filter(Flow.id == flow_id).filter(Flow.user_id == api_key_user.id).first() # flow = session.query(Flow).filter(Flow.id == flow_id).filter(Flow.user_id == api_key_user.id).first()
flow = session.exec( flow = session.exec(select(Flow).where(Flow.id == flow_id).where(Flow.user_id == api_key_user.id)).first()
select(Flow)
.where(Flow.id == flow_id)
.where(Flow.user_id == api_key_user.id)
).first()
if flow is None: if flow is None:
raise ValueError(f"Flow {flow_id} not found") raise ValueError(f"Flow {flow_id} not found")
@ -116,18 +108,12 @@ async def run_flow_with_caching(
# StatementError('(builtins.ValueError) badly formed hexadecimal UUID string') # StatementError('(builtins.ValueError) badly formed hexadecimal UUID string')
if "badly formed hexadecimal UUID string" in str(exc): if "badly formed hexadecimal UUID string" in str(exc):
# This means the Flow ID is not a valid UUID which means it can't find the flow # This means the Flow ID is not a valid UUID which means it can't find the flow
raise HTTPException( raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc
status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)
) from exc
except ValueError as exc: except ValueError as exc:
if f"Flow {flow_id} not found" in str(exc): if f"Flow {flow_id} not found" in str(exc):
raise HTTPException( raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc
status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)
) from exc
else: else:
raise HTTPException( raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=str(exc)) from exc
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=str(exc)
) from exc
@router.post( @router.post(
@ -156,8 +142,7 @@ async def process(
""" """
# Raise a depreciation warning # Raise a depreciation warning
logger.warning( logger.warning(
"The /process endpoint is deprecated and will be removed in a future version. " "The /process endpoint is deprecated and will be removed in a future version. " "Please use /run instead."
"Please use /run instead."
) )
raise HTTPException( raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, status_code=status.HTTP_400_BAD_REQUEST,
@ -229,16 +214,12 @@ async def custom_component(
built_frontend_node = build_custom_component_template(component, user_id=user.id) built_frontend_node = build_custom_component_template(component, user_id=user.id)
built_frontend_node = update_frontend_node_with_template_values( built_frontend_node = update_frontend_node_with_template_values(built_frontend_node, raw_code.frontend_node)
built_frontend_node, raw_code.frontend_node
)
return built_frontend_node return built_frontend_node
@router.post("/custom_component/reload", status_code=HTTPStatus.OK) @router.post("/custom_component/reload", status_code=HTTPStatus.OK)
async def reload_custom_component( async def reload_custom_component(path: str, user: User = Depends(get_current_active_user)):
path: str, user: User = Depends(get_current_active_user)
):
from langflow.interface.custom.utils import build_custom_component_template from langflow.interface.custom.utils import build_custom_component_template
try: try:
@ -260,8 +241,6 @@ async def custom_component_update(
): ):
component = CustomComponent(code=raw_code.code) component = CustomComponent(code=raw_code.code)
component_node = build_custom_component_template( component_node = build_custom_component_template(component, user_id=user.id, update_field=raw_code.field)
component, user_id=user.id, update_field=raw_code.field
)
# Update the field # Update the field
return component_node return component_node

View file

@ -21,26 +21,6 @@ class BuildStatus(Enum):
IN_PROGRESS = "in_progress" IN_PROGRESS = "in_progress"
class GraphData(BaseModel):
"""Data inside the exported flow."""
nodes: List[Dict[str, Any]]
edges: List[Dict[str, Any]]
class ExportedFlow(BaseModel):
"""Exported flow from Langflow."""
description: str
name: str
id: str
data: GraphData
class InputRequest(BaseModel):
input: dict
class TweaksRequest(BaseModel): class TweaksRequest(BaseModel):
tweaks: Optional[Dict[str, Dict[str, str]]] = Field(default_factory=dict) tweaks: Optional[Dict[str, Dict[str, str]]] = Field(default_factory=dict)
@ -178,9 +158,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

@ -17,7 +17,7 @@ class ConversationalAgent(CustomComponent):
display_name: str = "OpenAI Conversational Agent" display_name: str = "OpenAI Conversational Agent"
description: str = "Conversational Agent that can use OpenAI's function calling API" description: str = "Conversational Agent that can use OpenAI's function calling API"
icon = "OpenAI" icon = "OpenAI"
def build_config(self): def build_config(self):
openai_function_models = [ openai_function_models = [
"gpt-4-turbo-preview", "gpt-4-turbo-preview",

View file

@ -36,10 +36,8 @@ class ConversationChainComponent(CustomComponent):
# We need to check if it is a string or a BaseMessage # We need to check if it is a string or a BaseMessage
result_str: Text = "" result_str: Text = ""
if hasattr(result, "content") and isinstance(result.content, str): if hasattr(result, "content") and isinstance(result.content, str):
result_str = result.content result_str = result.content
elif isinstance(result, str): elif isinstance(result, str):
result_str = result result_str = result
else: else:
# is dict # is dict

View file

@ -7,9 +7,7 @@ from langflow.field_typing import BaseLanguageModel, Text
class LLMCheckerChainComponent(CustomComponent): class LLMCheckerChainComponent(CustomComponent):
display_name = "LLMCheckerChain" display_name = "LLMCheckerChain"
description = "" description = ""
documentation = ( documentation = "https://python.langchain.com/docs/modules/chains/additional/llm_checker"
"https://python.langchain.com/docs/modules/chains/additional/llm_checker"
)
def build_config(self): def build_config(self):
return { return {
@ -21,7 +19,6 @@ class LLMCheckerChainComponent(CustomComponent):
input_value: Text, input_value: Text,
llm: BaseLanguageModel, llm: BaseLanguageModel,
) -> Text: ) -> Text:
chain = LLMCheckerChain.from_llm(llm=llm) chain = LLMCheckerChain.from_llm(llm=llm)
response = chain.invoke({chain.input_key: input_value}) response = chain.invoke({chain.input_key: input_value})
result = response.get(chain.output_key, "") result = response.get(chain.output_key, "")

View file

@ -9,9 +9,7 @@ from langflow.field_typing import BaseLanguageModel, BaseMemory, Text
class LLMMathChainComponent(CustomComponent): class LLMMathChainComponent(CustomComponent):
display_name = "LLMMathChain" display_name = "LLMMathChain"
description = "Chain that interprets a prompt and executes python code to do math." description = "Chain that interprets a prompt and executes python code to do math."
documentation = ( documentation = "https://python.langchain.com/docs/modules/chains/additional/llm_math"
"https://python.langchain.com/docs/modules/chains/additional/llm_math"
)
def build_config(self): def build_config(self):
return { return {

View file

@ -47,19 +47,10 @@ class SQLGeneratorComponent(CustomComponent):
sql_query_chain = create_sql_query_chain(llm=llm, db=db, k=top_k) sql_query_chain = create_sql_query_chain(llm=llm, db=db, k=top_k)
else: else:
# Check if {question} is in the prompt # Check if {question} is in the prompt
if ( if "{question}" not in prompt_template.template or "question" not in prompt_template.input_variables:
"{question}" not in prompt_template.template raise ValueError("Prompt must contain `{question}` to be used with Natural Language to SQL.")
or "question" not in prompt_template.input_variables sql_query_chain = create_sql_query_chain(llm=llm, db=db, prompt=prompt_template, k=top_k)
): query_writer: Runnable = sql_query_chain | {"query": lambda x: x.replace("SQLQuery:", "").strip()}
raise ValueError(
"Prompt must contain `{question}` to be used with Natural Language to SQL."
)
sql_query_chain = create_sql_query_chain(
llm=llm, db=db, prompt=prompt_template, k=top_k
)
query_writer: Runnable = sql_query_chain | {
"query": lambda x: x.replace("SQLQuery:", "").strip()
}
response = query_writer.invoke({"question": input_value}) response = query_writer.invoke({"question": input_value})
query = response.get("query") query = response.get("query")
self.status = query self.status = query

View file

@ -70,17 +70,11 @@ class GatherRecordsComponent(CustomComponent):
glob = "**/*" if recursive else "*" glob = "**/*" if recursive else "*"
paths = walk_level(path_obj, depth) if depth else path_obj.glob(glob) paths = walk_level(path_obj, depth) if depth else path_obj.glob(glob)
file_paths = [ file_paths = [Text(p) for p in paths if p.is_file() and match_types(p) and is_not_hidden(p)]
Text(p)
for p in paths
if p.is_file() and match_types(p) and is_not_hidden(p)
]
return file_paths return file_paths
def parse_file_to_record( def parse_file_to_record(self, file_path: str, silent_errors: bool) -> Optional[Record]:
self, file_path: str, silent_errors: bool
) -> Optional[Record]:
# Use the partition function to load the file # Use the partition function to load the file
from unstructured.partition.auto import partition # type: ignore from unstructured.partition.auto import partition # type: ignore
@ -106,14 +100,9 @@ class GatherRecordsComponent(CustomComponent):
use_multithreading: bool, use_multithreading: bool,
) -> List[Optional[Record]]: ) -> List[Optional[Record]]:
if use_multithreading: if use_multithreading:
records = self.parallel_load_records( records = self.parallel_load_records(file_paths, silent_errors, max_concurrency)
file_paths, silent_errors, max_concurrency
)
else: else:
records = [ records = [self.parse_file_to_record(file_path, silent_errors) for file_path in file_paths]
self.parse_file_to_record(file_path, silent_errors)
for file_path in file_paths
]
records = list(filter(None, records)) records = list(filter(None, records))
return records return records
@ -142,20 +131,13 @@ class GatherRecordsComponent(CustomComponent):
if types is None: if types is None:
types = [] types = []
resolved_path = self.resolve_path(path) resolved_path = self.resolve_path(path)
file_paths = self.retrieve_file_paths( file_paths = self.retrieve_file_paths(resolved_path, types, load_hidden, recursive, depth)
resolved_path, types, load_hidden, recursive, depth
)
loaded_records = [] loaded_records = []
if use_multithreading: if use_multithreading:
loaded_records = self.parallel_load_records( loaded_records = self.parallel_load_records(file_paths, silent_errors, max_concurrency)
file_paths, silent_errors, max_concurrency
)
else: else:
loaded_records = [ loaded_records = [self.parse_file_to_record(file_path, silent_errors) for file_path in file_paths]
self.parse_file_to_record(file_path, silent_errors)
for file_path in file_paths
]
loaded_records = list(filter(None, loaded_records)) loaded_records = list(filter(None, loaded_records))
self.status = loaded_records self.status = loaded_records
return loaded_records return loaded_records

View file

@ -9,7 +9,7 @@ class HuggingFaceEmbeddingsComponent(CustomComponent):
documentation = ( documentation = (
"https://python.langchain.com/docs/modules/data_connection/text_embedding/integrations/sentence_transformers" "https://python.langchain.com/docs/modules/data_connection/text_embedding/integrations/sentence_transformers"
) )
icon="HuggingFace" icon = "HuggingFace"
def build_config(self): def build_config(self):
return { return {

View file

@ -9,8 +9,7 @@ class HuggingFaceInferenceAPIEmbeddingsComponent(CustomComponent):
display_name = "HuggingFaceInferenceAPIEmbeddings" display_name = "HuggingFaceInferenceAPIEmbeddings"
description = "HuggingFace sentence_transformers embedding models, API version." description = "HuggingFace sentence_transformers embedding models, API version."
documentation = "https://github.com/huggingface/text-embeddings-inference" documentation = "https://github.com/huggingface/text-embeddings-inference"
icon="HuggingFace" icon = "HuggingFace"
def build_config(self): def build_config(self):
return { return {

View file

@ -54,9 +54,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

@ -45,9 +45,7 @@ class ChatComponent(CustomComponent):
return [] return []
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("All of session_id, sender, and sender_name must be provided.")
"All of session_id, sender, and sender_name must be provided."
)
if isinstance(message, Record): if isinstance(message, Record):
record = message record = message
record.data.update( record.data.update(

View file

@ -12,7 +12,6 @@ class AmazonBedrockComponent(CustomComponent):
description: str = "LLM model from Amazon Bedrock." description: str = "LLM model from Amazon Bedrock."
icon = "Amazon" icon = "Amazon"
def build_config(self): def build_config(self):
return { return {
"model_id": { "model_id": {

View file

@ -10,7 +10,7 @@ from langflow import CustomComponent
class AnthropicLLM(CustomComponent): class AnthropicLLM(CustomComponent):
display_name: str = "AnthropicLLM" display_name: str = "AnthropicLLM"
description: str = "Anthropic Chat&Completion large language models." description: str = "Anthropic Chat&Completion large language models."
icon ="Anthropic" icon = "Anthropic"
def build_config(self): def build_config(self):
return { return {

View file

@ -10,8 +10,7 @@ from langflow.field_typing import BaseLanguageModel, NestedDict
class AnthropicComponent(CustomComponent): class AnthropicComponent(CustomComponent):
display_name = "Anthropic" display_name = "Anthropic"
description = "Anthropic large language models." description = "Anthropic large language models."
icon ="Anthropic" icon = "Anthropic"
def build_config(self): def build_config(self):
return { return {

View file

@ -9,7 +9,7 @@ class ChatAnthropicComponent(CustomComponent):
display_name = "ChatAnthropic" display_name = "ChatAnthropic"
description = "`Anthropic` chat large language models." description = "`Anthropic` chat large language models."
documentation = "https://python.langchain.com/docs/modules/model_io/models/chat/integrations/anthropic" documentation = "https://python.langchain.com/docs/modules/model_io/models/chat/integrations/anthropic"
icon ="Anthropic" icon = "Anthropic"
def build_config(self): def build_config(self):
return { return {

View file

@ -127,8 +127,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

@ -10,8 +10,7 @@ from langflow.field_typing import BaseLanguageModel
class ChatVertexAIComponent(CustomComponent): class ChatVertexAIComponent(CustomComponent):
display_name = "ChatVertexAI" display_name = "ChatVertexAI"
description = "`Vertex AI` Chat large language models API." description = "`Vertex AI` Chat large language models API."
icon="VertexAI" icon = "VertexAI"
def build_config(self): def build_config(self):
return { return {

View file

@ -8,8 +8,7 @@ from langflow import CustomComponent
class HuggingFaceEndpointsComponent(CustomComponent): class HuggingFaceEndpointsComponent(CustomComponent):
display_name: str = "Hugging Face Inference API" display_name: str = "Hugging Face Inference API"
description: str = "LLM model from Hugging Face Inference API." description: str = "LLM model from Hugging Face Inference API."
icon="HuggingFace" icon = "HuggingFace"
def build_config(self): def build_config(self):
return { return {

View file

@ -7,7 +7,7 @@ from langchain_community.llms.vertexai import VertexAI
class VertexAIComponent(CustomComponent): class VertexAIComponent(CustomComponent):
display_name = "VertexAI" display_name = "VertexAI"
description = "Google Vertex AI large language models" description = "Google Vertex AI large language models"
icon="VertexAI" icon = "VertexAI"
def build_config(self): def build_config(self):
return { return {

View file

@ -9,9 +9,7 @@ from langflow.field_typing import Text
class AnthropicLLM(LCModelComponent): class AnthropicLLM(LCModelComponent):
display_name: str = "AnthropicModel" display_name: str = "AnthropicModel"
description: str = ( description: str = "Generate text using Anthropic Chat&Completion large language models."
"Generate text using Anthropic Chat&Completion large language models."
)
icon = "Anthropic" icon = "Anthropic"
def build_config(self): def build_config(self):
@ -74,9 +72,7 @@ class AnthropicLLM(LCModelComponent):
try: try:
output = ChatAnthropic( output = ChatAnthropic(
model_name=model, model_name=model,
anthropic_api_key=( anthropic_api_key=(SecretStr(anthropic_api_key) if anthropic_api_key else None),
SecretStr(anthropic_api_key) if anthropic_api_key else None
),
max_tokens_to_sample=max_tokens, # type: ignore max_tokens_to_sample=max_tokens, # type: ignore
temperature=temperature, temperature=temperature,
anthropic_api_url=api_endpoint, anthropic_api_url=api_endpoint,

View file

@ -11,9 +11,7 @@ from langflow.field_typing import Text
class AzureChatOpenAIComponent(LCModelComponent): class AzureChatOpenAIComponent(LCModelComponent):
display_name: str = "AzureOpenAI Model" display_name: str = "AzureOpenAI Model"
description: str = "Generate text using LLM model from Azure OpenAI." description: str = "Generate text using LLM model from Azure OpenAI."
documentation: str = ( documentation: str = "https://python.langchain.com/docs/integrations/llms/azure_openai"
"https://python.langchain.com/docs/integrations/llms/azure_openai"
)
beta = False beta = False
icon = "Azure" icon = "Azure"

View file

@ -17,8 +17,7 @@ class VectaraSelfQueryRetriverComponent(CustomComponent):
description: str = "Implementation of Vectara Self Query Retriever" description: str = "Implementation of Vectara Self Query Retriever"
documentation = "https://python.langchain.com/docs/integrations/retrievers/self_query/vectara_self_query" documentation = "https://python.langchain.com/docs/integrations/retrievers/self_query/vectara_self_query"
beta = True beta = True
icon="Vectara" icon = "Vectara"
field_config = { field_config = {
"code": {"show": True}, "code": {"show": True},

View file

@ -32,9 +32,7 @@ class GetRequest(CustomComponent):
}, },
} }
def get_document( def get_document(self, session: requests.Session, url: str, headers: Optional[dict], timeout: int) -> Document:
self, session: requests.Session, url: str, headers: Optional[dict], timeout: int
) -> Document:
try: try:
response = session.get(url, headers=headers, timeout=int(timeout)) response = session.get(url, headers=headers, timeout=int(timeout))
try: try:

View file

@ -1,4 +1,5 @@
import uuid import uuid
from typing import Text
from langflow import CustomComponent from langflow import CustomComponent

View file

@ -67,16 +67,12 @@ class PostRequest(CustomComponent):
if not isinstance(document, list) and isinstance(document, Document): if not isinstance(document, list) and isinstance(document, Document):
documents: list[Document] = [document] documents: list[Document] = [document]
elif isinstance(document, list) and all( elif isinstance(document, list) and all(isinstance(doc, Document) for doc in document):
isinstance(doc, Document) for doc in document
):
documents = document documents = document
else: else:
raise ValueError("document must be a Document or a list of Documents") raise ValueError("document must be a Document or a list of Documents")
with requests.Session() as session: with requests.Session() as session:
documents = [ documents = [self.post_document(session, doc, url, headers) for doc in documents]
self.post_document(session, doc, url, headers) for doc in documents
]
self.repr_value = documents self.repr_value = documents
return documents return documents

View file

@ -27,10 +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, data=record.data, **record.data) for record in records]
template.format(text=record.text, data=record.data, **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

@ -42,9 +42,7 @@ class ShouldRunNext(CustomComponent):
result = result.get("response") result = result.get("response")
if result.lower() not in ["true", "false"]: if result.lower() not in ["true", "false"]:
raise ValueError( raise ValueError("The prompt should generate a boolean response (True or False).")
"The prompt should generate a boolean response (True or False)."
)
# The string should be the words true or false # The string should be the words true or false
# if not raise an error # if not raise an error
bool_result = result.lower() == "true" bool_result = result.lower() == "true"

View file

@ -41,9 +41,7 @@ class UpdateRequest(CustomComponent):
) -> Document: ) -> Document:
try: try:
if method == "PATCH": if method == "PATCH":
response = session.patch( response = session.patch(url, headers=headers, data=document.page_content)
url, headers=headers, data=document.page_content
)
elif method == "PUT": elif method == "PUT":
response = session.put(url, headers=headers, data=document.page_content) response = session.put(url, headers=headers, data=document.page_content)
else: else:
@ -80,17 +78,12 @@ class UpdateRequest(CustomComponent):
if not isinstance(document, list) and isinstance(document, Document): if not isinstance(document, list) and isinstance(document, Document):
documents: list[Document] = [document] documents: list[Document] = [document]
elif isinstance(document, list) and all( elif isinstance(document, list) and all(isinstance(doc, Document) for doc in document):
isinstance(doc, Document) for doc in document
):
documents = document documents = document
else: else:
raise ValueError("document must be a Document or a list of Documents") raise ValueError("document must be a Document or a list of Documents")
with requests.Session() as session: with requests.Session() as session:
documents = [ documents = [self.update_document(session, doc, url, headers, method) for doc in documents]
self.update_document(session, doc, url, headers, method)
for doc in documents
]
self.repr_value = documents self.repr_value = documents
return documents return documents

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,

View file

@ -93,8 +93,7 @@ class ChromaSearchComponent(LCVectorStoreComponent):
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

@ -34,9 +34,7 @@ class FAISSSearchComponent(LCVectorStoreComponent):
if not folder_path: if not folder_path:
raise ValueError("Folder path is required to save the FAISS index.") raise ValueError("Folder path is required to save the FAISS index.")
path = self.resolve_path(folder_path) path = self.resolve_path(folder_path)
vector_store = FAISS.load_local( vector_store = FAISS.load_local(folder_path=Text(path), embeddings=embedding, index_name=index_name)
folder_path=Text(path), embeddings=embedding, index_name=index_name
)
if not vector_store: if not vector_store:
raise ValueError("Failed to load the FAISS index.") raise ValueError("Failed to load the FAISS index.")

View file

@ -8,10 +8,8 @@ from langflow.field_typing import Document, Embeddings, NestedDict
class MongoDBAtlasComponent(CustomComponent): class MongoDBAtlasComponent(CustomComponent):
display_name = "MongoDB Atlas" display_name = "MongoDB Atlas"
description = ( description = "Construct a `MongoDB Atlas Vector Search` vector store from raw documents."
"Construct a `MongoDB Atlas Vector Search` vector store from raw documents." icon = "MongoDB"
)
icon="MongoDB"
def build_config(self): def build_config(self):
return { return {
@ -38,9 +36,7 @@ class MongoDBAtlasComponent(CustomComponent):
try: try:
from pymongo import MongoClient from pymongo import MongoClient
except ImportError: except ImportError:
raise ImportError( raise ImportError("Please install pymongo to use MongoDB Atlas Vector Store")
"Please install pymongo to use MongoDB Atlas Vector Store"
)
try: try:
mongo_client: MongoClient = MongoClient(mongodb_atlas_cluster_uri) mongo_client: MongoClient = MongoClient(mongodb_atlas_cluster_uri)
collection = mongo_client[db_name][collection_name] collection = mongo_client[db_name][collection_name]

View file

@ -10,7 +10,7 @@ from langflow.field_typing import Document, Embeddings, NestedDict
class QdrantComponent(CustomComponent): class QdrantComponent(CustomComponent):
display_name = "Qdrant" display_name = "Qdrant"
description = "Construct Qdrant wrapper from a list of texts." description = "Construct Qdrant wrapper from a list of texts."
icon="Qdrant" icon = "Qdrant"
def build_config(self): def build_config(self):
return { return {

View file

@ -38,9 +38,7 @@ class SupabaseSearchComponent(LCVectorStoreComponent):
supabase_url: str = "", supabase_url: str = "",
table_name: str = "", table_name: str = "",
) -> List[Record]: ) -> List[Record]:
supabase: Client = create_client( supabase: Client = create_client(supabase_url, supabase_key=supabase_service_key)
supabase_url, supabase_key=supabase_service_key
)
vector_store = SupabaseVectorStore( vector_store = SupabaseVectorStore(
client=supabase, client=supabase,
embedding=embedding, embedding=embedding,

View file

@ -14,9 +14,7 @@ from langflow.field_typing import BaseRetriever, Document
class VectaraComponent(CustomComponent): class VectaraComponent(CustomComponent):
display_name: str = "Vectara" display_name: str = "Vectara"
description: str = "Implementation of Vector Store using Vectara" description: str = "Implementation of Vector Store using Vectara"
documentation = ( documentation = "https://python.langchain.com/docs/integrations/vectorstores/vectara"
"https://python.langchain.com/docs/integrations/vectorstores/vectara"
)
beta = True beta = True
icon = "Vectara" icon = "Vectara"
field_config = { field_config = {

View file

@ -11,9 +11,7 @@ from langflow.schema import Record
class VectaraSearchComponent(VectaraComponent, LCVectorStoreComponent): class VectaraSearchComponent(VectaraComponent, LCVectorStoreComponent):
display_name: str = "Vectara Search" display_name: str = "Vectara Search"
description: str = "Search a Vectara Vector Store for similar documents." description: str = "Search a Vectara Vector Store for similar documents."
documentation = ( documentation = "https://python.langchain.com/docs/integrations/vectorstores/vectara"
"https://python.langchain.com/docs/integrations/vectorstores/vectara"
)
beta = True beta = True
icon = "Vectara" icon = "Vectara"

View file

@ -11,9 +11,7 @@ from langflow import CustomComponent
class WeaviateVectorStoreComponent(CustomComponent): class WeaviateVectorStoreComponent(CustomComponent):
display_name: str = "Weaviate" display_name: str = "Weaviate"
description: str = "Implementation of Vector Store using Weaviate" description: str = "Implementation of Vector Store using Weaviate"
documentation = ( documentation = "https://python.langchain.com/docs/integrations/vectorstores/weaviate"
"https://python.langchain.com/docs/integrations/vectorstores/weaviate"
)
beta = True beta = True
field_config = { field_config = {
"url": {"display_name": "Weaviate URL", "value": "http://localhost:8080"}, "url": {"display_name": "Weaviate URL", "value": "http://localhost:8080"},

View file

@ -11,9 +11,7 @@ from langflow.schema import Record
class WeaviateSearchVectorStore(WeaviateVectorStoreComponent, LCVectorStoreComponent): class WeaviateSearchVectorStore(WeaviateVectorStoreComponent, LCVectorStoreComponent):
display_name: str = "Weaviate Search" display_name: str = "Weaviate Search"
description: str = "Search a Weaviate Vector Store for similar documents." description: str = "Search a Weaviate Vector Store for similar documents."
documentation = ( documentation = "https://python.langchain.com/docs/integrations/vectorstores/weaviate"
"https://python.langchain.com/docs/integrations/vectorstores/weaviate"
)
beta = True beta = True
icon = "Weaviate" icon = "Weaviate"

View file

@ -10,7 +10,6 @@ from langflow.schema import Record
class LCVectorStoreComponent(CustomComponent): class LCVectorStoreComponent(CustomComponent):
display_name: str = "LC Vector Store" display_name: str = "LC Vector Store"
description: str = "Search a LC Vector Store for similar documents." description: str = "Search a LC Vector Store for similar documents."
beta: bool = True beta: bool = True
@ -37,14 +36,8 @@ class LCVectorStoreComponent(CustomComponent):
""" """
docs: List[Document] = [] docs: List[Document] = []
if ( if input_value and isinstance(input_value, str) and hasattr(vector_store, "search"):
input_value docs = vector_store.search(query=input_value, search_type=search_type.lower())
and isinstance(input_value, str)
and hasattr(vector_store, "search")
):
docs = vector_store.search(
query=input_value, search_type=search_type.lower()
)
else: else:
raise ValueError("Invalid inputs provided.") raise ValueError("Invalid inputs provided.")
return docs_to_records(docs) return docs_to_records(docs)

View file

@ -15,9 +15,7 @@ class PGVectorSearchComponent(PGVectorComponent, LCVectorStoreComponent):
display_name: str = "PGVector Search" display_name: str = "PGVector Search"
description: str = "Search a PGVector Store for similar documents." description: str = "Search a PGVector Store for similar documents."
documentation = ( documentation = "https://python.langchain.com/docs/integrations/vectorstores/pgvector"
"https://python.langchain.com/docs/integrations/vectorstores/pgvector"
)
def build_config(self): def build_config(self):
""" """

View file

@ -13,9 +13,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.")
@ -23,9 +21,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.")
@ -54,24 +50,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"]
@ -88,11 +76,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(
@ -103,10 +87,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 (
@ -118,11 +99,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):
@ -179,9 +156,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

@ -75,9 +75,7 @@ class Graph:
if getattr(vertex, attribute): if getattr(vertex, attribute):
getattr(self, f"_{attribute}_vertices").append(vertex.id) getattr(self, f"_{attribute}_vertices").append(vertex.id)
async def _run( async def _run(self, inputs: Dict[str, str], stream: bool) -> List[Optional["ResultData"]]:
self, inputs: Dict[str, str], stream: bool
) -> List[Optional["ResultData"]]:
"""Runs the graph with the given inputs.""" """Runs the graph with the given inputs."""
for vertex_id in self._is_input_vertices: for vertex_id in self._is_input_vertices:
vertex = self.get_vertex(vertex_id) vertex = self.get_vertex(vertex_id)
@ -100,9 +98,7 @@ class Graph:
outputs.append(vertex.result) outputs.append(vertex.result)
return outputs return outputs
async def run( async def run(self, inputs: Dict[str, Union[str, list[str]]], stream: bool) -> List[Optional["ResultData"]]:
self, inputs: Dict[str, Union[str, list[str]]], stream: bool
) -> List[Optional["ResultData"]]:
"""Runs the graph with the given inputs.""" """Runs the graph with the given inputs."""
# inputs is {"message": "Hello, world!"} # inputs is {"message": "Hello, world!"}
@ -114,9 +110,7 @@ class Graph:
if not isinstance(inputs_values, list): if not isinstance(inputs_values, list):
inputs_values = [inputs_values] inputs_values = [inputs_values]
for input_value in inputs_values: for input_value in inputs_values:
run_outputs = await self._run( run_outputs = await self._run({INPUT_FIELD_NAME: input_value}, stream=stream)
{INPUT_FIELD_NAME: input_value}, stream=stream
)
logger.debug(f"Run outputs: {run_outputs}") logger.debug(f"Run outputs: {run_outputs}")
outputs.extend(run_outputs) outputs.extend(run_outputs)
return outputs return outputs
@ -156,9 +150,7 @@ class Graph:
def build_parent_child_map(self): def build_parent_child_map(self):
parent_child_map = defaultdict(list) parent_child_map = defaultdict(list)
for vertex in self.vertices: for vertex in self.vertices:
parent_child_map[vertex.id] = [ parent_child_map[vertex.id] = [child.id for child in self.get_successors(vertex)]
child.id for child in self.get_successors(vertex)
]
return parent_child_map return parent_child_map
def increment_run_count(self): def increment_run_count(self):
@ -333,11 +325,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."""
@ -358,9 +346,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.display_name} is not connected to any other components")
f"{vertex.display_name} 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."""
@ -417,9 +403,7 @@ class Graph:
tasks = [] tasks = []
for vertex_id in layer: for vertex_id in layer:
vertex = self.get_vertex(vertex_id) vertex = self.get_vertex(vertex_id)
task = asyncio.create_task( task = asyncio.create_task(vertex.build(), name=f"layer-{layer_index}-vertex-{vertex_id}")
vertex.build(), name=f"layer-{layer_index}-vertex-{vertex_id}"
)
tasks.append(task) tasks.append(task)
logger.debug(f"Running layer {layer_index} with {len(tasks)} tasks") logger.debug(f"Running layer {layer_index} with {len(tasks)} tasks")
await self._execute_tasks(tasks) await self._execute_tasks(tasks)
@ -458,9 +442,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:
@ -484,17 +466,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."""
@ -533,9 +509,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]
@ -566,18 +540,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"]]
@ -589,9 +559,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) -> List[Vertex]: def sort_up_to_vertex(self, vertex_id: str) -> List[Vertex]:
@ -622,9 +590,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: List[List[str]] = [] layers: List[List[str]] = []
current_layer = 0 current_layer = 0
@ -680,9 +646,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:
@ -711,15 +675,11 @@ class Graph:
self._sorted_vertices_layers = vertices_layers self._sorted_vertices_layers = vertices_layers
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 = [
@ -731,22 +691,16 @@ class Graph:
] ]
return sorted_vertices return sorted_vertices
def sort_by_avg_build_time( def sort_by_avg_build_time(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 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."""
def sort_layer_by_avg_build_time(vertices_ids: List[str]) -> List[str]: def sort_layer_by_avg_build_time(vertices_ids: List[str]) -> List[str]:
"""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

@ -47,13 +47,8 @@ class Vertex:
self.will_stream = False self.will_stream = False
self.updated_raw_params = False self.updated_raw_params = False
self.id: str = data["id"] self.id: str = data["id"]
self.is_input = any( self.is_input = any(input_component_name in self.id for input_component_name in INPUT_COMPONENTS)
input_component_name in self.id for input_component_name in INPUT_COMPONENTS self.is_output = any(output_component_name in self.id for output_component_name in OUTPUT_COMPONENTS)
)
self.is_output = any(
output_component_name in self.id
for output_component_name in OUTPUT_COMPONENTS
)
self.has_session_id = None self.has_session_id = None
self._custom_component = None self._custom_component = None
self.has_external_input = False self.has_external_input = False
@ -87,17 +82,11 @@ class Vertex:
def set_state(self, state: str): def set_state(self, state: str):
self.state = VertexStates[state] self.state = VertexStates[state]
if ( if self.state == VertexStates.INACTIVE and self.graph.in_degree_map[self.id] < 2:
self.state == VertexStates.INACTIVE
and self.graph.in_degree_map[self.id] < 2
):
# If the vertex is inactive and has only one in degree # If the vertex is inactive and has only one in degree
# it means that it is not a merge point in the graph # it means that it is not a merge point in the graph
self.graph.inactive_vertices.add(self.id) self.graph.inactive_vertices.add(self.id)
elif ( elif self.state == VertexStates.ACTIVE and self.id in self.graph.inactive_vertices:
self.state == VertexStates.ACTIVE
and self.id in self.graph.inactive_vertices
):
self.graph.inactive_vertices.remove(self.id) self.graph.inactive_vertices.remove(self.id)
@property @property
@ -114,9 +103,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
@ -126,11 +113,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
@ -196,31 +179,19 @@ class Vertex:
self.selected_output_type = self.data["node"].get("selected_output_type") self.selected_output_type = self.data["node"].get("selected_output_type")
self.is_input = self.data["node"].get("is_input") or self.is_input self.is_input = self.data["node"].get("is_input") or self.is_input
self.is_output = self.data["node"].get("is_output") or self.is_output self.is_output = self.data["node"].get("is_output") or self.is_output
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.has_session_id = "session_id" in template_dicts self.has_session_id = "session_id" in template_dicts
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"]
@ -267,11 +238,7 @@ class Vertex:
self.updated_raw_params = False self.updated_raw_params = False
return return
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:
@ -322,11 +289,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:
@ -419,9 +382,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):
@ -451,9 +412,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:
@ -487,9 +446,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.
""" """
@ -532,7 +489,6 @@ class Vertex:
if self.base_type is None: if self.base_type is None:
raise ValueError(f"Base type for node {self.display_name} not found") raise ValueError(f"Base type for node {self.display_name} not found")
try: try:
result = await loading.instantiate_class( result = await loading.instantiate_class(
node_type=self.vertex_type, node_type=self.vertex_type,
base_type=self.base_type, base_type=self.base_type,
@ -544,9 +500,7 @@ class Vertex:
except Exception as exc: except Exception as exc:
logger.exception(exc) logger.exception(exc)
raise ValueError( raise ValueError(f"Error building node {self.display_name}: {str(exc)}") from exc
f"Error building node {self.display_name}: {str(exc)}"
) from exc
def _update_built_object_and_artifacts(self, result): def _update_built_object_and_artifacts(self, result):
""" """
@ -626,24 +580,16 @@ 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:
self.edges.append(edge) self.edges.append(edge)
def __repr__(self) -> str: def __repr__(self) -> str:
return ( return f"Vertex(display_name={self.display_name}, id={self.id}, data={self.data})"
f"Vertex(display_name={self.display_name}, id={self.id}, data={self.data})"
)
def __eq__(self, __o: object) -> bool: def __eq__(self, __o: object) -> bool:
try: try:
@ -656,11 +602,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

@ -123,11 +123,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.display_name}({len(self._built_object)} documents) return f"""{self.display_name}({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]}..."""
@ -200,9 +198,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]}..."""
@ -249,27 +245,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"]
@ -279,20 +266,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()
@ -304,9 +285,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
@ -314,11 +293,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)
@ -483,7 +458,6 @@ class RoutingVertex(StatelessVertex):
def dict_to_codeblock(d: dict) -> str: def dict_to_codeblock(d: dict) -> str:
serialized = {key: serialize_field(val) for key, val in d.items()} serialized = {key: serialize_field(val) for key, val in d.items()}
json_str = json.dumps(serialized, indent=4) json_str = json.dumps(serialized, indent=4)
return f"```json\n{json_str}\n```" return f"```json\n{json_str}\n```"

View file

@ -95,9 +95,7 @@ class CodeParser:
elif isinstance(node, ast.ImportFrom): elif isinstance(node, ast.ImportFrom):
for alias in node.names: for alias in node.names:
if alias.asname: if alias.asname:
self.data["imports"].append( self.data["imports"].append((node.module, f"{alias.name} as {alias.asname}"))
(node.module, f"{alias.name} as {alias.asname}")
)
else: else:
self.data["imports"].append((node.module, alias.name)) self.data["imports"].append((node.module, alias.name))
@ -146,9 +144,7 @@ class CodeParser:
return_type = None return_type = None
if node.returns: if node.returns:
return_type_str = ast.unparse(node.returns) return_type_str = ast.unparse(node.returns)
eval_env = self.construct_eval_env( eval_env = self.construct_eval_env(return_type_str, tuple(self.data["imports"]))
return_type_str, tuple(self.data["imports"])
)
try: try:
return_type = eval(return_type_str, eval_env) return_type = eval(return_type_str, eval_env)
@ -190,22 +186,14 @@ class CodeParser:
num_defaults = len(node.args.defaults) num_defaults = len(node.args.defaults)
num_missing_defaults = num_args - num_defaults num_missing_defaults = num_args - num_defaults
missing_defaults = [None] * num_missing_defaults missing_defaults = [None] * num_missing_defaults
default_values = [ default_values = [ast.unparse(default).strip("'") if default else None for default in node.args.defaults]
ast.unparse(default).strip("'") if default else None
for default in node.args.defaults
]
# Now check all default values to see if there # Now check all default values to see if there
# are any "None" values in the middle # are any "None" values in the middle
default_values = [ default_values = [None if value == "None" else value for value in default_values]
None if value == "None" else value for value in default_values
]
defaults = missing_defaults + default_values defaults = missing_defaults + default_values
args = [ args = [self.parse_arg(arg, default) for arg, default in zip(node.args.args, defaults)]
self.parse_arg(arg, default)
for arg, default in zip(node.args.args, defaults)
]
return args return args
def parse_varargs(self, node: ast.FunctionDef) -> List[Dict[str, Any]]: def parse_varargs(self, node: ast.FunctionDef) -> List[Dict[str, Any]]:
@ -223,17 +211,11 @@ class CodeParser:
""" """
Parses the keyword-only arguments of a function or method node. Parses the keyword-only arguments of a function or method node.
""" """
kw_defaults = [None] * ( kw_defaults = [None] * (len(node.args.kwonlyargs) - len(node.args.kw_defaults)) + [
len(node.args.kwonlyargs) - len(node.args.kw_defaults) ast.unparse(default) if default else None for default in node.args.kw_defaults
) + [
ast.unparse(default) if default else None
for default in node.args.kw_defaults
] ]
args = [ args = [self.parse_arg(arg, default) for arg, default in zip(node.args.kwonlyargs, kw_defaults)]
self.parse_arg(arg, default)
for arg, default in zip(node.args.kwonlyargs, kw_defaults)
]
return args return args
def parse_kwargs(self, node: ast.FunctionDef) -> List[Dict[str, Any]]: def parse_kwargs(self, node: ast.FunctionDef) -> List[Dict[str, Any]]:
@ -337,9 +319,7 @@ class CodeParser:
Extracts global variables from the code. Extracts global variables from the code.
""" """
global_var = { global_var = {
"targets": [ "targets": [t.id if hasattr(t, "id") else ast.dump(t) for t in node.targets],
t.id if hasattr(t, "id") else ast.dump(t) for t in node.targets
],
"value": ast.unparse(node.value), "value": ast.unparse(node.value),
} }
self.data["global_vars"].append(global_var) self.data["global_vars"].append(global_var)

View file

@ -119,9 +119,7 @@ class CustomComponent(Component):
def tree(self): def tree(self):
return self.get_code_tree(self.code or "") return self.get_code_tree(self.code or "")
def to_records( def to_records(self, data: Any, text_key: str = "text", data_key: str = "data") -> List[Record]:
self, data: Any, text_key: str = "text", data_key: str = "data"
) -> List[Record]:
""" """
Convert data into a list of records. Convert data into a list of records.
@ -148,9 +146,7 @@ class CustomComponent(Component):
return records return records
def create_references_from_records( def create_references_from_records(self, records: List[Record], include_data: bool = False) -> str:
self, records: List[Record], include_data: bool = False
) -> str:
""" """
Create references from a list of records. Create references from a list of records.
@ -185,8 +181,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."
), ),
}, },
) )
@ -200,20 +195,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 {}
@ -270,9 +259,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
@ -282,9 +269,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."""
@ -328,9 +313,7 @@ class CustomComponent(Component):
get_session = get_session or session_getter get_session = get_session or session_getter
db_service = get_db_service() db_service = get_db_service()
with get_session(db_service) as session: with get_session(db_service) as session:
flows = session.exec( flows = session.exec(select(Flow).where(Flow.user_id == self._user_id)).all()
select(Flow).where(Flow.user_id == self._user_id)
).all()
return flows return flows
except Exception as e: except Exception as e:
raise ValueError("Session is invalid") from e raise ValueError("Session is invalid") from e

View file

@ -80,13 +80,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):
@ -119,9 +115,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 = []
safe_path_obj = Path(safe_path) safe_path_obj = Path(safe_path)
@ -131,11 +125,7 @@ class DirectoryReader:
# any folders below [folder] will be ignored # any folders below [folder] will be ignored
# basically the parent folder of the file should be a # basically the parent folder of the file should be a
# folder in the safe_path # folder in the safe_path
if ( if file_path.is_file() and file_path.parent.parent == safe_path_obj and not file_path.name.startswith("__"):
file_path.is_file()
and file_path.parent.parent == safe_path_obj
and not file_path.name.startswith("__")
):
file_list.append(str(file_path)) file_list.append(str(file_path))
return file_list return file_list
@ -173,9 +163,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
@ -193,16 +181,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
@ -223,9 +209,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.",
@ -241,18 +227,14 @@ 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))
filename = os.path.basename(file_path) filename = os.path.basename(file_path)
validation_result, result_content = self.process_file(file_path) validation_result, result_content = self.process_file(file_path)
if not validation_result: if not validation_result:
logger.error( logger.error(f"Error while processing file {file_path}: {result_content}")
f"Error while processing file {file_path}: {result_content}"
)
menu_result = self.find_menu(response, menu_name) or { menu_result = self.find_menu(response, menu_name) or {
"name": menu_name, "name": menu_name,
@ -265,9 +247,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
@ -275,9 +255,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]
@ -293,9 +271,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

@ -8,11 +8,7 @@ from langflow.template.frontend_node.custom_components import (
def merge_nested_dicts_with_renaming(dict1, dict2): def merge_nested_dicts_with_renaming(dict1, dict2):
for key, value in dict2.items(): for key, value in dict2.items():
if ( if key in dict1 and isinstance(value, dict) and isinstance(dict1.get(key), dict):
key in dict1
and isinstance(value, dict)
and isinstance(dict1.get(key), dict)
):
for sub_key, sub_value in value.items(): for sub_key, sub_value in value.items():
# if sub_key in dict1[key]: # if sub_key in dict1[key]:
# new_key = get_new_key(dict1[key], sub_key) # new_key = get_new_key(dict1[key], sub_key)
@ -69,9 +65,7 @@ def build_custom_component_list_from_path(path: str):
file_list = load_files_from_path(path) file_list = load_files_from_path(path)
reader = DirectoryReader(path, False) reader = DirectoryReader(path, False)
valid_components, invalid_components = build_and_validate_all_files( valid_components, invalid_components = build_and_validate_all_files(reader, file_list)
reader, file_list
)
valid_menu = build_valid_menu(valid_components) valid_menu = build_valid_menu(valid_components)
invalid_menu = build_invalid_menu(invalid_components) invalid_menu = build_invalid_menu(invalid_components)
@ -118,9 +112,7 @@ def build_invalid_menu_items(menu_item):
menu_items[component_name] = component_template menu_items[component_name] = component_template
logger.debug(f"Added {component_name} to invalid menu.") logger.debug(f"Added {component_name} to invalid menu.")
except Exception as exc: except Exception as exc:
logger.exception( logger.exception(f"Error while creating custom component [{component_name}]: {str(exc)}")
f"Error while creating custom component [{component_name}]: {str(exc)}"
)
return menu_items return menu_items
@ -154,7 +146,5 @@ def build_menu_items(menu_item):
menu_items[component_name] = component_template menu_items[component_name] = component_template
except Exception as exc: except Exception as exc:
logger.error(f"Error loading Component: {component['output_types']}") logger.error(f"Error loading Component: {component['output_types']}")
logger.exception( logger.exception(f"Error while building custom component {component['output_types']}: {exc}")
f"Error while building custom component {component['output_types']}: {exc}"
)
return menu_items return menu_items

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
@ -316,24 +298,16 @@ def build_custom_component_template(
try: try:
frontend_node = build_frontend_node(custom_component.template_config) frontend_node = build_frontend_node(custom_component.template_config)
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
)
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
)
reorder_fields(frontend_node, custom_instance._get_field_order()) reorder_fields(frontend_node, custom_instance._get_field_order())
@ -344,9 +318,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
@ -370,9 +342,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:
@ -383,9 +353,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

@ -143,13 +143,9 @@ async def instantiate_based_on_type(
return class_object(**params) return class_object(**params)
async def instantiate_custom_component( async def instantiate_custom_component(node_type, class_object, params, user_id, vertex):
node_type, class_object, params, user_id, vertex
):
params_copy = params.copy() params_copy = params.copy()
class_object: Type["CustomComponent"] = eval_custom_component_code( class_object: Type["CustomComponent"] = eval_custom_component_code(params_copy.pop("code"))
params_copy.pop("code")
)
custom_component: "CustomComponent" = class_object( custom_component: "CustomComponent" = class_object(
user_id=user_id, user_id=user_id,
parameters=params_copy, parameters=params_copy,
@ -223,9 +219,7 @@ def instantiate_memory(node_type, class_object, params):
# I want to catch a specific attribute error that happens # I want to catch a specific attribute error that happens
# when the object does not have a cursor attribute # when the object does not have a cursor attribute
except Exception as exc: except Exception as exc:
if "object has no attribute 'cursor'" in str( if "object has no attribute 'cursor'" in str(exc) or 'object has no field "conn"' in str(exc):
exc
) or 'object has no field "conn"' in str(exc):
raise AttributeError( raise AttributeError(
( (
"Failed to build connection to database." "Failed to build connection to database."
@ -268,9 +262,7 @@ def instantiate_agent(node_type, class_object: Type[agent_module.Agent], params:
if class_method := getattr(class_object, method, None): if class_method := getattr(class_object, method, None):
agent = class_method(**params) agent = class_method(**params)
tools = params.get("tools", []) tools = params.get("tools", [])
return AgentExecutor.from_agent_and_tools( return AgentExecutor.from_agent_and_tools(agent=agent, tools=tools, handle_parsing_errors=True)
agent=agent, tools=tools, handle_parsing_errors=True
)
return load_agent_executor(class_object, params) return load_agent_executor(class_object, params)
@ -326,11 +318,7 @@ def instantiate_embedding(node_type, class_object, params: Dict):
try: try:
return class_object(**params) return class_object(**params)
except ValidationError: except ValidationError:
params = { params = {key: value for key, value in params.items() if key in class_object.model_fields}
key: value
for key, value in params.items()
if key in class_object.model_fields
}
return class_object(**params) return class_object(**params)
@ -342,9 +330,7 @@ def instantiate_vectorstore(class_object: Type[VectorStore], params: Dict):
if "texts" in params: if "texts" in params:
params["documents"] = params.pop("texts") params["documents"] = params.pop("texts")
if "documents" in params: if "documents" in params:
params["documents"] = [ params["documents"] = [doc for doc in params["documents"] if isinstance(doc, Document)]
doc for doc in params["documents"] if isinstance(doc, Document)
]
if initializer := vecstore_initializer.get(class_object.__name__): if initializer := vecstore_initializer.get(class_object.__name__):
vecstore = initializer(class_object, params) vecstore = initializer(class_object, params)
else: else:
@ -359,9 +345,7 @@ def instantiate_vectorstore(class_object: Type[VectorStore], params: Dict):
return vecstore return vecstore
def instantiate_documentloader( def instantiate_documentloader(node_type: str, class_object: Type[BaseLoader], params: Dict):
node_type: str, class_object: Type[BaseLoader], params: Dict
):
if "file_filter" in params: if "file_filter" in params:
# file_filter will be a string but we need a function # file_filter will be a string but we need a function
# that will be used to filter the files using file_filter # that will be used to filter the files using file_filter
@ -370,17 +354,13 @@ def instantiate_documentloader(
# in x and if it is, we will return True # in x and if it is, we will return True
file_filter = params.pop("file_filter") file_filter = params.pop("file_filter")
extensions = file_filter.split(",") extensions = file_filter.split(",")
params["file_filter"] = lambda x: any( params["file_filter"] = lambda x: any(extension.strip() in x for extension in extensions)
extension.strip() in x for extension in extensions
)
metadata = params.pop("metadata", None) metadata = params.pop("metadata", None)
if metadata and isinstance(metadata, str): if metadata and isinstance(metadata, str):
try: try:
metadata = orjson.loads(metadata) metadata = orjson.loads(metadata)
except json.JSONDecodeError as exc: except json.JSONDecodeError as exc:
raise ValueError( raise ValueError("The metadata you provided is not a valid JSON string.") from exc
"The metadata you provided is not a valid JSON string."
) from exc
if node_type == "WebBaseLoader": if node_type == "WebBaseLoader":
if web_path := params.pop("web_path", None): if web_path := params.pop("web_path", None):
@ -413,16 +393,12 @@ def instantiate_textsplitter(
"Try changing the chunk_size of the Text Splitter." "Try changing the chunk_size of the Text Splitter."
) from exc ) from exc
if ( if ("separator_type" in params and params["separator_type"] == "Text") or "separator_type" not in params:
"separator_type" in params and params["separator_type"] == "Text"
) or "separator_type" not in params:
params.pop("separator_type", None) params.pop("separator_type", None)
# separators might come in as an escaped string like \\n # separators might come in as an escaped string like \\n
# so we need to convert it to a string # so we need to convert it to a string
if "separators" in params: if "separators" in params:
params["separators"] = ( params["separators"] = params["separators"].encode().decode("unicode-escape")
params["separators"].encode().decode("unicode-escape")
)
text_splitter = class_object(**params) text_splitter = class_object(**params)
else: else:
from langchain.text_splitter import Language from langchain.text_splitter import Language
@ -449,8 +425,7 @@ def replace_zero_shot_prompt_with_prompt_template(nodes):
tools = [ tools = [
tool tool
for tool in nodes for tool in nodes
if tool["type"] != "chatOutputNode" if tool["type"] != "chatOutputNode" and "Tool" in tool["data"]["node"]["base_classes"]
and "Tool" in tool["data"]["node"]["base_classes"]
] ]
node["data"] = build_prompt_template(prompt=node["data"], tools=tools) node["data"] = build_prompt_template(prompt=node["data"], tools=tools)
break break
@ -464,9 +439,7 @@ def load_agent_executor(agent_class: type[agent_module.Agent], params, **kwargs)
# agent has hidden args for memory. might need to be support # agent has hidden args for memory. might need to be support
# memory = params["memory"] # memory = params["memory"]
# if allowed_tools is not a list or set, make it a list # if allowed_tools is not a list or set, make it a list
if not isinstance(allowed_tools, (list, set)) and isinstance( if not isinstance(allowed_tools, (list, set)) and isinstance(allowed_tools, BaseTool):
allowed_tools, BaseTool
):
allowed_tools = [allowed_tools] allowed_tools = [allowed_tools]
tool_names = [tool.name for tool in allowed_tools] tool_names = [tool.name for tool in allowed_tools]
# Agent class requires an output_parser but Agent classes # Agent class requires an output_parser but Agent classes
@ -494,10 +467,7 @@ def build_prompt_template(prompt, tools):
format_instructions = prompt["node"]["template"]["format_instructions"]["value"] format_instructions = prompt["node"]["template"]["format_instructions"]["value"]
tool_strings = "\n".join( tool_strings = "\n".join(
[ [f"{tool['data']['node']['name']}: {tool['data']['node']['description']}" for tool in tools]
f"{tool['data']['node']['name']}: {tool['data']['node']['description']}"
for tool in tools
]
) )
tool_names = ", ".join([tool["data"]["node"]["name"] for tool in tools]) tool_names = ", ".join([tool["data"]["node"]["name"] for tool in tools])
format_instructions = format_instructions.format(tool_names=tool_names) format_instructions = format_instructions.format(tool_names=tool_names)

View file

@ -18,9 +18,7 @@ from langflow.utils.logger import configure
def get_lifespan(fix_migration=False, socketio_server=None): def get_lifespan(fix_migration=False, socketio_server=None):
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI): async def lifespan(app: FastAPI):
initialize_services( initialize_services(fix_migration=fix_migration, socketio_server=socketio_server)
fix_migration=fix_migration, socketio_server=socketio_server
)
setup_llm_caching() setup_llm_caching()
LangfuseInstance.update() LangfuseInstance.update()
yield yield
@ -33,9 +31,7 @@ def create_app():
"""Create the FastAPI app and include the router.""" """Create the FastAPI app and include the router."""
configure() configure()
socketio_server = socketio.AsyncServer( socketio_server = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", logger=True)
async_mode="asgi", cors_allowed_origins="*", logger=True
)
lifespan = get_lifespan(socketio_server=socketio_server) lifespan = get_lifespan(socketio_server=socketio_server)
app = FastAPI(lifespan=lifespan) app = FastAPI(lifespan=lifespan)
origins = ["*"] origins = ["*"]
@ -102,9 +98,7 @@ def get_static_files_dir():
return frontend_path / "frontend" return frontend_path / "frontend"
def setup_app( def setup_app(static_files_dir: Optional[Path] = None, backend_only: bool = False) -> FastAPI:
static_files_dir: Optional[Path] = None, backend_only: bool = False
) -> FastAPI:
"""Setup the FastAPI app.""" """Setup the FastAPI app."""
# get the directory of the current file # get the directory of the current file
if not static_files_dir: if not static_files_dir:

View file

@ -36,9 +36,7 @@ def get_langfuse_callback(trace_id):
return None return None
def flush_langfuse_callback_if_present( def flush_langfuse_callback_if_present(callbacks: List[Union[BaseCallbackHandler, "CallbackHandler"]]):
callbacks: List[Union[BaseCallbackHandler, "CallbackHandler"]]
):
""" """
If langfuse callback is present, run callback.langfuse.flush() If langfuse callback is present, run callback.langfuse.flush()
""" """
@ -79,15 +77,9 @@ async def get_result_and_steps(langchain_object, inputs: Union[dict, str], **kwa
# if langfuse callback is present, run callback.langfuse.flush() # if langfuse callback is present, run callback.langfuse.flush()
flush_langfuse_callback_if_present(callbacks) flush_langfuse_callback_if_present(callbacks)
intermediate_steps = ( intermediate_steps = output.get("intermediate_steps", []) if isinstance(output, dict) else []
output.get("intermediate_steps", []) if isinstance(output, dict) else []
)
result = ( result = output.get(langchain_object.output_keys[0]) if isinstance(output, dict) else output
output.get(langchain_object.output_keys[0])
if isinstance(output, dict)
else output
)
try: try:
thought = format_actions(intermediate_steps) if intermediate_steps else "" thought = format_actions(intermediate_steps) if intermediate_steps else ""
except Exception as exc: except Exception as exc:

View file

@ -126,9 +126,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]
@ -137,9 +135,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")
@ -174,9 +170,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):
@ -214,9 +208,7 @@ async def run_graph(
else: else:
graph_data = graph._graph_data graph_data = graph._graph_data
if not session_id and session_service is not None: if not session_id and session_service is not None:
session_id = session_service.generate_key( session_id = session_service.generate_key(session_id=flow_id, data_graph=graph_data)
session_id=flow_id, data_graph=graph_data
)
if inputs is None: if inputs is None:
inputs = {} inputs = {}
@ -226,18 +218,14 @@ async def run_graph(
return outputs, session_id return outputs, 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
@ -246,9 +234,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():
@ -263,9 +249,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.
@ -286,9 +270,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
@ -300,8 +282,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

@ -23,9 +23,7 @@ async def process_graph(
if build_result is None: if build_result is None:
# Raise user facing error # Raise user facing error
raise ValueError( raise ValueError("There was an error loading the langchain_object. Please, check all the nodes and try again.")
"There was an error loading the langchain_object. Please, check all the nodes and try again."
)
# Generate result and thought # Generate result and thought
try: try:
@ -52,7 +50,5 @@ async def process_graph(
raise e raise e
async def run_build_result( async def run_build_result(build_result: Any, chat_inputs: ChatMessage, client_id: str, session_id: str):
build_result: Any, chat_inputs: ChatMessage, client_id: str, session_id: str
):
return build_result(inputs=chat_inputs.message) return build_result(inputs=chat_inputs.message)

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
@ -53,12 +51,8 @@ class MessageModel(BaseModel):
@classmethod @classmethod
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 not record.data or ( if not record.data or ("sender" not in record.data and "sender_name" not in record.data):
"sender" not in record.data and "sender_name" not in record.data raise ValueError("The record does not have the required fields 'sender' and 'sender_name' in the data.")
):
raise ValueError(
"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"],

View file

@ -44,9 +44,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

@ -23,10 +23,7 @@ def get_table_schema_as_dict(conn: duckdb.DuckDBPyConnection, table_name: str) -
def model_to_sql_column_definitions(model: Type[BaseModel]) -> dict: def model_to_sql_column_definitions(model: Type[BaseModel]) -> dict:
columns = {} columns = {}
for field_name, field_type in model.model_fields.items(): for field_name, field_type in model.model_fields.items():
if ( if hasattr(field_type.annotation, "__args__") and field_type.annotation is not None:
hasattr(field_type.annotation, "__args__")
and field_type.annotation is not None
):
field_args = field_type.annotation.__args__ field_args = field_type.annotation.__args__
else: else:
field_args = [] field_args = []
@ -49,9 +46,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:
@ -72,12 +67,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

@ -32,9 +32,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

@ -96,9 +96,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

@ -25,9 +25,7 @@ class S3StorageService(StorageService):
:raises Exception: If an error occurs during file saving. :raises Exception: If an error occurs during file saving.
""" """
try: try:
self.s3_client.put_object( self.s3_client.put_object(Bucket=self.bucket, Key=f"{folder}/{file_name}", Body=data)
Bucket=self.bucket, Key=f"{folder}/{file_name}", Body=data
)
logger.info(f"File {file_name} saved successfully in folder {folder}.") logger.info(f"File {file_name} saved successfully in folder {folder}.")
except NoCredentialsError: except NoCredentialsError:
logger.error("Credentials not available for AWS S3.") logger.error("Credentials not available for AWS S3.")
@ -46,12 +44,8 @@ class S3StorageService(StorageService):
:raises Exception: If an error occurs during file retrieval. :raises Exception: If an error occurs during file retrieval.
""" """
try: try:
response = self.s3_client.get_object( response = self.s3_client.get_object(Bucket=self.bucket, Key=f"{folder}/{file_name}")
Bucket=self.bucket, Key=f"{folder}/{file_name}" logger.info(f"File {file_name} retrieved successfully from folder {folder}.")
)
logger.info(
f"File {file_name} retrieved successfully from folder {folder}."
)
return response["Body"].read() return response["Body"].read()
except ClientError as e: except ClientError as e:
logger.error(f"Error retrieving file {file_name} from folder {folder}: {e}") logger.error(f"Error retrieving file {file_name} from folder {folder}: {e}")
@ -67,11 +61,7 @@ class S3StorageService(StorageService):
""" """
try: try:
response = self.s3_client.list_objects_v2(Bucket=self.bucket, Prefix=folder) response = self.s3_client.list_objects_v2(Bucket=self.bucket, Prefix=folder)
files = [ files = [item["Key"] for item in response.get("Contents", []) if "/" not in item["Key"][len(folder) :]]
item["Key"]
for item in response.get("Contents", [])
if "/" not in item["Key"][len(folder) :]
]
logger.info(f"{len(files)} files listed in folder {folder}.") logger.info(f"{len(files)} files listed in folder {folder}.")
return files return files
except ClientError as e: except ClientError as e:
@ -87,9 +77,7 @@ class S3StorageService(StorageService):
:raises Exception: If an error occurs during file deletion. :raises Exception: If an error occurs during file deletion.
""" """
try: try:
self.s3_client.delete_object( self.s3_client.delete_object(Bucket=self.bucket, Key=f"{folder}/{file_name}")
Bucket=self.bucket, Key=f"{folder}/{file_name}"
)
logger.info(f"File {file_name} deleted successfully from folder {folder}.") logger.info(f"File {file_name} deleted successfully from folder {folder}.")
except ClientError as e: except ClientError as e:
logger.error(f"Error deleting file {file_name} from folder {folder}: {e}") logger.error(f"Error deleting file {file_name} from folder {folder}: {e}")

View file

@ -11,9 +11,7 @@ if TYPE_CHECKING:
class StorageService(Service): class StorageService(Service):
name = "storage_service" name = "storage_service"
def __init__( def __init__(self, session_service: "SessionService", settings_service: "SettingsService"):
self, session_service: "SessionService", settings_service: "SettingsService"
):
self.settings_service = settings_service self.settings_service = settings_service
self.session_service = session_service self.session_service = session_service
self.set_ready() self.set_ready()

View file

@ -68,9 +68,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 = False title_case: bool = False
@ -119,10 +117,6 @@ class TemplateField(BaseModel):
if not isinstance(value, list): if not isinstance(value, list):
raise ValueError("file_types must be a list") raise ValueError("file_types must be a list")
return [ return [
( (f".{file_type}" if isinstance(file_type, str) and not file_type.startswith(".") else file_type)
f".{file_type}"
if isinstance(file_type, str) and not file_type.startswith(".")
else file_type
)
for file_type in value for file_type in value
] ]

View file

@ -171,9 +171,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
@ -184,11 +182,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
@ -199,9 +193,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>"}"""
@ -209,9 +201,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,
@ -224,9 +214,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"
@ -266,10 +254,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

@ -216,9 +216,7 @@ class MidJourneyPromptChainNode(FrontendNode):
), ),
], ],
) )
description: str = ( description: str = "MidJourneyPromptChain is a chain you can use to generate new MidJourney prompts."
"MidJourneyPromptChain is a chain you can use to generate new MidJourney prompts."
)
base_classes: list[str] = [ base_classes: list[str] = [
"LLMChain", "LLMChain",
"BaseCustomChain", "BaseCustomChain",

View file

@ -25,10 +25,7 @@ def patching(record):
def configure(log_level: Optional[str] = None, log_file: Optional[Path] = None): def configure(log_level: Optional[str] = None, log_file: Optional[Path] = None):
if ( if os.getenv("LANGFLOW_LOG_LEVEL", "").upper() in VALID_LOG_LEVELS and log_level is None:
os.getenv("LANGFLOW_LOG_LEVEL", "").upper() in VALID_LOG_LEVELS
and log_level is None
):
log_level = os.getenv("LANGFLOW_LOG_LEVEL") log_level = os.getenv("LANGFLOW_LOG_LEVEL")
if log_level is None: if log_level is None:
log_level = "INFO" log_level = "INFO"

View file

@ -45,9 +45,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:
@ -91,23 +89,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:
@ -134,23 +124,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]
@ -212,13 +194,9 @@ 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)
@ -239,11 +217,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
@ -256,9 +230,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
@ -302,9 +274,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

@ -24,9 +24,7 @@ def build_vertex(self, vertex: "Vertex") -> "Vertex":
async_to_sync(vertex.build)() async_to_sync(vertex.build)()
return vertex return vertex
except SoftTimeLimitExceeded as e: except SoftTimeLimitExceeded as e:
raise self.retry( raise self.retry(exc=SoftTimeLimitExceeded("Task took too long"), countdown=2) from e
exc=SoftTimeLimitExceeded("Task took too long"), countdown=2
) from e
@celery_app.task(acks_late=True) @celery_app.task(acks_late=True)

View file

@ -29,10 +29,7 @@ def poll_task_status(client, headers, href, max_attempts=20, sleep_time=1):
href, href,
headers=headers, headers=headers,
) )
if ( if task_status_response.status_code == 200 and task_status_response.json()["status"] == "SUCCESS":
task_status_response.status_code == 200
and task_status_response.json()["status"] == "SUCCESS"
):
return task_status_response.json() return task_status_response.json()
time.sleep(sleep_time) time.sleep(sleep_time)
return None # Return None if task did not complete in time return None # Return None if task did not complete in time
@ -126,11 +123,7 @@ def created_api_key(active_user):
) )
db_manager = get_db_service() db_manager = get_db_service()
with session_getter(db_manager) as session: with session_getter(db_manager) as session:
if ( if existing_api_key := session.query(ApiKey).filter(ApiKey.api_key == api_key.api_key).first():
existing_api_key := session.query(ApiKey)
.filter(ApiKey.api_key == api_key.api_key)
.first()
):
return existing_api_key return existing_api_key
session.add(api_key) session.add(api_key)
session.commit() session.commit()
@ -296,11 +289,7 @@ def test_get_all(client: TestClient, logged_in_headers):
dir_reader = DirectoryReader(settings.COMPONENTS_PATH[0]) dir_reader = DirectoryReader(settings.COMPONENTS_PATH[0])
files = dir_reader.get_files() files = dir_reader.get_files()
# json_response is a dict of dicts # json_response is a dict of dicts
all_names = [ all_names = [component_name for _, components in response.json().items() for component_name in components]
component_name
for _, components in response.json().items()
for component_name in components
]
json_response = response.json() json_response = response.json()
# We need to test the custom nodes # We need to test the custom nodes
assert len(all_names) > len(files) assert len(all_names) > len(files)
@ -425,19 +414,13 @@ def test_various_prompts(client, prompt, expected_input_variables):
def test_get_vertices_flow_not_found(client, logged_in_headers): def test_get_vertices_flow_not_found(client, logged_in_headers):
response = client.get( response = client.get("/api/v1/build/nonexistent_id/vertices", headers=logged_in_headers)
"/api/v1/build/nonexistent_id/vertices", headers=logged_in_headers assert response.status_code == 500 # Or whatever status code you've set for invalid ID
)
assert (
response.status_code == 500
) # Or whatever status code you've set for invalid ID
def test_get_vertices(client, added_flow_with_prompt_and_history, logged_in_headers): def test_get_vertices(client, added_flow_with_prompt_and_history, logged_in_headers):
flow_id = added_flow_with_prompt_and_history["id"] flow_id = added_flow_with_prompt_and_history["id"]
response = client.get( response = client.get(f"/api/v1/build/{flow_id}/vertices", headers=logged_in_headers)
f"/api/v1/build/{flow_id}/vertices", headers=logged_in_headers
)
assert response.status_code == 200 assert response.status_code == 200
assert "ids" in response.json() assert "ids" in response.json()
# The response should contain the list in this order # The response should contain the list in this order
@ -453,19 +436,13 @@ def test_get_vertices(client, added_flow_with_prompt_and_history, logged_in_head
def test_build_vertex_invalid_flow_id(client, logged_in_headers): def test_build_vertex_invalid_flow_id(client, logged_in_headers):
response = client.post( response = client.post("/api/v1/build/nonexistent_id/vertices/vertex_id", headers=logged_in_headers)
"/api/v1/build/nonexistent_id/vertices/vertex_id", headers=logged_in_headers
)
assert response.status_code == 500 assert response.status_code == 500
def test_build_vertex_invalid_vertex_id( def test_build_vertex_invalid_vertex_id(client, added_flow_with_prompt_and_history, logged_in_headers):
client, added_flow_with_prompt_and_history, logged_in_headers
):
flow_id = added_flow_with_prompt_and_history["id"] flow_id = added_flow_with_prompt_and_history["id"]
response = client.post( response = client.post(f"/api/v1/build/{flow_id}/vertices/invalid_vertex_id", headers=logged_in_headers)
f"/api/v1/build/{flow_id}/vertices/invalid_vertex_id", headers=logged_in_headers
)
assert response.status_code == 500 assert response.status_code == 500