diff --git a/core/src/main/java/org/springframework/ws/server/endpoint/support/PayloadRootUtils.java b/core/src/main/java/org/springframework/ws/server/endpoint/support/PayloadRootUtils.java
index f5b831d3..19e91bb0 100644
--- a/core/src/main/java/org/springframework/ws/server/endpoint/support/PayloadRootUtils.java
+++ b/core/src/main/java/org/springframework/ws/server/endpoint/support/PayloadRootUtils.java
@@ -16,22 +16,27 @@
package org.springframework.ws.server.endpoint.support;
+import java.io.InputStream;
+import java.io.Reader;
import javax.xml.namespace.QName;
+import javax.xml.stream.XMLEventReader;
import javax.xml.stream.XMLStreamConstants;
import javax.xml.stream.XMLStreamException;
import javax.xml.stream.XMLStreamReader;
+import javax.xml.stream.events.XMLEvent;
import javax.xml.transform.Source;
-import javax.xml.transform.Transformer;
import javax.xml.transform.TransformerException;
import javax.xml.transform.TransformerFactory;
import javax.xml.transform.dom.DOMResult;
-import javax.xml.transform.dom.DOMSource;
-import org.springframework.util.xml.StaxUtils;
import org.springframework.xml.namespace.QNameUtils;
+import org.springframework.xml.transform.TransformerHelper;
+import org.springframework.xml.transform.TraxUtils;
import org.w3c.dom.Document;
import org.w3c.dom.Node;
+import org.xml.sax.InputSource;
+import org.xml.sax.XMLReader;
/**
* Helper class for determining the root qualified name of a Web Service payload.
@@ -39,11 +44,9 @@ import org.w3c.dom.Node;
* @author Arjen Poutsma
* @since 1.0.0
*/
-@SuppressWarnings("Since15")
public abstract class PayloadRootUtils {
private PayloadRootUtils() {
-
}
/**
@@ -55,43 +58,91 @@ public abstract class PayloadRootUtils {
*/
public static QName getPayloadRootQName(Source source, TransformerFactory transformerFactory)
throws TransformerException {
+ return getPayloadRootQName(source, new TransformerHelper(transformerFactory));
+ }
+
+ public static QName getPayloadRootQName(Source source, TransformerHelper transformerHelper)
+ throws TransformerException {
if (source == null) {
return null;
}
- else if (source instanceof DOMSource) {
- DOMSource domSource = (DOMSource) source;
- Node node = domSource.getNode();
- if (node.getNodeType() == Node.ELEMENT_NODE) {
- return QNameUtils.getQNameForNode(node);
+ try {
+ PayloadRootSourceCallback callback = new PayloadRootSourceCallback();
+ TraxUtils.doWithSource(source, callback);
+ if (callback.result != null) {
+ return callback.result;
}
- else if (node.getNodeType() == Node.DOCUMENT_NODE) {
- Document document = (Document) node;
+ else {
+ // we have no other option than to transform
+ DOMResult domResult = new DOMResult();
+ transformerHelper.transform(source, domResult);
+ Document document = (Document) domResult.getNode();
return QNameUtils.getQNameForNode(document.getDocumentElement());
}
}
- else if (StaxUtils.isStaxSource(source)) {
- XMLStreamReader streamReader = StaxUtils.getXMLStreamReader(source);
- if (streamReader != null) {
- if (streamReader.getEventType() == XMLStreamConstants.START_DOCUMENT) {
- try {
- streamReader.nextTag();
- }
- catch (XMLStreamException ex) {
- throw new IllegalStateException("Could not read next tag: " + ex.getMessage(), ex);
- }
+ catch (TransformerException ex) {
+ throw ex;
+ }
+ catch (Exception ex) {
+ return null;
+ }
+ }
+
+ private static class PayloadRootSourceCallback implements TraxUtils.SourceCallback {
+
+ private QName result;
+
+ public void domSource(Node node) throws Exception {
+ if (node.getNodeType() == Node.ELEMENT_NODE) {
+ result = QNameUtils.getQNameForNode(node);
+ }
+ else if (node.getNodeType() == Node.DOCUMENT_NODE) {
+ Document document = (Document) node;
+ result = QNameUtils.getQNameForNode(document.getDocumentElement());
+ }
+ }
+
+ public void staxSource(XMLEventReader eventReader) throws Exception {
+ XMLEvent event = eventReader.peek();
+ if (event != null && event.isStartDocument()) {
+ event = eventReader.nextTag();
+ }
+ if (event != null) {
+ if (event.isStartElement()) {
+ result = event.asStartElement().getName();
}
- if (streamReader.getEventType() == XMLStreamConstants.START_ELEMENT ||
- streamReader.getEventType() == XMLStreamConstants.END_ELEMENT) {
- return streamReader.getName();
+ else if (event.isEndElement()) {
+ result = event.asEndElement().getName();
}
}
}
- // we have no other option than to transform
- Transformer transformer = transformerFactory.newTransformer();
- DOMResult domResult = new DOMResult();
- transformer.transform(source, domResult);
- Document document = (Document) domResult.getNode();
- return QNameUtils.getQNameForNode(document.getDocumentElement());
+
+ public void staxSource(XMLStreamReader streamReader) throws Exception {
+ if (streamReader.getEventType() == XMLStreamConstants.START_DOCUMENT) {
+ try {
+ streamReader.nextTag();
+ }
+ catch (XMLStreamException ex) {
+ throw new IllegalStateException("Could not read next tag: " + ex.getMessage(), ex);
+ }
+ }
+ if (streamReader.getEventType() == XMLStreamConstants.START_ELEMENT ||
+ streamReader.getEventType() == XMLStreamConstants.END_ELEMENT) {
+ result = streamReader.getName();
+ }
+ }
+
+ public void saxSource(XMLReader reader, InputSource inputSource) throws Exception {
+ // Do nothing
+ }
+
+ public void streamSource(InputStream inputStream) throws Exception {
+ // Do nothing
+ }
+
+ public void streamSource(Reader reader) throws Exception {
+ // Do nothing
+ }
}
diff --git a/core/src/test/java/org/springframework/ws/server/endpoint/support/PayloadRootUtilsTest.java b/core/src/test/java/org/springframework/ws/server/endpoint/support/PayloadRootUtilsTest.java
index f3babeba..9194325a 100644
--- a/core/src/test/java/org/springframework/ws/server/endpoint/support/PayloadRootUtilsTest.java
+++ b/core/src/test/java/org/springframework/ws/server/endpoint/support/PayloadRootUtilsTest.java
@@ -20,6 +20,7 @@ import java.io.StringReader;
import javax.xml.namespace.QName;
import javax.xml.parsers.DocumentBuilder;
import javax.xml.parsers.DocumentBuilderFactory;
+import javax.xml.stream.XMLEventReader;
import javax.xml.stream.XMLInputFactory;
import javax.xml.stream.XMLStreamReader;
import javax.xml.transform.Source;
@@ -36,7 +37,6 @@ import org.w3c.dom.Document;
import org.w3c.dom.Element;
import org.xml.sax.InputSource;
-@SuppressWarnings("Since15")
public class PayloadRootUtilsTest {
@Test
@@ -55,7 +55,7 @@ public class PayloadRootUtilsTest {
}
@Test
- public void testGetQNameForStaxSource() throws Exception {
+ public void testGetQNameForStaxSourceStreamReader() throws Exception {
String contents = "";
XMLInputFactory inputFactory = XMLInputFactory.newInstance();
XMLStreamReader streamReader = inputFactory.createXMLStreamReader(new StringReader(contents));
@@ -67,6 +67,19 @@ public class PayloadRootUtilsTest {
Assert.assertEquals("Qname has invalid prefix", "prefix", qName.getPrefix());
}
+ @Test
+ public void testGetQNameForStaxSourceEventReader() throws Exception {
+ String contents = "";
+ XMLInputFactory inputFactory = XMLInputFactory.newInstance();
+ XMLEventReader eventReader = inputFactory.createXMLEventReader(new StringReader(contents));
+ Source source = new StaxSource(eventReader);
+ QName qName = PayloadRootUtils.getPayloadRootQName(source, TransformerFactory.newInstance());
+ Assert.assertNotNull("getQNameForNode returns null", qName);
+ Assert.assertEquals("QName has invalid localname", "localname", qName.getLocalPart());
+ Assert.assertEquals("Qname has invalid namespace", "namespace", qName.getNamespaceURI());
+ Assert.assertEquals("Qname has invalid prefix", "prefix", qName.getPrefix());
+ }
+
@Test
public void testGetQNameForStreamSource() throws Exception {
String contents = "";