From a67c3946f29fb7dbd20b38eb3bc8b45f0a48c8f7 Mon Sep 17 00:00:00 2001 From: Chris Speciale Date: Thu, 2 Jul 2026 22:02:14 -0400 Subject: [PATCH] Add eval socket compatibility shim --- src/snake/Run.hx | 2 +- src/snake/_internal/net/Socket.hx | 383 +++++++++++++++++++++++ src/snake/http/BaseHTTPRequestHandler.hx | 3 +- src/snake/socket/BaseServer.hx | 2 +- src/snake/socket/TCPServer.hx | 2 +- test-socket.hxml | 4 + tests/TestSocket.hx | 139 ++++++++ 7 files changed, 531 insertions(+), 4 deletions(-) create mode 100644 src/snake/_internal/net/Socket.hx create mode 100644 test-socket.hxml create mode 100644 tests/TestSocket.hx diff --git a/src/snake/Run.hx b/src/snake/Run.hx index 13e9a4c..77ff96a 100644 --- a/src/snake/Run.hx +++ b/src/snake/Run.hx @@ -6,7 +6,7 @@ import snake.http.HTTPServer; import snake.http.SimpleHTTPRequestHandler; import snake.socket.BaseRequestHandler; import sys.net.Host; -import sys.net.Socket; +import snake._internal.net.Socket; /** Run `haxelib run snake-server` to start a local HTTP server that serves diff --git a/src/snake/_internal/net/Socket.hx b/src/snake/_internal/net/Socket.hx new file mode 100644 index 0000000..7975c51 --- /dev/null +++ b/src/snake/_internal/net/Socket.hx @@ -0,0 +1,383 @@ +package snake._internal.net; + +#if (eval && (haxe_ver >= 4.2)) +import eval.luv.Buffer; +import eval.luv.Handle; +import eval.luv.Loop; +import eval.luv.Loop.RunMode; +import eval.luv.SockAddr; +import eval.luv.Stream; +import eval.luv.Tcp; +import haxe.Exception; +import haxe.io.Bytes; +import haxe.io.BytesBuffer; +import haxe.io.Eof; +import haxe.io.Input; +import haxe.io.Output; +import sys.net.Host; +import sys.net.Socket as SysSocket; + +using eval.luv.Result.ResultTools; + +class Socket extends SysSocket { + private static final POLL_INTERVAL = 0.001; + + private var loop:Loop; + private var tcp:Tcp; + private var pending:Array; + private var readBuffer:BytesBuffer; + private var readOffset = 0; + private var listening:Bool; + private var reading:Bool; + private var closed:Bool; + private var readClosed:Bool; + private var timeout:Null = null; + private var localAddress:{host:Host, port:Int}; + private var peerAddress:{host:Host, port:Int}; + + public function new() { + super(); + super.close(); + var loop = Loop.defaultLoop(); + initLuv(loop, Tcp.init(loop).resolve()); + } + + private function initLuv(loop:Loop, tcp:Tcp):Void { + this.loop = loop; + this.tcp = tcp; + pending = []; + readBuffer = new BytesBuffer(); + readOffset = 0; + listening = false; + reading = false; + closed = false; + readClosed = false; + input = new LuvSocketInput(this); + output = new LuvSocketOutput(this); + } + + private static function fromTcp(loop:Loop, tcp:Tcp):Socket { + var socket = Type.createEmptyInstance(Socket); + socket.initLuv(loop, tcp); + return socket; + } + + override public function close():Void { + if (closed) { + return; + } + closed = true; + Handle.close(tcp, () -> {}); + } + + override public function read():String { + return input.readAll().toString(); + } + + override public function write(content:String):Void { + output.writeString(content); + } + + override public function connect(host:Host, port:Int):Void { + var done = false; + var failure:Dynamic = null; + tcp.connect(sockAddr(host, port), result -> { + try { + result.resolve(); + } catch (e:Dynamic) { + failure = e; + } + done = true; + }); + waitUntil(() -> done, timeout); + if (failure != null) { + throw failure; + } + localAddress = null; + peerAddress = {host: host, port: port}; + startRead(); + } + + override public function listen(connections:Int):Void { + listening = true; + Stream.listen(tcp, result -> { + result.resolve(); + var client = Socket.fromTcp(loop, Tcp.init(loop).resolve()); + Stream.accept(tcp, client.tcp).resolve(); + client.localAddress = client.readLocalAddress(); + client.peerAddress = client.readPeerAddress(); + client.startRead(); + pending.push(client); + }, connections); + } + + override public function shutdown(read:Bool, write:Bool):Void { + if (write) { + Stream.shutdown(tcp, _ -> close()); + } + if (read) { + readClosed = true; + } + } + + override public function bind(host:Host, port:Int):Void { + tcp.bind(sockAddr(host, port)).resolve(); + localAddress = readLocalAddress(); + } + + override public function accept():Socket { + if (pending.length == 0) { + select([this], null, null, timeout); + } + if (pending.length == 0) { + throw new Exception("No pending connection"); + } + return pending.shift(); + } + + override public function peer():{host:Host, port:Int} { + if (peerAddress == null) { + peerAddress = readPeerAddress(); + } + return peerAddress; + } + + override public function host():{host:Host, port:Int} { + if (localAddress == null) { + localAddress = readLocalAddress(); + } + return localAddress; + } + + override public function setTimeout(timeout:Float):Void { + this.timeout = timeout; + } + + override public function waitForRead():Void { + select([this], null, null, timeout); + } + + override public function setBlocking(b:Bool):Void {} + + override public function setFastSend(b:Bool):Void { + tcp.noDelay(b).resolve(); + } + + public static function select(read:Array, write:Array, others:Array, + ?timeout:Float):{read:Array, write:Array, others:Array} { + var loop = findLoop(read, write, others); + var deadline = timeout == null || timeout < 0 ? -1.0 : haxe.Timer.stamp() + timeout; + while (true) { + var readyRead = ready(read); + var readyWrite = writable(write); + var readyOthers:Array = []; + var nativeRead = nativeSockets(read); + var nativeWrite = nativeSockets(write); + var nativeOthers = nativeSockets(others); + if (nativeRead.length > 0 || nativeWrite.length > 0 || nativeOthers.length > 0) { + var selected = sys.net.Socket.select(nativeRead, nativeWrite, nativeOthers, 0); + readyRead = readyRead.concat(selected.read); + readyWrite = readyWrite.concat(selected.write); + readyOthers = selected.others; + } + if (readyRead.length > 0 || readyWrite.length > 0 || readyOthers.length > 0) { + return {read: readyRead, write: readyWrite, others: readyOthers}; + } + if (timeout == 0 || (deadline >= 0 && haxe.Timer.stamp() >= deadline)) { + return {read: [], write: [], others: []}; + } + if (loop != null) { + loop.run(RunMode.NOWAIT); + } + Sys.sleep(POLL_INTERVAL); + } + } + + private static function ready(sockets:Array):Array { + if (sockets == null) { + return []; + } + return sockets.filter(socket -> isLuvSocket(socket) && luv(socket).isReadyToRead()); + } + + private static function writable(sockets:Array):Array { + if (sockets == null) { + return []; + } + return sockets.filter(socket -> isLuvSocket(socket) && !luv(socket).listening && !luv(socket).closed); + } + + private static function nativeSockets(sockets:Array):Array { + if (sockets == null) { + return []; + } + return sockets.filter(socket -> !isLuvSocket(socket)); + } + + private static function findLoop(read:Array, write:Array, others:Array):Loop { + for (group in [read, write, others]) { + if (group != null && group.length > 0) { + for (socket in group) { + if (isLuvSocket(socket)) { + return luv(socket).loop; + } + } + } + } + return null; + } + + private static function isLuvSocket(socket:sys.net.Socket):Bool { + return Std.isOfType(socket, Socket); + } + + private static function luv(socket:sys.net.Socket):Socket { + return cast socket; + } + + private function isReadyToRead():Bool { + pump(); + return listening ? pending.length > 0 : available() > 0 || readClosed; + } + + private function startRead():Void { + if (reading) { + return; + } + reading = true; + Stream.readStart(tcp, result -> { + switch (result) { + case Ok(buffer): + if (buffer.size() > 0) { + readBuffer.add(buffer.toBytes()); + } + case Error(_): + readClosed = true; + } + }); + } + + public function readBytesInto(bytes:Bytes, pos:Int, len:Int):Int { + waitUntil(() -> available() > 0 || readClosed, timeout); + var availableBytes = available(); + if (availableBytes == 0) { + throw new Eof(); + } + var count = len < availableBytes ? len : availableBytes; + bytes.blit(pos, readBuffer.getBytes(), readOffset, count); + readOffset += count; + return count; + } + + public function writeBytesFrom(bytes:Bytes, pos:Int, len:Int):Int { + var chunk = bytes.sub(pos, len); + var done = false; + var failure:Dynamic = null; + var queued = Stream.write(tcp, [chunk], (result, _) -> { + switch (result) { + case Error(e): + failure = e; + case Ok(_) | null: + } + done = true; + }); + switch (queued) { + case Error(e): + throw e; + case Ok(_) | null: + } + waitUntil(() -> done, timeout); + if (failure != null) { + throw failure; + } + return len; + } + + private function waitUntil(condition:() -> Bool, timeout:Null):Void { + var deadline = timeout == null || timeout < 0 ? -1.0 : haxe.Timer.stamp() + timeout; + while (!condition()) { + if (deadline >= 0 && haxe.Timer.stamp() >= deadline) { + return; + } + pump(); + Sys.sleep(POLL_INTERVAL); + } + } + + private function pump():Void { + if (!closed) { + loop.run(RunMode.NOWAIT); + } + } + + private function available():Int { + return readBuffer.length - readOffset; + } + + private function readLocalAddress():{host:Host, port:Int} { + return socketAddress(tcp.getSockName().resolve()); + } + + private function readPeerAddress():{host:Host, port:Int} { + return socketAddress(tcp.getPeerName().resolve()); + } + + private static function sockAddr(host:Host, port:Int):SockAddr { + var address:String = host.toString(); + if (address.indexOf(":") == -1) { + var ipv4:String = address; + return SockAddr.ipv4(ipv4, port).resolve(); + } + var ipv6:String = address; + return SockAddr.ipv6(ipv6, port).resolve(); + } + + private static function socketAddress(address:SockAddr):{host:Host, port:Int} { + var text = address.toString(); + var colon = text.lastIndexOf(":"); + var host = colon == -1 ? text : text.substr(0, colon); + var port = address.port == null ? 0 : address.port; + return {host: new Host(host), port: port}; + } +} + +private class LuvSocketInput extends Input { + private final socket:Socket; + + public function new(socket:Socket) { + this.socket = socket; + } + + override public function readBytes(buf:Bytes, pos:Int, len:Int):Int { + return socket.readBytesInto(buf, pos, len); + } + + override public function close():Void { + socket.close(); + } +} + +private class LuvSocketOutput extends Output { + private final socket:Socket; + + public function new(socket:Socket) { + this.socket = socket; + } + + override public function writeBytes(buf:Bytes, pos:Int, len:Int):Int { + return socket.writeBytesFrom(buf, pos, len); + } + + override public function writeByte(c:Int):Void { + var bytes = Bytes.alloc(1); + bytes.set(0, c); + socket.writeBytesFrom(bytes, 0, 1); + } + + override public function close():Void { + socket.close(); + } +} +#else +typedef Socket = sys.net.Socket; +#end diff --git a/src/snake/http/BaseHTTPRequestHandler.hx b/src/snake/http/BaseHTTPRequestHandler.hx index 3787b05..b6d4f1b 100644 --- a/src/snake/http/BaseHTTPRequestHandler.hx +++ b/src/snake/http/BaseHTTPRequestHandler.hx @@ -9,6 +9,7 @@ import snake.socket.BaseServer; import snake.socket.StreamRequestHandler; import sys.io.File; import sys.net.Host; +import snake._internal.net.Socket as InternalSocket; import sys.net.Socket; class BaseHTTPRequestHandler extends StreamRequestHandler { @@ -231,7 +232,7 @@ class BaseHTTPRequestHandler extends StreamRequestHandler { **/ private function handleOneRequest():Void { try { - var selected = Socket.select([connection], null, null, 5); + var selected = InternalSocket.select([connection], null, null, 5); if (selected.read.length == 0) { closeConnection = true; return; diff --git a/src/snake/socket/BaseServer.hx b/src/snake/socket/BaseServer.hx index d522438..d0ee6ae 100644 --- a/src/snake/socket/BaseServer.hx +++ b/src/snake/socket/BaseServer.hx @@ -2,7 +2,7 @@ package snake.socket; import haxe.Exception; import sys.net.Host; -import sys.net.Socket; +import snake._internal.net.Socket; import sys.thread.Mutex; import sys.thread.Thread; diff --git a/src/snake/socket/TCPServer.hx b/src/snake/socket/TCPServer.hx index 884e06c..46af113 100644 --- a/src/snake/socket/TCPServer.hx +++ b/src/snake/socket/TCPServer.hx @@ -2,7 +2,7 @@ package snake.socket; import haxe.Exception; import sys.net.Host; -import sys.net.Socket; +import snake._internal.net.Socket; /** Base class for various socket-based server classes. diff --git a/test-socket.hxml b/test-socket.hxml new file mode 100644 index 0000000..70a8f1c --- /dev/null +++ b/test-socket.hxml @@ -0,0 +1,4 @@ +--main TestSocket +-cp src +-cp tests +--interp diff --git a/tests/TestSocket.hx b/tests/TestSocket.hx new file mode 100644 index 0000000..f33dce6 --- /dev/null +++ b/tests/TestSocket.hx @@ -0,0 +1,139 @@ +import snake.http.BaseHTTPRequestHandler; +import snake._internal.net.Socket as InternalSocket; +import snake.socket.BaseRequestHandler; +import snake.socket.BaseServer; +import sys.net.Host; +import sys.net.Socket; + +class TestSocket { + public static function main():Void { + #if sys + testInternalSocketIsSysSocket(); + testSocketBehavior(); + testImmediateRebindAfterConnection(); + testRealSocketHandler(); + #else + assert(true, "non-sys target"); + #end + } + + private static function testInternalSocketIsSysSocket():Void { + var socket:Socket = new InternalSocket(); + socket.close(); + } + + private static function testSocketBehavior():Void { + var server = new InternalSocket(); + var host = new Host("127.0.0.1"); + server.bind(host, 0); + server.listen(1); + var port = server.host().port; + assert(port > 0, "bound port"); + + var client = new InternalSocket(); + client.connect(host, port); + assert(client.input != null, "client input"); + assert(client.output != null, "client output"); + + var selected = InternalSocket.select([server], [server], [server], 0.01); + equals(1, selected.read.length, "server select read"); + equals(0, selected.write.length, "server select write"); + equals(0, selected.others.length, "server select others"); + + selected = InternalSocket.select([server], [server], [server], 0.01); + equals(1, selected.read.length, "server select read repeat"); + equals(0, selected.write.length, "server select write repeat"); + equals(0, selected.others.length, "server select others repeat"); + + var accepted = server.accept(); + assert(accepted != null, "accepted socket"); + assert(accepted.input != null, "accepted input"); + assert(accepted.output != null, "accepted output"); + accepted.setFastSend(true); + server.setBlocking(false); + + selected = InternalSocket.select([server], [server], [server], 0.01); + equals(0, selected.read.length, "server select read after accept"); + equals(0, selected.write.length, "server select write after accept"); + equals(0, selected.others.length, "server select others after accept"); + + accepted.output.writeByte(97); + accepted.output.writeByte(98); + accepted.output.writeByte(99); + accepted.close(); + + client.waitForRead(); + selected = InternalSocket.select([client], [client], [client]); + equals(1, selected.read.length, "client select read"); + equals(1, selected.write.length, "client select write"); + equals(0, selected.others.length, "client select others"); + equals("abc", client.read(), "client read"); + + client.close(); + server.close(); + } + + private static function testImmediateRebindAfterConnection():Void { + var host = new Host("127.0.0.1"); + var first = new InternalSocket(); + first.bind(host, 0); + first.listen(1); + var port = first.host().port; + + var client = new InternalSocket(); + client.connect(host, port); + var accepted = first.accept(); + client.close(); + accepted.close(); + first.close(); + + var second = new InternalSocket(); + second.bind(host, port); + second.listen(1); + second.close(); + } + + private static function testRealSocketHandler():Void { + var host = new Host("127.0.0.1"); + var server = new Socket(); + server.bind(host, 0); + server.listen(1); + + var client = new Socket(); + client.connect(host, server.host().port); + client.output.writeString("GET / HTTP/1.0\r\n\r\n"); + client.output.flush(); + + var accepted = server.accept(); + new RealSocketHTTPHandler(accepted, {host: host, port: 0}, null); + + client.close(); + server.close(); + } + + private static function assert(condition:Bool, label:String):Void { + if (!condition) { + throw label; + } + } + + private static function equals(expected:T, actual:T, label:String):Void { + if (actual != expected) { + throw '$label: expected $expected, got $actual'; + } + } +} + +private class SysSocketHandler extends BaseRequestHandler { + public function new(request:Socket, clientAddress:{host:Host, port:Int}, server:BaseServer) { + super(request, clientAddress, server); + } +} + +private class RealSocketHTTPHandler extends BaseHTTPRequestHandler { + public function new(request:Socket, clientAddress:{host:Host, port:Int}, server:BaseServer) { + super(request, clientAddress, server); + } + + override private function logMessage(message:String):Void {} +}