Skip to content
File

Blob: src/workerd/api/hibernatable-web-socket.c++

11.4 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 "hibernatable-web-socket.h"
6 
7#include <workerd/api/global-scope.h>
8#include <workerd/io/hibernation-manager.h>
9#include <workerd/io/tracer.h>
10#include <workerd/jsg/ser.h>
11 
12namespace workerd::api {
13 
14HibernatableWebSocketEvent::HibernatableWebSocketEvent(): ExtendableEvent("webSocketMessage") {};
15 
16Worker::Actor::HibernationManager& HibernatableWebSocketEvent::getHibernationManager(
17 jsg::Lock& lock) {
18 auto& actor = KJ_REQUIRE_NONNULL(IoContext::current().getActor());
19 return KJ_REQUIRE_NONNULL(actor.getHibernationManager());
20}
21 
22HibernatableWebSocketEvent::ItemsForRelease HibernatableWebSocketEvent::prepareForRelease(
23 jsg::Lock& lock, kj::StringPtr websocketId) {
24 auto& manager = kj::downcast<HibernationManagerImpl>(getHibernationManager(lock));
25 auto& hibernatableWebSocket =
26 KJ_REQUIRE_NONNULL(manager.webSocketsForEventHandler.findEntry(websocketId));
27 
28 // Note that we don't call `claimWebSocket()` to get this, since we would lose our reference to
29 // the HibernatableWebSocket (it removes it from `webSocketsForEventHandler`).
30 auto websocketRef = hibernatableWebSocket.value->getActiveOrUnhibernate(lock);
31 auto ownedWebSocket = kj::mv(KJ_REQUIRE_NONNULL(hibernatableWebSocket.value->ws));
32 auto tags = hibernatableWebSocket.value->cloneTags();
33 
34 // Now that we've obtained the websocket for the event, let's free up the slots we had allocated.
35 manager.webSocketsForEventHandler.erase(hibernatableWebSocket);
36 
37 return ItemsForRelease(kj::mv(websocketRef), kj::mv(ownedWebSocket), kj::mv(tags));
38}
39 
40jsg::Ref<WebSocket> HibernatableWebSocketEvent::claimWebSocket(
41 jsg::Lock& lock, kj::StringPtr websocketId) {
42 // Should only be called once per event since it removes the HibernatableWebSocket from the
43 // webSocketsForEventHandler collection.
44 auto& manager = kj::downcast<HibernationManagerImpl>(getHibernationManager(lock));
45 
46 // Grab it from our collection.
47 auto& hibernatableWebSocket =
48 KJ_REQUIRE_NONNULL(manager.webSocketsForEventHandler.findEntry(websocketId));
49 
50 // Get the reference.
51 auto websocket = hibernatableWebSocket.value->getActiveOrUnhibernate(lock);
52 
53 // Now that we've obtained the websocket, we need to remove the entry from the map and make the
54 // key available again.
55 manager.webSocketsForEventHandler.erase(hibernatableWebSocket);
56 
57 return kj::mv(websocket);
58}
59 
60kj::Promise<WorkerInterface::CustomEvent::Result> HibernatableWebSocketCustomEvent::run(
61 kj::Own<IoContext_IncomingRequest> incomingRequest,
62 kj::Maybe<kj::StringPtr> entrypointName,
63 kj::Maybe<Worker::VersionInfo> versionInfo,
64 Frankenvalue props,
65 kj::TaskSet& waitUntilTasks,
66 bool isDynamicDispatch) {
67 // Mark the request as delivered because we're about to run some JS.
68 auto& context = incomingRequest->getContext();
69 incomingRequest->delivered();
70 
71 KJ_DEFER({ waitUntilTasks.add(incomingRequest->drain().attach(kj::mv(incomingRequest))); });
72 
73 EventOutcome outcome = EventOutcome::OK;
74 
75 // We definitely have an actor by this point. Let's set the hibernation manager on the actor
76 // before we start running any events that might need to access it.
77 auto& a = KJ_REQUIRE_NONNULL(context.getActor());
78 if (a.getHibernationManager() == kj::none) {
79 a.setHibernationManager(kj::addRef(KJ_REQUIRE_NONNULL(manager)));
80 }
81 
82 auto eventParameters = consumeParams();
83 
84 try {
85 co_await context.run(
86 [entrypointName = entrypointName, &context, eventParameters = kj::mv(eventParameters),
87 versionInfo = kj::mv(versionInfo), props = kj::mv(props),
88 isDynamicDispatch](Worker::Lock& lock) mutable {
89 KJ_SWITCH_ONEOF(eventParameters.eventType) {
90 KJ_CASE_ONEOF(text, HibernatableSocketParams::Text) {
91 return lock.getGlobalScope().sendHibernatableWebSocketMessage(context,
92 kj::mv(text.message), eventParameters.eventTimeoutMs,
93 kj::mv(eventParameters.websocketId), lock,
94 lock.getExportedHandler(entrypointName, kj::mv(versionInfo), kj::mv(props),
95 context.getActor(), isDynamicDispatch));
96 }
97 KJ_CASE_ONEOF(data, HibernatableSocketParams::Data) {
98 return lock.getGlobalScope().sendHibernatableWebSocketMessage(context,
99 kj::mv(data.message), eventParameters.eventTimeoutMs,
100 kj::mv(eventParameters.websocketId), lock,
101 lock.getExportedHandler(entrypointName, kj::mv(versionInfo), kj::mv(props),
102 context.getActor(), isDynamicDispatch));
103 }
104 KJ_CASE_ONEOF(close, HibernatableSocketParams::Close) {
105 return lock.getGlobalScope().sendHibernatableWebSocketClose(context, kj::mv(close),
106 eventParameters.eventTimeoutMs, kj::mv(eventParameters.websocketId), lock,
107 lock.getExportedHandler(entrypointName, kj::mv(versionInfo), kj::mv(props),
108 context.getActor(), isDynamicDispatch));
109 }
110 KJ_CASE_ONEOF(e, HibernatableSocketParams::Error) {
111 return lock.getGlobalScope().sendHibernatableWebSocketError(context, kj::mv(e.error),
112 eventParameters.eventTimeoutMs, kj::mv(eventParameters.websocketId), lock,
113 lock.getExportedHandler(entrypointName, kj::mv(versionInfo), kj::mv(props),
114 context.getActor(), isDynamicDispatch));
115 }
116 KJ_UNREACHABLE;
117 }
118 });
119 } catch (kj::Exception& e) {
120 if (auto desc = e.getDescription();
121 !jsg::isTunneledException(desc) && !jsg::isDoNotLogException(desc)) {
122 LOG_EXCEPTION("HibernatableWebSocketCustomEvent"_kj, e);
123 }
124 outcome = EventOutcome::EXCEPTION;
125 }
126 
127 co_return Result{
128 .outcome = outcome,
129 };
130}
131 
132kj::Promise<WorkerInterface::CustomEvent::Result> HibernatableWebSocketCustomEvent::sendRpc(
133 capnp::HttpOverCapnpFactory& httpOverCapnpFactory,
134 capnp::ByteStreamFactory& byteStreamFactory,
135 rpc::EventDispatcher::Client dispatcher) {
136 auto req = dispatcher.castAs<rpc::HibernatableWebSocketEventDispatcher>()
137 .hibernatableWebSocketEventRequest();
138 
139 KJ_IF_SOME(rpcParameters, params.tryGet<kj::Own<HibernationReader>>()) {
140 req.setMessage(rpcParameters->getMessage());
141 } else {
142 auto message = req.initMessage();
143 auto payload = message.initPayload();
144 auto& eventParameters = KJ_REQUIRE_NONNULL(params.tryGet<HibernatableSocketParams>());
145 KJ_SWITCH_ONEOF(eventParameters.eventType) {
146 KJ_CASE_ONEOF(text, HibernatableSocketParams::Text) {
147 payload.setText(kj::mv(text.message));
148 }
149 KJ_CASE_ONEOF(data, HibernatableSocketParams::Data) {
150 payload.setData(kj::mv(data.message));
151 }
152 KJ_CASE_ONEOF(close, HibernatableSocketParams::Close) {
153 auto closeBuilder = payload.initClose();
154 closeBuilder.setCode(close.code);
155 closeBuilder.setReason(kj::mv(close.reason));
156 closeBuilder.setWasClean(close.wasClean);
157 }
158 KJ_CASE_ONEOF(e, HibernatableSocketParams::Error) {
159 payload.setError(e.error.getDescription());
160 }
161 KJ_UNREACHABLE;
162 }
163 message.setWebsocketId(kj::mv(eventParameters.websocketId));
164 KJ_IF_SOME(t, eventParameters.eventTimeoutMs) {
165 message.setEventTimeoutMs(t);
166 }
167 }
168 
169 return req.send().then([](auto resp) {
170 auto respResult = resp.getResult();
171 return WorkerInterface::CustomEvent::Result{
172 .outcome = respResult.getOutcome(),
173 };
174 });
175}
176 
177HibernatableWebSocketEvent::ItemsForRelease::ItemsForRelease(
178 jsg::Ref<WebSocket> ref, kj::Own<kj::WebSocket> owned, kj::Array<kj::String> tags)
179 : webSocketRef(kj::mv(ref)),
180 ownedWebSocket(kj::mv(owned)),
181 tags(kj::mv(tags)) {}
182 
183HibernatableWebSocketCustomEvent::HibernatableWebSocketCustomEvent(uint16_t typeId,
184 kj::Own<HibernationReader> params,
185 kj::Maybe<Worker::Actor::HibernationManager&> manager)
186 : typeId(typeId),
187 params(kj::mv(params)) {}
188HibernatableWebSocketCustomEvent::HibernatableWebSocketCustomEvent(
189 uint16_t typeId, HibernatableSocketParams params, Worker::Actor::HibernationManager& manager)
190 : typeId(typeId),
191 params(kj::mv(params)),
192 manager(manager) {}
193 
194// Try to extract event type from params if available
195tracing::HibernatableWebSocketEventInfo::Type HibernatableWebSocketCustomEvent::getEventType()
196 const {
197 KJ_SWITCH_ONEOF(params) {
198 KJ_CASE_ONEOF(socketParams, HibernatableSocketParams) {
199 KJ_SWITCH_ONEOF(socketParams.eventType) {
200 KJ_CASE_ONEOF(_, HibernatableSocketParams::Text) {
201 return tracing::HibernatableWebSocketEventInfo::Message{};
202 }
203 KJ_CASE_ONEOF(_, HibernatableSocketParams::Data) {
204 return tracing::HibernatableWebSocketEventInfo::Message{};
205 }
206 KJ_CASE_ONEOF(close, HibernatableSocketParams::Close) {
207 return tracing::HibernatableWebSocketEventInfo::Close{close.code, close.wasClean};
208 }
209 KJ_CASE_ONEOF(_, HibernatableSocketParams::Error) {
210 return tracing::HibernatableWebSocketEventInfo::Error{};
211 }
212 }
213 }
214 KJ_CASE_ONEOF(reader, kj::Own<HibernationReader>) {
215 // Parse the HibernationReader to determine the actual event type
216 auto payload = reader->getMessage().getPayload();
217 switch (payload.which()) {
218 case rpc::HibernatableWebSocketEventMessage::Payload::TEXT:
219 case rpc::HibernatableWebSocketEventMessage::Payload::DATA:
220 return tracing::HibernatableWebSocketEventInfo::Message{};
221 case rpc::HibernatableWebSocketEventMessage::Payload::CLOSE: {
222 auto close = payload.getClose();
223 return tracing::HibernatableWebSocketEventInfo::Close{
224 close.getCode(), close.getWasClean()};
225 }
226 case rpc::HibernatableWebSocketEventMessage::Payload::ERROR:
227 return tracing::HibernatableWebSocketEventInfo::Error{};
228 }
229 }
230 }
231 KJ_UNREACHABLE;
232}
233 
234tracing::EventInfo HibernatableWebSocketCustomEvent::getEventInfo() const {
235 return tracing::HibernatableWebSocketEventInfo(getEventType());
236}
237 
238HibernatableSocketParams HibernatableWebSocketCustomEvent::consumeParams() {
239 KJ_IF_SOME(p, params.tryGet<kj::Own<HibernationReader>>()) {
240 kj::Maybe<HibernatableSocketParams> eventParameters;
241 auto websocketId = kj::str(p->getMessage().getWebsocketId());
242 auto payload = p->getMessage().getPayload();
243 switch (payload.which()) {
244 case rpc::HibernatableWebSocketEventMessage::Payload::TEXT: {
245 eventParameters.emplace(kj::str(payload.getText()), kj::mv(websocketId));
246 break;
247 }
248 case rpc::HibernatableWebSocketEventMessage::Payload::DATA: {
249 kj::Array<byte> b = kj::heapArray(payload.getData().asBytes());
250 eventParameters.emplace(kj::mv(b), kj::mv(websocketId));
251 break;
252 }
253 case rpc::HibernatableWebSocketEventMessage::Payload::CLOSE: {
254 auto close = payload.getClose();
255 eventParameters.emplace(
256 close.getCode(), kj::str(close.getReason()), close.getWasClean(), kj::mv(websocketId));
257 break;
258 }
259 case rpc::HibernatableWebSocketEventMessage::Payload::ERROR: {
260 eventParameters.emplace(
261 KJ_EXCEPTION(FAILED, kj::str(payload.getError())), kj::mv(websocketId));
262 break;
263 }
264 }
265 KJ_REQUIRE_NONNULL(eventParameters).setTimeout(p->getMessage().getEventTimeoutMs());
266 return kj::mv(KJ_REQUIRE_NONNULL(eventParameters));
267 }
268 return kj::mv(KJ_REQUIRE_NONNULL(params.tryGet<HibernatableSocketParams>()));
269}
270 
271} // namespace workerd::api