| 1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798 |
- import unittest
- from io import StringIO
- from app import GatewayApp
- from public_gateway import PublicGatewayApp
- from tools.query_customer_list import QueryCustomerListTool
- class RecordingApiClient:
- def __init__(self):
- self.last_call = None
- def call_tool(self, tool_code, route_path, payload, request_id):
- self.last_call = {
- 'tool_code': tool_code, 'route_path': route_path,
- 'payload': payload, 'request_id': request_id,
- }
- return {'code': 'MCP_0000'}
- def list_enabled_tools(self, request_id=''):
- return {'code': 'MCP_0000', 'data': {'tool_codes': ['query_customer_list']}}
- class CustomerQueryToolTest(unittest.TestCase):
- def test_schema_is_closed_and_uses_strict_positive_integers(self):
- metadata = QueryCustomerListTool().metadata()
- schema = metadata['input_schema']
- self.assertIn('不得合并或缩写', metadata['description'])
- self.assertIn('完整逐条展示17个字段', metadata['description'])
- self.assertFalse(schema['additionalProperties'])
- self.assertEqual([], schema['required'])
- for field in ('customer_id', 'department_id', 'sales_id', 'merchandiser_id'):
- self.assertEqual({'type': 'integer', 'minimum': 1}, schema['properties'][field])
- for field, default in (('page', 1), ('limit', 20)):
- self.assertEqual(1, schema['properties'][field]['minimum'])
- self.assertEqual(100, schema['properties'][field]['maximum'])
- self.assertEqual(default, schema['properties'][field]['default'])
- def test_call_forwards_all_filters(self):
- client = RecordingApiClient()
- result = QueryCustomerListTool(client).call(
- customer_id=1, department_id=2, sales_id=3,
- merchandiser_id=4, page=5, limit=6, request_id='rq_customer',
- )
- self.assertEqual({'code': 'MCP_0000'}, result)
- self.assertEqual('query_customer_list', client.last_call['tool_code'])
- self.assertEqual('/mcp/tools/queryCustomerList', client.last_call['route_path'])
- self.assertEqual({
- 'customer_id': 1, 'department_id': 2, 'sales_id': 3,
- 'merchandiser_id': 4, 'page': 5, 'limit': 6,
- }, client.last_call['payload'])
- def test_call_rejects_non_strict_and_out_of_range_integers(self):
- tool = QueryCustomerListTool(RecordingApiClient())
- for field, value in (
- ('customer_id', True), ('department_id', '2'), ('sales_id', 1.5),
- ('merchandiser_id', 0), ('page', 101), ('limit', -1),
- ):
- with self.subTest(field=field, value=value):
- with self.assertRaisesRegex(ValueError, field):
- tool.call(**{field: value})
- def test_call_requires_api_client(self):
- with self.assertRaisesRegex(RuntimeError, 'api client is required'):
- QueryCustomerListTool().call()
- def test_local_and_public_registries_are_identical_and_have_17_tools(self):
- local = GatewayApp().registered_tool_names()
- public = PublicGatewayApp(None, None).registered_tool_names()
- self.assertEqual(local, public)
- self.assertEqual(17, len(local))
- self.assertIn('query_customer_list', local)
- self.assertIn('list_customer_filter_options', local)
- def test_cli_forwards_four_customer_filters(self):
- client = RecordingApiClient()
- output = StringIO()
- GatewayApp(api_client=client).run_cli([
- 'call', '--tool', 'query_customer_list', '--customer-id', '11',
- '--department-id', '12', '--sales-id', '13',
- '--merchandiser-id', '14', '--page', '2', '--limit', '30',
- ], stdout=output)
- self.assertEqual({
- 'customer_id': 11, 'department_id': 12, 'sales_id': 13,
- 'merchandiser_id': 14, 'page': 2, 'limit': 30,
- }, client.last_call['payload'])
- def test_cli_allows_customer_query_without_filters(self):
- client = RecordingApiClient()
- GatewayApp(api_client=client).run_cli([
- 'call', '--tool', 'query_customer_list',
- ], stdout=StringIO())
- self.assertEqual({'page': 1, 'limit': 20}, client.last_call['payload'])
- if __name__ == '__main__':
- unittest.main()
|