feat: update output attribute handling (#3329)
* refactor: improve function structure in _build_results * fix: remove function call that overrides Component.outputs * refactor: Add handling for outputs in __getattr__ method.
This commit is contained in:
parent
93537933c3
commit
bb923cf9e1
4 changed files with 126 additions and 116 deletions
|
|
@ -27,7 +27,7 @@ def ingestion_graph():
|
|||
# Ingestion Graph
|
||||
file_component = FileComponent(_id="file-123")
|
||||
file_component.set(path="test.txt")
|
||||
file_component.set_output_value("data", Data(text="This is a test file."))
|
||||
file_component.set_on_output("data", value=Data(text="This is a test file."))
|
||||
text_splitter = SplitTextComponent(_id="text-splitter-123")
|
||||
text_splitter.set(data_inputs=file_component.load_file)
|
||||
openai_embeddings = OpenAIEmbeddingsComponent(_id="openai-embeddings-123")
|
||||
|
|
@ -41,10 +41,9 @@ def ingestion_graph():
|
|||
api_endpoint="https://astra.example.com",
|
||||
token="token",
|
||||
)
|
||||
vector_store.set_output_value("vector_store", "mock_vector_store")
|
||||
vector_store.set_output_value("base_retriever", "mock_retriever")
|
||||
vector_store.set_output_value("search_results", [Data(text="This is a test file.")])
|
||||
|
||||
vector_store.set_on_output("vector_store", value="mock_vector_store")
|
||||
vector_store.set_on_output("base_retriever", value="mock_retriever")
|
||||
vector_store.set_on_output("search_results", value=[Data(text="This is a test file.")])
|
||||
ingestion_graph = Graph(file_component, vector_store)
|
||||
return ingestion_graph
|
||||
|
||||
|
|
@ -63,10 +62,15 @@ def rag_graph():
|
|||
embedding=openai_embeddings.build_embeddings,
|
||||
)
|
||||
# Mock search_documents
|
||||
rag_vector_store.get_output("search_results").value = [
|
||||
Data(data={"text": "Hello, world!"}),
|
||||
Data(data={"text": "Goodbye, world!"}),
|
||||
]
|
||||
rag_vector_store.set_on_output(
|
||||
"search_results",
|
||||
value=[
|
||||
Data(data={"text": "Hello, world!"}),
|
||||
Data(data={"text": "Goodbye, world!"}),
|
||||
],
|
||||
)
|
||||
rag_vector_store.set_on_output("base_retriever", value="mock_retriever")
|
||||
rag_vector_store.set_on_output("vector_store", value="mock_vector_store")
|
||||
parse_data = ParseDataComponent(_id="parse-data-123")
|
||||
parse_data.set(data=rag_vector_store.search_documents)
|
||||
prompt_component = PromptComponent(_id="prompt-123")
|
||||
|
|
@ -82,7 +86,7 @@ def rag_graph():
|
|||
|
||||
openai_component = OpenAIModelComponent(_id="openai-123")
|
||||
openai_component.set(api_key="sk-123", openai_api_base="https://api.openai.com/v1")
|
||||
openai_component.set_output_value("text_output", "Hello, world!")
|
||||
openai_component.set_on_output("text_output", value="Hello, world!")
|
||||
openai_component.set(input_value=prompt_component.build_prompt)
|
||||
|
||||
chat_output = ChatOutput(_id="chatoutput-123")
|
||||
|
|
@ -112,7 +116,7 @@ def test_vector_store_rag(ingestion_graph, rag_graph):
|
|||
]
|
||||
for ids, graph, len_results in zip([ingestion_ids, rag_ids], [ingestion_graph, rag_graph], [5, 8]):
|
||||
results = []
|
||||
for result in graph.start():
|
||||
for result in graph.start(config={"output": {"cache": True}}):
|
||||
results.append(result)
|
||||
|
||||
assert len(results) == len_results
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue