code format

This commit is contained in:
anovazzi1 2023-12-01 18:19:11 -03:00
commit bc7e612cfb
9 changed files with 50 additions and 56 deletions

View file

@ -12,8 +12,7 @@ from dotenv import load_dotenv
from langflow.main import setup_app from langflow.main import setup_app
from langflow.services.database.utils import session_getter from langflow.services.database.utils import session_getter
from langflow.services.deps import get_db_service, get_settings_service from langflow.services.deps import get_db_service, get_settings_service
from langflow.services.utils import (initialize_services, from langflow.services.utils import initialize_services, initialize_settings_service
initialize_settings_service)
from langflow.utils.logger import configure, logger from langflow.utils.logger import configure, logger
from multiprocess import Process, cpu_count # type: ignore from multiprocess import Process, cpu_count # type: ignore
from rich import box from rich import box
@ -328,18 +327,22 @@ def superuser(
@app.command() @app.command()
def migration(test: bool = typer.Option(True, help="Run migrations in test mode."), def migration(
fix: bool = typer.Option(False, help="Fix migrations. This is a destructive operation, and should only be used if you know what you are doing.") test: bool = typer.Option(True, help="Run migrations in test mode."),
fix: bool = typer.Option(
False,
help="Fix migrations. This is a destructive operation, and should only be used if you know what you are doing.",
),
): ):
""" """
Run or test migrations. Run or test migrations.
""" """
if fix: if fix:
if not typer.confirm("This will delete all data necessary to fix migrations. Are you sure you want to continue?"): if not typer.confirm(
"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()
if not test: if not test:

View file

@ -1,4 +1,3 @@
from typing import TYPE_CHECKING, List from typing import TYPE_CHECKING, List
if TYPE_CHECKING: if TYPE_CHECKING:
@ -83,8 +82,8 @@ def update_frontend_node_with_template_values(frontend_node, raw_template_data):
return frontend_node return frontend_node
def validate_is_component(flows: List["Flow"]):
def validate_is_component(flows: List["Flow"]):
for flow in flows: for flow in flows:
if not flow.data or flow.is_component is not None: if not flow.data or flow.is_component is not None:
continue continue
@ -97,7 +96,6 @@ def validate_is_component(flows: List["Flow"]):
return flows return flows
def get_is_component_from_data(data: dict): def get_is_component_from_data(data: dict):
"""Returns True if the data is a component.""" """Returns True if the data is a component."""
return data.get("is_component") return data.get("is_component")

View file

@ -1,9 +1,27 @@
from .constants import (AgentExecutor, BaseChatMemory, BaseLanguageModel, from .constants import (
BaseLLM, BaseLoader, BaseMemory, BaseOutputParser, AgentExecutor,
BasePromptTemplate, BaseRetriever, Callable, Chain, BaseChatMemory,
ChatPromptTemplate, Data, Document, Embeddings, BaseLanguageModel,
NestedDict, Object, Prompt, PromptTemplate, BaseLLM,
TextSplitter, Tool, VectorStore) BaseLoader,
BaseMemory,
BaseOutputParser,
BasePromptTemplate,
BaseRetriever,
Callable,
Chain,
ChatPromptTemplate,
Data,
Document,
Embeddings,
NestedDict,
Object,
Prompt,
PromptTemplate,
TextSplitter,
Tool,
VectorStore,
)
__all__ = [ __all__ = [
"NestedDict", "NestedDict",
@ -27,5 +45,5 @@ __all__ = [
"Callable", "Callable",
"BasePromptTemplate", "BasePromptTemplate",
"ChatPromptTemplate", "ChatPromptTemplate",
"Prompt" "Prompt",
] ]

View file

@ -5,8 +5,7 @@ from langchain.chains.base import Chain
from langchain.document_loaders.base import BaseLoader from langchain.document_loaders.base import BaseLoader
from langchain.llms.base import BaseLLM from langchain.llms.base import BaseLLM
from langchain.memory.chat_memory import BaseChatMemory from langchain.memory.chat_memory import BaseChatMemory
from langchain.prompts import (BasePromptTemplate, ChatPromptTemplate, from langchain.prompts import BasePromptTemplate, ChatPromptTemplate, PromptTemplate
PromptTemplate)
from langchain.schema import BaseOutputParser, BaseRetriever, Document from langchain.schema import BaseOutputParser, BaseRetriever, Document
from langchain.schema.embeddings import Embeddings from langchain.schema.embeddings import Embeddings
from langchain.schema.language_model import BaseLanguageModel from langchain.schema.language_model import BaseLanguageModel
@ -26,6 +25,7 @@ class Object:
class Data: class Data:
pass pass
class Prompt: class Prompt:
pass pass
@ -48,7 +48,6 @@ LANGCHAIN_BASE_TYPES = {
"BaseOutputParser": BaseOutputParser, "BaseOutputParser": BaseOutputParser,
"BaseMemory": BaseMemory, "BaseMemory": BaseMemory,
"BaseChatMemory": BaseChatMemory, "BaseChatMemory": BaseChatMemory,
} }
# Langchain base types plus Python base types # Langchain base types plus Python base types
CUSTOM_COMPONENT_SUPPORTED_TYPES = { CUSTOM_COMPONENT_SUPPORTED_TYPES = {

View file

@ -6,8 +6,7 @@ from typing import Any, Dict, List, Type, Union
from cachetools import TTLCache, cachedmethod, keys from cachetools import TTLCache, cachedmethod, keys
from fastapi import HTTPException from fastapi import HTTPException
from langflow.interface.custom.schema import (CallableCodeDetails, from langflow.interface.custom.schema import CallableCodeDetails, ClassCodeDetails
ClassCodeDetails)
class CodeSyntaxError(HTTPException): class CodeSyntaxError(HTTPException):
@ -57,9 +56,6 @@ class CodeParser:
ast.Assign: self.parse_global_vars, ast.Assign: self.parse_global_vars,
} }
def __get_tree(self): def __get_tree(self):
""" """
Parses the provided code to validate its syntax. Parses the provided code to validate its syntax.
@ -83,7 +79,6 @@ class CodeParser:
if handler := self.handlers.get(type(node)): # type: ignore if handler := self.handlers.get(type(node)): # type: ignore
handler(node) # type: ignore handler(node) # type: ignore
def parse_imports(self, node: Union[ast.Import, ast.ImportFrom]) -> None: def parse_imports(self, node: Union[ast.Import, ast.ImportFrom]) -> None:
""" """
Extracts "imports" from the code, including aliases. Extracts "imports" from the code, including aliases.
@ -154,7 +149,6 @@ class CodeParser:
# Handle cases where the type is not found in the constructed environment # Handle cases where the type is not found in the constructed environment
pass pass
func = CallableCodeDetails( func = CallableCodeDetails(
name=node.name, name=node.name,
doc=ast.get_docstring(node), doc=ast.get_docstring(node),
@ -164,7 +158,6 @@ class CodeParser:
has_return=self.parse_return_statement(node), has_return=self.parse_return_statement(node),
) )
return func.model_dump() return func.model_dump()
def parse_function_args(self, node: ast.FunctionDef) -> List[Dict[str, Any]]: def parse_function_args(self, node: ast.FunctionDef) -> List[Dict[str, Any]]:
@ -246,7 +239,6 @@ class CodeParser:
return any(isinstance(n, ast.Return) for n in node.body) return any(isinstance(n, ast.Return) for n in node.body)
def parse_assign(self, stmt): def parse_assign(self, stmt):
""" """
Parses an Assign statement and returns a dictionary Parses an Assign statement and returns a dictionary

View file

@ -9,7 +9,8 @@ from langflow.interface.custom.component import Component
from langflow.interface.custom.directory_reader import DirectoryReader from langflow.interface.custom.directory_reader import DirectoryReader
from langflow.interface.custom.utils import ( from langflow.interface.custom.utils import (
extract_inner_type_from_generic_alias, extract_inner_type_from_generic_alias,
extract_union_types_from_generic_alias) extract_union_types_from_generic_alias,
)
from langflow.services.database.models.flow import Flow from langflow.services.database.models.flow import Flow
from langflow.services.database.utils import session_getter from langflow.services.database.utils import session_getter
from langflow.services.deps import get_credential_service, get_db_service from langflow.services.deps import get_credential_service, get_db_service
@ -33,8 +34,6 @@ class CustomComponent(Component):
self.cache = TTLCache(maxsize=1024, ttl=60) self.cache = TTLCache(maxsize=1024, ttl=60)
super().__init__(**data) super().__init__(**data)
def custom_repr(self): def custom_repr(self):
if self.repr_value == "": if self.repr_value == "":
self.repr_value = self.status self.repr_value = self.status
@ -75,8 +74,6 @@ class CustomComponent(Component):
def validate(self) -> bool: def validate(self) -> bool:
return self._class_template_validation(self.code) if self.code else False return self._class_template_validation(self.code) if self.code else False
@property @property
def tree(self): def tree(self):
return self.get_code_tree(self.code) return self.get_code_tree(self.code)
@ -109,8 +106,6 @@ class CustomComponent(Component):
if not self.code: if not self.code:
return [] return []
component_classes = [cls for cls in self.tree["classes"] if self.code_class_base_inheritance in cls["bases"]] component_classes = [cls for cls in self.tree["classes"] if self.code_class_base_inheritance in cls["bases"]]
if not component_classes: if not component_classes:
return [] return []
@ -124,7 +119,6 @@ class CustomComponent(Component):
if not build_methods: if not build_methods:
return [] return []
return build_methods[0] return build_methods[0]
@property @property
@ -156,7 +150,6 @@ class CustomComponent(Component):
if not self.code: if not self.code:
return "" return ""
base_name = self.code_class_base_inheritance base_name = self.code_class_base_inheritance
method_name = self.function_entrypoint_name method_name = self.function_entrypoint_name
@ -175,7 +168,6 @@ class CustomComponent(Component):
if not self.code: if not self.code:
return {} return {}
attributes = [ attributes = [
main_class["attributes"] main_class["attributes"]
for main_class in self.tree.get("classes", []) for main_class in self.tree.get("classes", [])
@ -222,8 +214,7 @@ class CustomComponent(Component):
return validate.create_function(self.code, self.function_entrypoint_name) return validate.create_function(self.code, self.function_entrypoint_name)
async def load_flow(self, flow_id: str, tweaks: Optional[dict] = None) -> Any: async def load_flow(self, flow_id: str, tweaks: Optional[dict] = None) -> Any:
from langflow.processing.process import (build_sorted_vertices, from langflow.processing.process import build_sorted_vertices, process_tweaks
process_tweaks)
db_service = get_db_service() db_service = get_db_service()
with session_getter(db_service) as session: with session_getter(db_service) as session:

View file

@ -29,8 +29,7 @@ from langflow.interface.vector_store.base import vectorstore_creator
from langflow.interface.wrappers.base import wrapper_creator from langflow.interface.wrappers.base import wrapper_creator
from langflow.template.field.base import TemplateField from langflow.template.field.base import TemplateField
from langflow.template.frontend_node.constants import CLASSES_TO_REMOVE from langflow.template.frontend_node.constants import CLASSES_TO_REMOVE
from langflow.template.frontend_node.custom_components import \ from langflow.template.frontend_node.custom_components import CustomComponentFrontendNode
CustomComponentFrontendNode
from langflow.utils.util import get_base_classes from langflow.utils.util import get_base_classes
from loguru import logger from loguru import logger

View file

@ -153,7 +153,7 @@ class DatabaseService(Service):
try: try:
command.check(alembic_cfg) command.check(alembic_cfg)
except util.exc.AutogenerateDiffsDetected as exc: except util.exc.AutogenerateDiffsDetected:
logger.exception("AutogenerateDiffsDetected: {exc}") logger.exception("AutogenerateDiffsDetected: {exc}")
if not fix: if not fix:
raise RuntimeError("Something went wrong running migrations. Please, run `langflow migration --fix`") raise RuntimeError("Something went wrong running migrations. Please, run `langflow migration --fix`")
@ -169,7 +169,6 @@ class DatabaseService(Service):
command.check(alembic_cfg) command.check(alembic_cfg)
break break
except util.exc.AutogenerateDiffsDetected as exc: except util.exc.AutogenerateDiffsDetected as exc:
# downgrade to base and upgrade again # downgrade to base and upgrade again
logger.warning(f"AutogenerateDiffsDetected: {exc}") logger.warning(f"AutogenerateDiffsDetected: {exc}")
command.downgrade(alembic_cfg, f"-{i}") command.downgrade(alembic_cfg, f"-{i}")
@ -177,8 +176,6 @@ class DatabaseService(Service):
time.sleep(3) time.sleep(3)
command.upgrade(alembic_cfg, "head") command.upgrade(alembic_cfg, "head")
def run_migrations_test(self): def run_migrations_test(self):
# This method is used for testing purposes only # This method is used for testing purposes only
# We will check that all models are in the database # We will check that all models are in the database

View file

@ -2,8 +2,7 @@ from langflow.services.auth.utils import create_super_user, verify_password
from langflow.services.database.utils import initialize_database from langflow.services.database.utils import initialize_database
from langflow.services.manager import service_manager from langflow.services.manager import service_manager
from langflow.services.schema import ServiceType from langflow.services.schema import ServiceType
from langflow.services.settings.constants import (DEFAULT_SUPERUSER, from langflow.services.settings.constants import DEFAULT_SUPERUSER, DEFAULT_SUPERUSER_PASSWORD
DEFAULT_SUPERUSER_PASSWORD)
from loguru import logger from loguru import logger
from sqlmodel import Session from sqlmodel import Session
@ -16,8 +15,7 @@ def get_factories_and_deps():
from langflow.services.chat import factory as chat_factory from langflow.services.chat import factory as chat_factory
from langflow.services.credentials import factory as credentials_factory from langflow.services.credentials import factory as credentials_factory
from langflow.services.database import factory as database_factory from langflow.services.database import factory as database_factory
from langflow.services.session import \ from langflow.services.session import factory as session_service_factory # type: ignore
factory as session_service_factory # type: ignore
from langflow.services.settings import factory as settings_factory from langflow.services.settings import factory as settings_factory
from langflow.services.store import factory as store_factory from langflow.services.store import factory as store_factory
from langflow.services.task import factory as task_factory from langflow.services.task import factory as task_factory
@ -173,8 +171,7 @@ def initialize_session_service():
Initialize the session manager. Initialize the session manager.
""" """
from langflow.services.cache import factory as cache_factory from langflow.services.cache import factory as cache_factory
from langflow.services.session import \ from langflow.services.session import factory as session_service_factory # type: ignore
factory as session_service_factory # type: ignore
initialize_settings_service() initialize_settings_service()