File
Blob: src/workerd/util/stream-utils.c++
| 1 | #include "stream-utils.h" |
| 2 | |
| 3 | #include <kj/common.h> |
| 4 | #include <kj/debug.h> |
| 5 | #include <kj/exception.h> |
| 6 | #include <kj/one-of.h> |
| 7 | |
| 8 | namespace workerd { |
| 9 | |
| 10 | namespace { |
| 11 | |
| 12 | // An AsyncInputStream implementation that reads from an in-memory buffer. |
| 13 | // This is optimized for the case where the entire contents are available |
| 14 | // up-front, so it doesn't do any dynamic memory allocation or copying. |
| 15 | // It also supports optimized teeing when the backing storage is provided. |
| 16 | class MemoryInputStream final: public kj::AsyncInputStream { |
| 17 | private: |
| 18 | struct OwnedBacking: public kj::Refcounted { |
| 19 | kj::Own<void> backing; |
| 20 | OwnedBacking(kj::Own<void>&& backing): backing(kj::mv(backing)) {} |
| 21 | }; |
| 22 | |
| 23 | public: |
| 24 | MemoryInputStream( |
| 25 | kj::ArrayPtr<const kj::byte> data, kj::Maybe<kj::Own<void>> maybeBacking = kj::none) |
| 26 | : data(data), |
| 27 | // Note that we don't actually check that maybeBacking actually owns the |
| 28 | // memory that `data` points to. It is the caller's responsibility to ensure |
| 29 | // this is the case if they want teeing to be safely supported. |
| 30 | ownedBacking(maybeBacking.map( |
| 31 | [](kj::Own<void>& backing) mutable { return kj::rc<OwnedBacking>(kj::mv(backing)); })) { |
| 32 | } |
| 33 | MemoryInputStream(kj::ArrayPtr<const kj::byte> data, kj::Rc<OwnedBacking> ownedBacking) |
| 34 | : data(data), |
| 35 | ownedBacking(kj::mv(ownedBacking)) {} |
| 36 | |
| 37 | kj::Promise<size_t> tryRead(void* buffer, size_t minBytes, size_t maxBytes) override { |
| 38 | auto ptr = kj::arrayPtr<kj::byte>(static_cast<kj::byte*>(buffer), maxBytes); |
| 39 | size_t toRead = kj::min(data.size(), ptr.size()); |
| 40 | if (toRead == 0) return toRead; |
| 41 | ptr.first(toRead).copyFrom(data.first(toRead)); |
| 42 | data = data.slice(toRead); |
| 43 | return toRead; |
| 44 | } |
| 45 | |
| 46 | kj::Maybe<uint64_t> tryGetLength() override { |
| 47 | return data.size(); |
| 48 | } |
| 49 | |
| 50 | kj::Promise<uint64_t> pumpTo(kj::AsyncOutputStream& output, uint64_t amount) override { |
| 51 | // An optimized pumpTo... we know we have all the data right here. We can |
| 52 | // just write it all at once up to `amount`. |
| 53 | uint64_t toRead = kj::min(data.size(), amount); |
| 54 | if (toRead == 0) { |
| 55 | co_return toRead; |
| 56 | } |
| 57 | co_await output.write(data.first(toRead)); |
| 58 | data = data.slice(toRead); |
| 59 | co_return toRead; |
| 60 | } |
| 61 | |
| 62 | kj::Maybe<kj::Own<AsyncInputStream>> tryTee(uint64_t limit = kj::maxValue) override { |
| 63 | // If a MemoryInputStream is holding onto backing storage, then we can safely |
| 64 | // tee it here, allowing us to avoid the default tee implementation which needs |
| 65 | // additional buffering. Tee'ing just becomes a matter of sharing the backing |
| 66 | // storage and the data slice directly. This allows us to avoid any additional |
| 67 | // buffering in unread stream branches since all of the data is already in memory |
| 68 | // anyway. If we're not holding onto the backing storage, then we cannot safely |
| 69 | // assume that the a tee branch will safely be able to read the data, so we'll |
| 70 | // fall back to the default kj::newTee implementation. |
| 71 | KJ_IF_SOME(owned, ownedBacking) { |
| 72 | return kj::heap<MemoryInputStream>(data, owned.addRef()); |
| 73 | } |
| 74 | return kj::none; |
| 75 | } |
| 76 | |
| 77 | private: |
| 78 | kj::ArrayPtr<const kj::byte> data; |
| 79 | kj::Maybe<kj::Rc<OwnedBacking>> ownedBacking; |
| 80 | }; |
| 81 | |
| 82 | class NeuterableInputStreamImpl final: public NeuterableInputStream { |
| 83 | public: |
| 84 | NeuterableInputStreamImpl(kj::AsyncInputStream& inner): inner(&inner) {} |
| 85 | |
| 86 | void neuter(kj::Exception exception) override { |
| 87 | if (inner.is<kj::AsyncInputStream*>()) { |
| 88 | inner = exception.clone(); |
| 89 | if (!canceler.isEmpty()) { |
| 90 | canceler.cancel(kj::mv(exception)); |
| 91 | } |
| 92 | } |
| 93 | } |
| 94 | |
| 95 | kj::Promise<size_t> tryRead(void* buffer, size_t minBytes, size_t maxBytes) override { |
| 96 | return canceler.wrap(getStream().tryRead(buffer, minBytes, maxBytes)); |
| 97 | } |
| 98 | kj::Maybe<uint64_t> tryGetLength() override { |
| 99 | return getStream().tryGetLength(); |
| 100 | } |
| 101 | kj::Promise<uint64_t> pumpTo(kj::AsyncOutputStream& output, uint64_t amount) override { |
| 102 | return canceler.wrap(getStream().pumpTo(output, amount)); |
| 103 | } |
| 104 | |
| 105 | private: |
| 106 | kj::OneOf<kj::AsyncInputStream*, kj::Exception> inner; |
| 107 | kj::Canceler canceler; |
| 108 | |
| 109 | kj::AsyncInputStream& getStream() { |
| 110 | KJ_SWITCH_ONEOF(inner) { |
| 111 | KJ_CASE_ONEOF(stream, kj::AsyncInputStream*) { |
| 112 | return *stream; |
| 113 | } |
| 114 | KJ_CASE_ONEOF(exception, kj::Exception) { |
| 115 | kj::throwFatalException(exception.clone()); |
| 116 | } |
| 117 | } |
| 118 | KJ_UNREACHABLE; |
| 119 | } |
| 120 | }; |
| 121 | |
| 122 | class NeuterableIoStreamImpl final: public NeuterableIoStream { |
| 123 | public: |
| 124 | NeuterableIoStreamImpl(kj::AsyncIoStream& inner): inner(&inner) {} |
| 125 | |
| 126 | void neuter(kj::Exception reason) override { |
| 127 | if (inner.is<kj::AsyncIoStream*>()) { |
| 128 | inner = reason.clone(); |
| 129 | if (!canceler.isEmpty()) { |
| 130 | canceler.cancel(kj::mv(reason)); |
| 131 | } |
| 132 | } |
| 133 | } |
| 134 | |
| 135 | // AsyncInputStream |
| 136 | |
| 137 | kj::Promise<size_t> tryRead(void* buffer, size_t minBytes, size_t maxBytes) override { |
| 138 | return canceler.wrap(getStream().tryRead(buffer, minBytes, maxBytes)); |
| 139 | } |
| 140 | kj::Maybe<uint64_t> tryGetLength() override { |
| 141 | return getStream().tryGetLength(); |
| 142 | } |
| 143 | kj::Promise<uint64_t> pumpTo(kj::AsyncOutputStream& output, uint64_t amount) override { |
| 144 | return canceler.wrap(getStream().pumpTo(output, amount)); |
| 145 | } |
| 146 | |
| 147 | // AsyncOutputStream |
| 148 | |
| 149 | kj::Promise<void> write(kj::ArrayPtr<const kj::byte> buffer) override { |
| 150 | return canceler.wrap(getStream().write(buffer)); |
| 151 | } |
| 152 | kj::Promise<void> write(kj::ArrayPtr<const kj::ArrayPtr<const kj::byte>> pieces) override { |
| 153 | return canceler.wrap(getStream().write(pieces)); |
| 154 | } |
| 155 | kj::Maybe<kj::Promise<uint64_t>> tryPumpFrom( |
| 156 | kj::AsyncInputStream& input, uint64_t amount) override { |
| 157 | return getStream().tryPumpFrom(input, amount).map([this](kj::Promise<uint64_t> promise) { |
| 158 | return canceler.wrap(kj::mv(promise)); |
| 159 | }); |
| 160 | } |
| 161 | kj::Promise<void> whenWriteDisconnected() override { |
| 162 | return canceler.wrap(getStream().whenWriteDisconnected()); |
| 163 | } |
| 164 | |
| 165 | // AsyncIoStream |
| 166 | |
| 167 | void shutdownWrite() override { |
| 168 | getStream().shutdownWrite(); |
| 169 | }; |
| 170 | void abortRead() override { |
| 171 | getStream().abortRead(); |
| 172 | } |
| 173 | void getsockopt(int level, int option, void* value, kj::uint* length) override { |
| 174 | getStream().getsockopt(level, option, value, length); |
| 175 | } |
| 176 | void setsockopt(int level, int option, const void* value, kj::uint length) override { |
| 177 | getStream().setsockopt(level, option, value, length); |
| 178 | } |
| 179 | void getsockname(struct sockaddr* addr, kj::uint* length) override { |
| 180 | getStream().getsockname(addr, length); |
| 181 | } |
| 182 | void getpeername(struct sockaddr* addr, kj::uint* length) override { |
| 183 | getStream().getpeername(addr, length); |
| 184 | } |
| 185 | virtual kj::Maybe<int> getFd() const override { |
| 186 | return getStream().getFd(); |
| 187 | } |
| 188 | |
| 189 | private: |
| 190 | kj::OneOf<kj::AsyncIoStream*, kj::Exception> inner; |
| 191 | kj::Canceler canceler; |
| 192 | |
| 193 | kj::AsyncIoStream& getStream() { |
| 194 | KJ_IF_SOME(stream, inner.tryGet<kj::AsyncIoStream*>()) { |
| 195 | return *stream; |
| 196 | } |
| 197 | kj::throwFatalException(inner.get<kj::Exception>().clone()); |
| 198 | } |
| 199 | kj::AsyncIoStream& getStream() const { |
| 200 | KJ_IF_SOME(stream, inner.tryGet<kj::AsyncIoStream*>()) { |
| 201 | return *stream; |
| 202 | } |
| 203 | kj::throwFatalException(inner.get<kj::Exception>().clone()); |
| 204 | } |
| 205 | }; |
| 206 | |
| 207 | // The kj::NullStream instance is stateless, discards all writes, and returns |
| 208 | // EOF on all reads. We can, therefore, safely share a single static global |
| 209 | // instance instead of allocating a new one each time. |
| 210 | static kj::NullStream nullStream{}; |
| 211 | |
| 212 | } // namespace |
| 213 | |
| 214 | kj::AsyncOutputStream& getGlobalNullOutputStream() { |
| 215 | return nullStream; |
| 216 | } |
| 217 | |
| 218 | kj::Own<kj::AsyncIoStream> newNullIoStream() { |
| 219 | return kj::Own<kj::AsyncIoStream>(&nullStream, kj::NullDisposer::instance); |
| 220 | } |
| 221 | |
| 222 | kj::Own<kj::AsyncInputStream> newNullInputStream() { |
| 223 | return kj::Own<kj::AsyncInputStream>(&nullStream, kj::NullDisposer::instance); |
| 224 | } |
| 225 | |
| 226 | kj::Own<kj::AsyncOutputStream> newNullOutputStream() { |
| 227 | return kj::Own<kj::AsyncOutputStream>(&nullStream, kj::NullDisposer::instance); |
| 228 | } |
| 229 | |
| 230 | kj::Own<kj::AsyncInputStream> newMemoryInputStream( |
| 231 | kj::ArrayPtr<const kj::byte> data, kj::Maybe<kj::Own<void>> maybeBacking) { |
| 232 | return kj::heap<MemoryInputStream>(data, kj::mv(maybeBacking)); |
| 233 | } |
| 234 | |
| 235 | kj::Own<kj::AsyncInputStream> newMemoryInputStream( |
| 236 | kj::StringPtr data, kj::Maybe<kj::Own<void>> maybeBacking) { |
| 237 | return kj::heap<MemoryInputStream>(data.asBytes(), kj::mv(maybeBacking)); |
| 238 | } |
| 239 | |
| 240 | kj::Own<NeuterableInputStream> newNeuterableInputStream(kj::AsyncInputStream& inner) { |
| 241 | return kj::refcounted<NeuterableInputStreamImpl>(inner); |
| 242 | } |
| 243 | |
| 244 | kj::Own<NeuterableIoStream> newNeuterableIoStream(kj::AsyncIoStream& inner) { |
| 245 | return kj::heap<NeuterableIoStreamImpl>(inner); |
| 246 | } |
| 247 | |
| 248 | } // namespace workerd |