diff --git a/.changes/next-release/bugfix-AWSSDKforJavav2-40c3de6.json b/.changes/next-release/bugfix-AWSSDKforJavav2-40c3de6.json new file mode 100644 index 000000000000..b0efb699a569 --- /dev/null +++ b/.changes/next-release/bugfix-AWSSDKforJavav2-40c3de6.json @@ -0,0 +1,6 @@ +{ + "type": "bugfix", + "category": "AWS SDK for Java v2", + "contributor": "", + "description": "Eliminate per-operation lambdas for endpoint and auth scheme resolution in generated clients, permanently fixing constant pool overflow for large services like EC2 Internal." +} diff --git a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/AsyncClientClass.java b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/AsyncClientClass.java index 045040d12f45..33553cd7313f 100644 --- a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/AsyncClientClass.java +++ b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/AsyncClientClass.java @@ -152,6 +152,9 @@ protected void addFields(Builder type) { .addField(protocolSpec.protocolFactory(model)) .addField(SdkClientConfiguration.class, "clientConfiguration", PRIVATE, FINAL); + type.addField(ClientClassUtils.authSchemeOptionsResolverField()); + type.addField(ClientClassUtils.endpointResolverField()); + // Kinesis doesn't support CBOR for STS yet so need another protocol factory for JSON if (model.getMetadata().isCborProtocol()) { type.addField(AwsJsonProtocolFactory.class, "jsonProtocolFactory", PRIVATE, FINAL); @@ -173,9 +176,7 @@ protected void addAdditionalMethods(TypeSpec.Builder type) { .addMethod(protocolSpec.initProtocolFactory(model)) .addMethod(resolveMetricPublishersMethod()) .addMethod(ClientClassUtils.resolveAuthSchemeOptionsMethod(authSchemeSpecUtils, endpointRulesSpecUtils)) - .addMethod(ClientClassUtils.resolveEndpointMethod(authSchemeSpecUtils, endpointRulesSpecUtils)) - .addMethod(ClientClassUtils.authSchemeResolverFactoryMethod()) - .addMethod(ClientClassUtils.endpointResolverFactoryMethod()); + .addMethod(ClientClassUtils.resolveEndpointMethod(authSchemeSpecUtils, endpointRulesSpecUtils)); type.addMethod(ClientClassUtils.updateRetryStrategyClientConfigurationMethod()); type.addMethod(updateSdkClientConfigurationMethod(configurationUtils.serviceClientConfigurationBuilderClassName(), diff --git a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/ClientClassUtils.java b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/ClientClassUtils.java index 5b0f5ab6ae98..e5147d9e1ce8 100644 --- a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/ClientClassUtils.java +++ b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/ClientClassUtils.java @@ -20,6 +20,7 @@ import com.squareup.javapoet.ClassName; import com.squareup.javapoet.CodeBlock; +import com.squareup.javapoet.FieldSpec; import com.squareup.javapoet.MethodSpec; import com.squareup.javapoet.ParameterSpec; import com.squareup.javapoet.ParameterizedTypeName; @@ -59,13 +60,12 @@ import software.amazon.awssdk.core.client.config.ClientOverrideConfiguration; import software.amazon.awssdk.core.client.config.SdkClientConfiguration; import software.amazon.awssdk.core.client.config.SdkClientOption; -import software.amazon.awssdk.core.endpoint.EndpointResolver; import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.retry.RetryMode; import software.amazon.awssdk.core.signer.Signer; -import software.amazon.awssdk.core.spi.identity.AuthSchemeOptionsResolver; import software.amazon.awssdk.endpoints.Endpoint; import software.amazon.awssdk.http.auth.spi.scheme.AuthSchemeOption; import software.amazon.awssdk.retries.api.RetryStrategy; @@ -368,11 +368,13 @@ static MethodSpec resolveAuthSchemeOptionsMethod(AuthSchemeSpecUtils authSchemeS .addModifiers(PRIVATE) .returns(ParameterizedTypeName.get(ClassName.get(List.class), ClassName.get(AuthSchemeOption.class))) .addParameter(SdkRequest.class, "request") - .addParameter(String.class, "operationName") - .addParameter(SdkClientConfiguration.class, "clientConfiguration"); + .addParameter(ExecutionAttributes.class, "executionAttributes"); ClassName providerInterface = authSchemeSpecUtils.providerInterfaceName(); + builder.addStatement("String operationName = executionAttributes.getAttribute($T.OPERATION_NAME)", + SdkExecutionAttribute.class); + // Check for request-level authSchemeProvider override builder.addStatement("$T requestAuthSchemeProvider = request.overrideConfiguration()" + ".flatMap(c -> c.authSchemeProvider())" @@ -383,8 +385,9 @@ static MethodSpec resolveAuthSchemeOptionsMethod(AuthSchemeSpecUtils authSchemeS builder.addStatement("$T authSchemeProvider = requestAuthSchemeProvider != null " + "? requestAuthSchemeProvider " + ": $T.isInstanceOf($T.class, " - + "clientConfiguration.option($T.AUTH_SCHEME_PROVIDER), $S)", - providerInterface, Validate.class, providerInterface, SdkClientOption.class, + + "executionAttributes.getAttribute($T.AUTH_SCHEME_RESOLVER), $S)", + providerInterface, Validate.class, providerInterface, + SdkInternalExecutionAttribute.class, "Expected an instance of " + authSchemeSpecUtils.providerInterfaceName().simpleName()); if (authSchemeSpecUtils.useEndpointBasedAuthProvider()) { @@ -395,8 +398,8 @@ static MethodSpec resolveAuthSchemeOptionsMethod(AuthSchemeSpecUtils authSchemeS if (endpointRulesSpecUtils.isS3()) { ClassName sdkIdentityProperty = ClassName.get("software.amazon.awssdk.core.identity", "SdkIdentityProperty"); - builder.addStatement("$T sdkClient = clientConfiguration.option($T.SDK_CLIENT)", - SdkClient.class, SdkClientOption.class); + builder.addStatement("$T sdkClient = executionAttributes.getAttribute($T.SDK_CLIENT)", + SdkClient.class, SdkInternalExecutionAttribute.class); builder.addStatement("return options.stream().map(o -> o.toBuilder()" + ".putIdentityProperty($T.SDK_CLIENT, sdkClient).build())" + ".collect($T.toList())", @@ -411,19 +414,21 @@ static MethodSpec resolveAuthSchemeOptionsMethod(AuthSchemeSpecUtils authSchemeS private static void addSimpleAuthSchemeResolution(MethodSpec.Builder builder, AuthSchemeSpecUtils authSchemeSpecUtils) { ClassName paramsInterface = authSchemeSpecUtils.parametersInterfaceName(); - ClassName awsClientOption = ClassName.get("software.amazon.awssdk.awscore.client.config", "AwsClientOption"); + ClassName awsExecutionAttribute = ClassName.get("software.amazon.awssdk.awscore", "AwsExecutionAttribute"); builder.addStatement("$T.Builder paramsBuilder = $T.builder().operation(operationName)", paramsInterface, paramsInterface); if (authSchemeSpecUtils.usesSigV4()) { - builder.addStatement("paramsBuilder.region(clientConfiguration.option($T.AWS_REGION))", awsClientOption); + builder.addStatement("paramsBuilder.region(executionAttributes.getAttribute($T.AWS_REGION))", + awsExecutionAttribute); } if (authSchemeSpecUtils.hasSigV4aSupport()) { ClassName regionSet = ClassName.get("software.amazon.awssdk.http.auth.aws.signer", "RegionSet"); - builder.addStatement("$T sigv4aRegionSet = clientConfiguration.option($T.AWS_SIGV4A_SIGNING_REGION_SET)", - ClassName.get(Set.class), awsClientOption); + builder.addStatement("$T sigv4aRegionSet = executionAttributes" + + ".getAttribute($T.AWS_SIGV4A_SIGNING_REGION_SET)", + ClassName.get(Set.class), awsExecutionAttribute); builder.beginControlFlow("if (!$T.isNullOrEmpty(sigv4aRegionSet))", CollectionUtils.class); builder.addStatement("paramsBuilder.regionSet($T.create(sigv4aRegionSet))", regionSet); builder.endControlFlow(); @@ -437,32 +442,12 @@ private static void addEndpointBasedAuthSchemeResolution(MethodSpec.Builder buil AuthSchemeSpecUtils authSchemeSpecUtils, EndpointRulesSpecUtils endpointRulesSpecUtils) { ClassName paramsInterface = authSchemeSpecUtils.parametersInterfaceName(); - ClassName awsClientOption = ClassName.get("software.amazon.awssdk.awscore.client.config", "AwsClientOption"); + ClassName awsExecutionAttribute = ClassName.get("software.amazon.awssdk.awscore", "AwsExecutionAttribute"); ClassName endpointParamsClass = endpointRulesSpecUtils.parametersClassName(); ClassName endpointResolverUtils = endpointRulesSpecUtils.endpointResolverUtilsName(); - ClassName executionAttributesClass = ClassName.get("software.amazon.awssdk.core.interceptor", "ExecutionAttributes"); - ClassName awsExecutionAttribute = ClassName.get("software.amazon.awssdk.awscore", "AwsExecutionAttribute"); - ClassName sdkExecutionAttribute = ClassName.get("software.amazon.awssdk.core.interceptor", "SdkExecutionAttribute"); ClassName sdkInternalExecutionAttribute = ClassName.get("software.amazon.awssdk.core.interceptor", "SdkInternalExecutionAttribute"); - builder.addStatement("$T executionAttributes = new $T()", executionAttributesClass, executionAttributesClass); - builder.addStatement("executionAttributes.putAttribute($T.AWS_REGION, clientConfiguration.option($T.AWS_REGION))", - awsExecutionAttribute, awsClientOption); - builder.addStatement("executionAttributes.putAttribute($T.DUALSTACK_ENDPOINT_ENABLED, " - + "clientConfiguration.option($T.DUALSTACK_ENDPOINT_ENABLED))", - awsExecutionAttribute, awsClientOption); - builder.addStatement("executionAttributes.putAttribute($T.FIPS_ENDPOINT_ENABLED, " - + "clientConfiguration.option($T.FIPS_ENDPOINT_ENABLED))", - awsExecutionAttribute, awsClientOption); - builder.addStatement("executionAttributes.putAttribute($T.OPERATION_NAME, operationName)", sdkExecutionAttribute); - builder.addStatement("executionAttributes.putAttribute($T.CLIENT_ENDPOINT_PROVIDER, " - + "clientConfiguration.option($T.CLIENT_ENDPOINT_PROVIDER))", - sdkInternalExecutionAttribute, SdkClientOption.class); - builder.addStatement("executionAttributes.putAttribute($T.CLIENT_CONTEXT_PARAMS, " - + "clientConfiguration.option($T.CLIENT_CONTEXT_PARAMS))", - sdkInternalExecutionAttribute, SdkClientOption.class); - builder.addStatement("$T endpointParams = $T.ruleParams(request, executionAttributes)", endpointParamsClass, endpointResolverUtils); @@ -481,13 +466,15 @@ private static void addEndpointBasedAuthSchemeResolution(MethodSpec.Builder buil builder.addStatement("paramsBuilder.operation(operationName)"); if (authSchemeSpecUtils.usesSigV4() && !regionIncluded) { - builder.addStatement("paramsBuilder.region(clientConfiguration.option($T.AWS_REGION))", awsClientOption); + builder.addStatement("paramsBuilder.region(executionAttributes.getAttribute($T.AWS_REGION))", + awsExecutionAttribute); } if (authSchemeSpecUtils.hasSigV4aSupport()) { ClassName regionSet = ClassName.get("software.amazon.awssdk.http.auth.aws.signer", "RegionSet"); - builder.addStatement("$T sigv4aRegionSet = clientConfiguration.option($T.AWS_SIGV4A_SIGNING_REGION_SET)", - ClassName.get(Set.class), awsClientOption); + builder.addStatement("$T sigv4aRegionSet = executionAttributes" + + ".getAttribute($T.AWS_SIGV4A_SIGNING_REGION_SET)", + ClassName.get(Set.class), awsExecutionAttribute); builder.beginControlFlow("if (!$T.isNullOrEmpty(sigv4aRegionSet))", CollectionUtils.class); builder.addStatement("paramsBuilder.regionSet($T.create(sigv4aRegionSet))", regionSet); builder.endControlFlow(); @@ -497,8 +484,10 @@ private static void addEndpointBasedAuthSchemeResolution(MethodSpec.Builder buil ClassName endpointProviderInterface = endpointRulesSpecUtils.providerInterfaceName(); builder.beginControlFlow("if (paramsBuilder instanceof $T)", paramsBuilderClass); - builder.addStatement("$T endpointProvider = clientConfiguration.option($T.ENDPOINT_PROVIDER)", - ClassName.get("software.amazon.awssdk.endpoints", "EndpointProvider"), SdkClientOption.class); + builder.addStatement("$T endpointProvider = ($T) executionAttributes.getAttribute($T.ENDPOINT_PROVIDER)", + ClassName.get("software.amazon.awssdk.endpoints", "EndpointProvider"), + ClassName.get("software.amazon.awssdk.endpoints", "EndpointProvider"), + sdkInternalExecutionAttribute); builder.beginControlFlow("if (endpointProvider instanceof $T)", endpointProviderInterface); builder.addStatement("(($T) paramsBuilder).endpointProvider(($T) endpointProvider)", paramsBuilderClass, endpointProviderInterface); @@ -509,39 +498,6 @@ private static void addEndpointBasedAuthSchemeResolution(MethodSpec.Builder buil List.class, AuthSchemeOption.class); } - /** - * Generates a factory method that creates an {@code AuthSchemeOptionsResolver} for a given operation name. - * This avoids creating a new lambda per operation in the generated client class, which reduces constant pool - * pressure for services with a large number of operations (e.g., EC2). - */ - static MethodSpec authSchemeResolverFactoryMethod() { - ClassName authSchemeOptionsResolver = ClassName.get(AuthSchemeOptionsResolver.class); - - return MethodSpec.methodBuilder("authSchemeResolver") - .addModifiers(PRIVATE) - .returns(authSchemeOptionsResolver) - .addParameter(String.class, "operationName") - .addParameter(SdkClientConfiguration.class, "clientConfiguration") - .addStatement("return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration)") - .build(); - } - - /** - * Generates a factory method that creates an {@code EndpointResolver} for a given operation name. - * This avoids creating a new lambda per operation in the generated client class, which reduces constant pool - * pressure for services with a large number of operations (e.g., EC2). - */ - static MethodSpec endpointResolverFactoryMethod() { - ClassName endpointResolver = ClassName.get(EndpointResolver.class); - - return MethodSpec.methodBuilder("endpointResolver") - .addModifiers(PRIVATE) - .returns(endpointResolver) - .addParameter(String.class, "operationName") - .addStatement("return (r, a) -> resolveEndpoint(r, a, operationName)") - .build(); - } - static MethodSpec resolveEndpointMethod(AuthSchemeSpecUtils authSchemeSpecUtils, EndpointRulesSpecUtils endpointRulesSpecUtils) { ClassName utilsClass = endpointRulesSpecUtils.endpointResolverUtilsName(); @@ -555,8 +511,9 @@ static MethodSpec resolveEndpointMethod(AuthSchemeSpecUtils authSchemeSpecUtils, .addModifiers(PRIVATE) .returns(Endpoint.class) .addParameter(SdkRequest.class, "request") - .addParameter(ExecutionAttributes.class, "executionAttributes") - .addParameter(String.class, "operationName"); + .addParameter(ExecutionAttributes.class, "executionAttributes"); + + b.addStatement("String operationName = executionAttributes.getAttribute($T.OPERATION_NAME)", SdkExecutionAttribute.class); b.addStatement("$1T provider = ($1T) executionAttributes.getAttribute($2T.ENDPOINT_PROVIDER)", providerInterface, SdkInternalExecutionAttribute.class); @@ -623,4 +580,49 @@ static MethodSpec resolveEndpointMethod(AuthSchemeSpecUtils authSchemeSpecUtils, return b.build(); } + + static FieldSpec authSchemeOptionsResolverField() { + ClassName resolverType = ClassName.get("software.amazon.awssdk.core.spi.identity", "AuthSchemeOptionsResolver"); + ClassName sdkRequest = ClassName.get(SdkRequest.class); + ClassName executionAttributes = ClassName.get(ExecutionAttributes.class); + ClassName authSchemeOption = ClassName.get(AuthSchemeOption.class); + + return FieldSpec.builder(resolverType, "authSchemeOptionsResolver", PRIVATE, Modifier.FINAL) + .initializer("new $T() {\n" + + " @Override\n" + + " public $T<$T> resolve($T request) {\n" + + " throw new $T(\"Use resolve(SdkRequest, ExecutionAttributes) instead\");\n" + + " }\n" + + "\n" + + " @Override\n" + + " public $T<$T> resolve($T request, $T executionAttributes) {\n" + + " return resolveAuthSchemeOptions(request, executionAttributes);\n" + + " }\n" + + "}", + resolverType, + ClassName.get(List.class), authSchemeOption, sdkRequest, + ClassName.get(UnsupportedOperationException.class), + ClassName.get(List.class), authSchemeOption, + sdkRequest, executionAttributes) + .build(); + } + + static FieldSpec endpointResolverField() { + ClassName resolverType = ClassName.get("software.amazon.awssdk.core.endpoint", "EndpointResolver"); + ClassName sdkRequest = ClassName.get(SdkRequest.class); + ClassName executionAttributes = ClassName.get(ExecutionAttributes.class); + ClassName endpoint = ClassName.get(Endpoint.class); + + return FieldSpec.builder(resolverType, "endpointResolverInstance", PRIVATE, Modifier.FINAL) + .initializer("new $T() {\n" + + " @Override\n" + + " public $T resolve($T request, $T executionAttributes) {\n" + + " return resolveEndpoint(request, executionAttributes);\n" + + " }\n" + + "}", + resolverType, + endpoint, + sdkRequest, executionAttributes) + .build(); + } } diff --git a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/SyncClientClass.java b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/SyncClientClass.java index 99e37ea13aea..d23e8ee97357 100644 --- a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/SyncClientClass.java +++ b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/SyncClientClass.java @@ -129,6 +129,9 @@ protected void addFields(TypeSpec.Builder type) { .addField(SyncClientHandler.class, "clientHandler", PRIVATE, FINAL) .addField(protocolSpec.protocolFactory(model)) .addField(SdkClientConfiguration.class, "clientConfiguration", PRIVATE, FINAL); + + type.addField(ClientClassUtils.authSchemeOptionsResolverField()); + type.addField(ClientClassUtils.endpointResolverField()); } @Override @@ -141,9 +144,7 @@ protected void addAdditionalMethods(TypeSpec.Builder type) { .addMethods(protocolSpec.additionalMethods()) .addMethod(resolveMetricPublishersMethod()) .addMethod(ClientClassUtils.resolveAuthSchemeOptionsMethod(authSchemeSpecUtils, endpointRulesSpecUtils)) - .addMethod(ClientClassUtils.resolveEndpointMethod(authSchemeSpecUtils, endpointRulesSpecUtils)) - .addMethod(ClientClassUtils.authSchemeResolverFactoryMethod()) - .addMethod(ClientClassUtils.endpointResolverFactoryMethod()); + .addMethod(ClientClassUtils.resolveEndpointMethod(authSchemeSpecUtils, endpointRulesSpecUtils)); protocolSpec.createErrorResponseHandler().ifPresent(type::addMethod); type.addMethod(ClientClassUtils.updateRetryStrategyClientConfigurationMethod()); diff --git a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/JsonProtocolSpec.java b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/JsonProtocolSpec.java index d5937678ac21..39d7dc7f2efe 100644 --- a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/JsonProtocolSpec.java +++ b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/JsonProtocolSpec.java @@ -224,10 +224,8 @@ public CodeBlock executionHandler(OperationModel opModel) { .add(".withRequestConfiguration(clientConfiguration)") .add(".withInput($L)\n", opModel.getInput().getVariableName()) .add(".withMetricCollector(apiCallMetricCollector)\n") - .add(".withAuthSchemeOptionsResolver(authSchemeResolver($S, clientConfiguration))\n", - opModel.getOperationName()) - .add(".withEndpointResolver(endpointResolver($S))\n", - opModel.getOperationName()) + .add(".withAuthSchemeOptionsResolver(authSchemeOptionsResolver)\n") + .add(".withEndpointResolver(endpointResolverInstance)\n") .add(HttpChecksumRequiredTrait.putHttpChecksumAttribute(opModel)) .add(HttpChecksumTrait.create(opModel)); @@ -302,10 +300,8 @@ public CodeBlock asyncExecutionHandler(IntermediateModel intermediateModel, Oper .add(".withErrorResponseHandler(errorResponseHandler)\n") .add(".withRequestConfiguration(clientConfiguration)") .add(".withMetricCollector(apiCallMetricCollector)\n") - .add(".withAuthSchemeOptionsResolver(authSchemeResolver($S, clientConfiguration))\n", - opModel.getOperationName()) - .add(".withEndpointResolver(endpointResolver($S))\n", - opModel.getOperationName()) + .add(".withAuthSchemeOptionsResolver(authSchemeOptionsResolver)\n") + .add(".withEndpointResolver(endpointResolverInstance)\n") .add(hostPrefixExpression(opModel)) .add(discoveredEndpoint(opModel)) .add(credentialType(opModel, model)) diff --git a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/QueryProtocolSpec.java b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/QueryProtocolSpec.java index 302b99c9cec4..a78cda7e98c5 100644 --- a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/QueryProtocolSpec.java +++ b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/QueryProtocolSpec.java @@ -116,10 +116,8 @@ public CodeBlock executionHandler(OperationModel opModel) { .add(".withRequestConfiguration(clientConfiguration)") .add(".withInput($L)", opModel.getInput().getVariableName()) .add(".withMetricCollector(apiCallMetricCollector)") - .add(".withAuthSchemeOptionsResolver(authSchemeResolver($S, clientConfiguration))\n", - opModel.getOperationName()) - .add(".withEndpointResolver(endpointResolver($S))\n", - opModel.getOperationName()) + .add(".withAuthSchemeOptionsResolver(authSchemeOptionsResolver)\n") + .add(".withEndpointResolver(endpointResolverInstance)\n") .add(HttpChecksumRequiredTrait.putHttpChecksumAttribute(opModel)) .add(HttpChecksumTrait.create(opModel)); @@ -159,10 +157,8 @@ public CodeBlock asyncExecutionHandler(IntermediateModel intermediateModel, Oper .add(credentialType(opModel, intermediateModel)) .add(".withRequestConfiguration(clientConfiguration)") .add(".withMetricCollector(apiCallMetricCollector)\n") - .add(".withAuthSchemeOptionsResolver(authSchemeResolver($S, clientConfiguration))\n", - opModel.getOperationName()) - .add(".withEndpointResolver(endpointResolver($S))\n", - opModel.getOperationName()) + .add(".withAuthSchemeOptionsResolver(authSchemeOptionsResolver)\n") + .add(".withEndpointResolver(endpointResolverInstance)\n") .add(HttpChecksumRequiredTrait.putHttpChecksumAttribute(opModel)) .add(HttpChecksumTrait.create(opModel)); diff --git a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/XmlProtocolSpec.java b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/XmlProtocolSpec.java index fb421b0ccb77..d7bed6355d22 100644 --- a/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/XmlProtocolSpec.java +++ b/codegen/src/main/java/software/amazon/awssdk/codegen/poet/client/specs/XmlProtocolSpec.java @@ -135,10 +135,8 @@ public CodeBlock executionHandler(OperationModel opModel) { .add(credentialType(opModel, model)) .add(".withRequestConfiguration(clientConfiguration)") .add(".withInput($L)", opModel.getInput().getVariableName()) - .add(".withAuthSchemeOptionsResolver(authSchemeResolver($S, clientConfiguration))\n", - opModel.getOperationName()) - .add(".withEndpointResolver(endpointResolver($S))\n", - opModel.getOperationName()) + .add(".withAuthSchemeOptionsResolver(authSchemeOptionsResolver)\n") + .add(".withEndpointResolver(endpointResolverInstance)\n") .add(HttpChecksumRequiredTrait.putHttpChecksumAttribute(opModel)) .add(HttpChecksumTrait.create(opModel)); @@ -217,10 +215,8 @@ public CodeBlock asyncExecutionHandler(IntermediateModel intermediateModel, Oper builder.add(hostPrefixExpression(opModel)) .add(credentialType(opModel, model)) .add(".withMetricCollector(apiCallMetricCollector)\n") - .add(".withAuthSchemeOptionsResolver(authSchemeResolver($S, clientConfiguration))\n", - opModel.getOperationName()) - .add(".withEndpointResolver(endpointResolver($S))\n", - opModel.getOperationName()) + .add(".withAuthSchemeOptionsResolver(authSchemeOptionsResolver)\n") + .add(".withEndpointResolver(endpointResolverInstance)\n") .add(asyncRequestBody(opModel)) .add(HttpChecksumRequiredTrait.putHttpChecksumAttribute(opModel)) .add(HttpChecksumTrait.create(opModel)); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-json-async-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-json-async-client-class.java index 33f5c7b612e0..f05a7bf80cb9 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-json-async-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-json-async-client-class.java @@ -16,7 +16,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.client.handler.AwsClientHandlerUtils; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; @@ -50,6 +50,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.interceptor.trait.HttpChecksumRequired; import software.amazon.awssdk.core.internal.interceptor.trait.RequestCompression; @@ -156,6 +157,25 @@ final class DefaultJsonAsyncClient implements JsonAsyncClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + private final Executor executor; protected DefaultJsonAsyncClient(SdkClientConfiguration clientConfiguration) { @@ -238,8 +258,8 @@ public CompletableFuture aPostOperation(APostOperationRe .withMarshaller(new APostOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .hostPrefixExpression(resolvedHostExpression).withInput(aPostOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -318,8 +338,8 @@ public CompletableFuture aPostOperationWithOut .withMarshaller(new APostOperationWithOutputRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperationWithOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperationWithOutput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(aPostOperationWithOutputRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -428,8 +448,8 @@ public CompletableFuture eventStreamOperation(EventStreamOperationRequest .withInitialRequestEvent(true).withResponseHandler(voidResponseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("EventStreamOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("EventStreamOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(eventStreamOperationRequest), asyncResponseTransformer); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { if (e != null) { @@ -525,9 +545,8 @@ public CompletableFuture eventStreamO .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("EventStreamOperationWithOnlyInput", clientConfiguration)) - .withEndpointResolver(endpointResolver("EventStreamOperationWithOnlyInput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(eventStreamOperationWithOnlyInputRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -634,9 +653,8 @@ public CompletableFuture eventStreamOperationWithOnlyOutput( .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("EventStreamOperationWithOnlyOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("EventStreamOperationWithOnlyOutput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(eventStreamOperationWithOnlyOutputRequest), asyncResponseTransformer); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { if (e != null) { @@ -723,8 +741,8 @@ public CompletableFuture getWithoutRequiredMe .withMarshaller(new GetWithoutRequiredMembersRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetWithoutRequiredMembers", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetWithoutRequiredMembers")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(getWithoutRequiredMembersRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -802,9 +820,8 @@ public CompletableFuture operationWithChe .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithChecksumRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithChecksumRequired")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.HTTP_CHECKSUM_REQUIRED, HttpChecksumRequired.create()).withInput(operationWithChecksumRequiredRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { @@ -879,8 +896,8 @@ public CompletableFuture operationWithNoneAut .withMarshaller(new OperationWithNoneAuthTypeRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithNoneAuthType", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithNoneAuthType")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithNoneAuthTypeRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -958,9 +975,8 @@ public CompletableFuture operationWithR .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithRequestCompression", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithRequestCompression")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.REQUEST_COMPRESSION, RequestCompression.builder().encodings("gzip").isStreaming(false).build()) .withInput(operationWithRequestCompressionRequest)); @@ -1040,9 +1056,8 @@ public CompletableFuture paginatedOpera .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("PaginatedOperationWithResultKey", clientConfiguration)) - .withEndpointResolver(endpointResolver("PaginatedOperationWithResultKey")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(paginatedOperationWithResultKeyRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1120,9 +1135,8 @@ public CompletableFuture paginatedOp .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("PaginatedOperationWithoutResultKey", clientConfiguration)) - .withEndpointResolver(endpointResolver("PaginatedOperationWithoutResultKey")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(paginatedOperationWithoutResultKeyRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1205,8 +1219,8 @@ public CompletableFuture streamingInputOperatio .asyncRequestBody(requestBody).build()).withResponseHandler(responseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingInputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOperation")).withAsyncRequestBody(requestBody) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withAsyncRequestBody(requestBody) .withInput(streamingInputOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1303,9 +1317,8 @@ public CompletableFuture streamingInputOutputOperation( .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("StreamingInputOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOutputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withAsyncRequestBody(requestBody).withAsyncResponseTransformer(asyncResponseTransformer) .withInput(streamingInputOutputOperationRequest), asyncResponseTransformer); AsyncResponseTransformer finalAsyncResponseTransformer = asyncResponseTransformer; @@ -1400,8 +1413,8 @@ public CompletableFuture streamingOutputOperation( .withMarshaller(new StreamingOutputOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingOutputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withAsyncResponseTransformer(asyncResponseTransformer).withInput(streamingOutputOperationRequest), asyncResponseTransformer); AsyncResponseTransformer finalAsyncResponseTransformer = asyncResponseTransformer; @@ -1455,8 +1468,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -1464,15 +1478,16 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri .isInstanceOf(JsonAuthSchemeProvider.class, p, "Expected an instance of JsonAuthSchemeProvider")) .orElse(null); JsonAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(JsonAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(JsonAuthSchemeProvider.class, executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of JsonAuthSchemeProvider"); JsonAuthSchemeParams.Builder paramsBuilder = JsonAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonEndpointProvider provider = (JsonEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -1503,14 +1518,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-query-compatible-json-async-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-query-compatible-json-async-client-class.java index de37281ad310..387bd3bddef2 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-query-compatible-json-async-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-query-compatible-json-async-client-class.java @@ -13,7 +13,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -35,6 +35,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -85,6 +86,25 @@ final class DefaultQueryToJsonCompatibleAsyncClient implements QueryToJsonCompat private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultQueryToJsonCompatibleAsyncClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsAsyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -156,8 +176,8 @@ public CompletableFuture aPostOperation(APostOperationRe .withMarshaller(new APostOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .hostPrefixExpression(resolvedHostExpression).withInput(aPostOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -201,8 +221,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); QueryToJsonCompatibleAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -210,16 +231,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of QueryToJsonCompatibleAuthSchemeProvider")).orElse(null); QueryToJsonCompatibleAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(QueryToJsonCompatibleAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of QueryToJsonCompatibleAuthSchemeProvider"); QueryToJsonCompatibleAuthSchemeParams.Builder paramsBuilder = QueryToJsonCompatibleAuthSchemeParams.builder().operation( operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); QueryToJsonCompatibleEndpointProvider provider = (QueryToJsonCompatibleEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -251,14 +273,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-query-compatible-json-sync-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-query-compatible-json-sync-client-class.java index ead8f03c7bfd..ff6bdd51e1e3 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-query-compatible-json-sync-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-aws-query-compatible-json-sync-client-class.java @@ -8,7 +8,7 @@ import java.util.function.Function; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -30,6 +30,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -80,6 +81,25 @@ final class DefaultQueryToJsonCompatibleClient implements QueryToJsonCompatibleC private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultQueryToJsonCompatibleClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsSyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -147,8 +167,8 @@ public APostOperationResponse aPostOperation(APostOperationRequest aPostOperatio .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .hostPrefixExpression(resolvedHostExpression).withRequestConfiguration(clientConfiguration) .withInput(aPostOperationRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new APostOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -175,8 +195,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); QueryToJsonCompatibleAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -184,16 +205,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of QueryToJsonCompatibleAuthSchemeProvider")).orElse(null); QueryToJsonCompatibleAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(QueryToJsonCompatibleAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of QueryToJsonCompatibleAuthSchemeProvider"); QueryToJsonCompatibleAuthSchemeParams.Builder paramsBuilder = QueryToJsonCompatibleAuthSchemeParams.builder().operation( operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); QueryToJsonCompatibleEndpointProvider provider = (QueryToJsonCompatibleEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -225,14 +247,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private HttpResponseHandler createErrorResponseHandler(BaseAwsJsonProtocolFactory protocolFactory, JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { return protocolFactory.createErrorResponseHandler(operationMetadata, exceptionMetadataMapper); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-batchmanager-async.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-batchmanager-async.java index a047a35d87d5..d5972a6ddd60 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-batchmanager-async.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-batchmanager-async.java @@ -14,7 +14,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -36,6 +36,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -85,6 +86,25 @@ final class DefaultBatchManagerTestAsyncClient implements BatchManagerTestAsyncC private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + private final ScheduledExecutorService executorService; protected DefaultBatchManagerTestAsyncClient(SdkClientConfiguration clientConfiguration) { @@ -148,8 +168,8 @@ public CompletableFuture sendRequest(SendRequestRequest sen .withMarshaller(new SendRequestRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("SendRequest", clientConfiguration)) - .withEndpointResolver(endpointResolver("SendRequest")).withInput(sendRequestRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(sendRequestRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -197,8 +217,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); BatchManagerTestAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -206,16 +227,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of BatchManagerTestAuthSchemeProvider")).orElse(null); BatchManagerTestAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(BatchManagerTestAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of BatchManagerTestAuthSchemeProvider"); BatchManagerTestAuthSchemeParams.Builder paramsBuilder = BatchManagerTestAuthSchemeParams.builder().operation( operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); BatchManagerTestEndpointProvider provider = (BatchManagerTestEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -247,14 +269,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-cbor-async-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-cbor-async-client-class.java index 56c3e6c6c46a..129116f85b0a 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-cbor-async-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-cbor-async-client-class.java @@ -16,7 +16,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.client.handler.AwsClientHandlerUtils; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; @@ -50,6 +50,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.interceptor.trait.HttpChecksumRequired; import software.amazon.awssdk.core.internal.interceptor.trait.RequestCompression; @@ -157,6 +158,25 @@ final class DefaultJsonAsyncClient implements JsonAsyncClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + private final AwsJsonProtocolFactory jsonProtocolFactory; private final Executor executor; @@ -242,8 +262,8 @@ public CompletableFuture aPostOperation(APostOperationRe .withMarshaller(new APostOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .hostPrefixExpression(resolvedHostExpression).withInput(aPostOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -322,8 +342,8 @@ public CompletableFuture aPostOperationWithOut .withMarshaller(new APostOperationWithOutputRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperationWithOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperationWithOutput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(aPostOperationWithOutputRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -432,8 +452,8 @@ public CompletableFuture eventStreamOperation(EventStreamOperationRequest .withInitialRequestEvent(true).withResponseHandler(voidResponseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("EventStreamOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("EventStreamOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(eventStreamOperationRequest), asyncResponseTransformer); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { if (e != null) { @@ -529,9 +549,8 @@ public CompletableFuture eventStreamO .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("EventStreamOperationWithOnlyInput", clientConfiguration)) - .withEndpointResolver(endpointResolver("EventStreamOperationWithOnlyInput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(eventStreamOperationWithOnlyInputRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -638,9 +657,8 @@ public CompletableFuture eventStreamOperationWithOnlyOutput( .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("EventStreamOperationWithOnlyOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("EventStreamOperationWithOnlyOutput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(eventStreamOperationWithOnlyOutputRequest), asyncResponseTransformer); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { if (e != null) { @@ -727,8 +745,8 @@ public CompletableFuture getWithoutRequiredMe .withMarshaller(new GetWithoutRequiredMembersRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetWithoutRequiredMembers", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetWithoutRequiredMembers")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(getWithoutRequiredMembersRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -806,9 +824,8 @@ public CompletableFuture operationWithChe .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithChecksumRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithChecksumRequired")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.HTTP_CHECKSUM_REQUIRED, HttpChecksumRequired.create()).withInput(operationWithChecksumRequiredRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { @@ -883,8 +900,8 @@ public CompletableFuture operationWithNoneAut .withMarshaller(new OperationWithNoneAuthTypeRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithNoneAuthType", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithNoneAuthType")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithNoneAuthTypeRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -962,9 +979,8 @@ public CompletableFuture operationWithR .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithRequestCompression", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithRequestCompression")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.REQUEST_COMPRESSION, RequestCompression.builder().encodings("gzip").isStreaming(false).build()) .withInput(operationWithRequestCompressionRequest)); @@ -1044,9 +1060,8 @@ public CompletableFuture paginatedOpera .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("PaginatedOperationWithResultKey", clientConfiguration)) - .withEndpointResolver(endpointResolver("PaginatedOperationWithResultKey")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(paginatedOperationWithResultKeyRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1124,9 +1139,8 @@ public CompletableFuture paginatedOp .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("PaginatedOperationWithoutResultKey", clientConfiguration)) - .withEndpointResolver(endpointResolver("PaginatedOperationWithoutResultKey")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(paginatedOperationWithoutResultKeyRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1209,8 +1223,8 @@ public CompletableFuture streamingInputOperatio .asyncRequestBody(requestBody).build()).withResponseHandler(responseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingInputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOperation")).withAsyncRequestBody(requestBody) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withAsyncRequestBody(requestBody) .withInput(streamingInputOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1307,9 +1321,8 @@ public CompletableFuture streamingInputOutputOperation( .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("StreamingInputOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOutputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withAsyncRequestBody(requestBody).withAsyncResponseTransformer(asyncResponseTransformer) .withInput(streamingInputOutputOperationRequest), asyncResponseTransformer); AsyncResponseTransformer finalAsyncResponseTransformer = asyncResponseTransformer; @@ -1404,8 +1417,8 @@ public CompletableFuture streamingOutputOperation( .withMarshaller(new StreamingOutputOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingOutputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withAsyncResponseTransformer(asyncResponseTransformer).withInput(streamingOutputOperationRequest), asyncResponseTransformer); AsyncResponseTransformer finalAsyncResponseTransformer = asyncResponseTransformer; @@ -1459,8 +1472,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -1468,15 +1482,16 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri .isInstanceOf(JsonAuthSchemeProvider.class, p, "Expected an instance of JsonAuthSchemeProvider")) .orElse(null); JsonAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(JsonAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(JsonAuthSchemeProvider.class, executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of JsonAuthSchemeProvider"); JsonAuthSchemeParams.Builder paramsBuilder = JsonAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonEndpointProvider provider = (JsonEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -1507,14 +1522,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-cbor-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-cbor-client-class.java index 3198fa31df25..645560cb8cfa 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-cbor-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-cbor-client-class.java @@ -8,7 +8,7 @@ import java.util.function.Function; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -30,6 +30,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.interceptor.trait.HttpChecksumRequired; import software.amazon.awssdk.core.internal.interceptor.trait.RequestCompression; @@ -116,6 +117,25 @@ final class DefaultJsonClient implements JsonClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultJsonClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsSyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -186,8 +206,7 @@ public APostOperationResponse aPostOperation(APostOperationRequest aPostOperatio .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .hostPrefixExpression(resolvedHostExpression).withRequestConfiguration(clientConfiguration) .withInput(aPostOperationRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver).withEndpointResolver(endpointResolverInstance) .withMarshaller(new APostOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -256,9 +275,8 @@ public APostOperationWithOutputResponse aPostOperationWithOutput( .withOperationName("APostOperationWithOutput").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(aPostOperationWithOutputRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperationWithOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperationWithOutput")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new APostOperationWithOutputRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -326,9 +344,8 @@ public GetWithoutRequiredMembersResponse getWithoutRequiredMembers( .withOperationName("GetWithoutRequiredMembers").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(getWithoutRequiredMembersRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetWithoutRequiredMembers", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetWithoutRequiredMembers")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new GetWithoutRequiredMembersRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -396,9 +413,8 @@ public OperationWithChecksumRequiredResponse operationWithChecksumRequired( .withRequestConfiguration(clientConfiguration) .withInput(operationWithChecksumRequiredRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithChecksumRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithChecksumRequired")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.HTTP_CHECKSUM_REQUIRED, HttpChecksumRequired.create()) .withMarshaller(new OperationWithChecksumRequiredRequestMarshaller(protocolFactory))); @@ -464,9 +480,8 @@ public OperationWithNoneAuthTypeResponse operationWithNoneAuthType( .withOperationName("OperationWithNoneAuthType").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(operationWithNoneAuthTypeRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithNoneAuthType", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithNoneAuthType")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OperationWithNoneAuthTypeRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -534,9 +549,8 @@ public OperationWithRequestCompressionResponse operationWithRequestCompression( .withRequestConfiguration(clientConfiguration) .withInput(operationWithRequestCompressionRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithRequestCompression", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithRequestCompression")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.REQUEST_COMPRESSION, RequestCompression.builder().encodings("gzip").isStreaming(false).build()) .withMarshaller(new OperationWithRequestCompressionRequestMarshaller(protocolFactory))); @@ -599,16 +613,11 @@ public PaginatedOperationWithResultKeyResponse paginatedOperationWithResultKey( return clientHandler .execute(new ClientExecutionParams() - .withOperationName("PaginatedOperationWithResultKey") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withInput(paginatedOperationWithResultKeyRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("PaginatedOperationWithResultKey", clientConfiguration)) - .withEndpointResolver(endpointResolver("PaginatedOperationWithResultKey")) + .withOperationName("PaginatedOperationWithResultKey").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withInput(paginatedOperationWithResultKeyRequest) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new PaginatedOperationWithResultKeyRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -669,16 +678,11 @@ public PaginatedOperationWithoutResultKeyResponse paginatedOperationWithoutResul return clientHandler .execute(new ClientExecutionParams() - .withOperationName("PaginatedOperationWithoutResultKey") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withInput(paginatedOperationWithoutResultKeyRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("PaginatedOperationWithoutResultKey", clientConfiguration)) - .withEndpointResolver(endpointResolver("PaginatedOperationWithoutResultKey")) + .withOperationName("PaginatedOperationWithoutResultKey").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withInput(paginatedOperationWithoutResultKeyRequest) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new PaginatedOperationWithoutResultKeyRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -756,8 +760,8 @@ public StreamingInputOperationResponse streamingInputOperation(StreamingInputOpe .withRequestConfiguration(clientConfiguration) .withInput(streamingInputOperationRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingInputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withRequestBody(requestBody) .withMarshaller( StreamingRequestMarshaller.builder() @@ -848,9 +852,8 @@ public ReturnT streamingInputOutputOperation( .withRequestConfiguration(clientConfiguration) .withInput(streamingInputOutputOperationRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("StreamingInputOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOutputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withResponseTransformer(responseTransformer) .withRequestBody(requestBody) .withMarshaller( @@ -928,10 +931,8 @@ public ReturnT streamingOutputOperation(StreamingOutputOperationReques .withOperationName("StreamingOutputOperation").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(streamingOutputOperationRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingOutputOperation")) - .withResponseTransformer(responseTransformer) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withResponseTransformer(responseTransformer) .withMarshaller(new StreamingOutputOperationRequestMarshaller(protocolFactory)), responseTransformer); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -966,8 +967,8 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -975,15 +976,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri .isInstanceOf(JsonAuthSchemeProvider.class, p, "Expected an instance of JsonAuthSchemeProvider")) .orElse(null); JsonAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(JsonAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(JsonAuthSchemeProvider.class, + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of JsonAuthSchemeProvider"); JsonAuthSchemeParams.Builder paramsBuilder = JsonAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonEndpointProvider provider = (JsonEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -1014,14 +1017,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private HttpResponseHandler createErrorResponseHandler(BaseAwsJsonProtocolFactory protocolFactory, JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { return protocolFactory.createErrorResponseHandler(operationMetadata, exceptionMetadataMapper); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custom-context-params-async-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custom-context-params-async-client-class.java index 818e13b3663e..d30aac083af0 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custom-context-params-async-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custom-context-params-async-client-class.java @@ -14,7 +14,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -36,6 +36,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -78,7 +79,7 @@ final class DefaultFooBarAsyncClient implements FooBarAsyncClient { private static final Logger log = LoggerFactory.getLogger(DefaultFooBarAsyncClient.class); private static final AwsProtocolMetadata protocolMetadata = AwsProtocolMetadata.builder() - .serviceProtocol(AwsServiceProtocol.REST_JSON).build(); + .serviceProtocol(AwsServiceProtocol.REST_JSON).build(); private final AsyncClientHandler clientHandler; @@ -86,10 +87,29 @@ final class DefaultFooBarAsyncClient implements FooBarAsyncClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultFooBarAsyncClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsAsyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) - .option(SdkClientOption.API_METADATA, "Foo_Bar" + "#" + ServiceVersionInfo.VERSION).build(); + .option(SdkClientOption.API_METADATA, "Foo_Bar" + "#" + ServiceVersionInfo.VERSION).build(); this.protocolFactory = init(AwsJsonProtocolFactory.builder()).build(); } @@ -118,39 +138,39 @@ protected DefaultFooBarAsyncClient(SdkClientConfiguration clientConfiguration) { @Override public CompletableFuture getDatabaseVersion(GetDatabaseVersionRequest getDatabaseVersionRequest) { SdkClientConfiguration clientConfiguration = updateSdkClientConfiguration(getDatabaseVersionRequest, - this.clientConfiguration); + this.clientConfiguration); List metricPublishers = resolveMetricPublishers(clientConfiguration, getDatabaseVersionRequest - .overrideConfiguration().orElse(null)); + .overrideConfiguration().orElse(null)); MetricCollector apiCallMetricCollector = metricPublishers.isEmpty() ? NoOpMetricCollector.create() : MetricCollector - .create("ApiCall"); + .create("ApiCall"); try { apiCallMetricCollector.reportMetric(CoreMetric.SERVICE_ID, "Foo Bar"); apiCallMetricCollector.reportMetric(CoreMetric.OPERATION_NAME, "GetDatabaseVersion"); JsonOperationMetadata operationMetadata = JsonOperationMetadata.builder().hasStreamingSuccessResponse(false) - .isPayloadJson(true).build(); + .isPayloadJson(true).build(); HttpResponseHandler responseHandler = protocolFactory.createResponseHandler( - operationMetadata, GetDatabaseVersionResponse::builder); + operationMetadata, GetDatabaseVersionResponse::builder); Function> exceptionMetadataMapper = errorCode -> { if (errorCode == null) { return Optional.empty(); } switch (errorCode) { - default: - return Optional.empty(); + default: + return Optional.empty(); } }; HttpResponseHandler errorResponseHandler = createErrorResponseHandler(protocolFactory, - operationMetadata, exceptionMetadataMapper); + operationMetadata, exceptionMetadataMapper); CompletableFuture executeFuture = clientHandler - .execute(new ClientExecutionParams() - .withOperationName("GetDatabaseVersion").withProtocolMetadata(protocolMetadata) - .withMarshaller(new GetDatabaseVersionRequestMarshaller(protocolFactory)) - .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetDatabaseVersion", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetDatabaseVersion")).withInput(getDatabaseVersionRequest)); + .execute(new ClientExecutionParams() + .withOperationName("GetDatabaseVersion").withProtocolMetadata(protocolMetadata) + .withMarshaller(new GetDatabaseVersionRequestMarshaller(protocolFactory)) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(getDatabaseVersionRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -174,11 +194,11 @@ public final String serviceName() { private > T init(T builder) { return builder.clientConfiguration(clientConfiguration).defaultServiceExceptionSupplier(FooBarException::builder) - .protocol(AwsJsonProtocol.REST_JSON).protocolVersion("1.1"); + .protocol(AwsJsonProtocol.REST_JSON).protocolVersion("1.1"); } private static List resolveMetricPublishers(SdkClientConfiguration clientConfiguration, - RequestOverrideConfiguration requestOverrideConfiguration) { + RequestOverrideConfiguration requestOverrideConfiguration) { List publishers = null; if (requestOverrideConfiguration != null) { publishers = requestOverrideConfiguration.metricPublishers(); @@ -192,25 +212,27 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); FooBarAuthSchemeProvider requestAuthSchemeProvider = request - .overrideConfiguration() - .flatMap(c -> c.authSchemeProvider()) - .map(p -> Validate.isInstanceOf(FooBarAuthSchemeProvider.class, p, - "Expected an instance of FooBarAuthSchemeProvider")).orElse(null); + .overrideConfiguration() + .flatMap(c -> c.authSchemeProvider()) + .map(p -> Validate.isInstanceOf(FooBarAuthSchemeProvider.class, p, + "Expected an instance of FooBarAuthSchemeProvider")).orElse(null); FooBarAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(FooBarAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), - "Expected an instance of FooBarAuthSchemeProvider"); + .isInstanceOf(FooBarAuthSchemeProvider.class, + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), + "Expected an instance of FooBarAuthSchemeProvider"); FooBarAuthSchemeParams.Builder paramsBuilder = FooBarAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); FooBarEndpointProvider provider = (FooBarEndpointProvider) executionAttributes - .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); + .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { FooBarEndpointParams endpointParams = FooBarEndpointResolverUtils.ruleParams(request, executionAttributes); Endpoint endpoint = provider.resolveEndpoint(endpointParams).join(); @@ -222,10 +244,10 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } List endpointAuthSchemes = endpoint.attribute(AwsEndpointAttribute.AUTH_SCHEMES); SelectedAuthScheme selectedAuthScheme = executionAttributes - .getAttribute(SdkInternalExecutionAttribute.SELECTED_AUTH_SCHEME); + .getAttribute(SdkInternalExecutionAttribute.SELECTED_AUTH_SCHEME); if (endpointAuthSchemes != null && selectedAuthScheme != null) { selectedAuthScheme = FooBarEndpointResolverUtils.authSchemeWithEndpointSignerProperties(endpointAuthSchemes, - selectedAuthScheme); + selectedAuthScheme); executionAttributes.putAttribute(SdkInternalExecutionAttribute.SELECTED_AUTH_SCHEME, selectedAuthScheme); } FooBarEndpointResolverUtils.setMetricValues(endpoint, executionAttributes); @@ -239,14 +261,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); @@ -285,15 +299,15 @@ private SdkClientConfiguration updateSdkClientConfiguration(SdkRequest request, newContextParams = (newContextParams != null) ? newContextParams : AttributeMap.empty(); originalContextParams = originalContextParams != null ? originalContextParams : AttributeMap.empty(); Validate.validState( - Objects.equals(originalContextParams.get(FooBarClientContextParams.CROSS_REGION_ACCESS_ENABLED), - newContextParams.get(FooBarClientContextParams.CROSS_REGION_ACCESS_ENABLED)), - "CROSS_REGION_ACCESS_ENABLED cannot be modified by request level plugins"); + Objects.equals(originalContextParams.get(FooBarClientContextParams.CROSS_REGION_ACCESS_ENABLED), + newContextParams.get(FooBarClientContextParams.CROSS_REGION_ACCESS_ENABLED)), + "CROSS_REGION_ACCESS_ENABLED cannot be modified by request level plugins"); updateRetryStrategyClientConfiguration(configuration); return configuration.build(); } private HttpResponseHandler createErrorResponseHandler(BaseAwsJsonProtocolFactory protocolFactory, - JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { + JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { return protocolFactory.createErrorResponseHandler(operationMetadata, exceptionMetadataMapper); } diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custom-context-params-sync-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custom-context-params-sync-client-class.java index 6e29b234a04a..abb095a80e2c 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custom-context-params-sync-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custom-context-params-sync-client-class.java @@ -9,7 +9,7 @@ import java.util.function.Function; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -31,6 +31,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -81,6 +82,25 @@ final class DefaultFooBarClient implements FooBarClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultFooBarClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsSyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -140,8 +160,8 @@ public GetDatabaseVersionResponse getDatabaseVersion(GetDatabaseVersionRequest g .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(getDatabaseVersionRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetDatabaseVersion", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetDatabaseVersion")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new GetDatabaseVersionRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -168,23 +188,25 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); FooBarAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) .map(p -> Validate.isInstanceOf(FooBarAuthSchemeProvider.class, p, "Expected an instance of FooBarAuthSchemeProvider")).orElse(null); FooBarAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(FooBarAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(FooBarAuthSchemeProvider.class, executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of FooBarAuthSchemeProvider"); FooBarAuthSchemeParams.Builder paramsBuilder = FooBarAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); FooBarEndpointProvider provider = (FooBarEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -215,14 +237,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private HttpResponseHandler createErrorResponseHandler(BaseAwsJsonProtocolFactory protocolFactory, JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { return protocolFactory.createErrorResponseHandler(operationMetadata, exceptionMetadataMapper); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custompackage-async.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custompackage-async.java index aa23153cdd20..edf79c53abe2 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custompackage-async.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custompackage-async.java @@ -24,7 +24,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -46,6 +46,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -83,6 +84,25 @@ final class DefaultProtocolRestJsonWithCustomPackageAsyncClient implements Proto private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultProtocolRestJsonWithCustomPackageAsyncClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsAsyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration @@ -146,8 +166,8 @@ public CompletableFuture oneOperation(OneOperationRequest .withMarshaller(new OneOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OneOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("OneOperation")).withInput(oneOperationRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(oneOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -191,8 +211,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); ProtocolRestJsonWithCustomPackageAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -200,16 +221,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of ProtocolRestJsonWithCustomPackageAuthSchemeProvider")).orElse(null); ProtocolRestJsonWithCustomPackageAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(ProtocolRestJsonWithCustomPackageAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of ProtocolRestJsonWithCustomPackageAuthSchemeProvider"); ProtocolRestJsonWithCustomPackageAuthSchemeParams.Builder paramsBuilder = ProtocolRestJsonWithCustomPackageAuthSchemeParams .builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); ProtocolRestJsonWithCustomPackageEndpointProvider provider = (ProtocolRestJsonWithCustomPackageEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -242,14 +264,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custompackage-sync.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custompackage-sync.java index 24811ad10405..535b68770d32 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custompackage-sync.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-custompackage-sync.java @@ -19,7 +19,7 @@ import java.util.function.Function; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -41,6 +41,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -78,6 +79,25 @@ final class DefaultProtocolRestJsonWithCustomPackageClient implements ProtocolRe private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultProtocolRestJsonWithCustomPackageClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsSyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration @@ -137,8 +157,8 @@ public OneOperationResponse oneOperation(OneOperationRequest oneOperationRequest .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(oneOperationRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OneOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("OneOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OneOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -165,8 +185,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); ProtocolRestJsonWithCustomPackageAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -174,16 +195,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of ProtocolRestJsonWithCustomPackageAuthSchemeProvider")).orElse(null); ProtocolRestJsonWithCustomPackageAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(ProtocolRestJsonWithCustomPackageAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of ProtocolRestJsonWithCustomPackageAuthSchemeProvider"); ProtocolRestJsonWithCustomPackageAuthSchemeParams.Builder paramsBuilder = ProtocolRestJsonWithCustomPackageAuthSchemeParams .builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); ProtocolRestJsonWithCustomPackageEndpointProvider provider = (ProtocolRestJsonWithCustomPackageEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -216,14 +238,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private HttpResponseHandler createErrorResponseHandler(BaseAwsJsonProtocolFactory protocolFactory, JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { return protocolFactory.createErrorResponseHandler(operationMetadata, exceptionMetadataMapper); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-customservicemetadata-async.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-customservicemetadata-async.java index a4956ab331f3..ec00ec9dc270 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-customservicemetadata-async.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-customservicemetadata-async.java @@ -13,7 +13,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -35,6 +35,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -83,6 +84,25 @@ final class DefaultProtocolRestJsonWithCustomContentTypeAsyncClient implements P private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultProtocolRestJsonWithCustomContentTypeAsyncClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsAsyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration @@ -146,8 +166,8 @@ public CompletableFuture oneOperation(OneOperationRequest .withMarshaller(new OneOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OneOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("OneOperation")).withInput(oneOperationRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(oneOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -191,8 +211,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); ProtocolRestJsonWithCustomContentTypeAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -200,16 +221,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of ProtocolRestJsonWithCustomContentTypeAuthSchemeProvider")).orElse(null); ProtocolRestJsonWithCustomContentTypeAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(ProtocolRestJsonWithCustomContentTypeAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of ProtocolRestJsonWithCustomContentTypeAuthSchemeProvider"); ProtocolRestJsonWithCustomContentTypeAuthSchemeParams.Builder paramsBuilder = ProtocolRestJsonWithCustomContentTypeAuthSchemeParams .builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); ProtocolRestJsonWithCustomContentTypeEndpointProvider provider = (ProtocolRestJsonWithCustomContentTypeEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -242,14 +264,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-customservicemetadata-sync.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-customservicemetadata-sync.java index 31f56fe814ef..0a3a4e5d0670 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-customservicemetadata-sync.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-customservicemetadata-sync.java @@ -8,7 +8,7 @@ import java.util.function.Function; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -30,6 +30,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -78,6 +79,25 @@ final class DefaultProtocolRestJsonWithCustomContentTypeClient implements Protoc private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultProtocolRestJsonWithCustomContentTypeClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsSyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration @@ -137,8 +157,8 @@ public OneOperationResponse oneOperation(OneOperationRequest oneOperationRequest .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(oneOperationRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OneOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("OneOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OneOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -165,8 +185,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); ProtocolRestJsonWithCustomContentTypeAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -174,16 +195,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of ProtocolRestJsonWithCustomContentTypeAuthSchemeProvider")).orElse(null); ProtocolRestJsonWithCustomContentTypeAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(ProtocolRestJsonWithCustomContentTypeAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of ProtocolRestJsonWithCustomContentTypeAuthSchemeProvider"); ProtocolRestJsonWithCustomContentTypeAuthSchemeParams.Builder paramsBuilder = ProtocolRestJsonWithCustomContentTypeAuthSchemeParams .builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); ProtocolRestJsonWithCustomContentTypeEndpointProvider provider = (ProtocolRestJsonWithCustomContentTypeEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -216,14 +238,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private HttpResponseHandler createErrorResponseHandler(BaseAwsJsonProtocolFactory protocolFactory, JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { return protocolFactory.createErrorResponseHandler(operationMetadata, exceptionMetadataMapper); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-endpoint-discovery-async.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-endpoint-discovery-async.java index d043a064e0ee..fa454028fd42 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-endpoint-discovery-async.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-endpoint-discovery-async.java @@ -14,6 +14,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.AwsRequestOverrideConfiguration; import software.amazon.awssdk.awscore.client.config.AwsClientOption; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; @@ -39,6 +40,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -97,6 +99,25 @@ final class DefaultEndpointDiscoveryTestAsyncClient implements EndpointDiscovery private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + private EndpointDiscoveryRefreshCache endpointDiscoveryCache; protected DefaultEndpointDiscoveryTestAsyncClient(SdkClientConfiguration clientConfiguration) { @@ -162,8 +183,8 @@ public CompletableFuture describeEndpoints(DescribeEn .withMarshaller(new DescribeEndpointsRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("DescribeEndpoints", clientConfiguration)) - .withEndpointResolver(endpointResolver("DescribeEndpoints")).withInput(describeEndpointsRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(describeEndpointsRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -254,17 +275,13 @@ public CompletableFuture testDiscovery CompletableFuture executeFuture = endpointFuture .thenCompose(cachedEndpoint -> clientHandler .execute(new ClientExecutionParams() - .withOperationName("TestDiscoveryIdentifiersRequired") - .withProtocolMetadata(protocolMetadata) + .withOperationName("TestDiscoveryIdentifiersRequired").withProtocolMetadata(protocolMetadata) .withMarshaller(new TestDiscoveryIdentifiersRequiredRequestMarshaller(protocolFactory)) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("TestDiscoveryIdentifiersRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("TestDiscoveryIdentifiersRequired")) - .discoveredEndpoint(cachedEndpoint).withInput(testDiscoveryIdentifiersRequiredRequest))); + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).discoveredEndpoint(cachedEndpoint) + .withInput(testDiscoveryIdentifiersRequiredRequest))); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -345,17 +362,13 @@ public CompletableFuture testDiscoveryOptional( CompletableFuture executeFuture = endpointFuture .thenCompose(cachedEndpoint -> clientHandler .execute(new ClientExecutionParams() - .withOperationName("TestDiscoveryOptional") - .withProtocolMetadata(protocolMetadata) + .withOperationName("TestDiscoveryOptional").withProtocolMetadata(protocolMetadata) .withMarshaller(new TestDiscoveryOptionalRequestMarshaller(protocolFactory)) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("TestDiscoveryOptional", clientConfiguration)) - .withEndpointResolver(endpointResolver("TestDiscoveryOptional")) - .discoveredEndpoint(cachedEndpoint).withInput(testDiscoveryOptionalRequest))); + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).discoveredEndpoint(cachedEndpoint) + .withInput(testDiscoveryOptionalRequest))); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -444,17 +457,13 @@ public CompletableFuture testDiscoveryRequired( CompletableFuture executeFuture = endpointFuture .thenCompose(cachedEndpoint -> clientHandler .execute(new ClientExecutionParams() - .withOperationName("TestDiscoveryRequired") - .withProtocolMetadata(protocolMetadata) + .withOperationName("TestDiscoveryRequired").withProtocolMetadata(protocolMetadata) .withMarshaller(new TestDiscoveryRequiredRequestMarshaller(protocolFactory)) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("TestDiscoveryRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("TestDiscoveryRequired")) - .discoveredEndpoint(cachedEndpoint).withInput(testDiscoveryRequiredRequest))); + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).discoveredEndpoint(cachedEndpoint) + .withInput(testDiscoveryRequiredRequest))); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -497,8 +506,8 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); EndpointDiscoveryTestAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -506,16 +515,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of EndpointDiscoveryTestAuthSchemeProvider")).orElse(null); EndpointDiscoveryTestAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(EndpointDiscoveryTestAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of EndpointDiscoveryTestAuthSchemeProvider"); EndpointDiscoveryTestAuthSchemeParams.Builder paramsBuilder = EndpointDiscoveryTestAuthSchemeParams.builder().operation( operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); EndpointDiscoveryTestEndpointProvider provider = (EndpointDiscoveryTestEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -547,14 +557,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-endpoint-discovery-sync.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-endpoint-discovery-sync.java index e20ee0df3404..c4c57d8b5705 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-endpoint-discovery-sync.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-endpoint-discovery-sync.java @@ -10,6 +10,7 @@ import java.util.function.Function; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.AwsRequestOverrideConfiguration; import software.amazon.awssdk.awscore.client.config.AwsClientOption; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; @@ -35,6 +36,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -94,6 +96,25 @@ final class DefaultEndpointDiscoveryTestClient implements EndpointDiscoveryTestC private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + private EndpointDiscoveryRefreshCache endpointDiscoveryCache; protected DefaultEndpointDiscoveryTestClient(SdkClientConfiguration clientConfiguration) { @@ -154,9 +175,8 @@ public DescribeEndpointsResponse describeEndpoints(DescribeEndpointsRequest desc .withOperationName("DescribeEndpoints").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(describeEndpointsRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("DescribeEndpoints", clientConfiguration)) - .withEndpointResolver(endpointResolver("DescribeEndpoints")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new DescribeEndpointsRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -230,17 +250,12 @@ public TestDiscoveryIdentifiersRequiredResponse testDiscoveryIdentifiersRequired return clientHandler .execute(new ClientExecutionParams() - .withOperationName("TestDiscoveryIdentifiersRequired") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .discoveredEndpoint(cachedEndpoint) - .withRequestConfiguration(clientConfiguration) - .withInput(testDiscoveryIdentifiersRequiredRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("TestDiscoveryIdentifiersRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("TestDiscoveryIdentifiersRequired")) + .withOperationName("TestDiscoveryIdentifiersRequired").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .discoveredEndpoint(cachedEndpoint).withRequestConfiguration(clientConfiguration) + .withInput(testDiscoveryIdentifiersRequiredRequest).withMetricCollector(apiCallMetricCollector) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new TestDiscoveryIdentifiersRequiredRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -308,8 +323,7 @@ public TestDiscoveryOptionalResponse testDiscoveryOptional(TestDiscoveryOptional .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .discoveredEndpoint(cachedEndpoint).withRequestConfiguration(clientConfiguration) .withInput(testDiscoveryOptionalRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("TestDiscoveryOptional", clientConfiguration)) - .withEndpointResolver(endpointResolver("TestDiscoveryOptional")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver).withEndpointResolver(endpointResolverInstance) .withMarshaller(new TestDiscoveryOptionalRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -385,8 +399,7 @@ public TestDiscoveryRequiredResponse testDiscoveryRequired(TestDiscoveryRequired .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .discoveredEndpoint(cachedEndpoint).withRequestConfiguration(clientConfiguration) .withInput(testDiscoveryRequiredRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("TestDiscoveryRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("TestDiscoveryRequired")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver).withEndpointResolver(endpointResolverInstance) .withMarshaller(new TestDiscoveryRequiredRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -413,8 +426,8 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); EndpointDiscoveryTestAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -422,16 +435,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of EndpointDiscoveryTestAuthSchemeProvider")).orElse(null); EndpointDiscoveryTestAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(EndpointDiscoveryTestAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of EndpointDiscoveryTestAuthSchemeProvider"); EndpointDiscoveryTestAuthSchemeParams.Builder paramsBuilder = EndpointDiscoveryTestAuthSchemeParams.builder().operation( operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); EndpointDiscoveryTestEndpointProvider provider = (EndpointDiscoveryTestEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -463,14 +477,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private HttpResponseHandler createErrorResponseHandler(BaseAwsJsonProtocolFactory protocolFactory, JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { return protocolFactory.createErrorResponseHandler(operationMetadata, exceptionMetadataMapper); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-json-async-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-json-async-client-class.java index 552f5f7056a9..e0d41cee1f32 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-json-async-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-json-async-client-class.java @@ -17,7 +17,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.client.handler.AwsClientHandlerUtils; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; @@ -54,6 +54,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.interceptor.trait.HttpChecksum; import software.amazon.awssdk.core.interceptor.trait.HttpChecksumRequired; @@ -167,6 +168,25 @@ final class DefaultJsonAsyncClient implements JsonAsyncClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + private final ScheduledExecutorService executorService; private final Executor executor; @@ -249,9 +269,9 @@ public CompletableFuture aPostOperation(APostOperationRe .withMarshaller(new APostOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")) - .hostPrefixExpression(resolvedHostExpression).withInput(aPostOperationRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).hostPrefixExpression(resolvedHostExpression) + .withInput(aPostOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -325,9 +345,8 @@ public CompletableFuture aPostOperationWithOut .withMarshaller(new APostOperationWithOutputRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperationWithOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperationWithOutput")) - .withInput(aPostOperationWithOutputRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(aPostOperationWithOutputRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -397,8 +416,8 @@ public CompletableFuture bearerAuthOperation( .withMarshaller(new BearerAuthOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("BearerAuthOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("BearerAuthOperation")).credentialType(CredentialType.TOKEN) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).credentialType(CredentialType.TOKEN) .withInput(bearerAuthOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -513,9 +532,9 @@ public CompletableFuture eventStreamOperation(EventStreamOperationRequest .withAsyncRequestBody(AsyncRequestBody.fromPublisher(adapted)).withFullDuplex(true) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("EventStreamOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("EventStreamOperation")) - .withInput(eventStreamOperationRequest), restAsyncResponseTransformer); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(eventStreamOperationRequest), + restAsyncResponseTransformer); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { if (e != null) { try { @@ -598,18 +617,13 @@ public CompletableFuture eventStreamO CompletableFuture executeFuture = clientHandler .execute(new ClientExecutionParams() - .withOperationName("EventStreamOperationWithOnlyInput") - .withProtocolMetadata(protocolMetadata) + .withOperationName("EventStreamOperationWithOnlyInput").withProtocolMetadata(protocolMetadata) .withMarshaller(new EventStreamOperationWithOnlyInputRequestMarshaller(protocolFactory)) - .withAsyncRequestBody(AsyncRequestBody.fromPublisher(adapted)) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) + .withAsyncRequestBody(AsyncRequestBody.fromPublisher(adapted)).withResponseHandler(responseHandler) + .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("EventStreamOperationWithOnlyInput", clientConfiguration)) - .withEndpointResolver(endpointResolver("EventStreamOperationWithOnlyInput")) - .withInput(eventStreamOperationWithOnlyInputRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(eventStreamOperationWithOnlyInputRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -719,13 +733,10 @@ public CompletableFuture eventStreamOperationWithOnlyOutput( .withOperationName("EventStreamOperationWithOnlyOutput") .withProtocolMetadata(protocolMetadata) .withMarshaller(new EventStreamOperationWithOnlyOutputRequestMarshaller(protocolFactory)) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("EventStreamOperationWithOnlyOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("EventStreamOperationWithOnlyOutput")) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(eventStreamOperationWithOnlyOutputRequest), restAsyncResponseTransformer); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { if (e != null) { @@ -808,8 +819,8 @@ public CompletableFuture getOperationWithCheck .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum.builder().requestChecksumRequired(true).isRequestStreaming(false) @@ -889,9 +900,8 @@ public CompletableFuture getWithoutRequiredMe .withMarshaller(new GetWithoutRequiredMembersRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetWithoutRequiredMembers", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetWithoutRequiredMembers")) - .withInput(getWithoutRequiredMembersRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(getWithoutRequiredMembersRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -965,9 +975,8 @@ public CompletableFuture operationWithChe .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithChecksumRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithChecksumRequired")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.HTTP_CHECKSUM_REQUIRED, HttpChecksumRequired.create()).withInput(operationWithChecksumRequiredRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { @@ -1043,9 +1052,8 @@ public CompletableFuture operationWithR .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithRequestCompression", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithRequestCompression")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.REQUEST_COMPRESSION, RequestCompression.builder().encodings("gzip").isStreaming(false).build()) .withInput(operationWithRequestCompressionRequest)); @@ -1115,17 +1123,12 @@ public CompletableFuture paginatedOpera CompletableFuture executeFuture = clientHandler .execute(new ClientExecutionParams() - .withOperationName("PaginatedOperationWithResultKey") - .withProtocolMetadata(protocolMetadata) + .withOperationName("PaginatedOperationWithResultKey").withProtocolMetadata(protocolMetadata) .withMarshaller(new PaginatedOperationWithResultKeyRequestMarshaller(protocolFactory)) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("PaginatedOperationWithResultKey", clientConfiguration)) - .withEndpointResolver(endpointResolver("PaginatedOperationWithResultKey")) - .withInput(paginatedOperationWithResultKeyRequest)); + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(paginatedOperationWithResultKeyRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -1192,17 +1195,12 @@ public CompletableFuture paginatedOp CompletableFuture executeFuture = clientHandler .execute(new ClientExecutionParams() - .withOperationName("PaginatedOperationWithoutResultKey") - .withProtocolMetadata(protocolMetadata) + .withOperationName("PaginatedOperationWithoutResultKey").withProtocolMetadata(protocolMetadata) .withMarshaller(new PaginatedOperationWithoutResultKeyRequestMarshaller(protocolFactory)) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("PaginatedOperationWithoutResultKey", clientConfiguration)) - .withEndpointResolver(endpointResolver("PaginatedOperationWithoutResultKey")) - .withInput(paginatedOperationWithoutResultKeyRequest)); + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(paginatedOperationWithoutResultKeyRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -1301,8 +1299,8 @@ public CompletableFuture putOperationWithChecksum( .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("PutOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("PutOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withAsyncRequestBody(requestBody) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, @@ -1409,8 +1407,8 @@ public CompletableFuture streamingInputOperatio .asyncRequestBody(requestBody).build()).withResponseHandler(responseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingInputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOperation")).withAsyncRequestBody(requestBody) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withAsyncRequestBody(requestBody) .withInput(streamingInputOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1500,14 +1498,11 @@ public CompletableFuture streamingInputOutputOperation( .delegateMarshaller( new StreamingInputOutputOperationRequestMarshaller(protocolFactory)) .asyncRequestBody(requestBody).transferEncoding(true).build()) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("StreamingInputOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOutputOperation")) - .withAsyncRequestBody(requestBody).withAsyncResponseTransformer(asyncResponseTransformer) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withAsyncRequestBody(requestBody) + .withAsyncResponseTransformer(asyncResponseTransformer) .withInput(streamingInputOutputOperationRequest), asyncResponseTransformer); AsyncResponseTransformer finalAsyncResponseTransformer = asyncResponseTransformer; CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { @@ -1598,10 +1593,9 @@ public CompletableFuture streamingOutputOperation( .withMarshaller(new StreamingOutputOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingOutputOperation")) - .withAsyncResponseTransformer(asyncResponseTransformer).withInput(streamingOutputOperationRequest), - asyncResponseTransformer); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withAsyncResponseTransformer(asyncResponseTransformer) + .withInput(streamingOutputOperationRequest), asyncResponseTransformer); AsyncResponseTransformer finalAsyncResponseTransformer = asyncResponseTransformer; CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { if (e != null) { @@ -1658,8 +1652,8 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -1667,15 +1661,16 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri .isInstanceOf(JsonAuthSchemeProvider.class, p, "Expected an instance of JsonAuthSchemeProvider")) .orElse(null); JsonAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(JsonAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(JsonAuthSchemeProvider.class, executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of JsonAuthSchemeProvider"); JsonAuthSchemeParams.Builder paramsBuilder = JsonAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonEndpointProvider provider = (JsonEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -1706,14 +1701,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); @@ -1760,4 +1747,4 @@ private HttpResponseHandler createErrorResponseHandler(Base public void close() { clientHandler.close(); } -} +} \ No newline at end of file diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-json-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-json-client-class.java index a557baa0385d..720b097bed52 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-json-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-json-client-class.java @@ -8,7 +8,7 @@ import java.util.function.Function; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -32,6 +32,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.interceptor.trait.HttpChecksum; import software.amazon.awssdk.core.interceptor.trait.HttpChecksumRequired; @@ -124,6 +125,25 @@ final class DefaultJsonClient implements JsonClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultJsonClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsSyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -191,8 +211,7 @@ public APostOperationResponse aPostOperation(APostOperationRequest aPostOperatio .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .hostPrefixExpression(resolvedHostExpression).withRequestConfiguration(clientConfiguration) .withInput(aPostOperationRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver).withEndpointResolver(endpointResolverInstance) .withMarshaller(new APostOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -257,9 +276,8 @@ public APostOperationWithOutputResponse aPostOperationWithOutput( .withOperationName("APostOperationWithOutput").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(aPostOperationWithOutputRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperationWithOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperationWithOutput")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new APostOperationWithOutputRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -319,8 +337,7 @@ public BearerAuthOperationResponse bearerAuthOperation(BearerAuthOperationReques .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .credentialType(CredentialType.TOKEN).withRequestConfiguration(clientConfiguration) .withInput(bearerAuthOperationRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("BearerAuthOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("BearerAuthOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver).withEndpointResolver(endpointResolverInstance) .withMarshaller(new BearerAuthOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -385,8 +402,8 @@ public GetOperationWithChecksumResponse getOperationWithChecksum( .withRequestConfiguration(clientConfiguration) .withInput(getOperationWithChecksumRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum.builder().requestChecksumRequired(true).isRequestStreaming(false) @@ -456,9 +473,8 @@ public GetWithoutRequiredMembersResponse getWithoutRequiredMembers( .withOperationName("GetWithoutRequiredMembers").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(getWithoutRequiredMembersRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetWithoutRequiredMembers", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetWithoutRequiredMembers")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new GetWithoutRequiredMembersRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -523,9 +539,8 @@ public OperationWithChecksumRequiredResponse operationWithChecksumRequired( .withRequestConfiguration(clientConfiguration) .withInput(operationWithChecksumRequiredRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithChecksumRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithChecksumRequired")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.HTTP_CHECKSUM_REQUIRED, HttpChecksumRequired.create()) .withMarshaller(new OperationWithChecksumRequiredRequestMarshaller(protocolFactory))); @@ -592,9 +607,8 @@ public OperationWithRequestCompressionResponse operationWithRequestCompression( .withRequestConfiguration(clientConfiguration) .withInput(operationWithRequestCompressionRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithRequestCompression", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithRequestCompression")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.REQUEST_COMPRESSION, RequestCompression.builder().encodings("gzip").isStreaming(false).build()) .withMarshaller(new OperationWithRequestCompressionRequestMarshaller(protocolFactory))); @@ -654,16 +668,11 @@ public PaginatedOperationWithResultKeyResponse paginatedOperationWithResultKey( return clientHandler .execute(new ClientExecutionParams() - .withOperationName("PaginatedOperationWithResultKey") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withInput(paginatedOperationWithResultKeyRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("PaginatedOperationWithResultKey", clientConfiguration)) - .withEndpointResolver(endpointResolver("PaginatedOperationWithResultKey")) + .withOperationName("PaginatedOperationWithResultKey").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withInput(paginatedOperationWithResultKeyRequest) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new PaginatedOperationWithResultKeyRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -721,16 +730,11 @@ public PaginatedOperationWithoutResultKeyResponse paginatedOperationWithoutResul return clientHandler .execute(new ClientExecutionParams() - .withOperationName("PaginatedOperationWithoutResultKey") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withInput(paginatedOperationWithoutResultKeyRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("PaginatedOperationWithoutResultKey", clientConfiguration)) - .withEndpointResolver(endpointResolver("PaginatedOperationWithoutResultKey")) + .withOperationName("PaginatedOperationWithoutResultKey").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withInput(paginatedOperationWithoutResultKeyRequest) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new PaginatedOperationWithoutResultKeyRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -821,8 +825,8 @@ public ReturnT putOperationWithChecksum(PutOperationWithChecksumReques .withRequestConfiguration(clientConfiguration) .withInput(putOperationWithChecksumRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("PutOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("PutOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum @@ -917,8 +921,8 @@ public StreamingInputOperationResponse streamingInputOperation(StreamingInputOpe .withRequestConfiguration(clientConfiguration) .withInput(streamingInputOperationRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingInputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withRequestBody(requestBody) .withMarshaller( StreamingRequestMarshaller.builder() @@ -1006,9 +1010,8 @@ public ReturnT streamingInputOutputOperation( .withRequestConfiguration(clientConfiguration) .withInput(streamingInputOutputOperationRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("StreamingInputOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOutputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withResponseTransformer(responseTransformer) .withRequestBody(requestBody) .withMarshaller( @@ -1083,10 +1086,8 @@ public ReturnT streamingOutputOperation(StreamingOutputOperationReques .withOperationName("StreamingOutputOperation").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(streamingOutputOperationRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingOutputOperation")) - .withResponseTransformer(responseTransformer) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withResponseTransformer(responseTransformer) .withMarshaller(new StreamingOutputOperationRequestMarshaller(protocolFactory)), responseTransformer); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1121,8 +1122,8 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -1130,15 +1131,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri .isInstanceOf(JsonAuthSchemeProvider.class, p, "Expected an instance of JsonAuthSchemeProvider")) .orElse(null); JsonAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(JsonAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(JsonAuthSchemeProvider.class, + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of JsonAuthSchemeProvider"); JsonAuthSchemeParams.Builder paramsBuilder = JsonAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonEndpointProvider provider = (JsonEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -1169,14 +1172,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private HttpResponseHandler createErrorResponseHandler(BaseAwsJsonProtocolFactory protocolFactory, JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { return protocolFactory.createErrorResponseHandler(operationMetadata, exceptionMetadataMapper); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-presignedurl-async.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-presignedurl-async.java index 5e9fdee4bcc6..28d12dca6318 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-presignedurl-async.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-presignedurl-async.java @@ -13,7 +13,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -35,6 +35,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -86,6 +87,25 @@ final class DefaultJsonAsyncClient implements JsonAsyncClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultJsonAsyncClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsAsyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -153,8 +173,8 @@ public CompletableFuture aPostOperation(APostOperationRe .withMarshaller(new APostOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")).withInput(aPostOperationRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(aPostOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -201,8 +221,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -210,15 +231,16 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri .isInstanceOf(JsonAuthSchemeProvider.class, p, "Expected an instance of JsonAuthSchemeProvider")) .orElse(null); JsonAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(JsonAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(JsonAuthSchemeProvider.class, executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of JsonAuthSchemeProvider"); JsonAuthSchemeParams.Builder paramsBuilder = JsonAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); JsonEndpointProvider provider = (JsonEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -249,14 +271,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-query-async-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-query-async-client-class.java index 0a6706911756..edd6f97af768 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-query-async-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-query-async-client-class.java @@ -13,7 +13,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -41,6 +41,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.interceptor.trait.HttpChecksum; import software.amazon.awssdk.core.interceptor.trait.HttpChecksumRequired; @@ -138,6 +139,25 @@ final class DefaultQueryAsyncClient implements QueryAsyncClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + private final ScheduledExecutorService executorService; protected DefaultQueryAsyncClient(SdkClientConfiguration clientConfiguration) { @@ -196,8 +216,8 @@ public CompletableFuture aPostOperation(APostOperationRe .withMarshaller(new APostOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .hostPrefixExpression(resolvedHostExpression).withInput(aPostOperationRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -258,8 +278,8 @@ public CompletableFuture aPostOperationWithOut .withMarshaller(new APostOperationWithOutputRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperationWithOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperationWithOutput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(aPostOperationWithOutputRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -317,8 +337,8 @@ public CompletableFuture bearerAuthOperation( .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .credentialType(CredentialType.TOKEN).withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("BearerAuthOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("BearerAuthOperation")).withInput(bearerAuthOperationRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(bearerAuthOperationRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -377,8 +397,8 @@ public CompletableFuture getOperationWithCheck .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum.builder().requestChecksumRequired(true).isRequestStreaming(false) @@ -444,9 +464,8 @@ public CompletableFuture operationWithChe .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithChecksumRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithChecksumRequired")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.HTTP_CHECKSUM_REQUIRED, HttpChecksumRequired.create()).withInput(operationWithChecksumRequiredRequest)); CompletableFuture whenCompleteFuture = null; @@ -504,8 +523,8 @@ public CompletableFuture operationWithContext .withMarshaller(new OperationWithContextParamRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithContextParam", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithContextParam")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithContextParamRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -563,8 +582,8 @@ public CompletableFuture operationWithCustomM .withMarshaller(new OperationWithCustomMemberRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithCustomMember", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithCustomMember")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithCustomMemberRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -626,9 +645,8 @@ public CompletableFuture o .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithCustomizedOperationContextParam", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithCustomizedOperationContextParam")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithCustomizedOperationContextParamRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -690,9 +708,8 @@ public CompletableFuture operatio .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithMapOperationContextParam", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithMapOperationContextParam")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithMapOperationContextParamRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -749,8 +766,8 @@ public CompletableFuture operationWithNoneAut .withMarshaller(new OperationWithNoneAuthTypeRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithNoneAuthType", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithNoneAuthType")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithNoneAuthTypeRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -812,9 +829,8 @@ public CompletableFuture operationWi .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithOperationContextParam", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithOperationContextParam")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithOperationContextParamRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -875,9 +891,8 @@ public CompletableFuture operationWithR .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithRequestCompression", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithRequestCompression")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.REQUEST_COMPRESSION, RequestCompression.builder().encodings("gzip").isStreaming(false).build()) .withInput(operationWithRequestCompressionRequest)); @@ -940,9 +955,8 @@ public CompletableFuture operationWith .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithStaticContextParams", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithStaticContextParams")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithStaticContextParamsRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -1028,8 +1042,8 @@ public CompletableFuture putOperationWithChecksum( .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("PutOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("PutOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum @@ -1119,8 +1133,8 @@ public CompletableFuture streamingInputOperatio .asyncRequestBody(requestBody).build()).withResponseHandler(responseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingInputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOperation")).withAsyncRequestBody(requestBody) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withAsyncRequestBody(requestBody) .withInput(streamingInputOperationRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -1187,8 +1201,8 @@ public CompletableFuture streamingOutputOperation( .withMarshaller(new StreamingOutputOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingOutputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withAsyncResponseTransformer(asyncResponseTransformer).withInput(streamingOutputOperationRequest), asyncResponseTransformer); CompletableFuture whenCompleteFuture = null; @@ -1251,23 +1265,25 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); QueryAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) .map(p -> Validate.isInstanceOf(QueryAuthSchemeProvider.class, p, "Expected an instance of QueryAuthSchemeProvider")).orElse(null); QueryAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(QueryAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(QueryAuthSchemeProvider.class, executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of QueryAuthSchemeProvider"); QueryAuthSchemeParams.Builder paramsBuilder = QueryAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); QueryEndpointProvider provider = (QueryEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -1298,14 +1314,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-query-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-query-client-class.java index 004e8f563a94..5cb8efc9f1d7 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-query-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-query-client-class.java @@ -7,7 +7,7 @@ import java.util.function.Consumer; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -32,6 +32,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.interceptor.trait.HttpChecksum; import software.amazon.awssdk.core.interceptor.trait.HttpChecksumRequired; @@ -130,6 +131,25 @@ final class DefaultQueryClient implements QueryClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultQueryClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsSyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -181,8 +201,7 @@ public APostOperationResponse aPostOperation(APostOperationRequest aPostOperatio .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .hostPrefixExpression(resolvedHostExpression).withRequestConfiguration(clientConfiguration) .withInput(aPostOperationRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver).withEndpointResolver(endpointResolverInstance) .withMarshaller(new APostOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -233,9 +252,8 @@ public APostOperationWithOutputResponse aPostOperationWithOutput( .withOperationName("APostOperationWithOutput").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(aPostOperationWithOutputRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperationWithOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperationWithOutput")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new APostOperationWithOutputRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -281,8 +299,7 @@ public BearerAuthOperationResponse bearerAuthOperation(BearerAuthOperationReques .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .credentialType(CredentialType.TOKEN).withRequestConfiguration(clientConfiguration) .withInput(bearerAuthOperationRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("BearerAuthOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("BearerAuthOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver).withEndpointResolver(endpointResolverInstance) .withMarshaller(new BearerAuthOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -333,8 +350,8 @@ public GetOperationWithChecksumResponse getOperationWithChecksum( .withRequestConfiguration(clientConfiguration) .withInput(getOperationWithChecksumRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum.builder().requestChecksumRequired(true).isRequestStreaming(false) @@ -390,9 +407,8 @@ public OperationWithChecksumRequiredResponse operationWithChecksumRequired( .withRequestConfiguration(clientConfiguration) .withInput(operationWithChecksumRequiredRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithChecksumRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithChecksumRequired")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.HTTP_CHECKSUM_REQUIRED, HttpChecksumRequired.create()) .withMarshaller(new OperationWithChecksumRequiredRequestMarshaller(protocolFactory))); @@ -441,9 +457,8 @@ public OperationWithContextParamResponse operationWithContextParam( .withOperationName("OperationWithContextParam").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(operationWithContextParamRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithContextParam", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithContextParam")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OperationWithContextParamRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -491,9 +506,8 @@ public OperationWithCustomMemberResponse operationWithCustomMember( .withOperationName("OperationWithCustomMember").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(operationWithCustomMemberRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithCustomMember", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithCustomMember")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OperationWithCustomMemberRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -539,15 +553,11 @@ public OperationWithCustomizedOperationContextParamResponse operationWithCustomi return clientHandler .execute(new ClientExecutionParams() .withOperationName("OperationWithCustomizedOperationContextParam") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) + .withProtocolMetadata(protocolMetadata).withResponseHandler(responseHandler) + .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withInput(operationWithCustomizedOperationContextParamRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithCustomizedOperationContextParam", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithCustomizedOperationContextParam")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OperationWithCustomizedOperationContextParamRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -592,16 +602,12 @@ public OperationWithMapOperationContextParamResponse operationWithMapOperationCo return clientHandler .execute(new ClientExecutionParams() - .withOperationName("OperationWithMapOperationContextParam") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) + .withOperationName("OperationWithMapOperationContextParam").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) - .withInput(operationWithMapOperationContextParamRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithMapOperationContextParam", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithMapOperationContextParam")) + .withInput(operationWithMapOperationContextParamRequest).withMetricCollector(apiCallMetricCollector) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OperationWithMapOperationContextParamRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -648,9 +654,8 @@ public OperationWithNoneAuthTypeResponse operationWithNoneAuthType( .withOperationName("OperationWithNoneAuthType").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(operationWithNoneAuthTypeRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithNoneAuthType", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithNoneAuthType")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OperationWithNoneAuthTypeRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -695,16 +700,11 @@ public OperationWithOperationContextParamResponse operationWithOperationContextP return clientHandler .execute(new ClientExecutionParams() - .withOperationName("OperationWithOperationContextParam") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withInput(operationWithOperationContextParamRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithOperationContextParam", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithOperationContextParam")) + .withOperationName("OperationWithOperationContextParam").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withInput(operationWithOperationContextParamRequest) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OperationWithOperationContextParamRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -755,9 +755,8 @@ public OperationWithRequestCompressionResponse operationWithRequestCompression( .withRequestConfiguration(clientConfiguration) .withInput(operationWithRequestCompressionRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithRequestCompression", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithRequestCompression")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.REQUEST_COMPRESSION, RequestCompression.builder().encodings("gzip").isStreaming(false).build()) .withMarshaller(new OperationWithRequestCompressionRequestMarshaller(protocolFactory))); @@ -803,16 +802,11 @@ public OperationWithStaticContextParamsResponse operationWithStaticContextParams return clientHandler .execute(new ClientExecutionParams() - .withOperationName("OperationWithStaticContextParams") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withInput(operationWithStaticContextParamsRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithStaticContextParams", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithStaticContextParams")) + .withOperationName("OperationWithStaticContextParams").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withInput(operationWithStaticContextParamsRequest) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OperationWithStaticContextParamsRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -889,8 +883,8 @@ public ReturnT putOperationWithChecksum(PutOperationWithChecksumReques .withRequestConfiguration(clientConfiguration) .withInput(putOperationWithChecksumRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("PutOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("PutOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum @@ -969,8 +963,8 @@ public StreamingInputOperationResponse streamingInputOperation(StreamingInputOpe .withRequestConfiguration(clientConfiguration) .withInput(streamingInputOperationRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingInputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withRequestBody(requestBody) .withMarshaller( StreamingRequestMarshaller.builder() @@ -1028,10 +1022,8 @@ public ReturnT streamingOutputOperation(StreamingOutputOperationReques .withOperationName("StreamingOutputOperation").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(streamingOutputOperationRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingOutputOperation")) - .withResponseTransformer(responseTransformer) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withResponseTransformer(responseTransformer) .withMarshaller(new StreamingOutputOperationRequestMarshaller(protocolFactory)), responseTransformer); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1071,23 +1063,25 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); QueryAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) .map(p -> Validate.isInstanceOf(QueryAuthSchemeProvider.class, p, "Expected an instance of QueryAuthSchemeProvider")).orElse(null); QueryAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(QueryAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(QueryAuthSchemeProvider.class, + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of QueryAuthSchemeProvider"); QueryAuthSchemeParams.Builder paramsBuilder = QueryAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); QueryEndpointProvider provider = (QueryEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -1118,14 +1112,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-rpcv2-async-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-rpcv2-async-client-class.java index dddf865558a1..3f268f7698ac 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-rpcv2-async-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-rpcv2-async-client-class.java @@ -13,7 +13,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -35,6 +35,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -122,6 +123,25 @@ final class DefaultSmithyRpcV2ProtocolAsyncClient implements SmithyRpcV2Protocol private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultSmithyRpcV2ProtocolAsyncClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsAsyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -192,8 +212,8 @@ public CompletableFuture emptyInputOutput(EmptyInputOu .withMarshaller(new EmptyInputOutputRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("EmptyInputOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("EmptyInputOutput")).withInput(emptyInputOutputRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(emptyInputOutputRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -266,8 +286,8 @@ public CompletableFuture float16(Float16Request float16Request) .withProtocolMetadata(protocolMetadata).withMarshaller(new Float16RequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("Float16", clientConfiguration)) - .withEndpointResolver(endpointResolver("Float16")).withInput(float16Request)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(float16Request)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -342,8 +362,8 @@ public CompletableFuture fractionalSeconds(Fractional .withMarshaller(new FractionalSecondsRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("FractionalSeconds", clientConfiguration)) - .withEndpointResolver(endpointResolver("FractionalSeconds")).withInput(fractionalSecondsRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(fractionalSecondsRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -420,8 +440,8 @@ public CompletableFuture greetingWithErrors(Greeting .withMarshaller(new GreetingWithErrorsRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GreetingWithErrors", clientConfiguration)) - .withEndpointResolver(endpointResolver("GreetingWithErrors")).withInput(greetingWithErrorsRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(greetingWithErrorsRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -495,8 +515,8 @@ public CompletableFuture noInputOutput(NoInputOutputReque .withMarshaller(new NoInputOutputRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("NoInputOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("NoInputOutput")).withInput(noInputOutputRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(noInputOutputRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -573,8 +593,8 @@ public CompletableFuture operationWithDefaults( .withMarshaller(new OperationWithDefaultsRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithDefaults", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithDefaults")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithDefaultsRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -651,8 +671,8 @@ public CompletableFuture optionalInputOutput( .withMarshaller(new OptionalInputOutputRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OptionalInputOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("OptionalInputOutput")).withInput(optionalInputOutputRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(optionalInputOutputRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -727,8 +747,8 @@ public CompletableFuture recursiveShapes(RecursiveShape .withMarshaller(new RecursiveShapesRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("RecursiveShapes", clientConfiguration)) - .withEndpointResolver(endpointResolver("RecursiveShapes")).withInput(recursiveShapesRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(recursiveShapesRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -804,8 +824,8 @@ public CompletableFuture rpcV2CborDenseMaps(RpcV2Cbo .withMarshaller(new RpcV2CborDenseMapsRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("RpcV2CborDenseMaps", clientConfiguration)) - .withEndpointResolver(endpointResolver("RpcV2CborDenseMaps")).withInput(rpcV2CborDenseMapsRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(rpcV2CborDenseMapsRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -880,8 +900,8 @@ public CompletableFuture rpcV2CborLists(RpcV2CborListsRe .withMarshaller(new RpcV2CborListsRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("RpcV2CborLists", clientConfiguration)) - .withEndpointResolver(endpointResolver("RpcV2CborLists")).withInput(rpcV2CborListsRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(rpcV2CborListsRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -958,8 +978,8 @@ public CompletableFuture rpcV2CborSparseMaps( .withMarshaller(new RpcV2CborSparseMapsRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("RpcV2CborSparseMaps", clientConfiguration)) - .withEndpointResolver(endpointResolver("RpcV2CborSparseMaps")).withInput(rpcV2CborSparseMapsRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(rpcV2CborSparseMapsRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -1035,8 +1055,8 @@ public CompletableFuture simpleScalarProperties( .withMarshaller(new SimpleScalarPropertiesRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("SimpleScalarProperties", clientConfiguration)) - .withEndpointResolver(endpointResolver("SimpleScalarProperties")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(simpleScalarPropertiesRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1113,8 +1133,8 @@ public CompletableFuture sparseNullsOperation( .withMarshaller(new SparseNullsOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("SparseNullsOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("SparseNullsOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(sparseNullsOperationRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1158,8 +1178,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); SmithyRpcV2ProtocolAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -1167,16 +1188,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of SmithyRpcV2ProtocolAuthSchemeProvider")).orElse(null); SmithyRpcV2ProtocolAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(SmithyRpcV2ProtocolAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of SmithyRpcV2ProtocolAuthSchemeProvider"); SmithyRpcV2ProtocolAuthSchemeParams.Builder paramsBuilder = SmithyRpcV2ProtocolAuthSchemeParams.builder().operation( operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); SmithyRpcV2ProtocolEndpointProvider provider = (SmithyRpcV2ProtocolEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -1208,14 +1230,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-rpcv2-sync.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-rpcv2-sync.java index 8beb520ef006..6f053d91f5d2 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-rpcv2-sync.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-rpcv2-sync.java @@ -8,7 +8,7 @@ import java.util.function.Function; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -30,6 +30,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -117,6 +118,25 @@ final class DefaultSmithyRpcV2ProtocolClient implements SmithyRpcV2ProtocolClien private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultSmithyRpcV2ProtocolClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsSyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -183,8 +203,8 @@ public EmptyInputOutputResponse emptyInputOutput(EmptyInputOutputRequest emptyIn .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(emptyInputOutputRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("EmptyInputOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("EmptyInputOutput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new EmptyInputOutputRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -248,8 +268,8 @@ public Float16Response float16(Float16Request float16Request) throws AwsServiceE .withOperationName("Float16").withProtocolMetadata(protocolMetadata).withResponseHandler(responseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withInput(float16Request).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("Float16", clientConfiguration)) - .withEndpointResolver(endpointResolver("Float16")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new Float16RequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -315,8 +335,8 @@ public FractionalSecondsResponse fractionalSeconds(FractionalSecondsRequest frac .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(fractionalSecondsRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("FractionalSeconds", clientConfiguration)) - .withEndpointResolver(endpointResolver("FractionalSeconds")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new FractionalSecondsRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -385,8 +405,8 @@ public GreetingWithErrorsResponse greetingWithErrors(GreetingWithErrorsRequest g .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(greetingWithErrorsRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GreetingWithErrors", clientConfiguration)) - .withEndpointResolver(endpointResolver("GreetingWithErrors")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new GreetingWithErrorsRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -451,8 +471,8 @@ public NoInputOutputResponse noInputOutput(NoInputOutputRequest noInputOutputReq .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(noInputOutputRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("NoInputOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("NoInputOutput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new NoInputOutputRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -519,8 +539,8 @@ public OperationWithDefaultsResponse operationWithDefaults(OperationWithDefaults .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(operationWithDefaultsRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithDefaults", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithDefaults")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OperationWithDefaultsRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -586,8 +606,8 @@ public OptionalInputOutputResponse optionalInputOutput(OptionalInputOutputReques .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(optionalInputOutputRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OptionalInputOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("OptionalInputOutput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OptionalInputOutputRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -653,8 +673,8 @@ public RecursiveShapesResponse recursiveShapes(RecursiveShapesRequest recursiveS .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(recursiveShapesRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("RecursiveShapes", clientConfiguration)) - .withEndpointResolver(endpointResolver("RecursiveShapes")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new RecursiveShapesRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -721,8 +741,8 @@ public RpcV2CborDenseMapsResponse rpcV2CborDenseMaps(RpcV2CborDenseMapsRequest r .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(rpcV2CborDenseMapsRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("RpcV2CborDenseMaps", clientConfiguration)) - .withEndpointResolver(endpointResolver("RpcV2CborDenseMaps")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new RpcV2CborDenseMapsRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -788,8 +808,8 @@ public RpcV2CborListsResponse rpcV2CborLists(RpcV2CborListsRequest rpcV2CborList .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(rpcV2CborListsRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("RpcV2CborLists", clientConfiguration)) - .withEndpointResolver(endpointResolver("RpcV2CborLists")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new RpcV2CborListsRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -856,8 +876,8 @@ public RpcV2CborSparseMapsResponse rpcV2CborSparseMaps(RpcV2CborSparseMapsReques .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(rpcV2CborSparseMapsRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("RpcV2CborSparseMaps", clientConfiguration)) - .withEndpointResolver(endpointResolver("RpcV2CborSparseMaps")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new RpcV2CborSparseMapsRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -924,8 +944,8 @@ public SimpleScalarPropertiesResponse simpleScalarProperties(SimpleScalarPropert .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(simpleScalarPropertiesRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("SimpleScalarProperties", clientConfiguration)) - .withEndpointResolver(endpointResolver("SimpleScalarProperties")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new SimpleScalarPropertiesRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -991,8 +1011,8 @@ public SparseNullsOperationResponse sparseNullsOperation(SparseNullsOperationReq .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(sparseNullsOperationRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("SparseNullsOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("SparseNullsOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new SparseNullsOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -1019,8 +1039,9 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); SmithyRpcV2ProtocolAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) @@ -1028,16 +1049,17 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri "Expected an instance of SmithyRpcV2ProtocolAuthSchemeProvider")).orElse(null); SmithyRpcV2ProtocolAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate.isInstanceOf(SmithyRpcV2ProtocolAuthSchemeProvider.class, - clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of SmithyRpcV2ProtocolAuthSchemeProvider"); SmithyRpcV2ProtocolAuthSchemeParams.Builder paramsBuilder = SmithyRpcV2ProtocolAuthSchemeParams.builder().operation( operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); SmithyRpcV2ProtocolEndpointProvider provider = (SmithyRpcV2ProtocolEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -1069,14 +1091,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private HttpResponseHandler createErrorResponseHandler(BaseAwsJsonProtocolFactory protocolFactory, JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { return protocolFactory.createErrorResponseHandler(operationMetadata, exceptionMetadataMapper); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-unsigned-payload-trait-async-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-unsigned-payload-trait-async-client-class.java index 12bceafed07b..9cea417841cb 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-unsigned-payload-trait-async-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-unsigned-payload-trait-async-client-class.java @@ -14,7 +14,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -37,6 +37,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -121,6 +122,25 @@ final class DefaultDatabaseAsyncClient implements DatabaseAsyncClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultDatabaseAsyncClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsAsyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -188,8 +208,8 @@ public CompletableFuture deleteRow(DeleteRowRequest deleteRow .withMarshaller(new DeleteRowRequestMarshaller(protocolFactory)).withResponseHandler(responseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("DeleteRow", clientConfiguration)) - .withEndpointResolver(endpointResolver("DeleteRow")).withInput(deleteRowRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(deleteRowRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -260,8 +280,8 @@ public CompletableFuture getRow(GetRowRequest getRowRequest) { .withProtocolMetadata(protocolMetadata).withMarshaller(new GetRowRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetRow", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetRow")).withInput(getRowRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(getRowRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -340,9 +360,8 @@ public CompletableFuture opWithSigv .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("opWithSigv4AndSigv4aUnSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4AndSigv4aUnSignedPayload")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(opWithSigv4AndSigv4AUnSignedPayloadRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -417,8 +436,8 @@ public CompletableFuture opWithSigv4SignedPayl .withMarshaller(new OpWithSigv4SignedPayloadRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("opWithSigv4SignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4SignedPayload")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(opWithSigv4SignedPayloadRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -493,8 +512,8 @@ public CompletableFuture opWithSigv4UnSigned .withMarshaller(new OpWithSigv4UnSignedPayloadRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("opWithSigv4UnSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4UnSignedPayload")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(opWithSigv4UnSignedPayloadRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -585,9 +604,8 @@ public CompletableFuture opWithS .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("opWithSigv4UnSignedPayloadAndStreaming", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4UnSignedPayloadAndStreaming")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withAsyncRequestBody(requestBody).withInput(opWithSigv4UnSignedPayloadAndStreamingRequest)); CompletableFuture whenCompleted = executeFuture .whenComplete((r, e) -> { @@ -663,8 +681,8 @@ public CompletableFuture opWithSigv4aSignedPa .withMarshaller(new OpWithSigv4ASignedPayloadRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("opWithSigv4aSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4aSignedPayload")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(opWithSigv4ASignedPayloadRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -742,8 +760,8 @@ public CompletableFuture opWithSigv4aUnSign .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("opWithSigv4aUnSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4aUnSignedPayload")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(opWithSigv4AUnSignedPayloadRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -823,9 +841,8 @@ public CompletableFuture opsWithSigv .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("opsWithSigv4andSigv4aSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opsWithSigv4andSigv4aSignedPayload")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(opsWithSigv4AndSigv4ASignedPayloadRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -897,8 +914,8 @@ public CompletableFuture putRow(PutRowRequest putRowRequest) { .withProtocolMetadata(protocolMetadata).withMarshaller(new PutRowRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("PutRow", clientConfiguration)) - .withEndpointResolver(endpointResolver("PutRow")).withInput(putRowRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(putRowRequest)); CompletableFuture whenCompleted = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); }); @@ -977,9 +994,8 @@ public CompletableFuture secon .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("secondOpsWithSigv4andSigv4aSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("secondOpsWithSigv4andSigv4aSignedPayload")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(secondOpsWithSigv4AndSigv4ASignedPayloadRequest)); CompletableFuture whenCompleted = executeFuture .whenComplete((r, e) -> { @@ -1023,19 +1039,20 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); DatabaseAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) .map(p -> Validate.isInstanceOf(DatabaseAuthSchemeProvider.class, p, "Expected an instance of DatabaseAuthSchemeProvider")).orElse(null); DatabaseAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(DatabaseAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(DatabaseAuthSchemeProvider.class, executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of DatabaseAuthSchemeProvider"); DatabaseAuthSchemeParams.Builder paramsBuilder = DatabaseAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); - Set sigv4aRegionSet = clientConfiguration.option(AwsClientOption.AWS_SIGV4A_SIGNING_REGION_SET); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); + Set sigv4aRegionSet = executionAttributes.getAttribute(AwsExecutionAttribute.AWS_SIGV4A_SIGNING_REGION_SET); if (!CollectionUtils.isNullOrEmpty(sigv4aRegionSet)) { paramsBuilder.regionSet(RegionSet.create(sigv4aRegionSet)); } @@ -1043,7 +1060,8 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); DatabaseEndpointProvider provider = (DatabaseEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -1083,14 +1101,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-unsigned-payload-trait-sync-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-unsigned-payload-trait-sync-client-class.java index 299f08752a84..c7e03150bbe1 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-unsigned-payload-trait-sync-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-unsigned-payload-trait-sync-client-class.java @@ -9,7 +9,7 @@ import java.util.function.Function; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -31,6 +31,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.metrics.CoreMetric; import software.amazon.awssdk.core.retry.RetryMode; @@ -116,6 +117,25 @@ final class DefaultDatabaseClient implements DatabaseClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultDatabaseClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsSyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -178,8 +198,7 @@ public DeleteRowResponse deleteRow(DeleteRowRequest deleteRowRequest) throws Inv .withOperationName("DeleteRow").withProtocolMetadata(protocolMetadata).withResponseHandler(responseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withInput(deleteRowRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("DeleteRow", clientConfiguration)) - .withEndpointResolver(endpointResolver("DeleteRow")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver).withEndpointResolver(endpointResolverInstance) .withMarshaller(new DeleteRowRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -241,8 +260,7 @@ public GetRowResponse getRow(GetRowRequest getRowRequest) throws InvalidInputExc .withProtocolMetadata(protocolMetadata).withResponseHandler(responseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withInput(getRowRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetRow", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetRow")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver).withEndpointResolver(endpointResolverInstance) .withMarshaller(new GetRowRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -304,8 +322,7 @@ public PutRowResponse putRow(PutRowRequest putRowRequest) throws InvalidInputExc .withProtocolMetadata(protocolMetadata).withResponseHandler(responseHandler) .withErrorResponseHandler(errorResponseHandler).withRequestConfiguration(clientConfiguration) .withInput(putRowRequest).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("PutRow", clientConfiguration)) - .withEndpointResolver(endpointResolver("PutRow")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver).withEndpointResolver(endpointResolverInstance) .withMarshaller(new PutRowRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -368,16 +385,11 @@ public OpWithSigv4AndSigv4AUnSignedPayloadResponse opWithSigv4AndSigv4aUnSignedP return clientHandler .execute(new ClientExecutionParams() - .withOperationName("opWithSigv4AndSigv4aUnSignedPayload") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withInput(opWithSigv4AndSigv4AUnSignedPayloadRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("opWithSigv4AndSigv4aUnSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4AndSigv4aUnSignedPayload")) + .withOperationName("opWithSigv4AndSigv4aUnSignedPayload").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withInput(opWithSigv4AndSigv4AUnSignedPayloadRequest) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OpWithSigv4AndSigv4AUnSignedPayloadRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -442,9 +454,8 @@ public OpWithSigv4SignedPayloadResponse opWithSigv4SignedPayload( .withOperationName("opWithSigv4SignedPayload").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(opWithSigv4SignedPayloadRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("opWithSigv4SignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4SignedPayload")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OpWithSigv4SignedPayloadRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -509,9 +520,8 @@ public OpWithSigv4UnSignedPayloadResponse opWithSigv4UnSignedPayload( .withOperationName("opWithSigv4UnSignedPayload").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(opWithSigv4UnSignedPayloadRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("opWithSigv4UnSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4UnSignedPayload")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OpWithSigv4UnSignedPayloadRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -592,9 +602,8 @@ public OpWithSigv4UnSignedPayloadAndStreamingResponse opWithSigv4UnSignedPayload .withRequestConfiguration(clientConfiguration) .withInput(opWithSigv4UnSignedPayloadAndStreamingRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("opWithSigv4UnSignedPayloadAndStreaming", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4UnSignedPayloadAndStreaming")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withRequestBody(requestBody) .withMarshaller( StreamingRequestMarshaller @@ -665,9 +674,8 @@ public OpWithSigv4ASignedPayloadResponse opWithSigv4aSignedPayload( .withOperationName("opWithSigv4aSignedPayload").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(opWithSigv4ASignedPayloadRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("opWithSigv4aSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4aSignedPayload")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OpWithSigv4ASignedPayloadRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -729,15 +737,11 @@ public OpWithSigv4AUnSignedPayloadResponse opWithSigv4aUnSignedPayload( return clientHandler .execute(new ClientExecutionParams() - .withOperationName("opWithSigv4aUnSignedPayload") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withInput(opWithSigv4AUnSignedPayloadRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("opWithSigv4aUnSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opWithSigv4aUnSignedPayload")) + .withOperationName("opWithSigv4aUnSignedPayload").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withInput(opWithSigv4AUnSignedPayloadRequest) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OpWithSigv4AUnSignedPayloadRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -800,16 +804,11 @@ public OpsWithSigv4AndSigv4ASignedPayloadResponse opsWithSigv4andSigv4aSignedPay return clientHandler .execute(new ClientExecutionParams() - .withOperationName("opsWithSigv4andSigv4aSignedPayload") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) - .withRequestConfiguration(clientConfiguration) - .withInput(opsWithSigv4AndSigv4ASignedPayloadRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("opsWithSigv4andSigv4aSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("opsWithSigv4andSigv4aSignedPayload")) + .withOperationName("opsWithSigv4andSigv4aSignedPayload").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) + .withRequestConfiguration(clientConfiguration).withInput(opsWithSigv4AndSigv4ASignedPayloadRequest) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OpsWithSigv4AndSigv4ASignedPayloadRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -872,16 +871,12 @@ public SecondOpsWithSigv4AndSigv4ASignedPayloadResponse secondOpsWithSigv4andSig return clientHandler .execute(new ClientExecutionParams() - .withOperationName("secondOpsWithSigv4andSigv4aSignedPayload") - .withProtocolMetadata(protocolMetadata) - .withResponseHandler(responseHandler) - .withErrorResponseHandler(errorResponseHandler) + .withOperationName("secondOpsWithSigv4andSigv4aSignedPayload").withProtocolMetadata(protocolMetadata) + .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withInput(secondOpsWithSigv4AndSigv4ASignedPayloadRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("secondOpsWithSigv4andSigv4aSignedPayload", clientConfiguration)) - .withEndpointResolver(endpointResolver("secondOpsWithSigv4andSigv4aSignedPayload")) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new SecondOpsWithSigv4AndSigv4ASignedPayloadRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -908,19 +903,20 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); DatabaseAuthSchemeProvider requestAuthSchemeProvider = request .overrideConfiguration() .flatMap(c -> c.authSchemeProvider()) .map(p -> Validate.isInstanceOf(DatabaseAuthSchemeProvider.class, p, "Expected an instance of DatabaseAuthSchemeProvider")).orElse(null); DatabaseAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(DatabaseAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(DatabaseAuthSchemeProvider.class, + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of DatabaseAuthSchemeProvider"); DatabaseAuthSchemeParams.Builder paramsBuilder = DatabaseAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); - Set sigv4aRegionSet = clientConfiguration.option(AwsClientOption.AWS_SIGV4A_SIGNING_REGION_SET); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); + Set sigv4aRegionSet = executionAttributes.getAttribute(AwsExecutionAttribute.AWS_SIGV4A_SIGNING_REGION_SET); if (!CollectionUtils.isNullOrEmpty(sigv4aRegionSet)) { paramsBuilder.regionSet(RegionSet.create(sigv4aRegionSet)); } @@ -928,7 +924,8 @@ private List resolveAuthSchemeOptions(SdkRequest request, Stri return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); DatabaseEndpointProvider provider = (DatabaseEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -968,14 +965,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private HttpResponseHandler createErrorResponseHandler(BaseAwsJsonProtocolFactory protocolFactory, JsonOperationMetadata operationMetadata, Function> exceptionMetadataMapper) { return protocolFactory.createErrorResponseHandler(operationMetadata, exceptionMetadataMapper); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-xml-async-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-xml-async-client-class.java index 55a64c1b706a..aff36a86530a 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-xml-async-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-xml-async-client-class.java @@ -13,7 +13,7 @@ import org.slf4j.LoggerFactory; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsAsyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -46,6 +46,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.interceptor.trait.HttpChecksum; import software.amazon.awssdk.core.interceptor.trait.HttpChecksumRequired; @@ -130,6 +131,25 @@ final class DefaultXmlAsyncClient implements XmlAsyncClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + private final Executor executor; protected DefaultXmlAsyncClient(SdkClientConfiguration clientConfiguration) { @@ -188,8 +208,8 @@ public CompletableFuture aPostOperation(APostOperationRe .withMarshaller(new APostOperationRequestMarshaller(protocolFactory)) .withCombinedResponseHandler(responseHandler).hostPrefixExpression(resolvedHostExpression) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")).withInput(aPostOperationRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(aPostOperationRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -249,8 +269,8 @@ public CompletableFuture aPostOperationWithOut .withProtocolMetadata(protocolMetadata) .withMarshaller(new APostOperationWithOutputRequestMarshaller(protocolFactory)) .withCombinedResponseHandler(responseHandler).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperationWithOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperationWithOutput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(aPostOperationWithOutputRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -308,8 +328,8 @@ public CompletableFuture bearerAuthOperation( .withMarshaller(new BearerAuthOperationRequestMarshaller(protocolFactory)) .withCombinedResponseHandler(responseHandler).credentialType(CredentialType.TOKEN) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("BearerAuthOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("BearerAuthOperation")).withInput(bearerAuthOperationRequest)); + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withInput(bearerAuthOperationRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -383,8 +403,8 @@ public CompletableFuture eventStreamOperation(EventStreamOperationRequest .withMarshaller(new EventStreamOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("EventStreamOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("EventStreamOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(eventStreamOperationRequest), restAsyncResponseTransformer); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -450,8 +470,8 @@ public CompletableFuture getOperationWithCheck .withMarshaller(new GetOperationWithChecksumRequestMarshaller(protocolFactory)) .withCombinedResponseHandler(responseHandler) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum.builder().requestChecksumRequired(true).isRequestStreaming(false) @@ -516,9 +536,8 @@ public CompletableFuture operationWithChe .withMarshaller(new OperationWithChecksumRequiredRequestMarshaller(protocolFactory)) .withCombinedResponseHandler(responseHandler) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithChecksumRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithChecksumRequired")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.HTTP_CHECKSUM_REQUIRED, HttpChecksumRequired.create()).withInput(operationWithChecksumRequiredRequest)); CompletableFuture whenCompleteFuture = null; @@ -576,8 +595,8 @@ public CompletableFuture operationWithNoneAut .withProtocolMetadata(protocolMetadata) .withMarshaller(new OperationWithNoneAuthTypeRequestMarshaller(protocolFactory)) .withCombinedResponseHandler(responseHandler).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithNoneAuthType", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithNoneAuthType")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withInput(operationWithNoneAuthTypeRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -637,9 +656,8 @@ public CompletableFuture operationWithR .withMarshaller(new OperationWithRequestCompressionRequestMarshaller(protocolFactory)) .withCombinedResponseHandler(responseHandler) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithRequestCompression", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithRequestCompression")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.REQUEST_COMPRESSION, RequestCompression.builder().encodings("gzip").isStreaming(false).build()) .withInput(operationWithRequestCompressionRequest)); @@ -728,8 +746,8 @@ public CompletableFuture putOperationWithChecksum( .withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("PutOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("PutOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum @@ -818,8 +836,8 @@ public CompletableFuture streamingInputOperatio .delegateMarshaller(new StreamingInputOperationRequestMarshaller(protocolFactory)) .asyncRequestBody(requestBody).build()).withCombinedResponseHandler(responseHandler) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingInputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOperation")).withAsyncRequestBody(requestBody) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withAsyncRequestBody(requestBody) .withInput(streamingInputOperationRequest)); CompletableFuture whenCompleteFuture = null; whenCompleteFuture = executeFuture.whenComplete((r, e) -> { @@ -887,8 +905,8 @@ public CompletableFuture streamingOutputOperation( .withMarshaller(new StreamingOutputOperationRequestMarshaller(protocolFactory)) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingOutputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withAsyncResponseTransformer(asyncResponseTransformer).withInput(streamingOutputOperationRequest), asyncResponseTransformer); CompletableFuture whenCompleteFuture = null; @@ -946,21 +964,23 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, + ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); XmlAuthSchemeProvider requestAuthSchemeProvider = request.overrideConfiguration().flatMap(c -> c.authSchemeProvider()) .map(p -> Validate.isInstanceOf(XmlAuthSchemeProvider.class, p, "Expected an instance of XmlAuthSchemeProvider")) .orElse(null); XmlAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(XmlAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(XmlAuthSchemeProvider.class, executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of XmlAuthSchemeProvider"); XmlAuthSchemeParams.Builder paramsBuilder = XmlAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); XmlEndpointProvider provider = (XmlEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -991,14 +1011,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-xml-client-class.java b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-xml-client-class.java index 4e48a80ee74f..1af1d40a7da7 100644 --- a/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-xml-client-class.java +++ b/codegen/src/test/resources/software/amazon/awssdk/codegen/poet/client/test-xml-client-class.java @@ -7,7 +7,7 @@ import java.util.function.Consumer; import software.amazon.awssdk.annotations.Generated; import software.amazon.awssdk.annotations.SdkInternalApi; -import software.amazon.awssdk.awscore.client.config.AwsClientOption; +import software.amazon.awssdk.awscore.AwsExecutionAttribute; import software.amazon.awssdk.awscore.client.handler.AwsSyncClientHandler; import software.amazon.awssdk.awscore.endpoints.AwsEndpointAttribute; import software.amazon.awssdk.awscore.endpoints.AwsEndpointProviderUtils; @@ -32,6 +32,7 @@ import software.amazon.awssdk.core.exception.SdkClientException; import software.amazon.awssdk.core.http.HttpResponseHandler; import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.SdkExecutionAttribute; import software.amazon.awssdk.core.interceptor.SdkInternalExecutionAttribute; import software.amazon.awssdk.core.interceptor.trait.HttpChecksum; import software.amazon.awssdk.core.interceptor.trait.HttpChecksumRequired; @@ -112,6 +113,25 @@ final class DefaultXmlClient implements XmlClient { private final SdkClientConfiguration clientConfiguration; + private final AuthSchemeOptionsResolver authSchemeOptionsResolver = new AuthSchemeOptionsResolver() { + @Override + public List resolve(SdkRequest request) { + throw new UnsupportedOperationException("Use resolve(SdkRequest, ExecutionAttributes) instead"); + } + + @Override + public List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveAuthSchemeOptions(request, executionAttributes); + } + }; + + private final EndpointResolver endpointResolverInstance = new EndpointResolver() { + @Override + public Endpoint resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolveEndpoint(request, executionAttributes); + } + }; + protected DefaultXmlClient(SdkClientConfiguration clientConfiguration) { this.clientHandler = new AwsSyncClientHandler(clientConfiguration); this.clientConfiguration = clientConfiguration.toBuilder().option(SdkClientOption.SDK_CLIENT, this) @@ -160,9 +180,8 @@ public APostOperationResponse aPostOperation(APostOperationRequest aPostOperatio .withOperationName("APostOperation").withProtocolMetadata(protocolMetadata) .withCombinedResponseHandler(responseHandler).withMetricCollector(apiCallMetricCollector) .hostPrefixExpression(resolvedHostExpression).withRequestConfiguration(clientConfiguration) - .withInput(aPostOperationRequest) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperation")) + .withInput(aPostOperationRequest).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new APostOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -212,8 +231,8 @@ public APostOperationWithOutputResponse aPostOperationWithOutput( .withOperationName("APostOperationWithOutput").withProtocolMetadata(protocolMetadata) .withCombinedResponseHandler(responseHandler).withMetricCollector(apiCallMetricCollector) .withRequestConfiguration(clientConfiguration).withInput(aPostOperationWithOutputRequest) - .withAuthSchemeOptionsResolver(authSchemeResolver("APostOperationWithOutput", clientConfiguration)) - .withEndpointResolver(endpointResolver("APostOperationWithOutput")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new APostOperationWithOutputRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -257,9 +276,8 @@ public BearerAuthOperationResponse bearerAuthOperation(BearerAuthOperationReques .withOperationName("BearerAuthOperation").withProtocolMetadata(protocolMetadata) .withCombinedResponseHandler(responseHandler).withMetricCollector(apiCallMetricCollector) .credentialType(CredentialType.TOKEN).withRequestConfiguration(clientConfiguration) - .withInput(bearerAuthOperationRequest) - .withAuthSchemeOptionsResolver(authSchemeResolver("BearerAuthOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("BearerAuthOperation")) + .withInput(bearerAuthOperationRequest).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new BearerAuthOperationRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -308,8 +326,8 @@ public GetOperationWithChecksumResponse getOperationWithChecksum( .withMetricCollector(apiCallMetricCollector) .withRequestConfiguration(clientConfiguration) .withInput(getOperationWithChecksumRequest) - .withAuthSchemeOptionsResolver(authSchemeResolver("GetOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("GetOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum.builder().requestChecksumRequired(true).isRequestStreaming(false) @@ -363,9 +381,8 @@ public OperationWithChecksumRequiredResponse operationWithChecksumRequired( .withMetricCollector(apiCallMetricCollector) .withRequestConfiguration(clientConfiguration) .withInput(operationWithChecksumRequiredRequest) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithChecksumRequired", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithChecksumRequired")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.HTTP_CHECKSUM_REQUIRED, HttpChecksumRequired.create()) .withMarshaller(new OperationWithChecksumRequiredRequestMarshaller(protocolFactory))); @@ -413,8 +430,8 @@ public OperationWithNoneAuthTypeResponse operationWithNoneAuthType( .withOperationName("OperationWithNoneAuthType").withProtocolMetadata(protocolMetadata) .withCombinedResponseHandler(responseHandler).withMetricCollector(apiCallMetricCollector) .withRequestConfiguration(clientConfiguration).withInput(operationWithNoneAuthTypeRequest) - .withAuthSchemeOptionsResolver(authSchemeResolver("OperationWithNoneAuthType", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithNoneAuthType")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withMarshaller(new OperationWithNoneAuthTypeRequestMarshaller(protocolFactory))); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -463,9 +480,8 @@ public OperationWithRequestCompressionResponse operationWithRequestCompression( .withMetricCollector(apiCallMetricCollector) .withRequestConfiguration(clientConfiguration) .withInput(operationWithRequestCompressionRequest) - .withAuthSchemeOptionsResolver( - authSchemeResolver("OperationWithRequestCompression", clientConfiguration)) - .withEndpointResolver(endpointResolver("OperationWithRequestCompression")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute(SdkInternalExecutionAttribute.REQUEST_COMPRESSION, RequestCompression.builder().encodings("gzip").isStreaming(false).build()) .withMarshaller(new OperationWithRequestCompressionRequestMarshaller(protocolFactory))); @@ -544,8 +560,8 @@ public ReturnT putOperationWithChecksum(PutOperationWithChecksumReques .withRequestConfiguration(clientConfiguration) .withInput(putOperationWithChecksumRequest) .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("PutOperationWithChecksum", clientConfiguration)) - .withEndpointResolver(endpointResolver("PutOperationWithChecksum")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .putExecutionAttribute( SdkInternalExecutionAttribute.HTTP_CHECKSUM, HttpChecksum @@ -622,8 +638,8 @@ public StreamingInputOperationResponse streamingInputOperation(StreamingInputOpe .withMetricCollector(apiCallMetricCollector) .withRequestConfiguration(clientConfiguration) .withInput(streamingInputOperationRequest) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingInputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingInputOperation")) + .withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance) .withRequestBody(requestBody) .withMarshaller( StreamingRequestMarshaller.builder() @@ -681,10 +697,8 @@ public ReturnT streamingOutputOperation(StreamingOutputOperationReques .withOperationName("StreamingOutputOperation").withProtocolMetadata(protocolMetadata) .withResponseHandler(responseHandler).withErrorResponseHandler(errorResponseHandler) .withRequestConfiguration(clientConfiguration).withInput(streamingOutputOperationRequest) - .withMetricCollector(apiCallMetricCollector) - .withAuthSchemeOptionsResolver(authSchemeResolver("StreamingOutputOperation", clientConfiguration)) - .withEndpointResolver(endpointResolver("StreamingOutputOperation")) - .withResponseTransformer(responseTransformer) + .withMetricCollector(apiCallMetricCollector).withAuthSchemeOptionsResolver(authSchemeOptionsResolver) + .withEndpointResolver(endpointResolverInstance).withResponseTransformer(responseTransformer) .withMarshaller(new StreamingOutputOperationRequestMarshaller(protocolFactory)), responseTransformer); } finally { metricPublishers.forEach(p -> p.publish(apiCallMetricCollector.collect())); @@ -711,21 +725,23 @@ private static List resolveMetricPublishers(SdkClientConfigurat return publishers; } - private List resolveAuthSchemeOptions(SdkRequest request, String operationName, - SdkClientConfiguration clientConfiguration) { + private List resolveAuthSchemeOptions(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); XmlAuthSchemeProvider requestAuthSchemeProvider = request.overrideConfiguration().flatMap(c -> c.authSchemeProvider()) .map(p -> Validate.isInstanceOf(XmlAuthSchemeProvider.class, p, "Expected an instance of XmlAuthSchemeProvider")) .orElse(null); XmlAuthSchemeProvider authSchemeProvider = requestAuthSchemeProvider != null ? requestAuthSchemeProvider : Validate - .isInstanceOf(XmlAuthSchemeProvider.class, clientConfiguration.option(SdkClientOption.AUTH_SCHEME_PROVIDER), + .isInstanceOf(XmlAuthSchemeProvider.class, + executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_RESOLVER), "Expected an instance of XmlAuthSchemeProvider"); XmlAuthSchemeParams.Builder paramsBuilder = XmlAuthSchemeParams.builder().operation(operationName); - paramsBuilder.region(clientConfiguration.option(AwsClientOption.AWS_REGION)); + paramsBuilder.region(executionAttributes.getAttribute(AwsExecutionAttribute.AWS_REGION)); List options = authSchemeProvider.resolveAuthScheme(paramsBuilder.build()); return options; } - private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes, String operationName) { + private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executionAttributes) { + String operationName = executionAttributes.getAttribute(SdkExecutionAttribute.OPERATION_NAME); XmlEndpointProvider provider = (XmlEndpointProvider) executionAttributes .getAttribute(SdkInternalExecutionAttribute.ENDPOINT_PROVIDER); try { @@ -756,14 +772,6 @@ private Endpoint resolveEndpoint(SdkRequest request, ExecutionAttributes executi } } - private AuthSchemeOptionsResolver authSchemeResolver(String operationName, SdkClientConfiguration clientConfiguration) { - return r -> resolveAuthSchemeOptions(r, operationName, clientConfiguration); - } - - private EndpointResolver endpointResolver(String operationName) { - return (r, a) -> resolveEndpoint(r, a, operationName); - } - private void updateRetryStrategyClientConfiguration(SdkClientConfiguration.Builder configuration) { ClientOverrideConfiguration.Builder builder = configuration.asOverrideConfigurationBuilder(); RetryMode retryMode = builder.retryMode(); diff --git a/core/sdk-core/src/main/java/software/amazon/awssdk/core/endpoint/EndpointResolver.java b/core/sdk-core/src/main/java/software/amazon/awssdk/core/endpoint/EndpointResolver.java index c17b6f96d8e6..fb93e96a273d 100644 --- a/core/sdk-core/src/main/java/software/amazon/awssdk/core/endpoint/EndpointResolver.java +++ b/core/sdk-core/src/main/java/software/amazon/awssdk/core/endpoint/EndpointResolver.java @@ -23,7 +23,6 @@ /** * Callback interface for resolving endpoints from the request and execution context. */ -@FunctionalInterface @SdkProtectedApi public interface EndpointResolver { /** diff --git a/core/sdk-core/src/main/java/software/amazon/awssdk/core/http/auth/AuthSchemeResolver.java b/core/sdk-core/src/main/java/software/amazon/awssdk/core/http/auth/AuthSchemeResolver.java index ac795a34c494..1b5d64324c61 100644 --- a/core/sdk-core/src/main/java/software/amazon/awssdk/core/http/auth/AuthSchemeResolver.java +++ b/core/sdk-core/src/main/java/software/amazon/awssdk/core/http/auth/AuthSchemeResolver.java @@ -82,7 +82,7 @@ public static SelectedAuthScheme resolveAuthScheme( identityProviders = resolver.resolve(request, identityProviders, executionAttributes); } - List authOptions = optionsResolver.resolve(request); + List authOptions = optionsResolver.resolve(request, executionAttributes); return selectAuthScheme(authOptions, authSchemes, identityProviders, null); } diff --git a/core/sdk-core/src/main/java/software/amazon/awssdk/core/internal/http/pipeline/stages/AuthSchemeResolutionStage.java b/core/sdk-core/src/main/java/software/amazon/awssdk/core/internal/http/pipeline/stages/AuthSchemeResolutionStage.java index 831340b29d16..4c1f5d3b4f8d 100644 --- a/core/sdk-core/src/main/java/software/amazon/awssdk/core/internal/http/pipeline/stages/AuthSchemeResolutionStage.java +++ b/core/sdk-core/src/main/java/software/amazon/awssdk/core/internal/http/pipeline/stages/AuthSchemeResolutionStage.java @@ -114,7 +114,7 @@ private List resolveAuthSchemeOptions(ExecutionAttributes exec if (resolver == null) { return null; } - return resolver.resolve(request); + return resolver.resolve(request, executionAttributes); } /** diff --git a/core/sdk-core/src/main/java/software/amazon/awssdk/core/spi/identity/AuthSchemeOptionsResolver.java b/core/sdk-core/src/main/java/software/amazon/awssdk/core/spi/identity/AuthSchemeOptionsResolver.java index 7a37d7c3dd6b..759fe016904b 100644 --- a/core/sdk-core/src/main/java/software/amazon/awssdk/core/spi/identity/AuthSchemeOptionsResolver.java +++ b/core/sdk-core/src/main/java/software/amazon/awssdk/core/spi/identity/AuthSchemeOptionsResolver.java @@ -18,6 +18,7 @@ import java.util.List; import software.amazon.awssdk.annotations.SdkProtectedApi; import software.amazon.awssdk.core.SdkRequest; +import software.amazon.awssdk.core.interceptor.ExecutionAttributes; import software.amazon.awssdk.http.auth.spi.scheme.AuthSchemeOption; /** @@ -26,7 +27,6 @@ * This allows auth scheme resolution to happen after interceptors have modified the request, * ensuring that any request modifications affecting auth scheme selection are respected. */ -@FunctionalInterface @SdkProtectedApi public interface AuthSchemeOptionsResolver { /** @@ -36,4 +36,19 @@ public interface AuthSchemeOptionsResolver { * @return List of auth scheme options in priority order */ List resolve(SdkRequest request); + + /** + * Resolves auth scheme options for the given request and execution attributes. + *

+ * This is the method called by the SDK core pipeline. The default implementation delegates + * to {@link #resolve(SdkRequest)} for backward compatibility with service modules that + * only implement the 1-arg method. + * + * @param request The request (after interceptors have modified it) + * @param executionAttributes The execution attributes for the current request execution + * @return List of auth scheme options in priority order + */ + default List resolve(SdkRequest request, ExecutionAttributes executionAttributes) { + return resolve(request); + } } diff --git a/core/sdk-core/src/test/java/software/amazon/awssdk/core/internal/http/pipeline/stages/AuthSchemeResolutionStageTest.java b/core/sdk-core/src/test/java/software/amazon/awssdk/core/internal/http/pipeline/stages/AuthSchemeResolutionStageTest.java index 6ab7325a3093..1b1ace94846e 100644 --- a/core/sdk-core/src/test/java/software/amazon/awssdk/core/internal/http/pipeline/stages/AuthSchemeResolutionStageTest.java +++ b/core/sdk-core/src/test/java/software/amazon/awssdk/core/internal/http/pipeline/stages/AuthSchemeResolutionStageTest.java @@ -17,6 +17,7 @@ import static org.assertj.core.api.Assertions.assertThat; import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; import static org.mockito.Mockito.doReturn; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.verify; @@ -100,7 +101,7 @@ void execute_noResolver_returnsRequestUnchanged() throws Exception { void execute_resolverReturnsEmpty_returnsRequestUnchanged() throws Exception { executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEMES, createAuthSchemes()); executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_OPTIONS_RESOLVER, - (AuthSchemeOptionsResolver) req -> Collections.emptyList()); + resolverReturning(Collections.emptyList())); SdkHttpFullRequest.Builder result = stage.execute(httpRequestBuilder, context); @@ -126,7 +127,7 @@ void execute_resolverReceivesRequestFromInterceptorContext() throws Exception { .build(); AuthSchemeOptionsResolver resolver = mock(AuthSchemeOptionsResolver.class); - doReturn(createAuthOptions()).when(resolver).resolve(modifiedRequest); + doReturn(createAuthOptions()).when(resolver).resolve(eq(modifiedRequest), any(ExecutionAttributes.class)); executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEMES, createAuthSchemes()); executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_OPTIONS_RESOLVER, resolver); @@ -135,7 +136,7 @@ void execute_resolverReceivesRequestFromInterceptorContext() throws Exception { stage.execute(httpRequestBuilder, context); // Verify resolver was called with the MODIFIED request, not originalRequest - verify(resolver).resolve(modifiedRequest); + verify(resolver).resolve(eq(modifiedRequest), any(ExecutionAttributes.class)); } @Test @@ -156,7 +157,7 @@ void execute_withRequestIdentityProviderResolver_callsUpdaterWithRequest() throw executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEMES, authSchemes); executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_OPTIONS_RESOLVER, - (AuthSchemeOptionsResolver) req -> createAuthOptions()); + resolverReturning(createAuthOptions())); executionAttributes.putAttribute(SdkInternalExecutionAttribute.IDENTITY_PROVIDERS, baseProviders); executionAttributes.putAttribute(SdkInternalExecutionAttribute.IDENTITY_PROVIDER_RESOLVER, resolver); @@ -169,7 +170,7 @@ void execute_withRequestIdentityProviderResolver_callsUpdaterWithRequest() throw void execute_withoutRequestIdentityProviderResolver_doesNotFail() throws Exception { executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEMES, createAuthSchemes()); executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_OPTIONS_RESOLVER, - (AuthSchemeOptionsResolver) req -> createAuthOptions()); + resolverReturning(createAuthOptions())); executionAttributes.putAttribute(SdkInternalExecutionAttribute.IDENTITY_PROVIDERS, createIdentityProviders()); // No IDENTITY_PROVIDER_RESOLVER set @@ -183,7 +184,7 @@ void execute_withoutRequestIdentityProviderResolver_doesNotFail() throws Excepti void execute_happyPath_setsSelectedAuthScheme() throws Exception { executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEMES, createAuthSchemes()); executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_OPTIONS_RESOLVER, - (AuthSchemeOptionsResolver) req -> createAuthOptions()); + resolverReturning(createAuthOptions())); executionAttributes.putAttribute(SdkInternalExecutionAttribute.IDENTITY_PROVIDERS, createIdentityProviders()); SdkHttpFullRequest.Builder result = stage.execute(httpRequestBuilder, context); @@ -229,6 +230,15 @@ private List createAuthOptions() { ); } + /** + * Creates a mock AuthSchemeOptionsResolver that returns the given options when the 2-arg resolve is called. + */ + private AuthSchemeOptionsResolver resolverReturning(List options) { + AuthSchemeOptionsResolver resolver = mock(AuthSchemeOptionsResolver.class); + doReturn(options).when(resolver).resolve(any(SdkRequest.class), any(ExecutionAttributes.class)); + return resolver; + } + @Test void execute_authSchemeAlreadyResolved_skipsResolution() throws Exception { // Simulate old service interceptor already resolved auth scheme @@ -240,7 +250,7 @@ void execute_authSchemeAlreadyResolved_skipsResolution() throws Exception { executionAttributes.putAttribute(SdkInternalExecutionAttribute.SELECTED_AUTH_SCHEME, alreadyResolved); executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEMES, createAuthSchemes()); executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_OPTIONS_RESOLVER, - (AuthSchemeOptionsResolver) req -> createAuthOptions()); + resolverReturning(createAuthOptions())); executionAttributes.putAttribute(SdkInternalExecutionAttribute.IDENTITY_PROVIDERS, createIdentityProviders()); SdkHttpFullRequest.Builder result = stage.execute(httpRequestBuilder, context); @@ -262,7 +272,7 @@ void execute_authSchemeUnset_proceedsWithResolution() throws Exception { executionAttributes.putAttribute(SdkInternalExecutionAttribute.SELECTED_AUTH_SCHEME, unsetPlaceholder); executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEMES, createAuthSchemes()); executionAttributes.putAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_OPTIONS_RESOLVER, - (AuthSchemeOptionsResolver) req -> createAuthOptions()); + resolverReturning(createAuthOptions())); executionAttributes.putAttribute(SdkInternalExecutionAttribute.IDENTITY_PROVIDERS, createIdentityProviders()); stage.execute(httpRequestBuilder, context); diff --git a/services/s3/src/main/java/software/amazon/awssdk/services/s3/internal/s3express/S3ExpressUtils.java b/services/s3/src/main/java/software/amazon/awssdk/services/s3/internal/s3express/S3ExpressUtils.java index 42c7928db434..adddf4270c22 100644 --- a/services/s3/src/main/java/software/amazon/awssdk/services/s3/internal/s3express/S3ExpressUtils.java +++ b/services/s3/src/main/java/software/amazon/awssdk/services/s3/internal/s3express/S3ExpressUtils.java @@ -55,7 +55,7 @@ public static boolean isS3ExpressAuthRequest(SdkRequest request, ExecutionAttrib AuthSchemeOptionsResolver resolver = executionAttributes.getAttribute(SdkInternalExecutionAttribute.AUTH_SCHEME_OPTIONS_RESOLVER); if (resolver != null) { - List options = resolver.resolve(request); + List options = resolver.resolve(request, executionAttributes); return options.stream().anyMatch(o -> S3ExpressAuthScheme.SCHEME_ID.equals(o.schemeId())); } return false; diff --git a/services/s3/src/test/java/software/amazon/awssdk/services/s3/AuthSchemeProviderOverrideTest.java b/services/s3/src/test/java/software/amazon/awssdk/services/s3/AuthSchemeProviderOverrideTest.java new file mode 100644 index 000000000000..447c21641fb1 --- /dev/null +++ b/services/s3/src/test/java/software/amazon/awssdk/services/s3/AuthSchemeProviderOverrideTest.java @@ -0,0 +1,150 @@ +/* + * Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"). + * You may not use this file except in compliance with the License. + * A copy of the License is located at + * + * http://aws.amazon.com/apache2.0 + * + * or in the "license" file accompanying this file. This file is distributed + * on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either + * express or implied. See the License for the specific language governing + * permissions and limitations under the License. + */ + +package software.amazon.awssdk.services.s3; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.util.concurrent.atomic.AtomicReference; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; +import software.amazon.awssdk.auth.credentials.AwsBasicCredentials; +import software.amazon.awssdk.auth.credentials.StaticCredentialsProvider; +import software.amazon.awssdk.core.SdkPlugin; +import software.amazon.awssdk.core.interceptor.Context; +import software.amazon.awssdk.core.interceptor.ExecutionAttributes; +import software.amazon.awssdk.core.interceptor.ExecutionInterceptor; +import software.amazon.awssdk.regions.Region; +import software.amazon.awssdk.services.s3.auth.scheme.S3AuthSchemeProvider; + +/** + * Verifies that request-level plugin overrides for auth scheme provider and region + * are respected during auth scheme resolution. + */ +class AuthSchemeProviderOverrideTest { + + @Test + void requestPluginOverridesAuthSchemeProvider_isUsed() { + AtomicInteger defaultProviderCallCount = new AtomicInteger(0); + AtomicInteger overrideProviderCallCount = new AtomicInteger(0); + + S3AuthSchemeProvider defaultProvider = params -> { + defaultProviderCallCount.incrementAndGet(); + return S3AuthSchemeProvider.defaultProvider().resolveAuthScheme(params); + }; + + S3AuthSchemeProvider overrideProvider = params -> { + overrideProviderCallCount.incrementAndGet(); + return S3AuthSchemeProvider.defaultProvider().resolveAuthScheme(params); + }; + + S3Client client = S3Client.builder() + .region(Region.US_EAST_1) + .credentialsProvider(StaticCredentialsProvider.create( + AwsBasicCredentials.create("akid", "skid"))) + .addPlugin(authSchemeProviderPlugin(defaultProvider)) + .overrideConfiguration(c -> c.addExecutionInterceptor(new AbortInterceptor())) + .build(); + + // Request with plugin that overrides auth scheme provider + assertThatThrownBy(() -> client.getObject(r -> { + r.overrideConfiguration(c -> c.addPlugin(authSchemeProviderPlugin(overrideProvider))); + r.key("key").bucket("bucket"); + })).hasMessageContaining("aborted"); + + // The override provider should have been called, not the default + assertThat(overrideProviderCallCount.get()).isGreaterThan(0); + assertThat(defaultProviderCallCount.get()).isEqualTo(0); + } + + @Test + void noRequestPluginOverride_usesClientProvider() { + AtomicInteger defaultProviderCallCount = new AtomicInteger(0); + + S3AuthSchemeProvider defaultProvider = params -> { + defaultProviderCallCount.incrementAndGet(); + return S3AuthSchemeProvider.defaultProvider().resolveAuthScheme(params); + }; + + S3Client client = S3Client.builder() + .region(Region.US_EAST_1) + .credentialsProvider(StaticCredentialsProvider.create( + AwsBasicCredentials.create("akid", "skid"))) + .addPlugin(authSchemeProviderPlugin(defaultProvider)) + .overrideConfiguration(c -> c.addExecutionInterceptor(new AbortInterceptor())) + .build(); + + // Request without override — should use client-level provider + assertThatThrownBy(() -> client.getObject(r -> { + r.key("key").bucket("bucket"); + })).hasMessageContaining("aborted"); + + assertThat(defaultProviderCallCount.get()).isGreaterThan(0); + } + + @Test + void requestPluginOverridesRegion_isUsedInAuthSchemeResolution() { + AtomicReference capturedRegion = new AtomicReference<>(); + + S3AuthSchemeProvider capturingProvider = params -> { + capturedRegion.set(params.region()); + return S3AuthSchemeProvider.defaultProvider().resolveAuthScheme(params); + }; + + S3Client client = S3Client.builder() + .region(Region.US_EAST_1) + .credentialsProvider(StaticCredentialsProvider.create( + AwsBasicCredentials.create("akid", "skid"))) + .authSchemeProvider(capturingProvider) + .overrideConfiguration(c -> c.addExecutionInterceptor(new AbortInterceptor())) + .build(); + + // Request with plugin that overrides region to us-west-2 + assertThatThrownBy(() -> client.getObject(r -> { + r.overrideConfiguration(c -> c.addPlugin(regionOverridePlugin(Region.US_WEST_2))); + r.key("key").bucket("bucket"); + })).hasMessageContaining("aborted"); + + // The auth scheme params should have received the overridden region + assertThat(capturedRegion.get()).isEqualTo(Region.US_WEST_2); + } + + private SdkPlugin regionOverridePlugin(Region region) { + return config -> { + S3ServiceClientConfiguration.Builder s3Config = + (S3ServiceClientConfiguration.Builder) config; + s3Config.region(region); + }; + } + + private SdkPlugin authSchemeProviderPlugin(S3AuthSchemeProvider provider) { + return config -> { + S3ServiceClientConfiguration.Builder s3Config = + (S3ServiceClientConfiguration.Builder) config; + s3Config.authSchemeProvider(provider); + }; + } + + /** + * Interceptor that aborts the request before transmission so we don't need a real endpoint. + */ + static class AbortInterceptor implements ExecutionInterceptor { + @Override + public void beforeTransmission(Context.BeforeTransmission context, ExecutionAttributes executionAttributes) { + throw new RuntimeException("aborted"); + } + } +} diff --git a/test/old-client-version-compatibility-test/pom.xml b/test/old-client-version-compatibility-test/pom.xml index b8cda7a029c3..c7f979a519f5 100644 --- a/test/old-client-version-compatibility-test/pom.xml +++ b/test/old-client-version-compatibility-test/pom.xml @@ -105,6 +105,11 @@ s3 2.20.136 + + software.amazon.awssdk + polly + 2.48.0 + diff --git a/test/old-client-version-compatibility-test/src/test/java/OldAuthSchemeOptionsResolverCompatibilityTest.java b/test/old-client-version-compatibility-test/src/test/java/OldAuthSchemeOptionsResolverCompatibilityTest.java new file mode 100644 index 000000000000..e05ad605f2c2 --- /dev/null +++ b/test/old-client-version-compatibility-test/src/test/java/OldAuthSchemeOptionsResolverCompatibilityTest.java @@ -0,0 +1,70 @@ +/* + * Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"). + * You may not use this file except in compliance with the License. + * A copy of the License is located at + * + * + * http://aws.amazon.com/apache2.0 + * + * or in the "license" file accompanying this file. This file is distributed + * on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either + * express or implied. See the License for the specific language governing + * permissions and limitations under the License. + */ + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatNoException; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import software.amazon.awssdk.auth.credentials.AwsBasicCredentials; +import software.amazon.awssdk.auth.credentials.StaticCredentialsProvider; +import software.amazon.awssdk.regions.Region; +import software.amazon.awssdk.services.polly.PollyClient; +import software.amazon.awssdk.testutils.service.http.MockSyncHttpClient; + +/** + * Verifies that an old service module (Polly 2.48.0, compiled with the 1-arg AuthSchemeOptionsResolver) + * works correctly with the new core (which calls the 2-arg resolve method). + *

+ * The default method on AuthSchemeOptionsResolver delegates the 2-arg call to the 1-arg implementation, + * preventing AbstractMethodError in mixed-version scenarios. + */ +public class OldAuthSchemeOptionsResolverCompatibilityTest { + + private MockSyncHttpClient httpClient; + private PollyClient pollyClient; + + @BeforeEach + public void setup() { + this.httpClient = new MockSyncHttpClient(); + this.httpClient.stubNextResponse200(); + + this.pollyClient = PollyClient.builder() + .region(Region.US_EAST_1) + .credentialsProvider(StaticCredentialsProvider.create( + AwsBasicCredentials.create("akid", "skid"))) + .httpClient(httpClient) + .build(); + } + + @AfterEach + public void teardown() { + httpClient.close(); + pollyClient.close(); + } + + /** + * Old Polly client (2.48.0) uses a 1-arg AuthSchemeOptionsResolver lambda. + * New core calls the 2-arg resolve() method. + * The default method bridge should prevent AbstractMethodError. + */ + @Test + public void oldServiceModule_withNewCore_doesNotThrowAbstractMethodError() { + assertThatNoException().isThrownBy(() -> pollyClient.listLexicons()); + assertThat(httpClient.getLastRequest()).isNotNull(); + } +}