exceptionHandler) {
- super(requestReader, responseWriter, securityContextWriter, exceptionHandler);
+ super(requestTypeClass, responseTypeClass, requestReader, responseWriter, securityContextWriter, exceptionHandler);
+ // set the default log formatter for servlet implementations
+ setLogFormatter(new ApacheCombinedServletLogFormatter<>());
+ setServletContext(new AwsServletContext(this));
+ }
+
+ //-------------------------------------------------------------
+ // Methods - Public
+ //-------------------------------------------------------------
+
+ /**
+ * You can use the onStartup to intercept the ServletContext as the Spring application is
+ * initialized and inject custom values. The StartupHandler is called after the onStartup method
+ * of the LambdaSpringApplicationinitializer implementation. For example, you can use this method to
+ * add custom filters to the servlet context:
+ *
+ *
+ * {@code
+ * handler = SpringLambdaContainerHandler.getAwsProxyHandler(EchoSpringAppConfig.class);
+ * handler.onStartup(c -> {
+ * // the "c" parameter to this function is the initialized servlet context
+ * c.addFilter("CustomHeaderFilter", CustomHeaderFilter.class);
+ * });
+ * }
+ *
+ * @param h A lambda expression that implements the StartupHandler functional interface
+ */
+ public void onStartup(final StartupHandler h) {
+ startupHandler = h;
+ startupHandler.onStartup(getServletContext());
}
@@ -86,29 +125,11 @@ protected void setServletContext(final ServletContext context) {
servletContext = context;
// We assume custom implementations of the RequestWriter for HttpServletRequest will reuse
// the existing AwsServletContext object since it has no dependencies other than the Lambda context
- filterChainManager = new AwsFilterChainManager((AwsServletContext)context);
+ filterChainManager = new AwsFilterChainManager((AwsServletContext)servletContext);
}
-
- /**
- * You can use the onStartup to intercept the ServletContext as the Spring application is
- * initialized and inject custom values. The StartupHandler is called after the onStartup method
- * of the LambdaSpringApplicationinitializer implementation. For example, you can use this method to
- * add custom filters to the servlet context:
- *
- *
- * {@code
- * handler = SpringLambdaContainerHandler.getAwsProxyHandler(EchoSpringAppConfig.class);
- * handler.onStartup(c -> {
- * // the "c" parameter to this function is the initialized servlet context
- * c.addFilter("CustomHeaderFilter", CustomHeaderFilter.class);
- * });
- * }
- *
- * @param h A lambda expression that implements the StartupHandler functional interface
- */
- public void onStartup(final StartupHandler h) {
- startupHandler = h;
+ protected FilterChain getFilterChain(HttpServletRequest req, Servlet servlet) {
+ return filterChainManager.getFilterChain(req, servlet);
}
@@ -120,12 +141,54 @@ public void onStartup(final StartupHandler h) {
* Applies the filter chain in the request lifecycle
* @param request The Request object. This must be an implementation of HttpServletRequest
* @param response The response object. This must be an implementation of HttpServletResponse
+ * @param servlet Servlet at the end of the chain (optional).
+ * @throws IOException
+ * @throws ServletException
*/
- protected void doFilter(ContainerRequestType request, ContainerResponseType response) throws IOException, ServletException {
- FilterChainHolder chain = filterChainManager.getFilterChain(request);
+ protected void doFilter(HttpServletRequest request, HttpServletResponse response, Servlet servlet) throws IOException, ServletException {
+ if (AwsHttpServletRequest.class.isAssignableFrom(request.getClass())) {
+ ((AwsHttpServletRequest)request).setContainerHandler(this);
+ }
+
+ FilterChain chain = getFilterChain(request, servlet);
chain.doFilter(request, response);
+ if(requiresAsyncReDispatch(request)) {
+ chain = getFilterChain(request, servlet);
+ chain.doFilter(request, response);
+ }
+ // if for some reason the response wasn't flushed yet, we force it here unless it's being processed asynchronously (WebFlux)
+ if (!response.isCommitted() && request.getDispatcherType() != DispatcherType.ASYNC) {
+ response.flushBuffer();
+ }
}
+ private boolean requiresAsyncReDispatch(HttpServletRequest request) {
+ if (request.isAsyncStarted()) {
+ AsyncContext asyncContext = request.getAsyncContext();
+ return asyncContext instanceof AwsAsyncContext
+ && ((AwsAsyncContext) asyncContext).isDispatchStarted();
+ }
+ return false;
+ }
+
+ @Override
+ public void initialize() throws ContainerInitializationException {
+ // we expect all servlets to be wrapped in an AwsServletRegistration
+ ArrayList registrations = new ArrayList<>((Collection)getServletContext().getServletRegistrations().values());
+ registrations.sort(AwsServletRegistration::compareTo);
+ for (AwsServletRegistration r : registrations) {
+ if (r.getLoadOnStartup() == -1) { // skip Servlets that can be lazily loaded
+ continue;
+ }
+ try {
+ if (r.getServlet() != null) {
+ r.getServlet().init(r.getServletConfig());
+ }
+ } catch (ServletException e) {
+ throw new ContainerInitializationException("Could not initialize servlet " + r.getName(), e);
+ }
+ }
+ }
//-------------------------------------------------------------
// Inner Class -
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletRequest.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletRequest.java
index c11111a7..c2a257d3 100644
--- a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletRequest.java
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletRequest.java
@@ -13,51 +13,37 @@
package com.amazonaws.serverless.proxy.internal.servlet;
-import com.amazonaws.serverless.proxy.internal.model.AwsProxyRequest;
+import com.amazonaws.serverless.proxy.internal.HttpUtils;
+import com.amazonaws.serverless.proxy.internal.LambdaContainerHandler;
+import com.amazonaws.serverless.proxy.internal.SecurityUtils;
+import com.amazonaws.serverless.proxy.model.AwsProxyRequest;
+import com.amazonaws.serverless.proxy.model.ContainerConfig;
+import com.amazonaws.serverless.proxy.model.Headers;
+import com.amazonaws.serverless.proxy.model.RequestSource;
import com.amazonaws.services.lambda.runtime.Context;
-
-import org.apache.commons.fileupload.FileItem;
-import org.apache.commons.fileupload.FileUploadException;
-import org.apache.commons.fileupload.servlet.ServletFileUpload;
-
-import javax.servlet.AsyncContext;
-import javax.servlet.ReadListener;
-import javax.servlet.RequestDispatcher;
-import javax.servlet.ServletException;
-import javax.servlet.ServletInputStream;
-import javax.servlet.ServletRequest;
-import javax.servlet.ServletResponse;
-import javax.servlet.http.Cookie;
-import javax.servlet.http.HttpServletResponse;
-import javax.servlet.http.HttpUpgradeHandler;
-import javax.servlet.http.Part;
-import javax.ws.rs.core.HttpHeaders;
-import javax.ws.rs.core.MediaType;
-import javax.ws.rs.core.SecurityContext;
+import edu.umd.cs.findbugs.annotations.SuppressFBWarnings;
+import jakarta.servlet.*;
+import jakarta.servlet.http.Cookie;
+import jakarta.servlet.http.HttpServletRequest;
+import jakarta.servlet.http.HttpServletResponse;
+import jakarta.servlet.http.HttpUpgradeHandler;
+import jakarta.ws.rs.core.HttpHeaders;
+import jakarta.ws.rs.core.SecurityContext;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
import java.io.BufferedReader;
-import java.io.ByteArrayInputStream;
import java.io.IOException;
import java.io.StringReader;
import java.io.UnsupportedEncodingException;
-import java.net.URLDecoder;
+import java.nio.charset.Charset;
import java.security.Principal;
-import java.text.ParseException;
-import java.text.SimpleDateFormat;
-import java.util.ArrayList;
-import java.util.Arrays;
-import java.util.Base64;
-import java.util.Collection;
-import java.util.Collections;
-import java.util.Date;
-import java.util.Enumeration;
-import java.util.HashMap;
-import java.util.Iterator;
-import java.util.List;
-import java.util.Locale;
-import java.util.Map;
-import java.util.TreeMap;
-
+import java.time.Instant;
+import java.time.ZonedDateTime;
+import java.time.format.DateTimeParseException;
+import java.util.*;
+import java.util.stream.Collectors;
+import java.util.stream.Stream;
/**
* Implementation of the HttpServletRequest interface that supports AwsProxyRequest object.
@@ -72,28 +58,36 @@ public class AwsProxyHttpServletRequest extends AwsHttpServletRequest {
private AwsProxyRequest request;
private SecurityContext securityContext;
- private Map> urlEncodedFormParameters;
- private Map multipartFormParameters;
-
+ private AwsAsyncContext asyncContext;
+ private static Logger log = LoggerFactory.getLogger(AwsProxyHttpServletRequest.class);
+ private ContainerConfig config;
//-------------------------------------------------------------
// Constructors
//-------------------------------------------------------------
+
public AwsProxyHttpServletRequest(AwsProxyRequest awsProxyRequest, Context lambdaContext, SecurityContext awsSecurityContext) {
+ this(awsProxyRequest, lambdaContext, awsSecurityContext, LambdaContainerHandler.getContainerConfig());
+ }
+
+
+ public AwsProxyHttpServletRequest(AwsProxyRequest awsProxyRequest, Context lambdaContext, SecurityContext awsSecurityContext, ContainerConfig config) {
super(lambdaContext);
this.request = awsProxyRequest;
this.securityContext = awsSecurityContext;
-
- this.urlEncodedFormParameters = getFormUrlEncodedParametersMap();
- this.multipartFormParameters = getMultipartFormParametersMap();
+ this.config = config;
}
+ public AwsProxyRequest getAwsProxyRequest() {
+ return this.request;
+ }
//-------------------------------------------------------------
// Implementation - HttpServletRequest
//-------------------------------------------------------------
+
@Override
public String getAuthType() {
return securityContext.getAuthenticationScheme();
@@ -102,7 +96,10 @@ public String getAuthType() {
@Override
public Cookie[] getCookies() {
- String cookieHeader = getHeaderCaseInsensitive(HttpHeaders.COOKIE);
+ if (request.getMultiValueHeaders() == null) {
+ return new Cookie[0];
+ }
+ String cookieHeader = request.getMultiValueHeaders().getFirst(HttpHeaders.COOKIE);
if (cookieHeader == null) {
return new Cookie[0];
}
@@ -112,50 +109,56 @@ public Cookie[] getCookies() {
@Override
public long getDateHeader(String s) {
- String dateString = getHeaderCaseInsensitive(HttpHeaders.DATE);
+ if (request.getMultiValueHeaders() == null) {
+ return -1L;
+ }
+ String dateString = request.getMultiValueHeaders().getFirst(s);
if (dateString == null) {
- return new Date().getTime();
+ return -1L;
}
- SimpleDateFormat dateFormatter = new SimpleDateFormat(HEADER_DATE_FORMAT);
try {
- return dateFormatter.parse(dateString).getTime();
- } catch (ParseException e) {
- e.printStackTrace();
- return new Date().getTime();
+ return Instant.from(ZonedDateTime.parse(dateString, dateFormatter)).toEpochMilli();
+ } catch (DateTimeParseException e) {
+ log.warn("Invalid date header in request" + SecurityUtils.crlf(dateString));
+ return -1L;
}
}
@Override
public String getHeader(String s) {
- return getHeaderCaseInsensitive(s);
+ List values = getHeaderValues(s);
+ if (values == null || values.size() == 0) {
+ return null;
+ }
+ return values.get(0);
}
@Override
public Enumeration getHeaders(String s) {
- String headerValue = getHeaderCaseInsensitive(s);
- if (headerValue == null) {
- return Collections.enumeration(new ArrayList());
+ if (request.getMultiValueHeaders() == null || request.getMultiValueHeaders().get(s) == null) {
+ return Collections.emptyEnumeration();
}
- List valueCollection = new ArrayList<>();
- valueCollection.add(headerValue);
- return Collections.enumeration(valueCollection);
+ return Collections.enumeration(request.getMultiValueHeaders().get(s));
}
@Override
public Enumeration getHeaderNames() {
- if (request.getHeaders() == null) {
+ if (request.getMultiValueHeaders() == null) {
return Collections.emptyEnumeration();
}
- return Collections.enumeration(request.getHeaders().keySet());
+ return Collections.enumeration(request.getMultiValueHeaders().keySet());
}
@Override
public int getIntHeader(String s) {
- String headerValue = getHeaderCaseInsensitive(s);
+ if (request.getMultiValueHeaders() == null) {
+ return -1;
+ }
+ String headerValue = request.getMultiValueHeaders().getFirst(s);
if (headerValue == null) {
return -1;
}
@@ -172,11 +175,8 @@ public String getMethod() {
@Override
public String getPathInfo() {
- String pathInfo = getServletPath().replace(getContextPath(), "");
- if (!pathInfo.startsWith("/")) {
- pathInfo = "/" + pathInfo;
- }
- return pathInfo;
+ String pathInfo = cleanUri(request.getPath());
+ return decodeRequestPath(pathInfo, LambdaContainerHandler.getContainerConfig());
}
@@ -189,13 +189,22 @@ public String getPathTranslated() {
@Override
public String getContextPath() {
- return request.getRequestContext().getStage();
+ return generateContextPath(config, request.getRequestContext().getStage());
}
@Override
public String getQueryString() {
- return this.generateQueryString(request.getQueryStringParameters());
+ try {
+ return this.generateQueryString(
+ request.getMultiValueQueryStringParameters(),
+ // ALB does not automatically decode parameters, so we don't want to re-encode them
+ request.getRequestSource() != RequestSource.ALB,
+ config.getUriEncoding());
+ } catch (ServletException e) {
+ log.error("Could not generate query string", e);
+ return null;
+ }
}
@@ -217,32 +226,19 @@ public Principal getUserPrincipal() {
return securityContext.getUserPrincipal();
}
+
@Override
public String getRequestURI() {
- return request.getPath();
+ return cleanUri(getContextPath()) + cleanUri(request.getPath());
}
@Override
public StringBuffer getRequestURL() {
- String url = "";
- url += getServerName();
- url += "/";
- url += getContextPath();
- url += "/";
- url += request.getPath();
-
- url = url.replaceAll("/+", "/");
-
- return new StringBuffer(getScheme() + "://" + url);
+ return generateRequestURL(request.getPath());
}
- @Override
- public String getServletPath() {
- return request.getPath();
- }
-
@Override
public boolean authenticate(HttpServletResponse httpServletResponse)
throws IOException, ServletException {
@@ -256,98 +252,53 @@ public void login(String s, String s1)
throw new UnsupportedOperationException();
}
+
@Override
public void logout()
throws ServletException {
throw new UnsupportedOperationException();
}
-
- @Override
- public Collection getParts()
- throws IOException, ServletException {
- return multipartFormParameters.values();
- }
-
-
- @Override
- public Part getPart(String s)
- throws IOException, ServletException {
- return multipartFormParameters.get(s);
- }
-
-
@Override
public T upgrade(Class aClass)
throws IOException, ServletException {
- return null;
+ throw new UnsupportedOperationException();
}
-
//-------------------------------------------------------------
// Implementation - ServletRequest
//-------------------------------------------------------------
+
@Override
public String getCharacterEncoding() {
- // we only look at content-type because content-encoding should only be used for
- // "binary" requests such as gzip/deflate.
- String contentTypeHeader = getHeaderCaseInsensitive(HttpHeaders.CONTENT_TYPE);
- if (contentTypeHeader == null) {
- return null;
- }
-
- String[] contentTypeValues = contentTypeHeader.split(HEADER_VALUE_SEPARATOR);
- if (contentTypeValues.length <= 1) {
- return null;
+ if (request.getMultiValueHeaders() == null) {
+ return config.getDefaultContentCharset();
}
-
- for (String contentTypeValue : contentTypeValues) {
- if (contentTypeValue.trim().startsWith(ENCODING_VALUE_KEY)) {
- String[] encodingValues = contentTypeValue.split(HEADER_KEY_VALUE_SEPARATOR);
- if (encodingValues.length <= 1) {
- return null;
- }
- return encodingValues[1];
- }
- }
- return null;
+ Charset charset = HttpUtils.parseCharacterEncoding(request.getMultiValueHeaders().getFirst(HttpHeaders.CONTENT_TYPE),null);
+ return charset != null ? charset.name() : null;
}
@Override
- public void setCharacterEncoding(String s) throws UnsupportedEncodingException {
- String currentContentType = request.getHeaders().get(HttpHeaders.CONTENT_TYPE);
- if (currentContentType == null) {
- request.getHeaders().put(
- HttpHeaders.CONTENT_TYPE,
- HEADER_VALUE_SEPARATOR + " " + ENCODING_VALUE_KEY + HEADER_KEY_VALUE_SEPARATOR + s);
+ public void setCharacterEncoding(String s)
+ throws UnsupportedEncodingException {
+ if (request.getMultiValueHeaders() == null) {
+ request.setMultiValueHeaders(new Headers());
+ }
+ String currentContentType = request.getMultiValueHeaders().getFirst(HttpHeaders.CONTENT_TYPE);
+ if (currentContentType == null || currentContentType.isEmpty()) {
+ log.debug("Called set character encoding to " + SecurityUtils.crlf(s) + " on a request without a content type. Character encoding will not be set");
return;
}
- if (currentContentType.contains(HEADER_VALUE_SEPARATOR)) {
- String[] contentTypeValues = currentContentType.split(HEADER_VALUE_SEPARATOR);
- StringBuilder contentType = new StringBuilder(contentTypeValues[0]);
-
- for (String contentTypeValue : contentTypeValues) {
- String contentTypeString = HEADER_VALUE_SEPARATOR + " " + contentTypeValue;
- if (contentTypeValue.trim().startsWith(ENCODING_VALUE_KEY)) {
- contentTypeString = HEADER_VALUE_SEPARATOR + " " + ENCODING_VALUE_KEY + HEADER_KEY_VALUE_SEPARATOR + s;
- }
- contentType.append(contentTypeString);
- }
-
- request.getHeaders().put(HttpHeaders.CONTENT_TYPE, contentType.toString());
- } else {
- request.getHeaders().put(
- HttpHeaders.CONTENT_TYPE,
- currentContentType + HEADER_VALUE_SEPARATOR + " " + ENCODING_VALUE_KEY + HEADER_KEY_VALUE_SEPARATOR + s);
- }
+ request.getMultiValueHeaders().putSingle(HttpHeaders.CONTENT_TYPE, HttpUtils.appendCharacterEncoding(currentContentType, s));
}
+
@Override
public int getContentLength() {
- String headerValue = getHeaderCaseInsensitive(HttpHeaders.CONTENT_LENGTH);
+ String headerValue = request.getMultiValueHeaders().getFirst(HttpHeaders.CONTENT_LENGTH);
if (headerValue == null) {
return -1;
}
@@ -357,7 +308,7 @@ public int getContentLength() {
@Override
public long getContentLengthLong() {
- String headerValue = getHeaderCaseInsensitive(HttpHeaders.CONTENT_LENGTH);
+ String headerValue = request.getMultiValueHeaders().getFirst(HttpHeaders.CONTENT_LENGTH);
if (headerValue == null) {
return -1;
}
@@ -367,162 +318,150 @@ public long getContentLengthLong() {
@Override
public String getContentType() {
- return getHeaderCaseInsensitive(HttpHeaders.CONTENT_TYPE);
- }
-
-
- @Override
- public ServletInputStream getInputStream() throws IOException {
- byte[] bodyBytes = request.getBody().getBytes();
- if (request.isBase64Encoded()) {
- bodyBytes = Base64.getDecoder().decode(request.getBody());
+ String contentTypeHeader = request.getMultiValueHeaders().getFirst(HttpHeaders.CONTENT_TYPE);
+ if (contentTypeHeader == null || "".equals(contentTypeHeader.trim())) {
+ return null;
}
- ByteArrayInputStream requestBodyStream = new ByteArrayInputStream(bodyBytes);
- return new ServletInputStream() {
-
- private ReadListener listener;
-
- @Override
- public boolean isFinished() {
- return true;
- }
-
-
- @Override
- public boolean isReady() {
- return true;
- }
-
-
- @Override
- public void setReadListener(ReadListener readListener) {
- listener = readListener;
- try {
- listener.onDataAvailable();
- } catch (IOException e) {
- e.printStackTrace();
- }
- }
-
- @Override
- public int read() throws IOException {
- int readByte = requestBodyStream.read();
- if (requestBodyStream.available() == 0 && listener != null) {
- listener.onAllDataRead();
- }
- return readByte;
- }
- };
+ return contentTypeHeader;
}
-
@Override
public String getParameter(String s) {
- String queryStringParameter = getQueryStringParameterCaseInsensitive(s);
- if (queryStringParameter != null) {
- return queryStringParameter;
- }
- String[] bodyParams = getFormBodyParameterCaseInsensitive(s);
- if (bodyParams == null || bodyParams.length == 0) {
- return null;
- } else {
- return bodyParams[0];
+ // decode key if ALB
+ if (request.getRequestSource() == RequestSource.ALB) {
+ s = decodeValueIfEncoded(s);
}
+
+ String queryStringParameter = getFirstQueryParamValue(request.getMultiValueQueryStringParameters(), s, config.isQueryStringCaseSensitive());
+ if (queryStringParameter != null) {
+ if (request.getRequestSource() == RequestSource.ALB) {
+ queryStringParameter = decodeValueIfEncoded(queryStringParameter);
+ }
+ return queryStringParameter;
+ }
+
+ String[] bodyParams = getFormBodyParameterCaseInsensitive(s);
+ if (bodyParams.length == 0) {
+ return null;
+ } else {
+ return bodyParams[0];
+ }
}
@Override
public Enumeration getParameterNames() {
- List paramNames = new ArrayList<>();
- if (request.getQueryStringParameters() != null) {
- paramNames.addAll(request.getQueryStringParameters().keySet());
+ Set formParameterNames = getFormUrlEncodedParametersMap().keySet();
+ if (request.getMultiValueQueryStringParameters() == null) {
+ return Collections.enumeration(formParameterNames);
+ }
+
+ Set paramNames = request.getMultiValueQueryStringParameters().keySet();
+ if (request.getRequestSource() == RequestSource.ALB) {
+ paramNames = paramNames.stream().map(AwsProxyHttpServletRequest::decodeValueIfEncoded).collect(Collectors.toSet());
}
- paramNames.addAll(urlEncodedFormParameters.keySet());
- return Collections.enumeration(paramNames);
+
+ return Collections.enumeration(
+ Stream.concat(formParameterNames.stream(), paramNames.stream())
+ .collect(Collectors.toSet()));
}
@Override
+ @SuppressFBWarnings("PZLA_PREFER_ZERO_LENGTH_ARRAYS") // suppressing this as according to the specs we should be returning null here if we can't find params
public String[] getParameterValues(String s) {
- List values = new ArrayList<>();
- String queryStringValue = getQueryStringParameterCaseInsensitive(s);
- if (queryStringValue != null) {
- values.add(queryStringValue);
+
+ // decode key if ALB
+ if (request.getRequestSource() == RequestSource.ALB) {
+ s = decodeValueIfEncoded(s);
}
- String[] formBodyValues = getFormBodyParameterCaseInsensitive(s);
- if (formBodyValues != null) {
- values.addAll(Arrays.asList(formBodyValues));
+ List values = getQueryParamValuesAsList(request.getMultiValueQueryStringParameters(), s, config.isQueryStringCaseSensitive());
+
+ // copy list so we don't modifying the underlying multi-value query params
+ if (values != null) {
+ values = new ArrayList<>(values);
+ } else {
+ values = new ArrayList<>();
+ }
+
+ // decode values if ALB
+ if (values != null && request.getRequestSource() == RequestSource.ALB) {
+ values = values.stream().map(AwsHttpServletRequest::decodeValueIfEncoded).collect(Collectors.toList());
}
+ values.addAll(Arrays.asList(getFormBodyParameterCaseInsensitive(s)));
+
if (values.size() == 0) {
return null;
} else {
- String[] valuesArray = new String[values.size()];
- valuesArray = values.toArray(valuesArray);
- return valuesArray;
+ return values.toArray(new String[0]);
}
}
@Override
public Map getParameterMap() {
- Map output = new HashMap<>();
-
- Map> params = urlEncodedFormParameters;
- if (params == null) {
- params = new HashMap<>();
- }
-
- if (request.getQueryStringParameters() != null) {
- for (Map.Entry entry : request.getQueryStringParameters().entrySet()) {
- if (params.containsKey(entry.getKey())) {
- params.get(entry.getKey()).add(entry.getValue());
- } else {
- List valueList = new ArrayList<>();
- valueList.add(entry.getValue());
- params.put(entry.getKey(), valueList);
- }
- }
- }
-
- for (Map.Entry> entry : params.entrySet()) {
- String[] valuesArray = new String[entry.getValue().size()];
- valuesArray = entry.getValue().toArray(valuesArray);
- output.put(entry.getKey(), valuesArray);
- }
- return output;
+ return generateParameterMap(request.getMultiValueQueryStringParameters(), config, request.getRequestSource() == RequestSource.ALB);
}
@Override
public String getProtocol() {
- // TODO: We should have a cloudfront protocol header
- return null;
+ return request.getRequestContext().getProtocol();
}
@Override
public String getScheme() {
- String headerValue = getHeaderCaseInsensitive(CF_PROTOCOL_HEADER_NAME);
- if (headerValue == null) {
- return "https";
- }
- return headerValue;
+ return getSchemeFromHeader(request.getMultiValueHeaders());
}
@Override
public String getServerName() {
- String name = getHeaderCaseInsensitive(HttpHeaders.HOST);
+ String region = System.getenv("AWS_REGION");
+ if (region == null) {
+ // this is not a critical failure, we just put a static region in the URI
+ region = "us-east-1";
+ }
+
+ if (request.getMultiValueHeaders() != null && request.getMultiValueHeaders().containsKey(HOST_HEADER_NAME)) {
+ String hostHeader = request.getMultiValueHeaders().getFirst(HOST_HEADER_NAME);
+ if (SecurityUtils.isValidHost(hostHeader, request.getRequestContext().getApiId(), request.getRequestContext().getElb(), region)) {
+ return hostHeader;
+ }
+ }
+
+ return new StringBuilder().append(request.getRequestContext().getApiId())
+ .append(".execute-api.")
+ .append(region)
+ .append(".amazonaws.com").toString();
+ }
+
+ @Override
+ public int getServerPort() {
+ if (request.getMultiValueHeaders() == null) {
+ return 443;
+ }
+ String port = request.getMultiValueHeaders().getFirst(PORT_HEADER_NAME);
+ if (SecurityUtils.isValidPort(port)) {
+ return Integer.parseInt(port);
+ } else {
+ return 443; // default port
+ }
+ }
- if (name == null || name.length() == 0) {
- name = "lambda.amazonaws.com";
+ @Override
+ public ServletInputStream getInputStream() throws IOException {
+ if (requestInputStream == null) {
+ requestInputStream = new AwsServletInputStream(bodyStringToInputStream(request.getBody(), request.isBase64Encoded()));
}
- return name;
+ return requestInputStream;
}
+
@Override
public BufferedReader getReader()
throws IOException {
@@ -532,45 +471,45 @@ public BufferedReader getReader()
@Override
public String getRemoteAddr() {
+ if (request.getRequestContext() == null || request.getRequestContext().getIdentity() == null) {
+ return "127.0.0.1";
+ }
+ if (request.getRequestSource().equals(RequestSource.ALB)) {
+ return Objects.nonNull(request.getHeaders()) ?
+ request.getHeaders().get(CLIENT_IP_HEADER) :
+ request.getMultiValueHeaders().getFirst(CLIENT_IP_HEADER);
+ }
return request.getRequestContext().getIdentity().getSourceIp();
}
@Override
public String getRemoteHost() {
- return getHeaderCaseInsensitive(HttpHeaders.HOST);
+ String hostHeader;
+ if (request.getRequestSource().equals(RequestSource.ALB)) {
+ hostHeader = Objects.nonNull(request.getHeaders()) ?
+ request.getHeaders().get(HttpHeaders.HOST) :
+ request.getMultiValueHeaders().getFirst(HttpHeaders.HOST);
+ } else {
+ hostHeader = request.getMultiValueHeaders().getFirst(HttpHeaders.HOST);
+ }
+ // the host header has the form host:port, so we split the string to get the host part
+ return Arrays.asList(hostHeader.split(":")).get(0);
}
+
@Override
public Locale getLocale() {
- List> values = this.parseHeaderValue(
- getHeaderCaseInsensitive(HttpHeaders.ACCEPT_LANGUAGE)
- );
- if (values.size() == 0) {
- return Locale.getDefault();
- }
- return new Locale(values.get(0).getValue());
+ List locales = parseAcceptLanguageHeader(request.getMultiValueHeaders().getFirst(HttpHeaders.ACCEPT_LANGUAGE));
+ return locales.size() == 0 ? Locale.getDefault() : locales.get(0);
}
-
@Override
public Enumeration getLocales() {
- List> values = this.parseHeaderValue(
- getHeaderCaseInsensitive(HttpHeaders.ACCEPT_LANGUAGE)
- );
- List locales = new ArrayList<>();
- if (values.size() == 0) {
- locales.add(Locale.getDefault());
- } else {
- for (Map.Entry locale : values) {
- locales.add(new Locale(locale.getValue()));
- }
- }
-
+ List locales = parseAcceptLanguageHeader(request.getMultiValueHeaders().getFirst(HttpHeaders.ACCEPT_LANGUAGE));
return Collections.enumeration(locales);
}
-
@Override
public boolean isSecure() {
return securityContext.isSecure();
@@ -579,141 +518,120 @@ public boolean isSecure() {
@Override
public RequestDispatcher getRequestDispatcher(String s) {
- return null;
+ return getServletContext().getRequestDispatcher(s);
}
- @Override
- public String getRealPath(String s) {
- // we are in an archive on a remote server
- return null;
- }
-
@Override
public int getRemotePort() {
+ if (request.getRequestSource().equals(RequestSource.ALB)) {
+ String portHeader;
+ portHeader = Objects.nonNull(request.getHeaders()) ?
+ request.getHeaders().get(PORT_HEADER_NAME) :
+ request.getMultiValueHeaders().getFirst(PORT_HEADER_NAME);
+ if (Objects.nonNull(portHeader)) {
+ return Integer.parseInt(portHeader);
+ }
+ }
return 0;
}
+
@Override
- public AsyncContext startAsync() throws IllegalStateException {
- return null;
+ public boolean isAsyncSupported() {
+ return true;
}
-
@Override
- public AsyncContext startAsync(ServletRequest servletRequest, ServletResponse servletResponse) throws IllegalStateException {
- return null;
+ public boolean isAsyncStarted() {
+ if (asyncContext == null) {
+ return false;
+ }
+ if (asyncContext.isCompleted() || asyncContext.isDispatched()) {
+ return false;
+ }
+ return true;
}
- //-------------------------------------------------------------
- // Methods - Private
- //-------------------------------------------------------------
- private String getHeaderCaseInsensitive(String key) {
- if (request.getHeaders() == null) {
- return null;
- }
- for (String requestHeaderKey : request.getHeaders().keySet()) {
- if (key.toLowerCase().equals(requestHeaderKey.toLowerCase())) {
- return request.getHeaders().get(requestHeaderKey);
- }
- }
- return null;
+ @Override
+ public AsyncContext startAsync()
+ throws IllegalStateException {
+ asyncContext = new AwsAsyncContext(this, response);
+ setAttribute(DISPATCHER_TYPE_ATTRIBUTE, DispatcherType.ASYNC);
+ log.debug("Starting async context for request: " + SecurityUtils.crlf(request.getRequestContext().getRequestId()));
+ return asyncContext;
}
- private String getQueryStringParameterCaseInsensitive(String key) {
- if (request.getQueryStringParameters() == null) {
- return null;
- }
+ @Override
+ public AsyncContext startAsync(ServletRequest servletRequest, ServletResponse servletResponse)
+ throws IllegalStateException {
+ servletRequest.setAttribute(DISPATCHER_TYPE_ATTRIBUTE, DispatcherType.ASYNC);
+ asyncContext = new AwsAsyncContext((HttpServletRequest) servletRequest, (HttpServletResponse) servletResponse);
+ log.debug("Starting async context for request: " + SecurityUtils.crlf(request.getRequestContext().getRequestId()));
+ return asyncContext;
+ }
- for (String requestParamKey : request.getQueryStringParameters().keySet()) {
- if (key.toLowerCase().equals(requestParamKey.toLowerCase())) {
- return request.getQueryStringParameters().get(requestParamKey);
- }
+ @Override
+ public AsyncContext getAsyncContext() {
+ if (asyncContext == null) {
+ throw new IllegalStateException("Request " + SecurityUtils.crlf(request.getRequestContext().getRequestId())
+ + " is not in asynchronous mode. Call startAsync before attempting to get the async context.");
}
- return null;
+ return asyncContext;
}
+ @Override
+ public String getRequestId() {
+ return request.getRequestContext().getRequestId();
+ }
- private String[] getFormBodyParameterCaseInsensitive(String key) {
- List values = urlEncodedFormParameters.get(key);
- if (values != null) {
- String[] valuesArray = new String[values.size()];
- valuesArray = values.toArray(valuesArray);
- return valuesArray;
- } else {
- return null;
- }
- }
+ @Override
+ public String getProtocolRequestId() {
+ return "";
+ }
+ @Override
+ public ServletConnection getServletConnection() {
+ return null;
+ }
- private Map getMultipartFormParametersMap() {
- if (!ServletFileUpload.isMultipartContent(this)) { // isMultipartContent also checks the content type
- return new HashMap<>();
- }
+ //-------------------------------------------------------------
+ // Methods - Private
+ //-------------------------------------------------------------
- Map output = new TreeMap<>(String.CASE_INSENSITIVE_ORDER);
+ private List getHeaderValues(String key) {
+ // special cases for referer and user agent headers
+ List values = new ArrayList<>();
- ServletFileUpload upload = new ServletFileUpload();
- try {
- List items = upload.parseRequest(this);
- for (FileItem item : items) {
- AwsProxyRequestPart newPart = new AwsProxyRequestPart(item.get());
- newPart.setName(item.getName());
- newPart.setSubmittedFileName(item.getFieldName());
- newPart.setContentType(item.getContentType());
- newPart.setSize(item.getSize());
-
- Iterator headerNamesIterator = item.getHeaders().getHeaderNames();
- while (headerNamesIterator.hasNext()) {
- String headerName = headerNamesIterator.next();
- Iterator headerValuesIterator = item.getHeaders().getHeaders(headerName);
- while (headerValuesIterator.hasNext()) {
- newPart.addHeader(headerName, headerValuesIterator.next());
+ if (request.getRequestSource() == RequestSource.API_GATEWAY) {
+ if ("referer".equals(key.toLowerCase(Locale.ENGLISH))) {
+ if (request.getRequestContext() != null && request.getRequestContext().getIdentity() != null) {
+ String caller = request.getRequestContext().getIdentity().getCaller();
+ if (caller != null) {
+ values.add(caller);
+ return values;
+ }
+ }
+ }
+ if ("user-agent".equals(key.toLowerCase(Locale.ENGLISH))) {
+ if (request.getRequestContext() != null && request.getRequestContext().getIdentity() != null) {
+ String userAgent = request.getRequestContext().getIdentity().getUserAgent();
+ if (userAgent != null) {
+ values.add(userAgent);
+ return values;
}
}
-
- output.put(item.getFieldName(), newPart);
}
- } catch (FileUploadException e) {
- // TODO: Should we swallaw this?
- e.printStackTrace();
- }
- return output;
- }
-
-
- private Map> getFormUrlEncodedParametersMap() {
- String contentType = getContentType();
- if (contentType == null) {
- return new HashMap<>();
- }
- if (!contentType.startsWith(MediaType.APPLICATION_FORM_URLENCODED) || !getMethod().toLowerCase().equals("post")) {
- return new HashMap<>();
- }
- String rawBodyContent;
- try {
- rawBodyContent = URLDecoder.decode(request.getBody(), DEFAULT_CHARACTER_ENCODING);
- } catch (UnsupportedEncodingException e) {
- e.printStackTrace();
- rawBodyContent = request.getBody();
}
- Map> output = new TreeMap<>(String.CASE_INSENSITIVE_ORDER);
- for (String parameter : rawBodyContent.split(FORM_DATA_SEPARATOR)) {
- String[] parameterKeyValue = parameter.split(HEADER_KEY_VALUE_SEPARATOR);
- if (parameterKeyValue.length < 2) {
- continue;
- }
- List values = new ArrayList<>();
- if (output.containsKey(parameterKeyValue[0])) {
- values = output.get(parameterKeyValue[0]);
- }
- values.add(parameterKeyValue[1]);
- output.put(parameterKeyValue[0], values);
+ if (request.getMultiValueHeaders() == null) {
+ return null;
}
- return output;
+ return request.getMultiValueHeaders().get(key);
}
+
+
}
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletRequestReader.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletRequestReader.java
index 3358e270..ec56285f 100644
--- a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletRequestReader.java
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletRequestReader.java
@@ -13,35 +13,87 @@
package com.amazonaws.serverless.proxy.internal.servlet;
import com.amazonaws.serverless.exceptions.InvalidRequestEventException;
-import com.amazonaws.serverless.proxy.internal.RequestReader;
-import com.amazonaws.serverless.proxy.internal.model.AwsProxyRequest;
+import com.amazonaws.serverless.proxy.RequestReader;
+import com.amazonaws.serverless.proxy.model.AwsProxyRequest;
+import com.amazonaws.serverless.proxy.model.ContainerConfig;
import com.amazonaws.services.lambda.runtime.Context;
-import javax.ws.rs.core.SecurityContext;
+import jakarta.servlet.ServletContext;
+import jakarta.servlet.http.HttpServletRequest;
+import jakarta.ws.rs.core.HttpHeaders;
+import jakarta.ws.rs.core.SecurityContext;
/**
* Simple implementation of the RequestReader interface that receives an AwsProxyRequest
* object and uses it to initialize a AwsProxyHttpServletRequest object.
*/
-public class AwsProxyHttpServletRequestReader extends RequestReader {
+public class AwsProxyHttpServletRequestReader extends RequestReader {
+ static final String INVALID_REQUEST_ERROR = "The incoming event is not a valid request from Amazon API Gateway or an Application Load Balancer";
+ private ServletContext servletContext;
//-------------------------------------------------------------
// Methods - Implementation
//-------------------------------------------------------------
+ public void setServletContext(ServletContext ctx) {
+ servletContext = ctx;
+ }
+
@Override
- public AwsProxyHttpServletRequest readRequest(AwsProxyRequest request, SecurityContext securityContext, Context lambdaContext)
+ public HttpServletRequest readRequest(AwsProxyRequest request, SecurityContext securityContext, Context lambdaContext, ContainerConfig config)
throws InvalidRequestEventException {
- AwsProxyHttpServletRequest servletRequest = new AwsProxyHttpServletRequest(request, lambdaContext, securityContext);
+ // Expect the HTTP method and context to be populated. If they are not, we are handling an
+ // unsupported event type.
+ if (request.getHttpMethod() == null || request.getHttpMethod().equals("") || request.getRequestContext() == null) {
+ throw new InvalidRequestEventException(INVALID_REQUEST_ERROR);
+ }
+
+ request.setPath(stripBasePath(request.getPath(), config));
+ if (request.getMultiValueHeaders() != null && request.getMultiValueHeaders().getFirst(HttpHeaders.CONTENT_TYPE) != null) {
+ String contentType = request.getMultiValueHeaders().getFirst(HttpHeaders.CONTENT_TYPE);
+ // put single as we always expect to have one and only one content type in a request.
+ request.getMultiValueHeaders().putSingle(HttpHeaders.CONTENT_TYPE, getContentTypeWithCharset(contentType, config));
+ }
+ AwsProxyHttpServletRequest servletRequest = new AwsProxyHttpServletRequest(request, lambdaContext, securityContext, config);
+ servletRequest.setServletContext(servletContext);
servletRequest.setAttribute(API_GATEWAY_CONTEXT_PROPERTY, request.getRequestContext());
servletRequest.setAttribute(API_GATEWAY_STAGE_VARS_PROPERTY, request.getStageVariables());
+ servletRequest.setAttribute(API_GATEWAY_EVENT_PROPERTY, request);
+ servletRequest.setAttribute(ALB_CONTEXT_PROPERTY, request.getRequestContext().getElb());
servletRequest.setAttribute(LAMBDA_CONTEXT_PROPERTY, lambdaContext);
+ servletRequest.setAttribute(JAX_SECURITY_CONTEXT_PROPERTY, securityContext);
+
return servletRequest;
}
+ //-------------------------------------------------------------
+ // Methods - Protected
+ //-------------------------------------------------------------
@Override
protected Class extends AwsProxyRequest> getRequestClass() {
return AwsProxyRequest.class;
}
+
+ //-------------------------------------------------------------
+ // Methods - Private
+ //-------------------------------------------------------------
+
+ private String getContentTypeWithCharset(String headerValue, ContainerConfig config) {
+ if (headerValue == null || "".equals(headerValue.trim())) {
+ return headerValue;
+ }
+
+ if (headerValue.contains("charset=")) {
+ return headerValue;
+ }
+
+ String newValue = headerValue;
+ if (!headerValue.trim().endsWith(";")) {
+ newValue += "; ";
+ }
+
+ newValue += "charset=" + config.getDefaultContentCharset();
+ return newValue;
+ }
}
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletResponseWriter.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletResponseWriter.java
index 232bd15e..0e230e9a 100644
--- a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletResponseWriter.java
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyHttpServletResponseWriter.java
@@ -14,12 +14,21 @@
import com.amazonaws.serverless.exceptions.InvalidResponseObjectException;
-import com.amazonaws.serverless.proxy.internal.ResponseWriter;
-import com.amazonaws.serverless.proxy.internal.model.AwsProxyResponse;
+import com.amazonaws.serverless.proxy.ResponseWriter;
+import com.amazonaws.serverless.proxy.internal.LambdaContainerHandler;
+import com.amazonaws.serverless.proxy.internal.testutils.Timer;
+import com.amazonaws.serverless.proxy.model.AwsProxyResponse;
+import com.amazonaws.serverless.proxy.model.Headers;
+import com.amazonaws.serverless.proxy.model.RequestSource;
import com.amazonaws.services.lambda.runtime.Context;
-import java.io.IOException;
+import jakarta.ws.rs.core.Response;
+import jakarta.ws.rs.core.Response.Status;
+
import java.util.Base64;
+import java.util.HashMap;
+import java.util.Map;
+
/**
* Creates an AwsProxyResponse object given an AwsHttpServletResponse object. If the
@@ -27,6 +36,16 @@
*/
public class AwsProxyHttpServletResponseWriter extends ResponseWriter {
+ private boolean writeSingleValueHeaders;
+
+ public AwsProxyHttpServletResponseWriter() {
+ this(false);
+ }
+
+ public AwsProxyHttpServletResponseWriter(boolean singleValueHeaders) {
+ writeSingleValueHeaders = singleValueHeaders;
+ }
+
//-------------------------------------------------------------
// Methods - Implementation
//-------------------------------------------------------------
@@ -34,27 +53,59 @@ public class AwsProxyHttpServletResponseWriter extends ResponseWriter toSingleValueHeaders(Headers h) {
+ Map out = new HashMap<>();
+ if (h == null || h.isEmpty()) {
+ return out;
+ }
+ for (String k : h.keySet()) {
+ out.put(k, h.getFirst(k));
+ }
+ return out;
+ }
+
+ private boolean isBinary(String contentType) {
+ if(contentType != null) {
+ int semidx = contentType.indexOf(';');
+ if(semidx >= 0) {
+ return LambdaContainerHandler.getContainerConfig().isBinaryContentType(contentType.substring(0, semidx));
+ }
+ else {
+ return LambdaContainerHandler.getContainerConfig().isBinaryContentType(contentType);
+ }
+ }
+ return false;
+ }
}
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyRequestDispatcher.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyRequestDispatcher.java
new file mode 100644
index 00000000..f314ba0d
--- /dev/null
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyRequestDispatcher.java
@@ -0,0 +1,152 @@
+/*
+ * Copyright 2020 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 com.amazonaws.serverless.proxy.internal.servlet;
+
+import com.amazonaws.serverless.proxy.internal.SecurityUtils;
+import com.amazonaws.serverless.proxy.model.AwsProxyRequest;
+import com.amazonaws.serverless.proxy.model.HttpApiV2ProxyRequest;
+import edu.umd.cs.findbugs.annotations.SuppressFBWarnings;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+
+import jakarta.servlet.*;
+import jakarta.servlet.http.HttpServletRequest;
+import jakarta.servlet.http.HttpServletResponse;
+
+import java.io.IOException;
+
+import static com.amazonaws.serverless.proxy.RequestReader.API_GATEWAY_EVENT_PROPERTY;
+import static com.amazonaws.serverless.proxy.RequestReader.HTTP_API_EVENT_PROPERTY;
+import static com.amazonaws.serverless.proxy.internal.servlet.AwsHttpServletRequest.DISPATCHER_TYPE_ATTRIBUTE;
+
+/**
+ * Default RequestDispatcher implementation for the AwsProxyHttpServletRequest type. A new
+ * instance of this object is created each time a framework gets the RequestDispatcher from a servlet request. Behind
+ * the scenes, this object uses the AwsLambdaServletContainerHandler to send FORWARD and INCLUDE requests
+ * to the framework.
+ */
+public class AwsProxyRequestDispatcher implements RequestDispatcher {
+
+ //-------------------------------------------------------------
+ // Variables - Private
+ //-------------------------------------------------------------
+ private static final Logger log = LoggerFactory.getLogger(AwsHttpSession.class);
+ private String dispatchTo;
+ private boolean isNamedDispatcher;
+ private AwsLambdaServletContainerHandler lambdaContainerHandler;
+
+ //-------------------------------------------------------------
+ // Constructors
+ //-------------------------------------------------------------
+
+
+ public AwsProxyRequestDispatcher(final String target, final boolean namedDispatcher, final AwsLambdaServletContainerHandler handler) {
+ isNamedDispatcher = namedDispatcher;
+ dispatchTo = target;
+ lambdaContainerHandler = handler;
+ }
+
+ //-------------------------------------------------------------
+ // Implementation - RequestDispatcher
+ //-------------------------------------------------------------
+
+
+ @Override
+ @SuppressWarnings("unchecked")
+ public void forward(ServletRequest servletRequest, ServletResponse servletResponse)
+ throws ServletException, IOException {
+ if (lambdaContainerHandler == null) {
+ throw new IllegalStateException("Null container handler in dispatcher");
+ }
+ if (servletResponse.isCommitted()) {
+ throw new IllegalStateException("Cannot forward request with committed response");
+ }
+
+ try {
+ // Reset any output that has been buffered, but keep headers/cookies
+ servletResponse.resetBuffer();
+ } catch (IllegalStateException e) {
+ throw e;
+ }
+
+ if (isNamedDispatcher) {
+ lambdaContainerHandler.doFilter((HttpServletRequest) servletRequest, (HttpServletResponse) servletResponse, ((AwsServletRegistration)servletRequest.getServletContext().getServletRegistration(dispatchTo)).getServlet());
+ return;
+ }
+
+ servletRequest.setAttribute(DISPATCHER_TYPE_ATTRIBUTE, DispatcherType.FORWARD);
+ setRequestPath(servletRequest, dispatchTo);
+ lambdaContainerHandler.doFilter((HttpServletRequest) servletRequest, (HttpServletResponse) servletResponse, getServlet((HttpServletRequest)servletRequest));
+ }
+
+
+ @Override
+ @SuppressWarnings("unchecked")
+ @SuppressFBWarnings("SERVLET_QUERY_STRING")
+ public void include(ServletRequest servletRequest, ServletResponse servletResponse)
+ throws ServletException, IOException {
+ if (lambdaContainerHandler == null) {
+ throw new IllegalStateException("Null container handler in dispatcher");
+ }
+ if (servletResponse.isCommitted()) {
+ throw new IllegalStateException("Cannot forward request with committed response");
+ }
+ servletRequest.setAttribute(DISPATCHER_TYPE_ATTRIBUTE, DispatcherType.INCLUDE);
+ if (!isNamedDispatcher) {
+ servletRequest.setAttribute("javax.servlet.include.request_uri", ((HttpServletRequest)servletRequest).getRequestURI());
+ servletRequest.setAttribute("javax.servlet.include.context_path", ((HttpServletRequest) servletRequest).getContextPath());
+ servletRequest.setAttribute("javax.servlet.include.servlet_path", ((HttpServletRequest) servletRequest).getServletPath());
+ servletRequest.setAttribute("javax.servlet.include.path_info", ((HttpServletRequest) servletRequest).getPathInfo());
+ servletRequest.setAttribute("javax.servlet.include.query_string",
+ SecurityUtils.encode(SecurityUtils.crlf(((HttpServletRequest) servletRequest).getQueryString())));
+ setRequestPath(servletRequest, dispatchTo);
+ }
+ lambdaContainerHandler.doFilter((HttpServletRequest) servletRequest, (HttpServletResponse) servletResponse, getServlet((HttpServletRequest)servletRequest));
+ }
+
+ /**
+ * Sets the destination path in the given request. Uses the AwsProxyRequest.setPath method which
+ * is in turn read by the HttpServletRequest implementation.
+ * @param req The request object to be modified
+ * @param destinationPath The new path for the request
+ * @throws IllegalStateException If the given request object does not include the API_GATEWAY_EVENT_PROPERTY
+ * attribute or the value for the attribute is not of the correct type: AwsProxyRequest.
+ */
+ void setRequestPath(ServletRequest req, final String destinationPath) {
+ if (req instanceof AwsProxyHttpServletRequest) {
+ ((AwsProxyHttpServletRequest) req).getAwsProxyRequest().setPath(dispatchTo);
+ return;
+ }
+ if (req instanceof AwsHttpApiV2ProxyHttpServletRequest) {
+ ((AwsHttpApiV2ProxyHttpServletRequest) req).getRequest().setRawPath(destinationPath);
+ return;
+ }
+
+ log.debug("Request is not an proxy request generated by this library, attempting to extract the proxy event type from the request attributes");
+ if (req.getAttribute(API_GATEWAY_EVENT_PROPERTY) != null && req.getAttribute(API_GATEWAY_EVENT_PROPERTY) instanceof AwsProxyRequest) {
+ ((AwsProxyRequest)req.getAttribute(API_GATEWAY_EVENT_PROPERTY)).setPath(dispatchTo);
+ return;
+ }
+ if (req.getAttribute(HTTP_API_EVENT_PROPERTY) != null && req.getAttribute(HTTP_API_EVENT_PROPERTY) instanceof HttpApiV2ProxyRequest) {
+ ((HttpApiV2ProxyRequest)req.getAttribute(HTTP_API_EVENT_PROPERTY)).setRawPath(destinationPath);
+ return;
+ }
+
+ throw new IllegalStateException("Could not set new target path for the given ServletRequest object");
+ }
+
+ private Servlet getServlet(HttpServletRequest req) {
+ return ((AwsServletContext)lambdaContainerHandler.getServletContext()).getServletForPath(req.getPathInfo());
+ }
+
+}
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyRequestPart.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyRequestPart.java
index 18012850..c2d1fa62 100644
--- a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyRequestPart.java
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsProxyRequestPart.java
@@ -12,13 +12,19 @@
*/
package com.amazonaws.serverless.proxy.internal.servlet;
-import javax.servlet.http.Part;
-import javax.ws.rs.core.MultivaluedHashMap;
+import com.amazonaws.serverless.proxy.internal.SecurityUtils;
+import com.amazonaws.serverless.proxy.model.MultiValuedTreeMap;
+
+import edu.umd.cs.findbugs.annotations.SuppressFBWarnings;
+
+import jakarta.servlet.http.Part;
import java.io.ByteArrayInputStream;
import java.io.FileOutputStream;
import java.io.IOException;
import java.io.InputStream;
import java.util.Collection;
+import java.util.Collections;
+
public class AwsProxyRequestPart
implements Part {
@@ -31,7 +37,7 @@ public class AwsProxyRequestPart
private String submittedFileName;
private long size;
private String contentType;
- private MultivaluedHashMap headers;
+ private MultiValuedTreeMap headers;
private byte[] content;
@@ -40,7 +46,7 @@ public class AwsProxyRequestPart
//-------------------------------------------------------------
public AwsProxyRequestPart(byte[] content) {
- this.content = content;
+ this.content = content.clone();
}
@@ -79,9 +85,11 @@ public long getSize() {
}
+ @SuppressFBWarnings("PATH_TRAVERSAL_OUT")
@Override
public void write(String s) throws IOException {
- FileOutputStream fos = new FileOutputStream(s);
+ String canonicalFilePath = SecurityUtils.getValidFilePath(s);
+ FileOutputStream fos = new FileOutputStream(canonicalFilePath);
try {
fos.write(content);
} finally {
@@ -98,18 +106,27 @@ public void delete() throws IOException {
@Override
public String getHeader(String s) {
+ if (headers == null) {
+ return null;
+ }
return headers.getFirst(s);
}
@Override
public Collection getHeaders(String s) {
+ if (headers == null) {
+ return Collections.emptyList();
+ }
return headers.get(s);
}
@Override
public Collection getHeaderNames() {
+ if (headers == null) {
+ return Collections.emptyList();
+ }
return headers.keySet();
}
@@ -120,7 +137,7 @@ public Collection getHeaderNames() {
public void addHeader(String key, String value) {
if (headers == null) {
- headers = new MultivaluedHashMap<>();
+ headers = new MultiValuedTreeMap<>(String.CASE_INSENSITIVE_ORDER);
}
if (headers.containsKey(key)) {
@@ -155,7 +172,7 @@ public void setContentType(String contentType) {
}
- public void setHeaders(MultivaluedHashMap headers) {
+ public void setHeaders(MultiValuedTreeMap headers) {
this.headers = headers;
}
}
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsServletContext.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsServletContext.java
index fdadce13..94dcaf44 100644
--- a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsServletContext.java
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsServletContext.java
@@ -14,18 +14,16 @@
import com.amazonaws.serverless.proxy.internal.LambdaContainerHandler;
-import com.amazonaws.services.lambda.runtime.Context;
-
-import javax.servlet.Filter;
-import javax.servlet.FilterRegistration;
-import javax.servlet.RequestDispatcher;
-import javax.servlet.Servlet;
-import javax.servlet.ServletContext;
-import javax.servlet.ServletException;
-import javax.servlet.ServletRegistration;
-import javax.servlet.SessionCookieConfig;
-import javax.servlet.SessionTrackingMode;
-import javax.servlet.descriptor.JspConfigDescriptor;
+import com.amazonaws.serverless.proxy.internal.SecurityUtils;
+
+import edu.umd.cs.findbugs.annotations.SuppressFBWarnings;
+
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+
+import jakarta.servlet.*;
+import jakarta.servlet.ServletContext;
+import jakarta.servlet.descriptor.JspConfigDescriptor;
import java.io.File;
import java.io.IOException;
@@ -33,15 +31,11 @@
import java.net.MalformedURLException;
import java.net.URISyntaxException;
import java.net.URL;
+import java.net.URLConnection;
import java.nio.file.Files;
+import java.nio.file.InvalidPathException;
import java.nio.file.Paths;
-import java.util.Collections;
-import java.util.Enumeration;
-import java.util.EventListener;
-import java.util.HashMap;
-import java.util.LinkedHashMap;
-import java.util.Map;
-import java.util.Set;
+import java.util.*;
/**
@@ -55,8 +49,8 @@ public class AwsServletContext
//-------------------------------------------------------------
// Constants - Public
// -------------------------------------------------------------
- public static final int SERVLET_API_MAJOR_VERSION = 3;
- public static final int SERVLET_API_MINOR_VERSION = 1;
+ public static final int SERVLET_API_MAJOR_VERSION = 6;
+ public static final int SERVLET_API_MINOR_VERSION = 0;
public static final String SERVER_INFO = LambdaContainerHandler.SERVER_INFO + "/" + SERVLET_API_MAJOR_VERSION + "." + SERVLET_API_MINOR_VERSION;
@@ -64,9 +58,11 @@ public class AwsServletContext
// Variables - Private
//-------------------------------------------------------------
private Map filters;
- private Context lambdaContext;
+ private Map servletRegistrations;
private Map attributes;
private Map initParameters;
+ private AwsLambdaServletContainerHandler containerHandler;
+ private Logger log = LoggerFactory.getLogger(AwsServletContext.class);
//-------------------------------------------------------------
@@ -79,38 +75,57 @@ public class AwsServletContext
// Constructors
//-------------------------------------------------------------
- private AwsServletContext(Context lambdaContext) {
- this.lambdaContext = lambdaContext;
-
+ public AwsServletContext(AwsLambdaServletContainerHandler containerHandler) {
+ this.containerHandler = containerHandler;
this.attributes = new HashMap<>();
this.initParameters = new HashMap<>();
this.filters = new LinkedHashMap<>();
+ this.servletRegistrations = new HashMap<>();
}
-
//-------------------------------------------------------------
// Implementation - ServletContext
//-------------------------------------------------------------
+ public static void clearServletContextCache() {
+ instance = null;
+ }
- public static ServletContext getInstance(Context lambdaContext) {
- if (instance == null) {
- instance = new AwsServletContext(lambdaContext);
- }
- return instance;
+ @Override
+ public String getContextPath() {
+ // servlets are always at the root.
+ return "";
}
+ @Override
+ public void setResponseCharacterEncoding(String encoding) {
+ throw new UnsupportedOperationException();
+ }
- public static void clearServletContextCache() {
- instance = null;
+ @Override
+ public String getResponseCharacterEncoding() {
+ return null;
}
+ @Override
+ public void setRequestCharacterEncoding(String encoding) {
+ throw new UnsupportedOperationException();
+ }
@Override
- public String getContextPath() {
- // servlets are always at the root.
- return "/";
+ public String getRequestCharacterEncoding() {
+ return null;
+ }
+
+ @Override
+ public void setSessionTimeout(int sessionTimeout) {
+ throw new UnsupportedOperationException();
+ }
+
+ @Override
+ public int getSessionTimeout() {
+ return 0;
}
@@ -146,13 +161,33 @@ public int getEffectiveMinorVersion() {
@Override
- public String getMimeType(String s) {
- try {
- return Files.probeContentType(Paths.get(s));
- } catch (IOException e) {
- e.printStackTrace();
+ @SuppressFBWarnings("PATH_TRAVERSAL_IN") // suppressing because we are using the getValidFilePath
+ public String getMimeType(String file) {
+ if (file == null || !file.contains(".")) {
return null;
}
+
+ String mimeType = null;
+
+ // may not work on Lambda until mailcap package is present https://github.com/aws/serverless-java-container/pull/504
+ try {
+ mimeType = Files.probeContentType(Paths.get(file));
+ } catch (IOException | InvalidPathException e) {
+ log("unable to probe for content type, will use fallback", e);
+ }
+
+ if (mimeType == null) {
+ try {
+ String mimeTypeGuess = URLConnection.guessContentTypeFromName(new File(file).getName());
+ if (mimeTypeGuess !=null) {
+ mimeType = mimeTypeGuess;
+ }
+ } catch (Exception e) {
+ log("couldn't find a better contentType than " + mimeType + " for file " + file, e);
+ }
+ }
+
+ return mimeType;
}
@@ -166,65 +201,61 @@ public Set getResourcePaths(String s) {
@Override
public URL getResource(String s) throws MalformedURLException {
- return getClass().getResource(s);
+ return AwsServletContext.class.getResource(s);
}
@Override
public InputStream getResourceAsStream(String s) {
- return getClass().getResourceAsStream(s);
+ return AwsServletContext.class.getResourceAsStream(s);
}
@Override
public RequestDispatcher getRequestDispatcher(String s) {
- // TODO: This should be part of the reader interface described in the getResourcePaths method
- return null;
+ return new AwsProxyRequestDispatcher(s, false, containerHandler);
}
@Override
public RequestDispatcher getNamedDispatcher(String s) {
- // TODO: This should be part of the reader interface described in the getResourcePaths method
- return null;
- }
-
-
- @Override
- public Servlet getServlet(String s) throws ServletException {
- return null;
- }
-
-
- @Override
- public Enumeration getServlets() {
- return null;
- }
-
-
- @Override
- public Enumeration getServletNames() {
+ return new AwsProxyRequestDispatcher(s, true, containerHandler);
+ }
+
+ public Servlet getServletForPath(String path) {
+ String[] pathParts = path.split("/");
+ for (AwsServletRegistration reg : servletRegistrations.values()) {
+ for (String p : reg.getMappings()) {
+ if ("".equals(p) || "/".equals(p) || "/*".equals(p)) {
+ return reg.getServlet();
+ }
+ // if I have no path and I haven't matched something now I'll just move on to the next
+ if ("".equals(path) || "/".equals(path)) {
+ continue;
+ }
+ String[] regParts = p.split("/");
+ for (int i = 0; i < regParts.length; i++) {
+ if (!regParts[i].equals(pathParts[i]) && !"*".equals(regParts[i])) {
+ break;
+ }
+ if (i == regParts.length - 1 && (regParts[i].equals(pathParts[i]) || "*".equals(regParts[i]))) {
+ return reg.getServlet();
+ }
+ }
+ }
+ }
return null;
}
-
@Override
public void log(String s) {
- lambdaContext.getLogger().log(s);
- }
-
-
- @Override
- public void log(Exception e, String s) {
- lambdaContext.getLogger().log(s);
- lambdaContext.getLogger().log(e.getMessage());
+ log.info(SecurityUtils.encode(s));
}
@Override
public void log(String s, Throwable throwable) {
- lambdaContext.getLogger().log(s);
- lambdaContext.getLogger().log(throwable.getMessage());
+ log.error(SecurityUtils.encode(s), throwable);
}
@@ -236,7 +267,7 @@ public String getRealPath(String s) {
try {
absPath = new File(fileUrl.toURI()).getAbsolutePath();
} catch (URISyntaxException e) {
- lambdaContext.getLogger().log("Error while looking for real path: " + s + "\n" + e.getMessage());
+ log.error("Error while looking for real path {}: {}", SecurityUtils.encode(s), SecurityUtils.encode(e.getMessage()));
}
}
return absPath;
@@ -294,54 +325,72 @@ public void removeAttribute(String s) {
@Override
public String getServletContextName() {
- // TODO: This can also come from a reader interface
- throw new UnsupportedOperationException();
+ return null;
}
@Override
public ServletRegistration.Dynamic addServlet(String s, String s1) {
- log("Called addServlet: " + s1);
- log("Implemented frameworks are responsible for registering servlets");
- throw new UnsupportedOperationException();
+ try {
+ Class extends Servlet> servletClass = (Class extends Servlet>) this.getClassLoader().loadClass(s1);
+ Servlet servlet = createServlet(servletClass);
+ servletRegistrations.put(s, new AwsServletRegistration(s, servlet, this));
+ return servletRegistrations.get(s);
+ } catch (ServletException | ClassNotFoundException e) {
+ throw new RuntimeException(e);
+ }
}
@Override
public ServletRegistration.Dynamic addServlet(String s, Servlet servlet) {
- log("Called addServlet: " + servlet.getClass().getName());
- log("Implemented frameworks are responsible for registering servlets");
- throw new UnsupportedOperationException();
+ servletRegistrations.put(s, new AwsServletRegistration(s, servlet, this));
+ return servletRegistrations.get(s);
}
@Override
public ServletRegistration.Dynamic addServlet(String s, Class extends Servlet> aClass) {
- log("Called addServlet: " + aClass.getName());
- log("Implemented frameworks are responsible for registering servlets");
+ try {
+ Servlet servlet = createServlet(aClass);
+ servletRegistrations.put(s, new AwsServletRegistration(s, servlet, this));
+ return servletRegistrations.get(s);
+ } catch (ServletException e) {
+ throw new RuntimeException(e);
+ }
+ }
+
+ @Override
+ public ServletRegistration.Dynamic addJspFile(String s, String s1) {
throw new UnsupportedOperationException();
}
@Override
public T createServlet(Class aClass) throws ServletException {
- log("Called createServlet: " + aClass.getName());
- log("Implemented frameworks are responsible for creating servlets");
- throw new UnsupportedOperationException();
+ /*log("Called createServlet: " + aClass.getName());
+ log("Implemented frameworks are responsible for creating servlets");*/
+ // TODO: This method introspects the given clazz for the following annotations: ServletSecurity, MultipartConfig,
+ // javax.annotation.security.RunAs, and javax.annotation.security.DeclareRoles. In addition, this method supports
+ // resource injection if the given clazz represents a Managed Bean. See the Java EE platform and JSR 299 specifications
+ // for additional details about Managed Beans and resource injection.
+ try {
+ return aClass.newInstance();
+ } catch (InstantiationException | IllegalAccessException e) {
+ throw new ServletException(e);
+ }
}
@Override
public ServletRegistration getServletRegistration(String s) {
- // TODO: This could come from the reader interface
- return null;
+ return servletRegistrations.get(s);
}
@Override
public Map getServletRegistrations() {
- // TODO: This could come from the reader interface
- return null;
+ return servletRegistrations;
}
@@ -349,7 +398,7 @@ public ServletRegistration getServletRegistration(String s) {
public FilterRegistration.Dynamic addFilter(String name, String filterClass) {
try {
Class> newFilterClass = getClassLoader().loadClass(filterClass);
- if (!newFilterClass.isAssignableFrom(Filter.class)) {
+ if (!Filter.class.isAssignableFrom(newFilterClass)) {
throw new IllegalArgumentException(filterClass + " does not implement Filter");
}
@SuppressWarnings("unchecked")
@@ -357,7 +406,7 @@ public FilterRegistration.Dynamic addFilter(String name, String filterClass) {
return addFilter(name, filterCastClass);
} catch (ClassNotFoundException e) {
- e.printStackTrace();
+ log.error("Could not find filter class", e);
throw new IllegalStateException("Filter class " + filterClass + " not found");
}
}
@@ -371,6 +420,8 @@ public FilterRegistration.Dynamic addFilter(String name, Filter filter) {
// filter already exists, we do nothing
if (filters.containsKey(name)) {
return null;
+ } else {
+ log.debug("Adding filter '{}' from {}", SecurityUtils.encode(name), SecurityUtils.encode(filter.toString()));
}
FilterHolder newFilter = new FilterHolder(name, filter, this);
@@ -383,12 +434,13 @@ public FilterRegistration.Dynamic addFilter(String name, Filter filter) {
@Override
public FilterRegistration.Dynamic addFilter(String name, Class extends Filter> filterClass) {
try {
+ log.debug("Adding filter '{}' from {}", SecurityUtils.encode(name), SecurityUtils.encode(filterClass.getName()));
Filter newFilter = createFilter(filterClass);
return addFilter(name, newFilter);
} catch (ServletException e) {
// TODO: There is no clear indication in the servlet specs on whether we should throw an exception here.
// See JavaDoc here: http://docs.oracle.com/javaee/7/api/javax/servlet/ServletContext.html#addFilter-java.lang.String-java.lang.Class-
- lambdaContext.getLogger().log("Could not register filter: " + e.getMessage());
+ log.error("Could not register filter: ", e);
}
return null;
}
@@ -399,7 +451,7 @@ public T createFilter(Class aClass) throws ServletExceptio
try {
return aClass.newInstance();
} catch (InstantiationException | IllegalAccessException e) {
- lambdaContext.getLogger().log("Could not initialize filter class " + aClass.getName() + "\n" + e.getMessage());
+ log.error("Could not initialize filter class " + aClass.getName(), e);
throw new ServletException();
}
}
@@ -418,9 +470,10 @@ public FilterRegistration getFilterRegistration(String s) {
@Override
public Map getFilterRegistrations() {
Map registrations = new LinkedHashMap<>();
- for (String filter : filters.keySet()) {
- registrations.put(filter, filters.get(filter).getRegistration());
+ for (Map.Entry entry : filters.entrySet()) {
+ registrations.put(entry.getKey(), entry.getValue().getRegistration());
}
+
return registrations;
}
@@ -486,9 +539,7 @@ public JspConfigDescriptor getJspConfigDescriptor() {
@Override
public ClassLoader getClassLoader() {
- // for the time being we return the default class loader. We may want to let developers override this int the
- // future.
- return ClassLoader.getSystemClassLoader();
+ return getClass().getClassLoader();
}
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsServletInputStream.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsServletInputStream.java
new file mode 100644
index 00000000..220ca602
--- /dev/null
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsServletInputStream.java
@@ -0,0 +1,96 @@
+/*
+ * Copyright 2020 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 com.amazonaws.serverless.proxy.internal.servlet;
+
+import org.apache.commons.io.input.NullInputStream;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+
+import jakarta.servlet.ReadListener;
+import jakarta.servlet.ServletInputStream;
+import java.io.IOException;
+import java.io.InputStream;
+import java.nio.ByteBuffer;
+
+public class AwsServletInputStream extends ServletInputStream {
+ private static Logger log = LoggerFactory.getLogger(AwsServletInputStream.class);
+ private InputStream bodyStream;
+ private ReadListener listener;
+ private boolean finished;
+
+ public AwsServletInputStream(InputStream body) {
+ bodyStream = body;
+ finished = false;
+ }
+
+ @Override
+ public boolean isFinished() {
+ return finished;
+ }
+
+ @Override
+ public boolean isReady() {
+ if (finished && listener != null) {
+ try {
+ listener.onAllDataRead();
+ } catch (IOException e) {
+ log.error("Could not notify listeners that input stream data is ready", e);
+ throw new RuntimeException(e);
+ }
+ }
+ return !finished;
+ }
+
+ @Override
+ public void setReadListener(ReadListener readListener) {
+ listener = readListener;
+ try {
+ listener.onDataAvailable();
+ } catch (IOException e) {
+ log.error("Could not notify listeners that data is available", e);
+ throw new RuntimeException(e);
+ }
+ }
+
+ @Override
+ public int read()
+ throws IOException {
+ if (bodyStream == null || bodyStream instanceof NullInputStream) {
+ return -1;
+ }
+ int readByte = bodyStream.read();
+ if (readByte == -1) {
+ finished = true;
+ }
+ return readByte;
+ }
+
+ @Override
+ public int read(ByteBuffer b) throws IOException {
+ if (bodyStream == null || bodyStream instanceof NullInputStream) {
+ return -1;
+ }
+ if (!b.hasRemaining()) {
+ return 0;
+ }
+ byte[] buf = new byte[b.remaining()];
+ int bytesRead = bodyStream.read(buf);
+ if (bytesRead > 0) {
+ b.put(buf, 0, bytesRead);
+ }
+ if (bytesRead == -1) {
+ finished = true;
+ }
+ return bytesRead;
+ }
+}
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsServletRegistration.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsServletRegistration.java
new file mode 100644
index 00000000..9fdc3b3b
--- /dev/null
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/AwsServletRegistration.java
@@ -0,0 +1,182 @@
+/*
+ * Copyright 2020 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 com.amazonaws.serverless.proxy.internal.servlet;
+
+import edu.umd.cs.findbugs.annotations.SuppressFBWarnings;
+
+import jakarta.servlet.*;
+import java.util.*;
+
+/**
+ * Stores information about a servlet registered with Serverless Java Container's ServletContext.
+ */
+public class AwsServletRegistration implements ServletRegistration, ServletRegistration.Dynamic, Comparable {
+ private String servletName;
+ private Servlet servlet;
+ private AwsServletContext ctx;
+ private Map initParameters;
+ private int loadOnStartup;
+ private String runAsRole;
+ private boolean asyncSupported;
+ private Map servletPathMappings;
+
+
+ public AwsServletRegistration(String name, Servlet s, AwsServletContext context) {
+ servletName = name;
+ servlet = s;
+ ctx = context;
+ initParameters = new HashMap<>();
+ servletPathMappings = new HashMap<>();
+ loadOnStartup = -1;
+ asyncSupported = true;
+ }
+
+ @Override
+ public Set addMapping(String... strings) {
+ Set failedMappings = new HashSet<>();
+ for (String s : strings) {
+ if (servletPathMappings.containsKey(s)) {
+ failedMappings.add(s);
+ continue;
+ }
+ servletPathMappings.put(s, this);
+ }
+ return failedMappings;
+ }
+
+ @Override
+ public Collection getMappings() {
+ return servletPathMappings.keySet();
+ }
+
+ @Override
+ public String getRunAsRole() {
+ return runAsRole;
+ }
+
+ @Override
+ public String getName() {
+ return servletName;
+ }
+
+ @Override
+ public String getClassName() {
+ return servlet.getClass().getName();
+ }
+
+ @Override
+ public boolean setInitParameter(String s, String s1) {
+ if (initParameters.containsKey(s)) {
+ return false;
+ }
+ initParameters.put(s, s1);
+ return true;
+ }
+
+ @Override
+ public String getInitParameter(String s) {
+ return initParameters.get(s);
+ }
+
+ @Override
+ public Set setInitParameters(Map map) {
+ Set failedParameters = new HashSet<>();
+ for (Map.Entry param : map.entrySet()) {
+ if (initParameters.containsKey(param.getKey())) {
+ failedParameters.add(param.getKey());
+ }
+ initParameters.put(param.getKey(), param.getValue());
+ }
+ return failedParameters;
+ }
+
+ @Override
+ public Map getInitParameters() {
+ return initParameters;
+ }
+
+ public Servlet getServlet() {
+ return servlet;
+ }
+
+ @Override
+ public void setLoadOnStartup(int i) {
+ loadOnStartup = i;
+ }
+
+ public int getLoadOnStartup() {
+ return loadOnStartup;
+ }
+
+ @Override
+ public Set setServletSecurity(ServletSecurityElement servletSecurityElement) {
+ return null;
+ }
+
+ @Override
+ public void setMultipartConfig(MultipartConfigElement multipartConfigElement) {
+
+ }
+
+ @Override
+ public void setRunAsRole(String s) {
+ runAsRole = s;
+ }
+
+ @Override
+ public void setAsyncSupported(boolean b) {
+ asyncSupported = b;
+ }
+
+ @Override
+ public int compareTo(AwsServletRegistration r) {
+ return Integer.compare(loadOnStartup, r.getLoadOnStartup());
+ }
+
+ @Override
+ @SuppressFBWarnings("HE_EQUALS_USE_HASHCODE")
+ public boolean equals(Object r) {
+ if (r == null || !AwsServletRegistration.class.isAssignableFrom(r.getClass())) {
+ return false;
+ }
+ return ((AwsServletRegistration)r).getName().equals(getName()) && ((AwsServletRegistration)r).getServlet() == getServlet();
+ }
+
+ public boolean isAsyncSupported() {
+ return asyncSupported;
+ }
+
+ public ServletConfig getServletConfig() {
+ return new ServletConfig() {
+ @Override
+ public String getServletName() {
+ return servletName;
+ }
+
+ @Override
+ public ServletContext getServletContext() {
+ return ctx;
+ }
+
+ @Override
+ public String getInitParameter(String s) {
+ return initParameters.get(s);
+ }
+
+ @Override
+ public Enumeration getInitParameterNames() {
+ return Collections.enumeration(initParameters.keySet());
+ }
+ };
+ }
+}
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/CookieProcessor.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/CookieProcessor.java
new file mode 100644
index 00000000..c59dc806
--- /dev/null
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/CookieProcessor.java
@@ -0,0 +1,23 @@
+package com.amazonaws.serverless.proxy.internal.servlet;
+
+import jakarta.servlet.http.Cookie;
+
+public interface CookieProcessor {
+ /**
+ * Parse the provided cookie header value into an array of Cookie objects.
+ *
+ * @param cookieHeader The cookie header value string to parse, e.g., "SID=31d4d96e407aad42; lang=en-US"
+ * @return An array of Cookie objects parsed from the cookie header value
+ */
+ Cookie[] parseCookieHeader(String cookieHeader);
+
+ /**
+ * Generate the Set-Cookie HTTP header value for the given Cookie.
+ *
+ * @param cookie The cookie for which the header will be generated
+ * @return The header value in a form that can be added directly to the response
+ */
+ String generateHeader(Cookie cookie);
+
+
+}
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/FilterChainHolder.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/FilterChainHolder.java
index caa460bf..b45eed4a 100644
--- a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/FilterChainHolder.java
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/FilterChainHolder.java
@@ -12,10 +12,11 @@
*/
package com.amazonaws.serverless.proxy.internal.servlet;
-import javax.servlet.FilterChain;
-import javax.servlet.ServletException;
-import javax.servlet.ServletRequest;
-import javax.servlet.ServletResponse;
+import edu.umd.cs.findbugs.annotations.SuppressFBWarnings;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+
+import jakarta.servlet.*;
import java.io.IOException;
import java.util.ArrayList;
@@ -23,7 +24,7 @@
/**
* Implementation of the FilterChain interface. FilterChainHolder objects should be accessed through the
- * FilterChainManager. Once a filter chain is loaded, use the doFilter emthod to run the chain
+ * FilterChainManager. Once a filter chain is loaded, use the doFilter method to run the chain
* during a request lifecycle
*/
public class FilterChainHolder implements FilterChain {
@@ -33,7 +34,9 @@ public class FilterChainHolder implements FilterChain {
//-------------------------------------------------------------
private List filters;
- private int currentFilter;
+ int currentFilter;
+
+ private Logger log = LoggerFactory.getLogger(FilterChainHolder.class);
//-------------------------------------------------------------
@@ -43,7 +46,7 @@ public class FilterChainHolder implements FilterChain {
/**
* Creates a new empty FilterChainHolder
*/
- public FilterChainHolder() {
+ FilterChainHolder() {
this(new ArrayList<>());
}
@@ -52,9 +55,9 @@ public FilterChainHolder() {
* Creates a new instance of a filter chain holder
* @param allFilters A populated list of FilterHolder objects
*/
- public FilterChainHolder(List allFilters) {
+ FilterChainHolder(List allFilters) {
filters = allFilters;
- currentFilter = -1;
+ resetHolder();
}
@@ -62,20 +65,32 @@ public FilterChainHolder(List allFilters) {
// Implementation - FilterChain
//-------------------------------------------------------------
+ @SuppressFBWarnings("CRLF_INJECTION_LOGS")
@Override
public void doFilter(ServletRequest servletRequest, ServletResponse servletResponse) throws IOException, ServletException {
currentFilter++;
- if (filters == null || filters.size() == 0 || currentFilter > filters.size() - 1) {
- return;
- }
// TODO: We do not check for async filters here
- FilterHolder holder = filters.get(currentFilter);
-
- if (!holder.isFilterInitialized()) {
- holder.init();
+ // if we still have filters, keep running through the chain
+ if (currentFilter <= filters.size() - 1) {
+ FilterHolder holder = filters.get(currentFilter);
+
+ // confirm that this filter needs to be executed
+ if (!holder.getRegistration().getDispatcherTypes().contains(servletRequest.getDispatcherType())) {
+ // skip to the next filter - we have already incremented the currentFilter
+ doFilter(servletRequest, servletResponse);
+ }
+
+ // lazily initialize filters when they are needed
+ if (!holder.isFilterInitialized()) {
+ holder.init();
+ }
+ log.debug("Starting {}: filter {}-{}", servletRequest.getDispatcherType(),
+ currentFilter, holder.getFilterName());
+ holder.getFilter().doFilter(servletRequest, servletResponse, this);
+ log.debug("Executed {}: filter {}-{}", servletRequest.getDispatcherType(),
+ currentFilter, holder.getFilterName());
}
- holder.getFilter().doFilter(servletRequest, servletResponse, this);
}
@@ -87,7 +102,7 @@ public void doFilter(ServletRequest servletRequest, ServletResponse servletRespo
* Add a filter to the chain.
* @param newFilter The filter to be added at the end of the chain
*/
- public void addFilter(FilterHolder newFilter) {
+ void addFilter(FilterHolder newFilter) {
filters.add(newFilter);
}
@@ -96,7 +111,7 @@ public void addFilter(FilterHolder newFilter) {
* Returns the number of filters loaded in the chain holder
* @return The number of filters in the chain holder. If the filter chain is null then this will return 0
*/
- public int filterCount() {
+ int filterCount() {
if (filters == null) {
return 0;
} else {
@@ -110,11 +125,34 @@ public int filterCount() {
* @param idx The index in the chain. Use the filterCount method to get the filter count
* @return A populated FilterHolder object
*/
- public FilterHolder getFilter(int idx) {
+ FilterHolder getFilter(int idx) {
if (filters == null) {
return null;
} else {
return filters.get(idx);
}
}
+
+
+ /**
+ * Returns the list of filters in this chain.
+ * @return The list of filters
+ */
+ public List getFilters() {
+ return filters;
+ }
+
+
+ /**
+ * Resets the chain holder to the beginning of the filter chain. This method is used from the constructor as well as when
+ * the {@link FilterChainManager} return a holder from the cache.
+ */
+ private void resetHolder() {
+ currentFilter = -1;
+ }
+
+ @Override
+ public String toString() {
+ return "filters=" + filters;
+ }
}
diff --git a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/FilterChainManager.java b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/FilterChainManager.java
index 2cffe10a..4917b4aa 100644
--- a/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/FilterChainManager.java
+++ b/aws-serverless-java-container-core/src/main/java/com/amazonaws/serverless/proxy/internal/servlet/FilterChainManager.java
@@ -12,27 +12,34 @@
*/
package com.amazonaws.serverless.proxy.internal.servlet;
-import javax.servlet.DispatcherType;
-import javax.servlet.http.HttpServletRequest;
+import edu.umd.cs.findbugs.annotations.SuppressFBWarnings;
+
+import jakarta.servlet.*;
+import jakarta.servlet.http.HttpServletRequest;
+
+import java.io.IOException;
import java.util.Collections;
import java.util.HashMap;
+import java.util.List;
+import java.util.Locale;
import java.util.Map;
+import java.util.function.Predicate;
/**
- * This object in in charge of matching a servlet request to a set of filter, creating the filter chain for a request,
+ * This object is in charge of matching a servlet request to a set of filters, creating the filter chain for a request,
* and cache filter chains that were already loaded for re-use. This object should be used by the framework-specific
* implementations that use the HttpServletRequest and HttpServletResponse objects.
*
* For example, the Spring implementation creates the ServletContext when the application is initialized the first time
- * and creates a FitlerChainManager to execute its filters for each request.
+ * and creates a FilterChainManager to execute its filters for each request.
*/
-public abstract class FilterChainManager {
+public abstract class FilterChainManager {
//-------------------------------------------------------------
// Variables - Protected
//-------------------------------------------------------------
- protected static final String PATH_PART_SEPARATOR = "/";
+ static final String PATH_PART_SEPARATOR = "/";
//-------------------------------------------------------------
@@ -41,7 +48,7 @@ public abstract class FilterChainManager