修复断点下载范围解析与读取越界

This commit is contained in:
root
2026-09-08 09:21:01 +08:00
parent f65c5ad860
commit 7f823b6150
2 changed files with 125 additions and 99 deletions
@@ -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());
}
}