| 1 | # -*- coding: utf-8 -*- |
| 2 | # Description: |
| 3 | # SPDX-License-Identifier: GPL-3.0-or-later |
| 4 | |
| 5 | import errno |
| 6 | import socket |
| 7 | |
| 8 | try: |
| 9 | import ssl |
| 10 | except ImportError: |
| 11 | _TLS_SUPPORT = False |
| 12 | else: |
| 13 | _TLS_SUPPORT = True |
| 14 | |
| 15 | if _TLS_SUPPORT: |
| 16 | try: |
| 17 | PROTOCOL_TLS = ssl.PROTOCOL_TLS |
| 18 | except AttributeError: |
| 19 | PROTOCOL_TLS = ssl.PROTOCOL_SSLv23 |
| 20 | |
| 21 | from bases.FrameworkServices.SimpleService import SimpleService |
| 22 | |
| 23 | |
| 24 | DEFAULT_CONNECT_TIMEOUT = 2.0 |
| 25 | DEFAULT_READ_TIMEOUT = 2.0 |
| 26 | DEFAULT_WRITE_TIMEOUT = 2.0 |
| 27 | |
| 28 | |
| 29 | class SocketService(SimpleService): |
| 30 | def __init__(self, configuration=None, name=None): |
| 31 | self._sock = None |
| 32 | self._keep_alive = False |
| 33 | self.host = 'localhost' |
| 34 | self.port = None |
| 35 | self.unix_socket = None |
| 36 | self.dgram_socket = False |
| 37 | self.request = '' |
| 38 | self.tls = False |
| 39 | self.cert = None |
| 40 | self.key = None |
| 41 | self.__socket_config = None |
| 42 | self.__empty_request = "".encode() |
| 43 | SimpleService.__init__(self, configuration=configuration, name=name) |
| 44 | self.connect_timeout = configuration.get('connect_timeout', DEFAULT_CONNECT_TIMEOUT) |
| 45 | self.read_timeout = configuration.get('read_timeout', DEFAULT_READ_TIMEOUT) |
| 46 | self.write_timeout = configuration.get('write_timeout', DEFAULT_WRITE_TIMEOUT) |
| 47 | |
| 48 | def _socket_error(self, message=None): |
| 49 | if self.unix_socket is not None: |
| 50 | self.error('unix socket "{socket}": {message}'.format(socket=self.unix_socket, |
| 51 | message=message)) |
| 52 | else: |
| 53 | if self.__socket_config is not None: |
| 54 | _, _, _, _, sa = self.__socket_config |
| 55 | self.error('socket to "{address}" port {port}: {message}'.format(address=sa[0], |
| 56 | port=sa[1], |
| 57 | message=message)) |
| 58 | else: |
| 59 | self.error('unknown socket: {0}'.format(message)) |
| 60 | |
| 61 | def _connect2socket(self, res=None): |
| 62 | """ |
| 63 | Connect to a socket, passing the result of getaddrinfo() |
| 64 | :return: boolean |
| 65 | """ |
| 66 | if res is None: |
| 67 | res = self.__socket_config |
| 68 | if res is None: |
| 69 | self.error("Cannot create socket to 'None':") |
| 70 | return False |
| 71 | |
| 72 | af, sock_type, proto, _, sa = res |
| 73 | try: |
| 74 | self.debug('Creating socket to "{address}", port {port}'.format(address=sa[0], port=sa[1])) |
| 75 | self._sock = socket.socket(af, sock_type, proto) |
| 76 | except socket.error as error: |
| 77 | self.error('Failed to create socket "{address}", port {port}, error: {error}'.format(address=sa[0], |
| 78 | port=sa[1], |
| 79 | error=error)) |
| 80 | self._sock = None |
| 81 | self.__socket_config = None |
| 82 | return False |
| 83 | |
| 84 | if self.tls: |
| 85 | try: |
| 86 | self.debug('Encapsulating socket with TLS') |
| 87 | self.debug('Using keyfile: {0}, certfile: {1}, cert_reqs: {2}, ssl_version: {3}'.format( |
| 88 | self.key, self.cert, ssl.CERT_NONE, PROTOCOL_TLS |
| 89 | )) |
| 90 | self._sock = ssl.wrap_socket(self._sock, |
| 91 | keyfile=self.key, |
| 92 | certfile=self.cert, |
| 93 | server_side=False, |
| 94 | cert_reqs=ssl.CERT_NONE, |
| 95 | ssl_version=PROTOCOL_TLS, |
| 96 | ) |
| 97 | except (socket.error, ssl.SSLError, IOError, OSError) as error: |
| 98 | self.error('failed to wrap socket : {0}'.format(repr(error))) |
| 99 | self._disconnect() |
| 100 | self.__socket_config = None |
| 101 | return False |
| 102 | |
| 103 | try: |
| 104 | self.debug('connecting socket to "{address}", port {port}'.format(address=sa[0], port=sa[1])) |
| 105 | self._sock.settimeout(self.connect_timeout) |
| 106 | self.debug('set socket connect timeout to: {0}'.format(self._sock.gettimeout())) |
| 107 | self._sock.connect(sa) |
| 108 | except (socket.error, ssl.SSLError) as error: |
| 109 | self.error('Failed to connect to "{address}", port {port}, error: {error}'.format(address=sa[0], |
| 110 | port=sa[1], |
| 111 | error=error)) |
| 112 | self._disconnect() |
| 113 | self.__socket_config = None |
| 114 | return False |
| 115 | |
| 116 | self.debug('connected to "{address}", port {port}'.format(address=sa[0], port=sa[1])) |
| 117 | self.__socket_config = res |
| 118 | return True |
| 119 | |
| 120 | def _connect2unixsocket(self): |
| 121 | """ |
| 122 | Connect to a unix socket, given its filename |
| 123 | :return: boolean |
| 124 | """ |
| 125 | if self.unix_socket is None: |
| 126 | self.error("cannot connect to unix socket 'None'") |
| 127 | return False |
| 128 | |
| 129 | try: |
| 130 | self.debug('attempting DGRAM unix socket "{0}"'.format(self.unix_socket)) |
| 131 | self._sock = socket.socket(socket.AF_UNIX, socket.SOCK_DGRAM) |
| 132 | self._sock.settimeout(self.connect_timeout) |
| 133 | self.debug('set socket connect timeout to: {0}'.format(self._sock.gettimeout())) |
| 134 | self._sock.connect(self.unix_socket) |
| 135 | self.debug('connected DGRAM unix socket "{0}"'.format(self.unix_socket)) |
| 136 | return True |
| 137 | except socket.error as error: |
| 138 | self.debug('Failed to connect DGRAM unix socket "{socket}": {error}'.format(socket=self.unix_socket, |
| 139 | error=error)) |
| 140 | |
| 141 | try: |
| 142 | self.debug('attempting STREAM unix socket "{0}"'.format(self.unix_socket)) |
| 143 | self._sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) |
| 144 | self._sock.settimeout(self.connect_timeout) |
| 145 | self.debug('set socket connect timeout to: {0}'.format(self._sock.gettimeout())) |
| 146 | self._sock.connect(self.unix_socket) |
| 147 | self.debug('connected STREAM unix socket "{0}"'.format(self.unix_socket)) |
| 148 | return True |
| 149 | except socket.error as error: |
| 150 | self.debug('Failed to connect STREAM unix socket "{socket}": {error}'.format(socket=self.unix_socket, |
| 151 | error=error)) |
| 152 | self._sock = None |
| 153 | return False |
| 154 | |
| 155 | def _connect(self): |
| 156 | """ |
| 157 | Recreate socket and connect to it since sockets cannot be reused after closing |
| 158 | Available configurations are IPv6, IPv4 or UNIX socket |
| 159 | :return: |
| 160 | """ |
| 161 | try: |
| 162 | if self.unix_socket is not None: |
| 163 | self._connect2unixsocket() |
| 164 | |
| 165 | else: |
| 166 | if self.__socket_config is not None: |
| 167 | self._connect2socket() |
| 168 | else: |
| 169 | if self.dgram_socket: |
| 170 | sock_type = socket.SOCK_DGRAM |
| 171 | else: |
| 172 | sock_type = socket.SOCK_STREAM |
| 173 | for res in socket.getaddrinfo(self.host, self.port, socket.AF_UNSPEC, sock_type): |
| 174 | if self._connect2socket(res): |
| 175 | break |
| 176 | |
| 177 | except Exception as error: |
| 178 | self.error('unhandled exception during connect : {0}'.format(repr(error))) |
| 179 | self._sock = None |
| 180 | self.__socket_config = None |
| 181 | |
| 182 | def _disconnect(self): |
| 183 | """ |
| 184 | Close socket connection |
| 185 | :return: |
| 186 | """ |
| 187 | if self._sock is not None: |
| 188 | try: |
| 189 | self.debug('closing socket') |
| 190 | self._sock.shutdown(2) # 0 - read, 1 - write, 2 - all |
| 191 | self._sock.close() |
| 192 | except Exception as error: |
| 193 | if not (hasattr(error, 'errno') and error.errno == errno.ENOTCONN): |
| 194 | self.error(error) |
| 195 | self._sock = None |
| 196 | |
| 197 | def _send(self, request=None): |
| 198 | """ |
| 199 | Send request. |
| 200 | :return: boolean |
| 201 | """ |
| 202 | # Send request if it is needed |
| 203 | if self.request != self.__empty_request: |
| 204 | try: |
| 205 | self.debug('set socket write timeout to: {0}'.format(self._sock.gettimeout())) |
| 206 | self._sock.settimeout(self.write_timeout) |
| 207 | self.debug('sending request: {0}'.format(request or self.request)) |
| 208 | self._sock.send(request or self.request) |
| 209 | except Exception as error: |
| 210 | self._socket_error('error sending request: {0}'.format(error)) |
| 211 | self._disconnect() |
| 212 | return False |
| 213 | return True |
| 214 | |
| 215 | def _receive(self, raw=False): |
| 216 | """ |
| 217 | Receive data from socket |
| 218 | :param raw: set `True` to return bytes |
| 219 | :type raw: bool |
| 220 | :return: decoded str or raw bytes |
| 221 | :rtype: str/bytes |
| 222 | """ |
| 223 | data = "" if not raw else b"" |
| 224 | while True: |
| 225 | self.debug('receiving response') |
| 226 | try: |
| 227 | self.debug('set socket read timeout to: {0}'.format(self._sock.gettimeout())) |
| 228 | self._sock.settimeout(self.read_timeout) |
| 229 | buf = self._sock.recv(4096) |
| 230 | except Exception as error: |
| 231 | self._socket_error('failed to receive response: {0}'.format(error)) |
| 232 | self._disconnect() |
| 233 | break |
| 234 | |
| 235 | if buf is None or len(buf) == 0: # handle server disconnect |
| 236 | if data == "" or data == b"": |
| 237 | self._socket_error('unexpectedly disconnected') |
| 238 | else: |
| 239 | self.debug('server closed the connection') |
| 240 | self._disconnect() |
| 241 | break |
| 242 | |
| 243 | self.debug('received data') |
| 244 | data += buf.decode('utf-8', 'ignore') if not raw else buf |
| 245 | if self._check_raw_data(data): |
| 246 | break |
| 247 | |
| 248 | self.debug(u'final response: {0}'.format(data if not raw else u'binary data')) |
| 249 | return data |
| 250 | |
| 251 | def _get_raw_data(self, raw=False, request=None): |
| 252 | """ |
| 253 | Get raw data with low-level "socket" module. |
| 254 | :param raw: set `True` to return bytes |
| 255 | :type raw: bool |
| 256 | :return: decoded data (str) or raw data (bytes) |
| 257 | :rtype: str/bytes |
| 258 | """ |
| 259 | if self._sock is None: |
| 260 | self._connect() |
| 261 | if self._sock is None: |
| 262 | return None |
| 263 | |
| 264 | # Send request if it is needed |
| 265 | if not self._send(request): |
| 266 | return None |
| 267 | |
| 268 | data = self._receive(raw) |
| 269 | |
| 270 | if not self._keep_alive: |
| 271 | self._disconnect() |
| 272 | |
| 273 | return data |
| 274 | |
| 275 | @staticmethod |
| 276 | def _check_raw_data(data): |
| 277 | """ |
| 278 | Check if all data has been gathered from socket |
| 279 | :param data: str |
| 280 | :return: boolean |
| 281 | """ |
| 282 | return bool(data) |
| 283 | |
| 284 | def _parse_config(self): |
| 285 | """ |
| 286 | Parse configuration data |
| 287 | :return: boolean |
| 288 | """ |
| 289 | try: |
| 290 | self.unix_socket = str(self.configuration['socket']) |
| 291 | except (KeyError, TypeError): |
| 292 | self.debug('No unix socket specified. Trying TCP/IP socket.') |
| 293 | self.unix_socket = None |
| 294 | try: |
| 295 | self.host = str(self.configuration['host']) |
| 296 | except (KeyError, TypeError): |
| 297 | self.debug('No host specified. Using: "{0}"'.format(self.host)) |
| 298 | try: |
| 299 | self.port = int(self.configuration['port']) |
| 300 | except (KeyError, TypeError): |
| 301 | self.debug('No port specified. Using: "{0}"'.format(self.port)) |
| 302 | |
| 303 | self.tls = bool(self.configuration.get('tls', self.tls)) |
| 304 | if self.tls and not _TLS_SUPPORT: |
| 305 | self.warning('TLS requested but no TLS module found, disabling TLS support.') |
| 306 | self.tls = False |
| 307 | if _TLS_SUPPORT and not self.tls: |
| 308 | self.debug('No TLS preference specified, not using TLS.') |
| 309 | |
| 310 | if self.tls and _TLS_SUPPORT: |
| 311 | self.key = self.configuration.get('tls_key_file') |
| 312 | self.cert = self.configuration.get('tls_cert_file') |
| 313 | if not self.cert: |
| 314 | # If there's not a valid certificate, clear the key too. |
| 315 | self.debug('No valid TLS client certificate configuration found.') |
| 316 | self.key = None |
| 317 | self.cert = None |
| 318 | elif not self.key: |
| 319 | # If a key isn't listed, the config may still be |
| 320 | # valid, because there may be a key attached to the |
| 321 | # certificate. |
| 322 | self.info('No TLS client key specified, assuming it\'s attached to the certificate.') |
| 323 | self.key = None |
| 324 | |
| 325 | try: |
| 326 | self.request = str(self.configuration['request']) |
| 327 | except (KeyError, TypeError): |
| 328 | self.debug('No request specified. Using: "{0}"'.format(self.request)) |
| 329 | |
| 330 | self.request = self.request.encode() |
| 331 | |
| 332 | def check(self): |
| 333 | self._parse_config() |
| 334 | return SimpleService.check(self) |