Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- import asyncio, time, datetime, re
- LOCALPORT = 8000
- taskList = {}
- BYTE_LIMIT = 1024
- CONNECT_STRING = 'HTTP/1.1 200 Continue\r\n\r\n'
- INVALID_REQUEST_STRING = 'HTTP/1.1 500 Error\r\n\r\n'
- HOST_REGEX = re.compile(r'\r\nHost: (.+)\r\n', re.IGNORECASE)
- CONTENT_LENGTH_REGEX = re.compile(r'\r\nContent-Length: (\d+)\r\n', re.IGNORECASE)
- def accept_requests(client_reader, client_writer):
- request_task = asyncio.async(
- ClientHandler(client_reader, client_writer).treat_request()
- )
- taskList[request_task] = (client_reader, client_writer)
- def task_done(request_task):
- del taskList[request_task]
- client_writer.close()
- request_task.add_done_callback(task_done)
- class ClientHandler():
- def __init__(self, client_reader, client_writer):
- self.client_reader = client_reader
- self.client_writer = client_writer
- @asyncio.coroutine
- def treat_request(self):
- self.header = yield from self.readout_header()
- if self.header == None: return
- method = self.header[:self.header.find(" ")]
- if method == "CONNECT":
- try:
- try:
- host, port = self.get_host_port()
- except: return
- try:
- self.request_reader, self.request_writer = yield from asyncio.open_connection(
- host, port
- )
- except: return
- self.client_writer.write(str.encode(CONNECT_STRING))
- task_list = [
- asyncio.async(self.relay_connection(self.client_reader, self.request_writer)),
- asyncio.async(self.relay_connection(self.request_reader, self.client_writer))
- ]
- yield from asyncio.wait(task_list)
- except: return
- else:
- try:
- host, port = self.get_host_port()
- except: return
- if port == 443: port = 80
- self.content_length = CONTENT_LENGTH_REGEX.search(self.header)
- if self.content_length: self.content_length = int(self.content_length.group(1))
- if self.content_length:
- request_body = b''
- while len(request_body) < self.content_length:
- try:
- data_read = self.client_reader.read(BYTE_LIMIT)
- if data_read: request_body += data_read
- else: break
- except: break
- if len(request_body) < self.content_length:
- self.client_writer.write(str.encode(INVALID_REQUEST_STRING))
- return
- try:
- request_reader, request_writer = yield from asyncio.open_connection(
- host, port
- )
- except: return
- try:
- request_writer.write(str.encode(self.header))
- except:
- request_writer.close()
- return
- try:
- request_writer.write(str.encode(request_body))
- except NameError: pass
- while 1:
- try:
- buf = yield from request_reader.read(BYTE_LIMIT)
- if buf: self.client_writer.write(buf)
- else: break
- except: break
- request_writer.close()
- self.client_writer.close()
- def get_host_port(self):
- host_port = HOST_REGEX.search(self.header)
- if not host_port:
- self.client_writer.write(str.encode(INVALID_REQUEST_STRING))
- return
- else: host_port = host_port.group(1)
- if host_port.find(":") != -1:
- host_port = host_port.split(":")
- host = host_port[0]
- port = host_port[1]
- else:
- host = host_port
- port = 443
- return host, port
- @asyncio.coroutine
- def process_get_or_post(self, host, port, body=None):
- request_reader, request_writer = yield from asyncio.open_connection(
- host, port
- )
- request_writer.write(str.encode(self.header))
- if body: request_writer.write(str.encode(body))
- while 1:
- try:
- buf = yield from request_reader.read(BYTE_LIMIT)
- if buf: self.client_writer.write(buf)
- else: return
- except: return
- @asyncio.coroutine
- def relay_connection(self, reader, writer, max_count=10):
- loop_count = 0
- while True:
- if loop_count >= max_count: break
- else: loop_count += 1
- try:
- data_read = yield from reader.read(BYTE_LIMIT)
- if data_read:
- writer.write(data_read)
- loop_count = 0
- else: continue
- except:
- return
- @asyncio.coroutine
- def readout_header(self) :
- header = ""
- while True:
- try:
- data_read = yield from self.client_reader.read(1)
- if not data_read:
- return
- header += data_read.decode()
- if header.find("\r\n\r\n") != -1: break
- if header.find("\n\n") != -1: break
- if header.find("\n\r\n") != -1: break
- except:
- return
- return header
- @asyncio.coroutine
- def start_server(host, port):
- server = asyncio.start_server(accept_requests, host, port)
- return server
- def entry_point():
- try:
- loop = asyncio.get_event_loop()
- loop.run_until_complete(start_server('', LOCALPORT))
- loop.run_forever()
- except KeyboardInterrupt: pass
- except: pass
- finally:
- loop.close()
- if __name__ == "__main__":
- exit(entry_point())
Advertisement
Add Comment
Please, Sign In to add comment