test_query_order_exact_tool.py 9.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258
  1. import inspect
  2. import unittest
  3. from io import StringIO
  4. import app as gateway_app_module
  5. from app import GatewayApp, parse_int_list
  6. from public_gateway import PublicGatewayApp
  7. from tools.query_order_exact import QueryOrderExactTool
  8. class RecordingApiClient:
  9. def __init__(self):
  10. self.last_call = None
  11. def call_tool(self, tool_code, route_path, payload, request_id):
  12. self.last_call = {
  13. 'tool_code': tool_code,
  14. 'route_path': route_path,
  15. 'payload': payload,
  16. 'request_id': request_id,
  17. }
  18. return {'code': 'MCP_0000'}
  19. def list_enabled_tools(self):
  20. return {
  21. 'code': 'MCP_0000',
  22. 'data': {'tool_codes': ['query_order_exact']},
  23. }
  24. class PublicSessionStore:
  25. def get(self, gateway_session_id):
  26. if gateway_session_id == 'GWS_test':
  27. return {'mcp_token': 'MT_test'}
  28. return None
  29. class PublicApiClient:
  30. def list_enabled_tools(self, token):
  31. return {
  32. 'code': 'MCP_0000',
  33. 'data': {'tool_codes': ['query_order_exact']},
  34. }
  35. class QueryOrderExactToolTest(unittest.TestCase):
  36. def test_local_and_public_gateways_register_tool(self):
  37. local_names = {
  38. tool['name']
  39. for tool in GatewayApp(api_client=RecordingApiClient()).list_tools()
  40. }
  41. public_names = {
  42. tool['name']
  43. for tool in PublicGatewayApp(
  44. PublicSessionStore(),
  45. PublicApiClient(),
  46. ).list_tools('GWS_test')
  47. }
  48. self.assertIn('query_order_exact', local_names)
  49. self.assertIn('query_order_exact', public_names)
  50. def test_parse_int_list_accepts_comma_separated_ids(self):
  51. self.assertEqual([9, 12], parse_int_list('9, 12'))
  52. self.assertEqual([], parse_int_list(''))
  53. def test_metadata_exposes_all_exact_filters(self):
  54. schema = QueryOrderExactTool().metadata()['input_schema']
  55. for field in (
  56. 'order_number', 'reference_number', 'tracking_number',
  57. 'outbound_number', 'container_code', 'so_number', 'shipment_id',
  58. 'receiver_country', 'product_ids', 'customer_ids', 'sales_id',
  59. 'warehouse_ids', 'department_id', 'inbound_date_start',
  60. 'inbound_date_end', 'outbound_date_start', 'outbound_date_end',
  61. 'page', 'limit',
  62. ):
  63. self.assertIn(field, schema['properties'])
  64. for field in (
  65. 'order_numbers', 'reference_numbers', 'tracking_numbers',
  66. 'outbound_numbers', 'container_codes', 'so_numbers',
  67. ):
  68. self.assertIn(field, schema['properties'])
  69. self.assertEqual('array', schema['properties'][field]['type'])
  70. self.assertEqual(
  71. 'string',
  72. schema['properties'][field]['items']['type'],
  73. )
  74. self.assertEqual(100, schema['properties']['limit']['maximum'])
  75. def test_metadata_guides_ai_to_explicit_fields_without_fallback(self):
  76. metadata = QueryOrderExactTool().metadata()
  77. description = metadata['description']
  78. properties = metadata['input_schema']['properties']
  79. for phrase in (
  80. '单号类型不明确', '先询问用户', '不得改用其他字段',
  81. '同一字段使用 IN', '不同字段使用 AND',
  82. ):
  83. self.assertIn(phrase, description)
  84. single_fields = (
  85. 'order_number', 'reference_number', 'tracking_number',
  86. 'outbound_number', 'container_code', 'so_number',
  87. )
  88. batch_fields = (
  89. 'order_numbers', 'reference_numbers', 'tracking_numbers',
  90. 'outbound_numbers', 'container_codes', 'so_numbers',
  91. )
  92. for field in single_fields:
  93. self.assertTrue(properties[field]['description'])
  94. self.assertIsInstance(properties[field]['examples'][0], str)
  95. for field in batch_fields:
  96. self.assertTrue(properties[field]['description'])
  97. self.assertIsInstance(properties[field]['examples'][0], list)
  98. self.assertGreater(len(properties[field]['examples'][0]), 1)
  99. self.assertIn('系统订单号', properties['order_number']['description'])
  100. self.assertIn(
  101. '客户参考号',
  102. properties['reference_number']['description'],
  103. )
  104. self.assertIn('承运商', properties['tracking_number']['description'])
  105. self.assertIn('出库单号', properties['outbound_number']['description'])
  106. self.assertIn('柜号', properties['container_code']['description'])
  107. self.assertIn('Shipping Order', properties['so_number']['description'])
  108. for field in (
  109. 'shipment_id', 'receiver_country', 'product_ids', 'customer_ids',
  110. 'sales_id', 'warehouse_ids', 'department_id',
  111. 'inbound_date_start', 'inbound_date_end',
  112. 'outbound_date_start', 'outbound_date_end', 'page', 'limit',
  113. ):
  114. self.assertTrue(properties[field]['description'])
  115. def test_call_forwards_normalized_number_arrays(self):
  116. tool = QueryOrderExactTool(api_client=RecordingApiClient())
  117. self.assertIn('order_numbers', inspect.signature(tool.call).parameters)
  118. tool.call(
  119. order_number=' A ',
  120. order_numbers=[' B ', 'A'],
  121. tracking_numbers=['T1', 'T2'],
  122. )
  123. self.assertEqual('A', tool.api_client.last_call['payload']['order_number'])
  124. self.assertEqual(
  125. ['B', 'A'],
  126. tool.api_client.last_call['payload']['order_numbers'],
  127. )
  128. self.assertEqual(
  129. ['T1', 'T2'],
  130. tool.api_client.last_call['payload']['tracking_numbers'],
  131. )
  132. def test_call_rejects_invalid_number_arrays(self):
  133. tool = QueryOrderExactTool(api_client=RecordingApiClient())
  134. self.assertIn('order_numbers', inspect.signature(tool.call).parameters)
  135. invalid_values = ('A,B', [1], [''], [' '], ['X' * 101])
  136. for values in invalid_values:
  137. with self.subTest(values=values):
  138. with self.assertRaises(ValueError):
  139. tool.call(order_numbers=values)
  140. with self.assertRaisesRegex(ValueError, 'must not exceed 200'):
  141. tool.call(
  142. order_numbers=['O{0}'.format(i) for i in range(200)],
  143. so_numbers=['S1'],
  144. )
  145. def test_call_normalizes_exact_filters_and_id_arrays(self):
  146. client = RecordingApiClient()
  147. tool = QueryOrderExactTool(api_client=client)
  148. result = tool.call(
  149. order_number=' USC001 ',
  150. customer_ids=['9', 9, 0, -1, '12'],
  151. warehouse_ids=['-1', '5'],
  152. sales_id='7',
  153. page=0,
  154. limit=200,
  155. request_id='rq_exact',
  156. )
  157. self.assertEqual({'code': 'MCP_0000'}, result)
  158. self.assertEqual('query_order_exact', client.last_call['tool_code'])
  159. self.assertEqual('/mcp/tools/queryOrderExact', client.last_call['route_path'])
  160. self.assertEqual('rq_exact', client.last_call['request_id'])
  161. self.assertEqual({
  162. 'order_number': 'USC001',
  163. 'customer_ids': [9, 12],
  164. 'sales_id': 7,
  165. 'warehouse_ids': [-1, 5],
  166. 'page': 1,
  167. 'limit': 100,
  168. }, client.last_call['payload'])
  169. def test_call_requires_one_business_filter(self):
  170. tool = QueryOrderExactTool(api_client=RecordingApiClient())
  171. with self.assertRaisesRegex(ValueError, 'at least one exact order filter'):
  172. tool.call(page=1, limit=20)
  173. def test_call_requires_api_client(self):
  174. with self.assertRaisesRegex(RuntimeError, 'api client is required'):
  175. QueryOrderExactTool().call(order_number='USC001')
  176. def test_cli_forwards_exact_fields_and_id_lists(self):
  177. client = RecordingApiClient()
  178. output = StringIO()
  179. exit_code = GatewayApp(api_client=client).run_cli([
  180. 'call',
  181. '--tool', 'query_order_exact',
  182. '--order-number', 'USC001',
  183. '--customer-ids', '9,12',
  184. '--warehouse-ids=-1,5',
  185. '--inbound-date-start', '2026-07-01',
  186. ], stdout=output)
  187. self.assertEqual(0, exit_code)
  188. self.assertEqual({
  189. 'order_number': 'USC001',
  190. 'customer_ids': [9, 12],
  191. 'warehouse_ids': [-1, 5],
  192. 'inbound_date_start': '2026-07-01',
  193. 'page': 1,
  194. 'limit': 20,
  195. }, client.last_call['payload'])
  196. def test_cli_forwards_batch_number_lists(self):
  197. self.assertTrue(hasattr(gateway_app_module, 'parse_string_list'))
  198. self.assertEqual(
  199. ['A', 'B'],
  200. gateway_app_module.parse_string_list(' A, B, A '),
  201. )
  202. client = RecordingApiClient()
  203. output = StringIO()
  204. exit_code = GatewayApp(api_client=client).run_cli([
  205. 'call',
  206. '--tool', 'query_order_exact',
  207. '--order-numbers', 'A,B,A',
  208. '--tracking-numbers', 'T1,T2',
  209. ], stdout=output)
  210. self.assertEqual(0, exit_code)
  211. self.assertEqual(['A', 'B'], client.last_call['payload']['order_numbers'])
  212. self.assertEqual(
  213. ['T1', 'T2'],
  214. client.last_call['payload']['tracking_numbers'],
  215. )
  216. if __name__ == '__main__':
  217. unittest.main()