ref: Add ALL ruff rules for tests (#4183)

Add ALL ruff rules for tests
This commit is contained in:
Christophe Bornet 2024-10-19 22:41:37 +02:00 • committed by GitHub
commit f96f2eaf8a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
67 changed files with 421 additions and 361 deletions

View file

@ -43,6 +43,7 @@ AI: """
return graph
@pytest.mark.usefixtures("client")
def test_memory_chatbot(memory_chatbot_graph):
# Now we run step by step
expected_order = deque(["chat_input", "chat_memory", "prompt", "openai", "chat_output"])

View file

@ -34,7 +34,7 @@ def ingestion_graph():
embedding=openai_embeddings.build_embeddings,
ingest_data=text_splitter.split_text,
api_endpoint="https://astra.example.com",
token="token",
token="token", # noqa: S106
)
vector_store.set_on_output(name="vector_store", value="mock_vector_store", cache=True)
vector_store.set_on_output(name="base_retriever", value="mock_retriever", cache=True)
@ -53,7 +53,7 @@ def rag_graph():
rag_vector_store.set(
search_input=chat_input.message_response,
api_endpoint="https://astra.example.com",
token="token",
token="token", # noqa: S106
embedding=openai_embeddings.build_embeddings,
)
# Mock search_documents
@ -110,9 +110,7 @@ def test_vector_store_rag(ingestion_graph, rag_graph):
"openai-embeddings-124",
]
for ids, graph, len_results in [(ingestion_ids, ingestion_graph, 5), (rag_ids, rag_graph, 8)]:
results = []
for result in graph.start():
results.append(result)
results = list(graph.start())
assert len(results) == len_results
vids = [result.vertex.id for result in results if hasattr(result, "vertex")]
@ -217,12 +215,14 @@ def test_vector_store_rag_add(ingestion_graph: Graph, rag_graph: Graph):
rag_graph_copy = copy.deepcopy(rag_graph)
ingestion_graph_copy += rag_graph_copy
assert (
len(ingestion_graph_copy.vertices) == len(ingestion_graph.vertices) + len(rag_graph.vertices)
), f"Vertices mismatch: {len(ingestion_graph_copy.vertices)} != {len(ingestion_graph.vertices)} + {len(rag_graph.vertices)}"
assert len(ingestion_graph_copy.edges) == len(ingestion_graph.edges) + len(
rag_graph.edges
), f"Edges mismatch: {len(ingestion_graph_copy.edges)} != {len(ingestion_graph.edges)} + {len(rag_graph.edges)}"
assert len(ingestion_graph_copy.vertices) == len(ingestion_graph.vertices) + len(rag_graph.vertices), (
f"Vertices mismatch: {len(ingestion_graph_copy.vertices)} "
f"!= {len(ingestion_graph.vertices)} + {len(rag_graph.vertices)}"
)
assert len(ingestion_graph_copy.edges) == len(ingestion_graph.edges) + len(rag_graph.edges), (
f"Edges mismatch: {len(ingestion_graph_copy.edges)} "
f"!= {len(ingestion_graph.edges)} + {len(rag_graph.edges)}"
)
combined_graph_dump = ingestion_graph_copy.dump(
name="Combined Graph", description="Graph for data ingestion and RAG", endpoint_name="combined"