|
| 1 | +"""Run a Worker configured for one of the sample MCP transports.""" |
| 2 | + |
| 3 | +import argparse |
| 4 | +import asyncio |
| 5 | +import sys |
| 6 | +from collections.abc import Callable |
| 7 | +from pathlib import Path |
| 8 | +from typing import cast |
| 9 | + |
| 10 | +from mcp import Client as MCPClient |
| 11 | +from mcp import StdioServerParameters, stdio_client |
| 12 | +from temporalio.client import Client as TemporalClient |
| 13 | +from temporalio.envconfig import ClientConfig |
| 14 | +from temporalio.mcp import MCPPlugin |
| 15 | +from temporalio.worker import Worker |
| 16 | + |
| 17 | +from server import create_server |
| 18 | +from workflow import TRANSPORTS, MCPDemoWorkflow, Transport, client_name, task_queue |
| 19 | + |
| 20 | +SERVER_PATH = Path(__file__).parent / "server.py" |
| 21 | +DEFAULT_HTTP_URL = "http://127.0.0.1:8000/mcp" |
| 22 | + |
| 23 | + |
| 24 | +def client_factory( |
| 25 | + transport: Transport, http_url: str = DEFAULT_HTTP_URL |
| 26 | +) -> Callable[[], MCPClient]: |
| 27 | + """Create the worker-side client factory for a transport.""" |
| 28 | + if transport == "in-process": |
| 29 | + return lambda: MCPClient(create_server()) |
| 30 | + if transport == "stdio": |
| 31 | + parameters = StdioServerParameters( |
| 32 | + command=sys.executable, |
| 33 | + args=[str(SERVER_PATH), "stdio"], |
| 34 | + ) |
| 35 | + return lambda: MCPClient(stdio_client(parameters)) |
| 36 | + return lambda: MCPClient(http_url) |
| 37 | + |
| 38 | + |
| 39 | +def create_plugin(transport: Transport, http_url: str = DEFAULT_HTTP_URL) -> MCPPlugin: |
| 40 | + return MCPPlugin({client_name(transport): client_factory(transport, http_url)}) |
| 41 | + |
| 42 | + |
| 43 | +async def main(transport: Transport, http_url: str) -> None: |
| 44 | + config = ClientConfig.load_client_connect_config() |
| 45 | + config.setdefault("target_host", "localhost:7233") |
| 46 | + client = await TemporalClient.connect( |
| 47 | + **config, |
| 48 | + plugins=[create_plugin(transport, http_url)], |
| 49 | + ) |
| 50 | + |
| 51 | + worker = Worker( |
| 52 | + client, |
| 53 | + task_queue=task_queue(transport), |
| 54 | + workflows=[MCPDemoWorkflow], |
| 55 | + ) |
| 56 | + print(f"Worker started for {transport}. Ctrl+C to exit.") |
| 57 | + await worker.run() |
| 58 | + |
| 59 | + |
| 60 | +if __name__ == "__main__": |
| 61 | + parser = argparse.ArgumentParser(description=__doc__) |
| 62 | + parser.add_argument("transport", choices=TRANSPORTS) |
| 63 | + parser.add_argument("--http-url", default=DEFAULT_HTTP_URL) |
| 64 | + args = parser.parse_args() |
| 65 | + asyncio.run(main(cast(Transport, args.transport), args.http_url)) |
0 commit comments