chore: refactor and add components integration tests (#3607)
* improve inegration tests * add fixes * [autofix.ci] apply automated fixes --------- Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
parent
4ee25359a5
commit
96872f3aa5
32 changed files with 523 additions and 131 deletions
|
|
@ -0,0 +1,16 @@
|
|||
import pytest
|
||||
from langflow.schema.message import Message
|
||||
from tests.api_keys import get_openai_api_key
|
||||
from tests.integration.utils import download_flow_from_github, run_json_flow
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.api_key_required
|
||||
async def test_1_0_15_basic_prompting():
|
||||
api_key = get_openai_api_key()
|
||||
json_flow = download_flow_from_github("Basic Prompting (Hello, World)", "1.0.15")
|
||||
json_flow.set_value(json_flow.get_component_by_type("OpenAIModel"), "api_key", api_key)
|
||||
outputs = await run_json_flow(json_flow, run_input="my name is bob, say hello!")
|
||||
assert isinstance(outputs["message"], Message)
|
||||
response = outputs["message"].text.lower()
|
||||
assert "arr" in response or "ahoy" in response
|
||||
|
|
@ -1,118 +1,129 @@
|
|||
import os
|
||||
from typing import List
|
||||
|
||||
from astrapy.db import AstraDB
|
||||
import pytest
|
||||
from integration.utils import MockEmbeddings, check_env_vars, valid_nvidia_vectorize_region
|
||||
|
||||
from langflow.components.embeddings import OpenAIEmbeddingsComponent
|
||||
from langflow.custom import Component
|
||||
from langflow.inputs import StrInput
|
||||
from langflow.template import Output
|
||||
from tests.api_keys import get_astradb_application_token, get_astradb_api_endpoint, get_openai_api_key
|
||||
from tests.integration.utils import ComponentInputHandle
|
||||
from langchain_core.documents import Document
|
||||
|
||||
# from langflow.components.memories.AstraDBMessageReader import AstraDBMessageReaderComponent
|
||||
# from langflow.components.memories.AstraDBMessageWriter import AstraDBMessageWriterComponent
|
||||
|
||||
from langflow.components.vectorstores.AstraDB import AstraVectorStoreComponent
|
||||
from langflow.schema.data import Data
|
||||
from tests.integration.utils import run_single_component
|
||||
|
||||
COLLECTION = "test_basic"
|
||||
BASIC_COLLECTION = "test_basic"
|
||||
SEARCH_COLLECTION = "test_search"
|
||||
# MEMORY_COLLECTION = "test_memory"
|
||||
VECTORIZE_COLLECTION = "test_vectorize"
|
||||
VECTORIZE_COLLECTION_OPENAI = "test_vectorize_openai"
|
||||
VECTORIZE_COLLECTION_OPENAI_WITH_AUTH = "test_vectorize_openai_auth"
|
||||
ALL_COLLECTIONS = [
|
||||
BASIC_COLLECTION,
|
||||
SEARCH_COLLECTION,
|
||||
# MEMORY_COLLECTION,
|
||||
VECTORIZE_COLLECTION,
|
||||
VECTORIZE_COLLECTION_OPENAI,
|
||||
VECTORIZE_COLLECTION_OPENAI_WITH_AUTH,
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def astra_fixture(request):
|
||||
"""
|
||||
Sets up the astra collection and cleans up after
|
||||
"""
|
||||
try:
|
||||
from langchain_astradb import AstraDBVectorStore
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"Could not import langchain Astra DB integration package. Please install it with `pip install langchain-astradb`."
|
||||
)
|
||||
def astradb_client(request):
|
||||
client = AstraDB(api_endpoint=get_astradb_api_endpoint(), token=get_astradb_application_token())
|
||||
yield client
|
||||
for collection in ALL_COLLECTIONS:
|
||||
client.delete_collection(collection)
|
||||
|
||||
store = AstraDBVectorStore(
|
||||
collection_name=request.param,
|
||||
embedding=MockEmbeddings(),
|
||||
api_endpoint=os.getenv("ASTRA_DB_API_ENDPOINT"),
|
||||
token=os.getenv("ASTRA_DB_APPLICATION_TOKEN"),
|
||||
|
||||
@pytest.mark.api_key_required
|
||||
@pytest.mark.asyncio
|
||||
async def test_base(astradb_client: AstraDB):
|
||||
from langflow.components.embeddings import OpenAIEmbeddingsComponent
|
||||
|
||||
application_token = get_astradb_application_token()
|
||||
api_endpoint = get_astradb_api_endpoint()
|
||||
|
||||
results = await run_single_component(
|
||||
AstraVectorStoreComponent,
|
||||
inputs={
|
||||
"token": application_token,
|
||||
"api_endpoint": api_endpoint,
|
||||
"collection_name": BASIC_COLLECTION,
|
||||
"embedding": ComponentInputHandle(
|
||||
clazz=OpenAIEmbeddingsComponent,
|
||||
inputs={"openai_api_key": get_openai_api_key()},
|
||||
output_name="embeddings",
|
||||
),
|
||||
},
|
||||
)
|
||||
from langchain_core.vectorstores import VectorStoreRetriever
|
||||
|
||||
yield
|
||||
|
||||
store.delete_collection()
|
||||
assert isinstance(results["base_retriever"], VectorStoreRetriever)
|
||||
assert results["vector_store"] is not None
|
||||
assert results["search_results"] == []
|
||||
assert astradb_client.collection(BASIC_COLLECTION)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not check_env_vars("ASTRA_DB_APPLICATION_TOKEN", "ASTRA_DB_API_ENDPOINT"),
|
||||
reason="missing astra env vars",
|
||||
)
|
||||
@pytest.mark.parametrize("astra_fixture", [COLLECTION], indirect=True)
|
||||
def test_astra_setup(astra_fixture):
|
||||
application_token = os.getenv("ASTRA_DB_APPLICATION_TOKEN")
|
||||
api_endpoint = os.getenv("ASTRA_DB_API_ENDPOINT")
|
||||
embedding = MockEmbeddings()
|
||||
class TextToData(Component):
|
||||
inputs = [StrInput(name="text_data", is_list=True)]
|
||||
outputs = [Output(name="data", display_name="Data", method="create_data")]
|
||||
|
||||
component = AstraVectorStoreComponent()
|
||||
component.build(
|
||||
token=application_token,
|
||||
api_endpoint=api_endpoint,
|
||||
collection_name=COLLECTION,
|
||||
embedding=embedding,
|
||||
def create_data(self) -> List[Data]:
|
||||
return [Data(text=t) for t in self.text_data]
|
||||
|
||||
|
||||
@pytest.mark.api_key_required
|
||||
@pytest.mark.asyncio
|
||||
async def test_astra_embeds_and_search():
|
||||
application_token = get_astradb_application_token()
|
||||
api_endpoint = get_astradb_api_endpoint()
|
||||
|
||||
results = await run_single_component(
|
||||
AstraVectorStoreComponent,
|
||||
inputs={
|
||||
"token": application_token,
|
||||
"api_endpoint": api_endpoint,
|
||||
"collection_name": BASIC_COLLECTION,
|
||||
"number_of_results": 1,
|
||||
"search_input": "test1",
|
||||
"ingest_data": ComponentInputHandle(
|
||||
clazz=TextToData, inputs={"text_data": ["test1", "test2"]}, output_name="data"
|
||||
),
|
||||
"embedding": ComponentInputHandle(
|
||||
clazz=OpenAIEmbeddingsComponent,
|
||||
inputs={"openai_api_key": get_openai_api_key()},
|
||||
output_name="embeddings",
|
||||
),
|
||||
},
|
||||
)
|
||||
component.build_vector_store()
|
||||
assert len(results["search_results"]) == 1
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not check_env_vars("ASTRA_DB_APPLICATION_TOKEN", "ASTRA_DB_API_ENDPOINT"),
|
||||
reason="missing astra env vars",
|
||||
)
|
||||
@pytest.mark.parametrize("astra_fixture", [SEARCH_COLLECTION], indirect=True)
|
||||
def test_astra_embeds_and_search(astra_fixture):
|
||||
application_token = os.getenv("ASTRA_DB_APPLICATION_TOKEN")
|
||||
api_endpoint = os.getenv("ASTRA_DB_API_ENDPOINT")
|
||||
embedding = MockEmbeddings()
|
||||
|
||||
documents = [Document(page_content="test1"), Document(page_content="test2")]
|
||||
records = [Data.from_document(d) for d in documents]
|
||||
|
||||
component = AstraVectorStoreComponent()
|
||||
component.build(
|
||||
token=application_token,
|
||||
api_endpoint=api_endpoint,
|
||||
collection_name=SEARCH_COLLECTION,
|
||||
embedding=embedding,
|
||||
ingest_data=records,
|
||||
search_input="test1",
|
||||
number_of_results=1,
|
||||
)
|
||||
component.build_vector_store()
|
||||
records = component.search_documents()
|
||||
|
||||
assert len(records) == 1
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not check_env_vars("ASTRA_DB_APPLICATION_TOKEN", "ASTRA_DB_API_ENDPOINT")
|
||||
or not valid_nvidia_vectorize_region(os.getenv("ASTRA_DB_API_ENDPOINT")),
|
||||
reason="missing env vars or invalid region for nvidia vectorize",
|
||||
)
|
||||
@pytest.mark.api_key_required
|
||||
def test_astra_vectorize():
|
||||
from langchain_astradb import AstraDBVectorStore, CollectionVectorServiceOptions
|
||||
|
||||
from langflow.components.embeddings.AstraVectorize import AstraVectorizeComponent
|
||||
|
||||
application_token = get_astradb_application_token()
|
||||
api_endpoint = get_astradb_api_endpoint()
|
||||
|
||||
store = None
|
||||
try:
|
||||
options = {"provider": "nvidia", "modelName": "NV-Embed-QA"}
|
||||
store = AstraDBVectorStore(
|
||||
collection_name=VECTORIZE_COLLECTION,
|
||||
api_endpoint=os.getenv("ASTRA_DB_API_ENDPOINT"),
|
||||
token=os.getenv("ASTRA_DB_APPLICATION_TOKEN"),
|
||||
api_endpoint=api_endpoint,
|
||||
token=application_token,
|
||||
collection_vector_service_options=CollectionVectorServiceOptions.from_dict(options),
|
||||
)
|
||||
|
||||
application_token = os.getenv("ASTRA_DB_APPLICATION_TOKEN")
|
||||
api_endpoint = os.getenv("ASTRA_DB_API_ENDPOINT")
|
||||
|
||||
documents = [Document(page_content="test1"), Document(page_content="test2")]
|
||||
records = [Data.from_document(d) for d in documents]
|
||||
|
||||
|
|
@ -139,20 +150,18 @@ def test_astra_vectorize():
|
|||
store.delete_collection()
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not check_env_vars("ASTRA_DB_APPLICATION_TOKEN", "ASTRA_DB_API_ENDPOINT", "OPENAI_API_KEY"),
|
||||
reason="missing env vars",
|
||||
)
|
||||
@pytest.mark.api_key_required
|
||||
def test_astra_vectorize_with_provider_api_key():
|
||||
"""tests vectorize using an openai api key"""
|
||||
from langchain_astradb import AstraDBVectorStore, CollectionVectorServiceOptions
|
||||
|
||||
from langflow.components.embeddings.AstraVectorize import AstraVectorizeComponent
|
||||
|
||||
application_token = get_astradb_application_token()
|
||||
api_endpoint = get_astradb_api_endpoint()
|
||||
|
||||
store = None
|
||||
try:
|
||||
application_token = os.getenv("ASTRA_DB_APPLICATION_TOKEN")
|
||||
api_endpoint = os.getenv("ASTRA_DB_API_ENDPOINT")
|
||||
options = {"provider": "openai", "modelName": "text-embedding-3-small", "parameters": {}, "authentication": {}}
|
||||
store = AstraDBVectorStore(
|
||||
collection_name=VECTORIZE_COLLECTION_OPENAI,
|
||||
|
|
@ -188,10 +197,7 @@ def test_astra_vectorize_with_provider_api_key():
|
|||
store.delete_collection()
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not check_env_vars("ASTRA_DB_APPLICATION_TOKEN", "ASTRA_DB_API_ENDPOINT"),
|
||||
reason="missing env vars",
|
||||
)
|
||||
@pytest.mark.api_key_required
|
||||
def test_astra_vectorize_passes_authentication():
|
||||
"""tests vectorize using the authentication parameter"""
|
||||
from langchain_astradb import AstraDBVectorStore, CollectionVectorServiceOptions
|
||||
|
|
@ -200,8 +206,8 @@ def test_astra_vectorize_passes_authentication():
|
|||
|
||||
store = None
|
||||
try:
|
||||
application_token = os.getenv("ASTRA_DB_APPLICATION_TOKEN")
|
||||
api_endpoint = os.getenv("ASTRA_DB_API_ENDPOINT")
|
||||
application_token = get_astradb_application_token()
|
||||
api_endpoint = get_astradb_api_endpoint()
|
||||
options = {
|
||||
"provider": "openai",
|
||||
"modelName": "text-embedding-3-small",
|
||||
|
|
@ -0,0 +1,55 @@
|
|||
from langflow.memory import get_messages
|
||||
from langflow.schema.message import Message
|
||||
from tests.integration.utils import run_single_component
|
||||
|
||||
from langflow.components.inputs import ChatInput
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default():
|
||||
outputs = await run_single_component(ChatInput, run_input="hello")
|
||||
assert isinstance(outputs["message"], Message)
|
||||
assert outputs["message"].text == "hello"
|
||||
assert outputs["message"].sender == "User"
|
||||
assert outputs["message"].sender_name == "User"
|
||||
|
||||
outputs = await run_single_component(ChatInput, run_input="")
|
||||
assert isinstance(outputs["message"], Message)
|
||||
assert outputs["message"].text == ""
|
||||
assert outputs["message"].sender == "User"
|
||||
assert outputs["message"].sender_name == "User"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sender():
|
||||
outputs = await run_single_component(
|
||||
ChatInput, inputs={"sender": "Machine", "sender_name": "AI"}, run_input="hello"
|
||||
)
|
||||
assert isinstance(outputs["message"], Message)
|
||||
assert outputs["message"].text == "hello"
|
||||
assert outputs["message"].sender == "Machine"
|
||||
assert outputs["message"].sender_name == "AI"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_do_not_store_messages():
|
||||
session_id = "test-session-id"
|
||||
outputs = await run_single_component(
|
||||
ChatInput, inputs={"should_store_message": True}, run_input="hello", session_id=session_id
|
||||
)
|
||||
assert isinstance(outputs["message"], Message)
|
||||
assert outputs["message"].text == "hello"
|
||||
assert outputs["message"].session_id == session_id
|
||||
|
||||
assert len(get_messages(session_id=session_id)) == 1
|
||||
|
||||
session_id = "test-session-id-another"
|
||||
outputs = await run_single_component(
|
||||
ChatInput, inputs={"should_store_message": False}, run_input="hello", session_id=session_id
|
||||
)
|
||||
assert isinstance(outputs["message"], Message)
|
||||
assert outputs["message"].text == "hello"
|
||||
assert outputs["message"].session_id == session_id
|
||||
|
||||
assert len(get_messages(session_id=session_id)) == 0
|
||||
|
|
@ -0,0 +1,45 @@
|
|||
from langflow.components.outputs import ChatOutput
|
||||
from langflow.memory import get_messages
|
||||
from langflow.schema.message import Message
|
||||
from tests.integration.utils import run_single_component
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_string():
|
||||
outputs = await run_single_component(ChatOutput, inputs={"input_value": "hello"})
|
||||
assert isinstance(outputs["message"], Message)
|
||||
assert outputs["message"].text == "hello"
|
||||
assert outputs["message"].sender == "Machine"
|
||||
assert outputs["message"].sender_name == "AI"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message():
|
||||
outputs = await run_single_component(ChatOutput, inputs={"input_value": Message(text="hello")})
|
||||
assert isinstance(outputs["message"], Message)
|
||||
assert outputs["message"].text == "hello"
|
||||
assert outputs["message"].sender == "Machine"
|
||||
assert outputs["message"].sender_name == "AI"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_do_not_store_message():
|
||||
session_id = "test-session-id"
|
||||
outputs = await run_single_component(
|
||||
ChatOutput, inputs={"input_value": "hello", "should_store_message": True}, session_id=session_id
|
||||
)
|
||||
assert isinstance(outputs["message"], Message)
|
||||
assert outputs["message"].text == "hello"
|
||||
|
||||
assert len(get_messages(session_id=session_id)) == 1
|
||||
session_id = "test-session-id-another"
|
||||
|
||||
outputs = await run_single_component(
|
||||
ChatOutput, inputs={"input_value": "hello", "should_store_message": False}, session_id=session_id
|
||||
)
|
||||
assert isinstance(outputs["message"], Message)
|
||||
assert outputs["message"].text == "hello"
|
||||
|
||||
assert len(get_messages(session_id=session_id)) == 0
|
||||
|
|
@ -0,0 +1,13 @@
|
|||
from langflow.components.prompts import PromptComponent
|
||||
from langflow.schema.message import Message
|
||||
from tests.integration.utils import run_single_component
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test():
|
||||
outputs = await run_single_component(PromptComponent, inputs={"template": "test {var1}", "var1": "from the var"})
|
||||
print(outputs)
|
||||
assert isinstance(outputs["prompt"], Message)
|
||||
assert outputs["prompt"].text == "test from the var"
|
||||
22
src/backend/tests/integration/flows/test_basic_prompting.py
Normal file
22
src/backend/tests/integration/flows/test_basic_prompting.py
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
import pytest
|
||||
|
||||
from langflow.components.inputs import ChatInput
|
||||
from langflow.components.outputs import ChatOutput
|
||||
from langflow.components.prompts import PromptComponent
|
||||
from langflow.graph import Graph
|
||||
from langflow.schema.message import Message
|
||||
from tests.integration.utils import run_flow
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_simple_no_llm():
|
||||
graph = Graph()
|
||||
input = graph.add_component(ChatInput())
|
||||
output = graph.add_component(ChatOutput())
|
||||
component = PromptComponent(template="This is the message: {var1}", var1="")
|
||||
prompt = graph.add_component(component)
|
||||
graph.add_component_edge(input, ("message", "var1"), prompt)
|
||||
graph.add_component_edge(prompt, ("prompt", "input_value"), output)
|
||||
outputs = await run_flow(graph, run_input="hello!")
|
||||
assert isinstance(outputs["message"], Message)
|
||||
assert outputs["message"].text == "This is the message: hello!"
|
||||
|
|
@ -1,7 +1,16 @@
|
|||
import dataclasses
|
||||
import os
|
||||
import uuid
|
||||
from typing import Optional, Any
|
||||
|
||||
from astrapy.admin import parse_api_endpoint
|
||||
|
||||
from langflow.api.v1.schemas import InputValueRequest
|
||||
from langflow.custom import Component
|
||||
from langflow.field_typing import Embeddings
|
||||
from langflow.graph import Graph
|
||||
from langflow.processing.process import run_graph_internal
|
||||
import requests
|
||||
|
||||
|
||||
def check_env_vars(*vars):
|
||||
|
|
@ -49,3 +58,115 @@ class MockEmbeddings(Embeddings):
|
|||
def embed_query(self, text: str) -> list[float]:
|
||||
self.embedded_query = text
|
||||
return self.mock_embedding(text)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class JSONFlow:
|
||||
json: dict
|
||||
|
||||
def get_components_by_type(self, component_type):
|
||||
result = []
|
||||
for node in self.json["data"]["nodes"]:
|
||||
if node["data"]["type"] == component_type:
|
||||
result.append(node["id"])
|
||||
if not result:
|
||||
raise ValueError(
|
||||
f"Component of type {component_type} not found, available types: {', '.join(set(node['data']['type'] for node in self.json['data']['nodes']))}"
|
||||
)
|
||||
return result
|
||||
|
||||
def get_component_by_type(self, component_type):
|
||||
components = self.get_components_by_type(component_type)
|
||||
if len(components) > 1:
|
||||
raise ValueError(f"Multiple components of type {component_type} found")
|
||||
return components[0]
|
||||
|
||||
def set_value(self, component_id, key, value):
|
||||
done = False
|
||||
for node in self.json["data"]["nodes"]:
|
||||
if node["id"] == component_id:
|
||||
if key not in node["data"]["node"]["template"]:
|
||||
raise ValueError(f"Component {component_id} does not have input {key}")
|
||||
node["data"]["node"]["template"][key]["value"] = value
|
||||
node["data"]["node"]["template"][key]["load_from_db"] = False
|
||||
done = True
|
||||
break
|
||||
if not done:
|
||||
raise ValueError(f"Component {component_id} not found")
|
||||
|
||||
|
||||
def download_flow_from_github(name: str, version: str) -> JSONFlow:
|
||||
response = requests.get(
|
||||
f"https://raw.githubusercontent.com/langflow-ai/langflow/v{version}/src/backend/base/langflow/initial_setup/starter_projects/{name}.json"
|
||||
)
|
||||
response.raise_for_status()
|
||||
as_json = response.json()
|
||||
return JSONFlow(json=as_json)
|
||||
|
||||
|
||||
async def run_json_flow(
|
||||
json_flow: JSONFlow, run_input: Optional[Any] = None, session_id: Optional[str] = None
|
||||
) -> dict[str, Any]:
|
||||
graph = Graph.from_payload(json_flow.json)
|
||||
return await run_flow(graph, run_input, session_id)
|
||||
|
||||
|
||||
async def run_flow(graph: Graph, run_input: Optional[Any] = None, session_id: Optional[str] = None) -> dict[str, Any]:
|
||||
graph.prepare()
|
||||
if run_input:
|
||||
graph_run_inputs = [InputValueRequest(input_value=run_input, type="chat")]
|
||||
else:
|
||||
graph_run_inputs = []
|
||||
|
||||
flow_id = str(uuid.uuid4())
|
||||
|
||||
results, _ = await run_graph_internal(graph, flow_id, session_id=session_id, inputs=graph_run_inputs)
|
||||
outputs = {}
|
||||
for r in results:
|
||||
for out in r.outputs:
|
||||
outputs |= out.results
|
||||
return outputs
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class ComponentInputHandle:
|
||||
clazz: type
|
||||
inputs: dict
|
||||
output_name: str
|
||||
|
||||
|
||||
async def run_single_component(
|
||||
clazz: type, inputs: dict = None, run_input: Optional[Any] = None, session_id: Optional[str] = None
|
||||
) -> dict[str, Any]:
|
||||
user_id = str(uuid.uuid4())
|
||||
flow_id = str(uuid.uuid4())
|
||||
graph = Graph(user_id=user_id, flow_id=flow_id)
|
||||
|
||||
def _add_component(clazz: type, inputs: Optional[dict] = None) -> str:
|
||||
raw_inputs = {}
|
||||
if inputs:
|
||||
for key, value in inputs.items():
|
||||
if not isinstance(value, ComponentInputHandle):
|
||||
raw_inputs[key] = value
|
||||
if isinstance(value, Component):
|
||||
raise ValueError("Component inputs must be wrapped in ComponentInputHandle")
|
||||
component = clazz(**raw_inputs, _user_id=user_id)
|
||||
component_id = graph.add_component(component)
|
||||
if inputs:
|
||||
for input_name, handle in inputs.items():
|
||||
if isinstance(handle, ComponentInputHandle):
|
||||
handle_component_id = _add_component(handle.clazz, handle.inputs)
|
||||
graph.add_component_edge(handle_component_id, (handle.output_name, input_name), component_id)
|
||||
return component_id
|
||||
|
||||
component_id = _add_component(clazz, inputs)
|
||||
graph.prepare()
|
||||
if run_input:
|
||||
graph_run_inputs = [InputValueRequest(input_value=run_input, type="chat")]
|
||||
else:
|
||||
graph_run_inputs = []
|
||||
|
||||
_, _ = await run_graph_internal(
|
||||
graph, flow_id, session_id=session_id, inputs=graph_run_inputs, outputs=[component_id]
|
||||
)
|
||||
return graph.get_vertex(component_id)._built_object
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue