diff --git a/core/src/main/java/org/springframework/ws/server/endpoint/AbstractStaxEventPayloadEndpoint.java b/core/src/main/java/org/springframework/ws/server/endpoint/AbstractStaxEventPayloadEndpoint.java index 7652e72d..dd4070dc 100644 --- a/core/src/main/java/org/springframework/ws/server/endpoint/AbstractStaxEventPayloadEndpoint.java +++ b/core/src/main/java/org/springframework/ws/server/endpoint/AbstractStaxEventPayloadEndpoint.java @@ -78,6 +78,9 @@ public abstract class AbstractStaxEventPayloadEndpoint extends AbstractStaxPaylo } private XMLEventReader getEventReader(Source source) throws XMLStreamException, TransformerException { + if (source == null) { + return null; + } XMLEventReader eventReader = null; if (source instanceof StaxSource) { StaxSource staxSource = (StaxSource) source; diff --git a/core/src/main/java/org/springframework/ws/server/endpoint/AbstractStaxStreamPayloadEndpoint.java b/core/src/main/java/org/springframework/ws/server/endpoint/AbstractStaxStreamPayloadEndpoint.java index 98f3f5a5..8897432d 100644 --- a/core/src/main/java/org/springframework/ws/server/endpoint/AbstractStaxStreamPayloadEndpoint.java +++ b/core/src/main/java/org/springframework/ws/server/endpoint/AbstractStaxStreamPayloadEndpoint.java @@ -54,6 +54,9 @@ public abstract class AbstractStaxStreamPayloadEndpoint extends AbstractStaxPayl } private XMLStreamReader getStreamReader(Source source) throws XMLStreamException, TransformerException { + if (source == null) { + return null; + } XMLStreamReader streamReader = null; if (source instanceof StaxSource) { streamReader = ((StaxSource) source).getXMLStreamReader(); diff --git a/core/src/main/java/org/springframework/ws/support/MarshallingUtils.java b/core/src/main/java/org/springframework/ws/support/MarshallingUtils.java index 2b13cbd0..12c4be54 100644 --- a/core/src/main/java/org/springframework/ws/support/MarshallingUtils.java +++ b/core/src/main/java/org/springframework/ws/support/MarshallingUtils.java @@ -18,6 +18,7 @@ package org.springframework.ws.support; import java.io.IOException; import javax.activation.DataHandler; +import javax.xml.transform.Source; import org.springframework.oxm.Marshaller; import org.springframework.oxm.Unmarshaller; @@ -41,6 +42,9 @@ public abstract class MarshallingUtils { /** * Unmarshals the payload of the given message using the provided {@link Unmarshaller}. + *
+ * If the request message has no payload (i.e. {@link WebServiceMessage#getPayloadSource()} returns + *null), this method will return null.
*
* @param unmarshaller the unmarshaller
* @param message the message of which the payload is to be unmarshalled
@@ -48,13 +52,17 @@ public abstract class MarshallingUtils {
* @throws IOException in case of I/O errors
*/
public static Object unmarshal(Unmarshaller unmarshaller, WebServiceMessage message) throws IOException {
- if (unmarshaller instanceof MimeUnmarshaller && message instanceof MimeMessage) {
+ Source payload = message.getPayloadSource();
+ if (payload == null) {
+ return null;
+ }
+ else if (unmarshaller instanceof MimeUnmarshaller && message instanceof MimeMessage) {
MimeUnmarshaller mimeUnmarshaller = (MimeUnmarshaller) unmarshaller;
MimeMessageContainer container = new MimeMessageContainer((MimeMessage) message);
- return mimeUnmarshaller.unmarshal(message.getPayloadSource(), container);
+ return mimeUnmarshaller.unmarshal(payload, container);
}
else {
- return unmarshaller.unmarshal(message.getPayloadSource());
+ return unmarshaller.unmarshal(payload);
}
}
diff --git a/core/src/test/java/org/springframework/ws/MockWebServiceMessage.java b/core/src/test/java/org/springframework/ws/MockWebServiceMessage.java
index 264588c2..34e5bad7 100644
--- a/core/src/test/java/org/springframework/ws/MockWebServiceMessage.java
+++ b/core/src/test/java/org/springframework/ws/MockWebServiceMessage.java
@@ -45,14 +45,13 @@ import org.springframework.xml.transform.StringSource;
*/
public class MockWebServiceMessage implements FaultAwareWebServiceMessage {
- private final StringBuffer content;
+ private StringBuffer content;
private boolean fault = false;
private String faultReason;
public MockWebServiceMessage() {
- content = new StringBuffer();
}
public MockWebServiceMessage(Source source) throws TransformerException {
@@ -71,14 +70,17 @@ public class MockWebServiceMessage implements FaultAwareWebServiceMessage {
}
public MockWebServiceMessage(String content) {
- this.content = new StringBuffer(content);
+ if (content != null) {
+ this.content = new StringBuffer(content);
+ }
}
public String getPayloadAsString() {
- return content.toString();
+ return content != null ? content.toString() : null;
}
public void setPayload(InputStreamSource inputStreamSource) throws IOException {
+ checkContent();
InputStream is = null;
try {
is = inputStreamSource.getInputStream();
@@ -93,16 +95,24 @@ public class MockWebServiceMessage implements FaultAwareWebServiceMessage {
}
public void setPayload(String content) {
+ checkContent();
this.content.replace(0, this.content.length(), content);
}
+ private void checkContent() {
+ if (content == null) {
+ content = new StringBuffer();
+ }
+ }
+
public Result getPayloadResult() {
+ checkContent();
content.setLength(0);
return new StreamResult(new StringBufferWriter());
}
public Source getPayloadSource() {
- return new StringSource(content.toString());
+ return content != null ? new StringSource(content.toString()) : null;
}
public boolean hasFault() {
@@ -122,13 +132,17 @@ public class MockWebServiceMessage implements FaultAwareWebServiceMessage {
}
public void writeTo(OutputStream outputStream) throws IOException {
- PrintWriter writer = new PrintWriter(outputStream);
- writer.write(content.toString());
+ if (content != null) {
+ PrintWriter writer = new PrintWriter(outputStream);
+ writer.write(content.toString());
+ }
}
public String toString() {
StringBuffer buffer = new StringBuffer("MockWebServiceMessage {");
- buffer.append(content);
+ if (content != null) {
+ buffer.append(content);
+ }
buffer.append('}');
return buffer.toString();
}
diff --git a/core/src/test/java/org/springframework/ws/server/endpoint/AbstractMessageEndpointTestCase.java b/core/src/test/java/org/springframework/ws/server/endpoint/AbstractMessageEndpointTestCase.java
index 8d24b00b..e1545b5f 100644
--- a/core/src/test/java/org/springframework/ws/server/endpoint/AbstractMessageEndpointTestCase.java
+++ b/core/src/test/java/org/springframework/ws/server/endpoint/AbstractMessageEndpointTestCase.java
@@ -42,6 +42,15 @@ public abstract class AbstractMessageEndpointTestCase extends AbstractEndpointTe
assertFalse("Response message created", context.hasResponse());
}
+ public void testNoRequestPayload() throws Exception {
+ endpoint = createNoRequestPayloadEndpoint();
+
+ MessageContext context = new DefaultMessageContext(new MockWebServiceMessage((StringBuffer) null),
+ new MockWebServiceMessageFactory());
+ endpoint.invoke(context);
+ assertFalse("Response message created", context.hasResponse());
+ }
+
protected final void testSource(Source requestSource) throws Exception {
MessageContext context =
new DefaultMessageContext(new MockWebServiceMessage(requestSource), new MockWebServiceMessageFactory());
@@ -52,6 +61,8 @@ public abstract class AbstractMessageEndpointTestCase extends AbstractEndpointTe
protected abstract MessageEndpoint createNoResponseEndpoint();
+ protected abstract MessageEndpoint createNoRequestPayloadEndpoint();
+
protected abstract MessageEndpoint createResponseEndpoint();
}
diff --git a/core/src/test/java/org/springframework/ws/server/endpoint/MarshallingPayloadEndpointTest.java b/core/src/test/java/org/springframework/ws/server/endpoint/MarshallingPayloadEndpointTest.java
index a70094da..69837761 100644
--- a/core/src/test/java/org/springframework/ws/server/endpoint/MarshallingPayloadEndpointTest.java
+++ b/core/src/test/java/org/springframework/ws/server/endpoint/MarshallingPayloadEndpointTest.java
@@ -143,6 +143,25 @@ public class MarshallingPayloadEndpointTest extends XMLTestCase {
factoryControl.verify();
}
+ public void testInvokeNoRequest() throws Exception {
+ MockWebServiceMessage request = new MockWebServiceMessage((StringBuffer) null);
+ context = new DefaultMessageContext(request, factoryMock);
+ AbstractMarshallingPayloadEndpoint endpoint = new AbstractMarshallingPayloadEndpoint() {
+
+ protected Object invokeInternal(Object requestObject) throws Exception {
+ assertNull("No request expected", requestObject);
+ return null;
+ }
+ };
+ endpoint.setMarshaller(new SimpleMarshaller());
+ endpoint.setUnmarshaller(new SimpleMarshaller());
+ endpoint.afterPropertiesSet();
+ factoryControl.replay();
+ endpoint.invoke(context);
+ assertFalse("Response created", context.hasResponse());
+ factoryControl.verify();
+ }
+
public void testInvokeMimeMarshaller() throws Exception {
MockControl unmarshallerControl = MockControl.createControl(MimeUnmarshaller.class);
MimeUnmarshaller unmarshaller = (MimeUnmarshaller) unmarshallerControl.getMock();
@@ -187,7 +206,7 @@ public class MarshallingPayloadEndpointTest extends XMLTestCase {
messageControl.verify();
}
- private abstract static class SimpleMarshaller implements Marshaller, Unmarshaller {
+ private static class SimpleMarshaller implements Marshaller, Unmarshaller {
public void marshal(Object graph, Result result) throws XmlMappingException, IOException {
fail("Not expected");
diff --git a/core/src/test/java/org/springframework/ws/server/endpoint/StaxEventPayloadEndpointTest.java b/core/src/test/java/org/springframework/ws/server/endpoint/StaxEventPayloadEndpointTest.java
index ca718617..7e41b6e1 100644
--- a/core/src/test/java/org/springframework/ws/server/endpoint/StaxEventPayloadEndpointTest.java
+++ b/core/src/test/java/org/springframework/ws/server/endpoint/StaxEventPayloadEndpointTest.java
@@ -37,6 +37,18 @@ public class StaxEventPayloadEndpointTest extends AbstractMessageEndpointTestCas
protected void invokeInternal(XMLEventReader eventReader,
XMLEventConsumer eventWriter,
XMLEventFactory eventFactory) throws Exception {
+ assertNotNull("No EventReader passed", eventReader);
+ }
+ };
+ }
+
+ protected MessageEndpoint createNoRequestPayloadEndpoint() {
+ return new AbstractStaxEventPayloadEndpoint() {
+
+ protected void invokeInternal(XMLEventReader eventReader,
+ XMLEventConsumer eventWriter,
+ XMLEventFactory eventFactory) throws Exception {
+ assertNull("EventReader passed", eventReader);
}
};
}
diff --git a/core/src/test/java/org/springframework/ws/server/endpoint/StaxStreamPayloadEndpointTest.java b/core/src/test/java/org/springframework/ws/server/endpoint/StaxStreamPayloadEndpointTest.java
index 04799b21..50cebb09 100644
--- a/core/src/test/java/org/springframework/ws/server/endpoint/StaxStreamPayloadEndpointTest.java
+++ b/core/src/test/java/org/springframework/ws/server/endpoint/StaxStreamPayloadEndpointTest.java
@@ -42,6 +42,15 @@ public class StaxStreamPayloadEndpointTest extends AbstractMessageEndpointTestCa
protected MessageEndpoint createNoResponseEndpoint() {
return new AbstractStaxStreamPayloadEndpoint() {
protected void invokeInternal(XMLStreamReader streamReader, XMLStreamWriter streamWriter) throws Exception {
+ assertNotNull("No StreamReader passed", streamReader);
+ }
+ };
+ }
+
+ protected MessageEndpoint createNoRequestPayloadEndpoint() {
+ return new AbstractStaxStreamPayloadEndpoint() {
+ protected void invokeInternal(XMLStreamReader streamReader, XMLStreamWriter streamWriter) throws Exception {
+ assertNull("StreamReader passed", streamReader);
}
};
}
diff --git a/core/src/test/java/org/springframework/ws/server/endpoint/adapter/MarshallingMethodEndpointAdapterTest.java b/core/src/test/java/org/springframework/ws/server/endpoint/adapter/MarshallingMethodEndpointAdapterTest.java
index ae151765..f130d5cd 100644
--- a/core/src/test/java/org/springframework/ws/server/endpoint/adapter/MarshallingMethodEndpointAdapterTest.java
+++ b/core/src/test/java/org/springframework/ws/server/endpoint/adapter/MarshallingMethodEndpointAdapterTest.java
@@ -6,6 +6,7 @@ import junit.framework.TestCase;
import org.easymock.MockControl;
import org.springframework.oxm.Marshaller;
import org.springframework.oxm.Unmarshaller;
+import org.springframework.ws.MockWebServiceMessage;
import org.springframework.ws.MockWebServiceMessageFactory;
import org.springframework.ws.context.DefaultMessageContext;
import org.springframework.ws.context.MessageContext;
@@ -38,7 +39,8 @@ public class MarshallingMethodEndpointAdapterTest extends TestCase {
unmarshallerMock = (Unmarshaller) unmarshallerControl.getMock();
adapter.setUnmarshaller(unmarshallerMock);
adapter.afterPropertiesSet();
- messageContext = new DefaultMessageContext(new MockWebServiceMessageFactory());
+ MockWebServiceMessage request = new MockWebServiceMessage("