test_customer_payment_records_tool.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255
  1. import copy
  2. import importlib.util
  3. import io
  4. import json
  5. import os
  6. import unittest
  7. from app import GatewayApp
  8. from public_gateway import PublicGatewayApp
  9. from services.output_presenter import OutputPresenter
  10. class RecordingApiClient:
  11. def __init__(self):
  12. self.calls = []
  13. def list_enabled_tools(self, request_id=''):
  14. return {
  15. 'code': 'MCP_0000',
  16. 'data': {'tool_codes': ['query_customer_payment_records']},
  17. }
  18. def call_tool(self, tool_code, route_path, payload, request_id):
  19. self.calls.append((tool_code, route_path, payload, request_id))
  20. return {'code': 'MCP_0000', 'data': {}}
  21. class CustomerPaymentRecordsToolTest(unittest.TestCase):
  22. def load_tool_class(self):
  23. path = os.path.join(
  24. os.path.dirname(os.path.dirname(__file__)),
  25. 'tools',
  26. 'query_customer_payment_records.py',
  27. )
  28. self.assertTrue(os.path.isfile(path), '客户回款记录 Gateway 工具尚未实现')
  29. spec = importlib.util.spec_from_file_location('payment_records_tool', path)
  30. module = importlib.util.module_from_spec(spec)
  31. spec.loader.exec_module(module)
  32. return module.QueryCustomerPaymentRecordsTool
  33. def test_closed_schema_requires_customer_and_forwards_dates(self):
  34. tool_class = self.load_tool_class()
  35. client = RecordingApiClient()
  36. tool = tool_class(client)
  37. metadata = tool.metadata()
  38. schema = metadata['input_schema']
  39. self.assertEqual('query_customer_payment_records', metadata['name'])
  40. self.assertEqual(['customer_id'], schema['required'])
  41. self.assertFalse(schema['additionalProperties'])
  42. self.assertEqual({
  43. 'customer_id', 'receive_date_start', 'receive_date_end',
  44. 'page', 'limit',
  45. }, set(schema['properties']))
  46. self.assertEqual('date', schema['properties']['receive_date_start']['format'])
  47. self.assertIn('list_customer_filter_options', metadata['description'])
  48. self.assertIn('不得根据名称猜测', metadata['description'])
  49. result = tool.call(
  50. customer_id=7,
  51. receive_date_start='2026-01-01',
  52. receive_date_end='2026-12-31',
  53. page=2,
  54. limit=30,
  55. request_id='rq_records',
  56. )
  57. self.assertEqual('MCP_0000', result['code'])
  58. self.assertEqual((
  59. 'query_customer_payment_records',
  60. '/mcp/tools/queryCustomerPaymentRecords',
  61. {
  62. 'customer_id': 7,
  63. 'receive_date_start': '2026-01-01',
  64. 'receive_date_end': '2026-12-31',
  65. 'page': 2,
  66. 'limit': 30,
  67. },
  68. 'rq_records',
  69. ), client.calls[0])
  70. def test_invalid_arguments_and_oversized_date_range_fail(self):
  71. tool_class = self.load_tool_class()
  72. tool = tool_class(RecordingApiClient())
  73. invalid = (
  74. {'customer_id': True}, {'customer_id': 0}, {'customer_id': '7'},
  75. {'customer_id': 7, 'receive_date_start': '2026-02-30'},
  76. {'customer_id': 7, 'receive_date_start': 20260101},
  77. {'customer_id': 7, 'receive_date_start': '20260101'},
  78. {'customer_id': 7, 'receive_date_start': '2026-02-02', 'receive_date_end': '2026-02-01'},
  79. {'customer_id': 7, 'receive_date_start': '2025-01-01', 'receive_date_end': '2026-01-02'},
  80. {'customer_id': 7, 'page': 0}, {'customer_id': 7, 'page': 101},
  81. {'customer_id': 7, 'limit': False}, {'customer_id': 7, 'limit': 101},
  82. )
  83. for arguments in invalid:
  84. with self.subTest(arguments=arguments):
  85. with self.assertRaises(ValueError):
  86. tool.call(**arguments)
  87. with self.assertRaises(RuntimeError):
  88. tool_class().call(customer_id=7)
  89. client = RecordingApiClient()
  90. tool = tool_class(client)
  91. tool.call(customer_id=7, receive_date_start='2026-01-01')
  92. self.assertEqual('2026-01-01', client.calls[-1][2]['receive_date_start'])
  93. self.assertNotIn('receive_date_end', client.calls[-1][2])
  94. tool.call(customer_id=7, receive_date_end='2026-12-31')
  95. self.assertEqual('2026-12-31', client.calls[-1][2]['receive_date_end'])
  96. self.assertNotIn('receive_date_start', client.calls[-1][2])
  97. tool.call(
  98. customer_id=7,
  99. receive_date_start='2025-01-01',
  100. receive_date_end='2026-01-01',
  101. )
  102. self.assertEqual('2026-01-01', client.calls[-1][2]['receive_date_end'])
  103. def test_registries_and_cli_have_eighteen_tools(self):
  104. client = RecordingApiClient()
  105. local = GatewayApp(api_client=client)
  106. public = PublicGatewayApp(None, None)
  107. self.assertEqual(local.registered_tool_names(), public.registered_tool_names())
  108. self.assertEqual(22, len(local.registered_tool_names()))
  109. self.assertIn('query_customer_payment_records', local.registered_tool_names())
  110. stdout = io.StringIO()
  111. local.run_cli([
  112. 'call', '--tool', 'query_customer_payment_records',
  113. '--customer-id', '7', '--receive-date-start', '2026-01-01',
  114. '--receive-date-end', '2026-12-31', '--page', '2', '--limit', '30',
  115. ], stdout=stdout)
  116. self.assertEqual('MCP_0000', json.loads(stdout.getvalue())['code'])
  117. self.assertEqual({
  118. 'customer_id': 7,
  119. 'receive_date_start': '2026-01-01',
  120. 'receive_date_end': '2026-12-31',
  121. 'page': 2,
  122. 'limit': 30,
  123. }, client.calls[-1][2])
  124. start_only = io.StringIO()
  125. local.run_cli([
  126. 'call', '--tool', 'query_customer_payment_records',
  127. '--customer-id', '7', '--receive-date-start', '2026-01-01',
  128. ], stdout=start_only)
  129. self.assertIn('receive_date_start', client.calls[-1][2])
  130. self.assertNotIn('receive_date_end', client.calls[-1][2])
  131. end_only = io.StringIO()
  132. local.run_cli([
  133. 'call', '--tool', 'query_customer_payment_records',
  134. '--customer-id', '7', '--receive-date-end', '2026-12-31',
  135. ], stdout=end_only)
  136. self.assertIn('receive_date_end', client.calls[-1][2])
  137. self.assertNotIn('receive_date_start', client.calls[-1][2])
  138. with self.assertRaises(ValueError):
  139. local.run_cli([
  140. 'call', '--tool', 'query_customer_payment_records',
  141. ], stdout=io.StringIO())
  142. class CustomerPaymentRecordsPresenterTest(unittest.TestCase):
  143. KEYS = [
  144. 'customer_name', 'payment_reference', 'original_received_amount',
  145. 'actual_received_amount', 'receive_date', 'verified_amount',
  146. 'unverified_amount', 'payment_approval_status',
  147. ]
  148. HEADERS = [
  149. '客户名称', '收款水单号', '原币到账金额', '实际收款金额',
  150. '收款日期', '已核销金额', '未核销金额', '收款审核状态',
  151. ]
  152. def payload(self):
  153. return {
  154. 'code': 'MCP_0000',
  155. 'data': {
  156. 'columns': [{'key': key, 'name': key} for key in self.KEYS],
  157. 'records': [{
  158. 'customer_name': '甲客户',
  159. 'payment_reference': 'BANK-260001',
  160. 'original_received_amount': 100.25,
  161. 'actual_received_amount': 700.5,
  162. 'receive_date': '2026-07-20 16:30:00',
  163. 'verified_amount': 500.25,
  164. 'unverified_amount': 200.25,
  165. 'payment_approval_status': '审核通过',
  166. }],
  167. },
  168. 'meta': {'page': 1, 'limit': 20, 'has_more': False},
  169. }
  170. def test_presenter_translates_exact_eight_columns(self):
  171. result = OutputPresenter().present(
  172. 'query_customer_payment_records', self.payload()
  173. )
  174. self.assertFalse(result['is_error'])
  175. self.assertEqual(
  176. [{'label': header} for header in self.HEADERS],
  177. result['structured_content']['headers'],
  178. )
  179. self.assertEqual(8, len(result['structured_content']['rows'][0]))
  180. serialized = json.dumps(result, ensure_ascii=False)
  181. for key in self.KEYS:
  182. self.assertNotIn(key, serialized)
  183. def test_unknown_missing_malformed_and_nonfinite_fields_fail_closed(self):
  184. cases = []
  185. unknown = self.payload()
  186. unknown['data']['records'][0]['secret'] = 'hidden'
  187. cases.append(unknown)
  188. missing = self.payload()
  189. del missing['data']['records'][0]['payment_reference']
  190. cases.append(missing)
  191. reordered = self.payload()
  192. reordered['data']['columns'].reverse()
  193. cases.append(reordered)
  194. bad_columns = self.payload()
  195. bad_columns['data']['columns'] = 'bad'
  196. cases.append(bad_columns)
  197. bad_column = self.payload()
  198. bad_column['data']['columns'][0]['internal'] = True
  199. cases.append(bad_column)
  200. bad_records = self.payload()
  201. bad_records['data']['records'] = 'bad'
  202. cases.append(bad_records)
  203. bad_text = self.payload()
  204. bad_text['data']['records'][0]['customer_name'] = []
  205. cases.append(bad_text)
  206. bool_amount = self.payload()
  207. bool_amount['data']['records'][0]['verified_amount'] = True
  208. cases.append(bool_amount)
  209. for amount in (float('inf'), float('-inf'), float('nan'), '100.00'):
  210. bad_amount = self.payload()
  211. bad_amount['data']['records'][0]['actual_received_amount'] = amount
  212. cases.append(bad_amount)
  213. extra_data = self.payload()
  214. extra_data['data']['internal'] = True
  215. cases.append(extra_data)
  216. bad_meta = self.payload()
  217. bad_meta['meta']['total'] = 1
  218. cases.append(bad_meta)
  219. presenter = OutputPresenter()
  220. for payload in cases:
  221. with self.subTest(payload=payload):
  222. self.assertTrue(presenter.present(
  223. 'query_customer_payment_records', payload
  224. )['is_error'])
  225. def test_payload_factory_does_not_share_nested_state(self):
  226. self.assertEqual(self.payload(), copy.deepcopy(self.payload()))
  227. if __name__ == '__main__':
  228. unittest.main()