feat: Add ruff rules SIM (#3979)

Add ruff rules SIM
This commit is contained in:
Christophe Bornet 2024-10-02 14:45:41 +02:00 • committed by GitHub
commit 942c8dca36
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
62 changed files with 228 additions and 374 deletions

View file

@ -457,11 +457,10 @@ def migration(
""" """
Run or test migrations. Run or test migrations.
""" """
if fix: if fix and not typer.confirm(
if not typer.confirm( "This will delete all data necessary to fix migrations. Are you sure you want to continue?"
"This will delete all data necessary to fix migrations. Are you sure you want to continue?" ):
): raise typer.Abort
raise typer.Abort
initialize_services(fix_migration=fix) initialize_services(fix_migration=fix)
db_service = get_db_service() db_service = get_db_service()

View file

@ -90,10 +90,7 @@ async def logs(
status_code=HTTPStatus.BAD_REQUEST, status_code=HTTPStatus.BAD_REQUEST,
detail="Timestamp is required when requesting logs after the timestamp", detail="Timestamp is required when requesting logs after the timestamp",
) )
if lines_before <= 0: content = log_buffer.get_last_n(10) if lines_before <= 0 else log_buffer.get_last_n(lines_before)
content = log_buffer.get_last_n(10)
else:
content = log_buffer.get_last_n(lines_before)
else: else:
if lines_before > 0: if lines_before > 0:
content = log_buffer.get_before_timestamp(timestamp=timestamp, lines=lines_before) content = log_buffer.get_before_timestamp(timestamp=timestamp, lines=lines_before)

View file

@ -344,20 +344,17 @@ async def build_flow(
raise ValueError(msg) from exc raise ValueError(msg) from exc
event_manager.on_end_vertex(data={"build_data": build_data}) event_manager.on_end_vertex(data={"build_data": build_data})
await client_consumed_queue.get() await client_consumed_queue.get()
if vertex_build_response.valid: if vertex_build_response.valid and vertex_build_response.next_vertices_ids:
if vertex_build_response.next_vertices_ids: tasks = []
tasks = [] for next_vertex_id in vertex_build_response.next_vertices_ids:
for next_vertex_id in vertex_build_response.next_vertices_ids: task = asyncio.create_task(build_vertices(next_vertex_id, graph, client_consumed_queue, event_manager))
task = asyncio.create_task( tasks.append(task)
build_vertices(next_vertex_id, graph, client_consumed_queue, event_manager) try:
) await asyncio.gather(*tasks)
tasks.append(task) except asyncio.CancelledError:
try: for task in tasks:
await asyncio.gather(*tasks) task.cancel()
except asyncio.CancelledError: return
for task in tasks:
task.cancel()
return
async def event_generator(event_manager: EventManager, client_consumed_queue: asyncio.Queue) -> None: async def event_generator(event_manager: EventManager, client_consumed_queue: asyncio.Queue) -> None:
if not data: if not data:

View file

@ -95,13 +95,12 @@ def validate_input_and_tweaks(input_request: SimplifiedAPIRequest):
if has_input_value and input_value_is_chat: if has_input_value and input_value_is_chat:
msg = "If you pass an input_value to the chat input, you cannot pass a tweak with the same name." msg = "If you pass an input_value to the chat input, you cannot pass a tweak with the same name."
raise InvalidChatInputException(msg) raise InvalidChatInputException(msg)
elif "Text Input" in key or "TextInput" in key: elif ("Text Input" in key or "TextInput" in key) and isinstance(value, dict):
if isinstance(value, dict): has_input_value = value.get("input_value") is not None
has_input_value = value.get("input_value") is not None input_value_is_text = input_request.input_value is not None and input_request.input_type == "text"
input_value_is_text = input_request.input_value is not None and input_request.input_type == "text" if has_input_value and input_value_is_text:
if has_input_value and input_value_is_text: msg = "If you pass an input_value to the text input, you cannot pass a tweak with the same name."
msg = "If you pass an input_value to the text input, you cannot pass a tweak with the same name." raise InvalidChatInputException(msg)
raise InvalidChatInputException(msg)
async def simple_run_flow( async def simple_run_flow(

View file

@ -310,10 +310,7 @@ async def upload_file(
contents = await file.read() contents = await file.read()
data = orjson.loads(contents) data = orjson.loads(contents)
response_list = [] response_list = []
if "flows" in data: flow_list = FlowListCreate(**data) if "flows" in data else FlowListCreate(flows=[FlowCreate(**data)])
flow_list = FlowListCreate(**data)
else:
flow_list = FlowListCreate(flows=[FlowCreate(**data)])
# Now we set the user_id for all flows # Now we set the user_id for all flows
for flow in flow_list.flows: for flow in flow_list.flows:
flow.user_id = current_user.id flow.user_id = current_user.id

View file

@ -53,10 +53,7 @@ class BaseCrewComponent(Component):
self, self,
) -> Callable: ) -> Callable:
def task_callback(task_output: TaskOutput): def task_callback(task_output: TaskOutput):
if self._vertex: vertex_id = self._vertex.id if self._vertex else self.display_name or self.__class__.__name__
vertex_id = self._vertex.id
else:
vertex_id = self.display_name or self.__class__.__name__
self.log(task_output.model_dump(), name=f"Task (Agent: {task_output.agent}) - {vertex_id}") self.log(task_output.model_dump(), name=f"Task (Agent: {task_output.agent}) - {vertex_id}")
return task_callback return task_callback

View file

@ -67,10 +67,7 @@ def build_data_from_result_data(result_data: ResultData, get_final_results_only:
if isinstance(result_data.results, dict): if isinstance(result_data.results, dict):
for name, result in result_data.results.items(): for name, result in result_data.results.items():
dataobj: Data | Message | None = None dataobj: Data | Message | None = None
if isinstance(result, Message): dataobj = result if isinstance(result, Message) else Data(data=result, text_key=name)
dataobj = result
else:
dataobj = Data(data=result, text_key=name)
data.append(dataobj) data.append(dataobj)
else: else:

View file

@ -25,12 +25,16 @@ class ChatComponent(Component):
msg = "Only one message can be stored at a time." msg = "Only one message can be stored at a time."
raise ValueError(msg) raise ValueError(msg)
stored_message = messages[0] stored_message = messages[0]
if hasattr(self, "_event_manager") and self._event_manager and stored_message.id: if (
if not isinstance(message.text, str): hasattr(self, "_event_manager")
complete_message = self._stream_message(message, stored_message.id) and self._event_manager
message_table = update_message(message_id=stored_message.id, message={"text": complete_message}) and stored_message.id
stored_message = Message(**message_table.model_dump()) and not isinstance(message.text, str)
self.vertex._added_message = stored_message ):
complete_message = self._stream_message(message, stored_message.id)
message_table = update_message(message_id=stored_message.id, message={"text": complete_message})
stored_message = Message(**message_table.model_dump())
self.vertex._added_message = stored_message
self.status = stored_message self.status = stored_message
return stored_message return stored_message
@ -77,8 +81,6 @@ class ChatComponent(Component):
session_id: str | None = None, session_id: str | None = None,
return_message: bool | None = False, return_message: bool | None = False,
) -> Message: ) -> Message:
message: Message | None = None
if isinstance(input_value, Data): if isinstance(input_value, Data):
# Update the data of the record # Update the data of the record
message = Message.from_data(input_value) message = Message.from_data(input_value)
@ -86,10 +88,7 @@ class ChatComponent(Component):
message = Message( message = Message(
text=input_value, sender=sender, sender_name=sender_name, files=files, session_id=session_id text=input_value, sender=sender, sender_name=sender_name, files=files, session_id=session_id
) )
if not return_message: message_text = message.text if not return_message else message
message_text = message.text
else:
message_text = message # type: ignore
self.status = message_text self.status = message_text
if session_id and isinstance(message, Message) and isinstance(message.text, str): if session_id and isinstance(message, Message) and isinstance(message.text, str):

View file

@ -130,10 +130,7 @@ class SequentialTaskAgentComponent(Component):
# If there's a previous task, create a list of tasks # If there's a previous task, create a list of tasks
if self.previous_task: if self.previous_task:
if isinstance(self.previous_task, list): tasks = self.previous_task + [task] if isinstance(self.previous_task, list) else [self.previous_task, task]
tasks = self.previous_task + [task]
else:
tasks = [self.previous_task, task]
else: else:
tasks = [task] tasks = [task]

View file

@ -28,10 +28,7 @@ class SQLGeneratorComponent(LCChainComponent):
outputs = [Output(display_name="Text", name="text", method="invoke_chain")] outputs = [Output(display_name="Text", name="text", method="invoke_chain")]
def invoke_chain(self) -> Message: def invoke_chain(self) -> Message:
if self.prompt: prompt_template = PromptTemplate.from_template(template=self.prompt) if self.prompt else None
prompt_template = PromptTemplate.from_template(template=self.prompt)
else:
prompt_template = None
if self.top_k < 1: if self.top_k < 1:
msg = "Top K must be greater than 0." msg = "Top K must be greater than 0."

View file

@ -96,10 +96,7 @@ class GmailLoaderComponent(Component):
msg = "From email not found." msg = "From email not found."
raise ValueError(msg) raise ValueError(msg)
if "parts" in msg["payload"]: parts = msg["payload"]["parts"] if "parts" in msg["payload"] else [msg["payload"]]
parts = msg["payload"]["parts"]
else:
parts = [msg["payload"]]
for part in parts: for part in parts:
if part["mimeType"] == "text/plain": if part["mimeType"] == "text/plain":

View file

@ -96,17 +96,13 @@ class GoogleDriveSearchComponent(Component):
""" """
Generates the appropriate Google Drive URL for a file based on its MIME type. Generates the appropriate Google Drive URL for a file based on its MIME type.
""" """
if mime_type == "application/vnd.google-apps.document": return {
return f"https://docs.google.com/document/d/{file_id}/edit" "application/vnd.google-apps.document": f"https://docs.google.com/document/d/{file_id}/edit",
if mime_type == "application/vnd.google-apps.spreadsheet": "application/vnd.google-apps.spreadsheet": f"https://docs.google.com/spreadsheets/d/{file_id}/edit",
return f"https://docs.google.com/spreadsheets/d/{file_id}/edit" "application/vnd.google-apps.presentation": f"https://docs.google.com/presentation/d/{file_id}/edit",
if mime_type == "application/vnd.google-apps.presentation": "application/vnd.google-apps.drawing": f"https://docs.google.com/drawings/d/{file_id}/edit",
return f"https://docs.google.com/presentation/d/{file_id}/edit" "application/pdf": f"https://drive.google.com/file/d/{file_id}/view?usp=drivesdk",
if mime_type == "application/vnd.google-apps.drawing": }.get(mime_type, f"https://drive.google.com/file/d/{file_id}/view?usp=drivesdk")
return f"https://docs.google.com/drawings/d/{file_id}/edit"
if mime_type == "application/pdf":
return f"https://drive.google.com/file/d/{file_id}/view?usp=drivesdk"
return f"https://drive.google.com/file/d/{file_id}/view?usp=drivesdk"
def search_files(self) -> dict: def search_files(self) -> dict:
# Load the token information from the JSON string # Load the token information from the JSON string

View file

@ -39,10 +39,7 @@ class TextEmbedderComponent(Component):
embeddings = embedding_model.embed_documents([text_content]) embeddings = embedding_model.embed_documents([text_content])
# Assuming the embedding model returns a list of embeddings, we take the first one # Assuming the embedding model returns a list of embeddings, we take the first one
if embeddings: embedding_vector = embeddings[0] if embeddings else []
embedding_vector = embeddings[0]
else:
embedding_vector = []
# Create a Data object to encapsulate the results # Create a Data object to encapsulate the results
result_data = Data(data={"text": text_content, "embeddings": embedding_vector}) result_data = Data(data={"text": text_content, "embeddings": embedding_vector})

View file

@ -21,25 +21,24 @@ class AIMLEmbeddingsImpl(BaseModel, Embeddings):
"Authorization": f"Bearer {self.api_key.get_secret_value()}", "Authorization": f"Bearer {self.api_key.get_secret_value()}",
} }
with httpx.Client() as client: with httpx.Client() as client, concurrent.futures.ThreadPoolExecutor() as executor:
with concurrent.futures.ThreadPoolExecutor() as executor: futures = []
futures = [] for i, text in enumerate(texts):
for i, text in enumerate(texts): futures.append((i, executor.submit(self._embed_text, client, headers, text)))
futures.append((i, executor.submit(self._embed_text, client, headers, text)))
for index, future in futures: for index, future in futures:
try: try:
result_data = future.result() result_data = future.result()
assert len(result_data["data"]) == 1, "Expected one embedding" assert len(result_data["data"]) == 1, "Expected one embedding"
embeddings[index] = result_data["data"][0]["embedding"] embeddings[index] = result_data["data"][0]["embedding"]
except ( except (
httpx.HTTPStatusError, httpx.HTTPStatusError,
httpx.RequestError, httpx.RequestError,
json.JSONDecodeError, json.JSONDecodeError,
KeyError, KeyError,
) as e: ) as e:
logger.error(f"Error occurred: {e}") logger.error(f"Error occurred: {e}")
raise raise
return embeddings # type: ignore return embeddings # type: ignore

