// Copyright (c) 2017-2022 Cloudflare, Inc. // Licensed under the Apache 2.0 license found in the LICENSE file or at: // https://opensource.org/licenses/Apache-2.0 #include "worker-interface.h" #include #include #include using kj::byte; using kj::uint; namespace workerd { namespace { // A WorkerInterface that delays requests until some promise resolves, then forwards them to the // interface the promise resolved to. class PromisedWorkerInterface final: public WorkerInterface { public: PromisedWorkerInterface(kj::Promise> promise) : promise(promise.then([this](kj::Own result) { worker = kj::mv(result); }) .fork()) {} kj::Promise request(kj::HttpMethod method, kj::StringPtr url, const kj::HttpHeaders& headers, kj::AsyncInputStream& requestBody, Response& response) override { KJ_IF_SOME(w, worker) { co_await w->request(method, url, headers, requestBody, response); } else { co_await promise; co_await KJ_ASSERT_NONNULL(worker)->request(method, url, headers, requestBody, response); } } kj::Promise connect(kj::StringPtr host, const kj::HttpHeaders& headers, kj::AsyncIoStream& connection, ConnectResponse& response, kj::HttpConnectSettings settings) override { KJ_IF_SOME(w, worker) { co_await w->connect(host, headers, connection, response, kj::mv(settings)); } else { co_await promise; co_await KJ_ASSERT_NONNULL(worker)->connect( host, headers, connection, response, kj::mv(settings)); } } kj::Promise prewarm(kj::StringPtr url) override { KJ_IF_SOME(w, worker) { co_return co_await w->prewarm(url); } else { co_await promise; co_return co_await KJ_ASSERT_NONNULL(worker)->prewarm(url); } } kj::Promise runScheduled(kj::Date scheduledTime, kj::StringPtr cron) override { KJ_IF_SOME(wrk, worker) { co_return co_await wrk->runScheduled(scheduledTime, cron); } else { co_await promise; co_return co_await KJ_ASSERT_NONNULL(worker)->runScheduled(scheduledTime, cron); } } kj::Promise runAlarm(kj::Date scheduledTime, uint32_t retryCount) override { KJ_IF_SOME(w, worker) { co_return co_await w->runAlarm(scheduledTime, retryCount); } else { co_await promise; co_return co_await KJ_ASSERT_NONNULL(worker)->runAlarm(scheduledTime, retryCount); } } kj::Promise> abandonAlarm(kj::Date scheduledTime) override { KJ_IF_SOME(w, worker) { co_return co_await w->abandonAlarm(scheduledTime); } else { co_await promise; co_return co_await KJ_ASSERT_NONNULL(worker)->abandonAlarm(scheduledTime); } } kj::Promise customEvent(kj::Own event) override { KJ_IF_SOME(w, worker) { co_return co_await w->customEvent(kj::mv(event)); } else { try { co_await promise; } catch (...) { // Due to the exception, we're going to discard our CustomEvent. But we should tell it // about why it failed first. This is important for JsRpcSessionCustomEvent in // particular, as it needs to resolve the RPC client to the correct error. auto exception = kj::getCaughtExceptionAsKj(); event->failed(exception); kj::throwFatalException(kj::mv(exception)); } co_return co_await KJ_ASSERT_NONNULL(worker)->customEvent(kj::mv(event)); } } private: kj::ForkedPromise promise; kj::Maybe> worker; }; } // namespace kj::Own newPromisedWorkerInterface(kj::Promise> promise) { return kj::heap(kj::mv(promise)); } kj::Own asHttpClient(kj::Own workerInterface) { return kj::newHttpClient(*workerInterface).attach(kj::mv(workerInterface)); } // ======================================================================================= namespace { // A Revocable WebSocket wrapper, revoked when revokeProm rejects class RevocableWebSocket final: public kj::WebSocket { public: RevocableWebSocket(kj::Own ws, kj::Promise revokeProm) : ws(kj::mv(ws)), revokeProm(revokeProm .catch_([this](kj::Exception&& e) -> kj::Promise { canceler.cancel(e.clone()); KJ_IF_SOME(ws, this->ws.tryGet>()) { (ws)->abort(); } this->ws = kj::mv(e); return kj::READY_NOW; }) .eagerlyEvaluate(nullptr)) {} kj::Promise send(kj::ArrayPtr message) override { return wrap(getInner().send(message)); } kj::Promise send(kj::ArrayPtr message) override { return wrap(getInner().send(message)); } kj::Promise close(uint16_t code, kj::StringPtr reason) override { return wrap(getInner().close(code, reason)); } void disconnect() override { KJ_IF_SOME(ws, this->ws.tryGet>()) { return (ws)->disconnect(); } } void abort() override { KJ_IF_SOME(ws, this->ws.tryGet>()) { return (ws)->abort(); } } kj::Promise whenAborted() override { return wrap(getInner().whenAborted()); } kj::Promise receive(size_t maxSize) override { return wrap(getInner().receive(maxSize)); } kj::Promise pumpTo(WebSocket& other) override { return wrap(getInner().pumpTo(other)); } kj::Maybe> tryPumpFrom(WebSocket& other) override { return wrap(other.pumpTo(getInner())); } kj::Maybe getPreferredExtensions(ExtensionsContext ctx) override { return getInner().getPreferredExtensions(ctx); }; uint64_t sentByteCount() override { return 0; } uint64_t receivedByteCount() override { return 0; } private: template kj::Promise wrap(kj::Promise prom) { // just to fix the revocation promise return type, serves no purpose otherwise return canceler.wrap(kj::mv(prom)); } kj::WebSocket& getInner() { KJ_SWITCH_ONEOF(ws) { KJ_CASE_ONEOF(e, kj::Exception) { kj::throwFatalException(e.clone()); } KJ_CASE_ONEOF(ws, kj::Own) { return *ws.get(); } } KJ_UNREACHABLE; } kj::OneOf> ws; kj::Promise revokeProm; kj::Canceler canceler; }; // A HttpResponse that can revoke long-running websocket connections started as part of the // response. Ordinary HTTP requests are not revoked. class RevocableWebSocketHttpResponse final: public kj::HttpService::Response { public: RevocableWebSocketHttpResponse(kj::HttpService::Response& inner, kj::Promise revokeProm) : inner(inner), revokeProm(revokeProm.fork()) {} kj::Own send(uint statusCode, kj::StringPtr statusText, const kj::HttpHeaders& headers, kj::Maybe expectedBodySize = kj::none) override { return inner.send(statusCode, statusText, headers, expectedBodySize); } kj::Own acceptWebSocket(const kj::HttpHeaders& headers) override { return kj::heap(inner.acceptWebSocket(headers), revokeProm.addBranch()); } private: kj::HttpService::Response& inner; kj::ForkedPromise revokeProm; }; // A WorkerInterface that cancels WebSockets when revokeProm is rejected. // Currently only supports cancelling for upgrades. class RevocableWebSocketWorkerInterface final: public WorkerInterface { public: RevocableWebSocketWorkerInterface(WorkerInterface& worker, kj::Promise revokeProm); kj::Promise request(kj::HttpMethod method, kj::StringPtr url, const kj::HttpHeaders& headers, kj::AsyncInputStream& requestBody, Response& response) override; kj::Promise connect(kj::StringPtr host, const kj::HttpHeaders& headers, kj::AsyncIoStream& connection, ConnectResponse& response, kj::HttpConnectSettings settings) override; kj::Promise prewarm(kj::StringPtr url) override; kj::Promise runScheduled(kj::Date scheduledTime, kj::StringPtr cron) override; kj::Promise runAlarm(kj::Date scheduledTime, uint32_t retryCount) override; kj::Promise customEvent(kj::Own event) override; private: WorkerInterface& worker; kj::ForkedPromise revokeProm; }; kj::Promise RevocableWebSocketWorkerInterface::request(kj::HttpMethod method, kj::StringPtr url, const kj::HttpHeaders& headers, kj::AsyncInputStream& requestBody, kj::HttpService::Response& response) { auto wrappedResponse = kj::heap(response, revokeProm.addBranch()); return worker.request(method, url, headers, requestBody, *wrappedResponse) .attach(kj::mv(wrappedResponse)); } kj::Promise RevocableWebSocketWorkerInterface::connect(kj::StringPtr host, const kj::HttpHeaders& headers, kj::AsyncIoStream& connection, ConnectResponse& response, kj::HttpConnectSettings settings) { // We give TCP sockets the same treatment as WebSockets because the purpose here is to // disconnect long-running connections, e.g. on a code update for a Durable Object, and that // applies equally to TCP sockets. auto wrappedConnection = newNeuterableIoStream(connection); auto* wrappedConnectionPtr = wrappedConnection.get(); auto revokeTask = revokeProm.addBranch() .catch_([&connection, wrappedConnectionPtr](kj::Exception&& e) -> kj::Promise { wrappedConnectionPtr->neuter(e.clone()); connection.abortWrite(kj::mv(e)); connection.abortRead(); return kj::READY_NOW; }).eagerlyEvaluate(nullptr); return worker.connect(host, headers, *wrappedConnection, response, kj::mv(settings)) .attach(kj::mv(wrappedConnection), kj::mv(revokeTask)); } RevocableWebSocketWorkerInterface::RevocableWebSocketWorkerInterface( WorkerInterface& worker, kj::Promise revokeProm) : worker(worker), revokeProm(revokeProm.fork()) {} kj::Promise RevocableWebSocketWorkerInterface::prewarm(kj::StringPtr url) { return worker.prewarm(url); } kj::Promise RevocableWebSocketWorkerInterface::runScheduled( kj::Date scheduledTime, kj::StringPtr cron) { return worker.runScheduled(scheduledTime, cron); } kj::Promise RevocableWebSocketWorkerInterface::runAlarm( kj::Date scheduledTime, uint32_t retryCount) { return worker.runAlarm(scheduledTime, retryCount); } kj::Promise RevocableWebSocketWorkerInterface::customEvent( kj::Own event) { return worker.customEvent(kj::mv(event)); } } // namespace kj::Own newRevocableWebSocketWorkerInterface( kj::Own worker, kj::Promise revokeProm) { return kj::heap(*worker, kj::mv(revokeProm)) .attach(kj::mv(worker)); } // ======================================================================================= namespace { class ErrorWorkerInterface final: public WorkerInterface { public: ErrorWorkerInterface(kj::Exception&& exception): exception(kj::mv(exception)) {} kj::Promise request(kj::HttpMethod method, kj::StringPtr url, const kj::HttpHeaders& headers, kj::AsyncInputStream& requestBody, Response& response) override { kj::throwFatalException(kj::mv(exception)); } kj::Promise connect(kj::StringPtr host, const kj::HttpHeaders& headers, kj::AsyncIoStream& connection, ConnectResponse& response, kj::HttpConnectSettings settings) override { kj::throwFatalException(kj::mv(exception)); } kj::Promise prewarm(kj::StringPtr url) override { // ignore return kj::READY_NOW; } kj::Promise runScheduled(kj::Date scheduledTime, kj::StringPtr cron) override { kj::throwFatalException(kj::mv(exception)); } kj::Promise runAlarm(kj::Date scheduledTime, uint32_t retryCount) override { kj::throwFatalException(kj::mv(exception)); } kj::Promise customEvent(kj::Own event) override { kj::throwFatalException(kj::mv(exception)); } private: kj::Exception exception; }; } // namespace kj::Own WorkerInterface::fromException(kj::Exception&& e) { return kj::heap(kj::mv(e)); } // ======================================================================================= RpcWorkerInterface::RpcWorkerInterface(capnp::HttpOverCapnpFactory& httpOverCapnpFactory, capnp::ByteStreamFactory& byteStreamFactory, rpc::EventDispatcher::Client dispatcher) : httpOverCapnpFactory(httpOverCapnpFactory), byteStreamFactory(byteStreamFactory), dispatcher(kj::mv(dispatcher)) {} kj::Promise RpcWorkerInterface::request(kj::HttpMethod method, kj::StringPtr url, const kj::HttpHeaders& headers, kj::AsyncInputStream& requestBody, Response& response) { auto inner = httpOverCapnpFactory.capnpToKj(dispatcher.getHttpServiceRequest().send().getHttp()); auto promise = inner->request(method, url, headers, requestBody, response); return promise.attach(kj::mv(inner)); } kj::Promise RpcWorkerInterface::connect(kj::StringPtr host, const kj::HttpHeaders& headers, kj::AsyncIoStream& connection, ConnectResponse& tunnel, kj::HttpConnectSettings settings) { auto inner = httpOverCapnpFactory.capnpToKj(dispatcher.getHttpServiceRequest().send().getHttp()); auto promise = inner->connect(host, headers, connection, tunnel, kj::mv(settings)); return promise.attach(kj::mv(inner)); } kj::Promise RpcWorkerInterface::prewarm(kj::StringPtr url) { auto req = dispatcher.prewarmRequest(capnp::MessageSize{url.size() / sizeof(capnp::word) + 4, 0}); req.setUrl(url); return req.sendIgnoringResult(); } kj::Promise RpcWorkerInterface::runScheduled( kj::Date scheduledTime, kj::StringPtr cron) { auto req = dispatcher.runScheduledRequest(); req.setScheduledTime((scheduledTime - kj::UNIX_EPOCH) / kj::SECONDS); req.setCron(cron); return req.send().then([](auto resp) { auto respResult = resp.getResult(); return WorkerInterface::ScheduledResult{ .retry = respResult.getRetry(), .outcome = respResult.getOutcome()}; }); } kj::Promise RpcWorkerInterface::runAlarm( kj::Date scheduledTime, uint32_t retryCount) { auto req = dispatcher.runAlarmRequest(); req.setScheduledTime((scheduledTime - kj::UNIX_EPOCH) / kj::MILLISECONDS); req.setRetryCount(retryCount); return req.send().then([](auto resp) { auto respResult = resp.getResult(); kj::Maybe errorDescription; if (respResult.hasErrorDescription()) { errorDescription = kj::str(respResult.getErrorDescription()); } return WorkerInterface::AlarmResult{.retry = respResult.getRetry(), .retryCountsAgainstLimit = respResult.getRetryCountsAgainstLimit(), .outcome = respResult.getOutcome(), .errorDescription = kj::mv(errorDescription)}; }); } kj::Promise> RpcWorkerInterface::abandonAlarm(kj::Date scheduledTime) { auto req = dispatcher.abandonAlarmRequest(); req.setScheduledTimeMs((scheduledTime - kj::UNIX_EPOCH) / kj::MILLISECONDS); auto response = co_await req.send(); auto storedAlarmTimeMs = response.getStoredAlarmTimeMs(); if (storedAlarmTimeMs != 0) { co_return kj::UNIX_EPOCH + storedAlarmTimeMs* kj::MILLISECONDS; } co_return kj::Maybe(kj::none); } kj::Promise RpcWorkerInterface::customEvent( kj::Own event) { return event->sendRpc(httpOverCapnpFactory, byteStreamFactory, dispatcher).attach(kj::mv(event)); } // ====================================================================================== WorkerInterface::AlarmFulfiller::AlarmFulfiller( kj::Own> fulfiller) : maybeFulfiller(kj::mv(fulfiller)) {} WorkerInterface::AlarmFulfiller::~AlarmFulfiller() noexcept(false) { KJ_IF_SOME(fulfiller, getFulfiller()) { fulfiller.reject(KJ_EXCEPTION(FAILED, "AlarmFulfiller destroyed without resolution")); } } void WorkerInterface::AlarmFulfiller::fulfill(const AlarmOutcome& result) { KJ_IF_SOME(fulfiller, getFulfiller()) { fulfiller.fulfill(kj::cp(result)); } } void WorkerInterface::AlarmFulfiller::reject(const kj::Exception& e) { KJ_IF_SOME(fulfiller, getFulfiller()) { fulfiller.reject(e.clone()); } } void WorkerInterface::AlarmFulfiller::cancel() { KJ_IF_SOME(fulfiller, getFulfiller()) { fulfiller.fulfill(AlarmOutcome{ .retry = false, .outcome = EventOutcome::CANCELED, }); } } kj::Maybe&> WorkerInterface::AlarmFulfiller:: getFulfiller() { KJ_IF_SOME(fulfiller, maybeFulfiller) { if (fulfiller.get()->isWaiting()) { return *fulfiller; } } return kj::none; } } // namespace workerd