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 = "";