File size: 4,032 Bytes
542be03
 
 
 
52bfd89
542be03
 
 
 
52bfd89
542be03
 
 
 
 
52bfd89
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
542be03
 
52bfd89
542be03
 
52bfd89
542be03
 
52bfd89
 
989566d
52bfd89
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
542be03
 
 
 
 
 
 
 
 
 
 
 
 
52bfd89
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
72a5949
52bfd89
 
 
 
 
 
 
 
 
 
 
54c36c0
542be03
 
 
 
 
 
 
 
 
 
52bfd89
542be03
 
 
 
 
 
 
 
989566d
542be03
989566d
542be03
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
import asyncio
import os
import signal
import sys
from urllib.parse import urlparse

PROXY_PORT = 7860


async def health_check(writer):
    writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 3\r\n\r\nOK\n")
    await writer.drain()
    writer.close()


async def error(writer, status, message):
    body = message.encode()
    writer.write(f"HTTP/1.1 {status} {message}\r\nContent-Length: {len(body)}\r\n\r\n{message}\n".encode())
    await writer.drain()
    writer.close()


async def forward(src, dst):
    try:
        while True:
            data = await src.read(65536)
            if not data:
                break
            dst.write(data)
            await dst.drain()
    except (ConnectionResetError, BrokenPipeError, OSError):
        pass
    finally:
        try:
            dst.close()
        except OSError:
            pass


async def handle_tunnel(target_host, target_port, reader, writer):
    try:
        backend_reader, backend_writer = await asyncio.wait_for(
            asyncio.open_connection(target_host, target_port), timeout=15
        )
    except (OSError, asyncio.TimeoutError):
        await error(writer, 502, "Cannot connect to target")
        return

    writer.write(b"HTTP/1.1 200 Connection Established\r\n\r\n")
    await writer.drain()

    await asyncio.gather(
        forward(reader, backend_writer),
        forward(backend_reader, writer),
    )


async def forward_request(host, port, data, reader, writer):
    try:
        backend_reader, backend_writer = await asyncio.wait_for(
            asyncio.open_connection(host, port), timeout=15
        )
    except (OSError, asyncio.TimeoutError):
        await error(writer, 502, "Cannot connect to target")
        return

    backend_writer.write(data)
    await backend_writer.drain()

    await asyncio.gather(
        forward(reader, backend_writer),
        forward(backend_reader, writer),
    )


async def handle_client(reader, writer):
    try:
        data = await asyncio.wait_for(reader.read(4096), timeout=10)
        if not data:
            writer.close()
            return

        parts = data.split(b" ", 2)
        if len(parts) < 2:
            writer.close()
            return

        method = parts[0]
        raw_path = parts[1]

        if method == b"CONNECT":
            target = raw_path.split(b":")
            if len(target) == 2:
                try:
                    await handle_tunnel(target[0].decode(), int(target[1]), reader, writer)
                    return
                except (ValueError, UnicodeDecodeError):
                    pass

        path = raw_path.split(b"?")[0]
        if path in (b"/", b"/health"):
            await health_check(writer)
            return

        if raw_path.startswith(b"http://") or raw_path.startswith(b"https://"):
            parsed = urlparse(raw_path.decode())
            host = parsed.hostname
            port = parsed.port or (443 if parsed.scheme == "https" else 80)
            if host:
                await forward_request(host, port, data, reader, writer)
                return

        await error(writer, 400, "Bad Request")
    except (ConnectionResetError, BrokenPipeError, OSError, asyncio.TimeoutError):
        pass
    finally:
        try:
            writer.close()
        except OSError:
            pass


async def main():
    print(f"starting proxy on port {PROXY_PORT}...", flush=True)
    server = await asyncio.start_server(handle_client, "0.0.0.0", PROXY_PORT)

    async def shutdown():
        server.close()

    loop = asyncio.get_event_loop()
    for sig in (signal.SIGTERM, signal.SIGINT):
        try:
            loop.add_signal_handler(sig, lambda: asyncio.create_task(shutdown()))
        except NotImplementedError:
            signal.signal(sig, lambda n, f: asyncio.run_coroutine_threadsafe(shutdown(), loop))

    async with server:
        print(f"proxy server ready on port {PROXY_PORT}", flush=True)
        await server.serve_forever()


if __name__ == "__main__":
    asyncio.run(main())