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