Browse Source

implement udp relay

1.4
clowwindy 11 years ago
parent
commit
0f23b58656
  1. 77
      shadowsocks/udprelay.py

77
shadowsocks/udprelay.py

@ -70,6 +70,7 @@ import threading
import socket import socket
import logging import logging
import struct import struct
import encrypt
import eventloop import eventloop
BUF_SIZE = 65536 BUF_SIZE = 65536
@ -92,7 +93,8 @@ def parse_header(data):
addrlen = ord(data[1]) addrlen = ord(data[1])
if len(data) >= 2 + addrlen: if len(data) >= 2 + addrlen:
dest_addr = data[2:2 + addrlen] dest_addr = data[2:2 + addrlen]
dest_port = struct.unpack('>H', data[2 + addrlen:4 + addrlen])[0] dest_port = struct.unpack('>H', data[2 + addrlen:4 +
addrlen])[0]
header_length = 4 + addrlen header_length = 4 + addrlen
else: else:
logging.warn('[udp] header is too short') logging.warn('[udp] header is too short')
@ -112,6 +114,10 @@ def parse_header(data):
return (addrtype, dest_addr, dest_port, header_length) return (addrtype, dest_addr, dest_port, header_length)
def client_key(a, b, c, d):
return '%s:%s:%s:%s' % (a, b, c, d)
class UDPRelay(object): class UDPRelay(object):
def __init__(self, listen_addr='127.0.0.1', listen_port=1080, def __init__(self, listen_addr='127.0.0.1', listen_port=1080,
remote_addr='127.0.0.1', remote_port=8387, password=None, remote_addr='127.0.0.1', remote_port=8387, password=None,
@ -125,24 +131,84 @@ class UDPRelay(object):
self._timeout = timeout self._timeout = timeout
self._is_local = is_local self._is_local = is_local
self._eventloop = eventloop.EventLoop() self._eventloop = eventloop.EventLoop()
self._cache = {} # TODO replace this dictionary with an LRU cache self._cache = {} # TODO replace ith an LRU cache
self._client_fd_to_server_addr = {} # TODO replace ith an LRU cache
def _handle_server(self): def _handle_server(self):
server = self._server_socket server = self._server_socket
data = server.recvfrom(BUF_SIZE) data, r_addr = server.recvfrom(BUF_SIZE)
if self._is_local: if self._is_local:
frag = ord(data[2]) frag = ord(data[2])
if frag != 0: if frag != 0:
logging.warn('drop a message since frag is not 0') logging.warn('drop a message since frag is not 0')
return
else: else:
data = data[3:] data = data[3:]
else: else:
decrypt # decrypt data
data = encrypt.encrypt_all(self._password, self._method, 0, data)
if not data:
return
header_result = parse_header(data)
if header_result is None:
return
addrtype, dest_addr, dest_port, header_length = header_result
if self._is_local:
server_addr, server_port = self._remote_addr, self._remote_port
else:
server_addr, server_port = dest_addr, dest_port
key = client_key(r_addr[0], r_addr[1], dest_addr, dest_port)
client = self._cache.get(key, None)
if not client:
# TODO async getaddrinfo
addrs = socket.getaddrinfo(server_addr, server_port, 0,
socket.SOCK_DGRAM, socket.SOL_UDP)
if addrs:
af, socktype, proto, canonname, sa = addrs[0]
client = socket.socket(af, socktype, proto)
client.setblocking(False)
self._cache[key] = client
self._client_fd_to_server_addr[client.fileno()] = r_addr
else:
# drop
return
self._eventloop.add(client, eventloop.MODE_IN)
if self._is_local:
data = encrypt.encrypt_all(self._password, self._method, 1, data)
if not data:
return
client.sendto(data, (server_addr, server_port))
def _handle_client(self, sock): def _handle_client(self, sock):
# TODO data, r_addr = sock.recvfrom(BUF_SIZE)
if not self._is_local:
addrlen = len(r_addr[0])
if addrlen > 255:
# drop
return
data = '\x03' + chr(addrlen) + r_addr[0] + \
struct.pack('>H', r_addr[1]) + data
response = encrypt.encrypt_all(self._password, self._method, 1,
data)
if not response:
return
else:
data = encrypt.encrypt_all(self._password, self._method, 0,
data)
if not data:
return
# addrtype, dest_addr, dest_port, header_length = parse_header(data)
response = '\x00\x00\0x00' + data
client_addr = self._client_fd_to_server_addr.get(sock.fileno(), None)
if client_addr:
self._server_socket.sendto(response, client_addr)
else:
pass pass
# self._eventloop.remove(sock)
# sock.close()
# TODO remove it from cache else we can't close it
def _run(self): def _run(self):
server_socket = self._server_socket server_socket = self._server_socket
@ -171,3 +237,4 @@ class UDPRelay(object):
t = threading.Thread(target=self._run) t = threading.Thread(target=self._run)
t.setDaemon(True) t.setDaemon(True)
t.start() t.start()
self._thread = t

Loading…
Cancel
Save