WebMvc support for SSE Fragment stream

Closes gh-33194
This commit is contained in:
rstoyanchev
2024-08-02 17:36:08 +03:00
parent 184bb7c23c
commit 622c1b9e8c
3 changed files with 404 additions and 15 deletions

View File

@@ -18,6 +18,7 @@ package org.springframework.web.servlet.mvc.method.annotation;
import java.lang.reflect.Method;
import java.util.ArrayList;
import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
@@ -32,7 +33,10 @@ import jakarta.servlet.http.HttpSession;
import org.springframework.beans.factory.BeanFactory;
import org.springframework.beans.factory.BeanFactoryAware;
import org.springframework.beans.factory.BeanFactoryUtils;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.beans.factory.ListableBeanFactory;
import org.springframework.beans.factory.NoSuchBeanDefinitionException;
import org.springframework.beans.factory.config.ConfigurableBeanFactory;
import org.springframework.core.DefaultParameterNameDiscoverer;
import org.springframework.core.KotlinDetector;
@@ -41,6 +45,7 @@ import org.springframework.core.MethodParameter;
import org.springframework.core.ParameterNameDiscoverer;
import org.springframework.core.ReactiveAdapterRegistry;
import org.springframework.core.annotation.AnnotatedElementUtils;
import org.springframework.core.annotation.AnnotationAwareOrderComparator;
import org.springframework.core.log.LogFormatUtils;
import org.springframework.core.task.AsyncTaskExecutor;
import org.springframework.core.task.SimpleAsyncTaskExecutor;
@@ -96,8 +101,11 @@ import org.springframework.web.method.support.HandlerMethodReturnValueHandler;
import org.springframework.web.method.support.HandlerMethodReturnValueHandlerComposite;
import org.springframework.web.method.support.InvocableHandlerMethod;
import org.springframework.web.method.support.ModelAndViewContainer;
import org.springframework.web.servlet.DispatcherServlet;
import org.springframework.web.servlet.LocaleResolver;
import org.springframework.web.servlet.ModelAndView;
import org.springframework.web.servlet.View;
import org.springframework.web.servlet.ViewResolver;
import org.springframework.web.servlet.mvc.annotation.ModelAndViewResolver;
import org.springframework.web.servlet.mvc.method.AbstractHandlerMethodAdapter;
import org.springframework.web.servlet.mvc.support.RedirectAttributes;
@@ -767,7 +775,8 @@ public class RequestMappingHandlerAdapter extends AbstractHandlerMethodAdapter
handlers.add(new ModelMethodProcessor());
handlers.add(new ViewMethodReturnValueHandler());
handlers.add(new ResponseBodyEmitterReturnValueHandler(getMessageConverters(),
this.reactiveAdapterRegistry, this.taskExecutor, this.contentNegotiationManager));
this.reactiveAdapterRegistry, this.taskExecutor, this.contentNegotiationManager,
initViewResolvers(), initLocaleResolver()));
handlers.add(new StreamingResponseBodyReturnValueHandler());
handlers.add(new HttpEntityMethodProcessor(getMessageConverters(),
this.contentNegotiationManager, this.requestResponseBodyAdvice, this.errorResponseInterceptors));
@@ -801,6 +810,33 @@ public class RequestMappingHandlerAdapter extends AbstractHandlerMethodAdapter
return handlers;
}
private List<ViewResolver> initViewResolvers() {
if (getBeanFactory() instanceof ListableBeanFactory lbf) {
Map<String, ViewResolver> matchingBeans =
BeanFactoryUtils.beansOfTypeIncludingAncestors(lbf, ViewResolver.class, true, false);
if (!matchingBeans.isEmpty()) {
List<ViewResolver> viewResolvers = new ArrayList<>(matchingBeans.values());
AnnotationAwareOrderComparator.sort(viewResolvers);
return viewResolvers;
}
}
return Collections.emptyList();
}
@Nullable
private LocaleResolver initLocaleResolver() {
if (getBeanFactory() != null) {
try {
return getBeanFactory().getBean(
DispatcherServlet.LOCALE_RESOLVER_BEAN_NAME, LocaleResolver.class);
}
catch (NoSuchBeanDefinitionException ex) {
// ignore
}
}
return null;
}
private static Predicate<MethodParameter> methodParamPredicate(
List<HandlerMethodArgumentResolver> resolvers, Class<?> resolverType) {

View File

@@ -16,19 +16,28 @@
package org.springframework.web.servlet.mvc.method.annotation;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.io.PrintWriter;
import java.nio.charset.Charset;
import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.Locale;
import java.util.Set;
import java.util.function.Consumer;
import jakarta.servlet.ServletRequest;
import jakarta.servlet.ServletOutputStream;
import jakarta.servlet.WriteListener;
import jakarta.servlet.http.HttpServletRequest;
import jakarta.servlet.http.HttpServletResponse;
import jakarta.servlet.http.HttpServletResponseWrapper;
import org.springframework.core.MethodParameter;
import org.springframework.core.ReactiveAdapterRegistry;
import org.springframework.core.ResolvableType;
import org.springframework.core.task.SyncTaskExecutor;
import org.springframework.core.task.TaskExecutor;
import org.springframework.http.HttpHeaders;
import org.springframework.http.MediaType;
@@ -40,13 +49,23 @@ import org.springframework.http.server.ServerHttpResponse;
import org.springframework.http.server.ServletServerHttpResponse;
import org.springframework.lang.Nullable;
import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
import org.springframework.web.accept.ContentNegotiationManager;
import org.springframework.web.context.request.NativeWebRequest;
import org.springframework.web.context.request.RequestContextHolder;
import org.springframework.web.context.request.ServletRequestAttributes;
import org.springframework.web.context.request.ServletWebRequest;
import org.springframework.web.context.request.async.DeferredResult;
import org.springframework.web.context.request.async.WebAsyncUtils;
import org.springframework.web.filter.ShallowEtagHeaderFilter;
import org.springframework.web.method.support.HandlerMethodReturnValueHandler;
import org.springframework.web.method.support.ModelAndViewContainer;
import org.springframework.web.servlet.LocaleResolver;
import org.springframework.web.servlet.ModelAndView;
import org.springframework.web.servlet.View;
import org.springframework.web.servlet.ViewResolver;
import org.springframework.web.servlet.i18n.AcceptHeaderLocaleResolver;
import org.springframework.web.servlet.view.FragmentsRendering;
/**
* Handler for return values of type:
@@ -78,6 +97,10 @@ public class ResponseBodyEmitterReturnValueHandler implements HandlerMethodRetur
private final ReactiveTypeHandler reactiveHandler;
private final List<ViewResolver> viewResolvers;
private final LocaleResolver localeResolver;
/**
* Simple constructor with reactive type support based on a default instance of
@@ -86,13 +109,13 @@ public class ResponseBodyEmitterReturnValueHandler implements HandlerMethodRetur
* {@link ContentNegotiationManager} with an Accept header strategy.
*/
public ResponseBodyEmitterReturnValueHandler(List<HttpMessageConverter<?>> messageConverters) {
Assert.notEmpty(messageConverters, "HttpMessageConverter List must not be empty");
this.sseMessageConverters = initSseConverters(messageConverters);
this.reactiveHandler = new ReactiveTypeHandler();
this(messageConverters,
ReactiveAdapterRegistry.getSharedInstance(), new SyncTaskExecutor(),
new ContentNegotiationManager());
}
/**
* Complete constructor with pluggable "reactive" type support.
* Constructor that with added arguments to customize "reactive" type support.
* @param messageConverters converters to write emitted objects with
* @param registry for reactive return value type support
* @param executor for blocking I/O writes of items emitted from reactive types
@@ -102,9 +125,29 @@ public class ResponseBodyEmitterReturnValueHandler implements HandlerMethodRetur
public ResponseBodyEmitterReturnValueHandler(List<HttpMessageConverter<?>> messageConverters,
ReactiveAdapterRegistry registry, TaskExecutor executor, ContentNegotiationManager manager) {
this(messageConverters, registry, executor, manager, Collections.emptyList(), null);
}
/**
* Constructor that with added arguments for view rendering.
* @param messageConverters converters to write emitted objects with
* @param registry for reactive return value type support
* @param executor for blocking I/O writes of items emitted from reactive types
* @param manager for detecting streaming media types
* @param viewResolvers resolvers for fragment stream rendering
* @param localeResolver localeResolver for fragment stream rendering
* @since 6.2
*/
public ResponseBodyEmitterReturnValueHandler(
List<HttpMessageConverter<?>> messageConverters,
ReactiveAdapterRegistry registry, TaskExecutor executor, ContentNegotiationManager manager,
List<ViewResolver> viewResolvers, @Nullable LocaleResolver localeResolver) {
Assert.notEmpty(messageConverters, "HttpMessageConverter List must not be empty");
this.sseMessageConverters = initSseConverters(messageConverters);
this.reactiveHandler = new ReactiveTypeHandler(registry, executor, manager, null);
this.viewResolvers = viewResolvers;
this.localeResolver = (localeResolver != null ? localeResolver : new AcceptHeaderLocaleResolver());
}
private static List<HttpMessageConverter<?>> initSseConverters(List<HttpMessageConverter<?>> converters) {
@@ -156,7 +199,7 @@ public class ResponseBodyEmitterReturnValueHandler implements HandlerMethodRetur
}
}
ServletRequest request = webRequest.getNativeRequest(ServletRequest.class);
HttpServletRequest request = webRequest.getNativeRequest(HttpServletRequest.class);
Assert.state(request != null, "No ServletRequest");
ResponseBodyEmitter emitter;
@@ -180,21 +223,22 @@ public class ResponseBodyEmitterReturnValueHandler implements HandlerMethodRetur
// We are streaming
ShallowEtagHeaderFilter.disableContentCaching(request);
// Ignore further header changes; response is committed after first event
// Suppress header updates from message converters
outputMessage = new StreamingServletServerHttpResponse(outputMessage);
HttpMessageConvertingHandler handler;
DefaultSseEmitterHandler emitterHandler;
try {
DeferredResult<?> result = new DeferredResult<>(emitter.getTimeout());
WebAsyncUtils.getAsyncManager(webRequest).startDeferredResultProcessing(result, mavContainer);
handler = new HttpMessageConvertingHandler(outputMessage, result);
FragmentHandler handler = new FragmentHandler(request, response, this.viewResolvers, this.localeResolver);
emitterHandler = new DefaultSseEmitterHandler(this.sseMessageConverters, handler, outputMessage, result);
}
catch (Throwable ex) {
emitter.initializeWithError(ex);
throw ex;
}
emitter.initialize(handler);
emitter.initialize(emitterHandler);
}
@@ -221,15 +265,24 @@ public class ResponseBodyEmitterReturnValueHandler implements HandlerMethodRetur
/**
* ResponseBodyEmitter.Handler that writes with HttpMessageConverter's.
*/
private class HttpMessageConvertingHandler implements ResponseBodyEmitter.Handler {
private static final class DefaultSseEmitterHandler implements ResponseBodyEmitter.Handler {
private final List<HttpMessageConverter<?>> messageConverters;
private final FragmentHandler fragmentHandler;
private final ServerHttpResponse outputMessage;
private final DeferredResult<?> deferredResult;
public HttpMessageConvertingHandler(ServerHttpResponse outputMessage, DeferredResult<?> deferredResult) {
public DefaultSseEmitterHandler(
List<HttpMessageConverter<?>> messageConverters, FragmentHandler fragmentHandler,
ServerHttpResponse outputMessage, DeferredResult<?> result) {
this.messageConverters = messageConverters;
this.fragmentHandler = fragmentHandler;
this.outputMessage = outputMessage;
this.deferredResult = deferredResult;
this.deferredResult = result;
}
@Override
@@ -248,7 +301,11 @@ public class ResponseBodyEmitterReturnValueHandler implements HandlerMethodRetur
@SuppressWarnings("unchecked")
private <T> void sendInternal(T data, @Nullable MediaType mediaType) throws IOException {
for (HttpMessageConverter<?> converter : ResponseBodyEmitterReturnValueHandler.this.sseMessageConverters) {
if (data instanceof ModelAndView mav) {
this.fragmentHandler.handle(mav);
return;
}
for (HttpMessageConverter<?> converter : this.messageConverters) {
if (converter.canWrite(data.getClass(), mediaType)) {
((HttpMessageConverter<T>) converter).write(data, mediaType, this.outputMessage);
return;
@@ -289,4 +346,173 @@ public class ResponseBodyEmitterReturnValueHandler implements HandlerMethodRetur
}
}
/**
* Handler that renders ModelAndView fragments via FragmentsRendering.
*/
private static final class FragmentHandler {
private final HttpServletRequest request;
private final HttpServletResponse response;
private final List<ViewResolver> viewResolvers;
private final Locale locale;
private final Charset charset;
private final ServletRequestAttributes requestAttributes;
public FragmentHandler(
HttpServletRequest request, HttpServletResponse response,
List<ViewResolver> viewResolvers, LocaleResolver localeResolver) {
this.request = request;
this.response = response;
this.viewResolvers = viewResolvers;
this.charset = initCharset(response);
this.locale = localeResolver.resolveLocale(request);
this.requestAttributes = new ServletWebRequest(this.request);
}
private static Charset initCharset(HttpServletResponse response) {
String s = response.getHeader("Content-Type");
if (StringUtils.hasText(s)) {
MediaType contentType = MediaType.valueOf(s);
if (contentType.getCharset() != null) {
return contentType.getCharset();
}
}
return StandardCharsets.UTF_8;
}
public void handle(ModelAndView modelAndView) throws IOException {
RequestContextHolder.setRequestAttributes(this.requestAttributes);
try {
FragmentHttpServletResponse fragmentResponse =
new FragmentHttpServletResponse(this.response, this.charset);
FragmentsRendering render = FragmentsRendering.with(List.of(modelAndView)).build();
render.resolveNestedViews(this::resolveViewName, this.locale);
render.render(modelAndView.getModel(), this.request, fragmentResponse);
byte[] content = fragmentResponse.getFragmentContent();
this.response.getOutputStream().write(content);
}
catch (IOException ex) {
throw ex;
}
catch (Exception ex) {
throw new RuntimeException("Failed to render " + modelAndView, ex);
}
finally {
RequestContextHolder.resetRequestAttributes();
}
}
@Nullable
public View resolveViewName(String viewName, Locale locale) throws Exception {
for (ViewResolver resolver : this.viewResolvers) {
View view = resolver.resolveViewName(viewName, locale);
if (view != null) {
return view;
}
}
return null;
}
}
/**
* HttpServletResponse wrapper for fragment rendering.
* Ignores calls for setting content-type and content-length per fragment.
* Caches written content and replaces new lines.
*/
private static final class FragmentHttpServletResponse extends HttpServletResponseWrapper {
private final FragmentServletOutputStream outputStream;
private final PrintWriter writer;
private final Charset charset;
public FragmentHttpServletResponse(HttpServletResponse delegate, Charset charset) {
super(delegate);
this.outputStream = new FragmentServletOutputStream();
this.writer = new PrintWriter(this.outputStream);
this.charset = charset;
}
@Override
public void setContentType(String type) {
// ignore
}
@Override
public void setCharacterEncoding(String charset) {
// ignore
}
@Override
public void setContentLength(int len) {
// ignore
}
@Override
public ServletOutputStream getOutputStream() {
return this.outputStream;
}
@Override
public PrintWriter getWriter() {
return this.writer;
}
public byte[] getFragmentContent() {
this.writer.flush();
String content = this.outputStream.toString(this.charset);
content = content.replace("\n", "\ndata:");
return content.getBytes(this.charset);
}
}
/**
* ServletOutputStream that caches written fragment content.
*/
private static final class FragmentServletOutputStream extends ServletOutputStream {
private final ByteArrayOutputStream outputStream = new ByteArrayOutputStream();
@Override
public void write(int b) {
this.outputStream.write(b);
}
@Override
public void write(byte[] b) throws IOException {
this.outputStream.write(b);
}
@Override
public void write(byte[] b, int off, int len) {
this.outputStream.write(b, off, len);
}
@Override
public boolean isReady() {
return false;
}
@Override
public void setWriteListener(WriteListener writeListener) {
throw new UnsupportedOperationException();
}
public String toString(Charset charset) {
return this.outputStream.toString(charset);
}
}
}

View File

@@ -0,0 +1,127 @@
/*
* Copyright 2002-2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License 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 org.springframework.web.servlet.mvc.method.annotation;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.context.annotation.AnnotationConfigApplicationContext;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.context.support.ResourceBundleMessageSource;
import org.springframework.core.MethodParameter;
import org.springframework.core.ReactiveAdapterRegistry;
import org.springframework.core.task.SyncTaskExecutor;
import org.springframework.http.converter.json.MappingJackson2HttpMessageConverter;
import org.springframework.web.accept.ContentNegotiationManager;
import org.springframework.web.context.request.NativeWebRequest;
import org.springframework.web.context.request.ServletWebRequest;
import org.springframework.web.context.request.async.AsyncWebRequest;
import org.springframework.web.context.request.async.StandardServletAsyncWebRequest;
import org.springframework.web.context.request.async.WebAsyncUtils;
import org.springframework.web.method.support.ModelAndViewContainer;
import org.springframework.web.servlet.ModelAndView;
import org.springframework.web.servlet.view.script.ScriptTemplateConfigurer;
import org.springframework.web.servlet.view.script.ScriptTemplateViewResolver;
import org.springframework.web.testfixture.servlet.MockHttpServletRequest;
import org.springframework.web.testfixture.servlet.MockHttpServletResponse;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.web.testfixture.method.ResolvableMethod.on;
/**
* Tests for streaming of {@link ModelAndView} fragments.
* @author Rossen Stoyanchev
*/
public class FragmentRenderingStreamTests {
@Test
void streamFragments() throws Exception {
AnnotationConfigApplicationContext context =
new AnnotationConfigApplicationContext(ScriptTemplatingConfiguration.class);
String prefix = "org/springframework/web/servlet/view/script/kotlin/";
ScriptTemplateViewResolver viewResolver = new ScriptTemplateViewResolver(prefix, ".kts");
viewResolver.setApplicationContext(context);
ResponseBodyEmitterReturnValueHandler handler = new ResponseBodyEmitterReturnValueHandler(
List.of(new MappingJackson2HttpMessageConverter()),
ReactiveAdapterRegistry.getSharedInstance(), new SyncTaskExecutor(),
new ContentNegotiationManager(),
List.of(viewResolver), null);
MockHttpServletRequest request = new MockHttpServletRequest();
MockHttpServletResponse response = new MockHttpServletResponse();
NativeWebRequest webRequest = new ServletWebRequest(request, response);
AsyncWebRequest asyncWebRequest = new StandardServletAsyncWebRequest(request, response);
WebAsyncUtils.getAsyncManager(webRequest).setAsyncWebRequest(asyncWebRequest);
request.setAsyncSupported(true);
MethodParameter type = on(TestController.class).resolveReturnType(SseEmitter.class);
SseEmitter emitter = new SseEmitter();
handler.handleReturnValue(emitter, type, new ModelAndViewContainer(), webRequest);
assertThat(request.isAsyncStarted()).isTrue();
assertThat(response.getStatus()).isEqualTo(200);
ModelAndView mav1 = new ModelAndView("fragment1", Map.of("foo", "Foo"));
ModelAndView mav2 = new ModelAndView("fragment2", Map.of("bar", "Bar"));
emitter.send(SseEmitter.event().data(mav1).data(mav2));
assertThat(response.getContentType()).isEqualTo("text/event-stream");
assertThat(response.getContentAsString()).isEqualTo(("""
data:<p>Hello Foo</p>
data:<p>Hello Bar</p>
"""));
}
private static class TestController {
SseEmitter handle() {
return null;
}
}
@Configuration
static class ScriptTemplatingConfiguration {
@Bean
ScriptTemplateConfigurer kotlinScriptConfigurer() {
ScriptTemplateConfigurer configurer = new ScriptTemplateConfigurer();
configurer.setEngineName("kotlin");
configurer.setScripts("org/springframework/web/servlet/view/script/kotlin/render.kts");
configurer.setRenderFunction("render");
return configurer;
}
@Bean
ResourceBundleMessageSource messageSource() {
ResourceBundleMessageSource messageSource = new ResourceBundleMessageSource();
messageSource.setBasename("org/springframework/web/servlet/view/script/messages");
return messageSource;
}
}
}