diff --git a/src/main/java/com/lion/lionwebsite/Util/FileDownload.java b/src/main/java/com/lion/lionwebsite/Util/FileDownload.java index 0969fde..7d61603 100644 --- a/src/main/java/com/lion/lionwebsite/Util/FileDownload.java +++ b/src/main/java/com/lion/lionwebsite/Util/FileDownload.java @@ -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 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); + } } } } diff --git a/src/test/java/com/lion/lionwebsite/Util/FileDownloadTest.java b/src/test/java/com/lion/lionwebsite/Util/FileDownloadTest.java new file mode 100644 index 0000000..a8378b0 --- /dev/null +++ b/src/test/java/com/lion/lionwebsite/Util/FileDownloadTest.java @@ -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()); + } +}