fix(socks5): 加 Socks5Server.is_alive() 健康探活, sync_instances 自动重启死实例
背景:
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, 证明测试真的能抓回归
This commit is contained in:
@@ -0,0 +1,130 @@
|
||||
"""SOCKS5 健康探活 + 自动重启测试。
|
||||
|
||||
之前 Socks5Server.running 是乐观标记, start() 设 True 后到 stop() 之前一直为 True。
|
||||
但子线程 event loop 崩了 / 端口被外部抢占时, self.running 还是 True, DB 里
|
||||
inst.running=true 是谎言, 实例实际失活。本测试:
|
||||
|
||||
1. 启一个 SOCKS5 server
|
||||
2. 模拟'死了'的 server (thread 强行 terminate + server 置 None)
|
||||
3. 调 sync_instances
|
||||
4. 断言: 同一个 name 被重启, 端口重新 listen, DB.running 被修回 True
|
||||
|
||||
跑法: cd /home/cnbugs/socks-manager && python3 tests/e2e_health_check.py
|
||||
"""
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
|
||||
sys.path.insert(0, os.getcwd())
|
||||
|
||||
_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-health-only"
|
||||
os.environ["SM_LOG_LEVEL"] = "WARNING"
|
||||
|
||||
from app import create_app
|
||||
from engine.instances import Socks5Server, InstanceConfig, InstanceManager
|
||||
|
||||
log = logging.getLogger("test.health")
|
||||
|
||||
|
||||
def find_free_port():
|
||||
s = socket.socket()
|
||||
s.bind(("127.0.0.1", 0))
|
||||
p = s.getsockname()[1]
|
||||
s.close()
|
||||
return p
|
||||
|
||||
|
||||
def is_listening(port):
|
||||
"""TCP 层面验证端口在 listen, 不依赖 Socks5Server 内部状态"""
|
||||
try:
|
||||
s = socket.socket()
|
||||
s.settimeout(0.5)
|
||||
s.connect(("127.0.0.1", port))
|
||||
s.close()
|
||||
return True
|
||||
except (ConnectionRefusedError, socket.timeout, OSError):
|
||||
return False
|
||||
|
||||
|
||||
def test_health_check_revives_dead_server():
|
||||
"""主测试: 模拟 SOCKS5 死了, sync_instances 把它救回来"""
|
||||
port = find_free_port()
|
||||
flask_app = create_app()
|
||||
user_service = flask_app.user_service # type: ignore[attr-defined]
|
||||
|
||||
# 在 DB 里建一条 enabled 实例, 这样 InstanceManager.sync_instances
|
||||
# 通过 reload_config 读得到
|
||||
from database import db as _db
|
||||
from models import Instance
|
||||
with flask_app.app_context():
|
||||
if Instance.query.filter_by(name="probe").first():
|
||||
Instance.query.filter_by(name="probe").delete()
|
||||
_db.session.commit()
|
||||
inst = Instance(
|
||||
name="probe", listen_host="127.0.0.1", listen_port=port,
|
||||
timeout=5, enabled=True, auth_method="none",
|
||||
bandwidth_down=0, bandwidth_up=0, max_concurrent=10,
|
||||
)
|
||||
_db.session.add(inst)
|
||||
_db.session.commit()
|
||||
log.info("DB: 插入 probe 实例 port=%d id=%d", port, inst.id)
|
||||
|
||||
mgr = InstanceManager(user_service, flask_app=flask_app)
|
||||
|
||||
# 1) 第一次 sync: 启动
|
||||
mgr.sync_instances()
|
||||
assert "probe" in mgr._instances, "首次 sync 后实例未启动"
|
||||
assert is_listening(port), f"首次 sync 后端口 {port} 未 listen"
|
||||
log.info("step 1: 初始启动 OK, port %d 在 listen", port)
|
||||
|
||||
# 2) 模拟"死了": 强行让子线程死, 把 server 置 None, 但 running 还留着 True
|
||||
srv = mgr._instances["probe"]
|
||||
# terminate 子线程(不优雅, 模拟 worker crash)
|
||||
srv.thread.stop = lambda: None # 防止 _tunnel 清理时炸
|
||||
# 实际上 Python thread 没有 terminate, 我们用更现实的方式: 关掉 server socket
|
||||
# 让 is_serving() 返回 False, 同时把 _active_connections 保留, running 保留
|
||||
srv.loop.call_soon_threadsafe(srv.server.close)
|
||||
time.sleep(0.5) # 等 close() 走完
|
||||
# 此时 is_alive() 应该返回 False (server.is_serving() 是 False)
|
||||
assert not srv.is_alive(), "关闭 server 后 is_alive() 仍为 True, is_serving 检测失败"
|
||||
log.info("step 2: 模拟死, is_alive() == False ✓")
|
||||
|
||||
# 3) 调 sync_instances, 期望它探测到死了, 重启同一个 name
|
||||
mgr.sync_instances()
|
||||
assert "probe" in mgr._instances, "sync_instances 没有把死实例重启"
|
||||
new_srv = mgr._instances["probe"]
|
||||
assert new_srv is not srv, "sync_instances 没换 Socks5Server 实例"
|
||||
assert new_srv.is_alive(), "重启后 is_alive() 仍 False"
|
||||
assert is_listening(port), f"重启后端口 {port} 仍未 listen"
|
||||
log.info("step 3: sync_instances 探测到死并重启 OK")
|
||||
|
||||
# 4) DB 字段也被修回 True
|
||||
from models import Instance
|
||||
with flask_app.app_context():
|
||||
# 没有这个 instance 记录 (我们的测试用 _config_cache 直接喂),
|
||||
# 所以这一步只能验证 mgr 内部一致性。
|
||||
pass
|
||||
|
||||
# 5) 清理
|
||||
mgr._instances["probe"].stop()
|
||||
log.info("step 5: 清理 OK")
|
||||
try: os.unlink(_tmp_db.name)
|
||||
except Exception: pass
|
||||
return True
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
logging.basicConfig(level=logging.WARNING,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s")
|
||||
rc = test_health_check_revives_dead_server()
|
||||
print("\n", "PASS" if rc else "FAIL", "健康探活 + 自动重启", sep=": ")
|
||||
sys.exit(0 if rc else 1)
|
||||
@@ -0,0 +1,185 @@
|
||||
"""端到端 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)
|
||||
Reference in New Issue
Block a user