fix: update type extraction to extract inner types correctly (#3446)

* refactor: Improve post-processing of return type in type extraction.

* feat: Add support for Union[int, Sequence[str]] in post_process_type.

* refactor(langflow): Simplify return type extraction in CustomComponent.

* Refactor: Remove 'Sequence' and 'list' type options in starter projects.

* test: fix assertion

---------

Co-authored-by: italojohnny <italojohnnydosanjos@gmail.com>
This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-08-20 09:48:18 -03:00 • committed by GitHub
commit 037c0902d5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 18 additions and 37 deletions

View file

@ -17,10 +17,7 @@ from langflow.services.deps import get_storage_service, get_variable_service, se
from langflow.services.storage.service import StorageService from langflow.services.storage.service import StorageService
from langflow.services.tracing.schema import Log from langflow.services.tracing.schema import Log
from langflow.template.utils import update_frontend_node_with_template_values from langflow.template.utils import update_frontend_node_with_template_values
from langflow.type_extraction.type_extraction import ( from langflow.type_extraction.type_extraction import post_process_type
extract_inner_type_from_generic_alias,
extract_union_types_from_generic_alias,
)
from langflow.utils import validate from langflow.utils import validate
if TYPE_CHECKING: if TYPE_CHECKING:
@ -365,19 +362,7 @@ class CustomComponent(BaseComponent):
return self.get_method_return_type(self._function_entrypoint_name) return self.get_method_return_type(self._function_entrypoint_name)
def _extract_return_type(self, return_type: Any) -> List[Any]: def _extract_return_type(self, return_type: Any) -> List[Any]:
if hasattr(return_type, "__origin__") and return_type.__origin__ in [ return post_process_type(return_type)
list,
List,
]:
return_type = extract_inner_type_from_generic_alias(return_type)
# If the return type is not a Union, then we just return it as a list
inner_type = return_type[0] if isinstance(return_type, list) else return_type
if not hasattr(inner_type, "__origin__") or inner_type.__origin__ != Union:
return return_type if isinstance(return_type, list) else [return_type]
# If the return type is a Union, then we need to parse it
return_type = extract_union_types_from_generic_alias(return_type)
return return_type
@property @property
def get_main_class_name(self): def get_main_class_name(self):

View file

@ -4037,8 +4037,7 @@
"name": "api_run_model", "name": "api_run_model",
"selected": "Data", "selected": "Data",
"types": [ "types": [
"Data", "Data"
"list"
], ],
"value": "__UNDEFINED__" "value": "__UNDEFINED__"
}, },
@ -4049,8 +4048,7 @@
"name": "api_build_tool", "name": "api_build_tool",
"selected": "Tool", "selected": "Tool",
"types": [ "types": [
"Tool", "Tool"
"Sequence"
], ],
"value": "__UNDEFINED__" "value": "__UNDEFINED__"
} }

View file

@ -2615,8 +2615,7 @@
"name": "api_run_model", "name": "api_run_model",
"selected": "Data", "selected": "Data",
"types": [ "types": [
"Data", "Data"
"list"
], ],
"value": "__UNDEFINED__" "value": "__UNDEFINED__"
}, },
@ -2627,8 +2626,7 @@
"name": "api_build_tool", "name": "api_build_tool",
"selected": "Tool", "selected": "Tool",
"types": [ "types": [
"Tool", "Tool"
"Sequence"
], ],
"value": "__UNDEFINED__" "value": "__UNDEFINED__"
} }

View file

@ -2953,8 +2953,7 @@
"name": "api_run_model", "name": "api_run_model",
"selected": "Data", "selected": "Data",
"types": [ "types": [
"Data", "Data"
"list"
], ],
"value": "__UNDEFINED__" "value": "__UNDEFINED__"
}, },
@ -2965,8 +2964,7 @@
"name": "api_build_tool", "name": "api_build_tool",
"selected": "Tool", "selected": "Tool",
"types": [ "types": [
"Tool", "Tool"
"Sequence"
], ],
"value": "__UNDEFINED__" "value": "__UNDEFINED__"
} }

View file

@ -1,4 +1,6 @@
import re import re
from collections.abc import Sequence as SequenceABC
from itertools import chain
from types import GenericAlias from types import GenericAlias
from typing import Any, List, Union from typing import Any, List, Union
@ -7,7 +9,7 @@ def extract_inner_type_from_generic_alias(return_type: GenericAlias) -> Any:
""" """
Extracts the inner type from a type hint that is a list or a Optional. Extracts the inner type from a type hint that is a list or a Optional.
""" """
if return_type.__origin__ == list: if return_type.__origin__ in [list, SequenceABC]:
return list(return_type.__args__) return list(return_type.__args__)
return return_type return return_type
@ -57,10 +59,7 @@ def post_process_type(_type):
Union[List[Any], Any]: The processed return type. Union[List[Any], Any]: The processed return type.
""" """
if hasattr(_type, "__origin__") and _type.__origin__ in [ if hasattr(_type, "__origin__") and _type.__origin__ in [list, List, SequenceABC]:
list,
List,
]:
_type = extract_inner_type_from_generic_alias(_type) _type = extract_inner_type_from_generic_alias(_type)
# If the return type is not a Union, then we just return it as a list # If the return type is not a Union, then we just return it as a list
@ -69,7 +68,8 @@ def post_process_type(_type):
return _type if isinstance(_type, list) else [_type] return _type if isinstance(_type, list) else [_type]
# If the return type is a Union, then we need to parse it # If the return type is a Union, then we need to parse it
_type = extract_union_types_from_generic_alias(_type) _type = extract_union_types_from_generic_alias(_type)
return _type _type = set(chain.from_iterable([post_process_type(t) for t in _type]))
return list(_type)
def extract_union_types_from_generic_alias(return_type: GenericAlias) -> list: def extract_union_types_from_generic_alias(return_type: GenericAlias) -> list:

View file

@ -1,4 +1,4 @@
from typing import Union from typing import Sequence, Union
import pytest import pytest
from pydantic import ValidationError from pydantic import ValidationError
@ -42,6 +42,8 @@ class TestInput:
assert post_process_type(int) == [int] assert post_process_type(int) == [int]
assert post_process_type(list[int]) == [int] assert post_process_type(list[int]) == [int]
assert post_process_type(Union[int, str]) == [int, str] assert post_process_type(Union[int, str]) == [int, str]
assert post_process_type(Union[int, Sequence[str]]) == [int, str]
assert post_process_type(Union[int, Sequence[int]]) == [int]
def test_input_to_dict(self): def test_input_to_dict(self):
input_obj = Input(field_type="str") input_obj = Input(field_type="str")
@ -126,4 +128,4 @@ class TestPostProcessType:
class CustomType: class CustomType:
pass pass
assert post_process_type(Union[CustomType, int]) == [CustomType, int] assert set(post_process_type(Union[CustomType, int])) == {CustomType, int}