-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathProtocolSupport.java
More file actions
486 lines (468 loc) · 21.6 KB
/
Copy pathProtocolSupport.java
File metadata and controls
486 lines (468 loc) · 21.6 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
package io.smithycpp.codegen;
import java.util.List;
import java.util.Map;
import java.util.TreeMap;
import software.amazon.smithy.codegen.core.CodegenException;
import software.amazon.smithy.model.shapes.MemberShape;
import software.amazon.smithy.model.shapes.OperationShape;
import software.amazon.smithy.model.shapes.ServiceShape;
import software.amazon.smithy.model.shapes.Shape;
import software.amazon.smithy.model.shapes.ShapeId;
import software.amazon.smithy.model.shapes.StructureShape;
import software.amazon.smithy.model.traits.IdempotencyTokenTrait;
import software.amazon.smithy.model.traits.RetryableTrait;
/** Emission helpers shared by the concrete protocol generators. */
final class ProtocolSupport {
private ProtocolSupport() {}
/**
* Emits the protocol's shared error-parsing helpers (SanitizeErrorCode, ParsedError, ParseError,
* GenericError); only the wire decode and the code-carrying header differ per protocol.
*/
static void writeErrorSupport(CppWriter w, String decodeStatement, boolean errorTypeHeader) {
w.write("// The error shape name arrives namespaced (\"ns#Shape\") and possibly");
w.write("// URI-qualified; modeled error codes keep only the shape name.");
w.openBlock("std::string SanitizeErrorCode(std::string_view raw) {");
w.write(
"if (const auto colon = raw.find(':'); colon != std::string_view::npos) "
+ "raw = raw.substr(0, colon);");
w.write(
"if (const auto hash = raw.find('#'); hash != std::string_view::npos) "
+ "raw = raw.substr(hash + 1);");
w.write("return std::string(raw);");
w.closeBlock("}");
w.write("");
w.openBlock("struct ParsedError {");
w.write("int status = 0;");
w.write("std::string code = \"UnknownError\";");
w.write("std::string message;");
w.write("smithy::Document doc;");
w.closeBlock("};");
w.write("");
w.openBlock("ParsedError ParseError(const smithy::http::HttpResponse& response) {");
w.write("ParsedError parsed;");
w.write("parsed.status = response.status;");
w.write("parsed.message = \"HTTP \" + std::to_string(response.status);");
w.write(decodeStatement);
w.write("if (doc.ok()) parsed.doc = *std::move(doc);");
if (errorTypeHeader) {
w.write("const auto type_header = response.headers.Get(\"x-amzn-errortype\");");
w.write("if (type_header.has_value()) parsed.code = SanitizeErrorCode(*type_header);");
}
w.openBlock("if (parsed.doc.is_map()) {");
w.write("const smithy::Document* type = parsed.doc.Find(\"__type\");");
w.write("if (type == nullptr) type = parsed.doc.Find(\"code\");");
w.write(
"if (parsed.code == \"UnknownError\" && type != nullptr && type->is_string()) "
+ "parsed.code = SanitizeErrorCode(type->as_string());");
w.write("const smithy::Document* text = parsed.doc.Find(\"message\");");
w.write("if (text != nullptr && text->is_string()) parsed.message = text->as_string();");
w.closeBlock("}");
w.write("return parsed;");
w.closeBlock("}");
w.write("");
w.openBlock("smithy::Error GenericError(ParsedError parsed) {");
w.write("const bool retryable = parsed.status >= 500;");
w.write(
"if (parsed.code == \"UnknownError\") return smithy::Error(smithy::ErrorKind::kUnknown, "
+ "std::move(parsed.code), std::move(parsed.message), retryable);");
w.write(
"return smithy::Error::Modeled(std::move(parsed.code), std::move(parsed.message), "
+ "retryable);");
w.closeBlock("}");
w.write("");
}
/** Text-to-number helpers used for header/label/query bindings. */
static void writeNumericParseHelpers(CppWriter w) {
w.addInclude("<algorithm>");
w.addInclude("<charconv>");
w.addInclude("<cmath>");
w.addInclude("<cstdlib>");
w.write("// Strict text parsing for label/query/header bindings ([[maybe_unused]]:");
w.write("// emitted for every service; not every service binds numeric values).");
w.write("// Trailing text, floats-for-ints, and out-of-range values are rejected");
w.write("// (the malformed-request suites pin this).");
w.openBlock(
"[[maybe_unused]] smithy::Outcome<std::int64_t> ParseInt64Text(const std::string& text, "
+ "std::int64_t min_value, std::int64_t max_value) {");
w.write("std::int64_t value = 0;");
w.write("const char* first = text.data();");
w.write("const char* last = first + text.size();");
w.write("const auto result = std::from_chars(first, last, value, 10);");
w.write(
"if (text.empty() || result.ec != std::errc() || result.ptr != last || "
+ "value < min_value || value > max_value) {");
w.indent();
w.write("return smithy::Error::Serialization(\"invalid integer: \" + text);");
w.dedent();
w.write("}");
w.write("return value;");
w.closeBlock("}");
w.write("");
// Floating-point std::from_chars is missing on libc++ (Apple), so doubles
// pair a strict character-set check (rejects hex, inf/nan spellings, and
// leading '+'/whitespace strtod would accept) with a fully-consuming strtod.
w.openBlock(
"[[maybe_unused]] smithy::Outcome<double> ParseDoubleText(const std::string& text) {");
w.write("if (text == \"NaN\") return std::numeric_limits<double>::quiet_NaN();");
w.write("if (text == \"Infinity\") return std::numeric_limits<double>::infinity();");
w.write("if (text == \"-Infinity\") return -std::numeric_limits<double>::infinity();");
w.openBlock("const auto valid_char = [](char c) {");
w.write(
"return (c >= '0' && c <= '9') || c == '.' || c == 'e' || c == 'E' || c == '+' || "
+ "c == '-';");
w.closeBlock("};");
w.openBlock(
"if (text.empty() || text.front() == '+' || "
+ "!std::all_of(text.begin(), text.end(), valid_char)) {");
w.write("return smithy::Error::Serialization(\"invalid number: \" + text);");
w.closeBlock("}");
w.write("char* parse_end = nullptr;");
w.write("const double value = std::strtod(text.c_str(), &parse_end);");
w.openBlock("if (parse_end != text.c_str() + text.size() || !std::isfinite(value)) {");
w.write("return smithy::Error::Serialization(\"invalid number: \" + text);");
w.closeBlock("}");
w.write("return value;");
w.closeBlock("}");
w.write("");
}
/** min/max arguments for ParseInt64Text per integer shape type. */
static String int64Bounds(software.amazon.smithy.model.shapes.ShapeType type) {
return switch (type) {
case BYTE -> "-128, 127";
case SHORT -> "-32768, 32767";
case INTEGER, INT_ENUM -> "-2147483648LL, 2147483647LL";
default ->
"std::numeric_limits<std::int64_t>::min(), std::numeric_limits<std::int64_t>::max()";
};
}
/**
* Client-side @requestCompression: gzip the request body once it reaches the configured minimum
* size, appending to any member-bound Content-Encoding header. Emitted after the body and every
* header binding; Send() computes content-length afterwards.
*/
static void writeRequestCompression(
CppWriter w, software.amazon.smithy.model.shapes.OperationShape operation) {
var trait =
operation.getTrait(software.amazon.smithy.model.traits.RequestCompressionTrait.class);
if (trait.isEmpty() || !trait.get().getEncodings().contains("gzip")) {
return;
}
w.addInclude("\"smithy/compression/gzip.h\"");
w.addInclude("<cstddef>");
w.write("// @requestCompression(gzip): applied last, appended to Content-Encoding.");
w.openBlock(
"if (request.body.size() >= "
+ "static_cast<std::size_t>(config_.request_min_compression_size_bytes)) {");
w.write("auto compressed = smithy::GzipCompress(request.body);");
w.write("if (!compressed) return std::move(compressed).error();");
w.write("request.body = *std::move(compressed);");
w.write("const auto existing_encoding = request.headers.Get(\"content-encoding\");");
w.write(
"request.headers.Set(\"content-encoding\", existing_encoding.has_value() && "
+ "!existing_encoding->empty() ? *existing_encoding + \", gzip\" : \"gzip\");");
w.closeBlock("}");
}
/**
* Server-side inverse: transparently gunzip request bodies for @requestCompression operations
* when the (final) Content-Encoding is gzip. Emitted at the top of the route lambda.
*/
static void writeRequestDecompression(
CppWriter w,
software.amazon.smithy.model.shapes.OperationShape operation,
String errorFn,
String errorCode) {
var trait =
operation.getTrait(software.amazon.smithy.model.traits.RequestCompressionTrait.class);
if (trait.isEmpty() || !trait.get().getEncodings().contains("gzip")) {
return;
}
w.addInclude("\"smithy/compression/gzip.h\"");
w.write("// @requestCompression(gzip): decode before parsing.");
w.openBlock(
"if (const auto request_encoding = request.headers.Get(\"content-encoding\"); "
+ "request_encoding.has_value() && (*request_encoding == \"gzip\" || "
+ "request_encoding->ends_with(\", gzip\"))) {");
w.write("auto decompressed = smithy::GzipDecompress(request.body);");
w.openBlock("if (!decompressed) {");
w.write("return $L(400, $S, \"invalid gzip request body\", {});", errorFn, errorCode);
w.closeBlock("}");
w.write("request.body = *std::move(decompressed);");
w.closeBlock("}");
}
/** HTTP status for a modeled error shape: @httpError, else @error class default. */
static int errorStatus(StructureShape shape) {
var httpError = shape.getTrait(software.amazon.smithy.model.traits.HttpErrorTrait.class);
if (httpError.isPresent()) {
return httpError.get().getCode();
}
var error = shape.expectTrait(software.amazon.smithy.model.traits.ErrorTrait.class);
return error.isClientError() ? 400 : 500;
}
/**
* Emits the server's ErrorToResponse over {@code errorBodyFn(status, code, message, body)}:
* modeled errors get their @httpError status and serialized detail; validation/serialization
* failures map to 400; anything else is a non-leaking 500.
*/
static void writeServerErrorToResponse(
CppWriter w,
CppContext context,
ServiceShape service,
List<OperationShape> operations,
String errorBodyFn,
boolean errortypeHeader) {
Map<String, StructureShape> errorShapes = new TreeMap<>();
for (OperationShape operation : operations) {
for (ShapeId errorId : operation.getErrors(service)) {
StructureShape shape =
context.model().expectShape(errorId).asStructureShape().orElseThrow();
errorShapes.put(context.cppSymbols().toSymbol(shape).getName(), shape);
}
}
w.openBlock("smithy::http::HttpResponse ErrorToResponse(const smithy::Error& error) {");
if (errortypeHeader) {
w.write("std::vector<std::pair<std::string, std::string>> header_values;");
w.write("(void)header_values;");
}
w.openBlock("if (error.kind() == smithy::ErrorKind::kModeled) {");
for (StructureShape shape : errorShapes.values()) {
String type = context.cppSymbols().toSymbol(shape).getName();
w.openBlock("if (error.code() == $S) {", shape.getId().getName());
w.write("smithy::DocumentMap body;");
w.openBlock("if (const auto* detail = error.detail<$L>()) {", type);
w.write(
"body = Serialize$L(*detail).as_map();",
SerdeCodeGen.serdeFunctionSuffix(context, shape));
w.closeBlock("}");
w.write("// The typed detail's own message member wins over the generic one.");
w.write(
"const bool has_message = body.count(\"message\") != 0 || "
+ "body.count(\"Message\") != 0;");
w.openBlock("if (!has_message && !error.message().empty()) {");
w.write("body.emplace(\"message\", smithy::Document(error.message()));");
w.closeBlock("}");
if (errortypeHeader) {
// restJson1: @httpHeader-bound error members travel as headers, and the
// error shape name in X-Amzn-Errortype rather than the body.
var index = software.amazon.smithy.model.knowledge.HttpBindingIndex.of(context.model());
for (var binding : new TreeMap<>(index.getResponseBindings(shape)).values()) {
if (binding.getLocation()
!= software.amazon.smithy.model.knowledge.HttpBinding.Location.HEADER) {
continue;
}
Shape target = context.model().expectShape(binding.getMember().getTarget());
if (!target.isStringShape()
|| target.hasTrait(software.amazon.smithy.model.traits.MediaTypeTrait.class)) {
// Non-string error headers are not serialized yet (exclusions document this).
continue;
}
w.openBlock(
"if (auto it = body.find($S); it != body.end()) {",
binding.getMember().getMemberName());
w.write(
"if (it->second.is_string()) header_values.emplace_back($S, "
+ "it->second.as_string());",
binding.getLocationName());
w.write("body.erase(it);");
w.closeBlock("}");
}
w.write(
"auto response = $L($L, \"\", \"\", std::move(body));",
errorBodyFn,
errorStatus(shape));
w.write("response.headers.Set(\"x-amzn-errortype\", error.code());");
w.write(
"for (const auto& [name, value] : header_values) response.headers.Set(name, value);");
w.write("return response;");
} else {
// rpcv2Cbor: __type carries the fully qualified shape id in the body.
w.write(
"return $L($L, $S, \"\", std::move(body));",
errorBodyFn,
errorStatus(shape),
shape.getId().toString());
}
w.closeBlock("}");
}
w.write("return $L(400, error.code(), error.message(), {});", errorBodyFn);
w.closeBlock("}");
if (errortypeHeader) {
// restJson1: parse failures answer 400 with the SerializationException
// error identity in the header (the malformed-request suite pins this).
w.openBlock(
"if (error.kind() == smithy::ErrorKind::kValidation || error.kind() == "
+ "smithy::ErrorKind::kSerialization) {");
w.write("auto response = $L(400, \"\", error.message(), {});", errorBodyFn);
w.write("response.headers.Set(\"x-amzn-errortype\", \"SerializationException\");");
w.write("return response;");
w.closeBlock("}");
} else {
w.write(
"if (error.kind() == smithy::ErrorKind::kValidation || error.kind() == "
+ "smithy::ErrorKind::kSerialization) return $L(400, \"SerializationException\", "
+ "error.message(), {});",
errorBodyFn);
}
w.write("// Never leak internal detail on unexpected failures.");
w.write("return $L(500, \"InternalFailure\", \"internal failure\", {});", errorBodyFn);
w.closeBlock("}");
w.write("");
}
static List<String> sharedServerIncludes(CppContext context) {
return List.of(
"\"" + context.settings().includePrefix() + "/server.h\"",
"\"" + context.settings().includePrefix() + "/serde.h\"",
"\"smithy/server/router.h\"",
"<memory>",
"<utility>",
"<vector>",
"<string>",
"<string_view>",
"<utility>");
}
/**
* Emits Make<Error>Error for every error shape any operation declares (typed detail via the
* shape's serde, @retryable honored) plus a Deserialize<Op>Error dispatcher per operation
* that has errors. Must run inside the client source's anonymous namespace, after {@link
* #writeErrorSupport}.
*/
static void writeOperationErrorDeserializers(
CppWriter w,
CppContext context,
ServiceShape service,
ProtocolGenerator protocol,
List<OperationShape> operations) {
Map<String, StructureShape> errorShapes = new TreeMap<>();
for (OperationShape operation : operations) {
for (ShapeId errorId : operation.getErrors(service)) {
StructureShape shape =
context.model().expectShape(errorId).asStructureShape().orElseThrow();
errorShapes.put(context.cppSymbols().toSymbol(shape).getName(), shape);
}
}
for (StructureShape shape : errorShapes.values()) {
String type = context.cppSymbols().toSymbol(shape).getName();
boolean retryable = shape.hasTrait(RetryableTrait.class);
w.openBlock(
"smithy::Error Make$LError(const smithy::http::HttpResponse& response, "
+ "ParsedError parsed) {",
type);
w.write("(void)response;");
if (retryable) {
w.write("const bool retryable = true; // @retryable");
} else {
w.write("const bool retryable = parsed.status >= 500;");
}
// code() carries the wire-level shape name, which can differ from the
// C++ type name when a foreign-namespace shape was disambiguated.
w.write(
"smithy::Error error = smithy::Error::Modeled($S, std::move(parsed.message), "
+ "retryable);",
shape.getId().getName());
// Errors with header-only payloads (or none at all) may have no body:
// deserialize from an empty map so the typed detail still attaches.
w.write("if (!parsed.doc.is_map()) parsed.doc = smithy::Document(smithy::DocumentMap{});");
w.write(
"auto detail = Deserialize$L(parsed.doc);",
SerdeCodeGen.serdeFunctionSuffix(context, shape));
w.openBlock("if (detail.ok()) {");
protocol.writeErrorDetailPatches(w, context, shape);
w.write("error.set_detail(*std::move(detail));");
w.closeBlock("}");
w.write("return error;");
w.closeBlock("}");
w.write("");
}
for (OperationShape operation : operations) {
List<ShapeId> errors = operation.getErrors(service);
if (errors.isEmpty()) {
continue;
}
Map<String, StructureShape> sorted = new TreeMap<>();
for (ShapeId errorId : errors) {
StructureShape shape =
context.model().expectShape(errorId).asStructureShape().orElseThrow();
sorted.put(context.cppSymbols().toSymbol(shape).getName(), shape);
}
w.openBlock(
"smithy::Error Deserialize$LError(const smithy::http::HttpResponse& response) {",
CppReservedWords.escape(operation.getId().getName()));
w.write("ParsedError parsed = ParseError(response);");
for (Map.Entry<String, StructureShape> entry : sorted.entrySet()) {
w.write(
"if (parsed.code == $S) return Make$LError(response, std::move(parsed));",
entry.getValue().getId().getName(),
entry.getKey());
}
w.write("return GenericError(std::move(parsed));");
w.closeBlock("}");
w.write("");
}
}
/** The expression a protocol's operation body returns for a non-success response. */
static String errorExpression(CppContext context, ServiceShape service, OperationShape op) {
if (op.getErrors(service).isEmpty()) {
return "GenericError(ParseError(*response))";
}
return "Deserialize" + CppReservedWords.escape(op.getId().getName()) + "Error(*response)";
}
/**
* If the input has @idempotencyToken members, emits a prepared copy with unset tokens filled and
* returns the expression to use for the input from then on.
*/
static String prepareIdempotencyTokens(
CppWriter w, CppContext context, StructureShape input, String inputType) {
boolean any = false;
for (MemberShape member : input.members()) {
if (!member.hasTrait(IdempotencyTokenTrait.class)) {
continue;
}
if (!any) {
w.write("$L prepared = input;", inputType);
any = true;
}
String field = "prepared." + context.cppSymbols().toMemberName(member);
if (member.isRequired()) {
w.write("if ($L.empty()) $L = smithy::GenerateUuidV4();", field, field);
} else {
w.write("if (!$L.has_value()) $L = smithy::GenerateUuidV4();", field, field);
}
}
return any ? "prepared" : "input";
}
/** String conversion for label/query/header values of simple types. */
static String toStringExpression(
CppContext context, MemberShape member, String valueExpr, String timestampFormat) {
Shape target = context.model().expectShape(member.getTarget());
return switch (target.getType()) {
case STRING -> valueExpr;
case ENUM -> "std::string(" + valueExpr + ".ToString())";
case BYTE, SHORT, INTEGER, LONG, INT_ENUM ->
"std::to_string(static_cast<std::int64_t>(" + valueExpr + "))";
case FLOAT -> "smithy::FormatFloat(" + valueExpr + ")";
case DOUBLE -> "smithy::FormatDouble(" + valueExpr + ")";
case BOOLEAN -> "(" + valueExpr + " ? \"true\" : \"false\")";
case TIMESTAMP -> valueExpr + ".Format(" + timestampFormat + ")";
default ->
throw new CodegenException("cpp-codegen: unsupported binding target " + target.getId());
};
}
static StructureShape inputShape(CppContext context, OperationShape operation) {
return context.model().expectShape(operation.getInputShape()).asStructureShape().orElseThrow();
}
static StructureShape outputShape(CppContext context, OperationShape operation) {
return context.model().expectShape(operation.getOutputShape()).asStructureShape().orElseThrow();
}
static List<String> sharedClientIncludes(CppContext context) {
return List.of(
"\"" + context.settings().includePrefix() + "/client.h\"",
"\"" + context.settings().includePrefix() + "/serde.h\"",
"\"smithy/core/blob.h\"",
"\"smithy/core/document_serde.h\"",
"\"smithy/core/uuid.h\"",
"\"smithy/http/socket_transport.h\"",
"\"smithy/http/uri.h\"",
"<string>",
"<string_view>",
"<utility>");
}
}