Skip to content
File

Blob: src/workerd/util/abortable.h

cpp152 lines
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#pragma once
6#include "canceler.h"
7 
8#include <kj/compat/http.h>
9 
10namespace workerd {
11 
12template <typename T>
13class AbortableImpl final {
14 public:
15 AbortableImpl(kj::Own<T> inner, RefcountedCanceler& canceler)
16 : canceler(kj::addRef(canceler)),
17 inner(kj::mv(inner)),
18 onCancel(*(this->canceler), [this]() { this->inner = kj::none; }) {}
19 
20 template <typename V, typename... Args, typename... ArgsT>
21 kj::Promise<V> wrap(kj::Promise<V> (T::*fn)(ArgsT...), Args&&... args) {
22 return wrap([&](T& inner) { return (inner.*fn)(kj::fwd<ArgsT>(args)...); });
23 }
24 
25 template <typename Func>
26 auto wrap(Func fn) -> decltype(fn(kj::instance<T&>())) {
27 // Be aware that the getInner() here can throw synchronously if the
28 // canceler has already been tripped.
29 return canceler->wrap(fn(getInner()));
30 }
31 
32 T& getInner() {
33 canceler->throwIfCanceled();
34 // If we get past throwIfCanceled successfully, inner should still
35 // be set. If it's not, then we've got a bug somewhere and we need
36 // to know about it.
37 return *(KJ_ASSERT_NONNULL(inner));
38 }
39 
40 kj::Maybe<T&> tryGetInner() {
41 return inner;
42 }
43 
44 private:
45 kj::Own<RefcountedCanceler> canceler;
46 kj::Maybe<kj::Own<T>> inner;
47 RefcountedCanceler::Listener onCancel;
48};
49 
50// An InputStream that can be disconnected in response to RefcountedCanceler.
51// This is similar to NeuterableInputStream in global-scope.c++ but uses an
52// external kj::Canceler to trigger the disconnect.
53// This is currently only used in fetch() requests that use an AbortSignal.
54// The AbortableInputStream is created using a RefcountedCanceler,
55// which will be triggered when the AbortSignal is triggered.
56// TODO(later): It would be good to see if both this and NeuterableInputStream
57// could be combined into a single utility.
58class AbortableInputStream final: public kj::AsyncInputStream, public kj::Refcounted {
59 public:
60 AbortableInputStream(kj::Own<kj::AsyncInputStream> inner, RefcountedCanceler& canceler)
61 : impl(kj::mv(inner), canceler) {}
62 
63 kj::Promise<size_t> tryRead(void* buffer, size_t minBytes, size_t maxBytes) override {
64 kj::Promise<size_t> (kj::AsyncInputStream::*tryRead)(void*, size_t, size_t) =
65 &kj::AsyncInputStream::tryRead;
66 return impl.wrap(tryRead, buffer, minBytes, maxBytes);
67 }
68 
69 kj::Maybe<uint64_t> tryGetLength() override {
70 return impl.getInner().tryGetLength();
71 }
72 
73 kj::Promise<uint64_t> pumpTo(kj::AsyncOutputStream& output, uint64_t amount) override {
74 return impl.wrap(&kj::AsyncInputStream::pumpTo, output, amount);
75 }
76 
77 private:
78 AbortableImpl<kj::AsyncInputStream> impl;
79};
80 
81// A WebSocket wrapper that can be disconnected in response to a RefcountedCanceler.
82// This is currently only used when opening a WebSocket with a fetch() request that
83// is using an AbortSignal. The AbortableWebSocket is created using the AbortSignal's
84// RefcountedCanceler, which will be triggered when the AbortSignal is triggered.
85class AbortableWebSocket final: public kj::WebSocket, public kj::Refcounted {
86 public:
87 AbortableWebSocket(kj::Own<kj::WebSocket> inner, RefcountedCanceler& canceler)
88 : impl(kj::mv(inner), canceler) {}
89 
90 kj::Promise<void> send(kj::ArrayPtr<const kj::byte> message) override {
91 return impl.wrap(
92 static_cast<kj::Promise<void> (kj::WebSocket::*)(kj::ArrayPtr<const kj::byte>)>(
93 &kj::WebSocket::send),
94 message);
95 }
96 
97 kj::Promise<void> send(kj::ArrayPtr<const char> message) override {
98 return impl.wrap(static_cast<kj::Promise<void> (kj::WebSocket::*)(kj::ArrayPtr<const char>)>(
99 &kj::WebSocket::send),
100 message);
101 }
102 
103 kj::Promise<void> close(uint16_t code, kj::StringPtr reason) override {
104 return impl.wrap(&kj::WebSocket::close, code, reason);
105 }
106 
107 void disconnect() override {
108 KJ_IF_SOME(inner, impl.tryGetInner()) {
109 inner.disconnect();
110 }
111 }
112 
113 void abort() override {
114 KJ_IF_SOME(inner, impl.tryGetInner()) {
115 inner.abort();
116 }
117 }
118 
119 kj::Promise<void> whenAborted() override {
120 return impl.wrap(&kj::WebSocket::whenAborted);
121 }
122 
123 kj::Promise<Message> receive(size_t maxSize = SUGGESTED_MAX_MESSAGE_SIZE) override {
124 return impl.wrap(&kj::WebSocket::receive, maxSize);
125 }
126 
127 kj::Promise<void> pumpTo(kj::WebSocket& other) override {
128 return impl.wrap(&kj::WebSocket::pumpTo, other);
129 }
130 
131 kj::Maybe<kj::Promise<void>> tryPumpFrom(kj::WebSocket& other) override {
132 return impl.wrap([&other](auto& inner) -> kj::Promise<void> { return other.pumpTo(inner); });
133 }
134 
135 uint64_t sentByteCount() override {
136 return impl.getInner().sentByteCount();
137 }
138 
139 uint64_t receivedByteCount() override {
140 return impl.getInner().receivedByteCount();
141 }
142 
143 kj::Maybe<kj::String> getPreferredExtensions(ExtensionsContext ctx) override {
144 return impl.getInner().getPreferredExtensions(ctx);
145 };
146 
147 private:
148 AbortableImpl<kj::WebSocket> impl;
149};
150 
151} // namespace workerd