mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-25 13:11: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
|
||||
)
|
||||
|
||||
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]
|
||||
)
|
||||
flow.side_effect = lambda service_name: mock_prompt_service if service_name == "prompt-request" else flow_response if service_name == "response" else AsyncMock()
|
||||
|
||||
# Act
|
||||
await custom_processor.on_message(msg, consumer, flow)
|
||||
|
|
@ -374,7 +377,7 @@ class TestNLPQueryServiceIntegration:
|
|||
assert custom_processor.graphql_generation_template == "custom-graphql-generator"
|
||||
|
||||
# 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
|
||||
async def test_large_schema_set_integration(self, integration_processor):
|
||||
|
|
@ -466,30 +469,36 @@ class TestNLPQueryServiceIntegration:
|
|||
messages.append(msg)
|
||||
flows.append(flow)
|
||||
|
||||
# Mock responses for all requests
|
||||
mock_responses = []
|
||||
for i in range(10): # 2 calls per request (phase1 + phase2)
|
||||
if i % 2 == 0: # Phase 1 responses
|
||||
mock_responses.append(PromptResponse(
|
||||
text=json.dumps(["customers"]),
|
||||
error=None
|
||||
))
|
||||
else: # Phase 2 responses
|
||||
mock_responses.append(PromptResponse(
|
||||
text=json.dumps({
|
||||
"query": f"query {{ customers {{ id name }} }}",
|
||||
"variables": {},
|
||||
"confidence": 0.9
|
||||
}),
|
||||
error=None
|
||||
))
|
||||
|
||||
# Mock the flow context to return prompt service responses
|
||||
prompt_service = AsyncMock()
|
||||
prompt_service.request = AsyncMock(
|
||||
side_effect=mock_responses
|
||||
)
|
||||
flow.side_effect = lambda service_name: prompt_service if service_name == "prompt-request" else flow_response if service_name == "response" else AsyncMock()
|
||||
# Mock responses for all requests - create individual prompt services for each flow
|
||||
prompt_services = []
|
||||
for i in range(5): # 5 concurrent requests
|
||||
phase1_response = PromptResponse(
|
||||
text=json.dumps(["customers"]),
|
||||
error=None
|
||||
)
|
||||
phase2_response = PromptResponse(
|
||||
text=json.dumps({
|
||||
"query": f"query {{ customers {{ id name }} }}",
|
||||
"variables": {},
|
||||
"confidence": 0.9
|
||||
}),
|
||||
error=None
|
||||
)
|
||||
|
||||
# Create a prompt service for this request
|
||||
prompt_service = AsyncMock()
|
||||
prompt_service.request = AsyncMock(
|
||||
side_effect=[phase1_response, phase2_response]
|
||||
)
|
||||
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
|
||||
import asyncio
|
||||
|
|
@ -503,7 +512,8 @@ class TestNLPQueryServiceIntegration:
|
|||
await asyncio.gather(*tasks)
|
||||
|
||||
# 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:
|
||||
flow.return_value.send.assert_called_once()
|
||||
|
||||
|
|
|
|||
|
|
@ -562,34 +562,50 @@ class TestStructuredQueryServiceIntegration:
|
|||
messages.append(msg)
|
||||
flows.append(flow)
|
||||
|
||||
# Mock responses for all requests (6 total: 3 NLP + 3 Objects)
|
||||
mock_responses = []
|
||||
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={}
|
||||
))
|
||||
# Set up individual flow routing for each concurrent request
|
||||
service_call_count = 0
|
||||
|
||||
call_count = 0
|
||||
def mock_client_side_effect(name):
|
||||
nonlocal call_count
|
||||
client = AsyncMock()
|
||||
client.request.return_value = mock_responses[call_count]
|
||||
call_count += 1
|
||||
return client
|
||||
|
||||
integration_processor.client.side_effect = mock_client_side_effect
|
||||
for i in range(3): # 3 concurrent requests
|
||||
# Create NLP and Objects responses for this request
|
||||
nlp_response = QuestionToStructuredQueryResponse(
|
||||
error=None,
|
||||
graphql_query=f'query {{ test_{i} {{ id }} }}',
|
||||
variables={},
|
||||
detected_schemas=[f"test_{i}"],
|
||||
confidence=0.9
|
||||
)
|
||||
|
||||
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
|
||||
import asyncio
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue