test_query_customs_declaration_files_tool.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315
  1. import importlib
  2. import io
  3. import json
  4. import os
  5. import unittest
  6. from app import GatewayApp
  7. from public_gateway import PublicGatewayApp
  8. from services.output_presenter import OutputPresenter
  9. class RecordingApiClient:
  10. def __init__(self, enabled=True):
  11. self.enabled = enabled
  12. self.last_call = None
  13. def list_enabled_tools(self, request_id=''):
  14. return {
  15. 'code': 'MCP_0000',
  16. 'data': {'tool_codes': (
  17. ['query_customs_declaration_files'] if self.enabled else []
  18. )},
  19. }
  20. def call_tool(self, tool_code, route_path, payload, request_id):
  21. self.last_call = {
  22. 'tool_code': tool_code,
  23. 'route_path': route_path,
  24. 'payload': payload,
  25. 'request_id': request_id,
  26. }
  27. return {'code': 'MCP_0000', 'data': {'columns': [], 'records': []}}
  28. class PublicSessionStore:
  29. def get(self, gateway_session_id):
  30. if gateway_session_id == 'GWS_test':
  31. return {'mcp_token': 'MT_test', 'company_id': 7, 'admin_id': 9}
  32. return None
  33. class PublicApiClient:
  34. def __init__(self, enabled=True):
  35. self.enabled = enabled
  36. self.calls = []
  37. def list_enabled_tools(self, token, request_id=''):
  38. return {
  39. 'code': 'MCP_0000',
  40. 'data': {'tool_codes': (
  41. ['query_customs_declaration_files'] if self.enabled else []
  42. )},
  43. }
  44. def call_tool(self, **kwargs):
  45. self.calls.append(kwargs)
  46. return {'code': 'MCP_0000'}
  47. class QueryCustomsDeclarationFilesToolTest(unittest.TestCase):
  48. def tool_class(self):
  49. path = os.path.join(
  50. os.path.dirname(os.path.dirname(__file__)),
  51. 'tools',
  52. 'query_customs_declaration_files.py',
  53. )
  54. self.assertTrue(os.path.exists(path), path)
  55. module = importlib.import_module(
  56. 'tools.query_customs_declaration_files'
  57. )
  58. return module.QueryCustomsDeclarationFilesTool
  59. def test_metadata_teaches_ai_exact_order_number_semantics(self):
  60. metadata = self.tool_class()().metadata()
  61. schema = metadata['input_schema']
  62. properties = schema['properties']
  63. self.assertFalse(schema['additionalProperties'])
  64. self.assertEqual([
  65. {'required': ['outbound_numbers']},
  66. {'required': ['order_numbers']},
  67. ], schema['oneOf'])
  68. for field in ('outbound_numbers', 'order_numbers'):
  69. prop = properties[field]
  70. self.assertEqual('array', prop['type'])
  71. self.assertEqual('string', prop['items']['type'])
  72. self.assertEqual(1, prop['minItems'])
  73. self.assertEqual(100, prop['maxItems'])
  74. self.assertNotIn('uniqueItems', prop)
  75. order_description = properties['order_numbers']['description']
  76. for phrase in (
  77. '订单号',
  78. '后台订单列表',
  79. '排舱详情',
  80. '不是系统单号',
  81. 'order_id/id',
  82. '不是客户参考号',
  83. '快递单号',
  84. '排舱单号',
  85. ):
  86. self.assertIn(phrase, order_description)
  87. self.assertNotIn('系统订单号', order_description)
  88. description = metadata['description']
  89. for phrase in (
  90. '用户明确说“订单号”',
  91. '用户明确说“排舱单号”',
  92. '没有明确说明是订单号还是排舱单号',
  93. '必须先提问,让用户选择“订单号”或“排舱单号”',
  94. '只说“单号”时必须先追问',
  95. '不是系统单号',
  96. '用户说“系统单号”时也不得当作订单号',
  97. '不得根据号码格式猜测',
  98. '不得跨字段重试',
  99. '不得展示内部参数名',
  100. ):
  101. self.assertIn(phrase, description)
  102. self.assertNotIn('系统订单号', description)
  103. def test_outbound_batch_normalizes_and_forwards_supported_fields(self):
  104. client = RecordingApiClient()
  105. tool = self.tool_class()(api_client=client)
  106. result = tool.call(
  107. outbound_numbers=[' PC001 ', 'PC002', 'PC001'],
  108. page=2,
  109. limit=50,
  110. request_id='rq_customs',
  111. )
  112. self.assertEqual('MCP_0000', result['code'])
  113. self.assertEqual(
  114. 'query_customs_declaration_files',
  115. client.last_call['tool_code'],
  116. )
  117. self.assertEqual(
  118. '/mcp/tools/queryCustomsDeclarationFiles',
  119. client.last_call['route_path'],
  120. )
  121. self.assertEqual({
  122. 'outbound_numbers': ['PC001', 'PC002'],
  123. 'page': 2,
  124. 'limit': 50,
  125. }, client.last_call['payload'])
  126. def test_order_batch_uses_order_numbers_business_field(self):
  127. client = RecordingApiClient()
  128. tool = self.tool_class()(api_client=client)
  129. tool.call(order_numbers=[' ORD001 ', 'ORD002'])
  130. self.assertEqual({
  131. 'order_numbers': ['ORD001', 'ORD002'],
  132. 'page': 1,
  133. 'limit': 20,
  134. }, client.last_call['payload'])
  135. def test_call_rejects_mixed_missing_and_invalid_batches(self):
  136. tool = self.tool_class()(api_client=RecordingApiClient())
  137. invalid = (
  138. {},
  139. {'outbound_numbers': ['PC001'], 'order_numbers': ['ORD001']},
  140. {'order_numbers': []},
  141. {'order_numbers': 'ORD001'},
  142. {'order_numbers': ['']},
  143. {'order_numbers': ['ORD001', 2]},
  144. {'order_numbers': ['X' * 101]},
  145. {'order_numbers': ['ORD{0}'.format(i) for i in range(101)]},
  146. {'order_numbers': ['ORD001'], 'page': 0},
  147. {'order_numbers': ['ORD001'], 'page': 101},
  148. {'order_numbers': ['ORD001'], 'limit': 101},
  149. {'order_numbers': ['ORD001'], 'page': True},
  150. {'order_numbers': ['ORD001'], 'page': 'not-a-number'},
  151. )
  152. for arguments in invalid:
  153. with self.subTest(arguments=arguments):
  154. with self.assertRaises(ValueError):
  155. tool.call(**arguments)
  156. with self.assertRaisesRegex(RuntimeError, 'api client is required'):
  157. self.tool_class()().call(order_numbers=['ORD001'])
  158. def test_local_and_public_gateways_follow_dynamic_registry(self):
  159. local_enabled = {
  160. item['name']
  161. for item in GatewayApp(api_client=RecordingApiClient()).list_tools()
  162. }
  163. local_disabled = {
  164. item['name']
  165. for item in GatewayApp(
  166. api_client=RecordingApiClient(enabled=False)
  167. ).list_tools()
  168. }
  169. public_enabled = {
  170. item['name']
  171. for item in PublicGatewayApp(
  172. PublicSessionStore(),
  173. PublicApiClient(),
  174. ).list_tools('GWS_test')
  175. }
  176. self.assertIn('query_customs_declaration_files', local_enabled)
  177. self.assertNotIn('query_customs_declaration_files', local_disabled)
  178. self.assertIn('query_customs_declaration_files', public_enabled)
  179. def test_public_gateway_forwards_request_scoped_token_without_company_input(self):
  180. api_client = PublicApiClient()
  181. gateway = PublicGatewayApp(PublicSessionStore(), api_client)
  182. self.assertIn(
  183. 'query_customs_declaration_files',
  184. gateway.registered_tool_names(),
  185. )
  186. gateway.call_tool(
  187. 'GWS_test',
  188. 'query_customs_declaration_files',
  189. {'order_numbers': ['ORD001']},
  190. request_id='rq_public_customs',
  191. )
  192. call = api_client.calls[0]
  193. self.assertEqual('MT_test', call['token'])
  194. self.assertEqual(
  195. '/mcp/tools/queryCustomsDeclarationFiles',
  196. call['route_path'],
  197. )
  198. self.assertEqual({'order_numbers': ['ORD001']}, call['payload'])
  199. self.assertNotIn('company_id', call['payload'])
  200. def test_output_presenter_hides_internal_fields_and_keeps_links(self):
  201. result = OutputPresenter().present(
  202. 'query_customs_declaration_files',
  203. {
  204. 'code': 'MCP_0000',
  205. 'data': {
  206. 'summary': '当前页返回 1 个订单的报关资料',
  207. 'columns': [
  208. {'key': 'outbound_number', 'name': '排舱单号'},
  209. {'key': 'order_number', 'name': '订单号'},
  210. {'key': 'file_name', 'name': '文件名'},
  211. {'key': 'file_url', 'name': '文件链接'},
  212. ],
  213. 'records': [{
  214. 'outbound_number': 'PC001',
  215. 'order_number': 'ORD001',
  216. 'file_name': '报关单.pdf',
  217. 'file_url': 'https://files.test/a.pdf',
  218. 'order_id': 99,
  219. }],
  220. },
  221. 'meta': {
  222. 'page': 1,
  223. 'limit': 20,
  224. 'has_more': False,
  225. 'request_id': 'rq_present',
  226. },
  227. },
  228. )
  229. self.assertFalse(result['is_error'])
  230. self.assertEqual(
  231. ['排舱单号', '订单号', '文件名', '文件链接'],
  232. [header['label'] for header in result['structured_content']['headers']],
  233. )
  234. self.assertIn('https://files.test/a.pdf', result['text'])
  235. serialized = json.dumps(result, ensure_ascii=False)
  236. for internal in (
  237. 'outbound_number',
  238. 'order_number',
  239. 'file_name',
  240. 'file_url',
  241. 'order_id',
  242. ):
  243. self.assertNotIn(internal, serialized)
  244. def test_cli_forwards_batch_arguments(self):
  245. client = RecordingApiClient()
  246. stdout = io.StringIO()
  247. app = GatewayApp(api_client=client)
  248. self.assertIn(
  249. 'query_customs_declaration_files',
  250. app.registered_tool_names(),
  251. )
  252. result = app.run_cli([
  253. 'call',
  254. '--tool', 'query_customs_declaration_files',
  255. '--order-numbers', 'ORD001,ORD002',
  256. '--page', '2',
  257. '--limit', '10',
  258. ], stdout=stdout)
  259. self.assertEqual(0, result)
  260. self.assertEqual({
  261. 'order_numbers': ['ORD001', 'ORD002'],
  262. 'page': 2,
  263. 'limit': 10,
  264. }, client.last_call['payload'])
  265. app.run_cli([
  266. 'call',
  267. '--tool', 'query_customs_declaration_files',
  268. '--outbound-numbers', 'PC001,PC002',
  269. ], stdout=io.StringIO())
  270. self.assertEqual({
  271. 'outbound_numbers': ['PC001', 'PC002'],
  272. 'page': 1,
  273. 'limit': 20,
  274. }, client.last_call['payload'])
  275. if __name__ == '__main__':
  276. unittest.main()