Feat: Shared Component Cache Service (#4052)

split branch / PR
This commit is contained in:
Sebastián Estévez 2024-10-07 20:58:42 -04:00 • committed by GitHub
commit 9adf1ef2e5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
16 changed files with 818 additions and 531 deletions

View file

@ -97,7 +97,7 @@ dependencies = [
"yfinance>=0.2.40", "yfinance>=0.2.40",
"langchain-google-community==1.0.7", "langchain-google-community==1.0.7",
"wolframalpha>=5.1.3", "wolframalpha>=5.1.3",
"astra-assistants>=2.1.2", "astra-assistants>=2.1.4",
"composio-langchain==0.5.9", "composio-langchain==0.5.9",
"spider-client>=0.0.27", "spider-client>=0.0.27",
"nltk>=3.9.1", "nltk>=3.9.1",
@ -136,7 +136,6 @@ clickhouse-connect = [
[project.scripts] [project.scripts]
langflow = "langflow.__main__:main" langflow = "langflow.__main__:main"
[tool.uv] [tool.uv]
dev-dependencies = [ dev-dependencies = [
"pytest-instafail>=0.5.0", "pytest-instafail>=0.5.0",

View file

@ -1,3 +1,4 @@
from .astra_assistant_manager import AstraAssistantManager
from .create_assistant import AssistantsCreateAssistant from .create_assistant import AssistantsCreateAssistant
from .create_thread import AssistantsCreateThread from .create_thread import AssistantsCreateThread
from .dotenv import Dotenv from .dotenv import Dotenv
@ -7,6 +8,7 @@ from .list_assistants import AssistantsListAssistants
from .run import AssistantsRun from .run import AssistantsRun
__all__ = [ __all__ = [
"AstraAssistantManager",
"AssistantsCreateAssistant", "AssistantsCreateAssistant",
"AssistantsGetAssistantName", "AssistantsGetAssistantName",
"AssistantsListAssistants", "AssistantsListAssistants",

View file

@ -0,0 +1,134 @@
import asyncio
from astra_assistants.astra_assistants_manager import AssistantManager
from langflow.components.astra_assistants.util import (
get_patched_openai_client,
litellm_model_names,
tool_names,
tools_and_names,
)
from langflow.custom.custom_component.component_with_cache import ComponentWithCache
from langflow.inputs import DropdownInput, MultilineInput, StrInput
from langflow.schema.message import Message
from langflow.template import Output
class AstraAssistantManager(ComponentWithCache):
display_name = "Astra Assistant Manager"
description = "Manages Assistant Interactions"
icon = "bot"
inputs = [
StrInput(
name="instructions",
display_name="Instructions",
info="Instructions for the assistant, think of these as the system prompt.",
),
DropdownInput(
name="model_name",
display_name="Model Name",
advanced=False,
options=litellm_model_names,
value="gpt-4o-mini",
),
DropdownInput(
display_name="Tool",
name="tool",
options=tool_names,
),
MultilineInput(
name="user_message",
display_name="User Message",
info="User message to pass to the run.",
),
MultilineInput(
name="input_thread_id",
display_name="Thread ID (optional)",
info="ID of the thread",
),
MultilineInput(
name="input_assistant_id",
display_name="Assistant ID (optional)",
info="ID of the assistant",
),
MultilineInput(
name="env_set",
display_name="Environment Set",
info="Dummy input to allow chaining with Dotenv Component.",
),
]
outputs = [
Output(display_name="Assistant Response", name="assistant_response", method="get_assistant_response"),
Output(display_name="Tool output", name="tool_output", method="get_tool_output"),
Output(display_name="Thread Id", name="output_thread_id", method="get_thread_id"),
Output(display_name="Assistant Id", name="output_assistant_id", method="get_assistant_id"),
]
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.lock = asyncio.Lock()
self.initialized = False
self.assistant_response = None
self.tool_output = None
self.thread_id = None
self.assistant_id = None
self.client = get_patched_openai_client(self._shared_component_cache)
async def get_assistant_response(self) -> Message:
await self.initialize()
return self.assistant_response
async def get_tool_output(self) -> Message:
await self.initialize()
return self.tool_output
async def get_thread_id(self) -> Message:
await self.initialize()
return self.thread_id
async def get_assistant_id(self) -> Message:
await self.initialize()
return self.assistant_id
async def initialize(self):
async with self.lock:
if not self.initialized:
await self.process_inputs()
self.initialized = True
async def process_inputs(self):
print(f"env_set is {self.env_set}")
print(self.tool)
tools = []
tool_obj = None
if self.tool is not None and self.tool != "":
tool_cls = tools_and_names[self.tool]
tool_obj = tool_cls()
tools.append(tool_obj)
assistant_id = None
thread_id = None
if self.input_assistant_id:
assistant_id = self.input_assistant_id
if self.input_thread_id:
thread_id = self.input_thread_id
assistant_manager = AssistantManager(
instructions=self.instructions,
model=self.model_name,
name="managed_assistant",
tools=tools,
client=self.client,
thread_id=thread_id,
assistant_id=assistant_id,
)
content = self.user_message
result = await assistant_manager.run_thread(content=content, tool=tool_obj)
self.assistant_response = Message(text=result["text"])
if "decision" in result:
self.tool_output = Message(text=str(result["decision"].is_complete))
else:
self.tool_output = Message(text=result["text"])
self.thread_id = Message(text=assistant_manager.thread.id)
self.assistant_id = Message(text=assistant_manager.assistant.id)

View file

@ -1,17 +1,14 @@
from astra_assistants import patch # type: ignore from langflow.components.astra_assistants.util import get_patched_openai_client
from openai import OpenAI from langflow.custom.custom_component.component_with_cache import ComponentWithCache
from langflow.custom import Component
from langflow.inputs import MultilineInput, StrInput from langflow.inputs import MultilineInput, StrInput
from langflow.schema.message import Message from langflow.schema.message import Message
from langflow.template import Output from langflow.template import Output
class AssistantsCreateAssistant(Component): class AssistantsCreateAssistant(ComponentWithCache):
icon = "bot" icon = "bot"
display_name = "Create Assistant" display_name = "Create Assistant"
description = "Creates an Assistant and returns it's id" description = "Creates an Assistant and returns it's id"
client = patch(OpenAI())
inputs = [ inputs = [
StrInput( StrInput(
@ -46,6 +43,10 @@ class AssistantsCreateAssistant(Component):
Output(display_name="Assistant ID", name="assistant_id", method="process_inputs"), Output(display_name="Assistant ID", name="assistant_id", method="process_inputs"),
] ]
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.client = get_patched_openai_client(self._shared_component_cache)
def process_inputs(self) -> Message: def process_inputs(self) -> Message:
print(f"env_set is {self.env_set}") print(f"env_set is {self.env_set}")
assistant = self.client.beta.assistants.create( assistant = self.client.beta.assistants.create(

View file

@ -1,16 +1,13 @@
from astra_assistants import patch # type: ignore from langflow.components.astra_assistants.util import get_patched_openai_client
from openai import OpenAI from langflow.custom.custom_component.component_with_cache import ComponentWithCache
from langflow.custom import Component
from langflow.inputs import MultilineInput from langflow.inputs import MultilineInput
from langflow.schema.message import Message from langflow.schema.message import Message
from langflow.template import Output from langflow.template import Output
class AssistantsCreateThread(Component): class AssistantsCreateThread(ComponentWithCache):
display_name = "Create Assistant Thread" display_name = "Create Assistant Thread"
description = "Creates a thread and returns the thread id" description = "Creates a thread and returns the thread id"
client = patch(OpenAI())
inputs = [ inputs = [
MultilineInput( MultilineInput(
@ -24,6 +21,10 @@ class AssistantsCreateThread(Component):
Output(display_name="Thread ID", name="thread_id", method="process_inputs"), Output(display_name="Thread ID", name="thread_id", method="process_inputs"),
] ]
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.client = get_patched_openai_client(self._shared_component_cache)
def process_inputs(self) -> Message: def process_inputs(self) -> Message:
thread = self.client.beta.threads.create() thread = self.client.beta.threads.create()
thread_id = thread.id thread_id = thread.id

View file

@ -1,16 +1,13 @@
from astra_assistants import patch # type: ignore from langflow.components.astra_assistants.util import get_patched_openai_client
from openai import OpenAI from langflow.custom.custom_component.component_with_cache import ComponentWithCache
from langflow.custom import Component
from langflow.inputs import MultilineInput, StrInput from langflow.inputs import MultilineInput, StrInput
from langflow.schema.message import Message from langflow.schema.message import Message
from langflow.template import Output from langflow.template import Output
class AssistantsGetAssistantName(Component): class AssistantsGetAssistantName(ComponentWithCache):
display_name = "Get Assistant name" display_name = "Get Assistant name"
description = "Assistant by id" description = "Assistant by id"
client = patch(OpenAI())
inputs = [ inputs = [
StrInput( StrInput(
@ -29,6 +26,10 @@ class AssistantsGetAssistantName(Component):
Output(display_name="Assistant Name", name="assistant_name", method="process_inputs"), Output(display_name="Assistant Name", name="assistant_name", method="process_inputs"),
] ]
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.client = get_patched_openai_client(self._shared_component_cache)
def process_inputs(self) -> Message: def process_inputs(self) -> Message:
assistant = self.client.beta.assistants.retrieve( assistant = self.client.beta.assistants.retrieve(
assistant_id=self.assistant_id, assistant_id=self.assistant_id,

View file

@ -1,20 +1,21 @@
from astra_assistants import patch # type: ignore from langflow.components.astra_assistants.util import get_patched_openai_client
from openai import OpenAI from langflow.custom.custom_component.component_with_cache import ComponentWithCache
from langflow.custom import Component
from langflow.schema.message import Message from langflow.schema.message import Message
from langflow.template.field.base import Output from langflow.template.field.base import Output
class AssistantsListAssistants(Component): class AssistantsListAssistants(ComponentWithCache):
display_name = "List Assistants" display_name = "List Assistants"
description = "Returns a list of assistant id's" description = "Returns a list of assistant id's"
client = patch(OpenAI())
outputs = [ outputs = [
Output(display_name="Assistants", name="assistants", method="process_inputs"), Output(display_name="Assistants", name="assistants", method="process_inputs"),
] ]
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.client = get_patched_openai_client(self._shared_component_cache)
def process_inputs(self) -> Message: def process_inputs(self) -> Message:
assistants = self.client.beta.assistants.list().data assistants = self.client.beta.assistants.list().data
id_list = [assistant.id for assistant in assistants] id_list = [assistant.id for assistant in assistants]

View file

@ -4,17 +4,22 @@ from astra_assistants import patch # type: ignore
from openai import OpenAI from openai import OpenAI
from openai.lib.streaming import AssistantEventHandler from openai.lib.streaming import AssistantEventHandler
from langflow.custom import Component from langflow.components.astra_assistants.util import get_patched_openai_client
from langflow.custom.custom_component.component_with_cache import ComponentWithCache
from langflow.inputs import MultilineInput from langflow.inputs import MultilineInput
from langflow.schema import dotdict from langflow.schema import dotdict
from langflow.schema.message import Message from langflow.schema.message import Message
from langflow.template import Output from langflow.template import Output
class AssistantsRun(Component): class AssistantsRun(ComponentWithCache):
display_name = "Run Assistant" display_name = "Run Assistant"
description = "Executes an Assistant Run against a thread" description = "Executes an Assistant Run against a thread"
client = patch(OpenAI())
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.client = get_patched_openai_client(self._shared_component_cache)
self.thread_id = None
def update_build_config( def update_build_config(
self, self,

View file

@ -0,0 +1,60 @@
import importlib
import inspect
import json
import os
import pkgutil
import threading
import astra_assistants.tools as astra_assistants_tools
import requests
from astra_assistants import OpenAI, patch
from astra_assistants.tools.tool_interface import ToolInterface
client_lock = threading.Lock()
client = None
def get_patched_openai_client(shared_component_cache):
os.environ["ASTRA_ASSISTANTS_QUIET"] = "true"
client = shared_component_cache.get("client")
if client is None:
client = patch(OpenAI())
shared_component_cache.set("client", client)
return client
url = "https://raw.githubusercontent.com/BerriAI/litellm/refs/heads/main/model_prices_and_context_window.json"
response = requests.get(url)
data = json.loads(response.text)
# Extract the model names into a Python list
litellm_model_names = []
for model, _ in data.items():
if model != "sample_spec":
# litellm_model_names.append(f"{details['litellm_provider']}/{model}")
litellm_model_names.append(model)
# To store the class names that extend ToolInterface
tool_names = []
tools_and_names = {}
def tools_from_package(your_package):
# Iterate over all modules in the package
package_name = your_package.__name__
for module_info in pkgutil.iter_modules(your_package.__path__):
module_name = f"{package_name}.{module_info.name}"
# Dynamically import the module
module = importlib.import_module(module_name)
# Iterate over all members of the module
for name, obj in inspect.getmembers(module, inspect.isclass):
# Check if the class is a subclass of ToolInterface and is not ToolInterface itself
if issubclass(obj, ToolInterface) and obj is not ToolInterface:
tool_names.append(name)
tools_and_names[name] = obj
tools_from_package(astra_assistants_tools)

View file

@ -0,0 +1,8 @@
from langflow.custom import Component
from langflow.services.deps import get_shared_component_cache_service
class ComponentWithCache(Component):
def __init__(self, **data):
super().__init__(**data)
self._shared_component_cache = get_shared_component_cache_service()

View file

@ -210,6 +210,18 @@ def get_cache_service() -> CacheService:
return get_service(ServiceType.CACHE_SERVICE, CacheServiceFactory()) # type: ignore return get_service(ServiceType.CACHE_SERVICE, CacheServiceFactory()) # type: ignore
def get_shared_component_cache_service() -> CacheService:
"""
Retrieves the cache service from the service manager.
Returns:
The cache service instance.
"""
from langflow.services.shared_component_cache.factory import SharedComponentCacheServiceFactory
return get_service(ServiceType.SHARED_COMPONENT_CACHE_SERVICE, SharedComponentCacheServiceFactory()) # type: ignore
def get_session_service() -> SessionService: def get_session_service() -> SessionService:
""" """
Retrieves the session service from the service manager. Retrieves the session service from the service manager.

View file

@ -9,6 +9,7 @@ class ServiceType(str, Enum):
AUTH_SERVICE = "auth_service" AUTH_SERVICE = "auth_service"
CACHE_SERVICE = "cache_service" CACHE_SERVICE = "cache_service"
SHARED_COMPONENT_CACHE_SERVICE = "shared_component_cache_service"
SETTINGS_SERVICE = "settings_service" SETTINGS_SERVICE = "settings_service"
DATABASE_SERVICE = "database_service" DATABASE_SERVICE = "database_service"
CHAT_SERVICE = "chat_service" CHAT_SERVICE = "chat_service"

View file

@ -0,0 +1,15 @@
from typing import TYPE_CHECKING
from langflow.services.factory import ServiceFactory
from langflow.services.shared_component_cache.service import SharedComponentCacheService
if TYPE_CHECKING:
from langflow.services.settings.service import SettingsService
class SharedComponentCacheServiceFactory(ServiceFactory):
def __init__(self):
super().__init__(SharedComponentCacheService)
def create(self, settings_service: "SettingsService"):
return SharedComponentCacheService(expiration_time=settings_service.settings.cache_expire)

View file

@ -0,0 +1,9 @@
from langflow.services.cache import ThreadingInMemoryCache
class SharedComponentCacheService(ThreadingInMemoryCache):
"""
A caching service shared across components.
"""
name = "shared_component_cache_service"

1042
uv.lock generated

File diff suppressed because it is too large Load diff