fix: Update type hints in BatchRunComponent to avoid lint errors (#7567)

* fix: update type hints in BatchRunComponent for improved clarity and consistency

* fix: add type casting in BatchRunComponent for improved type safety

* [autofix.ci] apply automated fixes

---------

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Gabriel Luiz Freitas Almeida 2025-04-10 14:23:40 -03:00 • committed by GitHub
commit 4173656c89
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 19 additions and 17 deletions

View file

@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any, cast
import toml # type: ignore[import-untyped] import toml # type: ignore[import-untyped]
from loguru import logger from loguru import logger
@ -10,8 +10,6 @@ from langflow.io import BoolInput, DataFrameInput, HandleInput, MessageTextInput
from langflow.schema import DataFrame from langflow.schema import DataFrame
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Hashable
from langchain_core.runnables import Runnable from langchain_core.runnables import Runnable
@ -78,14 +76,14 @@ class BatchRunComponent(Component):
), ),
] ]
def _format_row_as_toml(self, row: dict[Hashable, Any]) -> str: def _format_row_as_toml(self, row: dict[str, Any]) -> str:
"""Convert a dictionary (row) into a TOML-formatted string.""" """Convert a dictionary (row) into a TOML-formatted string."""
formatted_dict = {str(col): {"value": str(val)} for col, val in row.items()} formatted_dict = {str(col): {"value": str(val)} for col, val in row.items()}
return toml.dumps(formatted_dict) return toml.dumps(formatted_dict)
def _create_base_row( def _create_base_row(
self, original_row: dict[Hashable, Any], model_response: str = "", batch_index: int = -1 self, original_row: dict[str, Any], model_response: str = "", batch_index: int = -1
) -> dict[Hashable, Any]: ) -> dict[str, Any]:
"""Create a base row with original columns and additional metadata.""" """Create a base row with original columns and additional metadata."""
row = original_row.copy() row = original_row.copy()
row[self.output_column_name] = model_response row[self.output_column_name] = model_response
@ -93,7 +91,7 @@ class BatchRunComponent(Component):
return row return row
def _add_metadata( def _add_metadata(
self, row: dict[Hashable, Any], *, success: bool = True, system_msg: str = "", error: str | None = None self, row: dict[str, Any], *, success: bool = True, system_msg: str = "", error: str | None = None
) -> None: ) -> None:
"""Add metadata to a row if enabled.""" """Add metadata to a row if enabled."""
if not self.enable_metadata: if not self.enable_metadata:
@ -145,7 +143,9 @@ class BatchRunComponent(Component):
if col_name: if col_name:
user_texts = df[col_name].astype(str).tolist() user_texts = df[col_name].astype(str).tolist()
else: else:
user_texts = [self._format_row_as_toml(row) for row in df.to_dict(orient="records")] user_texts = [
self._format_row_as_toml(cast(dict[str, Any], row)) for row in df.to_dict(orient="records")
]
total_rows = len(user_texts) total_rows = len(user_texts)
logger.info(f"Processing {total_rows} rows with batch run") logger.info(f"Processing {total_rows} rows with batch run")
@ -166,25 +166,27 @@ class BatchRunComponent(Component):
"callbacks": self.get_langchain_callbacks(), "callbacks": self.get_langchain_callbacks(),
} }
) )
# Process batches and track progress # Process batches and track progress
responses_with_idx = [ responses_with_idx = list(
(idx, response) zip(
for idx, response in zip( range(len(conversations)),
range(len(conversations)), await model.abatch(list(conversations)), strict=True await model.abatch(list(conversations)),
strict=True,
) )
] )
# Sort by index to maintain order # Sort by index to maintain order
responses_with_idx.sort(key=lambda x: x[0]) responses_with_idx.sort(key=lambda x: x[0])
# Build the final data with enhanced metadata # Build the final data with enhanced metadata
rows: list[dict[Hashable, Any]] = [] rows: list[dict[str, Any]] = []
for idx, (original_row, response) in enumerate( for idx, (original_row, response) in enumerate(
zip(df.to_dict(orient="records"), responses_with_idx, strict=False) zip(df.to_dict(orient="records"), responses_with_idx, strict=False)
): ):
response_text = response[1].content if hasattr(response[1], "content") else str(response[1]) response_text = response[1].content if hasattr(response[1], "content") else str(response[1])
row = self._create_base_row(original_row, model_response=response_text, batch_index=idx) row = self._create_base_row(
cast(dict[str, Any], original_row), model_response=response_text, batch_index=idx
)
self._add_metadata(row, success=True, system_msg=system_msg) self._add_metadata(row, success=True, system_msg=system_msg)
rows.append(row) rows.append(row)

File diff suppressed because one or more lines are too long