"""Tests for the MCP bridge:  python3 -m unittest scripts/test_mcp_bridge.py

Starts serve.py's server with the bridge on a free port, the demo stdio server, command-line
tools and a mock Streamable HTTP server, and talks to them like the page does.
"""
import base64
import http.client
import http.server
import json
import pathlib
import queue
import subprocess
import sys
import tempfile
import threading
import time
import unittest
import zlib

HERE = pathlib.Path(__file__).resolve().parent
sys.path.insert(0, str(HERE))
import mcp_bridge  # noqa: E402
import serve  # noqa: E402

PY = sys.executable
LONG_LINE_SERVER = r"""
import json, sys
for line in sys.stdin:
    msg = json.loads(line)
    if 'id' not in msg: continue
    sys.stdout.write('y' * 3_000_000 + '\n')
    sys.stdout.write(json.dumps({'jsonrpc': '2.0', 'id': msg['id'], 'result': {'ok': True}}) + '\n')
    sys.stdout.flush()
"""

# Claims to be another server inside its own messages; the bridge's tag must win
SPOOF_SERVER = r"""
import json, sys
for line in sys.stdin:
    msg = json.loads(line)
    if 'id' not in msg: continue
    sys.stdout.write(json.dumps({'jsonrpc': '2.0', 'method': 'notifications/message', 'server': 'tools',
                                 'params': {'level': 'info', 'data': 'spoofed', 'server': 'tools'}}) + '\n')
    sys.stdout.write(json.dumps({'jsonrpc': '2.0', 'id': msg['id'], 'result': {}}) + '\n')
    sys.stdout.flush()
"""


class MockHttpMcp(http.server.BaseHTTPRequestHandler):
    """A tiny Streamable HTTP MCP server: JSON for initialize, SSE for tools/list."""
    def log_message(self, *args):
        pass

    def do_POST(self):
        msg = json.loads(self.rfile.read(int(self.headers['Content-Length'])))
        if 'id' not in msg:
            self.send_response(202)
            self.end_headers()
            return
        if msg['method'] == 'initialize':
            body = json.dumps({'jsonrpc': '2.0', 'id': msg['id'], 'result': {'protocolVersion': '2025-06-18', 'capabilities': {}, 'serverInfo': {'name': 'mock', 'version': '1'}}}).encode()
            self.send_response(200)
            self.send_header('Content-Type', 'application/json')
            self.send_header('Mcp-Session-Id', 'sess-1')
            self.send_header('Content-Length', str(len(body)))
            self.end_headers()
            self.wfile.write(body)
            return
        if self.headers.get('Mcp-Session-Id') != 'sess-1' or self.headers.get('MCP-Protocol-Version') != '2025-06-18':
            self.send_response(400)
            self.end_headers()
            return
        events = [{'jsonrpc': '2.0', 'method': 'notifications/message', 'params': {'level': 'info', 'data': 'from sse'}},
                  {'jsonrpc': '2.0', 'id': msg['id'], 'result': {'tools': [{'name': 'remote', 'inputSchema': {'type': 'object'}}]}}]
        body = ''.join(f'data: {json.dumps(e)}\n\n' for e in events).encode()
        self.send_response(200)
        self.send_header('Content-Type', 'text/event-stream')
        self.send_header('Content-Length', str(len(body)))
        self.end_headers()
        self.wfile.write(body)


