| 1: | <?php |
| 2: | |
| 3: | declare(strict_types=1); |
| 4: | |
| 5: | |
| 6: | |
| 7: | |
| 8: | |
| 9: | |
| 10: | |
| 11: | |
| 12: | |
| 13: | |
| 14: | namespace Nexus\Mcp\Extension\Tasks\Client; |
| 15: | |
| 16: | use Amp\Cancellation; |
| 17: | use Nexus\Assert\Assert; |
| 18: | use Nexus\Mcp\Client\Client; |
| 19: | use Nexus\Mcp\Client\Time\CancellableDelayInterface; |
| 20: | use Nexus\Mcp\Client\Time\EventLoopDelay; |
| 21: | use Nexus\Mcp\Core\Schema\Request\CallToolRequest; |
| 22: | use Nexus\Mcp\Core\Schema\RequestParams\CallToolRequestParams; |
| 23: | use Nexus\Mcp\Core\Schema\Result\CallToolResult; |
| 24: | use Nexus\Mcp\Core\Schema\Result\InputRequiredResult; |
| 25: | use Nexus\Mcp\Core\Schema\ResultResponse\GenericResultResponse; |
| 26: | use Nexus\Mcp\Extension\Tasks\Client\Exception\StalledTaskException; |
| 27: | use Nexus\Mcp\Extension\Tasks\Schema\Enum\TaskStatus; |
| 28: | use Nexus\Mcp\Extension\Tasks\Schema\Request\CancelTaskRequest; |
| 29: | use Nexus\Mcp\Extension\Tasks\Schema\Request\GetTaskRequest; |
| 30: | use Nexus\Mcp\Extension\Tasks\Schema\Request\UpdateTaskRequest; |
| 31: | use Nexus\Mcp\Extension\Tasks\Schema\RequestParams\CancelTaskRequestParams; |
| 32: | use Nexus\Mcp\Extension\Tasks\Schema\RequestParams\GetTaskRequestParams; |
| 33: | use Nexus\Mcp\Extension\Tasks\Schema\RequestParams\UpdateTaskRequestParams; |
| 34: | use Nexus\Mcp\Extension\Tasks\Schema\Result\CreateTaskResult; |
| 35: | use Nexus\Mcp\Extension\Tasks\Schema\Result\GetTaskResult; |
| 36: | use Nexus\Mcp\Extension\Tasks\Schema\ResultResponse\GetTaskResultResponse; |
| 37: | use Nexus\Mcp\Extension\Tasks\Schema\ResultResponse\TaskCallToolResultResponse; |
| 38: | |
| 39: | |
| 40: | |
| 41: | |
| 42: | final readonly class TaskClient implements TaskClientInterface |
| 43: | { |
| 44: | public const int DEFAULT_MIN_POLL_INTERVAL_MS = 100; |
| 45: | |
| 46: | |
| 47: | |
| 48: | |
| 49: | public const int MAX_POLL_INTERVAL_MS = 3_600_000; |
| 50: | |
| 51: | |
| 52: | |
| 53: | |
| 54: | |
| 55: | |
| 56: | public function __construct( |
| 57: | private Client $client, |
| 58: | private int $stallCeiling = 60, |
| 59: | private CancellableDelayInterface $delay = new EventLoopDelay(), |
| 60: | private int $minPollIntervalMs = self::DEFAULT_MIN_POLL_INTERVAL_MS, |
| 61: | ) { |
| 62: | Assert::that($stallCeiling)->isPositiveInt('Task stall ceiling must be a positive integer, {value} given.'); |
| 63: | Assert::that($minPollIntervalMs)->isPositiveInt('Task minimum poll interval must be a positive integer, {value} given.'); |
| 64: | } |
| 65: | |
| 66: | #[\Override] |
| 67: | public function callToolAsTask(string $name, ?array $arguments = null, ?array $inputResponses = null, ?string $requestState = null): CallToolResult|CreateTaskResult|InputRequiredResult |
| 68: | { |
| 69: | $response = $this->client->sendRequest( |
| 70: | new CallToolRequest( |
| 71: | id: $this->client->mintRequestId(), |
| 72: | params: new CallToolRequestParams( |
| 73: | name: $name, |
| 74: | meta: $this->client->stampMeta(), |
| 75: | arguments: $arguments, |
| 76: | inputResponses: $inputResponses, |
| 77: | requestState: $requestState, |
| 78: | ), |
| 79: | ), |
| 80: | TaskCallToolResultResponse::class, |
| 81: | ); |
| 82: | |
| 83: | return $response->result; |
| 84: | } |
| 85: | |
| 86: | #[\Override] |
| 87: | public function getTask(string $taskId): GetTaskResult |
| 88: | { |
| 89: | $response = $this->client->sendRequest( |
| 90: | new GetTaskRequest( |
| 91: | id: $this->client->mintRequestId(), |
| 92: | params: new GetTaskRequestParams(taskId: $taskId, meta: $this->client->stampMeta()), |
| 93: | ), |
| 94: | GetTaskResultResponse::class, |
| 95: | ); |
| 96: | |
| 97: | return $response->result; |
| 98: | } |
| 99: | |
| 100: | #[\Override] |
| 101: | public function updateTask(string $taskId, array $inputResponses): void |
| 102: | { |
| 103: | $this->client->sendRequest( |
| 104: | new UpdateTaskRequest( |
| 105: | id: $this->client->mintRequestId(), |
| 106: | params: new UpdateTaskRequestParams(taskId: $taskId, inputResponses: $inputResponses, meta: $this->client->stampMeta()), |
| 107: | ), |
| 108: | GenericResultResponse::class, |
| 109: | ); |
| 110: | } |
| 111: | |
| 112: | #[\Override] |
| 113: | public function cancelTask(string $taskId): void |
| 114: | { |
| 115: | $this->client->sendRequest( |
| 116: | new CancelTaskRequest( |
| 117: | id: $this->client->mintRequestId(), |
| 118: | params: new CancelTaskRequestParams(taskId: $taskId, meta: $this->client->stampMeta()), |
| 119: | ), |
| 120: | GenericResultResponse::class, |
| 121: | ); |
| 122: | } |
| 123: | |
| 124: | #[\Override] |
| 125: | public function awaitTask(CreateTaskResult $task, ?\Closure $resolveInputRequests = null, ?Cancellation $cancellation = null): GetTaskResult |
| 126: | { |
| 127: | $intervalMs = $task->pollIntervalMs ?? 1_000; |
| 128: | $answeredKeys = []; |
| 129: | $stalledPolls = 0; |
| 130: | |
| 131: | while (true) { |
| 132: | $state = $this->getTask($task->taskId); |
| 133: | $intervalMs = $state->pollIntervalMs ?? $intervalMs; |
| 134: | |
| 135: | if (TaskStatus::InputRequired === $state->status) { |
| 136: | $unanswered = array_diff_key($state->inputRequests ?? [], $answeredKeys); |
| 137: | $responses = [] === $unanswered || null === $resolveInputRequests |
| 138: | ? [] |
| 139: | : array_intersect_key($resolveInputRequests($unanswered), $unanswered); |
| 140: | |
| 141: | if ([] !== $responses) { |
| 142: | $this->updateTask($task->taskId, $responses); |
| 143: | $answeredKeys += array_fill_keys(array_keys($responses), true); |
| 144: | $stalledPolls = 0; |
| 145: | } elseif (++$stalledPolls >= $this->stallCeiling) { |
| 146: | throw new StalledTaskException($task->taskId, $stalledPolls); |
| 147: | } |
| 148: | } elseif (TaskStatus::Working !== $state->status) { |
| 149: | return $state; |
| 150: | } else { |
| 151: | $answeredKeys = []; |
| 152: | $stalledPolls = 0; |
| 153: | } |
| 154: | |
| 155: | $this->delay->sleep( |
| 156: | min(max($intervalMs, $this->minPollIntervalMs), self::MAX_POLL_INTERVAL_MS) / 1_000.0, |
| 157: | $cancellation, |
| 158: | ); |
| 159: | } |
| 160: | } |
| 161: | } |
| 162: | |