Files
rigel/proxy_server.py

265 lines
8.8 KiB
Python

#!/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