simpleTaskInteractive.js 23 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598
  1. /**
  2. * Simple interactive task server demonstrating elicitation and sampling.
  3. *
  4. * This server demonstrates the task message queue pattern from the MCP Tasks spec:
  5. * - confirm_delete: Uses elicitation to ask the user for confirmation
  6. * - write_haiku: Uses sampling to request an LLM to generate content
  7. *
  8. * Both tools use the "call-now, fetch-later" pattern where the initial call
  9. * creates a task, and the result is fetched via tasks/result endpoint.
  10. */
  11. import { randomUUID } from 'node:crypto';
  12. import { Server } from '../../server/index.js';
  13. import { createMcpExpressApp } from '../../server/express.js';
  14. import { StreamableHTTPServerTransport } from '../../server/streamableHttp.js';
  15. import { RELATED_TASK_META_KEY, ListToolsRequestSchema, CallToolRequestSchema, GetTaskRequestSchema, GetTaskPayloadRequestSchema } from '../../types.js';
  16. import { isTerminal } from '../../experimental/tasks/interfaces.js';
  17. import { InMemoryTaskStore } from '../../experimental/tasks/stores/in-memory.js';
  18. // ============================================================================
  19. // Resolver - Promise-like for passing results between async operations
  20. // ============================================================================
  21. class Resolver {
  22. constructor() {
  23. this._done = false;
  24. this._promise = new Promise((resolve, reject) => {
  25. this._resolve = resolve;
  26. this._reject = reject;
  27. });
  28. }
  29. setResult(value) {
  30. if (this._done)
  31. return;
  32. this._done = true;
  33. this._resolve(value);
  34. }
  35. setException(error) {
  36. if (this._done)
  37. return;
  38. this._done = true;
  39. this._reject(error);
  40. }
  41. wait() {
  42. return this._promise;
  43. }
  44. done() {
  45. return this._done;
  46. }
  47. }
  48. class TaskMessageQueueWithResolvers {
  49. constructor() {
  50. this.queues = new Map();
  51. this.waitResolvers = new Map();
  52. }
  53. getQueue(taskId) {
  54. let queue = this.queues.get(taskId);
  55. if (!queue) {
  56. queue = [];
  57. this.queues.set(taskId, queue);
  58. }
  59. return queue;
  60. }
  61. async enqueue(taskId, message, _sessionId, maxSize) {
  62. const queue = this.getQueue(taskId);
  63. if (maxSize !== undefined && queue.length >= maxSize) {
  64. throw new Error(`Task message queue overflow: queue size (${queue.length}) exceeds maximum (${maxSize})`);
  65. }
  66. queue.push(message);
  67. // Notify any waiters
  68. this.notifyWaiters(taskId);
  69. }
  70. async enqueueWithResolver(taskId, message, resolver, originalRequestId) {
  71. const queue = this.getQueue(taskId);
  72. const queuedMessage = {
  73. type: 'request',
  74. message,
  75. timestamp: Date.now(),
  76. resolver,
  77. originalRequestId
  78. };
  79. queue.push(queuedMessage);
  80. this.notifyWaiters(taskId);
  81. }
  82. async dequeue(taskId, _sessionId) {
  83. const queue = this.getQueue(taskId);
  84. return queue.shift();
  85. }
  86. async dequeueAll(taskId, _sessionId) {
  87. const queue = this.queues.get(taskId) ?? [];
  88. this.queues.delete(taskId);
  89. return queue;
  90. }
  91. async waitForMessage(taskId) {
  92. // Check if there are already messages
  93. const queue = this.getQueue(taskId);
  94. if (queue.length > 0)
  95. return;
  96. // Wait for a message to be added
  97. return new Promise(resolve => {
  98. let waiters = this.waitResolvers.get(taskId);
  99. if (!waiters) {
  100. waiters = [];
  101. this.waitResolvers.set(taskId, waiters);
  102. }
  103. waiters.push(resolve);
  104. });
  105. }
  106. notifyWaiters(taskId) {
  107. const waiters = this.waitResolvers.get(taskId);
  108. if (waiters) {
  109. this.waitResolvers.delete(taskId);
  110. for (const resolve of waiters) {
  111. resolve();
  112. }
  113. }
  114. }
  115. cleanup() {
  116. this.queues.clear();
  117. this.waitResolvers.clear();
  118. }
  119. }
  120. // ============================================================================
  121. // Extended task store with wait functionality
  122. // ============================================================================
  123. class TaskStoreWithNotifications extends InMemoryTaskStore {
  124. constructor() {
  125. super(...arguments);
  126. this.updateResolvers = new Map();
  127. }
  128. async updateTaskStatus(taskId, status, statusMessage, sessionId) {
  129. await super.updateTaskStatus(taskId, status, statusMessage, sessionId);
  130. this.notifyUpdate(taskId);
  131. }
  132. async storeTaskResult(taskId, status, result, sessionId) {
  133. await super.storeTaskResult(taskId, status, result, sessionId);
  134. this.notifyUpdate(taskId);
  135. }
  136. async waitForUpdate(taskId) {
  137. return new Promise(resolve => {
  138. let waiters = this.updateResolvers.get(taskId);
  139. if (!waiters) {
  140. waiters = [];
  141. this.updateResolvers.set(taskId, waiters);
  142. }
  143. waiters.push(resolve);
  144. });
  145. }
  146. notifyUpdate(taskId) {
  147. const waiters = this.updateResolvers.get(taskId);
  148. if (waiters) {
  149. this.updateResolvers.delete(taskId);
  150. for (const resolve of waiters) {
  151. resolve();
  152. }
  153. }
  154. }
  155. }
  156. // ============================================================================
  157. // Task Result Handler - delivers queued messages and routes responses
  158. // ============================================================================
  159. class TaskResultHandler {
  160. constructor(store, queue) {
  161. this.store = store;
  162. this.queue = queue;
  163. this.pendingRequests = new Map();
  164. }
  165. async handle(taskId, server, _sessionId) {
  166. while (true) {
  167. // Get fresh task state
  168. const task = await this.store.getTask(taskId);
  169. if (!task) {
  170. throw new Error(`Task not found: ${taskId}`);
  171. }
  172. // Dequeue and send all pending messages
  173. await this.deliverQueuedMessages(taskId, server, _sessionId);
  174. // If task is terminal, return result
  175. if (isTerminal(task.status)) {
  176. const result = await this.store.getTaskResult(taskId);
  177. // Add related-task metadata per spec
  178. return {
  179. ...result,
  180. _meta: {
  181. ...(result._meta || {}),
  182. [RELATED_TASK_META_KEY]: { taskId }
  183. }
  184. };
  185. }
  186. // Wait for task update or new message
  187. await this.waitForUpdate(taskId);
  188. }
  189. }
  190. async deliverQueuedMessages(taskId, server, _sessionId) {
  191. while (true) {
  192. const message = await this.queue.dequeue(taskId);
  193. if (!message)
  194. break;
  195. console.log(`[Server] Delivering queued ${message.type} message for task ${taskId}`);
  196. if (message.type === 'request') {
  197. const reqMessage = message;
  198. // Send the request via the server
  199. // Store the resolver so we can route the response back
  200. if (reqMessage.resolver && reqMessage.originalRequestId) {
  201. this.pendingRequests.set(reqMessage.originalRequestId, reqMessage.resolver);
  202. }
  203. // Send the message - for elicitation/sampling, we use the server's methods
  204. // But since we're in tasks/result context, we need to send via transport
  205. // This is simplified - in production you'd use proper message routing
  206. try {
  207. const request = reqMessage.message;
  208. let response;
  209. if (request.method === 'elicitation/create') {
  210. // Send elicitation request to client
  211. const params = request.params;
  212. response = await server.elicitInput(params);
  213. }
  214. else if (request.method === 'sampling/createMessage') {
  215. // Send sampling request to client
  216. const params = request.params;
  217. response = await server.createMessage(params);
  218. }
  219. else {
  220. throw new Error(`Unknown request method: ${request.method}`);
  221. }
  222. // Route response back to resolver
  223. if (reqMessage.resolver) {
  224. reqMessage.resolver.setResult(response);
  225. }
  226. }
  227. catch (error) {
  228. if (reqMessage.resolver) {
  229. reqMessage.resolver.setException(error instanceof Error ? error : new Error(String(error)));
  230. }
  231. }
  232. }
  233. // For notifications, we'd send them too but this example focuses on requests
  234. }
  235. }
  236. async waitForUpdate(taskId) {
  237. // Race between store update and queue message
  238. await Promise.race([this.store.waitForUpdate(taskId), this.queue.waitForMessage(taskId)]);
  239. }
  240. routeResponse(requestId, response) {
  241. const resolver = this.pendingRequests.get(requestId);
  242. if (resolver && !resolver.done()) {
  243. this.pendingRequests.delete(requestId);
  244. resolver.setResult(response);
  245. return true;
  246. }
  247. return false;
  248. }
  249. routeError(requestId, error) {
  250. const resolver = this.pendingRequests.get(requestId);
  251. if (resolver && !resolver.done()) {
  252. this.pendingRequests.delete(requestId);
  253. resolver.setException(error);
  254. return true;
  255. }
  256. return false;
  257. }
  258. }
  259. // ============================================================================
  260. // Task Session - wraps server to enqueue requests during task execution
  261. // ============================================================================
  262. class TaskSession {
  263. constructor(server, taskId, store, queue) {
  264. this.server = server;
  265. this.taskId = taskId;
  266. this.store = store;
  267. this.queue = queue;
  268. this.requestCounter = 0;
  269. }
  270. nextRequestId() {
  271. return `task-${this.taskId}-${++this.requestCounter}`;
  272. }
  273. async elicit(message, requestedSchema) {
  274. // Update task status to input_required
  275. await this.store.updateTaskStatus(this.taskId, 'input_required');
  276. const requestId = this.nextRequestId();
  277. // Build the elicitation request with related-task metadata
  278. const params = {
  279. message,
  280. requestedSchema,
  281. mode: 'form',
  282. _meta: {
  283. [RELATED_TASK_META_KEY]: { taskId: this.taskId }
  284. }
  285. };
  286. const jsonrpcRequest = {
  287. jsonrpc: '2.0',
  288. id: requestId,
  289. method: 'elicitation/create',
  290. params
  291. };
  292. // Create resolver to wait for response
  293. const resolver = new Resolver();
  294. // Enqueue the request
  295. await this.queue.enqueueWithResolver(this.taskId, jsonrpcRequest, resolver, requestId);
  296. try {
  297. // Wait for response
  298. const response = await resolver.wait();
  299. // Update status back to working
  300. await this.store.updateTaskStatus(this.taskId, 'working');
  301. return response;
  302. }
  303. catch (error) {
  304. await this.store.updateTaskStatus(this.taskId, 'working');
  305. throw error;
  306. }
  307. }
  308. async createMessage(messages, maxTokens) {
  309. // Update task status to input_required
  310. await this.store.updateTaskStatus(this.taskId, 'input_required');
  311. const requestId = this.nextRequestId();
  312. // Build the sampling request with related-task metadata
  313. const params = {
  314. messages,
  315. maxTokens,
  316. _meta: {
  317. [RELATED_TASK_META_KEY]: { taskId: this.taskId }
  318. }
  319. };
  320. const jsonrpcRequest = {
  321. jsonrpc: '2.0',
  322. id: requestId,
  323. method: 'sampling/createMessage',
  324. params
  325. };
  326. // Create resolver to wait for response
  327. const resolver = new Resolver();
  328. // Enqueue the request
  329. await this.queue.enqueueWithResolver(this.taskId, jsonrpcRequest, resolver, requestId);
  330. try {
  331. // Wait for response
  332. const response = await resolver.wait();
  333. // Update status back to working
  334. await this.store.updateTaskStatus(this.taskId, 'working');
  335. return response;
  336. }
  337. catch (error) {
  338. await this.store.updateTaskStatus(this.taskId, 'working');
  339. throw error;
  340. }
  341. }
  342. }
  343. // ============================================================================
  344. // Server Setup
  345. // ============================================================================
  346. const PORT = process.env.PORT ? parseInt(process.env.PORT, 10) : 8000;
  347. // Create shared stores
  348. const taskStore = new TaskStoreWithNotifications();
  349. const messageQueue = new TaskMessageQueueWithResolvers();
  350. const taskResultHandler = new TaskResultHandler(taskStore, messageQueue);
  351. // Track active task executions
  352. const activeTaskExecutions = new Map();
  353. // Create the server
  354. const createServer = () => {
  355. const server = new Server({ name: 'simple-task-interactive', version: '1.0.0' }, {
  356. capabilities: {
  357. tools: {},
  358. tasks: {
  359. requests: {
  360. tools: { call: {} }
  361. }
  362. }
  363. }
  364. });
  365. // Register tools
  366. server.setRequestHandler(ListToolsRequestSchema, async () => {
  367. return {
  368. tools: [
  369. {
  370. name: 'confirm_delete',
  371. description: 'Asks for confirmation before deleting (demonstrates elicitation)',
  372. inputSchema: {
  373. type: 'object',
  374. properties: {
  375. filename: { type: 'string' }
  376. }
  377. },
  378. execution: { taskSupport: 'required' }
  379. },
  380. {
  381. name: 'write_haiku',
  382. description: 'Asks LLM to write a haiku (demonstrates sampling)',
  383. inputSchema: {
  384. type: 'object',
  385. properties: {
  386. topic: { type: 'string' }
  387. }
  388. },
  389. execution: { taskSupport: 'required' }
  390. }
  391. ]
  392. };
  393. });
  394. // Handle tool calls
  395. server.setRequestHandler(CallToolRequestSchema, async (request, extra) => {
  396. const { name, arguments: args } = request.params;
  397. const taskParams = (request.params._meta?.task || request.params.task);
  398. // Validate task mode - these tools require tasks
  399. if (!taskParams) {
  400. throw new Error(`Tool ${name} requires task mode`);
  401. }
  402. // Create task
  403. const taskOptions = {
  404. ttl: taskParams.ttl,
  405. pollInterval: taskParams.pollInterval ?? 1000
  406. };
  407. const task = await taskStore.createTask(taskOptions, extra.requestId, request, extra.sessionId);
  408. console.log(`\n[Server] ${name} called, task created: ${task.taskId}`);
  409. // Start background task execution
  410. const taskExecution = (async () => {
  411. try {
  412. const taskSession = new TaskSession(server, task.taskId, taskStore, messageQueue);
  413. if (name === 'confirm_delete') {
  414. const filename = args?.filename ?? 'unknown.txt';
  415. console.log(`[Server] confirm_delete: asking about '${filename}'`);
  416. console.log('[Server] Sending elicitation request to client...');
  417. const result = await taskSession.elicit(`Are you sure you want to delete '${filename}'?`, {
  418. type: 'object',
  419. properties: {
  420. confirm: { type: 'boolean' }
  421. },
  422. required: ['confirm']
  423. });
  424. console.log(`[Server] Received elicitation response: action=${result.action}, content=${JSON.stringify(result.content)}`);
  425. let text;
  426. if (result.action === 'accept' && result.content) {
  427. const confirmed = result.content.confirm;
  428. text = confirmed ? `Deleted '${filename}'` : 'Deletion cancelled';
  429. }
  430. else {
  431. text = 'Deletion cancelled';
  432. }
  433. console.log(`[Server] Completing task with result: ${text}`);
  434. await taskStore.storeTaskResult(task.taskId, 'completed', {
  435. content: [{ type: 'text', text }]
  436. });
  437. }
  438. else if (name === 'write_haiku') {
  439. const topic = args?.topic ?? 'nature';
  440. console.log(`[Server] write_haiku: topic '${topic}'`);
  441. console.log('[Server] Sending sampling request to client...');
  442. const result = await taskSession.createMessage([
  443. {
  444. role: 'user',
  445. content: { type: 'text', text: `Write a haiku about ${topic}` }
  446. }
  447. ], 50);
  448. let haiku = 'No response';
  449. if (result.content && 'text' in result.content) {
  450. haiku = result.content.text;
  451. }
  452. console.log(`[Server] Received sampling response: ${haiku.substring(0, 50)}...`);
  453. console.log('[Server] Completing task with haiku');
  454. await taskStore.storeTaskResult(task.taskId, 'completed', {
  455. content: [{ type: 'text', text: `Haiku:\n${haiku}` }]
  456. });
  457. }
  458. }
  459. catch (error) {
  460. console.error(`[Server] Task ${task.taskId} failed:`, error);
  461. await taskStore.storeTaskResult(task.taskId, 'failed', {
  462. content: [{ type: 'text', text: `Error: ${error}` }],
  463. isError: true
  464. });
  465. }
  466. finally {
  467. activeTaskExecutions.delete(task.taskId);
  468. }
  469. })();
  470. activeTaskExecutions.set(task.taskId, {
  471. promise: taskExecution,
  472. server,
  473. sessionId: extra.sessionId ?? ''
  474. });
  475. return { task };
  476. });
  477. // Handle tasks/get
  478. server.setRequestHandler(GetTaskRequestSchema, async (request) => {
  479. const { taskId } = request.params;
  480. const task = await taskStore.getTask(taskId);
  481. if (!task) {
  482. throw new Error(`Task ${taskId} not found`);
  483. }
  484. return task;
  485. });
  486. // Handle tasks/result
  487. server.setRequestHandler(GetTaskPayloadRequestSchema, async (request, extra) => {
  488. const { taskId } = request.params;
  489. console.log(`[Server] tasks/result called for task ${taskId}`);
  490. return taskResultHandler.handle(taskId, server, extra.sessionId ?? '');
  491. });
  492. return server;
  493. };
  494. // ============================================================================
  495. // Express App Setup
  496. // ============================================================================
  497. const app = createMcpExpressApp();
  498. // Map to store transports by session ID
  499. const transports = {};
  500. // Helper to check if request is initialize
  501. const isInitializeRequest = (body) => {
  502. return typeof body === 'object' && body !== null && 'method' in body && body.method === 'initialize';
  503. };
  504. // MCP POST endpoint
  505. app.post('/mcp', async (req, res) => {
  506. const sessionId = req.headers['mcp-session-id'];
  507. try {
  508. let transport;
  509. if (sessionId && transports[sessionId]) {
  510. transport = transports[sessionId];
  511. }
  512. else if (!sessionId && isInitializeRequest(req.body)) {
  513. transport = new StreamableHTTPServerTransport({
  514. sessionIdGenerator: () => randomUUID(),
  515. onsessioninitialized: sid => {
  516. console.log(`Session initialized: ${sid}`);
  517. transports[sid] = transport;
  518. }
  519. });
  520. transport.onclose = () => {
  521. const sid = transport.sessionId;
  522. if (sid && transports[sid]) {
  523. console.log(`Transport closed for session ${sid}`);
  524. delete transports[sid];
  525. }
  526. };
  527. const server = createServer();
  528. await server.connect(transport);
  529. await transport.handleRequest(req, res, req.body);
  530. return;
  531. }
  532. else {
  533. res.status(400).json({
  534. jsonrpc: '2.0',
  535. error: { code: -32000, message: 'Bad Request: No valid session ID' },
  536. id: null
  537. });
  538. return;
  539. }
  540. await transport.handleRequest(req, res, req.body);
  541. }
  542. catch (error) {
  543. console.error('Error handling MCP request:', error);
  544. if (!res.headersSent) {
  545. res.status(500).json({
  546. jsonrpc: '2.0',
  547. error: { code: -32603, message: 'Internal server error' },
  548. id: null
  549. });
  550. }
  551. }
  552. });
  553. // Handle GET requests for SSE streams
  554. app.get('/mcp', async (req, res) => {
  555. const sessionId = req.headers['mcp-session-id'];
  556. if (!sessionId || !transports[sessionId]) {
  557. res.status(400).send('Invalid or missing session ID');
  558. return;
  559. }
  560. const transport = transports[sessionId];
  561. await transport.handleRequest(req, res);
  562. });
  563. // Handle DELETE requests for session termination
  564. app.delete('/mcp', async (req, res) => {
  565. const sessionId = req.headers['mcp-session-id'];
  566. if (!sessionId || !transports[sessionId]) {
  567. res.status(400).send('Invalid or missing session ID');
  568. return;
  569. }
  570. console.log(`Session termination request: ${sessionId}`);
  571. const transport = transports[sessionId];
  572. await transport.handleRequest(req, res);
  573. });
  574. // Start server
  575. app.listen(PORT, () => {
  576. console.log(`Starting server on http://localhost:${PORT}/mcp`);
  577. console.log('\nAvailable tools:');
  578. console.log(' - confirm_delete: Demonstrates elicitation (asks user y/n)');
  579. console.log(' - write_haiku: Demonstrates sampling (requests LLM completion)');
  580. });
  581. // Handle shutdown
  582. process.on('SIGINT', async () => {
  583. console.log('\nShutting down server...');
  584. for (const sessionId of Object.keys(transports)) {
  585. try {
  586. await transports[sessionId].close();
  587. delete transports[sessionId];
  588. }
  589. catch (error) {
  590. console.error(`Error closing session ${sessionId}:`, error);
  591. }
  592. }
  593. taskStore.cleanup();
  594. messageQueue.cleanup();
  595. console.log('Server shutdown complete');
  596. process.exit(0);
  597. });
  598. //# sourceMappingURL=simpleTaskInteractive.js.map