codekingpro/portable-devtools
115k
1import base642import gzip3import importlib4import io5import logging6import secrets7import zlib8 9from . import packet10from . import payload11 12default_logger = logging.getLogger('engineio.server')13 14 15class BaseServer:16 compression_methods = ['gzip', 'deflate']17 event_names = ['connect', 'disconnect', 'message']18 valid_transports = ['polling', 'websocket']19 _default_monitor_clients = True20 sequence_number = 021 22 def __init__(self, async_mode=None, ping_interval=25, ping_timeout=20,23 max_http_buffer_size=1000000, allow_upgrades=True,24 http_compression=True, compression_threshold=1024,25 cookie=None, cors_allowed_origins=None,26 cors_credentials=True, logger=False, json=None,27 async_handlers=True, monitor_clients=None, transports=None,28 **kwargs):29 self.ping_timeout = ping_timeout30 if isinstance(ping_interval, tuple):31 self.ping_interval = ping_interval[0]32 self.ping_interval_grace_period = ping_interval[1]33 else:34 self.ping_interval = ping_interval35 self.ping_interval_grace_period = 036 self.max_http_buffer_size = max_http_buffer_size37 self.allow_upgrades = allow_upgrades38 self.http_compression = http_compression39 self.compression_threshold = compression_threshold40 self.cookie = cookie41 self.cors_allowed_origins = cors_allowed_origins42 self.cors_credentials = cors_credentials43 self.async_handlers = async_handlers44 self.sockets = {}45 self.handlers = {}46 self.log_message_keys = set()47 self.start_service_task = monitor_clients \48 if monitor_clients is not None else self._default_monitor_clients49 self.service_task_handle = None50 self.service_task_event = None51 if json is not None:52 packet.Packet.json = json53 if not isinstance(logger, bool):54 self.logger = logger55 else:56 self.logger = default_logger57 if self.logger.level == logging.NOTSET:58 if logger:59 self.logger.setLevel(logging.INFO)60 else:61 self.logger.setLevel(logging.ERROR)62 self.logger.addHandler(logging.StreamHandler())63 modes = self.async_modes()64 if async_mode is not None:65 modes = [async_mode] if async_mode in modes else []66 self._async = None67 self.async_mode = None68 for mode in modes:69 try:70 self._async = importlib.import_module(71 'engineio.async_drivers.' + mode)._async72 asyncio_based = self._async['asyncio'] \73 if 'asyncio' in self._async else False74 if asyncio_based != self.is_asyncio_based():75 continue # pragma: no cover76 self.async_mode = mode77 break78 except ImportError:79 pass80 if self.async_mode is None:81 raise ValueError('Invalid async_mode specified')82 if self.is_asyncio_based() and \83 ('asyncio' not in self._async or not84 self._async['asyncio']): # pragma: no cover85 raise ValueError('The selected async_mode is not asyncio '86 'compatible')87 if not self.is_asyncio_based() and 'asyncio' in self._async and \88 self._async['asyncio']: # pragma: no cover89 raise ValueError('The selected async_mode requires asyncio and '90 'must use the AsyncServer class')91 if transports is not None:92 if isinstance(transports, str):93 transports = [transports]94 transports = [transport for transport in transports95 if transport in self.valid_transports]96 if not transports:97 raise ValueError('No valid transports provided')98 self.transports = transports or self.valid_transports99 self.logger.info('Server initialized for %s.', self.async_mode)100 101 def is_asyncio_based(self):102 return False103 104 def async_modes(self):105 return ['eventlet', 'gevent_uwsgi', 'gevent', 'threading']106 107 def on(self, event, handler=None):108 """Register an event handler.109 110 :param event: The event name. Can be ``'connect'``, ``'message'`` or111 ``'disconnect'``.112 :param handler: The function that should be invoked to handle the113 event. When this parameter is not given, the method114 acts as a decorator for the handler function.115 116 Example usage::117 118 # as a decorator:119 @eio.on('connect')120 def connect_handler(sid, environ):121 print('Connection request')122 if environ['REMOTE_ADDR'] in blacklisted:123 return False # reject124 125 # as a method:126 def message_handler(sid, msg):127 print('Received message: ', msg)128 eio.send(sid, 'response')129 eio.on('message', message_handler)130 131 The handler function receives the ``sid`` (session ID) for the132 client as first argument. The ``'connect'`` event handler receives the133 WSGI environment as a second argument, and can return ``False`` to134 reject the connection. The ``'message'`` handler receives the message135 payload as a second argument. The ``'disconnect'`` handler does not136 take a second argument.137 """138 if event not in self.event_names:139 raise ValueError('Invalid event')140 141 def set_handler(handler):142 self.handlers[event] = handler143 return handler144 145 if handler is None:146 return set_handler147 set_handler(handler)148 149 def transport(self, sid):150 """Return the name of the transport used by the client.151 152 The two possible values returned by this function are ``'polling'``153 and ``'websocket'``.154 155 :param sid: The session of the client.156 """157 return 'websocket' if self._get_socket(sid).upgraded else 'polling'158 159 def create_queue(self, *args, **kwargs):160 """Create a queue object using the appropriate async model.161 162 This is a utility function that applications can use to create a queue163 without having to worry about using the correct call for the selected164 async mode.165 """166 return self._async['queue'](*args, **kwargs)167 168 def get_queue_empty_exception(self):169 """Return the queue empty exception for the appropriate async model.170 171 This is a utility function that applications can use to work with a172 queue without having to worry about using the correct call for the173 selected async mode.174 """175 return self._async['queue_empty']176 177 def create_event(self, *args, **kwargs):178 """Create an event object using the appropriate async model.179 180 This is a utility function that applications can use to create an181 event without having to worry about using the correct call for the182 selected async mode.183 """184 return self._async['event'](*args, **kwargs)185 186 def generate_id(self):187 """Generate a unique session id."""188 id = base64.b64encode(189 secrets.token_bytes(12) + self.sequence_number.to_bytes(3, 'big'))190 self.sequence_number = (self.sequence_number + 1) & 0xffffff191 return id.decode('utf-8').replace('/', '_').replace('+', '-')192 193 def _generate_sid_cookie(self, sid, attributes):194 """Generate the sid cookie."""195 cookie = attributes.get('name', 'io') + '=' + sid196 for attribute, value in attributes.items():197 if attribute == 'name':198 continue199 if callable(value):200 value = value()201 if value is True:202 cookie += '; ' + attribute203 else:204 cookie += '; ' + attribute + '=' + value205 return cookie206 207 def _upgrades(self, sid, transport):208 """Return the list of possible upgrades for a client connection."""209 if not self.allow_upgrades or self._get_socket(sid).upgraded or \210 transport == 'websocket':211 return []212 if self._async['websocket'] is None: # pragma: no cover213 self._log_error_once(214 'The WebSocket transport is not available, you must install a '215 'WebSocket server that is compatible with your async mode to '216 'enable it. See the documentation for details.',217 'no-websocket')218 return []219 return ['websocket']220 221 def _get_socket(self, sid):222 """Return the socket object for a given session."""223 try:224 s = self.sockets[sid]225 except KeyError:226 raise KeyError('Session not found')227 if s.closed:228 del self.sockets[sid]229 raise KeyError('Session is disconnected')230 return s231 232 def _ok(self, packets=None, headers=None, jsonp_index=None):233 """Generate a successful HTTP response."""234 if packets is not None:235 if headers is None:236 headers = []237 headers += [('Content-Type', 'text/plain; charset=UTF-8')]238 return {'status': '200 OK',239 'headers': headers,240 'response': payload.Payload(packets=packets).encode(241 jsonp_index=jsonp_index).encode('utf-8')}242 else:243 return {'status': '200 OK',244 'headers': [('Content-Type', 'text/plain')],245 'response': b'OK'}246 247 def _bad_request(self, message=None):248 """Generate a bad request HTTP error response."""249 if message is None:250 message = 'Bad Request'251 message = packet.Packet.json.dumps(message)252 return {'status': '400 BAD REQUEST',253 'headers': [('Content-Type', 'text/plain')],254 'response': message.encode('utf-8')}255 256 def _method_not_found(self):257 """Generate a method not found HTTP error response."""258 return {'status': '405 METHOD NOT FOUND',259 'headers': [('Content-Type', 'text/plain')],260 'response': b'Method Not Found'}261 262 def _unauthorized(self, message=None):263 """Generate a unauthorized HTTP error response."""264 if message is None:265 message = 'Unauthorized'266 message = packet.Packet.json.dumps(message)267 return {'status': '401 UNAUTHORIZED',268 'headers': [('Content-Type', 'application/json')],269 'response': message.encode('utf-8')}270 271 def _cors_allowed_origins(self, environ):272 default_origins = []273 if 'wsgi.url_scheme' in environ and 'HTTP_HOST' in environ:274 default_origins.append('{scheme}://{host}'.format(275 scheme=environ['wsgi.url_scheme'], host=environ['HTTP_HOST']))276 if 'HTTP_X_FORWARDED_PROTO' in environ or \277 'HTTP_X_FORWARDED_HOST' in environ:278 scheme = environ.get(279 'HTTP_X_FORWARDED_PROTO',280 environ['wsgi.url_scheme']).split(',')[0].strip()281 default_origins.append('{scheme}://{host}'.format(282 scheme=scheme, host=environ.get(283 'HTTP_X_FORWARDED_HOST', environ['HTTP_HOST']).split(284 ',')[0].strip()))285 if self.cors_allowed_origins is None:286 allowed_origins = default_origins287 elif self.cors_allowed_origins == '*':288 allowed_origins = None289 elif isinstance(self.cors_allowed_origins, str):290 allowed_origins = [self.cors_allowed_origins]291 elif callable(self.cors_allowed_origins):292 origin = environ.get('HTTP_ORIGIN')293 allowed_origins = [origin] \294 if self.cors_allowed_origins(origin) else []295 else:296 allowed_origins = self.cors_allowed_origins297 return allowed_origins298 299 def _cors_headers(self, environ):300 """Return the cross-origin-resource-sharing headers."""301 if self.cors_allowed_origins == []:302 # special case, CORS handling is completely disabled303 return []304 headers = []305 allowed_origins = self._cors_allowed_origins(environ)306 if 'HTTP_ORIGIN' in environ and \307 (allowed_origins is None or environ['HTTP_ORIGIN'] in308 allowed_origins):309 headers = [('Access-Control-Allow-Origin', environ['HTTP_ORIGIN'])]310 if environ['REQUEST_METHOD'] == 'OPTIONS':311 headers += [('Access-Control-Allow-Methods', 'OPTIONS, GET, POST')]312 if 'HTTP_ACCESS_CONTROL_REQUEST_HEADERS' in environ:313 headers += [('Access-Control-Allow-Headers',314 environ['HTTP_ACCESS_CONTROL_REQUEST_HEADERS'])]315 if self.cors_credentials:316 headers += [('Access-Control-Allow-Credentials', 'true')]317 return headers318 319 def _gzip(self, response):320 """Apply gzip compression to a response."""321 bytesio = io.BytesIO()322 with gzip.GzipFile(fileobj=bytesio, mode='w') as gz:323 gz.write(response)324 return bytesio.getvalue()325 326 def _deflate(self, response):327 """Apply deflate compression to a response."""328 return zlib.compress(response)329 330 def _log_error_once(self, message, message_key):331 """Log message with logging.ERROR level the first time, then log332 with given level."""333 if message_key not in self.log_message_keys:334 self.logger.error(message + ' (further occurrences of this error '335 'will be logged with level INFO)')336 self.log_message_keys.add(message_key)337 else:338 self.logger.info(message)339 