You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

SpringBoot单元测试:验证@Configuration类中RestTemplate的请求头

正确的SpringBoot单元测试:验证RestTemplate请求头包含x-correlation-id

原配置类代码

@Configuration
public class RestConfiguration {
    private static final Logger log = LoggerFactory.getLogger(RestConfiguration.class);

    @Bean @Primary
    public RestTemplate restTemplate(
            CorrelationComponent correlationComponent,
            CloudTokener cloudTokener,
            @Value("${tokens.aad.role-app.resource}") String resourceId,
            RestTemplateBuilder restTemplateBuilder,
            @Value("${rest.connect-timeout}") Long connectTimeout,
            @Value("${rest.read-timeout}") Long readTimeout
    ) {
        log.trace("restTemplate() start");
        List<ClientHttpRequestInterceptor> interceptors = Arrays.asList((request, body, execution) -> {
            // 添加关联ID请求头
            request.getHeaders().set(Constants.X_CORRELATION_ID, correlationComponent.getCorrelationId());
            return execution.execute(request, body);
        });
        RestTemplate bean = restTemplateBuilder
                .setConnectTimeout(Duration.ofMillis(connectTimeout))
                .setReadTimeout(Duration.ofMillis(readTimeout))
                .interceptors(interceptors)
                .build();
        log.trace("restTemplate() end");
        return bean;
    }
}

你原有测试的问题

  • 普通测试类中直接用@Value无法读取配置,因为未启动Spring上下文
  • 对RestTemplateBuilder的mock未设置正确行为,调用build()会返回null
  • 未触发RestTemplate的请求执行逻辑,无法验证拦截器是否添加了目标请求头
  • 变量名错误:azureTokener应为cloudTokener

正确的单元测试代码

import org.junit.jupiter.api.Test;
import org.springframework.boot.web.client.RestTemplateBuilder;
import org.springframework.http.HttpHeaders;
import org.springframework.http.HttpMethod;
import org.springframework.http.client.ClientHttpRequestExecution;
import org.springframework.http.client.ClientHttpRequestInterceptor;
import org.springframework.http.client.ClientHttpResponse;
import org.springframework.web.client.RestTemplate;

import java.io.IOException;
import java.util.List;

import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertNotNull;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;

class RestConfigurationTest {

    @Test
    void restTemplate_ShouldAddCorrelationIdHeader() {
        // 1. 模拟依赖组件,返回固定测试用关联ID
        CorrelationComponent mockCorrelationComponent = mock(CorrelationComponent.class);
        String testCorrelationId = "test-12345";
        when(mockCorrelationComponent.getCorrelationId()).thenReturn(testCorrelationId);

        CloudTokener mockCloudTokener = mock(CloudTokener.class);
        String resourceId = "dummy-resource-id";
        Long connectTimeout = 1000L;
        Long readTimeout = 2000L;

        // 2. 使用真实的RestTemplateBuilder,避免mock行为缺失问题
        RestTemplateBuilder restTemplateBuilder = new RestTemplateBuilder();

        // 3. 创建配置类实例并获取RestTemplate
        RestConfiguration restConfiguration = new RestConfiguration();
        RestTemplate restTemplate = restConfiguration.restTemplate(
                mockCorrelationComponent,
                mockCloudTokener,
                resourceId,
                restTemplateBuilder,
                connectTimeout,
                readTimeout
        );

        // 4. 获取拦截器,模拟请求执行过程
        List<ClientHttpRequestInterceptor> interceptors = restTemplate.getInterceptors();
        assertNotNull(interceptors);
        assertEquals(1, interceptors.size());

        ClientHttpRequestInterceptor interceptor = interceptors.get(0);
        MockClientHttpRequest request = new MockClientHttpRequest(HttpMethod.GET, "http://test.com");
        ClientHttpRequestExecution execution = mock(ClientHttpRequestExecution.class);
        when(execution.execute(request, new byte[0])).thenReturn(mock(ClientHttpResponse.class));

        try {
            interceptor.intercept(request, new byte[0], execution);
        } catch (IOException e) {
            e.printStackTrace();
        }

        // 5. 验证请求头是否包含正确的x-correlation-id
        HttpHeaders headers = request.getHeaders();
        assertEquals(testCorrelationId, headers.getFirst(Constants.X_CORRELATION_ID));
    }

    // 自定义Mock类,用于捕获请求头信息
    private static class MockClientHttpRequest extends org.springframework.http.client.AbstractClientHttpRequest {
        private final HttpMethod method;
        private final String uri;
        private final HttpHeaders headers = new HttpHeaders();

        public MockClientHttpRequest(HttpMethod method, String uri) {
            this.method = method;
            this.uri = uri;
        }

        @Override
        protected HttpHeaders getHeadersInternal() {
            return headers;
        }

        @Override
        protected ClientHttpResponse executeInternal(HttpHeaders headers, byte[] body) throws IOException {
            return null;
        }

        @Override
        public HttpMethod getMethod() {
            return method;
        }

        @Override
        public String getURI() {
            return uri;
        }
    }
}

关键说明

  • 使用真实的RestTemplateBuilder而非mock,避免因mock配置不全导致的异常
  • 给CorrelationComponent设置固定返回值,便于后续断言验证
  • 自定义MockClientHttpRequest捕获请求头,验证拦截器的逻辑是否生效
  • 手动触发拦截器的执行方法,模拟实际请求场景,确保请求头被正确添加

内容的提问来源于stack exchange,提问作者dp101

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.29 20:47:27