diff --git a/src/test/java/me/desair/tus/server/util/MockServletInputStream.java b/src/test/java/me/desair/tus/server/util/MockServletInputStream.java new file mode 100644 index 0000000..d335af2 --- /dev/null +++ b/src/test/java/me/desair/tus/server/util/MockServletInputStream.java @@ -0,0 +1,42 @@ +package me.desair.tus.server.util; + +import jakarta.servlet.ReadListener; +import jakarta.servlet.ServletInputStream; +import java.io.IOException; +import java.io.InputStream; + +public class MockServletInputStream extends ServletInputStream { + + private final InputStream delegate; + + public MockServletInputStream(InputStream delegate) { + this.delegate = delegate; + } + + @Override + public int read() throws IOException { + return delegate.read(); + } + + @Override + public int read(byte[] b, int off, int len) throws IOException { + return delegate.read(b, off, len); + } + + @Override + public boolean isFinished() { + try { + return delegate.available() == 0; + } catch (IOException e) { + return true; + } + } + + @Override + public boolean isReady() { + return true; + } + + @Override + public void setReadListener(ReadListener readListener) {} +} diff --git a/src/test/java/me/desair/tus/server/util/TusServletRequestTest.java b/src/test/java/me/desair/tus/server/util/TusServletRequestTest.java index 9b95c88..4ac1f62 100644 --- a/src/test/java/me/desair/tus/server/util/TusServletRequestTest.java +++ b/src/test/java/me/desair/tus/server/util/TusServletRequestTest.java @@ -1,7 +1,13 @@ package me.desair.tus.server.util; +import static org.hamcrest.MatcherAssert.assertThat; +import static org.hamcrest.Matchers.hasItems; +import static org.hamcrest.Matchers.is; +import static org.hamcrest.Matchers.notNullValue; +import static org.hamcrest.Matchers.nullValue; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNull; +import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; import jakarta.servlet.ReadListener; @@ -11,6 +17,10 @@ import java.io.IOException; import java.io.InputStream; import java.nio.charset.StandardCharsets; +import java.util.Set; +import me.desair.tus.server.HttpHeader; +import me.desair.tus.server.TusExtension; +import me.desair.tus.server.checksum.ChecksumAlgorithm; import org.apache.commons.io.IOUtils; import org.junit.Before; import org.junit.Test; @@ -23,18 +33,134 @@ public class TusServletRequestTest { @Mock private HttpServletRequest servletRequest; - private TusServletRequest tusServletRequest; + private TusServletRequest request; @Before public void setUp() { - tusServletRequest = new TusServletRequest(servletRequest, true); + request = new TusServletRequest(servletRequest, true); + } + + @Test + public void testGetContentInputStream() throws Exception { + byte[] data = "test data".getBytes(); + when(servletRequest.getInputStream()) + .thenReturn(new MockServletInputStream(new ByteArrayInputStream(data))); + + InputStream is = request.getContentInputStream(); + + assertThat(is, notNullValue()); + + // Read to the end to trigger counting + byte[] buffer = new byte[1024]; + int bytesRead = is.read(buffer); + + assertThat(bytesRead, is(9)); + assertThat(request.getBytesRead(), is(9L)); + } + + @Test + public void testGetContentInputStreamChunked() throws Exception { + TusServletRequest chunkedRequest = new TusServletRequest(servletRequest, true); + + byte[] data = "5\r\ntest \r\n4\r\ndata\r\n0\r\n\r\n".getBytes(); + when(servletRequest.getInputStream()) + .thenReturn(new MockServletInputStream(new ByteArrayInputStream(data))); + when(servletRequest.getHeader(HttpHeader.TRANSFER_ENCODING)).thenReturn("chunked"); + + InputStream is = chunkedRequest.getContentInputStream(); + + assertThat(is, notNullValue()); + + // Read to the end to trigger counting + byte[] buffer = new byte[1024]; + int bytesRead = 0; + int read; + while ((read = is.read(buffer)) != -1) { + bytesRead += read; + } + + assertThat(bytesRead, is(9)); + assertThat(chunkedRequest.getBytesRead(), is(9L)); + } + + @Test + public void testGetContentInputStreamWithChecksum() throws Exception { + byte[] data = "test data".getBytes(); + when(servletRequest.getInputStream()) + .thenReturn(new MockServletInputStream(new ByteArrayInputStream(data))); + when(servletRequest.getHeader(HttpHeader.UPLOAD_CHECKSUM)) + .thenReturn("sha1 9I3YU4IIYIFsddVND1hNyGMyenw="); + + InputStream is = request.getContentInputStream(); + assertThat(is, notNullValue()); + + byte[] buffer = new byte[1024]; + int read; + while ((read = is.read(buffer)) != -1) { + // Consume stream completely to calculate checksum + } + + assertThat(request.hasCalculatedChecksum(), is(true)); + Set algorithms = request.getEnabledChecksums(); + assertThat(algorithms, hasItems(ChecksumAlgorithm.SHA1)); + + assertThat( + request.getCalculatedChecksum(ChecksumAlgorithm.SHA1), is("9I3YU4IIYIFsddVND1hNyGMyenw=")); + } + + @Test + public void testGetContentInputStreamChunkedWithChecksum() throws Exception { + TusServletRequest chunkedRequest = new TusServletRequest(servletRequest, true); + + byte[] data = "5\r\ntest \r\n4\r\ndata\r\n0\r\n\r\n".getBytes(); + when(servletRequest.getInputStream()) + .thenReturn(new MockServletInputStream(new ByteArrayInputStream(data))); + when(servletRequest.getHeader(HttpHeader.TRANSFER_ENCODING)).thenReturn("chunked"); + + InputStream is = chunkedRequest.getContentInputStream(); + assertThat(is, notNullValue()); + + byte[] buffer = new byte[1024]; + int read; + while ((read = is.read(buffer)) != -1) { + // Consume stream completely to calculate checksum + } + + assertThat(chunkedRequest.hasCalculatedChecksum(), is(true)); + Set algorithms = chunkedRequest.getEnabledChecksums(); + // Since it's chunked and checksum can come at the end, it should keep track of all algorithms + assertThat(algorithms, hasItems(ChecksumAlgorithm.values())); + + assertThat( + chunkedRequest.getCalculatedChecksum(ChecksumAlgorithm.SHA1), + is("9I3YU4IIYIFsddVND1hNyGMyenw=")); + } + + @Test + public void testIsProcessedBy() { + TusExtension extension = mock(TusExtension.class); + when(extension.getName()).thenReturn("test"); + + assertThat(request.isProcessedBy(extension), is(false)); + + request.addProcessor(extension); + + assertThat(request.isProcessedBy(extension), is(true)); + } + + @Test + public void testGetHeader() { + when(servletRequest.getHeader("X-Custom-Header")).thenReturn("custom-value"); + + assertThat(request.getHeader("X-Custom-Header"), is("custom-value")); + assertThat(request.getHeader("X-Non-Existent"), is(nullValue())); } @Test public void getHeaderFromSuper() { when(servletRequest.getHeader("X-My-Header")).thenReturn("my-value"); - assertEquals("my-value", tusServletRequest.getHeader("X-My-Header")); + assertEquals("my-value", request.getHeader("X-My-Header")); } @Test @@ -68,11 +194,11 @@ public int read() throws IOException { }); // Read the whole input stream to parse trailers - InputStream contentInputStream = tusServletRequest.getContentInputStream(); + InputStream contentInputStream = request.getContentInputStream(); IOUtils.toByteArray(contentInputStream); // Verify trailer header is returned - assertEquals("trailer-value", tusServletRequest.getHeader("X-My-Trailer")); + assertEquals("trailer-value", request.getHeader("X-My-Trailer")); } @Test @@ -106,16 +232,16 @@ public int read() throws IOException { }); // Read the whole input stream to parse trailers - InputStream contentInputStream = tusServletRequest.getContentInputStream(); + InputStream contentInputStream = request.getContentInputStream(); IOUtils.toByteArray(contentInputStream); // Verify trailer header is returned because super returned a blank string - assertEquals("trailer-value", tusServletRequest.getHeader("X-My-Trailer")); + assertEquals("trailer-value", request.getHeader("X-My-Trailer")); } @Test public void getHeaderNotFound() { when(servletRequest.getHeader("X-My-Header")).thenReturn(null); - assertNull(tusServletRequest.getHeader("X-My-Header")); + assertNull(request.getHeader("X-My-Header")); } }