style: Fix type comparison issues in codebase (#3606)

style: fix E721
This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-08-29 08:48:15 -03:00 • committed by GitHub
commit e7a48d58e6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 9 additions and 9 deletions

View file

@ -44,7 +44,7 @@ def add_output_types(frontend_node: CustomComponentFrontendNode, return_types: l
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) )
if return_type == str: if return_type is str:
return_type = "Text" return_type = "Text"
elif hasattr(return_type, "__name__"): elif hasattr(return_type, "__name__"):
return_type = return_type.__name__ return_type = return_type.__name__
@ -85,7 +85,7 @@ def add_base_classes(frontend_node: CustomComponentFrontendNode, return_types: l
) )
base_classes = get_base_classes(return_type_instance) base_classes = get_base_classes(return_type_instance)
if return_type_instance == str: if return_type_instance is str:
base_classes.append("Text") base_classes.append("Text")
for base_class in base_classes: for base_class in base_classes:

View file

@ -2,7 +2,7 @@ from typing import Any
def format_type(type_: Any) -> str: def format_type(type_: Any) -> str:
if type_ == str: if type_ is str:
type_ = "Text" type_ = "Text"
elif hasattr(type_, "__name__"): elif hasattr(type_, "__name__"):
type_ = type_.__name__ type_ = type_.__name__

View file

@ -36,8 +36,8 @@ def is_list_of_any(field: FieldInfo) -> bool:
else: else:
union_args = [] union_args = []
return field.annotation.__origin__ == list or any( return field.annotation.__origin__ is list or any(
arg.__origin__ == list for arg in union_args if hasattr(arg, "__origin__") arg.__origin__ is list for arg in union_args if hasattr(arg, "__origin__")
) )
except AttributeError: except AttributeError:
return False return False

View file

@ -119,7 +119,7 @@ class Input(BaseModel):
@field_serializer("field_type") @field_serializer("field_type")
def serialize_field_type(self, value, _info): def serialize_field_type(self, value, _info):
if value == float and self.range_spec is None: if value is float and self.range_spec is None:
self.range_spec = RangeSpec() self.range_spec = RangeSpec()
return value return value

View file

@ -116,7 +116,7 @@ class TestCreateInputSchema:
input_instance = FileInput(name="file_field") input_instance = FileInput(name="file_field")
schema = create_input_schema([input_instance]) schema = create_input_schema([input_instance])
field_info = schema.model_fields["file_field"] field_info = schema.model_fields["file_field"]
assert field_info.annotation == str assert field_info.annotation is str
# Inputs with mixed required and optional fields are processed correctly # Inputs with mixed required and optional fields are processed correctly
def test_mixed_required_optional_fields_processing(self): def test_mixed_required_optional_fields_processing(self):
@ -184,7 +184,7 @@ class TestCreateInputSchema:
input_instance = StrInput(name="test_field", is_list=True) input_instance = StrInput(name="test_field", is_list=True)
schema = create_input_schema([input_instance]) schema = create_input_schema([input_instance])
field_info = schema.model_fields["test_field"] field_info = schema.model_fields["test_field"]
assert field_info.annotation == list[str] assert field_info.annotation == list[str] # type: ignore
# Converting FieldTypes to corresponding Python types # Converting FieldTypes to corresponding Python types
def test_field_types_conversion(self): def test_field_types_conversion(self):
@ -194,7 +194,7 @@ class TestCreateInputSchema:
input_instance = IntInput(name="int_field") input_instance = IntInput(name="int_field")
schema = create_input_schema([input_instance]) schema = create_input_schema([input_instance])
field_info = schema.model_fields["int_field"] field_info = schema.model_fields["int_field"]
assert field_info.annotation == int assert field_info.annotation is int # Use 'is' for type comparison
# Setting default values for non-required fields # Setting default values for non-required fields
def test_default_values_for_non_required_fields(self): def test_default_values_for_non_required_fields(self):