View file

@ -61,15 +61,9 @@ class FirecrawlCrawlApi(CustomComponent):
"Could not import firecrawl integration package. " "Please install it with `pip install firecrawl-py`." "Could not import firecrawl integration package. " "Please install it with `pip install firecrawl-py`."
) )
raise ImportError(msg) raise ImportError(msg)
if crawlerOptions: crawler_options_dict = crawlerOptions.__dict__["data"]["text"] if crawlerOptions else {}
crawler_options_dict = crawlerOptions.__dict__["data"]["text"]
else:
crawler_options_dict = {}
if pageOptions: page_options_dict = pageOptions.__dict__["data"]["text"] if pageOptions else {}
page_options_dict = pageOptions.__dict__["data"]["text"]
else:
page_options_dict = {}
if not idempotency_key: if not idempotency_key:
idempotency_key = str(uuid.uuid4()) idempotency_key = str(uuid.uuid4())

View file

@ -54,15 +54,9 @@ class FirecrawlScrapeApi(CustomComponent):
"Could not import firecrawl integration package. " "Please install it with `pip install firecrawl-py`." "Could not import firecrawl integration package. " "Please install it with `pip install firecrawl-py`."
) )
raise ImportError(msg) raise ImportError(msg)
if extractorOptions: extractor_options_dict = extractorOptions.__dict__["data"]["text"] if extractorOptions else {}
extractor_options_dict = extractorOptions.__dict__["data"]["text"]
else:
extractor_options_dict = {}
if pageOptions: page_options_dict = pageOptions.__dict__["data"]["text"] if pageOptions else {}
page_options_dict = pageOptions.__dict__["data"]["text"]
else:
page_options_dict = {}
app = FirecrawlApp(api_key=api_key) app = FirecrawlApp(api_key=api_key)
results = app.scrape_url( results = app.scrape_url(

View file

@ -79,10 +79,7 @@ class AIMLModelComponent(LCModelComponent):
aiml_api_base = self.aiml_api_base or "https://api.aimlapi.com" aiml_api_base = self.aiml_api_base or "https://api.aimlapi.com"
seed = self.seed seed = self.seed
if isinstance(aiml_api_key, SecretStr): openai_api_key = aiml_api_key.get_secret_value() if isinstance(aiml_api_key, SecretStr) else aiml_api_key
openai_api_key = aiml_api_key.get_secret_value()
else:
openai_api_key = aiml_api_key
return ChatOpenAI( return ChatOpenAI(
model=model_name, model=model_name,

View file

@ -36,10 +36,7 @@ class CohereComponent(LCModelComponent):
cohere_api_key = self.cohere_api_key cohere_api_key = self.cohere_api_key
temperature = self.temperature temperature = self.temperature
if cohere_api_key: api_key = SecretStr(cohere_api_key) if cohere_api_key else None
api_key = SecretStr(cohere_api_key)
else:
api_key = None
return ChatCohere( return ChatCohere(
temperature=temperature or 0.75, temperature=temperature or 0.75,

View file

@ -78,10 +78,7 @@ class MistralAIModelComponent(LCModelComponent):
random_seed = self.random_seed random_seed = self.random_seed
safe_mode = self.safe_mode safe_mode = self.safe_mode
if mistral_api_key: api_key = SecretStr(mistral_api_key) if mistral_api_key else None
api_key = SecretStr(mistral_api_key)
else:
api_key = None
return ChatMistralAI( return ChatMistralAI(
max_tokens=max_tokens or None, max_tokens=max_tokens or None,

View file

@ -102,10 +102,7 @@ class OpenAIModelComponent(LCModelComponent):
json_mode = bool(output_schema_dict) or self.json_mode json_mode = bool(output_schema_dict) or self.json_mode
seed = self.seed seed = self.seed
if openai_api_key: api_key = SecretStr(openai_api_key) if openai_api_key else None
api_key = SecretStr(openai_api_key)
else:
api_key = None
output = ChatOpenAI( output = ChatOpenAI(
max_tokens=max_tokens or None, max_tokens=max_tokens or None,
model_kwargs=model_kwargs, model_kwargs=model_kwargs,

View file

@ -57,7 +57,7 @@ class SubFlowComponent(Component):
for vertex in inputs_vertex: for vertex in inputs_vertex:
new_vertex_inputs = [] new_vertex_inputs = []
field_template = vertex.data["node"]["template"] field_template = vertex.data["node"]["template"]
for inp in field_template.keys(): for inp in field_template:
if inp not in ["code", "_type"]: if inp not in ["code", "_type"]:
field_template[inp]["display_name"] = ( field_template[inp]["display_name"] = (
vertex.display_name + " - " + field_template[inp]["display_name"] vertex.display_name + " - " + field_template[inp]["display_name"]
@ -84,10 +84,10 @@ class SubFlowComponent(Component):
async def generate_results(self) -> list[Data]: async def generate_results(self) -> list[Data]:
tweaks: dict = {} tweaks: dict = {}
for field in self._attributes.keys(): for field in self._attributes:
if field != "flow_name": if field != "flow_name":
[node, name] = field.split("|") [node, name] = field.split("|")
if node not in tweaks.keys(): if node not in tweaks:
tweaks[node] = {} tweaks[node] = {}
tweaks[node][name] = self._attributes[field] tweaks[node][name] = self._attributes[field]
flow_name = self._attributes.get("flow_name") flow_name = self._attributes.get("flow_name")

View file

@ -43,10 +43,7 @@ class CharacterTextSplitterComponent(LCTextSplitterComponent):
return self.data_input return self.data_input
def build_text_splitter(self) -> TextSplitter: def build_text_splitter(self) -> TextSplitter:
if self.separator: separator = unescape_string(self.separator) if self.separator else "\n\n"
separator = unescape_string(self.separator)
else:
separator = "\n\n"
return CharacterTextSplitter( return CharacterTextSplitter(
chunk_overlap=self.chunk_overlap, chunk_overlap=self.chunk_overlap,
chunk_size=self.chunk_size, chunk_size=self.chunk_size,

View file

@ -51,10 +51,7 @@ class NaturalLanguageTextSplitterComponent(LCTextSplitterComponent):
return self.data_input return self.data_input
def build_text_splitter(self) -> TextSplitter: def build_text_splitter(self) -> TextSplitter:
if self.separator: separator = unescape_string(self.separator) if self.separator else "\n\n"
separator = unescape_string(self.separator)
else:
separator = "\n\n"
return NLTKTextSplitter( return NLTKTextSplitter(
language=self.language.lower() if self.language else "english", language=self.language.lower() if self.language else "english",
separator=separator, separator=separator,

View file

@ -188,9 +188,7 @@ class PythonCodeStructuredTool(LCToolComponent):
schema_annotation = Any schema_annotation = Any
schema_fields[field_name] = ( schema_fields[field_name] = (
schema_annotation, schema_annotation,
Field( Field(default=func_arg.get("default", Undefined), description=field_description),
default=func_arg["default"] if "default" in func_arg else Undefined, description=field_description
),
) )
if "temp_annotation_type" in _globals: if "temp_annotation_type" in _globals:

View file

@ -174,10 +174,7 @@ class CassandraVectorStoreComponent(LCVectorStoreComponent):
else: else:
documents.append(_input) documents.append(_input)
if self.enable_body_search: body_index_options = [("index_analyzer", "STANDARD")] if self.enable_body_search else None
body_index_options = [("index_analyzer", "STANDARD")]
else:
body_index_options = None
if self.setup_mode == "Off": if self.setup_mode == "Off":
setup_mode = SetupMode.OFF setup_mode = SetupMode.OFF

View file

@ -161,10 +161,7 @@ class CassandraGraphVectorStoreComponent(LCVectorStoreComponent):
else: else:
documents.append(_input) documents.append(_input)
if self.setup_mode == "Off": setup_mode = SetupMode.OFF if self.setup_mode == "Off" else SetupMode.SYNC
setup_mode = SetupMode.OFF
else:
setup_mode = SetupMode.SYNC
if documents: if documents:
logger.debug(f"Adding {len(documents)} documents to the Vector Store.") logger.debug(f"Adding {len(documents)} documents to the Vector Store.")

View file

@ -125,10 +125,7 @@ class ChromaVectorStoreComponent(LCVectorStoreComponent):
client = Client(settings=chroma_settings) client = Client(settings=chroma_settings)
# Check persist_directory and expand it if it is a relative path # Check persist_directory and expand it if it is a relative path
if self.persist_directory is not None: persist_directory = self.resolve_path(self.persist_directory) if self.persist_directory is not None else None
persist_directory = self.resolve_path(self.persist_directory)
else:
persist_directory = None
chroma = Chroma( chroma = Chroma(
persist_directory=persist_directory, persist_directory=persist_directory,

View file

@ -1,4 +1,5 @@
import ast import ast
import contextlib
import inspect import inspect
import traceback import traceback
from typing import Any from typing import Any
@ -171,11 +172,9 @@ class CodeParser:
return_type_str = ast.unparse(node.returns) return_type_str = ast.unparse(node.returns)
eval_env = self.construct_eval_env(return_type_str, tuple(self.data["imports"])) eval_env = self.construct_eval_env(return_type_str, tuple(self.data["imports"]))
try: # Handle cases where the type is not found in the constructed environment
with contextlib.suppress(NameError):
return_type = eval(return_type_str, eval_env) return_type = eval(return_type_str, eval_env)
except NameError:
# Handle cases where the type is not found in the constructed environment
pass
func = CallableCodeDetails( func = CallableCodeDetails(
name=node.name, name=node.name,

View file

@ -81,7 +81,7 @@ class BaseComponent:
template_config[attribute] = func(value=value) template_config[attribute] = func(value=value)
for key in template_config.copy(): for key in template_config.copy():
if key not in ATTR_FUNC_MAPPING.keys(): if key not in ATTR_FUNC_MAPPING:
template_config.pop(key, None) template_config.pop(key, None)
return template_config return template_config

View file

@ -226,7 +226,7 @@ def add_extra_fields(frontend_node, field_config, function_args):
field_required, field_required,
config, config,
) )
if "kwargs" in function_args_names and not all(key in function_args_names for key in field_config.keys()): if "kwargs" in function_args_names and not all(key in function_args_names for key in field_config):
for field_name, field_config in _field_config.copy().items(): for field_name, field_config in _field_config.copy().items():
if "name" not in field_config or field_name == "code": if "name" not in field_config or field_name == "code":
continue continue
@ -503,35 +503,35 @@ def update_field_dict(
call: bool = False, call: bool = False,
): ):
"""Update the field dictionary by calling options() or value() if they are callable""" """Update the field dictionary by calling options() or value() if they are callable"""
if ("real_time_refresh" in field_dict or "refresh_button" in field_dict) and any( if (
( ("real_time_refresh" in field_dict or "refresh_button" in field_dict)
field_dict.get("real_time_refresh", False), and any(
field_dict.get("refresh_button", False), (
field_dict.get("real_time_refresh", False),
field_dict.get("refresh_button", False),
)
) )
and call
): ):
if call: try:
try: dd_build_config = dotdict(build_config)
dd_build_config = dotdict(build_config) custom_component_instance.update_build_config(
custom_component_instance.update_build_config( build_config=dd_build_config,
build_config=dd_build_config, field_value=update_field,
field_value=update_field, field_name=update_field_value,
field_name=update_field_value, )
) build_config = dd_build_config
build_config = dd_build_config except Exception as exc:
except Exception as exc: logger.error(f"Error while running update_build_config: {str(exc)}")
logger.error(f"Error while running update_build_config: {str(exc)}") msg = f"Error while running update_build_config: {str(exc)}"
msg = f"Error while running update_build_config: {str(exc)}" raise UpdateBuildConfigError(msg) from exc
raise UpdateBuildConfigError(msg) from exc
return build_config return build_config
def sanitize_field_config(field_config: dict | Input): def sanitize_field_config(field_config: dict | Input):
# If any of the already existing keys are in field_config, remove them # If any of the already existing keys are in field_config, remove them
if isinstance(field_config, Input): field_dict = field_config.to_dict() if isinstance(field_config, Input) else field_config
field_dict = field_config.to_dict()
else:
field_dict = field_config
for key in [ for key in [
"name", "name",
"field_type", "field_type",

View file

@ -244,10 +244,10 @@ class CycleEdge(Edge):
await self.honor(source, target) await self.honor(source, target)
# If the target vertex is a power component we log messages # If the target vertex is a power component we log messages
if target.vertex_type == "ChatOutput" and ( if (
isinstance(target.params.get(INPUT_FIELD_NAME), str) target.vertex_type == "ChatOutput"
or isinstance(target.params.get(INPUT_FIELD_NAME), dict) and isinstance(target.params.get(INPUT_FIELD_NAME), str | dict)
and target.params.get("message") == ""
): ):
if target.params.get("message") == "": return self.result
return self.result
return self.result return self.result

View file

@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import contextlib
import copy import copy
import json import json
import uuid import uuid
@ -1061,10 +1062,7 @@ class Graph:
same_length = len(vertex.edges) == len(other_vertex.edges) same_length = len(vertex.edges) == len(other_vertex.edges)
if not same_length: if not same_length:
return False return False
for edge in vertex.edges: return all(edge in other_vertex.edges for edge in vertex.edges)
if edge not in other_vertex.edges:
return False
return True
def update(self, other: Graph) -> Graph: def update(self, other: Graph) -> Graph:
# Existing vertices in self graph # Existing vertices in self graph
@ -1080,10 +1078,8 @@ class Graph:
# Remove vertices that are not in the other graph # Remove vertices that are not in the other graph
for vertex_id in removed_vertex_ids: for vertex_id in removed_vertex_ids:
try: with contextlib.suppress(ValueError):
self.remove_vertex(vertex_id) self.remove_vertex(vertex_id)
except ValueError:
pass
# The order here matters because adding the vertex is required # The order here matters because adding the vertex is required
# if any of them have edges that point to any of the new vertices # if any of them have edges that point to any of the new vertices

View file

@ -55,9 +55,7 @@ class RunnableVerticesManager:
return False return False
if vertex_id not in self.vertices_to_run: if vertex_id not in self.vertices_to_run:
return False return False
if not self.are_all_predecessors_fulfilled(vertex_id): return self.are_all_predecessors_fulfilled(vertex_id)
return False
return True
def are_all_predecessors_fulfilled(self, vertex_id: str) -> bool: def are_all_predecessors_fulfilled(self, vertex_id: str) -> bool:
return not any(self.run_predecessors.get(vertex_id, [])) return not any(self.run_predecessors.get(vertex_id, []))

View file

@ -354,12 +354,7 @@ def has_cycle(vertex_ids: list[str], edges: list[tuple[str, str]]) -> bool:
visited: set[str] = set() visited: set[str] = set()
rec_stack: set[str] = set() rec_stack: set[str] = set()
for vertex in vertex_ids: return any(vertex not in visited and dfs(vertex, visited, rec_stack) for vertex in vertex_ids)
if vertex not in visited:
if dfs(vertex, visited, rec_stack):
return True
return False
def find_cycle_edge(entry_point: str, edges: list[tuple[str, str]]) -> tuple[str, str]: def find_cycle_edge(entry_point: str, edges: list[tuple[str, str]]) -> tuple[str, str]:

View file

@ -103,11 +103,10 @@ def get_artifact_type(value, build_result) -> str:
case Message(): case Message():
result = ArtifactType.MESSAGE result = ArtifactType.MESSAGE
if result == ArtifactType.UNKNOWN: if result == ArtifactType.UNKNOWN and (
if isinstance(build_result, Generator): isinstance(build_result, Generator) or isinstance(value, Message) and isinstance(value.text, Generator)
result = ArtifactType.STREAM ):
elif isinstance(value, Message) and isinstance(value.text, Generator): result = ArtifactType.STREAM
result = ArtifactType.STREAM
return result.value return result.value

View file

@ -287,7 +287,7 @@ class Vertex:
if not param_dict or len(param_dict) != 1: if not param_dict or len(param_dict) != 1:
params[param_key] = self.graph.get_vertex(edge.source_id) params[param_key] = self.graph.get_vertex(edge.source_id)
else: else:
params[param_key] = {key: self.graph.get_vertex(edge.source_id) for key in param_dict.keys()} params[param_key] = {key: self.graph.get_vertex(edge.source_id) for key in param_dict}
else: else:
params[param_key] = self.graph.get_vertex(edge.source_id) params[param_key] = self.graph.get_vertex(edge.source_id)
@ -415,8 +415,6 @@ class Vertex:
elif val is not None and val != "": elif val is not None and val != "":
params[field_name] = val params[field_name] = val
elif val is not None and val != "":
params[field_name] = val
if field.get("load_from_db"): if field.get("load_from_db"):
load_from_db_fields.append(field_name) load_from_db_fields.append(field_name)
@ -534,10 +532,7 @@ class Vertex:
# to the frontend # to the frontend
self.set_artifacts() self.set_artifacts()
artifacts = self.artifacts_raw artifacts = self.artifacts_raw
if isinstance(artifacts, dict): messages = self.extract_messages_from_artifacts(artifacts) if isinstance(artifacts, dict) else []
messages = self.extract_messages_from_artifacts(artifacts)
else:
messages = []
result_dict = ResultData( result_dict = ResultData(
results=result_dict, results=result_dict,
artifacts=artifacts, artifacts=artifacts,

View file

@ -1,4 +1,5 @@
import asyncio import asyncio
import contextlib
import json import json
from collections.abc import AsyncIterator, Generator, Iterator from collections.abc import AsyncIterator, Generator, Iterator
from typing import TYPE_CHECKING, Any, cast from typing import TYPE_CHECKING, Any, cast
@ -167,7 +168,7 @@ class ComponentVertex(Vertex):
message_dict = artifact if isinstance(artifact, dict) else artifact.model_dump() message_dict = artifact if isinstance(artifact, dict) else artifact.model_dump()
if not message_dict.get("text"): if not message_dict.get("text"):
continue continue
try: with contextlib.suppress(KeyError):
messages.append( messages.append(
ChatOutputResponse( ChatOutputResponse(
message=message_dict["text"], message=message_dict["text"],
@ -182,8 +183,6 @@ class ComponentVertex(Vertex):
type=self.artifacts_type[key], type=self.artifacts_type[key],
).model_dump(exclude_none=True) ).model_dump(exclude_none=True)
) )
except KeyError:
pass
return messages return messages
def _finalize_build(self): def _finalize_build(self):
@ -440,11 +439,12 @@ class InterfaceVertex(ComponentVertex):
for key, value in origin_vertex.results.items(): for key, value in origin_vertex.results.items():
if isinstance(value, AsyncIterator | Iterator): if isinstance(value, AsyncIterator | Iterator):
origin_vertex.results[key] = complete_message origin_vertex.results[key] = complete_message
if self._custom_component: if (
if hasattr(self._custom_component, "should_store_message") and hasattr( self._custom_component
self._custom_component, "store_message" and hasattr(self._custom_component, "should_store_message")
): and hasattr(self._custom_component, "store_message")
self._custom_component.store_message(message) ):
self._custom_component.store_message(message)
log_vertex_build( log_vertex_build(
flow_id=self.graph.flow_id, flow_id=self.graph.flow_id,
vertex_id=self.id, vertex_id=self.id,

View file

@ -90,17 +90,19 @@ def update_projects_components_with_latest_component_versions(project_data, all_
) )
else: else:
for attr in NODE_FORMAT_ATTRIBUTES: for attr in NODE_FORMAT_ATTRIBUTES:
if attr in latest_node: if (
attr in latest_node
# Check if it needs to be updated # Check if it needs to be updated
if latest_node[attr] != node_data.get(attr): and latest_node[attr] != node_data.get(attr)
node_changes_log[node_data["display_name"]].append( ):
{ node_changes_log[node_data["display_name"]].append(
"attr": attr, {
"old_value": node_data.get(attr), "attr": attr,
"new_value": latest_node[attr], "old_value": node_data.get(attr),
} "new_value": latest_node[attr],
) }
node_data[attr] = latest_node[attr] )
node_data[attr] = latest_node[attr]
for field_name, field_dict in latest_template.items(): for field_name, field_dict in latest_template.items():
if field_name not in node_data["template"]: if field_name not in node_data["template"]:
@ -109,17 +111,20 @@ def update_projects_components_with_latest_component_versions(project_data, all_
# The idea here is to update some attributes of the field # The idea here is to update some attributes of the field
to_check_attributes = FIELD_FORMAT_ATTRIBUTES to_check_attributes = FIELD_FORMAT_ATTRIBUTES
for attr in to_check_attributes: for attr in to_check_attributes:
if attr in field_dict and attr in node_data["template"].get(field_name): if (
attr in field_dict
and attr in node_data["template"].get(field_name)
# Check if it needs to be updated # Check if it needs to be updated
if field_dict[attr] != node_data["template"][field_name][attr]: and field_dict[attr] != node_data["template"][field_name][attr]
node_changes_log[node_data["display_name"]].append( ):
{ node_changes_log[node_data["display_name"]].append(
"attr": f"{field_name}.{attr}", {
"old_value": node_data["template"][field_name][attr], "attr": f"{field_name}.{attr}",
"new_value": field_dict[attr], "old_value": node_data["template"][field_name][attr],
} "new_value": field_dict[attr],
) }
node_data["template"][field_name][attr] = field_dict[attr] )
node_data["template"][field_name][attr] = field_dict[attr]
# Remove fields that are not in the latest template # Remove fields that are not in the latest template
if node_data.get("display_name") != "Prompt": if node_data.get("display_name") != "Prompt":
for field_name in list(node_data["template"].keys()): for field_name in list(node_data["template"].keys()):
@ -274,16 +279,17 @@ def update_edges_with_latest_component_versions(project_data):
source_handle["output_types"] = new_output_types source_handle["output_types"] = new_output_types
field_name = target_handle.get("fieldName") field_name = target_handle.get("fieldName")
if field_name in target_node_data.get("template"): if field_name in target_node_data.get("template") and target_handle["inputTypes"] != target_node_data.get(
if target_handle["inputTypes"] != target_node_data.get("template").get(field_name).get("input_types"): "template"
edge_changes_log[target_node_data["display_name"]].append( ).get(field_name).get("input_types"):
{ edge_changes_log[target_node_data["display_name"]].append(
"attr": "inputTypes", {
"old_value": target_handle["inputTypes"], "attr": "inputTypes",
"new_value": target_node_data.get("template").get(field_name).get("input_types"), "old_value": target_handle["inputTypes"],
} "new_value": target_node_data.get("template").get(field_name).get("input_types"),
) }
target_handle["inputTypes"] = target_node_data.get("template").get(field_name).get("input_types") )
target_handle["inputTypes"] = target_node_data.get("template").get(field_name).get("input_types")
escaped_source_handle = escape_json_dump(source_handle) escaped_source_handle = escape_json_dump(source_handle)
escaped_target_handle = escape_json_dump(target_handle) escaped_target_handle = escape_json_dump(target_handle)
try: try:
@ -390,10 +396,7 @@ def get_project_data(project):
updated_at_datetime = datetime.strptime(project_updated_at, "%Y-%m-%dT%H:%M:%S.%f") updated_at_datetime = datetime.strptime(project_updated_at, "%Y-%m-%dT%H:%M:%S.%f")
project_data = project.get("data") project_data = project.get("data")
project_icon = project.get("icon") project_icon = project.get("icon")
if project_icon and purely_emoji(project_icon): project_icon = demojize(project_icon) if project_icon and purely_emoji(project_icon) else ""
project_icon = demojize(project_icon)
else:
project_icon = ""
project_icon_bg_color = project.get("icon_bg_color") project_icon_bg_color = project.get("icon_bg_color")
return ( return (
project_name, project_name,

View file

@ -131,12 +131,7 @@ class StrInput(BaseInputMixin, ListableInputMixin, DatabaseLoadMixin, MetadataTr
ValueError: If the value is not of a valid type or if the input is missing a required key. ValueError: If the value is not of a valid type or if the input is missing a required key.
""" """
is_list = _info.data["is_list"] is_list = _info.data["is_list"]
value = None return [cls._validate_value(vv, _info) for vv in v] if is_list else cls._validate_value(v, _info)
if is_list:
value = [cls._validate_value(vv, _info) for vv in v]
else:
value = cls._validate_value(v, _info)
return value
class MessageInput(StrInput, InputTraceMixin): class MessageInput(StrInput, InputTraceMixin):

View file

@ -89,12 +89,11 @@ def convert_kwargs(params):
# Loop through items to avoid repeated lookups # Loop through items to avoid repeated lookups
items_to_remove = [] items_to_remove = []
for key, value in params.items(): for key, value in params.items():
if "kwargs" in key or "config" in key: if ("kwargs" in key or "config" in key) and isinstance(value, str):
if isinstance(value, str): try:
try: params[key] = orjson.loads(value)
params[key] = orjson.loads(value) except orjson.JSONDecodeError:
except orjson.JSONDecodeError: items_to_remove.append(key)
items_to_remove.append(key)
# Remove invalid keys outside the loop to avoid modifying dict during iteration # Remove invalid keys outside the loop to avoid modifying dict during iteration
for key in items_to_remove: for key in items_to_remove:

View file

@ -69,11 +69,9 @@ def extract_input_variables_from_prompt(prompt: str) -> list[str]:
if not match: if not match:
break break
# Extract the variable name from either the single or double brace match # Extract the variable name from either the single or double brace match.
if match.group(1): # Match found in double braces # If match found in double braces, re-add single braces for JSON strings.
variable_name = "{{" + match.group(1) + "}}" # Re-add single braces for JSON strings variable_name = "{{" + match.group(1) + "}}" if match.group(1) else match.group(2)
else: # Match found in single braces
variable_name = match.group(2)
if variable_name is not None: if variable_name is not None:
# This means there is a match # This means there is a match
# but there is nothing inside the braces # but there is nothing inside the braces

View file

@ -77,22 +77,18 @@ class SizedLogBuffer:
try: try:
with self._wlock: with self._wlock:
as_list = list(self.buffer) as_list = list(self.buffer)
i = 0
max_index = -1 max_index = -1
for ts, msg in as_list: for i, (ts, msg) in enumerate(as_list):
if ts >= timestamp: if ts >= timestamp:
max_index = i max_index = i
break break
i += 1
if max_index == -1: if max_index == -1:
return self.get_last_n(lines) return self.get_last_n(lines)
rc = {} rc = {}
i = 0
start_from = max(max_index - lines, 0) start_from = max(max_index - lines, 0)
for ts, msg in as_list: for i, (ts, msg) in enumerate(as_list):
if start_from <= i < max_index: if start_from <= i < max_index:
rc[ts] = msg rc[ts] = msg
i += 1
return rc return rc
finally: finally:
self._rsemaphore.release() self._rsemaphore.release()

View file

@ -46,10 +46,7 @@ def get_messages(
if flow_id: if flow_id:
stmt = stmt.where(MessageTable.flow_id == flow_id) stmt = stmt.where(MessageTable.flow_id == flow_id)
if order_by: if order_by:
if order == "DESC": col = getattr(MessageTable, order_by).desc() if order == "DESC" else getattr(MessageTable, order_by).asc()
col = getattr(MessageTable, order_by).desc()
else:
col = getattr(MessageTable, order_by).asc()
stmt = stmt.order_by(col) stmt = stmt.order_by(col)
if limit: if limit:
stmt = stmt.limit(limit) stmt = stmt.limit(limit)

View file

@ -29,10 +29,7 @@ async def run_graph_internal(
) -> tuple[list[RunOutputs], str]: ) -> tuple[list[RunOutputs], str]:
"""Run the graph and generate the result""" """Run the graph and generate the result"""
inputs = inputs or [] inputs = inputs or []
if session_id is None: session_id_str = flow_id if session_id is None else session_id
session_id_str = flow_id
else:
session_id_str = session_id
components = [] components = []
inputs_list = [] inputs_list = []
types = [] types = []
@ -168,11 +165,7 @@ def process_tweaks(
:return: The modified graph_data dictionary. :return: The modified graph_data dictionary.
:raises ValueError: If the input is not in the expected format. :raises ValueError: If the input is not in the expected format.
""" """
tweaks_dict = {} tweaks_dict = cast(dict[str, Any], tweaks.model_dump()) if not isinstance(tweaks, dict) else tweaks
if not isinstance(tweaks, dict):
tweaks_dict = cast(dict[str, Any], tweaks.model_dump())
else:
tweaks_dict = tweaks
if "stream" not in tweaks_dict: if "stream" not in tweaks_dict:
tweaks_dict |= {"stream": stream} tweaks_dict |= {"stream": stream}
nodes = validate_input(graph_data, cast(dict[str, str | dict[str, Any]], tweaks_dict)) nodes = validate_input(graph_data, cast(dict[str, str | dict[str, Any]], tweaks_dict))
@ -182,9 +175,7 @@ def process_tweaks(
all_nodes_tweaks = {} all_nodes_tweaks = {}
for key, value in tweaks_dict.items(): for key, value in tweaks_dict.items():
if isinstance(value, dict): if isinstance(value, dict):
if node := nodes_map.get(key): if (node := nodes_map.get(key)) or (node := nodes_display_name_map.get(key)):
apply_tweaks(node, value)
elif node := nodes_display_name_map.get(key):
apply_tweaks(node, value) apply_tweaks(node, value)
else: else:
all_nodes_tweaks[key] = value all_nodes_tweaks[key] = value

View file

@ -41,11 +41,13 @@ def get_artifact_type(value, build_result=None) -> str:
case list(): case list():
result = ArtifactType.ARRAY result = ArtifactType.ARRAY
if result == ArtifactType.UNKNOWN: if result == ArtifactType.UNKNOWN and (
if build_result and isinstance(build_result, Generator): build_result
result = ArtifactType.STREAM and isinstance(build_result, Generator)
elif isinstance(value, Message) and isinstance(value.text, Generator): or isinstance(value, Message)
result = ArtifactType.STREAM and isinstance(value.text, Generator)
):
result = ArtifactType.STREAM
return result.value return result.value
@ -56,9 +58,7 @@ def post_process_raw(raw, artifact_type: str):
elif artifact_type == ArtifactType.ARRAY.value: elif artifact_type == ArtifactType.ARRAY.value:
_raw = [] _raw = []
for item in raw: for item in raw:
if hasattr(item, "dict"): if hasattr(item, "dict") or hasattr(item, "model_dump"):
_raw.append(recursive_serialize_or_str(item))
elif hasattr(item, "model_dump"):
_raw.append(recursive_serialize_or_str(item)) _raw.append(recursive_serialize_or_str(item))
else: else:
_raw.append(str(item)) _raw.append(str(item))

View file

@ -103,10 +103,7 @@ class Message(Data):
# they are: "text", "sender" # they are: "text", "sender"
if self.text is None or not self.sender: if self.text is None or not self.sender:
logger.warning("Missing required keys ('text', 'sender') in Message, defaulting to HumanMessage.") logger.warning("Missing required keys ('text', 'sender') in Message, defaulting to HumanMessage.")
if not isinstance(self.text, str): text = "" if not isinstance(self.text, str) else self.text
text = ""
else:
text = self.text
if self.sender == MESSAGE_SENDER_USER or not self.sender: if self.sender == MESSAGE_SENDER_USER or not self.sender:
if self.files: if self.files:
@ -160,9 +157,7 @@ class Message(Data):
@field_serializer("text", mode="plain") @field_serializer("text", mode="plain")
def serialize_text(self, value): def serialize_text(self, value):
if isinstance(value, AsyncIterator): if isinstance(value, AsyncIterator | Iterator):
return ""
if isinstance(value, Iterator):
return "" return ""
return value return value

View file

@ -56,12 +56,13 @@ def get_type(payload):
case str(): case str():
result = LogType.TEXT result = LogType.TEXT
if result == LogType.UNKNOWN: if result == LogType.UNKNOWN and (
if payload and isinstance(payload, Generator): payload
result = LogType.STREAM and isinstance(payload, Generator)
or isinstance(payload, Message)
elif isinstance(payload, Message) and isinstance(payload.text, Generator): and isinstance(payload.text, Generator)
result = LogType.STREAM ):
result = LogType.STREAM
return result return result

View file

@ -72,11 +72,7 @@ class ThreadingInMemoryCache(CacheService, Generic[LockType]): # type: ignore
# Move the key to the end to make it recently used # Move the key to the end to make it recently used
self._cache.move_to_end(key) self._cache.move_to_end(key)
# Check if the value is pickled # Check if the value is pickled
if isinstance(item["value"], bytes): return pickle.loads(item["value"]) if isinstance(item["value"], bytes) else item["value"]
value = pickle.loads(item["value"])
else:
value = item["value"]
return value
self.delete(key) self.delete(key)
return None return None

View file

@ -126,10 +126,7 @@ def save_uploaded_file(file: UploadFile, folder_name):
cache_path = Path(CACHE_DIR) cache_path = Path(CACHE_DIR)
folder_path = cache_path / folder_name folder_path = cache_path / folder_name
filename = file.filename filename = file.filename
if isinstance(filename, str) or isinstance(filename, Path): file_extension = Path(filename).suffix if isinstance(filename, str | Path) else ""
file_extension = Path(filename).suffix
else:
file_extension = ""
file_object = file.file file_object = file.file
# Create the folder if it doesn't exist # Create the folder if it doesn't exist

View file

@ -93,10 +93,7 @@ class CacheService(Subject, Service):
"image": "png", "image": "png",
"pandas": "csv", "pandas": "csv",
} }
if obj_type in object_extensions: _extension = object_extensions[obj_type] if obj_type in object_extensions else type(obj).__name__.lower()
_extension = object_extensions[obj_type]
else:
_extension = type(obj).__name__.lower()
self.current_cache[name] = { self.current_cache[name] = {
"obj": obj, "obj": obj,
"type": obj_type, "type": obj_type,

View file

@ -117,10 +117,10 @@ class FlowBase(SQLModel):
raise ValueError(msg) raise ValueError(msg)
# data must contain nodes and edges # data must contain nodes and edges
if "nodes" not in v.keys(): if "nodes" not in v:
msg = "Flow must have nodes" msg = "Flow must have nodes"
raise ValueError(msg) raise ValueError(msg)
if "edges" not in v.keys(): if "edges" not in v:
msg = "Flow must have edges" msg = "Flow must have edges"
raise ValueError(msg) raise ValueError(msg)

View file

@ -46,12 +46,9 @@ class MessageBase(SQLModel):
timestamp = message.timestamp timestamp = message.timestamp
if not flow_id and message.flow_id: if not flow_id and message.flow_id:
flow_id = message.flow_id flow_id = message.flow_id
if not isinstance(message.text, str): # If the text is not a string, it means it could be
# If the text is not a string, it means it could be # async iterator so we simply add it as an empty string
# async iterator so we simply add it as an empty string message_text = "" if not isinstance(message.text, str) else message.text
message_text = ""
else:
message_text = message.text
return cls( return cls(
sender=message.sender, sender=message.sender,
sender_name=message.sender_name, sender_name=message.sender_name,

View file

@ -31,10 +31,7 @@ def is_list_of_any(field: FieldInfo) -> bool:
if field.annotation is None: if field.annotation is None:
return False return False
try: try:
if hasattr(field.annotation, "__args__"): union_args = field.annotation.__args__ if hasattr(field.annotation, "__args__") else []
union_args = field.annotation.__args__
else:
union_args = []
return field.annotation.__origin__ is list or any( return field.annotation.__origin__ is list or any(
arg.__origin__ is list for arg in union_args if hasattr(arg, "__origin__") arg.__origin__ is list for arg in union_args if hasattr(arg, "__origin__")
@ -267,10 +264,7 @@ class Settings(BaseSettings):
final_path = new_path final_path = new_path
if final_path is None: if final_path is None:
if is_pre_release: final_path = new_pre_path if is_pre_release else new_path
final_path = new_pre_path
else:
final_path = new_path
value = f"sqlite:///{final_path}" value = f"sqlite:///{final_path}"
@ -370,7 +364,7 @@ def load_settings_from_yaml(file_path: str) -> Settings:
settings_dict = {k.upper(): v for k, v in settings_dict.items()} settings_dict = {k.upper(): v for k, v in settings_dict.items()}
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:
msg = f"Key {key} not found in settings" msg = f"Key {key} not found in settings"
raise KeyError(msg) raise KeyError(msg)
logger.debug(f"Loading {len(settings_dict[key])} {key} from {file_path}") logger.debug(f"Loading {len(settings_dict[key])} {key} from {file_path}")

View file

@ -30,7 +30,7 @@ class SettingsService(Service):
settings_dict = {k.upper(): v for k, v in settings_dict.items()} settings_dict = {k.upper(): v for k, v in settings_dict.items()}
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:
msg = f"Key {key} not found in settings" msg = f"Key {key} not found in settings"
raise KeyError(msg) raise KeyError(msg)
logger.debug(f"Loading {len(settings_dict[key])} {key} from {file_path}") logger.debug(f"Loading {len(settings_dict[key])} {key} from {file_path}")

View file

@ -126,10 +126,7 @@ class StoreService(Service):
self, url: str, api_key: str | None = None, params: dict[str, Any] | None = None self, url: str, api_key: str | None = None, params: dict[str, Any] | None = None
) -> tuple[list[dict[str, Any]], dict[str, Any]]: ) -> tuple[list[dict[str, Any]], dict[str, Any]]:
"""Utility method to perform GET requests.""" """Utility method to perform GET requests."""
if api_key: headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
headers = {"Authorization": f"Bearer {api_key}"}
else:
headers = {}
async with httpx.AsyncClient() as client: async with httpx.AsyncClient() as client:
try: try:
response = await client.get(url, headers=headers, params=params, timeout=self.timeout) response = await client.get(url, headers=headers, params=params, timeout=self.timeout)

View file

@ -1,6 +1,8 @@
from __future__ import annotations
import os import os
from datetime import datetime from datetime import datetime
from typing import TYPE_CHECKING, Any, Optional from typing import TYPE_CHECKING, Any
from uuid import UUID from uuid import UUID
from loguru import logger from loguru import logger
@ -10,6 +12,7 @@ from langflow.services.tracing.schema import Log
if TYPE_CHECKING: if TYPE_CHECKING:
from langchain.callbacks.base import BaseCallbackHandler from langchain.callbacks.base import BaseCallbackHandler
from langfuse.client import StatefulSpanClient
from langflow.graph.vertex.base import Vertex from langflow.graph.vertex.base import Vertex
@ -23,7 +26,7 @@ class LangFuseTracer(BaseTracer):
self.trace_type = trace_type self.trace_type = trace_type
self.trace_id = trace_id self.trace_id = trace_id
self.flow_id = trace_name.split(" - ")[-1] self.flow_id = trace_name.split(" - ")[-1]
self.last_span = None self.last_span: StatefulSpanClient | None = None
self.spans: dict = {} self.spans: dict = {}
self._ready: bool = self.setup_langfuse() self._ready: bool = self.setup_langfuse()
@ -68,7 +71,7 @@ class LangFuseTracer(BaseTracer):
trace_type: str, trace_type: str,
inputs: dict[str, Any], inputs: dict[str, Any],
metadata: dict[str, Any] | None = None, metadata: dict[str, Any] | None = None,
vertex: Optional["Vertex"] = None, vertex: Vertex | None = None,
): ):
start_time = datetime.utcnow() start_time = datetime.utcnow()
if not self._ready: if not self._ready:
@ -86,10 +89,7 @@ class LangFuseTracer(BaseTracer):
"start_time": start_time, "start_time": start_time,
} }
if self.last_span: span = self.last_span.span(**content_span) if self.last_span else self.trace.span(**content_span)
span = self.last_span.span(**content_span)
else:
span = self.trace.span(**content_span)
self.last_span = span self.last_span = span
self.spans[trace_id] = span self.spans[trace_id] = span
@ -127,7 +127,7 @@ class LangFuseTracer(BaseTracer):
self._client.flush() self._client.flush()
def get_langchain_callback(self) -> Optional["BaseCallbackHandler"]: def get_langchain_callback(self) -> BaseCallbackHandler | None:
if not self._ready: if not self._ready:
return None return None
return None # self._callback return None # self._callback

View file

@ -243,7 +243,7 @@ class TracingService(Service):
def _cleanup_inputs(self, inputs: dict[str, Any]): def _cleanup_inputs(self, inputs: dict[str, Any]):
inputs = inputs.copy() inputs = inputs.copy()
for key in inputs.keys(): for key in inputs:
if "api_key" in key: if "api_key" in key:
inputs[key] = "*****" # avoid logging api_keys for security reasons inputs[key] = "*****" # avoid logging api_keys for security reasons
return inputs return inputs

View file

@ -20,10 +20,7 @@ def convert_to_langchain_type(value):
else: else:
value = value.to_lc_document() value = value.to_lc_document()
elif isinstance(value, Data): elif isinstance(value, Data):
if "text" in value.data: value = value.to_lc_document() if "text" in value.data else value.data
value = value.to_lc_document()
else:
value = value.data
return value return value

View file

@ -93,7 +93,7 @@ class KubernetesSecretService(VariableService, Service):
return [] return []
names = [] names = []
for key in variables.keys(): for key in variables:
if key.startswith(CREDENTIAL_TYPE + "_"): if key.startswith(CREDENTIAL_TYPE + "_"):
names.append(key[len(CREDENTIAL_TYPE) + 1 :]) names.append(key[len(CREDENTIAL_TYPE) + 1 :])
else: else:

View file

@ -97,9 +97,8 @@ class Input(BaseModel):
def serialize_model(self, handler): def serialize_model(self, handler):
result = handler(self) result = handler(self)
# If the field is str, we add the Text input type # If the field is str, we add the Text input type
if self.field_type in ["str", "Text"]: if self.field_type in ["str", "Text"] and "input_types" not in result:
if "input_types" not in result: result["input_types"] = ["Text"]
result["input_types"] = ["Text"]
if self.field_type == Text: if self.field_type == Text:
result["type"] = "str" result["type"] = "str"
else: else:

View file

@ -56,9 +56,7 @@ def build_template_from_function(name: str, type_to_loader_dict: dict, add_funct
elif name_ not in ["name"]: elif name_ not in ["name"]:
variables[class_field_items][name_] = value_ variables[class_field_items][name_] = value_
variables[class_field_items]["placeholder"] = ( variables[class_field_items]["placeholder"] = docs.params.get(class_field_items, "")
docs.params[class_field_items] if class_field_items in docs.params else ""
)
# Adding function to base classes to allow # Adding function to base classes to allow
# the output to be a function # the output to be a function
base_classes = get_base_classes(_class) base_classes = get_base_classes(_class)

View file

@ -169,6 +169,7 @@ select = [
"Q", "Q",
"RET", "RET",
"RSE", "RSE",
"SIM",
"SLOT", "SLOT",
"T10", "T10",
"TID", "TID",