74a68fcbd0
背景:
Socks5Server.running 是乐观标记, start() 之后到 stop() 之前一直 True。
但子线程 event loop 崩了/端口被外部抢占/worker 被 gunicorn 杀 等情况,
self.running 仍为 True, DB 里 inst.running=true 是谎言。
上次 40080 实例莫名停止就是这种状态。
变更:
+ Socks5Server.is_alive() thread.is_alive() + server.is_serving()
M InstanceManager.sync_instances 启动前先探活, 死的从 _instances 摘掉
让'启动缺失'循环用 DB 配置重建它
M inst.running = srv.is_alive() 不再用乐观的 srv.running
测试:
+ tests/e2e_socks5_lifecycle.py 完整 SOCKS5 userpass 生命周期, 拦住
_record_stats 里 await 同步方法的回归
+ tests/e2e_health_check.py 模拟 Socks5Server 死了, sync_instances
自动重启; 反向验证: 注释掉健康检查就 fail
验证:
- 正向 (修复在): lifecycle PASS, health PASS, smoke tunnel PASS
- 反向 (回滚修复): 两个 e2e 都正确 FAIL, 证明测试真的能抓回归
186 lines
7.4 KiB
Python
186 lines
7.4 KiB
Python
"""端到端 SOCKS5 生命周期测试。
|
|
|
|
之前 _record_stats 里 await 了同步方法 user_service.add_traffic(),
|
|
被 FakeUserService.add_traffic (误写为 async def) 掩盖, 没在烟雾测试里发现,
|
|
结果生产机第一个真实连接断开时抛 TypeError, worker 卡死, 40080 端口失活。
|
|
|
|
本测试用真实的 services.user_service.UserService (不要 await 同步方法),
|
|
跑完整 SOCKS5 生命周期 (含 userpass 认证 + add_traffic 调用),
|
|
任何 await NoneType / 协程签名错配都会在这里被抓住。
|
|
捕获一个 StringIO 形式的 stderr, 检查 'NoneType' / '记录统计失败'
|
|
这类错误日志反向证明 add_traffic 路径没踩 await 协程错配。
|
|
|
|
跑法: cd /home/cnbugs/socks-manager && python3 tests/e2e_socks5_lifecycle.py
|
|
|
|
注意: 用临时 sqlite 文件, 不污染开发 DB。
|
|
"""
|
|
import asyncio
|
|
import io
|
|
import logging
|
|
import os
|
|
import socket
|
|
import struct
|
|
import sys
|
|
import tempfile
|
|
import time
|
|
|
|
sys.path.insert(0, os.getcwd())
|
|
|
|
# 用临时 DB 跑, 避免污染项目根目录的 socks_manager.db
|
|
_tmp_db = tempfile.NamedTemporaryFile(suffix=".db", delete=False, dir="/tmp")
|
|
_tmp_db.close()
|
|
os.environ["SM_DB_URI"] = f"sqlite:///{_tmp_db.name}"
|
|
os.environ["SM_ADMIN_PASSWORD"] = "test123"
|
|
os.environ["SM_SECRET_KEY"] = "test-secret-key-for-e2e-only"
|
|
os.environ["SM_LOG_LEVEL"] = "WARNING"
|
|
|
|
from app import create_app
|
|
from engine.instances import Socks5Server, InstanceConfig
|
|
from services.user_service import UserService
|
|
|
|
|
|
def find_free_port():
|
|
s = socket.socket()
|
|
s.bind(("127.0.0.1", 0))
|
|
p = s.getsockname()[1]
|
|
s.close()
|
|
return p
|
|
|
|
|
|
async def silent_listener(port):
|
|
async def cb(r, w):
|
|
await asyncio.sleep(3600)
|
|
return await asyncio.start_server(cb, "127.0.0.1", port)
|
|
|
|
|
|
def build_userpass_request(username: str, password: str) -> bytes:
|
|
"""RFC 1929: VER(1) ULEN(1) UNAME PLEN(1) PASSWD"""
|
|
ub = username.encode()
|
|
pb = password.encode()
|
|
assert 0 < len(ub) <= 255
|
|
assert 0 < len(pb) <= 255
|
|
return struct.pack("!BB", 1, len(ub)) + ub + struct.pack("!B", len(pb)) + pb
|
|
|
|
|
|
async def client_handshake_userpass(host, socks_port, target_port, username, password, send_data=b"hello"):
|
|
"""完整生命周期: 握手(userpass) + CONNECT + 发数据 + 断开"""
|
|
r, w = await asyncio.open_connection(host, socks_port)
|
|
# 1. method negotiation: 提出 userpass (0x02)
|
|
w.write(struct.pack("!BB", 5, 1) + bytes([0x02]))
|
|
await w.drain()
|
|
sel = await r.readexactly(2)
|
|
assert sel == bytes([5, 0x02]), f"method sel={sel!r}, want userpass"
|
|
# 2. userpass 认证
|
|
w.write(build_userpass_request(username, password))
|
|
await w.drain()
|
|
auth_resp = await r.readexactly(2)
|
|
assert auth_resp == bytes([1, 0]), f"auth resp={auth_resp!r}, want success"
|
|
# 3. CONNECT
|
|
req = struct.pack("!BBBB", 5, 1, 0, 1) + socket.inet_aton("127.0.0.1") + struct.pack("!H", target_port)
|
|
w.write(req)
|
|
await w.drain()
|
|
resp_head = await r.readexactly(4)
|
|
rep = resp_head[1]
|
|
atyp = resp_head[3]
|
|
assert rep == 0, f"REP={rep} expected SUCCEEDED"
|
|
extra_len = {1: 4, 3: 1, 4: 16}[atyp]
|
|
await r.readexactly(extra_len + 2)
|
|
# 4. 发数据 (走 tunnel)
|
|
w.write(send_data)
|
|
await w.drain()
|
|
# 5. 主动关 writer
|
|
w.close()
|
|
try:
|
|
await w.wait_closed()
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
async def run_test(timeout_sec: int):
|
|
socks_port = find_free_port()
|
|
target_port = find_free_port()
|
|
target_srv = await silent_listener(target_port)
|
|
flask_app = create_app()
|
|
user_service = flask_app.user_service # type: ignore[attr-defined]
|
|
|
|
# 在临时 DB 里建一个真用户, 这样 ConnectionHandler._record_stats 里
|
|
# add_traffic(self.username, ...) 这一行会被真正走到。
|
|
user_service.create_user("e2euser", "e2epass", max_concurrent=100,
|
|
enabled=True, banned=False)
|
|
|
|
# 接管 socks.engine logger, 把日志收集到 buffer, 用来检测
|
|
# "记录统计失败" 这类反向信号
|
|
log_buf = io.StringIO()
|
|
log_handler = logging.StreamHandler(log_buf)
|
|
log_handler.setLevel(logging.WARNING)
|
|
engine_log = logging.getLogger("socks.engine")
|
|
engine_log.addHandler(log_handler)
|
|
engine_log.setLevel(logging.WARNING)
|
|
|
|
cfg = InstanceConfig(
|
|
instance_id=999, name="e2e", listen_host="127.0.0.1", listen_port=socks_port,
|
|
timeout=timeout_sec, auth_method="userpass",
|
|
bandwidth_down=0, bandwidth_up=0, max_concurrent=100,
|
|
)
|
|
srv = Socks5Server(cfg, user_service, flask_app=flask_app)
|
|
srv.start()
|
|
for _ in range(50):
|
|
if srv.server: break
|
|
await asyncio.sleep(0.05)
|
|
assert srv.server, "SOCKS5 server not listening"
|
|
|
|
try:
|
|
# 跑 3 次完整生命周期, 触发 add_traffic 多次
|
|
for i in range(3):
|
|
await client_handshake_userpass("127.0.0.1", socks_port, target_port,
|
|
"e2euser", "e2epass",
|
|
send_data=f"msg{i}".encode())
|
|
# 等所有 _record_stats 协程结束 (每个连接断开都会触发, 在子线程
|
|
# event loop 里跑)。这里 sleep 要够, 否则 next assert 看到的是
|
|
# bytes_out=0 (统计还没落库)。srv.stop() 必须在最后。
|
|
await asyncio.sleep(4.0)
|
|
# 断言: 真用户的 bytes_in/out 应该累计了 3 次 msg 的字节数
|
|
u = user_service.get_user("e2euser")
|
|
assert u is not None
|
|
# 双向都过 tunnel: bytes_in 累计的是 dst->client (我们 server 收到的),
|
|
# 但 silent listener 不回数据, 所以实际只有 client->dst 方向有流量。
|
|
# 我们的 msg 是 client 发出来的, 走的是 client->dst, 对应 bytes_out。
|
|
assert u.bytes_out >= 12, f"bytes_out={u.bytes_out} 累计不够 3 条 msg (>=12), " \
|
|
f"add_traffic 路径有 bug 或 _record_stats 没跑"
|
|
# 关键反向断言: 业务路径上不能出现 add_traffic 协程错配特征
|
|
# (停止 server 时协程被 cancel 仍会冒 GeneratorExit 警告, 那是另一码事)
|
|
log_text = log_buf.getvalue()
|
|
bad_signs = ["object NoneType can't be used in 'await'"]
|
|
for bad in bad_signs:
|
|
assert bad not in log_text, (
|
|
f"反向信号: 业务路径日志里出现 '{bad}', 意味着 _record_stats 路径"
|
|
f"有 await 同步方法 bug\n--- 完整日志 ---\n{log_text}"
|
|
)
|
|
print(f"PASS: userpass 完整生命周期 OK, bytes_out={u.bytes_out}, "
|
|
f"add_traffic 调用成功, 无 NoneType await 错误")
|
|
return True
|
|
finally:
|
|
# 关键: 先停 socks server (关掉 listener, 不再 accept 新连接),
|
|
# 等所有现有 ConnectionHandler 的 _tunnel 协程自然结束 (客户端已关,
|
|
# 服务端读 EOF 后会走 _tunnel 的 finally, 然后 _record_stats 写 DB),
|
|
# 最后再 stop target。
|
|
srv.stop()
|
|
# 等 ConnectionHandler 把 _record_stats 跑完
|
|
await asyncio.sleep(0.5)
|
|
target_srv.close()
|
|
try:
|
|
await target_srv.wait_closed()
|
|
except Exception:
|
|
pass
|
|
engine_log.removeHandler(log_handler)
|
|
# 清理临时 DB
|
|
try: os.unlink(_tmp_db.name)
|
|
except Exception: pass
|
|
|
|
|
|
if __name__ == "__main__":
|
|
logging.basicConfig(level=logging.WARNING,
|
|
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s")
|
|
rc = asyncio.run(run_test(timeout_sec=2))
|
|
sys.exit(0 if rc else 1)
|