codeaddict

Untitled

Mar 24th, 2016
108
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 5.92 KB | None | 0 0
  1. import asyncio, time, datetime, re
  2. LOCALPORT = 8000
  3.  
  4. taskList = {}
  5. BYTE_LIMIT = 1024
  6.  
  7. CONNECT_STRING = 'HTTP/1.1 200 Continue\r\n\r\n'
  8. INVALID_REQUEST_STRING = 'HTTP/1.1 500 Error\r\n\r\n'
  9.  
  10. HOST_REGEX = re.compile(r'\r\nHost: (.+)\r\n', re.IGNORECASE)
  11. CONTENT_LENGTH_REGEX = re.compile(r'\r\nContent-Length: (\d+)\r\n', re.IGNORECASE)
  12.  
  13. def accept_requests(client_reader, client_writer):
  14.     request_task = asyncio.async(
  15.         ClientHandler(client_reader, client_writer).treat_request()
  16.     )
  17.     taskList[request_task] = (client_reader, client_writer)
  18.  
  19.     def task_done(request_task):
  20.         del taskList[request_task]
  21.         client_writer.close()
  22.  
  23.     request_task.add_done_callback(task_done)
  24.  
  25. class ClientHandler():
  26.     def __init__(self, client_reader, client_writer):
  27.         self.client_reader = client_reader
  28.         self.client_writer = client_writer
  29.  
  30.     @asyncio.coroutine
  31.     def treat_request(self):
  32.         self.header = yield from self.readout_header()
  33.         if self.header == None: return
  34.         method = self.header[:self.header.find(" ")]
  35.  
  36.         if method == "CONNECT":
  37.             try:
  38.                 try:
  39.                     host, port = self.get_host_port()
  40.                 except: return
  41.  
  42.                 try:
  43.                     self.request_reader, self.request_writer = yield from asyncio.open_connection(
  44.                         host, port
  45.                     )
  46.                 except: return
  47.  
  48.                 self.client_writer.write(str.encode(CONNECT_STRING))
  49.                 task_list = [
  50.                     asyncio.async(self.relay_connection(self.client_reader, self.request_writer)),
  51.                     asyncio.async(self.relay_connection(self.request_reader, self.client_writer))
  52.                 ]
  53.  
  54.                 yield from asyncio.wait(task_list)
  55.             except: return
  56.         else:
  57.             try:
  58.                 host, port = self.get_host_port()
  59.             except: return
  60.  
  61.             if port == 443: port = 80
  62.  
  63.             self.content_length = CONTENT_LENGTH_REGEX.search(self.header)
  64.  
  65.             if self.content_length: self.content_length = int(self.content_length.group(1))
  66.            
  67.             if self.content_length:
  68.                 request_body = b''
  69.                 while len(request_body) < self.content_length:
  70.                     try:
  71.                         data_read = self.client_reader.read(BYTE_LIMIT)
  72.                         if data_read: request_body += data_read
  73.                         else: break
  74.                     except: break
  75.  
  76.                 if len(request_body) < self.content_length:
  77.                     self.client_writer.write(str.encode(INVALID_REQUEST_STRING))
  78.                     return
  79.  
  80.             try:
  81.                 request_reader, request_writer = yield from asyncio.open_connection(
  82.                     host, port
  83.                 )
  84.             except: return
  85.  
  86.             try:
  87.                 request_writer.write(str.encode(self.header))
  88.             except:
  89.                 request_writer.close()
  90.                 return
  91.  
  92.             try:
  93.                 request_writer.write(str.encode(request_body))
  94.             except NameError: pass
  95.  
  96.             while 1:
  97.                 try:
  98.                     buf = yield from request_reader.read(BYTE_LIMIT)
  99.                     if buf: self.client_writer.write(buf)
  100.                     else: break
  101.                 except: break
  102.  
  103.             request_writer.close()
  104.             self.client_writer.close()
  105.  
  106.     def get_host_port(self):
  107.         host_port = HOST_REGEX.search(self.header)
  108.  
  109.         if not host_port:
  110.             self.client_writer.write(str.encode(INVALID_REQUEST_STRING))
  111.             return
  112.         else: host_port = host_port.group(1)
  113.  
  114.         if host_port.find(":") != -1:
  115.             host_port = host_port.split(":")
  116.             host = host_port[0]
  117.             port = host_port[1]
  118.         else:
  119.             host = host_port
  120.             port = 443
  121.        
  122.         return host, port
  123.  
  124.     @asyncio.coroutine
  125.     def process_get_or_post(self, host, port, body=None):
  126.         request_reader, request_writer = yield from asyncio.open_connection(
  127.             host, port
  128.         )
  129.  
  130.         request_writer.write(str.encode(self.header))
  131.         if body: request_writer.write(str.encode(body))
  132.         while 1:
  133.             try:
  134.                 buf = yield from request_reader.read(BYTE_LIMIT)
  135.                 if buf: self.client_writer.write(buf)
  136.                 else: return
  137.             except: return
  138.  
  139.     @asyncio.coroutine
  140.     def relay_connection(self, reader, writer, max_count=10):
  141.         loop_count = 0
  142.         while True:
  143.             if loop_count >= max_count: break
  144.             else: loop_count += 1
  145.             try:
  146.                 data_read = yield from reader.read(BYTE_LIMIT)
  147.                 if data_read:
  148.                     writer.write(data_read)
  149.                     loop_count = 0
  150.                 else: continue
  151.             except:
  152.                 return
  153.  
  154.     @asyncio.coroutine
  155.     def readout_header(self)    :
  156.         header = ""
  157.         while True:
  158.             try:
  159.                 data_read = yield from self.client_reader.read(1)
  160.  
  161.                 if not data_read:
  162.                     return
  163.  
  164.                 header += data_read.decode()
  165.                 if header.find("\r\n\r\n") != -1: break
  166.                 if header.find("\n\n") != -1: break
  167.                 if header.find("\n\r\n") != -1: break
  168.             except:
  169.                 return
  170.  
  171.         return header
  172.  
  173. @asyncio.coroutine
  174. def start_server(host, port):
  175.     server = asyncio.start_server(accept_requests, host, port)
  176.     return server
  177.  
  178.  
  179. def entry_point():
  180.     try:
  181.         loop = asyncio.get_event_loop()
  182.         loop.run_until_complete(start_server('', LOCALPORT))
  183.         loop.run_forever()
  184.     except KeyboardInterrupt: pass
  185.     except: pass
  186.     finally:
  187.         loop.close()
  188.  
  189.  
  190. if __name__ == "__main__":
  191.     exit(entry_point())
Advertisement
Add Comment
Please, Sign In to add comment