Files
TransPyC/includes/socket.py
2026-07-18 19:25:40 +08:00

451 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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 用 closesocketPOSIX 用 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 下调用 WSAStartupPOSIX 无操作。"""
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