diff --git a/src/main/java/net/forgecraft/serverpacklocator/client/MultiThreadedDownloader.java b/src/main/java/net/forgecraft/serverpacklocator/client/MultiThreadedDownloader.java index 9108788..7904bc8 100644 --- a/src/main/java/net/forgecraft/serverpacklocator/client/MultiThreadedDownloader.java +++ b/src/main/java/net/forgecraft/serverpacklocator/client/MultiThreadedDownloader.java @@ -4,23 +4,26 @@ import net.forgecraft.serverpacklocator.FileChecksumValidator; import net.forgecraft.serverpacklocator.ServerManifest; import net.forgecraft.serverpacklocator.secure.IConnectionSecurityManager; +import net.forgecraft.serverpacklocator.utils.CompressionUtils; import net.neoforged.fml.loading.progress.StartupNotificationManager; import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; import java.io.IOException; +import java.io.InputStream; import java.net.URI; import java.net.URLEncoder; import java.net.http.HttpClient; import java.net.http.HttpRequest; import java.net.http.HttpResponse; import java.nio.ByteBuffer; +import java.nio.channels.Channels; +import java.nio.channels.FileChannel; import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Path; import java.nio.file.StandardOpenOption; import java.util.ArrayList; -import java.util.Base64; import java.util.List; import java.util.Objects; import java.util.OptionalLong; @@ -64,10 +67,35 @@ private PreparedServerDownloadData downloadManifest() throws IOException, Interr authenticate(); var progressBar = StartupNotificationManager.addProgressBar("Requesting server manifest...", 1); try { - var response = makeRequest("servermanifest.json", true, HttpResponse.BodyHandlers.ofString(StandardCharsets.UTF_8)); + var response = makeRequest("servermanifest.json", true, responseInfo -> HttpResponse.BodySubscribers.mapping(HttpResponse.BodySubscribers.ofInputStream(), in -> { + Function decompressor = s -> s; + var encoding = responseInfo.headers().firstValue("Content-Encoding").orElse(null); + if (encoding != null) { + var methodNames = encoding.split(","); + for (int i = methodNames.length - 1; i >= 0; i--) { + var method = CompressionUtils.methodFromName(methodNames[i].trim()); + if (method != null) { + decompressor = decompressor.andThen(s -> { + try { + return method.decompress(s); + } catch (IOException e) { + LOGGER.error(e); + throw new RuntimeException(e); + } + }); + } + } + } + + return decompressor.apply(in); + })); + this.connectionSecurityManager.handleClientResponse(response); - var serverManifest = ServerManifest.fromString(response.body()); + ServerManifest serverManifest; + try (var stream = response.body()) { + serverManifest = ServerManifest.fromString(String.valueOf(StandardCharsets.UTF_8.decode(ByteBuffer.wrap(stream.readAllBytes())))); + } // Write the file to the client system for debugging if (serverManifest != null) { @@ -96,6 +124,7 @@ private HttpResponse makeRequest(String path, boolean authenticated, Http LOGGER.info("ServerPackLocator is requesting {}...", requestUri); var request = HttpRequest.newBuilder(requestUri); + request.header("Accept-Encoding", "gzip"); this.connectionSecurityManager.decorateClientRequest(request, authenticated); var response = httpClient.send(request.build(), bodyHandler); this.connectionSecurityManager.handleClientResponse(response); @@ -265,9 +294,37 @@ private void downloadFile(final FileToDownload fileToDownload, final ProgressLis makeRequest( "files/" + URLEncoder.encode(nextFile, StandardCharsets.UTF_8).replace("+", "%20"), true, - progressListener.trackBodyHandler( - HttpResponse.BodyHandlers.ofFile(destinationPath, StandardOpenOption.WRITE, StandardOpenOption.CREATE, StandardOpenOption.TRUNCATE_EXISTING) - ) + progressListener.trackBodyHandler((HttpResponse.BodyHandler) responseInfo -> { + Function decompressor = s -> s; + var encoding = responseInfo.headers().firstValue("Content-Encoding").orElse(null); + if (encoding != null) { + var methodNames = encoding.split(","); + for (int i = methodNames.length - 1; i >= 0; i--) { + var method = CompressionUtils.methodFromName(methodNames[i].trim()); + if (method != null) { + decompressor = decompressor.andThen(s -> { + try { + return method.decompress(s); + } catch (IOException e) { + LOGGER.error(e); + throw new RuntimeException(e); + } + }); + } + } + } + + var downstream = HttpResponse.BodySubscribers.mapping(HttpResponse.BodySubscribers.ofInputStream(), decompressor); + return HttpResponse.BodySubscribers.mapping(downstream, s -> { + try (var out = Channels.newOutputStream(FileChannel.open(fileToDownload.localFile(), StandardOpenOption.WRITE, StandardOpenOption.CREATE, StandardOpenOption.TRUNCATE_EXISTING))){ + s.transferTo(out); + } catch (IOException e) { + throw new RuntimeException(e); + } + + return fileToDownload.localFile(); + }); + }) ); // Validate that the downloaded file actually matches the expected checksum diff --git a/src/main/java/net/forgecraft/serverpacklocator/server/RequestHandler.java b/src/main/java/net/forgecraft/serverpacklocator/server/RequestHandler.java index d3a0a1c..22e6285 100644 --- a/src/main/java/net/forgecraft/serverpacklocator/server/RequestHandler.java +++ b/src/main/java/net/forgecraft/serverpacklocator/server/RequestHandler.java @@ -1,19 +1,32 @@ package net.forgecraft.serverpacklocator.server; -import net.forgecraft.serverpacklocator.ModAccessor; -import net.forgecraft.serverpacklocator.secure.IConnectionSecurityManager; import io.netty.buffer.ByteBuf; import io.netty.buffer.Unpooled; import io.netty.channel.ChannelHandlerContext; -import io.netty.channel.DefaultFileRegion; import io.netty.channel.SimpleChannelInboundHandler; -import io.netty.handler.codec.http.*; +import io.netty.handler.codec.http.DefaultFullHttpResponse; +import io.netty.handler.codec.http.DefaultHttpResponse; +import io.netty.handler.codec.http.FullHttpRequest; +import io.netty.handler.codec.http.FullHttpResponse; +import io.netty.handler.codec.http.HttpChunkedInput; +import io.netty.handler.codec.http.HttpHeaderNames; +import io.netty.handler.codec.http.HttpHeaderValues; +import io.netty.handler.codec.http.HttpMethod; +import io.netty.handler.codec.http.HttpResponse; +import io.netty.handler.codec.http.HttpResponseStatus; +import io.netty.handler.codec.http.HttpUtil; +import io.netty.handler.codec.http.HttpVersion; +import io.netty.handler.stream.ChunkedStream; +import net.forgecraft.serverpacklocator.ModAccessor; +import net.forgecraft.serverpacklocator.secure.IConnectionSecurityManager; import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; import javax.net.ssl.SSLException; +import java.io.IOException; import java.net.URLDecoder; import java.nio.charset.StandardCharsets; +import java.nio.file.Files; import java.util.Objects; class RequestHandler extends SimpleChannelInboundHandler { @@ -106,15 +119,19 @@ private void buildReply(final ChannelHandlerContext ctx, final FullHttpRequest m } private void buildFileReply(final ChannelHandlerContext ctx, final FullHttpRequest msg, final ServerFileManager.ExposedFile file) { - final HttpResponse response = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - HttpUtil.setKeepAlive(response, HttpUtil.isKeepAlive(msg)); - response.headers().set(HttpHeaderNames.CONTENT_TYPE, "application/octet-stream"); - response.headers().set("filename", file.name()); - HttpUtil.setContentLength(response, file.size()); - this.connectionSecurityManager.decorateServerResponse(ctx, msg, response); - - ctx.write(response); - ctx.write(new DefaultFileRegion(file.path().toFile(), 0, file.size())); - ctx.writeAndFlush(LastHttpContent.EMPTY_LAST_CONTENT); + try (var fileStream = Files.newInputStream(file.path())) { + var out = new HttpChunkedInput(new ChunkedStream(fileStream)); + + final HttpResponse response = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); + HttpUtil.setKeepAlive(response, HttpUtil.isKeepAlive(msg)); + response.headers().set(HttpHeaderNames.CONTENT_TYPE, "application/octet-stream"); + response.headers().set("filename", file.name()); + response.headers().set(HttpHeaderNames.TRANSFER_ENCODING, HttpHeaderValues.CHUNKED); + + ctx.write(response); + ctx.writeAndFlush(out); + } catch (IOException e) { + throw new RuntimeException(e); + } } } diff --git a/src/main/java/net/forgecraft/serverpacklocator/server/SimpleHttpServer.java b/src/main/java/net/forgecraft/serverpacklocator/server/SimpleHttpServer.java index 2d39773..6fe6ae2 100644 --- a/src/main/java/net/forgecraft/serverpacklocator/server/SimpleHttpServer.java +++ b/src/main/java/net/forgecraft/serverpacklocator/server/SimpleHttpServer.java @@ -1,16 +1,21 @@ package net.forgecraft.serverpacklocator.server; -import net.forgecraft.serverpacklocator.secure.IConnectionSecurityManager; import com.google.common.util.concurrent.ThreadFactoryBuilder; import com.mojang.logging.LogUtils; import io.netty.bootstrap.ServerBootstrap; -import io.netty.channel.*; +import io.netty.channel.ChannelHandlerContext; +import io.netty.channel.ChannelInboundHandlerAdapter; +import io.netty.channel.ChannelInitializer; +import io.netty.channel.ChannelOption; import io.netty.channel.nio.NioEventLoopGroup; import io.netty.channel.socket.ServerSocketChannel; import io.netty.channel.socket.SocketChannel; import io.netty.channel.socket.nio.NioServerSocketChannel; +import io.netty.handler.codec.http.HttpContentCompressor; import io.netty.handler.codec.http.HttpObjectAggregator; import io.netty.handler.codec.http.HttpServerCodec; +import io.netty.handler.stream.ChunkedWriteHandler; +import net.forgecraft.serverpacklocator.secure.IConnectionSecurityManager; import org.slf4j.Logger; /** @@ -53,6 +58,8 @@ public void channelActive(final ChannelHandlerContext ctx) { @Override protected void initChannel(final SocketChannel ch) { ch.pipeline().addLast("codec", new HttpServerCodec()); + ch.pipeline().addLast("deflater", new HttpContentCompressor()); + ch.pipeline().addLast("chunkedWriter", new ChunkedWriteHandler()); ch.pipeline().addLast("aggregator", new HttpObjectAggregator(MAX_CONTENT_LENGTH)); ch.pipeline().addLast("request", new RequestHandler( securityManager, fileManager diff --git a/src/main/java/net/forgecraft/serverpacklocator/utils/CompressionUtils.java b/src/main/java/net/forgecraft/serverpacklocator/utils/CompressionUtils.java new file mode 100644 index 0000000..3292918 --- /dev/null +++ b/src/main/java/net/forgecraft/serverpacklocator/utils/CompressionUtils.java @@ -0,0 +1,48 @@ +package net.forgecraft.serverpacklocator.utils; + +import org.jetbrains.annotations.NotNull; + +import java.io.IOException; +import java.io.InputStream; +import java.io.OutputStream; +import java.util.zip.GZIPInputStream; +import java.util.zip.GZIPOutputStream; + +public class CompressionUtils { + public abstract static class CompressionMethod { + CompressionMethod() {} + + public abstract @NotNull String name(); + + public abstract @NotNull InputStream decompress(InputStream input) throws IOException; + + public abstract @NotNull OutputStream compress(OutputStream input) throws IOException; + } + + public static class Gzip extends CompressionMethod { + public static final CompressionMethod INSTANCE = new Gzip(); + public static final String NAME = "gzip"; + + @Override + public @NotNull String name() { + return NAME; + } + + @Override + public @NotNull InputStream decompress(InputStream input) throws IOException { + return new GZIPInputStream(input); + } + + @Override + public @NotNull OutputStream compress(OutputStream input) throws IOException { + return new GZIPOutputStream(input); + } + } + + public static CompressionMethod methodFromName(String name) { + return switch (name) { + case Gzip.NAME -> Gzip.INSTANCE; + case null, default -> null; + }; + } +}