Skip to content
File

Blob: src/workerd/util/capnp-mock.h

cpp388 lines
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 
16namespace 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 
79const capnp::TextCodec TEXT_CODEC;
80 
81kj::String canonicalizeCapnpText(
82 capnp::StructSchema schema, kj::StringPtr text, kj::Maybe<kj::StringPtr> capName = kj::none);
83 
84class 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!
131class 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 
379template <typename Schema, typename InitFunc>
380kj::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