Skip to content
File

Blob: src/workerd/io/hibernation-manager.c++

16.2 KB
1// Copyright (c) 2017-2023 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 "hibernation-manager.h"
6 
7#include "io-channels.h"
8 
9#include <workerd/util/uuid.h>
10 
11namespace workerd {
12 
13HibernationManagerImpl::HibernatableWebSocket::HibernatableWebSocket(
14 jsg::Ref<api::WebSocket> websocket,
15 kj::ArrayPtr<kj::String> tags,
16 HibernationManagerImpl& manager)
17 : tagItems(kj::heapArray<TagListItem>(tags.size())),
18 activeOrPackage(kj::mv(websocket)),
19 // The `ws` starts off empty because we need to set up our tagging infrastructure before
20 // calling api::WebSocket::acceptAsHibernatable(). We will transfer ownership of the
21 // kj::WebSocket prior to starting the readLoop.
22 ws(kj::none),
23 manager(manager) {}
24 
25HibernationManagerImpl::HibernatableWebSocket::~HibernatableWebSocket() noexcept(false) {
26 // We expect this dtor to be called when we're removing a HibernatableWebSocket
27 // from our `allWs` collection in the HibernationManager.
28 
29 // This removal is fast because we have direct access to each kj::List, as well as direct
30 // access to each TagListItem we want to remove.
31 for (auto& item: tagItems) {
32 KJ_IF_SOME(list, item.list) {
33 // The list reference is non-null, so we still have a valid reference to this
34 // TagListItem in the list, which we will now remove.
35 list.remove(item);
36 if (list.empty()) {
37 // Remove the bucket in tagToWs if the tag has no more websockets.
38 manager.tagToWs.erase(kj::mv(item.tag));
39 }
40 }
41 item.hibWS = kj::none;
42 item.list = kj::none;
43 }
44}
45 
46kj::Array<kj::StringPtr> HibernationManagerImpl::HibernatableWebSocket::getTags() {
47 auto tags = kj::heapArray<kj::StringPtr>(tagItems.size());
48 for (auto i: kj::indices(tagItems)) {
49 tags[i] = tagItems[i].tag;
50 }
51 return tags;
52}
53 
54kj::Array<kj::String> HibernationManagerImpl::HibernatableWebSocket::cloneTags() {
55 auto tags = kj::heapArray<kj::String>(tagItems.size());
56 for (auto i: kj::indices(tagItems)) {
57 tags[i] = kj::str(tagItems[i].tag);
58 }
59 return tags;
60}
61 
62jsg::Ref<api::WebSocket> HibernationManagerImpl::HibernatableWebSocket::getActiveOrUnhibernate(
63 jsg::Lock& js) {
64 KJ_IF_SOME(package, activeOrPackage.tryGet<api::WebSocket::HibernationPackage>()) {
65 // Recreate our tags array for the api::WebSocket.
66 package.maybeTags = getTags();
67 
68 // Now that we unhibernated the WebSocket, we can set the last received autoResponse timestamp
69 // that was stored in the corresponding HibernatableWebSocket. We also move autoResponsePromise
70 // from the hibernation manager to api::websocket to prevent possible ws.send races.
71 activeOrPackage
72 .init<jsg::Ref<api::WebSocket>>(
73 api::WebSocket::hibernatableFromNative(js, *KJ_REQUIRE_NONNULL(ws), kj::mv(package)))
74 ->setAutoResponseStatus(autoResponseTimestamp, kj::mv(autoResponsePromise));
75 autoResponsePromise = kj::READY_NOW;
76 }
77 return activeOrPackage.get<jsg::Ref<api::WebSocket>>().addRef();
78}
79 
80HibernationManagerImpl::HibernationManagerImpl(
81 kj::Own<Worker::Actor::Loopback> loopback, uint16_t hibernationEventType)
82 : loopback(kj::mv(loopback)),
83 hibernationEventType(hibernationEventType),
84 onDisconnect(DisconnectHandler{}),
85 readLoopTasks(onDisconnect) {}
86 
87HibernationManagerImpl::~HibernationManagerImpl() noexcept(false) {
88 // Drop our outstanding tasks, the `readLoopTasks` have weak references to the
89 // `HibernatableWebSockets` in `allWs`, and since we're about to drop all of those WebSockets,
90 // we can't allow any more events to be delivered.
91 readLoopTasks.clear();
92 
93 // Note that the HibernatableWebSocket destructor handles removing any references to itself in
94 // `tagToWs`, and even removes the hashmap entry if there are no more entries in the bucket.
95 allWs.clear();
96 KJ_ASSERT(tagToWs.size() == 0, "tagToWs hashmap wasn't cleared.");
97}
98 
99kj::Own<Worker::Actor::HibernationManager> HibernationManagerImpl::addRef() {
100 return kj::addRef(*this);
101}
102 
103void HibernationManagerImpl::acceptWebSocket(
104 jsg::Ref<api::WebSocket> ws, kj::ArrayPtr<kj::String> tags) {
105 // First, we create the HibernatableWebSocket and add it to the collection where it'll stay
106 // until it's destroyed.
107 
108 JSG_REQUIRE(allWs.size() < ACTIVE_CONNECTION_LIMIT, Error, "only ", ACTIVE_CONNECTION_LIMIT,
109 " websockets can be accepted on a single Durable Object instance");
110 
111 auto hib = kj::heap<HibernatableWebSocket>(kj::mv(ws), tags, *this);
112 HibernatableWebSocket& refToHibernatable = *hib.get();
113 allWs.push_front(kj::mv(hib));
114 refToHibernatable.node = allWs.begin();
115 
116 // If the `tags` array is empty (i.e. user did not provide a tag), we skip the population of the
117 // `tagToWs` HashMap below and go straight to initiating the readLoop.
118 
119 // It is the caller's responsibility to ensure all elements of `tags` are unique.
120 // TODO(cleanup): Maybe we could enforce uniqueness by using an immutable type that
121 // can only be constructed if the elements in the collection are distinct, ex. "DistinctArray".
122 //
123 // We need to add the HibernatableWebSocket to each bucket in `tagToWs` corresponding to its tags.
124 // 1. Create the entry if it doesn't exist
125 // 2. Fill the TagListItem in the HibernatableWebSocket's tagItems array
126 size_t position = 0;
127 for (auto tag = tags.begin(); tag < tags.end(); tag++, position++) {
128 auto& tagCollection = tagToWs.findOrCreate(*tag, [&tag]() {
129 auto item = kj::heap<TagCollection>(
130 kj::mv(*tag), kj::heap<kj::List<TagListItem, &TagListItem::link>>());
131 return decltype(tagToWs)::Entry{item->tag, kj::mv(item)};
132 });
133 // This TagListItem sits in the HibernatableWebSocket's tagItems array.
134 auto& tagListItem = refToHibernatable.tagItems[position];
135 tagListItem.hibWS = refToHibernatable;
136 tagListItem.tag = tagCollection->tag.asPtr();
137 
138 auto& list = tagCollection->list;
139 list->add(tagListItem);
140 // We also give the TagListItem a reference to the list it was added to so the
141 // HibernatableWebSocket can quickly remove itself from the list without doing a lookup
142 // in `tagToWs`.
143 tagListItem.list = *list.get();
144 }
145 
146 // Before starting the readLoop, we need to move the kj::Own<kj::WebSocket> from the
147 // api::WebSocket into the HibernatableWebSocket and accept the api::WebSocket as "hibernatable".
148 refToHibernatable.ws =
149 refToHibernatable.activeOrPackage.get<jsg::Ref<api::WebSocket>>()->acceptAsHibernatable(
150 refToHibernatable.getTags());
151 
152 // Finally, we initiate the readloop for this HibernatableWebSocket and
153 // give the task to the HibernationManager so it lives long.
154 readLoopTasks.add(handleReadLoop(refToHibernatable).catch_([](kj::Exception&& e) {
155 if (isInterestingException(e)) {
156 LOG_EXCEPTION_IF_INTERNAL("HibernationManagerImpl::handleReadLoop", e);
157 }
158 }));
159}
160 
161kj::Promise<void> HibernationManagerImpl::handleReadLoop(HibernatableWebSocket& refToHibernatable) {
162 kj::Maybe<kj::Exception> maybeException;
163 try {
164 co_await readLoop(refToHibernatable);
165 } catch (...) {
166 maybeException = kj::getCaughtExceptionAsKj();
167 }
168 co_await handleSocketTermination(refToHibernatable, maybeException);
169}
170 
171kj::Vector<jsg::Ref<api::WebSocket>> HibernationManagerImpl::getWebSockets(
172 jsg::Lock& js, kj::Maybe<kj::StringPtr> maybeTag) {
173 kj::Vector<jsg::Ref<api::WebSocket>> matches;
174 KJ_IF_SOME(tag, maybeTag) {
175 KJ_IF_SOME(item, tagToWs.find(tag)) {
176 auto& list = *((item)->list);
177 for (auto& entry: list) {
178 auto& hibWS = KJ_REQUIRE_NONNULL(entry.hibWS);
179 matches.add(hibWS.getActiveOrUnhibernate(js));
180 }
181 }
182 } else {
183 // Add all websockets!
184 for (auto& hibWS: allWs) {
185 matches.add(hibWS->getActiveOrUnhibernate(js));
186 }
187 }
188 return kj::mv(matches);
189}
190 
191void HibernationManagerImpl::setWebSocketAutoResponse(
192 kj::Maybe<kj::StringPtr> request, kj::Maybe<kj::StringPtr> response) {
193 KJ_IF_SOME(req, request) {
194 // If we have a request, we must also have a response. If response is kj::none, we'll throw.
195 autoResponsePair->request = kj::str(req);
196 autoResponsePair->response = kj::str(KJ_REQUIRE_NONNULL(response));
197 return;
198 }
199 // If we don't have a request, we must unset both request and response.
200 autoResponsePair->request = kj::none;
201 autoResponsePair->response = kj::none;
202}
203 
204kj::Maybe<jsg::Ref<api::WebSocketRequestResponsePair>> HibernationManagerImpl::
205 getWebSocketAutoResponse(jsg::Lock& js) {
206 KJ_IF_SOME(req, autoResponsePair->request) {
207 // When getting the currently set auto-response pair, if we have a request we must have a response
208 // set. If not, we'll throw.
209 return api::WebSocketRequestResponsePair::constructor(
210 js, kj::str(req), kj::str(KJ_REQUIRE_NONNULL(autoResponsePair->response)));
211 }
212 return kj::none;
213}
214 
215void HibernationManagerImpl::setTimerChannel(TimerChannel& timerChannel) {
216 timer = timerChannel;
217}
218 
219void HibernationManagerImpl::hibernateWebSockets(Worker::Lock& lock) {
220 JSG_WITHIN_CONTEXT_SCOPE(lock, lock.getContext(), [&](jsg::Lock& js) {
221 for (auto& ws: allWs) {
222 KJ_IF_SOME(active, ws->activeOrPackage.tryGet<jsg::Ref<api::WebSocket>>()) {
223 // Transfers ownership of properties from api::WebSocket to HibernatableWebSocket via the
224 // HibernationPackage.
225 ws->activeOrPackage.init<api::WebSocket::HibernationPackage>(
226 active.get()->buildPackageForHibernation());
227 } else {
228 } // Here to quash compiler warning
229 }
230 });
231}
232 
233void HibernationManagerImpl::setEventTimeout(kj::Maybe<uint32_t> timeoutMs) {
234 eventTimeoutMs = timeoutMs;
235}
236 
237kj::Maybe<uint32_t> HibernationManagerImpl::getEventTimeout() {
238 return eventTimeoutMs;
239}
240 
241void HibernationManagerImpl::dropHibernatableWebSocket(HibernatableWebSocket& hib) {
242 removeFromAllWs(hib);
243}
244 
245inline void HibernationManagerImpl::removeFromAllWs(HibernatableWebSocket& hib) {
246 auto& node = KJ_REQUIRE_NONNULL(hib.node);
247 allWs.erase(node);
248}
249 
250kj::Promise<void> HibernationManagerImpl::handleSocketTermination(
251 HibernatableWebSocket& hib, kj::Maybe<kj::Exception>& maybeError) {
252 kj::Maybe<kj::Promise<void>> event;
253 KJ_IF_SOME(error, maybeError) {
254 auto websocketId = randomUUID(kj::none);
255 webSocketsForEventHandler.insert(kj::str(websocketId), &hib);
256 kj::Maybe<api::HibernatableSocketParams> params;
257 if (!hib.hasDispatchedClose && (error.getType() == kj::Exception::Type::DISCONNECTED)) {
258 // If premature disconnect/cancel, dispatch a close event if we haven't already.
259 hib.hasDispatchedClose = true;
260 params = api::HibernatableSocketParams(1006,
261 kj::str("WebSocket disconnected without sending Close frame."), false,
262 kj::mv(websocketId));
263 } else {
264 // Otherwise, we need to dispatch an error event!
265 params = api::HibernatableSocketParams(kj::mv(error), kj::mv(websocketId));
266 }
267 
268 KJ_REQUIRE_NONNULL(params).setTimeout(eventTimeoutMs);
269 // Dispatch the event.
270 auto workerInterface = loopback->getWorker(IoChannelFactory::SubrequestMetadata{});
271 event = workerInterface
272 ->customEvent(kj::heap<api::HibernatableWebSocketCustomEvent>(
273 hibernationEventType, kj::mv(KJ_REQUIRE_NONNULL(params)), *this))
274 .ignoreResult()
275 .attach(kj::mv(workerInterface));
276 }
277 
278 // Returning the event promise will store it in readLoopTasks.
279 // After the task completes, we want to drop the websocket since we've closed the connection.
280 KJ_IF_SOME(promise, event) {
281 co_await promise;
282 }
283 
284 dropHibernatableWebSocket(hib);
285}
286 
287kj::Promise<void> HibernationManagerImpl::readLoop(HibernatableWebSocket& hib) {
288 // Like the api::WebSocket readLoop(), but we dispatch different types of events.
289 auto& ws = *KJ_REQUIRE_NONNULL(hib.ws);
290 while (true) {
291 kj::WebSocket::Message message = co_await ws.receive();
292 // Note that errors are handled by the callee of `readLoop`, since we throw from `receive()`.
293 
294 auto skip = false;
295 
296 // If we have a request != kj::none, we can compare it the received message. This also implies
297 // that we have a response set in autoResponsePair.
298 KJ_IF_SOME(req, autoResponsePair->request) {
299 KJ_SWITCH_ONEOF(message) {
300 KJ_CASE_ONEOF(text, kj::String) {
301 if (text == req) {
302 // If the received message matches the one set for auto-response, we must
303 // short-circuit readLoop, store the current timestamp and and automatically respond
304 // with the expected response.
305 TimerChannel& timerChannel = KJ_REQUIRE_NONNULL(timer);
306 // This should count as a new IO event, hence we should call syncTime
307 // otherwise the autoResponseTimestamp wouldn't be accurate.
308 timerChannel.syncTime();
309 // We should have set the timerChannel previously in the hibernation manager.
310 // If we haven't, we aren't able to get the current time.
311 hib.autoResponseTimestamp = timerChannel.now();
312 // We'll store the current timestamp in the HibernatableWebSocket to assure it gets
313 // stored even if the WebSocket is currently hibernating. In that scenario, the timestamp
314 // value will be loaded into the WebSocket during unhibernation.
315 KJ_SWITCH_ONEOF(hib.activeOrPackage) {
316 KJ_CASE_ONEOF(apiWs, jsg::Ref<api::WebSocket>) {
317 // If the actor is not hibernated/If the WebSocket is active, we need to update
318 // autoResponseTimestamp on the active websocket.
319 apiWs->setAutoResponseStatus(hib.autoResponseTimestamp, kj::READY_NOW);
320 // Since we had a request set, we must have and response that's sent back using the
321 // same websocket here. The sending of response is managed in web-socket to avoid
322 // possible racing problems with regular websocket messages.
323 co_await apiWs->sendAutoResponse(
324 kj::str(KJ_REQUIRE_NONNULL(autoResponsePair->response).asArray()), ws);
325 }
326 KJ_CASE_ONEOF(package, api::WebSocket::HibernationPackage) {
327 if (!package.closedOutgoingConnection) {
328 // We need to store the autoResponsePromise because we may instantiate an api::websocket
329 // If we do that, we have to provide it with the promise to avoid races. This can
330 // happen if we have a websocket hibernating, that unhibernates and sends a
331 // message while ws.send() for auto-response is also sending.
332 auto p = ws.send(KJ_REQUIRE_NONNULL(autoResponsePair->response).asArray()).fork();
333 hib.autoResponsePromise = p.addBranch();
334 co_await p;
335 hib.autoResponsePromise = kj::READY_NOW;
336 }
337 }
338 }
339 // If we've sent an auto response message, we should not unhibernate or deliver the
340 // received message to the actor
341 skip = true;
342 }
343 }
344 KJ_CASE_ONEOF_DEFAULT {}
345 }
346 }
347 
348 if (skip) {
349 continue;
350 }
351 
352 auto websocketId = randomUUID(kj::none);
353 webSocketsForEventHandler.insert(kj::str(websocketId), &hib);
354 
355 // Build the event params depending on what type of message we got.
356 kj::Maybe<api::HibernatableSocketParams> maybeParams;
357 KJ_SWITCH_ONEOF(message) {
358 KJ_CASE_ONEOF(text, kj::String) {
359 maybeParams.emplace(kj::mv(text), kj::mv(websocketId));
360 }
361 KJ_CASE_ONEOF(data, kj::Array<kj::byte>) {
362 maybeParams.emplace(kj::mv(data), kj::mv(websocketId));
363 }
364 KJ_CASE_ONEOF(close, kj::WebSocket::Close) {
365 maybeParams.emplace(close.code, kj::mv(close.reason), true, kj::mv(websocketId));
366 // We'll dispatch the close event, so let's mark our websocket as having done so to
367 // prevent a situation where we dispatch it twice.
368 hib.hasDispatchedClose = true;
369 }
370 }
371 
372 auto params = kj::mv(KJ_REQUIRE_NONNULL(maybeParams));
373 params.setTimeout(eventTimeoutMs);
374 auto isClose = params.isCloseEvent();
375 // Dispatch the event.
376 auto workerInterface = loopback->getWorker(IoChannelFactory::SubrequestMetadata{});
377 co_await workerInterface->customEvent(kj::heap<api::HibernatableWebSocketCustomEvent>(
378 hibernationEventType, kj::mv(params), *this));
379 if (isClose) {
380 co_return;
381 }
382 }
383}
384 
385}; // namespace workerd