Skip to content
File

Blob: src/workerd/server/channel-token.c++

10.1 KB
1// Copyright (c) 2025 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 "channel-token.h"
6 
7#include <workerd/server/channel-token.capnp.h>
8#include <workerd/util/entropy.h>
9 
10#include <openssl/evp.h>
11#include <openssl/sha.h>
12 
13#include <capnp/serialize-packed.h>
14#include <kj/io.h>
15 
16// It's 2025, nobody uses big-endian anymore. But just in case someone tries, flag it here.
17// Specifically, the magic number in TokenHeader is encoded in host order.
18#if !defined(__LITTLE_ENDIAN__) || __BYTE_ORDER != __LITTLE_ENDIAN
19#error "This code assumes little-endian architecture."
20#endif
21 
22namespace workerd::server {
23 
24ChannelTokenHandler::ChannelTokenHandler(Resolver& resolver): resolver(resolver) {
25 getEntropy(tokenKey);
26 
27 SHA256_CTX ctx{};
28 KJ_ASSERT(SHA256_Init(&ctx));
29 KJ_ASSERT(SHA256_Update(&ctx, tokenKey, sizeof(tokenKey)));
30 
31 byte hash[SHA256_DIGEST_LENGTH]{};
32 KJ_ASSERT(SHA256_Final(hash, &ctx));
33 
34 static_assert(KEY_ID_SIZE <= SHA256_DIGEST_LENGTH);
35 kj::arrayPtr(keyId).copyFrom(kj::arrayPtr(hash).first(KEY_ID_SIZE));
36}
37 
38kj::Array<byte> ChannelTokenHandler::encodeChannelTokenImpl(ChannelToken::Type type,
39 IoChannelFactory::ChannelTokenUsage usage,
40 kj::StringPtr serviceName,
41 kj::Maybe<kj::StringPtr> entrypoint,
42 Frankenvalue& props) {
43 capnp::word scratch[128]{};
44 capnp::MallocMessageBuilder message(scratch);
45 auto builder = message.getRoot<ChannelToken>();
46 
47 builder.setType(type);
48 
49 builder.setName(serviceName);
50 
51 KJ_IF_SOME(e, entrypoint) {
52 builder.setEntrypoint(e);
53 }
54 
55 {
56 auto propsBuilder = builder.initProps();
57 props.toCapnp(propsBuilder);
58 
59 auto capTable = props.getCapTable();
60 if (capTable.size() > 0) {
61 auto tableBuilder = propsBuilder.initCapTable().initAs<ChannelToken::FrankenvalueCapTable>();
62 
63 auto caps = tableBuilder.initCaps(capTable.size());
64 
65 for (auto i: kj::indices(capTable)) {
66 KJ_IF_SOME(subreq, kj::tryDowncast<IoChannelFactory::SubrequestChannel>(*capTable[i])) {
67 caps[i].setSubrequestChannel(subreq.getToken(usage));
68 } else KJ_IF_SOME(actorClass,
69 kj::tryDowncast<IoChannelFactory::ActorClassChannel>(*capTable[i])) {
70 caps[i].setActorClassChannel(actorClass.getToken(usage));
71 } else {
72 KJ_FAIL_REQUIRE("unknown type in props");
73 }
74 }
75 }
76 }
77 
78 kj::VectorOutputStream out;
79 capnp::writePackedMessage(out, message);
80 
81 auto plaintext = out.getArray();
82 
83 switch (usage) {
84 case IoChannelFactory::ChannelTokenUsage::RPC: {
85 static_assert(alignof(TokenHeader) <= __STDCPP_DEFAULT_NEW_ALIGNMENT__);
86 auto result = kj::heapArray<byte>(sizeof(TokenHeader) + plaintext.size() + AES_MAC_SIZE);
87 auto& header = *reinterpret_cast<TokenHeader*>(result.begin());
88 
89 header.magic = ChannelToken::RPC_TOKEN_MAGIC;
90 getEntropy(header.iv);
91 kj::arrayPtr(header.keyId).copyFrom(keyId);
92 
93 EVP_CIPHER_CTX* aesCtx = EVP_CIPHER_CTX_new();
94 KJ_ASSERT(aesCtx != nullptr);
95 KJ_DEFER(EVP_CIPHER_CTX_free(aesCtx));
96 
97 KJ_ASSERT(EVP_EncryptInit(aesCtx, EVP_aes_256_gcm(), tokenKey, header.iv));
98 
99 // Add header as AAD first.
100 {
101 int outSize = 0;
102 KJ_ASSERT(
103 EVP_EncryptUpdate(aesCtx, nullptr, &outSize, result.begin(), sizeof(TokenHeader)));
104 KJ_ASSERT(outSize == sizeof(TokenHeader));
105 }
106 
107 // Encrypt the body.
108 {
109 int outSize = 0;
110 KJ_ASSERT(EVP_EncryptUpdate(aesCtx, result.begin() + sizeof(TokenHeader), &outSize,
111 plaintext.begin(), plaintext.size()));
112 KJ_ASSERT(outSize == plaintext.size()); // because AES-GCM is a stream cipher
113 }
114 
115 int out = 0;
116 KJ_ASSERT(EVP_EncryptFinal_ex(aesCtx, nullptr, &out));
117 KJ_ASSERT(out == 0); // No padding for stream ciphers like AES-GCM.
118 
119 // Get the MAC.
120 KJ_ASSERT(EVP_CIPHER_CTX_ctrl(
121 aesCtx, EVP_CTRL_GCM_GET_TAG, AES_MAC_SIZE, result.end() - AES_MAC_SIZE));
122 
123 return result;
124 }
125 
126 case IoChannelFactory::ChannelTokenUsage::STORAGE: {
127 auto magic = kj::asBytes(ChannelToken::STORAGE_TOKEN_MAGIC);
128 auto result = kj::heapArray<byte>(magic.size() + plaintext.size());
129 result.slice(0, magic.size()).copyFrom(magic);
130 result.slice(magic.size()).copyFrom(plaintext);
131 return result;
132 }
133 }
134 
135 KJ_UNREACHABLE;
136}
137 
138kj::Array<byte> ChannelTokenHandler::encodeSubrequestChannelToken(
139 IoChannelFactory::ChannelTokenUsage usage,
140 kj::StringPtr serviceName,
141 kj::Maybe<kj::StringPtr> entrypoint,
142 Frankenvalue& props) {
143 return encodeChannelTokenImpl(
144 ChannelToken::Type::SUBREQUEST, usage, serviceName, entrypoint, props);
145}
146 
147kj::Array<byte> ChannelTokenHandler::encodeActorClassChannelToken(
148 IoChannelFactory::ChannelTokenUsage usage,
149 kj::StringPtr serviceName,
150 kj::Maybe<kj::StringPtr> entrypoint,
151 Frankenvalue& props) {
152 return encodeChannelTokenImpl(
153 ChannelToken::Type::ACTOR_CLASS, usage, serviceName, entrypoint, props);
154}
155 
156kj::Own<Frankenvalue::CapTableEntry> ChannelTokenHandler::decodeChannelTokenImpl(
157 ChannelToken::Type type,
158 IoChannelFactory::ChannelTokenUsage usage,
159 kj::ArrayPtr<const byte> token) {
160 kj::ArrayPtr<const byte> plaintext;
161 kj::Array<byte> ownPlaintext;
162 
163 switch (usage) {
164 case IoChannelFactory::ChannelTokenUsage::RPC: {
165 TokenHeader header;
166 KJ_REQUIRE(token.size() >= sizeof(header) + AES_MAC_SIZE, "invalid channel token for RPC");
167 
168 kj::asBytes(header).copyFrom(token.first(sizeof(TokenHeader)));
169 KJ_REQUIRE(header.magic == ChannelToken::RPC_TOKEN_MAGIC, "invalid channel token for RPC");
170 
171 auto mac = token.slice(token.size() - AES_MAC_SIZE);
172 auto ciphertext = token.slice(sizeof(header), token.size() - AES_MAC_SIZE);
173 
174 EVP_CIPHER_CTX* aesCtx = EVP_CIPHER_CTX_new();
175 KJ_ASSERT(aesCtx != nullptr);
176 KJ_DEFER(EVP_CIPHER_CTX_free(aesCtx));
177 
178 KJ_ASSERT(EVP_DecryptInit(aesCtx, EVP_aes_256_gcm(), tokenKey, header.iv));
179 
180 // Add header as AAD first.
181 {
182 int outSize = 0;
183 KJ_ASSERT(EVP_DecryptUpdate(aesCtx, nullptr, &outSize, token.begin(), sizeof(header)));
184 KJ_ASSERT(outSize == sizeof(TokenHeader));
185 }
186 
187 // Decrypt the body.
188 ownPlaintext = kj::heapArray<byte>(ciphertext.size());
189 {
190 int outSize = 0;
191 KJ_ASSERT(EVP_DecryptUpdate(
192 aesCtx, ownPlaintext.begin(), &outSize, ciphertext.begin(), ciphertext.size()));
193 KJ_ASSERT(outSize == ownPlaintext.size());
194 }
195 plaintext = ownPlaintext;
196 
197 // Check MAC.
198 KJ_ASSERT(EVP_CIPHER_CTX_ctrl(aesCtx, EVP_CTRL_GCM_SET_TAG, AES_MAC_SIZE,
199 // const_cast needed for the EVP_CIPHER_CTX_ctrl() interface, but this won't actually
200 // modify the buffer.
201 const_cast<byte*>(mac.begin())));
202 
203 int out;
204 KJ_REQUIRE(EVP_DecryptFinal_ex(aesCtx, nullptr, &out), "channel token failed authentication");
205 KJ_ASSERT(out == 0);
206 
207 break;
208 }
209 
210 case IoChannelFactory::ChannelTokenUsage::STORAGE: {
211 uint32_t magic;
212 KJ_REQUIRE(token.size() >= sizeof(magic), "invalid channel token for storage");
213 
214 kj::asBytes(magic).copyFrom(token.first(sizeof(magic)));
215 KJ_REQUIRE(magic == ChannelToken::STORAGE_TOKEN_MAGIC, "invalid channel token for storage");
216 
217 plaintext = token.slice(sizeof(magic));
218 break;
219 }
220 }
221 
222 kj::ArrayInputStream input(plaintext);
223 capnp::word scratch[128]{};
224 capnp::PackedMessageReader message(input, {}, scratch);
225 auto reader = message.getRoot<ChannelToken>();
226 
227 KJ_REQUIRE(reader.getType() == type, "channel token type mismatch");
228 
229 kj::Maybe<kj::StringPtr> entrypoint;
230 if (reader.hasEntrypoint()) {
231 entrypoint = reader.getEntrypoint();
232 }
233 
234 Frankenvalue props;
235 if (reader.hasProps()) {
236 auto propsReader = reader.getProps();
237 auto tableReader = propsReader.getCapTable().getAs<ChannelToken::FrankenvalueCapTable>();
238 
239 kj::Vector<kj::Own<Frankenvalue::CapTableEntry>> capTable;
240 if (tableReader.hasCaps()) {
241 auto caps = tableReader.getCaps();
242 capTable.reserve(caps.size());
243 
244 for (auto cap: caps) {
245 switch (cap.which()) {
246 case ChannelToken::FrankenvalueCapTable::Cap::UNKNOWN:
247 break;
248 case ChannelToken::FrankenvalueCapTable::Cap::SUBREQUEST_CHANNEL:
249 capTable.add(decodeSubrequestChannelToken(usage, cap.getSubrequestChannel()));
250 continue;
251 case ChannelToken::FrankenvalueCapTable::Cap::ACTOR_CLASS_CHANNEL:
252 capTable.add(decodeActorClassChannelToken(usage, cap.getActorClassChannel()));
253 continue;
254 }
255 KJ_FAIL_REQUIRE("unknown cap table type", cap.which());
256 }
257 }
258 
259 props = Frankenvalue::fromCapnp(propsReader, kj::mv(capTable));
260 }
261 
262 // HACK: It would be more type-safe for us to return the (name, entrypoint, props) triplet and
263 // let the caller call the appropriate resolver method. However, this would require making
264 // heap string copies of the name and entrypoint which would just be thrown way immediately.
265 // Since both types happen to subclass Frankenvalue::CapTableEntry, we just make the resolver
266 // call here, return either type, and let the caller downcast to the right type.
267 switch (type) {
268 case ChannelToken::Type::SUBREQUEST:
269 return resolver.resolveEntrypoint(reader.getName(), entrypoint, kj::mv(props));
270 case ChannelToken::Type::ACTOR_CLASS:
271 return resolver.resolveActorClass(reader.getName(), entrypoint, kj::mv(props));
272 }
273 
274 KJ_UNREACHABLE;
275}
276 
277kj::Own<IoChannelFactory::SubrequestChannel> ChannelTokenHandler::decodeSubrequestChannelToken(
278 IoChannelFactory::ChannelTokenUsage usage, kj::ArrayPtr<const byte> token) {
279 return decodeChannelTokenImpl(ChannelToken::Type::SUBREQUEST, usage, token)
280 .downcast<IoChannelFactory::SubrequestChannel>();
281}
282 
283kj::Own<IoChannelFactory::ActorClassChannel> ChannelTokenHandler::decodeActorClassChannelToken(
284 IoChannelFactory::ChannelTokenUsage usage, kj::ArrayPtr<const byte> token) {
285 return decodeChannelTokenImpl(ChannelToken::Type::ACTOR_CLASS, usage, token)
286 .downcast<IoChannelFactory::ActorClassChannel>();
287}
288 
289} // namespace workerd::server