Loading codegen-server-test/build.gradle.kts +2 −1 Original line number Diff line number Diff line Loading @@ -25,7 +25,8 @@ data class CodegenTest(val service: String, val module: String, val extraConfig: val CodegenTests = listOf( CodegenTest("com.amazonaws.simple#SimpleService", "simple"), CodegenTest("aws.protocoltests.restjson#RestJson", "rest_json"), CodegenTest("com.amazonaws.ebs#Ebs", "ebs") CodegenTest("com.amazonaws.ebs#Ebs", "ebs"), CodegenTest("com.amazonaws.s3#AmazonS3", "s3") ) /** Loading codegen-server/src/main/kotlin/software/amazon/smithy/rust/codegen/server/smithy/generators/protocol/ServerProtocolTestGenerator.kt +4 −2 Original line number Diff line number Diff line Loading @@ -150,10 +150,12 @@ class ServerProtocolTestGenerator( testModuleWriter.write("Test ID: ${testCase.id}") testModuleWriter.setNewlinePrefix("") testModuleWriter.writeWithNoFormatting("#[tokio::test]") // TODO: this allows to check-in RestJson protocol tests without // TODO: this allows to check-in RestJson and RestXml protocol tests without // failures as the protocol is not fully implemented yet. // Remove it once the protocol is fully implemented. if (operationShape.id.getNamespace() == "aws.protocoltests.restjson") { if (operationShape.id.getNamespace() == "aws.protocoltests.restjson" || operationShape.id.getNamespace() == "com.amazonaws.s3" ) { testModuleWriter.writeWithNoFormatting("#[ignore]") } val Tokio = CargoDependency( Loading codegen-server/src/main/kotlin/software/amazon/smithy/rust/codegen/server/smithy/protocols/ServerHttpProtocolGenerator.kt +20 −2 Original line number Diff line number Diff line Loading @@ -5,6 +5,8 @@ package software.amazon.smithy.rust.codegen.server.smithy.protocols import software.amazon.smithy.aws.traits.protocols.RestJson1Trait import software.amazon.smithy.aws.traits.protocols.RestXmlTrait import software.amazon.smithy.codegen.core.Symbol import software.amazon.smithy.model.knowledge.HttpBindingIndex import software.amazon.smithy.model.node.ExpectationNotMetException Loading Loading @@ -40,6 +42,7 @@ import software.amazon.smithy.rust.codegen.smithy.generators.protocol.MakeOperat import software.amazon.smithy.rust.codegen.smithy.generators.protocol.ProtocolGenerator import software.amazon.smithy.rust.codegen.smithy.generators.protocol.ProtocolTraitImplGenerator import software.amazon.smithy.rust.codegen.smithy.generators.setterName import software.amazon.smithy.rust.codegen.smithy.makeOptional import software.amazon.smithy.rust.codegen.smithy.protocols.HttpBindingDescriptor import software.amazon.smithy.rust.codegen.smithy.protocols.HttpBoundProtocolBodyGenerator import software.amazon.smithy.rust.codegen.smithy.protocols.HttpLocation Loading Loading @@ -514,12 +517,13 @@ private class ServerHttpProtocolImplGenerator( rust("let mut input = #T::default();", inputShape.builderSymbol(symbolProvider)) val parser = structuredDataParser.serverInputParser(operationShape) if (parser != null) { val contentTypeCheck = getContentTypeCheck() rustTemplate( """ let body = request.take_body().ok_or(#{SmithyHttpServer}::rejection::BodyAlreadyExtracted)?; let bytes = #{Hyper}::body::to_bytes(body).await?; if !bytes.is_empty() { #{SmithyHttpServer}::protocols::check_json_content_type(request)?; #{SmithyHttpServer}::protocols::$contentTypeCheck(request)?; input = #{parser}(bytes.as_ref(), input)?; } """, Loading Loading @@ -869,7 +873,7 @@ private class ServerHttpProtocolImplGenerator( // TODO These functions can be replaced with the ones in https://docs.rs/aws-smithy-types/latest/aws_smithy_types/primitive/trait.Parse.html private fun generateParseStrAsPrimitiveFn(binding: HttpBindingDescriptor): RuntimeType { val output = symbolProvider.toSymbol(binding.member) val output = symbolProvider.toSymbol(binding.member).makeOptional() val fnName = generateParseStrFnName(binding) return RuntimeType.forInlineFun(fnName, operationDeserModule) { writer -> writer.rustBlockTemplate( Loading @@ -893,4 +897,18 @@ private class ServerHttpProtocolImplGenerator( val memberName = binding.memberName.toSnakeCase() return "parse_str_${containerName}_$memberName" } private fun getContentTypeCheck(): String { when (codegenContext.protocol) { RestJson1Trait.ID -> { return "check_json_content_type" } RestXmlTrait.ID -> { return "check_xml_content_type" } else -> { TODO("Protocol ${codegenContext.protocol} not supported yet") } } } } codegen-server/src/main/kotlin/software/amazon/smithy/rust/codegen/server/smithy/protocols/ServerProtocolLoader.kt +2 −0 Original line number Diff line number Diff line Loading @@ -6,6 +6,7 @@ package software.amazon.smithy.rust.codegen.server.smithy.protocols import software.amazon.smithy.aws.traits.protocols.RestJson1Trait import software.amazon.smithy.aws.traits.protocols.RestXmlTrait import software.amazon.smithy.codegen.core.CodegenException import software.amazon.smithy.model.Model import software.amazon.smithy.model.knowledge.ServiceIndex Loading Loading @@ -39,6 +40,7 @@ class ServerProtocolLoader(private val supportedProtocols: ProtocolMap) { val DefaultProtocols = mapOf( // TODO: support other protocols. RestJson1Trait.ID to ServerRestJsonFactory(), RestXmlTrait.ID to ServerRestXmlFactory(), ) } } codegen-server/src/main/kotlin/software/amazon/smithy/rust/codegen/server/smithy/protocols/ServerRustXml.kt 0 → 100644 +106 −0 Original line number Diff line number Diff line /* * Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. * SPDX-License-Identifier: Apache-2.0. */ package software.amazon.smithy.rust.codegen.server.smithy.protocols import software.amazon.smithy.aws.traits.protocols.RestXmlTrait import software.amazon.smithy.model.Model import software.amazon.smithy.model.shapes.OperationShape import software.amazon.smithy.model.traits.TimestampFormatTrait import software.amazon.smithy.rust.codegen.rustlang.CargoDependency import software.amazon.smithy.rust.codegen.rustlang.RustModule import software.amazon.smithy.rust.codegen.rustlang.asType import software.amazon.smithy.rust.codegen.rustlang.rust import software.amazon.smithy.rust.codegen.rustlang.rustBlockTemplate import software.amazon.smithy.rust.codegen.smithy.CodegenContext import software.amazon.smithy.rust.codegen.smithy.RuntimeType import software.amazon.smithy.rust.codegen.smithy.generators.protocol.ProtocolSupport import software.amazon.smithy.rust.codegen.smithy.protocols.HttpBindingResolver import software.amazon.smithy.rust.codegen.smithy.protocols.HttpTraitHttpBindingResolver import software.amazon.smithy.rust.codegen.smithy.protocols.Protocol import software.amazon.smithy.rust.codegen.smithy.protocols.ProtocolContentTypes import software.amazon.smithy.rust.codegen.smithy.protocols.ProtocolGeneratorFactory import software.amazon.smithy.rust.codegen.smithy.protocols.parse.RestXmlParserGenerator import software.amazon.smithy.rust.codegen.smithy.protocols.parse.StructuredDataParserGenerator import software.amazon.smithy.rust.codegen.smithy.protocols.serialize.StructuredDataSerializerGenerator import software.amazon.smithy.rust.codegen.smithy.protocols.serialize.XmlBindingTraitSerializerGenerator import software.amazon.smithy.rust.codegen.util.expectTrait class ServerRestXmlFactory(private val generator: (CodegenContext) -> Protocol = { ServerRestXml(it) }) : ProtocolGeneratorFactory<ServerHttpProtocolGenerator> { override fun protocol(codegenContext: CodegenContext): Protocol = generator(codegenContext) override fun buildProtocolGenerator(codegenContext: CodegenContext): ServerHttpProtocolGenerator = ServerHttpProtocolGenerator(codegenContext, ServerRestXml(codegenContext)) override fun transformModel(model: Model): Model = model override fun support(): ProtocolSupport { return ProtocolSupport( /* Client support */ requestSerialization = false, requestBodySerialization = false, responseDeserialization = false, errorDeserialization = false, /* Server support */ requestDeserialization = true, requestBodyDeserialization = true, responseSerialization = true, errorSerialization = true ) } } open class ServerRestXml(private val codegenContext: CodegenContext) : Protocol { private val restXml = codegenContext.serviceShape.expectTrait<RestXmlTrait>() private val runtimeConfig = codegenContext.runtimeConfig private val errorScope = arrayOf( "Bytes" to RuntimeType.Bytes, "Error" to RuntimeType.GenericError(runtimeConfig), "HeaderMap" to RuntimeType.http.member("HeaderMap"), "Response" to RuntimeType.http.member("Response"), "XmlError" to CargoDependency.smithyXml(runtimeConfig).asType().member("decode::XmlError") ) private val xmlDeserModule = RustModule.private("xml_deser") protected val restXmlErrors: RuntimeType = when (restXml.isNoErrorWrapping) { true -> RuntimeType.unwrappedXmlErrors(runtimeConfig) false -> RuntimeType.wrappedXmlErrors(runtimeConfig) } override val httpBindingResolver: HttpBindingResolver = HttpTraitHttpBindingResolver(codegenContext.model, ProtocolContentTypes.consistent("application/xml")) override val defaultTimestampFormat: TimestampFormatTrait.Format = TimestampFormatTrait.Format.DATE_TIME override fun structuredDataParser(operationShape: OperationShape): StructuredDataParserGenerator { return RestXmlParserGenerator(codegenContext, restXmlErrors) } override fun structuredDataSerializer(operationShape: OperationShape): StructuredDataSerializerGenerator { return XmlBindingTraitSerializerGenerator(codegenContext, httpBindingResolver) } override fun parseHttpGenericError(operationShape: OperationShape): RuntimeType = RuntimeType.forInlineFun("parse_http_generic_error", xmlDeserModule) { writer -> writer.rustBlockTemplate( "pub fn parse_http_generic_error(response: &#{Response}<#{Bytes}>) -> Result<#{Error}, #{XmlError}>", *errorScope ) { rust("#T::parse_generic_error(response.body().as_ref())", restXmlErrors) } } override fun parseEventStreamGenericError(operationShape: OperationShape): RuntimeType = RuntimeType.forInlineFun("parse_event_stream_generic_error", xmlDeserModule) { writer -> writer.rustBlockTemplate( "pub fn parse_event_stream_generic_error(payload: &#{Bytes}) -> Result<#{Error}, #{XmlError}>", *errorScope ) { rust("#T::parse_generic_error(payload.as_ref())", restXmlErrors) } } } Loading
codegen-server-test/build.gradle.kts +2 −1 Original line number Diff line number Diff line Loading @@ -25,7 +25,8 @@ data class CodegenTest(val service: String, val module: String, val extraConfig: val CodegenTests = listOf( CodegenTest("com.amazonaws.simple#SimpleService", "simple"), CodegenTest("aws.protocoltests.restjson#RestJson", "rest_json"), CodegenTest("com.amazonaws.ebs#Ebs", "ebs") CodegenTest("com.amazonaws.ebs#Ebs", "ebs"), CodegenTest("com.amazonaws.s3#AmazonS3", "s3") ) /** Loading
codegen-server/src/main/kotlin/software/amazon/smithy/rust/codegen/server/smithy/generators/protocol/ServerProtocolTestGenerator.kt +4 −2 Original line number Diff line number Diff line Loading @@ -150,10 +150,12 @@ class ServerProtocolTestGenerator( testModuleWriter.write("Test ID: ${testCase.id}") testModuleWriter.setNewlinePrefix("") testModuleWriter.writeWithNoFormatting("#[tokio::test]") // TODO: this allows to check-in RestJson protocol tests without // TODO: this allows to check-in RestJson and RestXml protocol tests without // failures as the protocol is not fully implemented yet. // Remove it once the protocol is fully implemented. if (operationShape.id.getNamespace() == "aws.protocoltests.restjson") { if (operationShape.id.getNamespace() == "aws.protocoltests.restjson" || operationShape.id.getNamespace() == "com.amazonaws.s3" ) { testModuleWriter.writeWithNoFormatting("#[ignore]") } val Tokio = CargoDependency( Loading
codegen-server/src/main/kotlin/software/amazon/smithy/rust/codegen/server/smithy/protocols/ServerHttpProtocolGenerator.kt +20 −2 Original line number Diff line number Diff line Loading @@ -5,6 +5,8 @@ package software.amazon.smithy.rust.codegen.server.smithy.protocols import software.amazon.smithy.aws.traits.protocols.RestJson1Trait import software.amazon.smithy.aws.traits.protocols.RestXmlTrait import software.amazon.smithy.codegen.core.Symbol import software.amazon.smithy.model.knowledge.HttpBindingIndex import software.amazon.smithy.model.node.ExpectationNotMetException Loading Loading @@ -40,6 +42,7 @@ import software.amazon.smithy.rust.codegen.smithy.generators.protocol.MakeOperat import software.amazon.smithy.rust.codegen.smithy.generators.protocol.ProtocolGenerator import software.amazon.smithy.rust.codegen.smithy.generators.protocol.ProtocolTraitImplGenerator import software.amazon.smithy.rust.codegen.smithy.generators.setterName import software.amazon.smithy.rust.codegen.smithy.makeOptional import software.amazon.smithy.rust.codegen.smithy.protocols.HttpBindingDescriptor import software.amazon.smithy.rust.codegen.smithy.protocols.HttpBoundProtocolBodyGenerator import software.amazon.smithy.rust.codegen.smithy.protocols.HttpLocation Loading Loading @@ -514,12 +517,13 @@ private class ServerHttpProtocolImplGenerator( rust("let mut input = #T::default();", inputShape.builderSymbol(symbolProvider)) val parser = structuredDataParser.serverInputParser(operationShape) if (parser != null) { val contentTypeCheck = getContentTypeCheck() rustTemplate( """ let body = request.take_body().ok_or(#{SmithyHttpServer}::rejection::BodyAlreadyExtracted)?; let bytes = #{Hyper}::body::to_bytes(body).await?; if !bytes.is_empty() { #{SmithyHttpServer}::protocols::check_json_content_type(request)?; #{SmithyHttpServer}::protocols::$contentTypeCheck(request)?; input = #{parser}(bytes.as_ref(), input)?; } """, Loading Loading @@ -869,7 +873,7 @@ private class ServerHttpProtocolImplGenerator( // TODO These functions can be replaced with the ones in https://docs.rs/aws-smithy-types/latest/aws_smithy_types/primitive/trait.Parse.html private fun generateParseStrAsPrimitiveFn(binding: HttpBindingDescriptor): RuntimeType { val output = symbolProvider.toSymbol(binding.member) val output = symbolProvider.toSymbol(binding.member).makeOptional() val fnName = generateParseStrFnName(binding) return RuntimeType.forInlineFun(fnName, operationDeserModule) { writer -> writer.rustBlockTemplate( Loading @@ -893,4 +897,18 @@ private class ServerHttpProtocolImplGenerator( val memberName = binding.memberName.toSnakeCase() return "parse_str_${containerName}_$memberName" } private fun getContentTypeCheck(): String { when (codegenContext.protocol) { RestJson1Trait.ID -> { return "check_json_content_type" } RestXmlTrait.ID -> { return "check_xml_content_type" } else -> { TODO("Protocol ${codegenContext.protocol} not supported yet") } } } }
codegen-server/src/main/kotlin/software/amazon/smithy/rust/codegen/server/smithy/protocols/ServerProtocolLoader.kt +2 −0 Original line number Diff line number Diff line Loading @@ -6,6 +6,7 @@ package software.amazon.smithy.rust.codegen.server.smithy.protocols import software.amazon.smithy.aws.traits.protocols.RestJson1Trait import software.amazon.smithy.aws.traits.protocols.RestXmlTrait import software.amazon.smithy.codegen.core.CodegenException import software.amazon.smithy.model.Model import software.amazon.smithy.model.knowledge.ServiceIndex Loading Loading @@ -39,6 +40,7 @@ class ServerProtocolLoader(private val supportedProtocols: ProtocolMap) { val DefaultProtocols = mapOf( // TODO: support other protocols. RestJson1Trait.ID to ServerRestJsonFactory(), RestXmlTrait.ID to ServerRestXmlFactory(), ) } }
codegen-server/src/main/kotlin/software/amazon/smithy/rust/codegen/server/smithy/protocols/ServerRustXml.kt 0 → 100644 +106 −0 Original line number Diff line number Diff line /* * Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. * SPDX-License-Identifier: Apache-2.0. */ package software.amazon.smithy.rust.codegen.server.smithy.protocols import software.amazon.smithy.aws.traits.protocols.RestXmlTrait import software.amazon.smithy.model.Model import software.amazon.smithy.model.shapes.OperationShape import software.amazon.smithy.model.traits.TimestampFormatTrait import software.amazon.smithy.rust.codegen.rustlang.CargoDependency import software.amazon.smithy.rust.codegen.rustlang.RustModule import software.amazon.smithy.rust.codegen.rustlang.asType import software.amazon.smithy.rust.codegen.rustlang.rust import software.amazon.smithy.rust.codegen.rustlang.rustBlockTemplate import software.amazon.smithy.rust.codegen.smithy.CodegenContext import software.amazon.smithy.rust.codegen.smithy.RuntimeType import software.amazon.smithy.rust.codegen.smithy.generators.protocol.ProtocolSupport import software.amazon.smithy.rust.codegen.smithy.protocols.HttpBindingResolver import software.amazon.smithy.rust.codegen.smithy.protocols.HttpTraitHttpBindingResolver import software.amazon.smithy.rust.codegen.smithy.protocols.Protocol import software.amazon.smithy.rust.codegen.smithy.protocols.ProtocolContentTypes import software.amazon.smithy.rust.codegen.smithy.protocols.ProtocolGeneratorFactory import software.amazon.smithy.rust.codegen.smithy.protocols.parse.RestXmlParserGenerator import software.amazon.smithy.rust.codegen.smithy.protocols.parse.StructuredDataParserGenerator import software.amazon.smithy.rust.codegen.smithy.protocols.serialize.StructuredDataSerializerGenerator import software.amazon.smithy.rust.codegen.smithy.protocols.serialize.XmlBindingTraitSerializerGenerator import software.amazon.smithy.rust.codegen.util.expectTrait class ServerRestXmlFactory(private val generator: (CodegenContext) -> Protocol = { ServerRestXml(it) }) : ProtocolGeneratorFactory<ServerHttpProtocolGenerator> { override fun protocol(codegenContext: CodegenContext): Protocol = generator(codegenContext) override fun buildProtocolGenerator(codegenContext: CodegenContext): ServerHttpProtocolGenerator = ServerHttpProtocolGenerator(codegenContext, ServerRestXml(codegenContext)) override fun transformModel(model: Model): Model = model override fun support(): ProtocolSupport { return ProtocolSupport( /* Client support */ requestSerialization = false, requestBodySerialization = false, responseDeserialization = false, errorDeserialization = false, /* Server support */ requestDeserialization = true, requestBodyDeserialization = true, responseSerialization = true, errorSerialization = true ) } } open class ServerRestXml(private val codegenContext: CodegenContext) : Protocol { private val restXml = codegenContext.serviceShape.expectTrait<RestXmlTrait>() private val runtimeConfig = codegenContext.runtimeConfig private val errorScope = arrayOf( "Bytes" to RuntimeType.Bytes, "Error" to RuntimeType.GenericError(runtimeConfig), "HeaderMap" to RuntimeType.http.member("HeaderMap"), "Response" to RuntimeType.http.member("Response"), "XmlError" to CargoDependency.smithyXml(runtimeConfig).asType().member("decode::XmlError") ) private val xmlDeserModule = RustModule.private("xml_deser") protected val restXmlErrors: RuntimeType = when (restXml.isNoErrorWrapping) { true -> RuntimeType.unwrappedXmlErrors(runtimeConfig) false -> RuntimeType.wrappedXmlErrors(runtimeConfig) } override val httpBindingResolver: HttpBindingResolver = HttpTraitHttpBindingResolver(codegenContext.model, ProtocolContentTypes.consistent("application/xml")) override val defaultTimestampFormat: TimestampFormatTrait.Format = TimestampFormatTrait.Format.DATE_TIME override fun structuredDataParser(operationShape: OperationShape): StructuredDataParserGenerator { return RestXmlParserGenerator(codegenContext, restXmlErrors) } override fun structuredDataSerializer(operationShape: OperationShape): StructuredDataSerializerGenerator { return XmlBindingTraitSerializerGenerator(codegenContext, httpBindingResolver) } override fun parseHttpGenericError(operationShape: OperationShape): RuntimeType = RuntimeType.forInlineFun("parse_http_generic_error", xmlDeserModule) { writer -> writer.rustBlockTemplate( "pub fn parse_http_generic_error(response: &#{Response}<#{Bytes}>) -> Result<#{Error}, #{XmlError}>", *errorScope ) { rust("#T::parse_generic_error(response.body().as_ref())", restXmlErrors) } } override fun parseEventStreamGenericError(operationShape: OperationShape): RuntimeType = RuntimeType.forInlineFun("parse_event_stream_generic_error", xmlDeserModule) { writer -> writer.rustBlockTemplate( "pub fn parse_event_stream_generic_error(payload: &#{Bytes}) -> Result<#{Error}, #{XmlError}>", *errorScope ) { rust("#T::parse_generic_error(payload.as_ref())", restXmlErrors) } } }