修复断点下载范围解析与读取越界
This commit is contained in:
@@ -1,119 +1,83 @@
|
||||
package com.lion.lionwebsite.Util;
|
||||
|
||||
import cn.hutool.core.util.StrUtil;
|
||||
import cn.hutool.core.util.URLUtil;
|
||||
import jakarta.servlet.http.HttpServletRequest;
|
||||
import jakarta.servlet.http.HttpServletResponse;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.apache.catalina.connector.ClientAbortException;
|
||||
import org.springframework.http.ContentDisposition;
|
||||
import org.springframework.http.HttpHeaders;
|
||||
import org.springframework.http.HttpRange;
|
||||
|
||||
import java.io.BufferedOutputStream;
|
||||
import java.io.File;
|
||||
import java.io.IOException;
|
||||
import java.io.RandomAccessFile;
|
||||
|
||||
import java.io.*;
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.util.List;
|
||||
|
||||
@Slf4j
|
||||
public class FileDownload {
|
||||
public static void export(HttpServletRequest request, HttpServletResponse response, String path) {
|
||||
File file = new File(path);
|
||||
|
||||
String fileName = file.getName();
|
||||
|
||||
String range = request.getHeader(HttpHeaders.RANGE);
|
||||
|
||||
String rangeSeparator = "-";
|
||||
// 开始下载位置
|
||||
long startByte = 0;
|
||||
// 结束下载位置
|
||||
long endByte = file.length() - 1;
|
||||
|
||||
// 如果是断点续传
|
||||
if (range != null && range.contains("bytes=") && range.contains(rangeSeparator)) {
|
||||
// 设置响应状态码为 206
|
||||
response.setStatus(HttpServletResponse.SC_PARTIAL_CONTENT);
|
||||
|
||||
range = range.substring(range.lastIndexOf("=") + 1).trim();
|
||||
String[] ranges = range.split(rangeSeparator);
|
||||
try {
|
||||
// 判断 range 的类型
|
||||
if (ranges.length == 1) {
|
||||
// 类型一:bytes=-2343
|
||||
if (range.startsWith(rangeSeparator)) {
|
||||
endByte = Long.parseLong(ranges[0]);
|
||||
}
|
||||
// 类型二:bytes=2343-
|
||||
else if (range.endsWith(rangeSeparator)) {
|
||||
startByte = Long.parseLong(ranges[0]);
|
||||
}
|
||||
}
|
||||
// 类型三:bytes=22-2343
|
||||
else if (ranges.length == 2) {
|
||||
startByte = Long.parseLong(ranges[0]);
|
||||
endByte = Long.parseLong(ranges[1]);
|
||||
}
|
||||
} catch (NumberFormatException e) {
|
||||
// 传参不规范,则直接返回所有内容
|
||||
startByte = 0;
|
||||
endByte = file.length() - 1;
|
||||
}
|
||||
} else {
|
||||
// 没有 ranges 即全部一次性传输,需要用 200 状态码,这一行应该可以省掉,因为默认返回是 200 状态码
|
||||
response.setStatus(HttpServletResponse.SC_OK);
|
||||
if (!file.isFile()) {
|
||||
response.setStatus(HttpServletResponse.SC_NOT_FOUND);
|
||||
return;
|
||||
}
|
||||
|
||||
//要下载的长度(endByte 为总长度 -1,这时候要加回去)
|
||||
long contentLength = endByte - startByte + 1;
|
||||
//文件类型
|
||||
String contentType = request.getServletContext().getMimeType(fileName);
|
||||
|
||||
if (StrUtil.isEmpty(contentType)) {
|
||||
contentType = "attachment";
|
||||
}
|
||||
|
||||
response.setHeader(HttpHeaders.ACCEPT_RANGES, "bytes");
|
||||
response.setHeader(HttpHeaders.CONTENT_TYPE, contentType);
|
||||
// 这里文件名换你想要的,inline 表示浏览器可以直接使用
|
||||
// 参考资料:https://developer.mozilla.org/zh-CN/docs/Web/HTTP/Headers/Content-Disposition
|
||||
response.setHeader(HttpHeaders.CONTENT_DISPOSITION, contentType + ";filename=\"" + URLUtil.encode(fileName) + "\"");
|
||||
response.setHeader(HttpHeaders.CONTENT_LENGTH, String.valueOf(contentLength));
|
||||
// [要下载的开始位置]-[结束位置]/[文件总大小]
|
||||
response.setHeader(HttpHeaders.CONTENT_RANGE, "bytes " + startByte + rangeSeparator + endByte + "/" + file.length());
|
||||
|
||||
BufferedOutputStream outputStream;
|
||||
//已传送数据大小
|
||||
long transmitted = 0;
|
||||
try (RandomAccessFile randomAccessFile = new RandomAccessFile(file, "r")) {
|
||||
try {
|
||||
outputStream = new BufferedOutputStream(response.getOutputStream());
|
||||
byte[] buff = new byte[4096];
|
||||
int len = 0;
|
||||
randomAccessFile.seek(startByte);
|
||||
while ((transmitted + len) <= contentLength && (len = randomAccessFile.read(buff)) != -1) {
|
||||
outputStream.write(buff, 0, len);
|
||||
transmitted += len;
|
||||
// 本地测试, 防止下载速度过快
|
||||
// Thread.sleep(1);
|
||||
// Size and content refer to the same opened file, even if a cache is replaced.
|
||||
try (RandomAccessFile input = new RandomAccessFile(file, "r")) {
|
||||
long size = input.length();
|
||||
long start = 0;
|
||||
long end = size - 1;
|
||||
boolean partial = false;
|
||||
String range = request.getHeader(HttpHeaders.RANGE);
|
||||
if (range != null && range.startsWith("bytes=")) {
|
||||
try {
|
||||
List<HttpRange> ranges = HttpRange.parseRanges(range);
|
||||
// Multiple ranges are intentionally ignored; send the full representation.
|
||||
if (ranges.size() == 1) {
|
||||
if (size == 0) throw new IllegalArgumentException("empty file");
|
||||
start = ranges.getFirst().getRangeStart(size);
|
||||
end = ranges.getFirst().getRangeEnd(size);
|
||||
if (start < 0 || start >= size || end < start)
|
||||
throw new IllegalArgumentException("unsatisfiable range");
|
||||
partial = true;
|
||||
}
|
||||
} catch (IllegalArgumentException e) {
|
||||
response.setStatus(HttpServletResponse.SC_REQUESTED_RANGE_NOT_SATISFIABLE);
|
||||
response.setHeader(HttpHeaders.CONTENT_RANGE, "bytes */" + size);
|
||||
response.setContentLengthLong(0);
|
||||
return;
|
||||
}
|
||||
// 处理不足 buff.length 部分
|
||||
if (transmitted < contentLength) {
|
||||
len = randomAccessFile.read(buff, 0, (int) (contentLength - transmitted));
|
||||
outputStream.write(buff, 0, len);
|
||||
}
|
||||
|
||||
outputStream.flush();
|
||||
response.flushBuffer();
|
||||
randomAccessFile.close();
|
||||
// log.trace("下载完毕: {}-{}, 已传输 {}", startByte, endByte, transmitted);
|
||||
} catch (ClientAbortException e) {
|
||||
// ignore 用户停止下载
|
||||
// log.trace("用户停止下载: {}-{}, 已传输 {}", startByte, endByte, transmitted);
|
||||
} catch (IOException e) {
|
||||
log.error("文件下载IO错误: {}", path, e);
|
||||
}
|
||||
long remaining = end - start + 1;
|
||||
response.setStatus(partial ? HttpServletResponse.SC_PARTIAL_CONTENT : HttpServletResponse.SC_OK);
|
||||
response.setHeader(HttpHeaders.ACCEPT_RANGES, "bytes");
|
||||
String mime = request.getServletContext().getMimeType(file.getName());
|
||||
response.setContentType(mime == null ? "application/octet-stream" : mime);
|
||||
response.setHeader(HttpHeaders.CONTENT_DISPOSITION,
|
||||
ContentDisposition.inline().filename(file.getName(), StandardCharsets.UTF_8).build().toString());
|
||||
response.setContentLengthLong(remaining);
|
||||
if (partial)
|
||||
response.setHeader(HttpHeaders.CONTENT_RANGE, "bytes " + start + "-" + end + "/" + size);
|
||||
if ("HEAD".equalsIgnoreCase(request.getMethod()))
|
||||
return;
|
||||
input.seek(start);
|
||||
BufferedOutputStream output = new BufferedOutputStream(response.getOutputStream());
|
||||
byte[] buffer = new byte[8192];
|
||||
while (remaining > 0) {
|
||||
int count = input.read(buffer, 0, (int) Math.min(buffer.length, remaining));
|
||||
if (count == -1)
|
||||
throw new EOFException("File changed during download");
|
||||
output.write(buffer, 0, count);
|
||||
remaining -= count;
|
||||
}
|
||||
output.flush();
|
||||
response.flushBuffer();
|
||||
} catch (ClientAbortException e) {
|
||||
// The client cancelled its download.
|
||||
} catch (IOException e) {
|
||||
log.warn("关闭RandomAccessFile失败: {}", path, e);
|
||||
log.warn("文件下载失败: {}", path, e);
|
||||
if (!response.isCommitted()) {
|
||||
response.reset();
|
||||
response.setStatus(HttpServletResponse.SC_INTERNAL_SERVER_ERROR);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
package com.lion.lionwebsite.Util;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.junit.jupiter.api.io.TempDir;
|
||||
import org.springframework.mock.web.MockHttpServletRequest;
|
||||
import org.springframework.mock.web.MockHttpServletResponse;
|
||||
import java.nio.file.*;
|
||||
import java.util.Arrays;
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
|
||||
class FileDownloadTest {
|
||||
@TempDir Path directory;
|
||||
|
||||
private MockHttpServletResponse download(String range, int size, String method) throws Exception {
|
||||
byte[] bytes = new byte[size];
|
||||
for (int i = 0; i < size; i++) bytes[i] = (byte) i;
|
||||
Path file = directory.resolve("sample.bin");
|
||||
Files.write(file, bytes);
|
||||
MockHttpServletRequest request = new MockHttpServletRequest(method, "/file");
|
||||
if (range != null) request.addHeader("Range", range);
|
||||
MockHttpServletResponse response = new MockHttpServletResponse();
|
||||
FileDownload.export(request, response, file.toString());
|
||||
return response;
|
||||
}
|
||||
|
||||
@Test void smallRangeDoesNotOverread() throws Exception {
|
||||
var response = download("bytes=10-109", 10_000, "GET");
|
||||
assertEquals(206, response.getStatus());
|
||||
assertEquals(100, response.getContentAsByteArray().length);
|
||||
assertEquals("bytes 10-109/10000", response.getHeader("Content-Range"));
|
||||
assertEquals(10, response.getContentAsByteArray()[0]);
|
||||
assertEquals(109, response.getContentAsByteArray()[99]);
|
||||
}
|
||||
|
||||
@Test void supportsSuffixAndOpenEndedRanges() throws Exception {
|
||||
var suffix = download("bytes=-10", 100, "GET");
|
||||
assertEquals("bytes 90-99/100", suffix.getHeader("Content-Range"));
|
||||
assertArrayEquals(download("bytes=90-", 100, "GET").getContentAsByteArray(), suffix.getContentAsByteArray());
|
||||
assertEquals(10, suffix.getContentAsByteArray().length);
|
||||
}
|
||||
|
||||
@Test void clampsEndAndRejectsInvalidRanges() throws Exception {
|
||||
assertEquals(10, download("bytes=90-999", 100, "GET").getContentAsByteArray().length);
|
||||
for (String range : Arrays.asList("bytes=100-", "bytes=9-2", "bytes=-0", "bytes=oops")) {
|
||||
var response = download(range, 100, "GET");
|
||||
assertEquals(416, response.getStatus(), range);
|
||||
assertEquals("bytes */100", response.getHeader("Content-Range"));
|
||||
assertEquals(0, response.getContentAsByteArray().length);
|
||||
}
|
||||
}
|
||||
|
||||
@Test void handlesFullEmptyHeadAndMultipleRanges() throws Exception {
|
||||
var full = download(null, 100, "GET");
|
||||
assertEquals(200, full.getStatus());
|
||||
assertNull(full.getHeader("Content-Range"));
|
||||
assertEquals(100, full.getContentAsByteArray().length);
|
||||
assertEquals(0, download(null, 0, "GET").getContentAsByteArray().length);
|
||||
assertEquals(416, download("bytes=0-", 0, "GET").getStatus());
|
||||
assertEquals(0, download(null, 100, "HEAD").getContentAsByteArray().length);
|
||||
assertEquals(200, download("bytes=0-1,5-6", 100, "GET").getStatus());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user