File
Blob: src/workerd/api/hyperdrive.c++
| 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 | |
| 16 | namespace workerd::api { |
| 17 | Hyperdrive::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 | |
| 29 | jsg::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 | |
| 54 | kj::StringPtr Hyperdrive::getDatabase() { |
| 55 | return this->database; |
| 56 | } |
| 57 | |
| 58 | kj::StringPtr Hyperdrive::getUser() { |
| 59 | return this->user; |
| 60 | } |
| 61 | kj::StringPtr Hyperdrive::getPassword() { |
| 62 | return this->password; |
| 63 | } |
| 64 | |
| 65 | kj::StringPtr Hyperdrive::getScheme() { |
| 66 | return this->scheme; |
| 67 | } |
| 68 | |
| 69 | kj::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 |
| 84 | uint16_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 | |
| 93 | kj::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 | |
| 101 | kj::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 |