451 lines
14 KiB
Python
451 lines
14 KiB
Python
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 |