test_customer_query_tools.py 4.2 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798
  1. import unittest
  2. from io import StringIO
  3. from app import GatewayApp
  4. from public_gateway import PublicGatewayApp
  5. from tools.query_customer_list import QueryCustomerListTool
  6. class RecordingApiClient:
  7. def __init__(self):
  8. self.last_call = None
  9. def call_tool(self, tool_code, route_path, payload, request_id):
  10. self.last_call = {
  11. 'tool_code': tool_code, 'route_path': route_path,
  12. 'payload': payload, 'request_id': request_id,
  13. }
  14. return {'code': 'MCP_0000'}
  15. def list_enabled_tools(self, request_id=''):
  16. return {'code': 'MCP_0000', 'data': {'tool_codes': ['query_customer_list']}}
  17. class CustomerQueryToolTest(unittest.TestCase):
  18. def test_schema_is_closed_and_uses_strict_positive_integers(self):
  19. metadata = QueryCustomerListTool().metadata()
  20. schema = metadata['input_schema']
  21. self.assertIn('不得合并或缩写', metadata['description'])
  22. self.assertIn('完整逐条展示17个字段', metadata['description'])
  23. self.assertFalse(schema['additionalProperties'])
  24. self.assertEqual([], schema['required'])
  25. for field in ('customer_id', 'department_id', 'sales_id', 'merchandiser_id'):
  26. self.assertEqual({'type': 'integer', 'minimum': 1}, schema['properties'][field])
  27. for field, default in (('page', 1), ('limit', 20)):
  28. self.assertEqual(1, schema['properties'][field]['minimum'])
  29. self.assertEqual(100, schema['properties'][field]['maximum'])
  30. self.assertEqual(default, schema['properties'][field]['default'])
  31. def test_call_forwards_all_filters(self):
  32. client = RecordingApiClient()
  33. result = QueryCustomerListTool(client).call(
  34. customer_id=1, department_id=2, sales_id=3,
  35. merchandiser_id=4, page=5, limit=6, request_id='rq_customer',
  36. )
  37. self.assertEqual({'code': 'MCP_0000'}, result)
  38. self.assertEqual('query_customer_list', client.last_call['tool_code'])
  39. self.assertEqual('/mcp/tools/queryCustomerList', client.last_call['route_path'])
  40. self.assertEqual({
  41. 'customer_id': 1, 'department_id': 2, 'sales_id': 3,
  42. 'merchandiser_id': 4, 'page': 5, 'limit': 6,
  43. }, client.last_call['payload'])
  44. def test_call_rejects_non_strict_and_out_of_range_integers(self):
  45. tool = QueryCustomerListTool(RecordingApiClient())
  46. for field, value in (
  47. ('customer_id', True), ('department_id', '2'), ('sales_id', 1.5),
  48. ('merchandiser_id', 0), ('page', 101), ('limit', -1),
  49. ):
  50. with self.subTest(field=field, value=value):
  51. with self.assertRaisesRegex(ValueError, field):
  52. tool.call(**{field: value})
  53. def test_call_requires_api_client(self):
  54. with self.assertRaisesRegex(RuntimeError, 'api client is required'):
  55. QueryCustomerListTool().call()
  56. def test_local_and_public_registries_are_identical_and_have_25_tools(self):
  57. local = GatewayApp().registered_tool_names()
  58. public = PublicGatewayApp(None, None).registered_tool_names()
  59. self.assertEqual(local, public)
  60. self.assertEqual(25, len(local))
  61. self.assertIn('query_customer_list', local)
  62. self.assertIn('list_customer_filter_options', local)
  63. def test_cli_forwards_four_customer_filters(self):
  64. client = RecordingApiClient()
  65. output = StringIO()
  66. GatewayApp(api_client=client).run_cli([
  67. 'call', '--tool', 'query_customer_list', '--customer-id', '11',
  68. '--department-id', '12', '--sales-id', '13',
  69. '--merchandiser-id', '14', '--page', '2', '--limit', '30',
  70. ], stdout=output)
  71. self.assertEqual({
  72. 'customer_id': 11, 'department_id': 12, 'sales_id': 13,
  73. 'merchandiser_id': 14, 'page': 2, 'limit': 30,
  74. }, client.last_call['payload'])
  75. def test_cli_allows_customer_query_without_filters(self):
  76. client = RecordingApiClient()
  77. GatewayApp(api_client=client).run_cli([
  78. 'call', '--tool', 'query_customer_list',
  79. ], stdout=StringIO())
  80. self.assertEqual({'page': 1, 'limit': 20}, client.last_call['payload'])
  81. if __name__ == '__main__':
  82. unittest.main()