Files
TransPyC/includes/socket.py

454 lines
14 KiB
Python

import t, c
from stdint import *
import string
import platmacro
import stdio
import w32.winsock2
# ============================================================
# socket.py — Cross-platform wrapper for Socket
# Python-style OOP API + Compatibility with legacy procedural functions
#
# Example:
# sock = socket.Socket(socket.AF_INET, socket.SOCK_STREAM)
# sock.connect("example.com", 80)
# sock.send(data, length)
# n = sock.recv(buf, bufsize)
# sock.close()
#
# Only 3 platform differences (via platmacro.IS_WINDOWS branch):
# 1. WSAStartup() / WSACleanup()
# 2. closesocket() vs close()
# 3. WSAGetLastError() vs errno
# ============================================================
# ============================================================
# Constants
# ============================================================
AF_INET: t.CDefine = 2
AF_INET6: t.CDefine = 10
SOCK_STREAM: t.CDefine = 1
SOCK_DGRAM: t.CDefine = 2
SOCK_RAW: t.CDefine = 3
IPPROTO_TCP: t.CDefine = 6
IPPROTO_UDP: t.CDefine = 17
SOL_SOCKET: t.CDefine = 1
SO_RCVTIMEO: t.CDefine = 20
SO_SNDTIMEO: t.CDefine = 21
SO_REUSEADDR: t.CDefine = 2
INADDR_ANY: t.CDefine = 0
SOCKET_ERROR: t.CDefine = -1
INVALID_SOCKET: t.CDefine = -1
SOCKET_DEFAULT_TIMEOUT: t.CDefine = 0
# ============================================================
# Structures
# ============================================================
class SocketAddr:
"""sockaddr_in — IPv4 address (16 bytes layout)"""
family: u16
port: u16
addr: u32
zero: u64
class HostEnt:
"""struct hostent — DNS resolution result"""
h_name: CHARPTR
h_aliases: CHARPTR
h_addrtype: INT
h_length: INT
h_addr_list: CHARPTR
class timeval:
"""struct timeval — Timeout setting"""
tv_sec: INT
tv_usec: INT
# ============================================================
# External C Functions — Using u64 Handle, Cross-platform Compatibility
# (Windows SOCKET=u64, POSIX fd=int → u64 Zero-Extension Compatible)
# ============================================================
def socket(family: INT, type: INT, protocol: INT) -> u64 | t.CExtern | t.CExport: pass
def connect(fd: u64, addr: SocketAddr | t.CPtr, addrlen: INT) -> INT | t.CExtern | t.CExport: pass
def send(fd: u64, buf: t.CVoid | t.CPtr, len: INT, flags: INT) -> INT | t.CExtern | t.CExport: pass
def recv(fd: u64, buf: t.CVoid | t.CPtr, len: INT, flags: INT) -> INT | t.CExtern | t.CExport: pass
def bind(fd: u64, addr: SocketAddr | t.CPtr, addrlen: INT) -> INT | t.CExtern | t.CExport: pass
def listen(fd: u64, backlog: INT) -> INT | t.CExtern | t.CExport: pass
def accept(fd: u64, addr: SocketAddr | t.CPtr, addrlen: INT | t.CPtr) -> u64 | t.CExtern | t.CExport: pass
def setsockopt(fd: u64, level: INT, optname: INT, optval: t.CVoid | t.CPtr, optlen: INT) -> INT | t.CExtern | t.CExport: pass
def gethostbyname(name: t.CChar | t.CConst | t.CPtr) -> HostEnt | t.CPtr | t.CExtern | t.CExport: pass
def ntohs(netshort: u16) -> u16 | t.CExtern | t.CExport: pass
def htons(hostshort: u16) -> u16 | t.CExtern | t.CExport: pass
def close(fd: u64) -> INT | t.CExtern | t.CExport: pass
# ============================================================
# Internal Auxiliary Functions
# ============================================================
def _CloseFd(fd: INT) -> INT:
"""Platform differences: Windows uses closesocket, POSIX uses close。"""
if platmacro.IS_WINDOWS:
return w32.winsock2.closesocket(u64(fd))
return close(fd)
def _MakeAddr(host: str, port: INT) -> SocketAddr:
"""Build SocketAddr. First try decimal IP, then DNS resolution.
If decimal IP fails, try DNS resolution.
If DNS resolution fails, return INADDR_ANY."""
addr: SocketAddr
string.memset(c.Addr(addr), 0, 16)
addr.family = u16(AF_INET)
addr.port = htons(u16(port))
ip: u32 = 0
if _ParseIPv4(host, c.Addr(ip)):
addr.addr = ip
return addr
# DNS resolution — gethostbyname returns struct hostent*
# h_addr_list is char** → h_addr_list[0] is char* → *(uint32_t*)h_addr_list[0] is IP
if platmacro.IS_WINDOWS:
he_raw: w32.winsock2.WSAHostEnt | t.CPtr = w32.winsock2.gethostbyname(host)
if he_raw is not None:
# he_raw.h_addr_list → he_raw->h_addr_list (char**)
# c.Deref(t.CUInt64T(he_raw.h_addr_list, t.CPtr)) → *(uint64_t*)(he_raw->h_addr_list) → h_addr_list[0]
ip_ptr: u64 = u64(c.Deref(t.CUInt64T(he_raw.h_addr_list, t.CPtr)))
if ip_ptr != 0:
addr.addr = u32(c.Deref(t.CUInt32T(ip_ptr, t.CPtr)))
else:
he: HostEnt | t.CPtr = gethostbyname(host)
if he is not None:
ip_ptr: u64 = u64(c.Deref(t.CUInt64T(he.h_addr_list, t.CPtr)))
if ip_ptr != 0:
addr.addr = u32(c.Deref(t.CUInt32T(ip_ptr, t.CPtr)))
return addr
def _ParseIPv4(ip_str: str, out_addr: u32 | t.CPtr) -> bool:
"""Parse IPv4 address."""
if ip_str is None or out_addr is None: return False
parts: t.CArray[INT, 4]
part_count: INT = 0
current: INT = 0
started: bool = False
for ch in ip_str:
if ch >= '0' and ch <= '9':
if not started: started = True
current = current * 10 + INT(ch - '0')
if current > 255: return False
elif ch == '.':
if not started: return False
if part_count >= 4: return False
parts[part_count] = current
part_count += 1
current = 0
started = False
else:
break
if started and part_count < 4:
parts[part_count] = current
part_count += 1
if part_count != 4: return False
result: u32 = 0
result = result | (u32(parts[0]) << 24)
result = result | (u32(parts[1]) << 16)
result = result | (u32(parts[2]) << 8)
result = result | u32(parts[3])
c.DerefAs(out_addr, result)
return True
# ============================================================
# Socket class — Python OOP API
# ============================================================
class Socket:
"""Cross-platform Socket class, similar to Python socket.socket.
Usage:
sock = Socket(AF_INET, SOCK_STREAM)
sock.connect("example.com", 80)
sock.send(data, length)
n = sock.recv(buf, bufsize)
sock.close()
"""
fd: INT
family: INT
type: INT
protocol: INT
def __init__(self, family: INT, type: INT, protocol: INT = 0):
"""Create Socket.
Args:
family: Address family (AF_INET)
type: Socket type (SOCK_STREAM)
protocol: Protocol (default 0, auto-match)
"""
self.family = family
self.type = type
self.protocol = protocol
if platmacro.IS_WINDOWS:
self.fd = INT(w32.winsock2.socket(family, type, protocol))
else:
self.fd = socket(family, type, protocol)
def connect(self, host: str, port: INT) -> INT:
"""Connect to remote host.
Args:
host: Hostname or IP
port: Port number
Returns:
0 success, non 0 otherwise
"""
if self.fd == INVALID_SOCKET: return -1
addr: SocketAddr = _MakeAddr(host, port)
result: INT = 0
if platmacro.IS_WINDOWS:
result = w32.winsock2.connect(u64(self.fd), c.Addr(addr), 16)
if result != 0:
err: INT = w32.winsock2.WSAGetLastError()
else:
result = connect(self.fd, c.Addr(addr), 16)
return result
def send(self, data: t.CVoid | t.CPtr, length: INT) -> INT:
"""Send data.
Returns:
Sent bytes, SOCKET_ERROR on error
"""
if self.fd == INVALID_SOCKET: return SOCKET_ERROR
if data is None or length <= 0: return 0
if platmacro.IS_WINDOWS:
return w32.winsock2.send(u64(self.fd), data, length, 0)
return send(self.fd, data, length, 0)
def recv(self, buf: t.CVoid | t.CPtr, bufsize: INT) -> INT:
"""Receive data.
Returns:
Received bytes, SOCKET_ERROR on error, 0 on connection close
"""
if self.fd == INVALID_SOCKET: return SOCKET_ERROR
if buf is None or bufsize <= 0: return 0
if platmacro.IS_WINDOWS:
return w32.winsock2.recv(u64(self.fd), buf, bufsize, 0)
return recv(self.fd, buf, bufsize, 0)
def sendall(self, data: t.CVoid | t.CPtr, length: INT) -> INT:
"""Send all data.
Returns:
0 success, SOCKET_ERROR on error
"""
if self.fd == INVALID_SOCKET: return SOCKET_ERROR
if data is None or length <= 0: return 0
total: INT = 0
while total < length:
n: INT = self.send(t.CVoid(t.CUInt64T(data) + total, t.CPtr), length - total)
if n <= 0: return SOCKET_ERROR
total += n
return 0
def close(self) -> INT:
"""Close Socket."""
if self.fd == INVALID_SOCKET: return 0
result: INT = _CloseFd(self.fd)
self.fd = INVALID_SOCKET
return result
def bind(self, host: str, port: INT) -> INT:
"""Bind address and port (server-side).
Returns:
0 success, non 0 otherwise
"""
if self.fd == INVALID_SOCKET: return -1
addr: SocketAddr = _MakeAddr(host, port)
if platmacro.IS_WINDOWS:
return w32.winsock2.bind(u64(self.fd), c.Addr(addr), 16)
return bind(self.fd, c.Addr(addr), 16)
def listen(self, backlog: INT) -> INT:
"""Start for incoming connections (server-side).
Returns:
0 success, non 0 otherwise
"""
if self.fd == INVALID_SOCKET: return -1
if platmacro.IS_WINDOWS:
return w32.winsock2.listen(u64(self.fd), backlog)
return listen(self.fd, backlog)
def accept(self, out_client: 'Socket' | t.CPtr) -> INT:
"""Accept a client connection.
Args:
out_client: Output parameter, receives client Socket
Returns:
Client fd (>= 0 on success), SOCKET_ERROR on error
"""
if self.fd == INVALID_SOCKET: return SOCKET_ERROR
if out_client is None: return SOCKET_ERROR
addr: SocketAddr
string.memset(c.Addr(addr), 0, 16)
addrlen: INT = 16
if platmacro.IS_WINDOWS:
client_fd: u64 = w32.winsock2.accept(u64(self.fd), c.Addr(addr), c.Addr(addrlen))
if client_fd == w32.winsock2.INVALID_SOCKET:
return SOCKET_ERROR
out_client.fd = INT(client_fd)
out_client.family = self.family
out_client.type = self.type
out_client.protocol = self.protocol
return INT(client_fd)
client_fd: INT = accept(self.fd, c.Addr(addr), c.Addr(addrlen))
if client_fd == SOCKET_ERROR:
return SOCKET_ERROR
out_client.fd = client_fd
out_client.family = self.family
out_client.type = self.type
out_client.protocol = self.protocol
return client_fd
def settimeout(self, timeout_ms: INT) -> INT:
"""Set timeout for receive and send operations (milliseconds).
Returns:
0 success, non 0 otherwise
"""
if self.fd == INVALID_SOCKET: return -1
tv: timeval
tv.tv_sec = timeout_ms / 1000
tv.tv_usec = (timeout_ms % 1000) * 1000
optlen: INT = 8
if platmacro.IS_WINDOWS:
r1: INT = w32.winsock2.setsockopt(u64(self.fd), SOL_SOCKET, SO_RCVTIMEO, c.Addr(tv), optlen)
r2: INT = w32.winsock2.setsockopt(u64(self.fd), SOL_SOCKET, SO_SNDTIMEO, c.Addr(tv), optlen)
else:
r1 = setsockopt(self.fd, SOL_SOCKET, SO_RCVTIMEO, c.Addr(tv), optlen)
r2 = setsockopt(self.fd, SOL_SOCKET, SO_SNDTIMEO, c.Addr(tv), optlen)
if r1 != 0: return r1
return r2
def fileno(self) -> INT:
"""Return underlying file descriptor / SOCKET handle."""
return self.fd
# ============================================================
# Platform initialization
# ============================================================
def SocketInit() -> INT:
"""Initialize the Socket library. On Windows, call WSAStartup; on POSIX, do nothing."""
if platmacro.IS_WINDOWS:
wsa_data: t.CArray[BYTE, 400]
string.memset(c.Addr(wsa_data), 0, 400)
return w32.winsock2.WSAStartup(w32.winsock2.WINSOCK_VERSION, c.Addr(wsa_data))
return 0
def SocketCleanup() -> INT:
"""Cleanup the Socket library."""
if platmacro.IS_WINDOWS:
return w32.winsock2.WSACleanup()
return 0
# ============================================================
# Compatibility with old procedural API (delegates to Socket methods)
# ============================================================
def SocketCreate(family: INT, type: INT, protocol: INT, out_sock: Socket | t.CPtr) -> INT:
"""Create a Socket handle (compatible with old API). It is recommended to use Socket(...) constructor."""
if out_sock is None: return -1
out_sock.family = family
out_sock.type = type
out_sock.protocol = protocol
if platmacro.IS_WINDOWS:
fd: u64 = w32.winsock2.socket(family, type, protocol)
out_sock.fd = INT(fd)
if fd == w32.winsock2.INVALID_SOCKET: return -1
else:
fd: INT = socket(family, type, protocol)
out_sock.fd = fd
if fd == INVALID_SOCKET: return -1
return 0
# But maybe using OOP would be better, though that can be wrapped up later.
def SocketConnect(sock: Socket | t.CPtr, host: str, port: INT) -> INT:
return sock.connect(host, port)
def SocketSend(sock: Socket | t.CPtr, data: t.CVoid | t.CPtr, length: INT) -> INT:
return sock.send(data, length)
def SocketRecv(sock: Socket | t.CPtr, buf: t.CVoid | t.CPtr, bufsize: INT) -> INT:
return sock.recv(buf, bufsize)
def SocketClose(sock: Socket | t.CPtr) -> INT:
return sock.close()
def SocketBind(sock: Socket | t.CPtr, host: str, port: INT) -> INT:
return sock.bind(host, port)
def SocketListen(sock: Socket | t.CPtr, backlog: INT) -> INT:
return sock.listen(backlog)
def SocketAccept(sock: Socket | t.CPtr, out_client: Socket | t.CPtr) -> INT:
return sock.accept(out_client)
def SocketSetTimeout(sock: Socket | t.CPtr, timeout_ms: INT) -> INT:
return sock.settimeout(timeout_ms)
def SocketIsValid(sock: Socket | t.CPtr) -> bool:
"""Check if a Socket is valid."""
if sock is None: return False
return sock.fd != INVALID_SOCKET