如何用Mockito测试含AWS SDK最终类的Java Lambda代码?
Let's break down your problem step by step, covering mocking AWS SDK final classes, optimizing your code for testability, and writing effective unit tests—all while keeping your private helper methods intact.
1. Mocking AWS SDK Final Classes with Mockito
The error you’re seeing (Mockito can’t mock final classes) is common with AWS SDK, since many of its builder classes (like AWSStepFunctionsClientBuilder) are marked final. Here’s how to fix this with just Mockito:
Enable Mockito's Inline Mock Maker
Mockito 3.4.0+ includes an inline mock maker that can mock final classes and methods. To enable it:
- Create a directory
src/test/resources/mockito-extensionsin your project - Add a file named
org.mockito.plugins.MockMakerwith the single line:mock-maker-inline - If you’re using Maven/Gradle, ensure you have the
mockito-inlinedependency (instead of plainmockito-core) in your test scope:
Maven:
Gradle:<dependency> <groupId>org.mockito</groupId> <artifactId>mockito-inline</artifactId> <version>4.11.0</version> <scope>test</scope> </dependency>testImplementation 'org.mockito:mockito-inline:4.11.0'
With this setup, you can mock final classes like AWSStepFunctionsClientBuilder—though we’ll cover a better approach below that avoids mocking builders entirely.
2. Optimize Your Code for Testability (Keep Private Helpers)
Your current code has tight coupling to AWS SDK creation and static time calls, which makes testing hard. Here’s how to refactor it without removing your private methods:
Refactor to Use Dependency Injection
Instead of creating AWSStepFunctions inside the handleRequest method, inject it via a constructor. This lets you pass a mock instance during tests while keeping the default constructor for Lambda runtime.
Extract Time Dependency
Replace the static LocalDateTime.now() with a Clock instance, so you can fix the time during tests (no more flakiness from changing system time).
Here’s the refactored code:
import java.time.Clock; import java.time.LocalDateTime; import java.time.ZoneId; import com.amazonaws.services.stepfunctions.AWSStepFunctions; import com.amazonaws.services.stepfunctions.AWSStepFunctionsClientBuilder; import com.amazonaws.services.stepfunctions.model.*; import com.amazonaws.services.lambda.runtime.Context; import com.amazonaws.services.lambda.runtime.RequestStreamHandler; import java.io.IOException; import java.io.InputStream; import java.io.OutputStream; import java.time.temporal.ChronoUnit; public class StepFunctionMetricsLambda implements RequestStreamHandler { private final AWSStepFunctions stepFunctionsClient; private final Clock clock; // Default constructor for Lambda runtime public StepFunctionMetricsLambda() { this(AWSStepFunctionsClientBuilder.standard().build(), Clock.systemDefaultZone()); } // Constructor for testing (inject mock client and fixed clock) public StepFunctionMetricsLambda(AWSStepFunctions stepFunctionsClient, Clock clock) { this.stepFunctionsClient = stepFunctionsClient; this.clock = clock; } @Override public void handleRequest(InputStream input, OutputStream output, Context context) throws IOException { int inProgressStateMachines = 0; LocalDateTime now = LocalDateTime.now(clock); // Use injected clock long alarmThreshold = getAlarmThreshold(input, context.getLogger()); ListStateMachinesRequest listStateMachinesRequest = new ListStateMachinesRequest(); ListStateMachinesResult listStateMachinesResult = stepFunctionsClient.listStateMachines(listStateMachinesRequest); for (StateMachineListItem stateMachineListItem : listStateMachinesResult.getStateMachines()) { ListExecutionsRequest listExecutionRequest = new ListExecutionsRequest() .withStateMachineArn(stateMachineListItem.getStateMachineArn()) .withStatusFilter(ExecutionStatus.RUNNING); ListExecutionsResult listExecutionsResult = stepFunctionsClient.listExecutions(listExecutionRequest); for (ExecutionListItem executionListItem : listExecutionsResult.getExecutions()) { LocalDateTime stateMachineStartTime = LocalDateTime.ofInstant( executionListItem.getStartDate().toInstant(), ZoneId.systemDefault()); long elapsedTime = ChronoUnit.SECONDS.between(stateMachineStartTime, now); if (elapsedTime > alarmThreshold){ inProgressStateMachines++; } } publishMetrics(inProgressStateMachines); } } // Your existing private helper methods stay here private long getAlarmThreshold(InputStream input, com.amazonaws.services.lambda.runtime.Logger logger) { // ... your existing logic } private void publishMetrics(int count) { // ... your existing logic } }
Key improvements:
- No more direct creation of AWS clients inside the handler—easy to mock
- Time is now controlled via a
Clockinstance, perfect for testing elapsed time logic - Private methods remain untouched, as requested
3. Writing Unit Tests for the Refactored Code
Now you can write clean, reliable unit tests using Mockito. Here’s an example using JUnit 5:
import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.InjectMocks; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import java.io.ByteArrayInputStream; import java.io.ByteArrayOutputStream; import java.time.Clock; import java.time.Instant; import java.time.ZoneId; import java.util.List; import com.amazonaws.services.stepfunctions.AWSStepFunctions; import com.amazonaws.services.stepfunctions.model.*; import com.amazonaws.services.lambda.runtime.Context; import static org.mockito.Mockito.*; @ExtendWith(MockitoExtension.class) class StepFunctionMetricsLambdaTest { @Mock private AWSStepFunctions mockStepFunctionsClient; private final Clock fixedClock = Clock.fixed(Instant.parse("2024-05-20T10:00:00Z"), ZoneId.systemDefault()); private final StepFunctionMetricsLambda lambda = new StepFunctionMetricsLambda(mockStepFunctionsClient, fixedClock); @Test void handleRequest_ShouldCountLongRunningExecutionsAndPublishMetrics() throws Exception { // 1. Setup mock AWS responses StateMachineListItem mockStateMachine = new StateMachineListItem().withStateMachineArn("arn:aws:states:us-east-1:123456789012:stateMachine:TestSM"); ListStateMachinesResult mockSMResult = new ListStateMachinesResult().withStateMachines(List.of(mockStateMachine)); when(mockStepFunctionsClient.listStateMachines(any(ListStateMachinesRequest.class))).thenReturn(mockSMResult); // Create an execution that's over the threshold (e.g., threshold is 300 seconds = 5 mins) ExecutionListItem longRunningExecution = new ExecutionListItem() .withStartDate(java.util.Date.from(fixedClock.instant().minusSeconds(360))); // 6 mins old ListExecutionsResult mockExecResult = new ListExecutionsResult().withExecutions(List.of(longRunningExecution)); when(mockStepFunctionsClient.listExecutions(any(ListExecutionsRequest.class))).thenReturn(mockExecResult); // 2. Call the handler ByteArrayInputStream input = new ByteArrayInputStream("{\"alarmThreshold\":300}".getBytes()); // Example input lambda.handleRequest(input, new ByteArrayOutputStream(), mock(Context.class)); // 3. Verify interactions verify(mockStepFunctionsClient, times(1)).listStateMachines(any()); verify(mockStepFunctionsClient, times(1)).listExecutions(any()); // To verify the private publishMetrics method was called with the right count: StepFunctionMetricsLambda lambdaSpy = spy(new StepFunctionMetricsLambda(mockStepFunctionsClient, fixedClock)); doNothing().when(lambdaSpy).publishMetrics(anyInt()); // Suppress actual metric publishing lambdaSpy.handleRequest(input, new ByteArrayOutputStream(), mock(Context.class)); verify(lambdaSpy, times(1)).publishMetrics(1); // Should count 1 long-running execution } }
Verifying Private Helper Method Parameters
Since you want to keep private methods, you can’t directly verify their parameters with plain Mockito. Instead:
- For
getAlarmThreshold: Test it indirectly by passing a known input stream and asserting that the threshold logic behaves as expected (e.g., if input has threshold 300, ensure only executions older than 300s are counted) - For
publishMetrics: Use a spy of your lambda class (as shown in the test above) to verify it’s called with the correct count. You can suppress the actual implementation withdoNothing()so tests don’t try to send real metrics.
内容的提问来源于stack exchange,提问作者Em Ae

