Skip to content
File

Blob: src/workerd/api/hyperdrive.c++

4.6 KB
1// Copyright (c) 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 "hyperdrive.h"
6 
7#include "sockets.h"
8 
9#include <workerd/api/global-scope.h>
10#include <workerd/util/entropy.h>
11 
12#include <kj/compat/http.h>
13#include <kj/encoding.h>
14#include <kj/string.h>
15 
16namespace workerd::api {
17Hyperdrive::Hyperdrive(
18 uint clientIndex, kj::String database, kj::String user, kj::String password, kj::String scheme)
19 : clientIndex(clientIndex),
20 database(kj::mv(database)),
21 user(kj::mv(user)),
22 password(kj::mv(password)),
23 scheme(kj::mv(scheme)) {
24 kj::FixedArray<kj::byte, 16> randomBytes;
25 getEntropy(randomBytes.asPtr());
26 randomHost = kj::str(kj::encodeHex(randomBytes), ".hyperdrive.local");
27}
28 
29jsg::Ref<Socket> Hyperdrive::connect(jsg::Lock& js) {
30 auto connPromise = connectToDb();
31 
32 auto paf = kj::newPromiseAndFulfiller<kj::Maybe<kj::Exception>>();
33 auto conn =
34 kj::newPromisedStream(connPromise
35 .then([&f = *paf.fulfiller](kj::Own<kj::AsyncIoStream> stream) {
36 f.fulfill(kj::none);
37 return kj::mv(stream);
38 }, [&f = *paf.fulfiller](kj::Exception e) {
39 KJ_LOG(WARNING, "failed to connect to local database", e);
40 f.fulfill(e.clone());
41 return kj::mv(e);
42 }).attach(kj::mv(paf.fulfiller)));
43 
44 // TODO(someday): Support TLS? It's not at all necessary since we're connecting locally, but
45 // some users may want it anyway.
46 auto nullTlsStarter = kj::heap<kj::TlsStarterCallback>();
47 auto sock = setupSocket(js, kj::mv(conn), kj::str(getHost(), ":", getPort()),
48 kj::none /* localAddress */, kj::none, kj::mv(nullTlsStarter), SecureTransportKind::OFF,
49 kj::str(this->randomHost), false, kj::none /* maybeOpenedPrPair */);
50 sock->handleProxyStatus(js, kj::mv(paf.promise));
51 return sock;
52}
53 
54kj::StringPtr Hyperdrive::getDatabase() {
55 return this->database;
56}
57 
58kj::StringPtr Hyperdrive::getUser() {
59 return this->user;
60}
61kj::StringPtr Hyperdrive::getPassword() {
62 return this->password;
63}
64 
65kj::StringPtr Hyperdrive::getScheme() {
66 return this->scheme;
67}
68 
69kj::StringPtr Hyperdrive::getHost() {
70 if (!registeredConnectOverride) {
71 // Returns the random hostname and ensures the connect override is registered on the
72 // ServiceWorkerGlobalScope for the Worker. This getter has a side effect: it registers (or
73 // re-registers) an entry in the ServiceWorkerGlobalScope's connectOverrides HashMap so that
74 // cloudflare:sockets's connect() will route connections to this magic hostname through Hyperdrive.
75 auto& globalScope = IoContext::current().getCurrentLock().getGlobalScope();
76 globalScope.setConnectOverride(kj::str(randomHost, ":", getPort()),
77 [self = JSG_THIS](jsg::Lock& js) mutable { return self->connect(js); });
78 registeredConnectOverride = true;
79 }
80 return randomHost;
81}
82 
83// We currently only support Postgres and MySQL
84uint16_t Hyperdrive::getPort() {
85 if (scheme == "mysql") {
86 return 3306;
87 }
88 
89 // We default to postgres if the scheme is not mysql
90 return 5432;
91}
92 
93kj::String Hyperdrive::getConnectionString() {
94 // MySQL: `?ssl-mode=disabled`
95 // PostgreSQL: `?sslmode=disable`
96 auto sslParameter = scheme == "mysql" ? "?ssl-mode=disabled" : "?sslmode=disable";
97 return kj::str(getScheme(), "://", getUser(), ":", getPassword(), "@", getHost(), ":", getPort(),
98 "/", getDatabase(), sslParameter);
99}
100 
101kj::Promise<kj::Own<kj::AsyncIoStream>> Hyperdrive::connectToDb() {
102 auto& context = IoContext::current();
103 auto service = context.getSubrequestChannel(
104 this->clientIndex, true, kj::none, kj::ConstString("hyperdrive_connect"_kjc));
105 
106 kj::HttpHeaderTable headerTable;
107 kj::HttpHeaders headers(headerTable);
108 
109 auto connectReq = kj::newHttpClient(*service)->connect(
110 kj::str(getHost(), ":", getPort()), headers, kj::HttpConnectSettings{});
111 
112 auto status = co_await connectReq.status;
113 
114 if (status.statusCode >= 200 && status.statusCode < 300) {
115 co_return kj::mv(connectReq.connection);
116 }
117 
118 KJ_IF_SOME(e, status.errorBody) {
119 try {
120 auto errorBody = co_await e->readAllText();
121 kj::throwFatalException(
122 KJ_EXCEPTION(FAILED, kj::str("unexpected error connecting to database: ", errorBody)));
123 } catch (const kj::Exception& e) {
124 kj::throwFatalException(KJ_EXCEPTION(FAILED,
125 kj::str("unexpected error connecting to database "
126 "and couldn't read error details: ",
127 e)));
128 }
129 } else {
130 kj::throwFatalException(KJ_EXCEPTION(
131 FAILED, kj::str("unexpected error connecting to database: ", status.statusText)));
132 }
133}
134} // namespace workerd::api