| import ast |
| import re |
| import unittest |
| from pathlib import Path |
|
|
|
|
| def load_sanitizer(): |
| source = Path(__file__).resolve().parents[1] / "ui_server.py" |
| tree = ast.parse(source.read_text(encoding="utf-8")) |
| function = next( |
| node |
| for node in tree.body |
| if isinstance(node, ast.FunctionDef) and node.name == "_sanitize_final_answer" |
| ) |
| namespace = {"re": re} |
| exec(compile(ast.Module(body=[function], type_ignores=[]), str(source), "exec"), namespace) |
| return namespace["_sanitize_final_answer"] |
|
|
|
|
| sanitize = load_sanitizer() |
|
|
|
|
| class OutputSanitizerTests(unittest.TestCase): |
| def test_rejects_literal_function_call_markup(self): |
| raw = '''我需要查询数据,让我调用工具。<function_calls> |
| <invoke name="mcp_marine_fisheries_inventory"> |
| <parameter name="query">IATTC</parameter> |
| </invoke> |
| </function_calls>''' |
| self.assertEqual(sanitize(raw), "") |
|
|
| def test_keeps_final_answer_and_removes_trailing_markup(self): |
| raw = '''【FINAL】IATTC 数据已确认存在。 |
| <function_calls><invoke name="x"></invoke></function_calls>''' |
| self.assertEqual(sanitize(raw), "IATTC 数据已确认存在。") |
|
|
| def test_keeps_normal_answer(self): |
| self.assertEqual(sanitize("【FINAL】正常回答"), "正常回答") |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|