File
Blob: src/workerd/util/capnp-mock.h
| 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 | #pragma once |
| 6 | |
| 7 | #include <capnp/dynamic.h> |
| 8 | #include <capnp/message.h> |
| 9 | #include <capnp/serialize-text.h> |
| 10 | #include <kj/debug.h> |
| 11 | #include <kj/list.h> |
| 12 | #include <kj/map.h> |
| 13 | #include <kj/refcount.h> |
| 14 | #include <kj/source-location.h> |
| 15 | |
| 16 | namespace workerd { |
| 17 | |
| 18 | // ======================================================================================= |
| 19 | // KJ assert macros that support specifying a SourceLocation. These allow our test functions |
| 20 | // below to capture the caller's SourceLocation for use in errors, which is nice. |
| 21 | // |
| 22 | // TODO(cleanup): Move this to KJ! |
| 23 | |
| 24 | #ifndef KJ_REQUIRE_AT |
| 25 | #define KJ_REQUIRE_AT(cond, location, ...) \ |
| 26 | if (auto _kjCondition = ::kj::_::MAGIC_ASSERT << cond) { \ |
| 27 | } else \ |
| 28 | for (::kj::_::Debug::Fault f(location.fileName, location.lineNumber, \ |
| 29 | ::kj::Exception::Type::FAILED, #cond, "_kjCondition," #__VA_ARGS__, _kjCondition, \ |
| 30 | ##__VA_ARGS__); \ |
| 31 | ; f.fatal()) |
| 32 | #endif |
| 33 | |
| 34 | #ifndef KJ_FAIL_REQUIRE_AT |
| 35 | #define KJ_FAIL_REQUIRE_AT(location, ...) \ |
| 36 | for (::kj::_::Debug::Fault f(location.fileName, location.lineNumber, \ |
| 37 | ::kj::Exception::Type::FAILED, nullptr, #__VA_ARGS__, ##__VA_ARGS__); \ |
| 38 | ; f.fatal()) |
| 39 | #endif |
| 40 | |
| 41 | #ifndef KJ_REQUIRE_NONNULL_AT |
| 42 | #define KJ_REQUIRE_NONNULL_AT(value, location, ...) \ |
| 43 | (*({ \ |
| 44 | auto _kj_result = ::kj::_::readMaybe(value); \ |
| 45 | if (KJ_UNLIKELY(!_kj_result)) { \ |
| 46 | ::kj::_::Debug::Fault(location.fileName, location.lineNumber, ::kj::Exception::Type::FAILED, \ |
| 47 | #value " != nullptr", #__VA_ARGS__, ##__VA_ARGS__) \ |
| 48 | .fatal(); \ |
| 49 | } \ |
| 50 | kj::mv(_kj_result); \ |
| 51 | })) |
| 52 | #endif |
| 53 | |
| 54 | #ifndef KJ_ASSERT_AT |
| 55 | #define KJ_ASSERT_AT KJ_REQUIRE_AT |
| 56 | #endif |
| 57 | |
| 58 | #ifndef KJ_FAIL_ASSERT_AT |
| 59 | #define KJ_FAIL_ASSERT_AT KJ_FAIL_REQUIRE_AT |
| 60 | #endif |
| 61 | |
| 62 | #ifndef KJ_ASSERT_NONNULL_AT |
| 63 | #define KJ_ASSERT_NONNULL_AT KJ_REQUIRE_NONNULL_AT |
| 64 | #endif |
| 65 | |
| 66 | #ifndef KJ_LOG_AT |
| 67 | #define KJ_LOG_AT(severity, location, ...) \ |
| 68 | for (bool _kj_shouldLog = ::kj::_::Debug::shouldLog(::kj::LogSeverity::severity); _kj_shouldLog; \ |
| 69 | _kj_shouldLog = false) \ |
| 70 | ::kj::_::Debug::log(location.fileName, location.lineNumber, ::kj::LogSeverity::severity, \ |
| 71 | #__VA_ARGS__, ##__VA_ARGS__) |
| 72 | #endif |
| 73 | |
| 74 | // ======================================================================================= |
| 75 | // Cap'n Proto mocking framework |
| 76 | // |
| 77 | // TODO(cleanup): Move this to Cap'n Proto! |
| 78 | |
| 79 | const capnp::TextCodec TEXT_CODEC; |
| 80 | |
| 81 | kj::String canonicalizeCapnpText( |
| 82 | capnp::StructSchema schema, kj::StringPtr text, kj::Maybe<kj::StringPtr> capName = kj::none); |
| 83 | |
| 84 | class MockClient: public capnp::DynamicCapability::Client { |
| 85 | public: |
| 86 | using capnp::DynamicCapability::Client::Client; |
| 87 | MockClient(capnp::DynamicCapability::Client&& client) |
| 88 | : capnp::DynamicCapability::Client(kj::mv(client)) {} |
| 89 | |
| 90 | class ExpectedCall { |
| 91 | public: |
| 92 | ExpectedCall(capnp::RemotePromise<capnp::DynamicStruct> promise): promise(kj::mv(promise)) {} |
| 93 | |
| 94 | void expectReturns( |
| 95 | kj::StringPtr resultsText, kj::WaitScope& ws, kj::SourceLocation location = {}) && { |
| 96 | kj::String expectedResults = canonicalizeCapnpText(promise.getSchema(), resultsText); |
| 97 | auto response = promise.wait(ws); |
| 98 | auto actualResults = TEXT_CODEC.encode(response); |
| 99 | KJ_ASSERT_AT(expectedResults == actualResults, location); |
| 100 | } |
| 101 | |
| 102 | void expectThrows(kj::Exception::Type expectedType, |
| 103 | kj::StringPtr expectedMessageSubstring, |
| 104 | kj::WaitScope& ws, |
| 105 | kj::SourceLocation location = {}) { |
| 106 | promise |
| 107 | .then([&](auto&&) { |
| 108 | KJ_FAIL_ASSERT_AT(location, "expected call to throw exception but instead it returned", |
| 109 | expectedType, expectedMessageSubstring); |
| 110 | }, [&](kj::Exception&& e) { |
| 111 | KJ_ASSERT_AT(e.getDescription().contains(expectedMessageSubstring), location, |
| 112 | expectedMessageSubstring, e); |
| 113 | KJ_ASSERT_AT(e.getType() == expectedType, location, e); |
| 114 | }).wait(ws); |
| 115 | } |
| 116 | |
| 117 | private: |
| 118 | capnp::RemotePromise<capnp::DynamicStruct> promise; |
| 119 | }; |
| 120 | |
| 121 | ExpectedCall call(kj::StringPtr methodName, kj::StringPtr params) { |
| 122 | auto req = newRequest(methodName); |
| 123 | TEXT_CODEC.decode(params, req); |
| 124 | return ExpectedCall(req.send()); |
| 125 | } |
| 126 | }; |
| 127 | |
| 128 | // Infrastructure to mock a capability! |
| 129 | // |
| 130 | // TODO(cleanup): This should obviously go in Cap'n Proto! |
| 131 | class MockServer: public kj::Refcounted { |
| 132 | struct ReceivedCall; |
| 133 | |
| 134 | public: |
| 135 | MockServer(capnp::InterfaceSchema schema): schema(schema) {} |
| 136 | |
| 137 | template <typename T> |
| 138 | struct Pair { |
| 139 | kj::Own<MockServer> mock; |
| 140 | T::Client client; |
| 141 | }; |
| 142 | |
| 143 | template <typename T> |
| 144 | static Pair<T> make() { |
| 145 | auto mock = kj::refcounted<MockServer>(capnp::Schema::from<T>()); |
| 146 | capnp::DynamicCapability::Client client = kj::heap<Server>(*mock); |
| 147 | return {kj::mv(mock), client.as<T>()}; |
| 148 | } |
| 149 | |
| 150 | class ExpectedCall { |
| 151 | public: |
| 152 | ExpectedCall(ReceivedCall& received): maybeReceived(received) { |
| 153 | received.expectedCall = this; |
| 154 | } |
| 155 | ExpectedCall(ExpectedCall&& other): maybeReceived(kj::mv(other.maybeReceived)) { |
| 156 | KJ_IF_SOME(r, maybeReceived) r.expectedCall = *this; |
| 157 | } |
| 158 | ~ExpectedCall() noexcept(false) { |
| 159 | KJ_IF_SOME(r, maybeReceived) { |
| 160 | KJ_ASSERT(&KJ_ASSERT_NONNULL(r.expectedCall) == this); |
| 161 | r.expectedCall = kj::none; |
| 162 | } |
| 163 | } |
| 164 | |
| 165 | ExpectedCall withParams(kj::StringPtr paramsText, |
| 166 | kj::Maybe<kj::StringPtr> capName = kj::none, |
| 167 | kj::SourceLocation location = {}) && |
| 168 | KJ_WARN_UNUSED_RESULT { |
| 169 | // Expect that the call had the given parameters. |
| 170 | |
| 171 | auto& received = getReceived(location); |
| 172 | |
| 173 | kj::String expectedParams = |
| 174 | canonicalizeCapnpText(received.method.getParamType(), paramsText, capName); |
| 175 | |
| 176 | auto actualParams = TEXT_CODEC.encode(received.context.getParams()); |
| 177 | KJ_ASSERT_AT(expectedParams == actualParams, location); |
| 178 | |
| 179 | return kj::mv(*this); |
| 180 | } |
| 181 | |
| 182 | // Helper for cases where the received call is expected to invoke some callback capability. |
| 183 | // |
| 184 | // Expect that the params contain a field named `callbackName` whose type is an interface. |
| 185 | // `func()` will be invoked and passed a `MockClient` representing this capability. It can |
| 186 | // then invoke the callback as it seems fit. |
| 187 | // |
| 188 | // Note that it's explicitly OK if `func` captures a `WaitScope` and uses it. In this way, |
| 189 | // the incoming call can be delayed from returning until the callback completes. |
| 190 | template <typename Func> |
| 191 | ExpectedCall useCallback( |
| 192 | kj::StringPtr callbackName, Func&& func, kj::SourceLocation location = {}) && |
| 193 | KJ_WARN_UNUSED_RESULT { |
| 194 | auto& received = getReceived(location); |
| 195 | func(received.context.getParams().get(callbackName).as<capnp::DynamicCapability>()); |
| 196 | return kj::mv(*this); |
| 197 | } |
| 198 | |
| 199 | // Causes the method to return the given result message, which is parsed from text. |
| 200 | void thenReturn(kj::StringPtr message, kj::SourceLocation location = {}) && { |
| 201 | auto& received = getReceived(location); |
| 202 | TEXT_CODEC.decode(message, received.context.getResults()); |
| 203 | received.fulfiller.fulfill(); |
| 204 | } |
| 205 | |
| 206 | // Causes the method to return the given result message, which is parsed from text. |
| 207 | // All capabilities in the result message will be filled in, with MockServer instances |
| 208 | // returned in the hashmap. |
| 209 | kj::HashMap<kj::String, kj::Own<MockServer>> thenReturnWithMocks( |
| 210 | kj::StringPtr message, kj::SourceLocation location = {}) && { |
| 211 | auto& received = getReceived(location); |
| 212 | auto callResults = received.context.getResults(); |
| 213 | auto results = kj::HashMap<kj::String, kj::Own<MockServer>>(); |
| 214 | TEXT_CODEC.decode(message, callResults); |
| 215 | for (const auto& field: received.method.getResultType().getFields()) { |
| 216 | if (field.getType().isInterface()) { |
| 217 | auto name = field.getProto().getName(); |
| 218 | auto mockServer = kj::refcounted<MockServer>(field.getType().asInterface()); |
| 219 | callResults.set(name, kj::heap<Server>(*mockServer)); |
| 220 | results.insert(kj::str(name), kj::mv(mockServer)); |
| 221 | } |
| 222 | } |
| 223 | |
| 224 | received.fulfiller.fulfill(); |
| 225 | return kj::mv(results); |
| 226 | } |
| 227 | |
| 228 | // Causes the method to throw an exception |
| 229 | void thenThrow(kj::Exception&& e, kj::SourceLocation location = {}) && { |
| 230 | auto& received = getReceived(location); |
| 231 | received.fulfiller.reject(kj::mv(e)); |
| 232 | } |
| 233 | |
| 234 | // Return a new mock capability. The method result type is expected to contain a single |
| 235 | // field with the given name whose type is an interface type. It will be filled in with a |
| 236 | // new mock object, and the MockServer is returned in order to set further expectations. |
| 237 | kj::Own<MockServer> returnMock(kj::StringPtr fieldName, kj::SourceLocation location = {}) && { |
| 238 | auto& received = getReceived(location); |
| 239 | auto field = received.method.getResultType().getFieldByName(fieldName); |
| 240 | auto result = kj::refcounted<MockServer>(field.getType().asInterface()); |
| 241 | received.context.getResults().set(field, kj::heap<Server>(*result)); |
| 242 | received.fulfiller.fulfill(); |
| 243 | return result; |
| 244 | } |
| 245 | |
| 246 | void expectCanceled(kj::SourceLocation location = {}) { |
| 247 | KJ_ASSERT_AT(maybeReceived == kj::none, location, "call has not been canceled"); |
| 248 | } |
| 249 | |
| 250 | private: |
| 251 | kj::Maybe<ReceivedCall&> maybeReceived; |
| 252 | ReceivedCall& getReceived(kj::SourceLocation location) { |
| 253 | return KJ_REQUIRE_NONNULL_AT(maybeReceived, location, "call was unexpectedly canceled"); |
| 254 | } |
| 255 | friend struct ReceivedCall; |
| 256 | }; |
| 257 | |
| 258 | ExpectedCall expectCall(kj::StringPtr methodName, |
| 259 | kj::WaitScope& waitScope, |
| 260 | kj::SourceLocation location = {}) KJ_WARN_UNUSED_RESULT { |
| 261 | auto expectedMethod = schema.getMethodByName(methodName); |
| 262 | |
| 263 | KJ_ASSERT_AT( |
| 264 | waitForEvent(waitScope), location, "no method call was received when expected", methodName); |
| 265 | |
| 266 | KJ_ASSERT_AT( |
| 267 | !dropped, location, "capability was dropped without making expected call", methodName); |
| 268 | |
| 269 | auto& received = receivedCalls.front(); |
| 270 | receivedCalls.remove(received); |
| 271 | |
| 272 | KJ_ASSERT_AT(received.method == expectedMethod, location, |
| 273 | "a different method was called than expected", received.method.getProto().getName(), |
| 274 | expectedMethod.getProto().getName()); |
| 275 | |
| 276 | return ExpectedCall(received); |
| 277 | } |
| 278 | |
| 279 | void expectDropped(kj::WaitScope& waitScope, kj::SourceLocation location = {}) { |
| 280 | KJ_ASSERT_AT(waitForEvent(waitScope), location, "capability was not dropped when expected"); |
| 281 | KJ_ASSERT_AT( |
| 282 | receivedCalls.empty(), location, receivedCalls.front().method.getProto().getName()); |
| 283 | |
| 284 | KJ_ASSERT(dropped); // should always be true if receivedCalls is empty |
| 285 | } |
| 286 | |
| 287 | void expectNoActivity(kj::WaitScope& waitScope, kj::SourceLocation location = {}) { |
| 288 | if (waitForEvent(waitScope)) { |
| 289 | if (!receivedCalls.empty()) { |
| 290 | KJ_FAIL_ASSERT_AT(location, "unexpected call received", |
| 291 | receivedCalls.front().method.getProto().getName()); |
| 292 | } |
| 293 | if (dropped) { |
| 294 | KJ_FAIL_ASSERT_AT(location, "mock capability unexpectedly dropped"); |
| 295 | } |
| 296 | } |
| 297 | } |
| 298 | |
| 299 | private: |
| 300 | capnp::InterfaceSchema schema; |
| 301 | kj::Maybe<kj::Own<kj::PromiseFulfiller<void>>> waiter; |
| 302 | |
| 303 | struct ReceivedCall { |
| 304 | ReceivedCall(kj::PromiseFulfiller<void>& fulfiller, |
| 305 | MockServer& mock, |
| 306 | capnp::InterfaceSchema::Method method, |
| 307 | capnp::CallContext<capnp::DynamicStruct, capnp::DynamicStruct> context) |
| 308 | : fulfiller(fulfiller), |
| 309 | mock(mock), |
| 310 | method(method), |
| 311 | context(kj::mv(context)) { |
| 312 | mock.receivedCalls.add(*this); |
| 313 | KJ_IF_SOME(w, mock.waiter) { |
| 314 | w.get()->fulfill(); |
| 315 | } |
| 316 | } |
| 317 | ~ReceivedCall() noexcept(false) { |
| 318 | if (link.isLinked()) { |
| 319 | mock.receivedCalls.remove(*this); |
| 320 | } |
| 321 | KJ_IF_SOME(e, expectedCall) { |
| 322 | e.maybeReceived = kj::none; |
| 323 | } |
| 324 | } |
| 325 | KJ_DISALLOW_COPY_AND_MOVE(ReceivedCall); |
| 326 | |
| 327 | kj::PromiseFulfiller<void>& fulfiller; |
| 328 | MockServer& mock; |
| 329 | capnp::InterfaceSchema::Method method; |
| 330 | capnp::CallContext<capnp::DynamicStruct, capnp::DynamicStruct> context; |
| 331 | kj::ListLink<ReceivedCall> link; |
| 332 | |
| 333 | kj::Maybe<ExpectedCall&> expectedCall; // if one is attached |
| 334 | }; |
| 335 | |
| 336 | kj::List<ReceivedCall, &ReceivedCall::link> receivedCalls; |
| 337 | bool dropped = false; |
| 338 | |
| 339 | bool waitForEvent(kj::WaitScope& waitScope) { |
| 340 | if (receivedCalls.empty() && !dropped) { |
| 341 | auto paf = kj::newPromiseAndFulfiller<void>(); |
| 342 | waiter = kj::mv(paf.fulfiller); |
| 343 | if (!paf.promise.poll(waitScope)) { |
| 344 | waiter = kj::none; |
| 345 | return false; |
| 346 | } |
| 347 | paf.promise.wait(waitScope); |
| 348 | } |
| 349 | return true; |
| 350 | } |
| 351 | |
| 352 | class Server final: public capnp::DynamicCapability::Server { |
| 353 | public: |
| 354 | Server(MockServer& mock) |
| 355 | : capnp::DynamicCapability::Server(mock.schema, {.allowCancellation = true}), |
| 356 | mock(kj::addRef(mock)) {} |
| 357 | ~Server() noexcept(false) { |
| 358 | mock->dropped = true; |
| 359 | KJ_IF_SOME(w, mock->waiter) { |
| 360 | w.get()->fulfill(); |
| 361 | } |
| 362 | } |
| 363 | |
| 364 | kj::Promise<void> call(capnp::InterfaceSchema::Method method, |
| 365 | capnp::CallContext<capnp::DynamicStruct, capnp::DynamicStruct> context) override { |
| 366 | return kj::newAdaptedPromise<void, ReceivedCall>(*mock, method, kj::mv(context)); |
| 367 | } |
| 368 | |
| 369 | private: |
| 370 | kj::Own<MockServer> mock; |
| 371 | }; |
| 372 | }; |
| 373 | |
| 374 | // Wraps a "capnp struct literal". This actually just stringifies the arguments, adding enclosing |
| 375 | // parentheses. The nice thing about it, though, is that you don't have to escape quotes inside |
| 376 | // the literal. |
| 377 | #define CAPNP(...) ("(" #__VA_ARGS__ ")"_kj) |
| 378 | |
| 379 | template <typename Schema, typename InitFunc> |
| 380 | kj::String Capnp(InitFunc func) { |
| 381 | capnp::MallocMessageBuilder message; |
| 382 | auto builder = message.initRoot<Schema>(); |
| 383 | func(builder); |
| 384 | return TEXT_CODEC.encode(builder.asReader()); |
| 385 | } |
| 386 | |
| 387 | } // namespace workerd |