Skip to content
File

Blob: src/workerd/io/worker-interface.c++

17.3 KB
1// Copyright (c) 2017-2022 Cloudflare, Inc.
2// Licensed under the Apache 2.0 license found in the LICENSE file or at:
3// https://opensource.org/licenses/Apache-2.0
4 
5#include "worker-interface.h"
6 
7#include <workerd/util/http-util.h>
8#include <workerd/util/stream-utils.h>
9 
10#include <kj/debug.h>
11 
12using kj::byte;
13using kj::uint;
14 
15namespace workerd {
16 
17namespace {
18// A WorkerInterface that delays requests until some promise resolves, then forwards them to the
19// interface the promise resolved to.
20class PromisedWorkerInterface final: public WorkerInterface {
21 public:
22 PromisedWorkerInterface(kj::Promise<kj::Own<WorkerInterface>> promise)
23 : promise(promise.then([this](kj::Own<WorkerInterface> result) { worker = kj::mv(result); })
24 .fork()) {}
25 
26 kj::Promise<void> request(kj::HttpMethod method,
27 kj::StringPtr url,
28 const kj::HttpHeaders& headers,
29 kj::AsyncInputStream& requestBody,
30 Response& response) override {
31 KJ_IF_SOME(w, worker) {
32 co_await w->request(method, url, headers, requestBody, response);
33 } else {
34 co_await promise;
35 co_await KJ_ASSERT_NONNULL(worker)->request(method, url, headers, requestBody, response);
36 }
37 }
38 
39 kj::Promise<void> connect(kj::StringPtr host,
40 const kj::HttpHeaders& headers,
41 kj::AsyncIoStream& connection,
42 ConnectResponse& response,
43 kj::HttpConnectSettings settings) override {
44 KJ_IF_SOME(w, worker) {
45 co_await w->connect(host, headers, connection, response, kj::mv(settings));
46 } else {
47 co_await promise;
48 co_await KJ_ASSERT_NONNULL(worker)->connect(
49 host, headers, connection, response, kj::mv(settings));
50 }
51 }
52 
53 kj::Promise<void> prewarm(kj::StringPtr url) override {
54 KJ_IF_SOME(w, worker) {
55 co_return co_await w->prewarm(url);
56 } else {
57 co_await promise;
58 co_return co_await KJ_ASSERT_NONNULL(worker)->prewarm(url);
59 }
60 }
61 
62 kj::Promise<ScheduledResult> runScheduled(kj::Date scheduledTime, kj::StringPtr cron) override {
63 KJ_IF_SOME(wrk, worker) {
64 co_return co_await wrk->runScheduled(scheduledTime, cron);
65 } else {
66 co_await promise;
67 co_return co_await KJ_ASSERT_NONNULL(worker)->runScheduled(scheduledTime, cron);
68 }
69 }
70 
71 kj::Promise<AlarmResult> runAlarm(kj::Date scheduledTime, uint32_t retryCount) override {
72 KJ_IF_SOME(w, worker) {
73 co_return co_await w->runAlarm(scheduledTime, retryCount);
74 } else {
75 co_await promise;
76 co_return co_await KJ_ASSERT_NONNULL(worker)->runAlarm(scheduledTime, retryCount);
77 }
78 }
79 
80 kj::Promise<kj::Maybe<kj::Date>> abandonAlarm(kj::Date scheduledTime) override {
81 KJ_IF_SOME(w, worker) {
82 co_return co_await w->abandonAlarm(scheduledTime);
83 } else {
84 co_await promise;
85 co_return co_await KJ_ASSERT_NONNULL(worker)->abandonAlarm(scheduledTime);
86 }
87 }
88 
89 kj::Promise<CustomEvent::Result> customEvent(kj::Own<CustomEvent> event) override {
90 KJ_IF_SOME(w, worker) {
91 co_return co_await w->customEvent(kj::mv(event));
92 } else {
93 try {
94 co_await promise;
95 } catch (...) {
96 // Due to the exception, we're going to discard our CustomEvent. But we should tell it
97 // about why it failed first. This is important for JsRpcSessionCustomEvent in
98 // particular, as it needs to resolve the RPC client to the correct error.
99 auto exception = kj::getCaughtExceptionAsKj();
100 event->failed(exception);
101 kj::throwFatalException(kj::mv(exception));
102 }
103 co_return co_await KJ_ASSERT_NONNULL(worker)->customEvent(kj::mv(event));
104 }
105 }
106 
107 private:
108 kj::ForkedPromise<void> promise;
109 kj::Maybe<kj::Own<WorkerInterface>> worker;
110};
111} // namespace
112 
113kj::Own<WorkerInterface> newPromisedWorkerInterface(kj::Promise<kj::Own<WorkerInterface>> promise) {
114 return kj::heap<PromisedWorkerInterface>(kj::mv(promise));
115}
116 
117kj::Own<kj::HttpClient> asHttpClient(kj::Own<WorkerInterface> workerInterface) {
118 return kj::newHttpClient(*workerInterface).attach(kj::mv(workerInterface));
119}
120 
121// =======================================================================================
122namespace {
123// A Revocable WebSocket wrapper, revoked when revokeProm rejects
124class RevocableWebSocket final: public kj::WebSocket {
125 public:
126 RevocableWebSocket(kj::Own<WebSocket> ws, kj::Promise<void> revokeProm)
127 : ws(kj::mv(ws)),
128 revokeProm(revokeProm
129 .catch_([this](kj::Exception&& e) -> kj::Promise<void> {
130 canceler.cancel(e.clone());
131 KJ_IF_SOME(ws, this->ws.tryGet<kj::Own<kj::WebSocket>>()) {
132 (ws)->abort();
133 }
134 this->ws = kj::mv(e);
135 return kj::READY_NOW;
136 })
137 .eagerlyEvaluate(nullptr)) {}
138 
139 kj::Promise<void> send(kj::ArrayPtr<const byte> message) override {
140 return wrap<void>(getInner().send(message));
141 }
142 kj::Promise<void> send(kj::ArrayPtr<const char> message) override {
143 return wrap<void>(getInner().send(message));
144 }
145 
146 kj::Promise<void> close(uint16_t code, kj::StringPtr reason) override {
147 return wrap<void>(getInner().close(code, reason));
148 }
149 
150 void disconnect() override {
151 KJ_IF_SOME(ws, this->ws.tryGet<kj::Own<kj::WebSocket>>()) {
152 return (ws)->disconnect();
153 }
154 }
155 
156 void abort() override {
157 KJ_IF_SOME(ws, this->ws.tryGet<kj::Own<kj::WebSocket>>()) {
158 return (ws)->abort();
159 }
160 }
161 
162 kj::Promise<void> whenAborted() override {
163 return wrap<void>(getInner().whenAborted());
164 }
165 
166 kj::Promise<Message> receive(size_t maxSize) override {
167 return wrap<Message>(getInner().receive(maxSize));
168 }
169 
170 kj::Promise<void> pumpTo(WebSocket& other) override {
171 return wrap<void>(getInner().pumpTo(other));
172 }
173 
174 kj::Maybe<kj::Promise<void>> tryPumpFrom(WebSocket& other) override {
175 return wrap<void>(other.pumpTo(getInner()));
176 }
177 
178 kj::Maybe<kj::String> getPreferredExtensions(ExtensionsContext ctx) override {
179 return getInner().getPreferredExtensions(ctx);
180 };
181 
182 uint64_t sentByteCount() override {
183 return 0;
184 }
185 uint64_t receivedByteCount() override {
186 return 0;
187 }
188 
189 private:
190 template <typename T>
191 kj::Promise<T> wrap(kj::Promise<T> prom) {
192 // just to fix the revocation promise return type, serves no purpose otherwise
193 return canceler.wrap(kj::mv(prom));
194 }
195 
196 kj::WebSocket& getInner() {
197 KJ_SWITCH_ONEOF(ws) {
198 KJ_CASE_ONEOF(e, kj::Exception) {
199 kj::throwFatalException(e.clone());
200 }
201 KJ_CASE_ONEOF(ws, kj::Own<kj::WebSocket>) {
202 return *ws.get();
203 }
204 }
205 KJ_UNREACHABLE;
206 }
207 
208 kj::OneOf<kj::Exception, kj::Own<kj::WebSocket>> ws;
209 kj::Promise<void> revokeProm;
210 kj::Canceler canceler;
211};
212 
213// A HttpResponse that can revoke long-running websocket connections started as part of the
214// response. Ordinary HTTP requests are not revoked.
215class RevocableWebSocketHttpResponse final: public kj::HttpService::Response {
216 public:
217 RevocableWebSocketHttpResponse(kj::HttpService::Response& inner, kj::Promise<void> revokeProm)
218 : inner(inner),
219 revokeProm(revokeProm.fork()) {}
220 
221 kj::Own<kj::AsyncOutputStream> send(uint statusCode,
222 kj::StringPtr statusText,
223 const kj::HttpHeaders& headers,
224 kj::Maybe<uint64_t> expectedBodySize = kj::none) override {
225 return inner.send(statusCode, statusText, headers, expectedBodySize);
226 }
227 
228 kj::Own<kj::WebSocket> acceptWebSocket(const kj::HttpHeaders& headers) override {
229 return kj::heap<RevocableWebSocket>(inner.acceptWebSocket(headers), revokeProm.addBranch());
230 }
231 
232 private:
233 kj::HttpService::Response& inner;
234 kj::ForkedPromise<void> revokeProm;
235};
236 
237// A WorkerInterface that cancels WebSockets when revokeProm is rejected.
238// Currently only supports cancelling for upgrades.
239class RevocableWebSocketWorkerInterface final: public WorkerInterface {
240 public:
241 RevocableWebSocketWorkerInterface(WorkerInterface& worker, kj::Promise<void> revokeProm);
242 kj::Promise<void> request(kj::HttpMethod method,
243 kj::StringPtr url,
244 const kj::HttpHeaders& headers,
245 kj::AsyncInputStream& requestBody,
246 Response& response) override;
247 kj::Promise<void> connect(kj::StringPtr host,
248 const kj::HttpHeaders& headers,
249 kj::AsyncIoStream& connection,
250 ConnectResponse& response,
251 kj::HttpConnectSettings settings) override;
252 kj::Promise<void> prewarm(kj::StringPtr url) override;
253 kj::Promise<ScheduledResult> runScheduled(kj::Date scheduledTime, kj::StringPtr cron) override;
254 kj::Promise<AlarmResult> runAlarm(kj::Date scheduledTime, uint32_t retryCount) override;
255 kj::Promise<CustomEvent::Result> customEvent(kj::Own<CustomEvent> event) override;
256 
257 private:
258 WorkerInterface& worker;
259 kj::ForkedPromise<void> revokeProm;
260};
261 
262kj::Promise<void> RevocableWebSocketWorkerInterface::request(kj::HttpMethod method,
263 kj::StringPtr url,
264 const kj::HttpHeaders& headers,
265 kj::AsyncInputStream& requestBody,
266 kj::HttpService::Response& response) {
267 auto wrappedResponse = kj::heap<RevocableWebSocketHttpResponse>(response, revokeProm.addBranch());
268 return worker.request(method, url, headers, requestBody, *wrappedResponse)
269 .attach(kj::mv(wrappedResponse));
270}
271 
272kj::Promise<void> RevocableWebSocketWorkerInterface::connect(kj::StringPtr host,
273 const kj::HttpHeaders& headers,
274 kj::AsyncIoStream& connection,
275 ConnectResponse& response,
276 kj::HttpConnectSettings settings) {
277 // We give TCP sockets the same treatment as WebSockets because the purpose here is to
278 // disconnect long-running connections, e.g. on a code update for a Durable Object, and that
279 // applies equally to TCP sockets.
280 auto wrappedConnection = newNeuterableIoStream(connection);
281 auto* wrappedConnectionPtr = wrappedConnection.get();
282 auto revokeTask =
283 revokeProm.addBranch()
284 .catch_([&connection, wrappedConnectionPtr](kj::Exception&& e) -> kj::Promise<void> {
285 wrappedConnectionPtr->neuter(e.clone());
286 connection.abortWrite(kj::mv(e));
287 connection.abortRead();
288 return kj::READY_NOW;
289 }).eagerlyEvaluate(nullptr);
290 
291 return worker.connect(host, headers, *wrappedConnection, response, kj::mv(settings))
292 .attach(kj::mv(wrappedConnection), kj::mv(revokeTask));
293}
294 
295RevocableWebSocketWorkerInterface::RevocableWebSocketWorkerInterface(
296 WorkerInterface& worker, kj::Promise<void> revokeProm)
297 : worker(worker),
298 revokeProm(revokeProm.fork()) {}
299 
300kj::Promise<void> RevocableWebSocketWorkerInterface::prewarm(kj::StringPtr url) {
301 return worker.prewarm(url);
302}
303 
304kj::Promise<WorkerInterface::ScheduledResult> RevocableWebSocketWorkerInterface::runScheduled(
305 kj::Date scheduledTime, kj::StringPtr cron) {
306 return worker.runScheduled(scheduledTime, cron);
307}
308 
309kj::Promise<WorkerInterface::AlarmResult> RevocableWebSocketWorkerInterface::runAlarm(
310 kj::Date scheduledTime, uint32_t retryCount) {
311 return worker.runAlarm(scheduledTime, retryCount);
312}
313 
314kj::Promise<WorkerInterface::CustomEvent::Result> RevocableWebSocketWorkerInterface::customEvent(
315 kj::Own<CustomEvent> event) {
316 return worker.customEvent(kj::mv(event));
317}
318 
319} // namespace
320 
321kj::Own<WorkerInterface> newRevocableWebSocketWorkerInterface(
322 kj::Own<WorkerInterface> worker, kj::Promise<void> revokeProm) {
323 return kj::heap<RevocableWebSocketWorkerInterface>(*worker, kj::mv(revokeProm))
324 .attach(kj::mv(worker));
325}
326 
327// =======================================================================================
328 
329namespace {
330 
331class ErrorWorkerInterface final: public WorkerInterface {
332 public:
333 ErrorWorkerInterface(kj::Exception&& exception): exception(kj::mv(exception)) {}
334 
335 kj::Promise<void> request(kj::HttpMethod method,
336 kj::StringPtr url,
337 const kj::HttpHeaders& headers,
338 kj::AsyncInputStream& requestBody,
339 Response& response) override {
340 kj::throwFatalException(kj::mv(exception));
341 }
342 
343 kj::Promise<void> connect(kj::StringPtr host,
344 const kj::HttpHeaders& headers,
345 kj::AsyncIoStream& connection,
346 ConnectResponse& response,
347 kj::HttpConnectSettings settings) override {
348 kj::throwFatalException(kj::mv(exception));
349 }
350 
351 kj::Promise<void> prewarm(kj::StringPtr url) override {
352 // ignore
353 return kj::READY_NOW;
354 }
355 
356 kj::Promise<ScheduledResult> runScheduled(kj::Date scheduledTime, kj::StringPtr cron) override {
357 kj::throwFatalException(kj::mv(exception));
358 }
359 
360 kj::Promise<AlarmResult> runAlarm(kj::Date scheduledTime, uint32_t retryCount) override {
361 kj::throwFatalException(kj::mv(exception));
362 }
363 
364 kj::Promise<CustomEvent::Result> customEvent(kj::Own<CustomEvent> event) override {
365 kj::throwFatalException(kj::mv(exception));
366 }
367 
368 private:
369 kj::Exception exception;
370};
371 
372} // namespace
373 
374kj::Own<WorkerInterface> WorkerInterface::fromException(kj::Exception&& e) {
375 return kj::heap<ErrorWorkerInterface>(kj::mv(e));
376}
377 
378// =======================================================================================
379 
380RpcWorkerInterface::RpcWorkerInterface(capnp::HttpOverCapnpFactory& httpOverCapnpFactory,
381 capnp::ByteStreamFactory& byteStreamFactory,
382 rpc::EventDispatcher::Client dispatcher)
383 : httpOverCapnpFactory(httpOverCapnpFactory),
384 byteStreamFactory(byteStreamFactory),
385 dispatcher(kj::mv(dispatcher)) {}
386 
387kj::Promise<void> RpcWorkerInterface::request(kj::HttpMethod method,
388 kj::StringPtr url,
389 const kj::HttpHeaders& headers,
390 kj::AsyncInputStream& requestBody,
391 Response& response) {
392 auto inner = httpOverCapnpFactory.capnpToKj(dispatcher.getHttpServiceRequest().send().getHttp());
393 auto promise = inner->request(method, url, headers, requestBody, response);
394 return promise.attach(kj::mv(inner));
395}
396 
397kj::Promise<void> RpcWorkerInterface::connect(kj::StringPtr host,
398 const kj::HttpHeaders& headers,
399 kj::AsyncIoStream& connection,
400 ConnectResponse& tunnel,
401 kj::HttpConnectSettings settings) {
402 auto inner = httpOverCapnpFactory.capnpToKj(dispatcher.getHttpServiceRequest().send().getHttp());
403 auto promise = inner->connect(host, headers, connection, tunnel, kj::mv(settings));
404 return promise.attach(kj::mv(inner));
405}
406 
407kj::Promise<void> RpcWorkerInterface::prewarm(kj::StringPtr url) {
408 auto req = dispatcher.prewarmRequest(capnp::MessageSize{url.size() / sizeof(capnp::word) + 4, 0});
409 req.setUrl(url);
410 return req.sendIgnoringResult();
411}
412 
413kj::Promise<WorkerInterface::ScheduledResult> RpcWorkerInterface::runScheduled(
414 kj::Date scheduledTime, kj::StringPtr cron) {
415 auto req = dispatcher.runScheduledRequest();
416 req.setScheduledTime((scheduledTime - kj::UNIX_EPOCH) / kj::SECONDS);
417 req.setCron(cron);
418 return req.send().then([](auto resp) {
419 auto respResult = resp.getResult();
420 return WorkerInterface::ScheduledResult{
421 .retry = respResult.getRetry(), .outcome = respResult.getOutcome()};
422 });
423}
424 
425kj::Promise<WorkerInterface::AlarmResult> RpcWorkerInterface::runAlarm(
426 kj::Date scheduledTime, uint32_t retryCount) {
427 auto req = dispatcher.runAlarmRequest();
428 req.setScheduledTime((scheduledTime - kj::UNIX_EPOCH) / kj::MILLISECONDS);
429 req.setRetryCount(retryCount);
430 return req.send().then([](auto resp) {
431 auto respResult = resp.getResult();
432 kj::Maybe<kj::String> errorDescription;
433 if (respResult.hasErrorDescription()) {
434 errorDescription = kj::str(respResult.getErrorDescription());
435 }
436 return WorkerInterface::AlarmResult{.retry = respResult.getRetry(),
437 .retryCountsAgainstLimit = respResult.getRetryCountsAgainstLimit(),
438 .outcome = respResult.getOutcome(),
439 .errorDescription = kj::mv(errorDescription)};
440 });
441}
442 
443kj::Promise<kj::Maybe<kj::Date>> RpcWorkerInterface::abandonAlarm(kj::Date scheduledTime) {
444 auto req = dispatcher.abandonAlarmRequest();
445 req.setScheduledTimeMs((scheduledTime - kj::UNIX_EPOCH) / kj::MILLISECONDS);
446 auto response = co_await req.send();
447 auto storedAlarmTimeMs = response.getStoredAlarmTimeMs();
448 if (storedAlarmTimeMs != 0) {
449 co_return kj::UNIX_EPOCH + storedAlarmTimeMs* kj::MILLISECONDS;
450 }
451 co_return kj::Maybe<kj::Date>(kj::none);
452}
453 
454kj::Promise<WorkerInterface::CustomEvent::Result> RpcWorkerInterface::customEvent(
455 kj::Own<CustomEvent> event) {
456 return event->sendRpc(httpOverCapnpFactory, byteStreamFactory, dispatcher).attach(kj::mv(event));
457}
458 
459// ======================================================================================
460WorkerInterface::AlarmFulfiller::AlarmFulfiller(
461 kj::Own<kj::PromiseFulfiller<AlarmOutcome>> fulfiller)
462 : maybeFulfiller(kj::mv(fulfiller)) {}
463 
464WorkerInterface::AlarmFulfiller::~AlarmFulfiller() noexcept(false) {
465 KJ_IF_SOME(fulfiller, getFulfiller()) {
466 fulfiller.reject(KJ_EXCEPTION(FAILED, "AlarmFulfiller destroyed without resolution"));
467 }
468}
469 
470void WorkerInterface::AlarmFulfiller::fulfill(const AlarmOutcome& result) {
471 KJ_IF_SOME(fulfiller, getFulfiller()) {
472 fulfiller.fulfill(kj::cp(result));
473 }
474}
475 
476void WorkerInterface::AlarmFulfiller::reject(const kj::Exception& e) {
477 KJ_IF_SOME(fulfiller, getFulfiller()) {
478 fulfiller.reject(e.clone());
479 }
480}
481 
482void WorkerInterface::AlarmFulfiller::cancel() {
483 KJ_IF_SOME(fulfiller, getFulfiller()) {
484 fulfiller.fulfill(AlarmOutcome{
485 .retry = false,
486 .outcome = EventOutcome::CANCELED,
487 });
488 }
489}
490 
491kj::Maybe<kj::PromiseFulfiller<WorkerInterface::AlarmOutcome>&> WorkerInterface::AlarmFulfiller::
492 getFulfiller() {
493 KJ_IF_SOME(fulfiller, maybeFulfiller) {
494 if (fulfiller.get()->isWaiting()) {
495 return *fulfiller;
496 }
497 }
498 
499 return kj::none;
500}
501 
502} // namespace workerd