mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-26 13:41:02 +02:00
Fix tests
This commit is contained in:
parent
86f0d41553
commit
49417eccab
2 changed files with 80 additions and 54 deletions
|
|
@ -362,9 +362,12 @@ class TestNLPQueryServiceIntegration:
|
||||||
error=None
|
error=None
|
||||||
)
|
)
|
||||||
|
|
||||||
custom_processor.client.return_value.request = AsyncMock(
|
# Mock flow context to return prompt service responses
|
||||||
|
mock_prompt_service = AsyncMock()
|
||||||
|
mock_prompt_service.request = AsyncMock(
|
||||||
side_effect=[phase1_response, phase2_response]
|
side_effect=[phase1_response, phase2_response]
|
||||||
)
|
)
|
||||||
|
flow.side_effect = lambda service_name: mock_prompt_service if service_name == "prompt-request" else flow_response if service_name == "response" else AsyncMock()
|
||||||
|
|
||||||
# Act
|
# Act
|
||||||
await custom_processor.on_message(msg, consumer, flow)
|
await custom_processor.on_message(msg, consumer, flow)
|
||||||
|
|
@ -374,7 +377,7 @@ class TestNLPQueryServiceIntegration:
|
||||||
assert custom_processor.graphql_generation_template == "custom-graphql-generator"
|
assert custom_processor.graphql_generation_template == "custom-graphql-generator"
|
||||||
|
|
||||||
# Verify the calls were made
|
# Verify the calls were made
|
||||||
assert custom_processor.client.return_value.request.call_count == 2
|
assert mock_prompt_service.request.call_count == 2
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_large_schema_set_integration(self, integration_processor):
|
async def test_large_schema_set_integration(self, integration_processor):
|
||||||
|
|
@ -466,30 +469,36 @@ class TestNLPQueryServiceIntegration:
|
||||||
messages.append(msg)
|
messages.append(msg)
|
||||||
flows.append(flow)
|
flows.append(flow)
|
||||||
|
|
||||||
# Mock responses for all requests
|
# Mock responses for all requests - create individual prompt services for each flow
|
||||||
mock_responses = []
|
prompt_services = []
|
||||||
for i in range(10): # 2 calls per request (phase1 + phase2)
|
for i in range(5): # 5 concurrent requests
|
||||||
if i % 2 == 0: # Phase 1 responses
|
phase1_response = PromptResponse(
|
||||||
mock_responses.append(PromptResponse(
|
text=json.dumps(["customers"]),
|
||||||
text=json.dumps(["customers"]),
|
error=None
|
||||||
error=None
|
)
|
||||||
))
|
phase2_response = PromptResponse(
|
||||||
else: # Phase 2 responses
|
text=json.dumps({
|
||||||
mock_responses.append(PromptResponse(
|
"query": f"query {{ customers {{ id name }} }}",
|
||||||
text=json.dumps({
|
"variables": {},
|
||||||
"query": f"query {{ customers {{ id name }} }}",
|
"confidence": 0.9
|
||||||
"variables": {},
|
}),
|
||||||
"confidence": 0.9
|
error=None
|
||||||
}),
|
)
|
||||||
error=None
|
|
||||||
))
|
|
||||||
|
|
||||||
# Mock the flow context to return prompt service responses
|
# Create a prompt service for this request
|
||||||
prompt_service = AsyncMock()
|
prompt_service = AsyncMock()
|
||||||
prompt_service.request = AsyncMock(
|
prompt_service.request = AsyncMock(
|
||||||
side_effect=mock_responses
|
side_effect=[phase1_response, phase2_response]
|
||||||
)
|
)
|
||||||
flow.side_effect = lambda service_name: prompt_service if service_name == "prompt-request" else flow_response if service_name == "response" else AsyncMock()
|
prompt_services.append(prompt_service)
|
||||||
|
|
||||||
|
# Set up the flow for this request
|
||||||
|
flow_response = flows[i].return_value
|
||||||
|
flows[i].side_effect = lambda service_name, ps=prompt_service, fr=flow_response: (
|
||||||
|
ps if service_name == "prompt-request" else
|
||||||
|
fr if service_name == "response" else
|
||||||
|
AsyncMock()
|
||||||
|
)
|
||||||
|
|
||||||
# Act - Process all messages concurrently
|
# Act - Process all messages concurrently
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
@ -503,7 +512,8 @@ class TestNLPQueryServiceIntegration:
|
||||||
await asyncio.gather(*tasks)
|
await asyncio.gather(*tasks)
|
||||||
|
|
||||||
# Assert - All requests should be processed
|
# Assert - All requests should be processed
|
||||||
assert prompt_service.request.call_count == 10
|
total_calls = sum(ps.request.call_count for ps in prompt_services)
|
||||||
|
assert total_calls == 10 # 2 calls per request (phase1 + phase2)
|
||||||
for flow in flows:
|
for flow in flows:
|
||||||
flow.return_value.send.assert_called_once()
|
flow.return_value.send.assert_called_once()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -562,34 +562,50 @@ class TestStructuredQueryServiceIntegration:
|
||||||
messages.append(msg)
|
messages.append(msg)
|
||||||
flows.append(flow)
|
flows.append(flow)
|
||||||
|
|
||||||
# Mock responses for all requests (6 total: 3 NLP + 3 Objects)
|
# Set up individual flow routing for each concurrent request
|
||||||
mock_responses = []
|
service_call_count = 0
|
||||||
for i in range(6):
|
|
||||||
if i % 2 == 0: # NLP responses
|
|
||||||
mock_responses.append(QuestionToStructuredQueryResponse(
|
|
||||||
error=None,
|
|
||||||
graphql_query=f'query {{ test_{i//2} {{ id }} }}',
|
|
||||||
variables={},
|
|
||||||
detected_schemas=[f"test_{i//2}"],
|
|
||||||
confidence=0.9
|
|
||||||
))
|
|
||||||
else: # Objects responses
|
|
||||||
mock_responses.append(ObjectsQueryResponse(
|
|
||||||
error=None,
|
|
||||||
data=f'{{"test_{i//2}": [{{"id": "{i//2}"}}]}}',
|
|
||||||
errors=None,
|
|
||||||
extensions={}
|
|
||||||
))
|
|
||||||
|
|
||||||
call_count = 0
|
for i in range(3): # 3 concurrent requests
|
||||||
def mock_client_side_effect(name):
|
# Create NLP and Objects responses for this request
|
||||||
nonlocal call_count
|
nlp_response = QuestionToStructuredQueryResponse(
|
||||||
client = AsyncMock()
|
error=None,
|
||||||
client.request.return_value = mock_responses[call_count]
|
graphql_query=f'query {{ test_{i} {{ id }} }}',
|
||||||
call_count += 1
|
variables={},
|
||||||
return client
|
detected_schemas=[f"test_{i}"],
|
||||||
|
confidence=0.9
|
||||||
|
)
|
||||||
|
|
||||||
integration_processor.client.side_effect = mock_client_side_effect
|
objects_response = ObjectsQueryResponse(
|
||||||
|
error=None,
|
||||||
|
data=f'{{"test_{i}": [{{"id": "{i}"}}]}}',
|
||||||
|
errors=None,
|
||||||
|
extensions={}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Create mock services for this request
|
||||||
|
mock_nlp_client = AsyncMock()
|
||||||
|
mock_nlp_client.request.return_value = nlp_response
|
||||||
|
|
||||||
|
mock_objects_client = AsyncMock()
|
||||||
|
mock_objects_client.request.return_value = objects_response
|
||||||
|
|
||||||
|
# Set up flow routing for this specific request
|
||||||
|
flow_response = flows[i].return_value
|
||||||
|
def create_flow_router(nlp_client, objects_client, response_producer):
|
||||||
|
def flow_router(service_name):
|
||||||
|
nonlocal service_call_count
|
||||||
|
service_call_count += 1
|
||||||
|
if service_name == "nlp-query-request":
|
||||||
|
return nlp_client
|
||||||
|
elif service_name == "objects-query-request":
|
||||||
|
return objects_client
|
||||||
|
elif service_name == "response":
|
||||||
|
return response_producer
|
||||||
|
else:
|
||||||
|
return AsyncMock()
|
||||||
|
return flow_router
|
||||||
|
|
||||||
|
flows[i].side_effect = create_flow_router(mock_nlp_client, mock_objects_client, flow_response)
|
||||||
|
|
||||||
# Act - Process all messages concurrently
|
# Act - Process all messages concurrently
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue