Fix imports and add RangeSpec class

This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-12-10 11:17:39 -03:00
commit 6da73895be
4 changed files with 38 additions and 0 deletions

View file

@ -22,6 +22,7 @@ from .constants import (
Tool, Tool,
VectorStore, VectorStore,
) )
from .range_spec import RangeSpec
__all__ = [ __all__ = [
"NestedDict", "NestedDict",
@ -46,4 +47,5 @@ __all__ = [
"BasePromptTemplate", "BasePromptTemplate",
"ChatPromptTemplate", "ChatPromptTemplate",
"Prompt", "Prompt",
"RangeSpec",
] ]

View file

@ -0,0 +1,21 @@
from pydantic import BaseModel, field_validator
class RangeSpec(BaseModel):
min: float = -1.0
max: float = 1.0
step: float = 0.1
@field_validator("max")
@classmethod
def max_must_be_greater_than_min(cls, v, values, **kwargs):
if "min" in values.data and v <= values.data["min"]:
raise ValueError("max must be greater than min")
return v
@field_validator("step")
@classmethod
def step_must_be_positive(cls, v):
if v <= 0:
raise ValueError("step must be positive")
return v

View file

@ -10,6 +10,7 @@ from cachetools import LRUCache, cached
from fastapi import HTTPException from fastapi import HTTPException
from loguru import logger from loguru import logger
from langflow.field_typing.range_spec import RangeSpec
from langflow.interface.agents.base import agent_creator from langflow.interface.agents.base import agent_creator
from langflow.interface.chains.base import chain_creator from langflow.interface.chains.base import chain_creator
from langflow.interface.custom.custom_component import CustomComponent from langflow.interface.custom.custom_component import CustomComponent
@ -256,6 +257,11 @@ def update_field_dict(field_dict):
field_dict["value"] = field_dict["value"](field_dict.get("options", [])) field_dict["value"] = field_dict["value"](field_dict.get("options", []))
field_dict["refresh"] = True field_dict["refresh"] = True
# Let's check if "range_spec" is a RangeSpec object
if "range_spec" in field_dict and isinstance(field_dict["range_spec"], RangeSpec):
field_dict["range_spec"] = field_dict["range_spec"].model_dump()
field_dict["refresh"] = True
def add_extra_fields(frontend_node, field_config, function_args): def add_extra_fields(frontend_node, field_config, function_args):
"""Add extra fields to the frontend node""" """Add extra fields to the frontend node"""

View file

@ -3,6 +3,8 @@ from typing import Any, Optional, Union
from pydantic import BaseModel from pydantic import BaseModel
from langflow.field_typing.range_spec import RangeSpec
class TemplateFieldCreator(BaseModel, ABC): class TemplateFieldCreator(BaseModel, ABC):
field_type: str = "str" field_type: str = "str"
@ -59,6 +61,9 @@ class TemplateFieldCreator(BaseModel, ABC):
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] = None
"""Range specification for the field. Defaults to None."""
def to_dict(self): def to_dict(self):
result = self.model_dump() result = self.model_dump()
# Remove key if it is None # Remove key if it is None
@ -73,6 +78,10 @@ class TemplateFieldCreator(BaseModel, ABC):
if self.field_type == "file": if self.field_type == "file":
result["file_path"] = self.file_path result["file_path"] = self.file_path
# If type is float but range_spec is not set, set it to default
if self.field_type == "float" and not self.range_spec:
result["range_spec"] = RangeSpec().model_dump()
return result return result