diff --git a/src/network/NetworkClient.java b/src/network/NetworkClient.java index ed429310..c676a15f 100644 --- a/src/network/NetworkClient.java +++ b/src/network/NetworkClient.java @@ -36,8 +36,11 @@ import java.nio.ByteBuffer; import java.nio.ByteOrder; import java.util.LinkedList; import java.util.List; +import java.util.Queue; +import java.util.concurrent.locks.ReentrantLock; import resources.control.Intent; +import resources.server_info.Log; import network.encryption.Compression; import network.packets.Packet; import network.packets.swg.SWGPacket; @@ -48,11 +51,13 @@ public class NetworkClient { private static final int DEFAULT_BUFFER = 128; private final Object prevPacketIntentMutex = new Object(); - private final Object outboundMutex = new Object(); private final Object bufferMutex = new Object(); + private final ReentrantLock inboundLock = new ReentrantLock(true); + private final ReentrantLock outboundLock = new ReentrantLock(true); private final InetSocketAddress address; private final long networkId; private final PacketSender packetSender; + private final Queue outboundQueue; private Intent prevPacketIntent; private ByteBuffer buffer; private long lastBufferSizeModification; @@ -62,10 +67,21 @@ public class NetworkClient { this.networkId = networkId; this.packetSender = packetSender; this.buffer = ByteBuffer.allocate(DEFAULT_BUFFER); + this.outboundQueue = new LinkedList<>(); lastBufferSizeModification = System.nanoTime(); prevPacketIntent = null; } + public void close() { + synchronized (bufferMutex) { + buffer = ByteBuffer.allocate(0); + } + synchronized (prevPacketIntentMutex) { + prevPacketIntent = null; + } + outboundQueue.clear(); + } + public InetSocketAddress getAddress() { return address; } @@ -81,27 +97,28 @@ public class NetworkClient { } } - public void sendPacket(Packet p) { - byte [] encoded = p.encode().array(); - int decompressedLength = encoded.length; - boolean compressed = encoded.length >= 16; - if (compressed) { - byte [] compressedData = Compression.compress(encoded); - if (compressedData.length >= encoded.length) - compressed = false; - else - encoded = compressedData; + public void processOutbound() { + if (!outboundLock.tryLock()) + return; + try { + Packet p; + while (!outboundQueue.isEmpty()) { + p = outboundQueue.poll(); + if (p == null) + break; + sendPacket(p); + } + } finally { + outboundLock.unlock(); } - ByteBuffer data = ByteBuffer.allocate(encoded.length + 5).order(ByteOrder.LITTLE_ENDIAN); - byte bitmask = 0; - bitmask |= (compressed?1:0) << 0; // Compressed - bitmask |= 1 << 1; // SWG - data.put(bitmask); - data.putShort((short) encoded.length); - data.putShort((short) decompressedLength); - data.put(encoded); - synchronized (outboundMutex) { - packetSender.sendPacket(address, data.array()); + } + + public void addToOutbound(Packet packet) { + outboundLock.lock(); + try { + outboundQueue.add(packet); + } finally { + outboundLock.unlock(); } } @@ -125,23 +142,30 @@ public class NetworkClient { } } - public boolean process() { - List packets; - synchronized (bufferMutex) { - buffer.flip(); - packets = processPackets(); - buffer.compact(); - } - synchronized (prevPacketIntentMutex) { - for (Packet p : packets) { - p.setAddress(address.getAddress()); - p.setPort(address.getPort()); - InboundPacketIntent i = new InboundPacketIntent(p, networkId); - i.broadcastAfterIntent(prevPacketIntent); - prevPacketIntent = i; + public boolean processInbound() { + if (!inboundLock.tryLock()) + return false; + try { + List packets; + synchronized (bufferMutex) { + buffer.flip(); + packets = processPackets(); + buffer.compact(); } + synchronized (prevPacketIntentMutex) { + for (Packet p : packets) { + p.setAddress(address.getAddress()); + p.setPort(address.getPort()); + Log.d("NetworkClient", "Inbound: %s", p.getClass().getSimpleName()); + InboundPacketIntent i = new InboundPacketIntent(p, networkId); + i.broadcastAfterIntent(prevPacketIntent); + prevPacketIntent = i; + } + } + return packets.size() > 0; + } finally { + inboundLock.unlock(); } - return packets.size() > 0; } private void shrinkBuffer() { @@ -219,6 +243,44 @@ public class NetworkClient { } } + private void sendPacket(Packet p) { + ByteBuffer encoded = p.encode(); + encoded.position(0); + int decompressedLength = encoded.remaining(); + boolean compressed = decompressedLength >= 50; + if (compressed) { + ByteBuffer compress = compress(encoded); + compressed = compress != encoded; + encoded = compress; + } + Log.d("NetworkClient", "Outbound: %s", p.getClass().getSimpleName()); + sendPacket(encoded, compressed, decompressedLength); + } + + private ByteBuffer compress(ByteBuffer data) { + ByteBuffer compressedBuffer = ByteBuffer.allocate(Compression.getMaxCompressedLength(data.remaining())); + int length = Compression.compress(data.array(), compressedBuffer.array()); + compressedBuffer.position(0); + compressedBuffer.limit(length); + if (length >= data.remaining()) + return data; + else + return compressedBuffer; + } + + private void sendPacket(ByteBuffer packet, boolean compressed, int rawLength) { + ByteBuffer data = ByteBuffer.allocate(packet.remaining() + 5).order(ByteOrder.LITTLE_ENDIAN); + byte bitmask = 0; + bitmask |= (compressed?1:0) << 0; // Compressed + bitmask |= 1 << 1; // SWG + data.put(bitmask); + data.putShort((short) packet.remaining()); + data.putShort((short) rawLength); + data.put(packet); + data.flip(); + packetSender.sendPacket(address, data); + } + public String toString() { return "NetworkClient["+address+"]"; } diff --git a/src/network/PacketSender.java b/src/network/PacketSender.java index c69ad37b..2b85a508 100644 --- a/src/network/PacketSender.java +++ b/src/network/PacketSender.java @@ -1,9 +1,10 @@ package network; import java.net.InetSocketAddress; +import java.nio.ByteBuffer; public interface PacketSender { - void sendPacket(InetSocketAddress sock, byte [] data); + void sendPacket(InetSocketAddress sock, ByteBuffer data); } diff --git a/src/network/encryption/Compression.java b/src/network/encryption/Compression.java index 0f0a5f48..5fe5c980 100644 --- a/src/network/encryption/Compression.java +++ b/src/network/encryption/Compression.java @@ -9,13 +9,12 @@ public class Compression { private static final LZ4Compressor COMPRESSOR = LZ4Factory.safeInstance().highCompressor(); private static final LZ4SafeDecompressor DECOMPRESSOR = LZ4Factory.safeInstance().safeDecompressor(); - public static byte [] compress(byte [] data) { - int maxCompressedLength = COMPRESSOR.maxCompressedLength(data.length); - byte[] compressed = new byte[maxCompressedLength]; - int length = COMPRESSOR.compress(data, compressed); - byte [] ret = new byte[length]; - System.arraycopy(compressed, 0, ret, 0, length); - return ret; + public static int getMaxCompressedLength(int len) { + return COMPRESSOR.maxCompressedLength(len); + } + + public static int compress(byte [] data, byte [] buffer) { + return COMPRESSOR.compress(data, buffer); } public static byte [] decompress(byte [] data) { diff --git a/src/resources/network/TCPServer.java b/src/resources/network/TCPServer.java index 620872ec..8d401216 100644 --- a/src/resources/network/TCPServer.java +++ b/src/resources/network/TCPServer.java @@ -40,14 +40,18 @@ import java.nio.channels.Selector; import java.nio.channels.ServerSocketChannel; import java.nio.channels.SocketChannel; import java.util.HashMap; -import java.util.Iterator; import java.util.Locale; import java.util.Map; -import java.util.Set; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; + +import resources.server_info.Log; +import utilities.ThreadUtilities; public class TCPServer { + private final ExecutorService callbackExecutor; private final Map sockets; private final InetAddress addr; private final int port; @@ -61,6 +65,7 @@ public class TCPServer { } public TCPServer(InetAddress addr, int port, int bufferSize) { + this.callbackExecutor = Executors.newSingleThreadExecutor(ThreadUtilities.newThreadFactory("tcp-server-callback-executor")); this.sockets = new HashMap<>(); this.addr = addr; this.port = port; @@ -95,11 +100,14 @@ public class TCPServer { public boolean disconnect(SocketAddress sock) { synchronized (sockets) { SocketChannel sc = sockets.get(sock); + if (sc == null) + return false; sockets.remove(sock); try { + Socket s = sc.socket(); sc.close(); if (callback != null) - callback.onConnectionDisconnect(sc.socket()); + callbackExecutor.execute(() -> callback.onConnectionDisconnect(s, sock)); return true; } catch (IOException e) { e.printStackTrace(); @@ -128,18 +136,17 @@ public class TCPServer { return false; } - public boolean send(InetSocketAddress sock, byte [] data) { + public boolean send(InetSocketAddress sock, ByteBuffer data) { synchronized (sockets) { SocketChannel sc = sockets.get(sock); try { if (sc != null && sc.isConnected()) { - ByteBuffer bb = ByteBuffer.wrap(data); - while (bb.hasRemaining()) - sc.write(bb); + while (data.hasRemaining()) + sc.write(data); return true; } } catch (IOException e) { - e.printStackTrace(); + Log.e("TCPServer", "Terminated connection with %s. Error: %s", sock.toString(), e.getMessage()); disconnect(sc); } } @@ -152,7 +159,7 @@ public class TCPServer { public interface TCPCallback { void onIncomingConnection(Socket s); - void onConnectionDisconnect(Socket s); + void onConnectionDisconnect(Socket s, SocketAddress addr); void onIncomingData(Socket s, byte [] data); } @@ -208,37 +215,43 @@ public class TCPServer { } private void processSelectionKeys(Selector selector) throws ClosedChannelException { - Set keys = selector.selectedKeys(); - Iterator it = keys.iterator(); - while (it.hasNext()) { - SelectionKey key = it.next(); + for (SelectionKey key : selector.selectedKeys()) { + if (!key.isValid()) + continue; if (key.isAcceptable()) { accept(selector); } else if (key.isReadable()) { SelectableChannel selectable = key.channel(); - if (selectable instanceof SocketChannel) - read(key, (SocketChannel) selectable); + if (selectable instanceof SocketChannel) { + boolean canRead = true; + while (canRead) + canRead = read(key, (SocketChannel) selectable); + } } - it.remove(); } } private void accept(Selector selector) { try { - SocketChannel sc = channel.accept(); - if (sc == null) - return; - sc.configureBlocking(false); - sc.register(selector, SelectionKey.OP_READ); - sockets.put(sc.getRemoteAddress(), sc); - if (callback != null) - callback.onIncomingConnection(sc.socket()); + while (channel.isOpen()) { + SocketChannel sc = channel.accept(); + if (sc == null) + break; + SocketChannel old = sockets.get(sc.getRemoteAddress()); + if (old != null) + disconnect(old); + sockets.put(sc.getRemoteAddress(), sc); + sc.configureBlocking(false); + sc.register(selector, SelectionKey.OP_READ); + if (callback != null) + callbackExecutor.execute(() -> callback.onIncomingConnection(sc.socket())); + } } catch (IOException e) { e.printStackTrace(); } } - private void read(SelectionKey key, SocketChannel s) { + private boolean read(SelectionKey key, SocketChannel s) { try { buffer.position(0); buffer.limit(bufferSize); @@ -251,17 +264,19 @@ public class TCPServer { ByteBuffer smaller = ByteBuffer.allocate(n); smaller.put(buffer); if (callback != null) - callback.onIncomingData(s.socket(), smaller.array()); + callbackExecutor.execute(() -> callback.onIncomingData(s.socket(), smaller.array())); + return true; } } catch (IOException e) { if (e.getMessage().toLowerCase(Locale.US).contains("connection reset")) - System.err.println("Connection Reset"); - else - e.printStackTrace(); - System.err.flush(); + Log.e("TCPServer", "Connection Reset with %s", s.socket().getRemoteSocketAddress()); + else { + Log.e("TCPServer", e); + } key.cancel(); disconnect(s); } + return false; } } diff --git a/src/services/network/NetworkClientManager.java b/src/services/network/NetworkClientManager.java index bfa61582..0ff496e8 100644 --- a/src/services/network/NetworkClientManager.java +++ b/src/services/network/NetworkClientManager.java @@ -35,6 +35,7 @@ import java.io.IOException; import java.net.InetSocketAddress; import java.net.Socket; import java.net.SocketAddress; +import java.nio.ByteBuffer; import java.util.HashMap; import java.util.Hashtable; import java.util.LinkedList; @@ -61,33 +62,26 @@ public class NetworkClientManager extends Manager implements TCPCallback, Packet private final Map sockets; private final Map clients; - private final Queue processQueue; - private final ExecutorService clientProcessor; + private final Queue inboundQueue; + private final Queue outboundQueue; + private final ExecutorService inboundProcessor; + private final ExecutorService outboundProcessor; private final Runnable processBufferRunnable; + private final Runnable processOutboundRunnable; private final AtomicLong networkIdCounter; private final TCPServer tcpServer; public NetworkClientManager() { + final int threadCount = getConfig(ConfigFile.NETWORK).getInt("PACKET-THREAD-COUNT", 10); sockets = new HashMap(); clients = new Hashtable(); - processQueue = new LinkedList<>(); + inboundQueue = new LinkedList<>(); + outboundQueue = new LinkedList<>(); networkIdCounter = new AtomicLong(1); - clientProcessor = Executors.newFixedThreadPool(Runtime.getRuntime().availableProcessors(), ThreadUtilities.newThreadFactory("packet-processor-%d")); - processBufferRunnable = new Runnable() { - public void run() { - try { - NetworkClient client; - synchronized (processQueue) { - client = processQueue.poll(); - if (client == null) - return; - } - client.process(); - } catch (Exception e) { - e.printStackTrace(); - } - } - }; + inboundProcessor = Executors.newFixedThreadPool(threadCount/10, ThreadUtilities.newThreadFactory("inbound-packet-processor-%d")); + outboundProcessor = Executors.newFixedThreadPool(threadCount, ThreadUtilities.newThreadFactory("outbound-packet-processor-%d")); + processBufferRunnable = () -> processBufferRunnable(); + processOutboundRunnable = () -> processOutboundRunnable(); tcpServer = new TCPServer(getBindPort(), getBufferSize()); registerForIntent(OutboundPacketIntent.TYPE); @@ -114,10 +108,12 @@ public class NetworkClientManager extends Manager implements TCPCallback, Packet @Override public boolean terminate() { - clientProcessor.shutdownNow(); + inboundProcessor.shutdownNow(); + outboundProcessor.shutdownNow(); boolean success = true; try { - success = clientProcessor.awaitTermination(5, TimeUnit.SECONDS); + success = inboundProcessor.awaitTermination(5, TimeUnit.SECONDS); + success = outboundProcessor.awaitTermination(5, TimeUnit.SECONDS) && success; } catch (InterruptedException e) { e.printStackTrace(); } @@ -146,8 +142,8 @@ public class NetworkClientManager extends Manager implements TCPCallback, Packet } @Override - public void onConnectionDisconnect(Socket s) { - SocketAddress addr = s.getRemoteSocketAddress(); + public void onConnectionDisconnect(Socket s, SocketAddress addr) { + Log.i(this, "Disconnected from %s", addr); if (addr instanceof InetSocketAddress) onSessionDisconnect((InetSocketAddress) addr); else if (addr != null) @@ -164,7 +160,7 @@ public class NetworkClientManager extends Manager implements TCPCallback, Packet } @Override - public void sendPacket(InetSocketAddress sock, byte[] data) { + public void sendPacket(InetSocketAddress sock, ByteBuffer data) { tcpServer.send(sock, data); } @@ -182,18 +178,20 @@ public class NetworkClientManager extends Manager implements TCPCallback, Packet sockets.put(address, networkId); clients.put(networkId, client); } + Log.d(this, "Created " + client.getAddress()); client.onConnected(); } private void onSessionDisconnect(InetSocketAddress address) { + Long networkId; synchronized (clients) { - Long networkId = sockets.get(address); - if (networkId != null) { - deleteSession(networkId); - new ConnectionClosedIntent(networkId, DisconnectReason.OTHER_SIDE_TERMINATED).broadcast(); - } else { - System.err.println("Network ID not found for " + address + "!"); - } + networkId = sockets.get(address); + } + if (networkId != null) { + deleteSession(networkId); + new ConnectionClosedIntent(networkId, DisconnectReason.OTHER_SIDE_TERMINATED).broadcast(); + } else { + System.err.println("Network ID not found for " + address + "!"); } } @@ -204,36 +202,80 @@ public class NetworkClientManager extends Manager implements TCPCallback, Packet System.err.println("No NetworkClient found for network id: " + networkId); return; } + Log.d(this, "Deleted " + client.getAddress()); sockets.remove(client.getAddress()); + synchronized (inboundQueue) { + inboundQueue.remove(client); + } + synchronized (outboundQueue) { + outboundQueue.remove(client); + } + client.close(); } } private void handleOutboundPacket(long networkId, Packet p) { + NetworkClient client; synchronized (clients) { - NetworkClient client = clients.get(networkId); - if (client != null) - client.sendPacket(p); - else - Log.w(this, "NetworkClient does not exist for ID: %d", networkId); + client = clients.get(networkId); + } + if (client != null) { + client.addToOutbound(p); + synchronized (outboundQueue) { + while (outboundQueue.remove(client)); + outboundQueue.add(client); + } + outboundProcessor.execute(processOutboundRunnable); } } private void handleIncomingData(InetSocketAddress addr, byte [] data) { + Long netId = sockets.get(addr); + if (netId == null) { + Log.w(this, "Unknown socket address! Address: %s", addr); + return; + } + NetworkClient client; synchronized (clients) { - Long netId = sockets.get(addr); - if (netId == null) { - Log.w(this, "Unknown socket address! Address: %s", addr); - return; + client = clients.get(netId); + } + if (client != null) { + client.addToBuffer(data); + synchronized (inboundQueue) { + inboundQueue.add(client); } - NetworkClient client = clients.get(netId); - if (client != null) { - client.addToBuffer(data); - synchronized (processQueue) { - processQueue.add(client); - } - clientProcessor.execute(processBufferRunnable); - } else - Log.w(this, "Unknown connection! Network ID: %d Address: %s", netId, addr); + inboundProcessor.execute(processBufferRunnable); + } else + Log.w(this, "Unknown connection! Network ID: %d Address: %s", netId, addr); + } + + private void processBufferRunnable() { + try { + NetworkClient client; + synchronized (inboundQueue) { + client = inboundQueue.poll(); + if (client == null) + return; + } + client.processInbound(); + } catch (Exception e) { + e.printStackTrace(); + Log.e(this, e); + } + } + + private void processOutboundRunnable() { + try { + NetworkClient client; + synchronized (outboundQueue) { + client = outboundQueue.poll(); + if (client == null) + return; + } + client.processOutbound(); + } catch (Exception e) { + e.printStackTrace(); + Log.e(this, e); } } diff --git a/test/network/encryption/TestCompression.java b/test/network/encryption/TestCompression.java index 65760e46..f18a27f7 100644 --- a/test/network/encryption/TestCompression.java +++ b/test/network/encryption/TestCompression.java @@ -10,30 +10,39 @@ public class TestCompression { @Test public void testSmall() { + byte [] buffer = new byte[Compression.getMaxCompressedLength(101)]; byte [] data = new byte[101]; for (int i = 0; i < data.length; i++) data[i] = (byte) (i % 100); - byte [] compressed = Compression.compress(data); + int len = Compression.compress(data, buffer); + byte [] compressed = new byte[len]; + System.arraycopy(buffer, 0, compressed, 0, len); byte [] decompressed = Compression.decompress(compressed); Assert.assertArrayEquals(data, decompressed); } @Test public void testMedium() { + byte [] buffer = new byte[Compression.getMaxCompressedLength(256)]; byte [] data = new byte[256]; for (int i = 0; i < data.length; i++) data[i] = (byte) (i % 100); - byte [] compressed = Compression.compress(data); + int len = Compression.compress(data, buffer); + byte [] compressed = new byte[len]; + System.arraycopy(buffer, 0, compressed, 0, len); byte [] decompressed = Compression.decompress(compressed); Assert.assertArrayEquals(data, decompressed); } @Test public void testLarge() { + byte [] buffer = new byte[Compression.getMaxCompressedLength(1024)]; byte [] data = new byte[1024]; for (int i = 0; i < data.length; i++) data[i] = (byte) (i % 100); - byte [] compressed = Compression.compress(data); + int len = Compression.compress(data, buffer); + byte [] compressed = new byte[len]; + System.arraycopy(buffer, 0, compressed, 0, len); Assert.assertTrue("Compressed should be less than actual. Compressed: "+compressed.length+" Data: "+data.length, compressed.length < data.length); byte [] decompressed = Compression.decompress(compressed); Assert.assertArrayEquals(data, decompressed);