diff --git a/core/src/main/java/org/springframework/ws/server/endpoint/interceptor/PayloadTransformingInterceptor.java b/core/src/main/java/org/springframework/ws/server/endpoint/interceptor/PayloadTransformingInterceptor.java index 926e278c..932c48d0 100644 --- a/core/src/main/java/org/springframework/ws/server/endpoint/interceptor/PayloadTransformingInterceptor.java +++ b/core/src/main/java/org/springframework/ws/server/endpoint/interceptor/PayloadTransformingInterceptor.java @@ -16,13 +16,21 @@ package org.springframework.ws.server.endpoint.interceptor; +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; import javax.xml.transform.Source; import javax.xml.transform.Templates; import javax.xml.transform.Transformer; +import javax.xml.transform.TransformerException; import javax.xml.transform.TransformerFactory; +import javax.xml.transform.stream.StreamResult; +import javax.xml.transform.stream.StreamSource; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; +import org.xml.sax.XMLReader; +import org.xml.sax.helpers.XMLReaderFactory; + import org.springframework.beans.factory.InitializingBean; import org.springframework.core.io.Resource; import org.springframework.util.Assert; @@ -30,8 +38,7 @@ import org.springframework.ws.WebServiceMessage; import org.springframework.ws.context.MessageContext; import org.springframework.ws.server.EndpointInterceptor; import org.springframework.xml.transform.ResourceSource; -import org.xml.sax.XMLReader; -import org.xml.sax.helpers.XMLReaderFactory; +import org.springframework.xml.transform.TransformerObjectSupport; /** * Interceptor that transforms the payload of WebServiceMessages using XSLT stylesheet. Allows for seperate @@ -47,7 +54,8 @@ import org.xml.sax.helpers.XMLReaderFactory; * @see #setResponseXslt(org.springframework.core.io.Resource) * @since 1.0.0 */ -public class PayloadTransformingInterceptor implements EndpointInterceptor, InitializingBean { +public class PayloadTransformingInterceptor extends TransformerObjectSupport + implements EndpointInterceptor, InitializingBean { private static final Log logger = LogFactory.getLog(PayloadTransformingInterceptor.class); @@ -81,7 +89,7 @@ public class PayloadTransformingInterceptor implements EndpointInterceptor, Init if (requestTemplates != null) { WebServiceMessage request = messageContext.getRequest(); Transformer transformer = requestTemplates.newTransformer(); - transformer.transform(request.getPayloadSource(), request.getPayloadResult()); + transformMessage(request, transformer); logger.debug("Request message transformed"); } return true; @@ -99,12 +107,19 @@ public class PayloadTransformingInterceptor implements EndpointInterceptor, Init if (responseTemplates != null) { WebServiceMessage response = messageContext.getResponse(); Transformer transformer = responseTemplates.newTransformer(); - transformer.transform(response.getPayloadSource(), response.getPayloadResult()); + transformMessage(response, transformer); logger.debug("Response message transformed"); } return true; } + private void transformMessage(WebServiceMessage message, Transformer transformer) throws TransformerException { + ByteArrayOutputStream os = new ByteArrayOutputStream(); + transformer.transform(message.getPayloadSource(), new StreamResult(os)); + ByteArrayInputStream is = new ByteArrayInputStream(os.toByteArray()); + transform(new StreamSource(is), message.getPayloadResult()); + } + /** Does nothing by default. Faults are not transformed. */ public boolean handleFault(MessageContext messageContext, Object endpoint) throws Exception { return true; diff --git a/core/src/test/java/org/springframework/ws/server/endpoint/interceptor/PayloadTransformingInterceptorTest.java b/core/src/test/java/org/springframework/ws/server/endpoint/interceptor/PayloadTransformingInterceptorTest.java index 16c5209f..2e3f6125 100644 --- a/core/src/test/java/org/springframework/ws/server/endpoint/interceptor/PayloadTransformingInterceptorTest.java +++ b/core/src/test/java/org/springframework/ws/server/endpoint/interceptor/PayloadTransformingInterceptorTest.java @@ -16,19 +16,27 @@ package org.springframework.ws.server.endpoint.interceptor; +import javax.xml.soap.MessageFactory; +import javax.xml.soap.SOAPMessage; import javax.xml.transform.Transformer; import javax.xml.transform.TransformerFactory; import javax.xml.transform.sax.SAXSource; import org.custommonkey.xmlunit.XMLTestCase; import org.custommonkey.xmlunit.XMLUnit; + import org.springframework.core.io.ClassPathResource; import org.springframework.core.io.Resource; import org.springframework.ws.MockWebServiceMessage; import org.springframework.ws.MockWebServiceMessageFactory; import org.springframework.ws.context.DefaultMessageContext; import org.springframework.ws.context.MessageContext; +import org.springframework.ws.pox.dom.DomPoxMessage; +import org.springframework.ws.pox.dom.DomPoxMessageFactory; +import org.springframework.ws.soap.saaj.SaajSoapMessage; +import org.springframework.ws.soap.saaj.SaajSoapMessageFactory; import org.springframework.xml.sax.SaxUtils; +import org.springframework.xml.transform.ResourceSource; import org.springframework.xml.transform.StringResult; public class PayloadTransformingInterceptorTest extends XMLTestCase { @@ -61,9 +69,9 @@ public class PayloadTransformingInterceptorTest extends XMLTestCase { boolean result = interceptor.handleRequest(context, null); assertTrue("Invalid interceptor result", result); - StringResult stringResult = new StringResult(); - transformer.transform(new SAXSource(SaxUtils.createInputSource(output)), stringResult); - assertXMLEqual(stringResult.toString(), request.getPayloadAsString()); + StringResult expected = new StringResult(); + transformer.transform(new SAXSource(SaxUtils.createInputSource(output)), expected); + assertXMLEqual(expected.toString(), request.getPayloadAsString()); } public void testHandleRequestNoXslt() throws Exception { @@ -74,9 +82,9 @@ public class PayloadTransformingInterceptorTest extends XMLTestCase { boolean result = interceptor.handleRequest(context, null); assertTrue("Invalid interceptor result", result); - StringResult stringResult = new StringResult(); - transformer.transform(new SAXSource(SaxUtils.createInputSource(input)), stringResult); - assertXMLEqual(stringResult.toString(), request.getPayloadAsString()); + StringResult expected = new StringResult(); + transformer.transform(new SAXSource(SaxUtils.createInputSource(input)), expected); + assertXMLEqual(expected.toString(), request.getPayloadAsString()); } public void testHandleResponse() throws Exception { @@ -89,9 +97,9 @@ public class PayloadTransformingInterceptorTest extends XMLTestCase { boolean result = interceptor.handleResponse(context, null); assertTrue("Invalid interceptor result", result); - StringResult stringResult = new StringResult(); - transformer.transform(new SAXSource(SaxUtils.createInputSource(output)), stringResult); - assertXMLEqual(stringResult.toString(), response.getPayloadAsString()); + StringResult expected = new StringResult(); + transformer.transform(new SAXSource(SaxUtils.createInputSource(output)), expected); + assertXMLEqual(expected.toString(), response.getPayloadAsString()); } public void testHandleResponseNoXslt() throws Exception { @@ -104,9 +112,44 @@ public class PayloadTransformingInterceptorTest extends XMLTestCase { boolean result = interceptor.handleResponse(context, null); assertTrue("Invalid interceptor result", result); - StringResult stringResult = new StringResult(); - transformer.transform(new SAXSource(SaxUtils.createInputSource(input)), stringResult); - assertXMLEqual(stringResult.toString(), response.getPayloadAsString()); + StringResult expected = new StringResult(); + transformer.transform(new SAXSource(SaxUtils.createInputSource(input)), expected); + assertXMLEqual(expected.toString(), response.getPayloadAsString()); + } + + public void testSaaj() throws Exception { + interceptor.setRequestXslt(xslt); + interceptor.afterPropertiesSet(); + MessageFactory messageFactory = MessageFactory.newInstance(); + SOAPMessage saajMessage = messageFactory.createMessage(); + SaajSoapMessage message = new SaajSoapMessage(saajMessage); + transformer.transform(new ResourceSource(input), message.getPayloadResult()); + MessageContext context = new DefaultMessageContext(message, new SaajSoapMessageFactory(messageFactory)); + + assertTrue("Invalid interceptor result", interceptor.handleRequest(context, null)); + StringResult expected = new StringResult(); + transformer.transform(new SAXSource(SaxUtils.createInputSource(output)), expected); + StringResult result = new StringResult(); + transformer.transform(message.getPayloadSource(), result); + assertXMLEqual(expected.toString(), result.toString()); + + } + + public void testPox() throws Exception { + interceptor.setRequestXslt(xslt); + interceptor.afterPropertiesSet(); + DomPoxMessageFactory factory = new DomPoxMessageFactory(); + DomPoxMessage message = (DomPoxMessage) factory.createWebServiceMessage(); + transformer.transform(new ResourceSource(input), message.getPayloadResult()); + MessageContext context = new DefaultMessageContext(message, factory); + + assertTrue("Invalid interceptor result", interceptor.handleRequest(context, null)); + StringResult expected = new StringResult(); + transformer.transform(new SAXSource(SaxUtils.createInputSource(output)), expected); + StringResult result = new StringResult(); + transformer.transform(message.getPayloadSource(), result); + assertXMLEqual(expected.toString(), result.toString()); + } public void testNoStylesheetsSet() throws Exception {