Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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) {}
}
142 changes: 134 additions & 8 deletions src/test/java/me/desair/tus/server/util/TusServletRequestTest.java
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -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;
Expand All @@ -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<ChecksumAlgorithm> 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<ChecksumAlgorithm> 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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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"));
}
}
Loading