class BridgeTest(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.tmp = tempfile.TemporaryDirectory()
        cls.mock = http.server.ThreadingHTTPServer(('127.0.0.1', 0), MockHttpMcp)
        threading.Thread(target=cls.mock.serve_forever, daemon=True).start()
        demo = [PY, str(HERE / 'mcp_demo_server.py')]
        config = {'servers': [
            {'id': 'demo', 'name': '示例', 'command': demo},
            {'id': 'quick', 'name': '一秒超时', 'command': demo, 'timeout_s': 1},
            {'id': 'remote', 'name': '远程', 'url': f'http://127.0.0.1:{cls.mock.server_address[1]}/mcp'},
            # Writes an over-long line (no newline for a while) before every proper response
            {'id': 'noisy', 'name': '超长输出', 'command': [PY, '-c', LONG_LINE_SERVER]},
            {'id': 'spoof', 'name': '冒充', 'command': [PY, '-c', SPOOF_SERVER]},
            {'id': 'tools', 'name': '命令行', 'cli': [
                {'name': 'length', 'description': 'Length of the text', 'command': [PY, '-c', 'import sys; print(len(sys.argv[1]))', '{text}'],
                 'input_schema': {'type': 'object', 'properties': {'text': {'type': 'string'}}, 'required': ['text']}},
                {'name': 'args', 'command': [PY, '-c', 'import sys; print(sys.argv[1:])', 'fixed', '--n={n}', '{flag}']},
                {'name': 'upper', 'command': [PY, '-c', 'import sys; print(sys.stdin.read().upper())'], 'stdin': 'text'},
                {'name': 'exit3', 'command': [PY, '-c', 'import sys; print("oops", file=sys.stderr); sys.exit(3)']},
                {'name': 'sleepy', 'command': [PY, '-c', 'import time; time.sleep(5)'], 'timeout_s': 0.5},
                {'name': 'missing', 'command': ['/nonexistent/program']},
                {'name': 'flood', 'command': [PY, '-c', 'import sys; sys.stdout.write("x" * 5_000_000)']},
                {'name': 'grep', 'command': [PY, '-c', 'import sys; print(sys.argv[1:])', '{pattern}']},
                {'name': 'calc', 'command': [PY, '-c', 'import sys; print(sys.argv[1:])', '{n}'], 'allow_option_values': True},
            ]},
        ]}
        path = pathlib.Path(cls.tmp.name) / 'mcp.config.json'
        path.write_text(json.dumps(config), encoding='utf-8')
        mcp_bridge.MAX_LINE = 1_000_000   # keep the over-long-line test small
        mcp_bridge.MAX_OUTPUT = 100_000
        cls.bridge = mcp_bridge.Bridge(path, 0)
        cls.httpd = serve.make_server(0, cls.bridge, host='127.0.0.1', quiet=True)
        cls.port = cls.httpd.server_address[1]
        cls.bridge.port = cls.port
        threading.Thread(target=cls.httpd.serve_forever, daemon=True).start()
        status, headers, _ = cls.request('GET', '/www/mcp/', cookie=False)
        cls.cookie = headers.get('Set-Cookie', '').split(';')[0]

    @classmethod
    def tearDownClass(cls):
        cls.httpd.shutdown()
        cls.bridge.shutdown()
        cls.mock.shutdown()
        cls.tmp.cleanup()

    @classmethod
    def request(cls, method, path, body=None, cookie=True, host=None, origin=None, timeout=10):
        conn = http.client.HTTPConnection('127.0.0.1', cls.port, timeout=timeout)
        headers = {'Host': host or f'localhost:{cls.port}'}
        if cookie and getattr(cls, 'cookie', None):
            headers['Cookie'] = cls.cookie
        if origin:
            headers['Origin'] = origin
        data = None
        if body is not None:
            data = (body if isinstance(body, str) else json.dumps(body)).encode()
            headers['Content-Type'] = 'application/json'
        conn.request(method, path, body=data, headers=headers)
        resp = conn.getresponse()
        text = resp.read().decode('utf-8', 'replace')
        conn.close()
        return resp.status, dict(resp.getheaders()), text

    def rpc(self, server, method, params=None, mid=1, notification=False):
        msg = {'jsonrpc': '2.0', 'method': method}
        if not notification:
            msg['id'] = mid
        if params is not None:
            msg['params'] = params
        status, _, text = self.request('POST', f'/api/mcp/{server}/rpc', msg)
        return status, (json.loads(text) if text else None)

    def initialize(self, server):
        status, reply = self.rpc(server, 'initialize', {'protocolVersion': '2025-06-18', 'capabilities': {}, 'clientInfo': {'name': 'test', 'version': '0'}})
        self.assertEqual(status, 200, reply)
        self.assertEqual(self.rpc(server, 'notifications/initialized', notification=True)[0], 202)
        return reply['result']

    def events(self, server=None):
        """Subscribe to a server's events (or, with no server, to every server's events on
        one stream); returns a queue of parsed events."""
        q = queue.Queue()
        conn = http.client.HTTPConnection('127.0.0.1', self.port, timeout=30)
        conn.request('GET', f'/api/mcp/{server}/events' if server else '/api/mcp/events', headers={'Host': f'localhost:{self.port}', 'Cookie': self.cookie})
        resp = conn.getresponse()
        self.assertEqual(resp.status, 200)

        def read():
            try:
                for raw in resp:
                    line = raw.decode().rstrip('\n')
                    if line.startswith('data: '):
                        q.put(json.loads(line[6:]))
            except OSError:
                pass
        threading.Thread(target=read, daemon=True).start()
        self.addCleanup(conn.close)
        return q

    def wait_for(self, q, pred, timeout=10):
        end = time.time() + timeout
        while time.time() < end:
            try:
                e = q.get(timeout=0.2)
            except queue.Empty:
                continue
            if pred(e):
                return e
        self.fail('expected event did not arrive')

    # ── security ──

    def test_token_cookie_is_issued_to_local_page_requests_only(self):
        self.assertTrue(self.cookie.startswith('mcp_token='))
        _, headers, _ = self.request('GET', '/www/mcp/', cookie=False, host='evil.example:80')
        self.assertNotIn('Set-Cookie', headers)

    def test_requests_without_the_token_are_refused(self):
        status, _, text = self.request('GET', '/api/mcp/servers', cookie=False)
        self.assertEqual(status, 403)
        self.assertIn('令牌', json.loads(text)['error'])

    def test_requests_from_other_origins_or_hosts_are_refused(self):
        self.assertEqual(self.request('GET', '/api/mcp/servers', origin='https://evil.example')[0], 403)
        self.assertEqual(self.request('GET', '/api/mcp/servers', host='evil.example')[0], 403, 'DNS rebinding: wrong Host')
        self.assertEqual(self.request('GET', '/api/mcp/servers', origin=f'http://localhost:{self.port}')[0], 200)

    def test_bad_requests(self):
        self.assertEqual(self.request('POST', '/api/mcp/demo/rpc', 'not json')[0], 400)
        self.assertEqual(self.request('POST', '/api/mcp/demo/rpc', '[]')[0], 400)
        self.assertEqual(self.request('POST', '/api/mcp/nope/rpc', {'jsonrpc': '2.0', 'id': 1, 'method': 'ping'})[0], 404)
        self.assertEqual(self.request('GET', '/api/mcp/demo/unknown')[0], 404)

    def test_malformed_content_length_gets_an_answer(self):
        conn = http.client.HTTPConnection('127.0.0.1', self.port, timeout=10)
        conn.putrequest('POST', '/api/mcp/demo/rpc', skip_host=True)
        for k, v in [('Host', f'localhost:{self.port}'), ('Cookie', self.cookie), ('Content-Length', 'notanumber')]:
            conn.putheader(k, v)
        conn.endheaders()
        self.assertEqual(conn.getresponse().status, 400)
        conn.close()

    def test_cookie_is_not_issued_to_cross_site_requests(self):
        conn = http.client.HTTPConnection('127.0.0.1', self.port, timeout=10)
        conn.request('GET', '/www/mcp/', headers={'Host': f'localhost:{self.port}', 'Sec-Fetch-Site': 'cross-site'})
        self.assertIsNone(conn.getresponse().getheader('Set-Cookie'))
        conn.close()
        conn = http.client.HTTPConnection('127.0.0.1', self.port, timeout=10)
        conn.request('GET', '/www/mcp/', headers={'Host': f'localhost:{self.port}', 'Origin': 'https://evil.example'})
        self.assertIsNone(conn.getresponse().getheader('Set-Cookie'))
        conn.close()

    def test_server_list(self):
        status, _, text = self.request('GET', '/api/mcp/servers')
        self.assertEqual(status, 200)
        self.assertEqual([(s['id'], s['kind']) for s in json.loads(text)], [('demo', 'stdio'), ('quick', 'stdio'), ('remote', 'http'), ('noisy', 'stdio'), ('spoof', 'stdio'), ('tools', 'cli')])

    # ── one event stream for all servers ──

    def test_one_stream_carries_every_servers_events_tagged_by_the_bridge(self):
        q = self.events()
        self.initialize('demo')
        self.rpc('demo', 'tools/call', {'name': 'notify', 'arguments': {'steps': 1}}, mid=2)
        demo = self.wait_for(q, lambda e: e['type'] == 'message' and e['message'].get('method') == 'notifications/message')
        self.assertEqual(demo['server'], 'demo')
        self.rpc('spoof', 'ping', mid=3)
        spoofed = self.wait_for(q, lambda e: e['type'] == 'message' and e['message'].get('params', {}).get('data') == 'spoofed')
        self.assertEqual(spoofed['server'], 'spoof', 'the tag comes from the bridge, not the message')
        self.assertEqual(spoofed['message']['server'], 'tools')

    def test_per_server_streams_are_tagged_too(self):
        q = self.events('spoof')
        self.rpc('spoof', 'ping', mid=4)
        self.assertEqual(self.wait_for(q, lambda e: e['type'] == 'message')['server'], 'spoof')

    def test_merged_stream_needs_the_token_and_this_origin(self):
        self.assertEqual(self.request('GET', '/api/mcp/events', cookie=False)[0], 403)
        self.assertEqual(self.request('GET', '/api/mcp/events', origin='https://evil.example')[0], 403)
        self.assertEqual(self.request('GET', '/api/mcp/events', host='evil.example')[0], 403)

    # ── stdio ──

    def test_stdio_handshake_pagination_and_call(self):
        result = self.initialize('demo')
        self.assertEqual(result['serverInfo']['name'], 'webtt-demo')
        _, page = self.rpc('demo', 'tools/list', mid=2)
        self.assertEqual(len(page['result']['tools']), 3)
        self.assertEqual(page['result']['nextCursor'], '3')
        _, added = self.rpc('demo', 'tools/call', {'name': 'add', 'arguments': {'a': 2, 'b': 40}}, mid=3)
        self.assertEqual(added['result']['content'][0]['text'], '42')

    def test_stdio_notifications_and_stderr_reach_the_event_stream(self):
        q = self.events('demo')
        self.request('POST', '/api/mcp/demo/restart')
        self.initialize('demo')
        self.wait_for(q, lambda e: e['type'] == 'stderr' and 'ready' in e['line'])
        self.rpc('demo', 'tools/call', {'name': 'notify', 'arguments': {'steps': 2}}, mid=5)
        self.wait_for(q, lambda e: e['type'] == 'message' and e['message'].get('method') == 'notifications/message')

    def test_server_to_client_requests_are_answered_through_the_bridge(self):
        self.initialize('demo')
        q = self.events('demo')
        box = {}
        caller = threading.Thread(target=lambda: box.update(r=self.rpc('demo', 'tools/call', {'name': 'ask_client', 'arguments': {}}, mid=9)))
        caller.start()
        ask = self.wait_for(q, lambda e: e['type'] == 'message' and e['message'].get('method') == 'roots/list')['message']
        answer = {'jsonrpc': '2.0', 'id': ask['id'], 'result': {'roots': []}}
        self.assertEqual(self.request('POST', '/api/mcp/demo/rpc', answer)[0], 202)
        caller.join(10)
        self.assertIn('client answered', box['r'][1]['result']['content'][0]['text'])

    def test_the_same_id_from_two_clients_does_not_mix_up_responses(self):
        # Two tabs, each starting their own ids at 1: both must get their own answer
        self.initialize('demo')
        box = {}
        slow = threading.Thread(target=lambda: box.update(slow=self.rpc('demo', 'tools/call', {'name': 'slow', 'arguments': {'seconds': 1}}, mid=1)))
        slow.start()
        time.sleep(0.2)
        status, fast = self.rpc('demo', 'tools/call', {'name': 'add', 'arguments': {'a': 1, 'b': 1}}, mid=1)
        slow.join(10)
        self.assertEqual((status, fast['id'], fast['result']['content'][0]['text']), (200, 1, '2'))
        self.assertEqual(box['slow'][0], 200)
        self.assertEqual((box['slow'][1]['id'], box['slow'][1]['result']['content'][0]['text']), (1, 'slept 1 s'))

    def test_over_long_lines_are_dropped_and_reported(self):
        q = self.events('noisy')
        status, reply = self.rpc('noisy', 'ping', mid=4)
        self.assertEqual((status, reply['result']), (200, {'ok': True}))
        self.wait_for(q, lambda e: e['type'] == 'stderr' and '超过' in e['line'])

    def test_slow_responses_time_out(self):
        self.initialize('quick')
        status, reply = self.rpc('quick', 'tools/call', {'name': 'slow', 'arguments': {'seconds': 3}}, mid=4)
        self.assertEqual(status, 504)
        self.assertIn('秒内没有收到响应', reply['error'])

    def test_crashed_server_is_restarted_on_the_next_request(self):
        self.initialize('demo')
        status, reply = self.rpc('demo', 'tools/call', {'name': 'crash', 'arguments': {}}, mid=6)
        self.assertEqual(status, 502, reply)
        self.assertEqual(self.initialize('demo')['serverInfo']['name'], 'webtt-demo')

    # ── command-line tools ──

    def call_cli(self, name, args, mid=1):
        status, reply = self.rpc('tools', 'tools/call', {'name': name, 'arguments': args}, mid=mid)
        self.assertEqual(status, 200)
        return reply['result']

    def test_cli_tools_are_listed_as_mcp_tools(self):
        self.initialize('tools')
        _, reply = self.rpc('tools', 'tools/list', mid=2)
        names = [t['name'] for t in reply['result']['tools']]
        self.assertEqual(names, ['length', 'args', 'upper', 'exit3', 'sleepy', 'missing', 'flood', 'grep', 'calc'])
        self.assertEqual(reply['result']['tools'][0]['inputSchema']['required'], ['text'])

    def test_cli_arguments_never_reach_a_shell(self):
        payload = 'a; echo pwned $(id) `id` | cat > /tmp/x'
        result = self.call_cli('length', {'text': payload})
        self.assertEqual(result['content'][0]['text'].strip(), str(len(payload)))
        self.assertFalse(result['isError'])

    def test_cli_placeholders_optional_args_and_stdin(self):
        self.assertEqual(self.call_cli('args', {'n': 5, 'flag': True})['content'][0]['text'].strip(), "['fixed', '--n=5', 'true']")
        self.assertEqual(self.call_cli('args', {})['content'][0]['text'].strip(), "['fixed']", 'elements with missing values are dropped')
        self.assertEqual(self.call_cli('upper', {'text': 'hello'})['content'][0]['text'].strip(), 'HELLO')

    def test_cli_failures_are_tool_errors(self):
        r = self.call_cli('exit3', {})
        self.assertTrue(r['isError'])
        self.assertEqual(r['structuredContent']['exitCode'], 3)
        self.assertIn('oops', r['content'][1]['text'])
        self.assertIn('超过', self.call_cli('sleepy', {})['content'][0]['text'])
        self.assertIn('找不到程序', self.call_cli('missing', {})['content'][0]['text'])

    def test_cli_output_is_capped(self):
        r = self.call_cli('flood', {})
        self.assertLessEqual(len(r['content'][0]['text']), 100_000 + 100)
        self.assertIn('截断', r['content'][-1]['text'])

    def test_option_like_values_are_refused_unless_allowed(self):
        r = self.call_cli('grep', {'pattern': '--upload-pack=evil'})
        self.assertTrue(r['isError'])
        self.assertIn('-', r['content'][0]['text'])
        self.assertEqual(self.call_cli('grep', {'pattern': 'a-b'})['content'][0]['text'].strip(), "['a-b']")
        self.assertEqual(self.call_cli('calc', {'n': '-5'})['content'][0]['text'].strip(), "['-5']", 'opted in')
        self.assertFalse(self.call_cli('grep', {'pattern': -3})['isError'], 'numbers are not options')

    # ── Streamable HTTP ──

    def test_http_server_session_and_sse_responses(self):
        q = self.events('remote')
        self.assertEqual(self.initialize('remote')['serverInfo']['name'], 'mock')
        status, reply = self.rpc('remote', 'tools/list', mid=2)
        self.assertEqual(status, 200, reply)
        self.assertEqual(reply['result']['tools'][0]['name'], 'remote')
        self.wait_for(q, lambda e: e['type'] == 'message' and e['message'].get('params', {}).get('data') == 'from sse')


TARGET = HERE.parent / 'target'


class StdioPeer:
    """Talk to an MCP server over stdio directly, one request at a time."""

    def __init__(self, argv):
        self.proc = subprocess.Popen(argv, stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, text=True, encoding='utf-8')

    def request(self, mid, method, params=None):
        msg = {'jsonrpc': '2.0', 'id': mid, 'method': method, **({'params': params} if params is not None else {})}
        self.proc.stdin.write(json.dumps(msg) + '\n')
        self.proc.stdin.flush()
        while True:  # skip notifications and server-to-client requests
            reply = json.loads(self.proc.stdout.readline())
            if reply.get('id') == mid and 'method' not in reply:
                return reply

    def close(self):
        self.proc.stdin.close()
        self.proc.wait(10)
        self.proc.stdout.close()


def png_pixels(b64):
    """The decompressed IDAT of a PNG (the two servers compress differently)."""
    data, pos, idat = base64.b64decode(b64), 8, b''
    while pos < len(data):
        n = int.from_bytes(data[pos:pos + 4], 'big')
        if data[pos + 4:pos + 8] == b'IDAT':
            idat += data[pos + 8:pos + 8 + n]
        pos += 12 + n
    return zlib.decompress(idat)


class DemoTwinChecks:
    """A demo server in another language answers like scripts/mcp_demo_server.py.
    Subclasses set BINARY (skipped until it is built) and NAME (its serverInfo name)."""
    BINARY = NAME = None

    REQUESTS = [
        ('tools/list', None), ('tools/list', {'cursor': '3'}), ('tools/list', {'cursor': '6'}),
        ('tools/call', {'name': 'echo', 'arguments': {'text': '你好', 'times': 2}}),
        ('tools/call', {'name': 'add', 'arguments': {'a': 40, 'b': 2.5}}),
        ('tools/call', {'name': 'add', 'arguments': {'a': -1, 'b': 1}}),
        ('tools/call', {'name': 'fail', 'arguments': {}}),
        ('tools/call', {'name': 'slow', 'arguments': {'seconds': 0.1}}),
        ('tools/call', {'name': 'notify', 'arguments': {'steps': 1}}),
        ('tools/call', {'name': 'nope', 'arguments': {}}),
        ('tools/call', {'name': 'echo', 'arguments': {}}),
        ('resources/list', None), ('resources/templates/list', None),
        ('resources/read', {'uri': 'demo://readme'}), ('resources/read', {'uri': 'demo://notes/x y'}),
        ('resources/read', {'uri': 'demo://missing'}),
        ('prompts/list', None), ('prompts/get', {'name': 'review', 'arguments': {'code': 'fn main() {}', 'language': 'Rust'}}),
        ('prompts/get', {'name': 'review', 'arguments': {}}),
        ('ping', None), ('logging/setLevel', {'level': 'info'}), ('no/such/method', None),
    ]

    def setUp(self):
        self.py = StdioPeer([PY, str(HERE / 'mcp_demo_server.py')])
        self.rs = StdioPeer([str(self.BINARY)])
        self.addCleanup(self.py.close)
        self.addCleanup(self.rs.close)

    def test_answers_match_the_python_server(self):
        init = {'protocolVersion': '2025-03-26', 'capabilities': {}, 'clientInfo': {'name': 't', 'version': '0'}}
        py, rs = self.py.request(1, 'initialize', init)['result'], self.rs.request(1, 'initialize', init)['result']
        self.assertEqual(rs['serverInfo']['name'], self.NAME)
        self.assertEqual((rs['protocolVersion'], rs['capabilities']), (py['protocolVersion'], py['capabilities']))
        for mid, (method, params) in enumerate(self.REQUESTS, start=2):
            py, rs = self.py.request(mid, method, params), self.rs.request(mid, method, params)
            if 'error' in py:  # messages are worded differently; codes must agree
                self.assertEqual(rs.get('error', {}).get('code'), py['error']['code'], (method, params, rs))
            else:
                self.assertEqual(rs, py, (method, params))

    def test_image_has_the_same_pixels(self):
        args = {'name': 'image', 'arguments': {'color': 'green'}}
        py, rs = self.py.request(1, 'tools/call', args)['result'], self.rs.request(1, 'tools/call', args)['result']
        self.assertEqual(png_pixels(rs['content'][0]['data']), png_pixels(py['content'][0]['data']))
        self.assertEqual(rs['content'][1], py['content'][1])

    def test_ask_client_round_trip_and_crash(self):
        self.rs.proc.stdin.write(json.dumps({'jsonrpc': '2.0', 'id': 1, 'method': 'tools/call', 'params': {'name': 'ask_client', 'arguments': {}}}) + '\n')
        self.rs.proc.stdin.flush()
        ask = json.loads(self.rs.proc.stdout.readline())
        self.assertEqual(ask['method'], 'roots/list')
        self.rs.proc.stdin.write(json.dumps({'jsonrpc': '2.0', 'id': ask['id'], 'result': {'roots': []}}) + '\n')
        self.rs.proc.stdin.flush()
        self.assertIn('client answered', json.loads(self.rs.proc.stdout.readline())['result']['content'][0]['text'])
        self.rs.proc.stdin.write(json.dumps({'jsonrpc': '2.0', 'id': 2, 'method': 'tools/call', 'params': {'name': 'crash', 'arguments': {}}}) + '\n')
        self.rs.proc.stdin.flush()
        self.assertEqual(self.rs.proc.wait(10), 3)


BUILD_HINT = 'build it first (docs/INSTALL.md §3.2)'


@unittest.skipUnless((TARGET / 'release' / 'examples' / 'demo_server').exists(), BUILD_HINT)
class RustDemoServerTest(DemoTwinChecks, unittest.TestCase):
    BINARY, NAME = TARGET / 'release' / 'examples' / 'demo_server', 'webtt-demo-rust'


@unittest.skipUnless((TARGET / 'mcp-demo' / 'demo-go').exists(), BUILD_HINT)
class GoDemoServerTest(DemoTwinChecks, unittest.TestCase):
    BINARY, NAME = TARGET / 'mcp-demo' / 'demo-go', 'webtt-demo-go'


@unittest.skipUnless((TARGET / 'mcp-demo' / 'demo-zig').exists(), BUILD_HINT)
class ZigDemoServerTest(DemoTwinChecks, unittest.TestCase):
    BINARY, NAME = TARGET / 'mcp-demo' / 'demo-zig', 'webtt-demo-zig'


class ConfigTest(unittest.TestCase):
    def check(self, servers):
        with tempfile.NamedTemporaryFile('w', suffix='.json', delete=False) as f:
            json.dump({'servers': servers}, f)
        try:
            return mcp_bridge.load_config(f.name)
        finally:
            pathlib.Path(f.name).unlink()

    def test_invalid_configs_are_rejected(self):
        for bad in ([{'id': 'a b', 'command': ['x']}], [{'id': 'a', 'command': ['x']}, {'id': 'a', 'command': ['y']}],
                    [{'id': 'a'}], [{'id': 'a', 'command': 'x'}], [{'id': 'a', 'url': 'ftp://x'}],
                    [{'id': 'a', 'command': ['x'], 'url': 'http://x'}], [{'id': 'a', 'cli': [{'name': 'bad name', 'command': ['x']}]}],
                    [{'id': 'a', 'cli': [{'name': 't', 'command': ['{program}', 'x']}]}]):
            with self.assertRaises(ValueError, msg=bad):
                self.check(bad)

    def test_argv_template(self):
        self.assertEqual(mcp_bridge.argv_for(['p', '{a}', '--x={b}', '{c}'], {'a': 'v w', 'b': 2}), ['p', 'v w', '--x=2'])


if __name__ == '__main__':
    unittest.main()
