import t, c from stdint import * import string import platmacro import stdio import w32.winsock2 # ============================================================ # socket.py — 跨平台 Socket 封装 # Python 风格 OOP API + 兼容旧版过程式函数 # # 用法示例: # 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() # # 仅 3 处平台差异(通过 platmacro.IS_WINDOWS 分支): # 1. WSAStartup() / WSACleanup() # 2. closesocket() vs close() # 3. WSAGetLastError() vs errno # ============================================================ # ============================================================ # 常量 # ============================================================ 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 # ============================================================ # 结构体 # ============================================================ class SocketAddr: """sockaddr_in — IPv4 地址(16 字节,POSIX 布局)""" family: u16 port: u16 addr: u32 zero: u64 class HostEnt: """struct hostent — DNS 解析结果""" h_name: CHARPTR h_aliases: CHARPTR h_addrtype: INT h_length: INT h_addr_list: CHARPTR class timeval: """struct timeval — 超时设置""" tv_sec: INT tv_usec: INT # ============================================================ # 外部 C 函数 — 使用 u64 句柄,跨平台兼容 # (Windows SOCKET=u64, POSIX fd=int → u64 零扩展兼容) # ============================================================ 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 # ============================================================ # 内部辅助函数 # ============================================================ def _CloseFd(fd: INT) -> INT: """平台差异:Windows 用 closesocket,POSIX 用 close。""" if platmacro.IS_WINDOWS: return w32.winsock2.closesocket(u64(fd)) return close(fd) def _MakeAddr(host: str, port: INT) -> SocketAddr: """构建 SocketAddr。先尝试点分十进制 IP,失败则 DNS 解析。""" 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 解析 — gethostbyname 返回 struct hostent* # h_addr_list 是 char** → h_addr_list[0] 是 char* → *(uint32_t*)h_addr_list[0] 是 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: """解析点分十进制 IPv4 地址。""" 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 类 — Python 风格 OOP API # ============================================================ class Socket: """跨平台 Socket 封装,贴近 Python socket.socket 用法。 用法: 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): """创建 Socket。 Args: family: 地址族 (AF_INET) type: 套接字类型 (SOCK_STREAM) protocol: 协议 (默认 0,自动匹配) """ 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: """连接到远程主机。 Args: host: 主机名或 IP port: 端口号 Returns: 0 成功,非 0 失败 """ 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: """发送数据。 Returns: 已发送字节数,SOCKET_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: """接收数据。 Returns: 已接收字节数,SOCKET_ERROR 失败,0 表示连接关闭 """ 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: """发送全部数据(循环发送直到完成或出错)。 Returns: 0 成功,SOCKET_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: """关闭 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: """绑定地址和端口(服务器端)。 Returns: 0 成功,非 0 失败 """ 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: """开始监听连接(服务器端)。 Returns: 0 成功,非 0 失败 """ 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: """接受一个客户端连接。 Args: out_client: 输出参数,接收客户端 Socket Returns: 客户端 fd(>= 0 成功),SOCKET_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: """设置接收和发送超时(毫秒)。0 表示无超时。 Returns: 0 成功,非 0 失败 """ 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: """返回底层文件描述符 / SOCKET 句柄。""" return self.fd # ============================================================ # 平台初始化 # ============================================================ def SocketInit() -> INT: """初始化 Socket 库。Windows 下调用 WSAStartup,POSIX 无操作。""" 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: """清理 Socket 库。""" if platmacro.IS_WINDOWS: return w32.winsock2.WSACleanup() return 0 # ============================================================ # 兼容旧版过程式 API(委托到 Socket 方法) # ============================================================ def SocketCreate(family: INT, type: INT, protocol: INT, out_sock: Socket | t.CPtr) -> INT: """创建 Socket 句柄(兼容旧版 API)。建议使用 Socket(...) 构造。""" 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 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: """检查 Socket 是否有效。""" if sock is None: return False return sock.fd != INVALID_SOCKET