Rigel: auth SOCKS5 frontend (atyp=3 length-prefix fix, UTC auth, per-upstream creds), full-creds dashboard, honest provision_proxy
This commit is contained in:
264
proxy_server.py
Normal file
264
proxy_server.py
Normal file
@@ -0,0 +1,264 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Rigel — authenticated SOCKS5 frontend (customer entry point).
|
||||
|
||||
Customers authenticate with the per-user credentials issued when their
|
||||
invoice settles (subscriptions.proxy_user / proxy_pass). We validate LIVE
|
||||
against rigel.db, resolve the subscription's location, and relay to that
|
||||
location's upstream SOCKS5 exit on the LAN.
|
||||
|
||||
Upstreams are derived from app.LOCATIONS so the two cannot drift. Upstreams
|
||||
that themselves need auth (IPRoyal residential/mobile) are supported via
|
||||
UPSTREAM_<LOCATION>_USER / UPSTREAM_<LOCATION>_PASS env vars.
|
||||
|
||||
Implements SOCKS5 RFC1928 + username/password auth RFC1929.
|
||||
"""
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import ipaddress
|
||||
import logging
|
||||
import os
|
||||
import sqlite3
|
||||
import struct
|
||||
from datetime import datetime
|
||||
|
||||
logging.basicConfig(level=logging.INFO,
|
||||
format="%(asctime)s %(levelname)s %(message)s")
|
||||
log = logging.getLogger("rigel-proxy")
|
||||
|
||||
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||
DB = os.environ.get("RIGEL_DB", os.path.join(HERE, "rigel.db"))
|
||||
LISTEN_HOST = os.environ.get("PROXY_LISTEN_HOST", "0.0.0.0")
|
||||
LISTEN_PORT = int(os.environ.get("PROXY_LISTEN_PORT", "1081"))
|
||||
|
||||
DEFAULT_LOCATIONS = {
|
||||
"tokyo": {"upstream": "10.30.20.154:1080"},
|
||||
"london": {"upstream": "10.30.20.71:1080"},
|
||||
"sydney": {"upstream": "10.30.20.189:1080"},
|
||||
}
|
||||
|
||||
|
||||
def _load_upstreams():
|
||||
"""Derive upstreams from app.LOCATIONS (single source of truth)."""
|
||||
locs = DEFAULT_LOCATIONS
|
||||
try:
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"rigel_app", os.path.join(HERE, "app.py"))
|
||||
mod = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(mod)
|
||||
locs = mod.LOCATIONS
|
||||
log.info("upstreams loaded from app.LOCATIONS")
|
||||
except Exception as exc: # keep serving on built-ins rather than dying
|
||||
log.warning("could not import app.LOCATIONS (%s); using defaults", exc)
|
||||
out = {}
|
||||
for key, val in locs.items():
|
||||
up = (val or {}).get("upstream", "")
|
||||
if ":" not in up:
|
||||
continue
|
||||
host, port = up.rsplit(":", 1)
|
||||
try:
|
||||
port = int(port)
|
||||
except ValueError:
|
||||
continue
|
||||
out[key] = {
|
||||
"host": host,
|
||||
"port": port,
|
||||
"user": os.environ.get(f"UPSTREAM_{key.upper()}_USER") or None,
|
||||
"pass": os.environ.get(f"UPSTREAM_{key.upper()}_PASS") or None,
|
||||
}
|
||||
return out
|
||||
|
||||
|
||||
UPSTREAMS = _load_upstreams()
|
||||
|
||||
|
||||
def authenticate(user, pw):
|
||||
"""Validate per-user creds live against the DB. -> (location, upstream)|None."""
|
||||
try:
|
||||
con = sqlite3.connect(DB, timeout=5)
|
||||
except sqlite3.Error as exc:
|
||||
log.error("db open failed: %s", exc)
|
||||
return None
|
||||
con.row_factory = sqlite3.Row
|
||||
try:
|
||||
row = con.execute(
|
||||
"SELECT location, proxy_pass, status, expires_at FROM subscriptions "
|
||||
"WHERE proxy_user=? ORDER BY id DESC LIMIT 1", (user,)
|
||||
).fetchone()
|
||||
except sqlite3.Error as exc:
|
||||
log.error("db query failed: %s", exc)
|
||||
return None
|
||||
finally:
|
||||
con.close()
|
||||
if not row or row["proxy_pass"] != pw:
|
||||
return None
|
||||
if row["status"] != "active":
|
||||
return None
|
||||
exp = row["expires_at"]
|
||||
if exp:
|
||||
try:
|
||||
# webhook stores naive UTC (datetime.utcnow().isoformat())
|
||||
if datetime.fromisoformat(exp) < datetime.utcnow():
|
||||
log.info("expired sub user=%s exp=%s", user, exp)
|
||||
return None
|
||||
except ValueError:
|
||||
pass
|
||||
up = UPSTREAMS.get(row["location"])
|
||||
if not up:
|
||||
return None
|
||||
return row["location"], up
|
||||
|
||||
|
||||
async def _upstream_open(up, atyp, addr_bytes, port):
|
||||
"""Open a SOCKS5 connection to an upstream exit for the client's target."""
|
||||
ur, uw = await asyncio.wait_for(
|
||||
asyncio.open_connection(up["host"], up["port"]), timeout=20)
|
||||
if up.get("user"):
|
||||
uw.write(b"\x05\x02\x00\x02")
|
||||
await uw.drain()
|
||||
if await ur.readexactly(2) != b"\x05\x02":
|
||||
raise IOError("upstream refused user/pass auth")
|
||||
u = up["user"].encode()
|
||||
p = (up["pass"] or "").encode()
|
||||
uw.write(b"\x01" + bytes([len(u)]) + u + bytes([len(p)]) + p)
|
||||
await uw.drain()
|
||||
if (await ur.readexactly(2))[1] != 0:
|
||||
raise IOError("upstream auth rejected")
|
||||
else:
|
||||
uw.write(b"\x05\x01\x00")
|
||||
await uw.drain()
|
||||
if await ur.readexactly(2) != b"\x05\x00":
|
||||
raise IOError("upstream refused no-auth")
|
||||
# NOTE: for ATYP=3 the domain MUST be length-prefixed; without it the
|
||||
# upstream reads the first domain byte ('a' = 0x61 = 97) as the length
|
||||
# and blocks forever waiting for a 97-byte hostname.
|
||||
addr_field = bytes([len(addr_bytes)]) + addr_bytes if atyp == 3 else addr_bytes
|
||||
uw.write(b"\x05\x01\x00" + bytes([atyp]) + addr_field + struct.pack(">H", port))
|
||||
await uw.drain()
|
||||
rep = await ur.readexactly(4)
|
||||
if rep[1] != 0:
|
||||
raise IOError(f"upstream CONNECT rep={rep[1]}")
|
||||
if rep[3] == 1:
|
||||
await ur.readexactly(4)
|
||||
elif rep[3] == 4:
|
||||
await ur.readexactly(16)
|
||||
elif rep[3] == 3:
|
||||
await ur.readexactly((await ur.readexactly(1))[0])
|
||||
await ur.readexactly(2)
|
||||
return ur, uw
|
||||
|
||||
|
||||
async def _pipe(reader, writer):
|
||||
try:
|
||||
while True:
|
||||
data = await reader.read(65536)
|
||||
if not data:
|
||||
break
|
||||
writer.write(data)
|
||||
await writer.drain()
|
||||
except (OSError, asyncio.IncompleteReadError):
|
||||
pass
|
||||
finally:
|
||||
try:
|
||||
writer.close()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _deny(writer, rep=0x01):
|
||||
writer.write(b"\x05" + bytes([rep]) + b"\x00\x01" + b"\x00" * 4 + b"\x00\x00")
|
||||
|
||||
|
||||
async def handle(reader, writer):
|
||||
peer = writer.get_extra_info("peername")
|
||||
ip = peer[0] if peer else "?"
|
||||
try:
|
||||
ver, nmeth = await reader.readexactly(2)
|
||||
if ver != 5:
|
||||
return
|
||||
methods = await reader.readexactly(nmeth)
|
||||
if 0x02 not in methods:
|
||||
writer.write(b"\x05\xff")
|
||||
await writer.drain()
|
||||
return
|
||||
writer.write(b"\x05\x02")
|
||||
await writer.drain()
|
||||
|
||||
if (await reader.readexactly(1))[0] != 1:
|
||||
return
|
||||
ulen = (await reader.readexactly(1))[0]
|
||||
uname = (await reader.readexactly(ulen)).decode(errors="replace")
|
||||
plen = (await reader.readexactly(1))[0]
|
||||
passwd = (await reader.readexactly(plen)).decode(errors="replace")
|
||||
|
||||
auth = authenticate(uname, passwd)
|
||||
if not auth:
|
||||
log.info("AUTH FAIL user=%r from %s", uname, ip)
|
||||
writer.write(b"\x01\x01")
|
||||
await writer.drain()
|
||||
return
|
||||
location, up = auth
|
||||
writer.write(b"\x01\x00")
|
||||
await writer.drain()
|
||||
|
||||
req = await reader.readexactly(4)
|
||||
cmd, atyp = req[1], req[3]
|
||||
if cmd != 1:
|
||||
_deny(writer, 0x07) # command not supported
|
||||
await writer.drain()
|
||||
return
|
||||
if atyp == 1:
|
||||
addr_bytes = await reader.readexactly(4)
|
||||
target = str(ipaddress.IPv4Address(addr_bytes))
|
||||
elif atyp == 3:
|
||||
n = (await reader.readexactly(1))[0]
|
||||
addr_bytes = await reader.readexactly(n)
|
||||
target = addr_bytes.decode(errors="replace")
|
||||
elif atyp == 4:
|
||||
addr_bytes = await reader.readexactly(16)
|
||||
target = str(ipaddress.IPv6Address(addr_bytes))
|
||||
else:
|
||||
_deny(writer, 0x08)
|
||||
await writer.drain()
|
||||
return
|
||||
port = struct.unpack(">H", await reader.readexactly(2))[0]
|
||||
|
||||
try:
|
||||
ur, uw = await _upstream_open(up, atyp, addr_bytes, port)
|
||||
except (OSError, asyncio.IncompleteReadError, asyncio.TimeoutError, IOError) as exc:
|
||||
log.warning("UPSTREAM FAIL user=%r loc=%s target=%s:%s err=%s",
|
||||
uname, location, target, port, exc)
|
||||
_deny(writer, 0x04) # host unreachable
|
||||
await writer.drain()
|
||||
return
|
||||
|
||||
writer.write(b"\x05\x00\x00\x01" + b"\x00" * 4 + b"\x00\x00")
|
||||
await writer.drain()
|
||||
log.info("OK user=%r loc=%s -> %s:%s from %s", uname, location, target, port, ip)
|
||||
|
||||
await asyncio.gather(_pipe(reader, uw), _pipe(ur, writer))
|
||||
except (asyncio.IncompleteReadError, ConnectionResetError, BrokenPipeError):
|
||||
pass
|
||||
except Exception as exc:
|
||||
log.exception("handler error from %s: %s", ip, exc)
|
||||
finally:
|
||||
try:
|
||||
writer.close()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
async def main():
|
||||
server = await asyncio.start_server(handle, LISTEN_HOST, LISTEN_PORT)
|
||||
addrs = ", ".join(str(s.getsockname()) for s in server.sockets)
|
||||
log.info("rigel proxy listening on %s | locations=%s | db=%s",
|
||||
addrs, sorted(UPSTREAMS), DB)
|
||||
async with server:
|
||||
await server.serve_forever()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
asyncio.run(main())
|
||||
except KeyboardInterrupt:
|
||||
pass
|
||||
Reference in New Issue
Block a user