File
Blob: src/workerd/io/hibernation-manager.c++
| 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 | |
| 11 | namespace workerd { |
| 12 | |
| 13 | HibernationManagerImpl::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 | |
| 25 | HibernationManagerImpl::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 | |
| 46 | kj::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 | |
| 54 | kj::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 | |
| 62 | jsg::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 | |
| 80 | HibernationManagerImpl::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 | |
| 87 | HibernationManagerImpl::~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 | |
| 99 | kj::Own<Worker::Actor::HibernationManager> HibernationManagerImpl::addRef() { |
| 100 | return kj::addRef(*this); |
| 101 | } |
| 102 | |
| 103 | void 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 | |
| 161 | kj::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 | |
| 171 | kj::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 | |
| 191 | void 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 | |
| 204 | kj::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 | |
| 215 | void HibernationManagerImpl::setTimerChannel(TimerChannel& timerChannel) { |
| 216 | timer = timerChannel; |
| 217 | } |
| 218 | |
| 219 | void 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 | |
| 233 | void HibernationManagerImpl::setEventTimeout(kj::Maybe<uint32_t> timeoutMs) { |
| 234 | eventTimeoutMs = timeoutMs; |
| 235 | } |
| 236 | |
| 237 | kj::Maybe<uint32_t> HibernationManagerImpl::getEventTimeout() { |
| 238 | return eventTimeoutMs; |
| 239 | } |
| 240 | |
| 241 | void HibernationManagerImpl::dropHibernatableWebSocket(HibernatableWebSocket& hib) { |
| 242 | removeFromAllWs(hib); |
| 243 | } |
| 244 | |
| 245 | inline void HibernationManagerImpl::removeFromAllWs(HibernatableWebSocket& hib) { |
| 246 | auto& node = KJ_REQUIRE_NONNULL(hib.node); |
| 247 | allWs.erase(node); |
| 248 | } |
| 249 | |
| 250 | kj::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 | |
| 287 | kj::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 |