From 8810a0392a7e3cc6d58b5fb9e8b7a7a7c6f716fe Mon Sep 17 00:00:00 2001 From: Gabriel Almeida Date: Mon, 8 May 2023 15:05:28 -0300 Subject: [PATCH] refactor(langflow): rename BaseLLM to BaseLanguageModel in multiple files fix(langflow): fix type hints in multiple files --- src/backend/langflow/interface/agents/custom.py | 10 ++++++---- src/backend/langflow/interface/importing/utils.py | 7 +++---- src/backend/langflow/interface/loading.py | 4 ++-- src/backend/langflow/interface/tools/base.py | 4 +++- 4 files changed, 14 insertions(+), 11 deletions(-) diff --git a/src/backend/langflow/interface/agents/custom.py b/src/backend/langflow/interface/agents/custom.py index 9e0052c8c..6d2efed68 100644 --- a/src/backend/langflow/interface/agents/custom.py +++ b/src/backend/langflow/interface/agents/custom.py @@ -27,7 +27,7 @@ from langchain.agents.agent_toolkits.vectorstore.prompt import ( from langchain.agents.mrkl.prompt import FORMAT_INSTRUCTIONS from langchain.agents.mrkl.prompt import FORMAT_INSTRUCTIONS as SQL_FORMAT_INSTRUCTIONS from langchain.base_language import BaseLanguageModel -from langchain.llms.base import BaseLLM + from langchain.memory.chat_memory import BaseChatMemory from langchain.sql_database import SQLDatabase from langchain.tools.python.tool import PythonAstREPLTool @@ -134,7 +134,7 @@ class VectorStoreAgent(AgentExecutor): @classmethod def from_toolkit_and_llm( - cls, llm: BaseLLM, vectorstoreinfo: VectorStoreInfo, **kwargs: Any + cls, llm: BaseLanguageModel, vectorstoreinfo: VectorStoreInfo, **kwargs: Any ): """Construct a vectorstore agent from an LLM and tools.""" @@ -171,7 +171,9 @@ class SQLAgent(AgentExecutor): super().__init__(*args, **kwargs) @classmethod - def from_toolkit_and_llm(cls, llm: BaseLLM, database_uri: str, **kwargs: Any): + def from_toolkit_and_llm( + cls, llm: BaseLanguageModel, database_uri: str, **kwargs: Any + ): """Construct a sql agent from an LLM and tools.""" db = SQLDatabase.from_uri(database_uri) toolkit = SQLDatabaseToolkit(db=db, llm=llm) @@ -275,7 +277,7 @@ class InitializeAgent(AgentExecutor): @classmethod def initialize( cls, - llm: BaseLLM, + llm: BaseLanguageModel, tools: List[Tool], agent: str, memory: Optional[BaseChatMemory] = None, diff --git a/src/backend/langflow/interface/importing/utils.py b/src/backend/langflow/interface/importing/utils.py index 499b70a65..0d14249cb 100644 --- a/src/backend/langflow/interface/importing/utils.py +++ b/src/backend/langflow/interface/importing/utils.py @@ -7,7 +7,7 @@ from langchain import PromptTemplate from langchain.agents import Agent from langchain.chains.base import Chain from langchain.chat_models.base import BaseChatModel -from langchain.llms.base import BaseLLM +from langchain.base_language import BaseLanguageModel from langchain.tools import BaseTool @@ -98,7 +98,7 @@ def import_agent(agent: str) -> Agent: return import_class(f"langchain.agents.{agent}") -def import_llm(llm: str) -> BaseLLM: +def import_llm(llm: str) -> BaseLanguageModel: """Import llm from llm name""" return import_class(f"langchain.llms.{llm}") @@ -106,9 +106,8 @@ def import_llm(llm: str) -> BaseLLM: def import_tool(tool: str) -> BaseTool: """Import tool from tool name""" from langflow.interface.tools.base import tool_creator - from langflow.interface.tools.constants import ALL_TOOLS_NAMES - if tool in ALL_TOOLS_NAMES: + if tool in tool_creator.type_to_loader_dict: return tool_creator.type_to_loader_dict[tool]["fcn"] return import_class(f"langchain.tools.{tool}") diff --git a/src/backend/langflow/interface/loading.py b/src/backend/langflow/interface/loading.py index 2b2d8e2c0..a2d9a799c 100644 --- a/src/backend/langflow/interface/loading.py +++ b/src/backend/langflow/interface/loading.py @@ -15,7 +15,7 @@ from langchain.agents.loading import load_agent_from_config from langchain.agents.tools import Tool from langchain.callbacks.base import BaseCallbackManager from langchain.chains.loading import load_chain_from_config -from langchain.llms.base import BaseLLM +from langchain.base_language import BaseLanguageModel from langchain.llms.loading import load_llm_from_config from pydantic import ValidationError @@ -186,7 +186,7 @@ def load_langchain_type_from_config(config: Dict[str, Any]): def load_agent_executor_from_config( config: dict, - llm: Optional[BaseLLM] = None, + llm: Optional[BaseLanguageModel] = None, tools: Optional[list[Tool]] = None, callback_manager: Optional[BaseCallbackManager] = None, **kwargs: Any, diff --git a/src/backend/langflow/interface/tools/base.py b/src/backend/langflow/interface/tools/base.py index 10eeead03..888accb29 100644 --- a/src/backend/langflow/interface/tools/base.py +++ b/src/backend/langflow/interface/tools/base.py @@ -29,7 +29,9 @@ TOOL_INPUTS = { placeholder="", value="", ), - "llm": TemplateField(field_type="BaseLLM", required=True, is_list=False, show=True), + "llm": TemplateField( + field_type="BaseLanguageModel", required=True, is_list=False, show=True + ), "func": TemplateField( field_type="function", required=True,