[style] update python style to conform to PEP 8 (#3951)

With the exception of line length set to 119 vs. 79.

Add tests/ and tools/ to py-pretty-check.
This commit is contained in:
Jonathan Hui
2019-07-11 11:45:55 -04:00
committed by GitHub
parent a938ee2845
commit f924adcb60
341 changed files with 8702 additions and 4460 deletions
+172 -107
View File
@@ -27,11 +27,9 @@
# POSSIBILITY OF SUCH DAMAGE.
#
from binascii import hexlify
from enum import IntEnum
from functools import reduce
import io
import math
import struct
from ipv6 import BuildableFromBytes
@@ -93,9 +91,9 @@ class AlertDescription(IntEnum):
class Record(ConvertibleToBytes, BuildableFromBytes):
def __init__(self, content_type, version, epoch,
sequence_number, length, fragment):
def __init__(
self, content_type, version, epoch, sequence_number, length, fragment
):
self.content_type = content_type
self.version = version
self.epoch = epoch
@@ -104,12 +102,14 @@ class Record(ConvertibleToBytes, BuildableFromBytes):
self.fragment = fragment
def to_bytes(self):
return (struct.pack(">B", self.content_type) +
self.version.to_bytes() +
struct.pack(">H", self.epoch) +
self.sequence_number.to_bytes(6, byteorder='big') +
struct.pack(">H", self.length) +
self.fragment)
return (
struct.pack(">B", self.content_type)
+ self.version.to_bytes()
+ struct.pack(">H", self.epoch)
+ self.sequence_number.to_bytes(6, byteorder='big')
+ struct.pack(">H", self.length)
+ self.fragment
)
@classmethod
def from_bytes(cls, data):
@@ -119,15 +119,16 @@ class Record(ConvertibleToBytes, BuildableFromBytes):
sequence_number = struct.unpack(">Q", b'\x00\x00' + data.read(6))[0]
length = struct.unpack(">H", data.read(2))[0]
fragment = bytes(data.read(length))
return cls(content_type, version, epoch, sequence_number, length, fragment)
return cls(
content_type, version, epoch, sequence_number, length, fragment
)
def __repr__(self):
return "Record(content_type={}, version={}, epoch={}, sequence_number={}, length={})".format(
str(self.content_type), self.version, self.epoch, self.sequence_number, self.length)
str(self.content_type), self.version, self.epoch, self.sequence_number, self.length, )
class Message(ConvertibleToBytes, BuildableFromBytes):
def __init__(self, content_type):
self.content_type = content_type
@@ -140,9 +141,15 @@ class Message(ConvertibleToBytes, BuildableFromBytes):
class HandshakeMessage(Message):
def __init__(self, handshake_type, length, message_seq,
fragment_offset, fragment_length, body):
def __init__(
self,
handshake_type,
length,
message_seq,
fragment_offset,
fragment_length,
body,
):
super(HandshakeMessage, self).__init__(ContentType.HANDSHAKE)
self.handshake_type = handshake_type
self.length = length
@@ -152,12 +159,14 @@ class HandshakeMessage(Message):
self.body = body
def to_bytes(self):
return (struct.pack(">B", self.handshake_type) +
struct.pack(">I", self.length)[1:] +
struct.pack(">H", self.message_seq) +
struct.pack(">I", self.fragment_offset)[1:] +
struct.pack(">I", self.fragment_length)[1:] +
self.body.to_bytes())
return (
struct.pack(">B", self.handshake_type)
+ struct.pack(">I", self.length)[1:]
+ struct.pack(">H", self.message_seq)
+ struct.pack(">I", self.fragment_offset)[1:]
+ struct.pack(">I", self.fragment_length)[1:]
+ self.body.to_bytes()
)
@classmethod
def from_bytes(cls, data):
@@ -177,23 +186,32 @@ class HandshakeMessage(Message):
body = bytes(data.read(fragment_length))
assert data.tell() == end_position
return cls(handshake_type, length, message_seq,
fragment_offset, fragment_length, body)
return cls(
handshake_type,
length,
message_seq,
fragment_offset,
fragment_length,
body,
)
def __repr__(self):
return "Handshake(type={}, length={})".format(str(self.handshake_type), self.length)
return "Handshake(type={}, length={})".format(
str(self.handshake_type), self.length
)
class ProtocolVersion(ConvertibleToBytes, BuildableFromBytes):
def __init__(self, major, minor):
self.major = major
self.minor = minor
def __eq__(self, other):
return (type(self) == type(other) and
self.major == other.major and
self.minor == other.minor)
return (
isinstance(self, type(other))
and self.major == other.major
and self.minor == other.minor
)
def to_bytes(self):
return struct.pack(">BB", self.major, self.minor)
@@ -204,7 +222,9 @@ class ProtocolVersion(ConvertibleToBytes, BuildableFromBytes):
return cls(major, minor)
def __repr__(self):
return "ProtocolVersion(major={}, minor={})".format(self.major, self.minor)
return "ProtocolVersion(major={}, minor={})".format(
self.major, self.minor
)
class Random(ConvertibleToBytes, BuildableFromBytes):
@@ -217,9 +237,11 @@ class Random(ConvertibleToBytes, BuildableFromBytes):
assert len(self.random_bytes) == Random.random_bytes_length
def __eq__(self, other):
return (type(self) == type(other) and
self.gmt_unix_time == other.gmt_unix_time and
self.random_bytes == other.random_bytes)
return (
isinstance(self, type(other))
and self.gmt_unix_time == other.gmt_unix_time
and self.random_bytes == other.random_bytes
)
def to_bytes(self):
return struct.pack(">I", self.gmt_unix_time) + (self.random_bytes)
@@ -232,7 +254,6 @@ class Random(ConvertibleToBytes, BuildableFromBytes):
class VariableVector(ConvertibleToBytes):
def __init__(self, subrange, ele_cls, elements):
self.subrange = subrange
self.ele_cls = ele_cls
@@ -243,10 +264,12 @@ class VariableVector(ConvertibleToBytes):
return len(self.elements)
def __eq__(self, other):
return (type(self) == type(other) and
self.subrange == other.subrange and
self.ele_cls == other.ele_cls and
self.elements == other.elements)
return (
isinstance(self, type(other))
and self.subrange == other.subrange
and self.ele_cls == other.ele_cls
and self.elements == other.elements
)
def to_bytes(self):
data = reduce(lambda ele, acc: acc + ele.to_bytes(), self.elements)
@@ -264,7 +287,11 @@ class VariableVector(ConvertibleToBytes):
@classmethod
def _decode_length(cls, subrange, data):
length_in_byte = cls._calc_length_in_byte(subrange[1])
return reduce(lambda acc, byte: (acc << 8) | byte, bytearray(data.read(length_in_byte)), 0)
return reduce(
lambda acc, byte: (acc << 8) | byte,
bytearray(data.read(length_in_byte)),
0,
)
@classmethod
def _encode_length(cls, length, subrange):
@@ -281,12 +308,11 @@ class VariableVector(ConvertibleToBytes):
class Opaque(ConvertibleToBytes, BuildableFromBytes):
def __init__(self, byte):
self.byte = byte
def __eq__(self, other):
return type(self) == type(other) and self.byte == other.byte
return isinstance(self, type(other)) and self.byte == other.byte
def to_bytes(self):
return struct.pack(">B", self.byte)
@@ -297,12 +323,11 @@ class Opaque(ConvertibleToBytes, BuildableFromBytes):
class CipherSuite(ConvertibleToBytes, BuildableFromBytes):
def __init__(self, cipher):
self.cipher = cipher
def __eq__(self, other):
return type(self) == type(other) and self.cipher == other.cipher
return isinstance(self, type(other)) and self.cipher == other.cipher
def to_bytes(self):
return struct.pack(">BB", self.cipher[0], self.cipher[1])
@@ -323,7 +348,7 @@ class CompressionMethod(ConvertibleToBytes, BuildableFromBytes):
pass
def __eq__(self, other):
return type(self) == type(other)
return isinstance(self, type(other))
def to_bytes(self):
return struct.pack(">B", CompressionMethod.NULL)
@@ -336,31 +361,43 @@ class CompressionMethod(ConvertibleToBytes, BuildableFromBytes):
class Extension(ConvertibleToBytes, BuildableFromBytes):
def __init__(self, extension_type, extension_data):
self.extension_type = extension_type
self.extension_data = extension_data
def __eq__(self, other):
return (type(self) == type(other) and
self.extension_type == other.extension_type and
self.extension_data == other.extension_data)
return (
isinstance(self, type(other))
and self.extension_type == other.extension_type
and self.extension_data == other.extension_data
)
def to_bytes(self):
return (struct.pack(">H", self.extension_type) +
self.extension_data.to_bytes())
return (
struct.pack(">H", self.extension_type)
+ self.extension_data.to_bytes()
)
@classmethod
def from_bytes(cls, data):
extension_type = struct.unpack(">H", data.read(2))[0]
extension_data = VariableVector.from_bytes(Opaque, (0, 2**16 - 1), data)
extension_data = VariableVector.from_bytes(
Opaque, (0, 2 ** 16 - 1), data
)
return cls(extension_type, extension_data)
class ClientHello(HandshakeMessage):
def __init__(self, client_version, random, session_id,
cookie, cipher_suites, compression_methods, extensions):
def __init__(
self,
client_version,
random,
session_id,
cookie,
cipher_suites,
compression_methods,
extensions,
):
self.client_version = client_version
self.random = random
self.session_id = session_id
@@ -370,31 +407,45 @@ class ClientHello(HandshakeMessage):
self.extensions = extensions
def to_bytes(self):
return (self.client_version.to_bytes() +
self.random.to_bytes() +
self.session_id.to_bytes() +
self.cookie.to_bytes() +
self.cipher_suites.to_bytes() +
self.compression_methods.to_bytes() +
self.extensions.to_bytes())
return (
self.client_version.to_bytes()
+ self.random.to_bytes()
+ self.session_id.to_bytes()
+ self.cookie.to_bytes()
+ self.cipher_suites.to_bytes()
+ self.compression_methods.to_bytes()
+ self.extensions.to_bytes()
)
@classmethod
def from_bytes(cls, data):
client_version = ProtocolVersion.from_bytes(data)
random = Random.from_bytes(data)
session_id = VariableVector.from_bytes(Opaque, (0, 32), data)
cookie = VariableVector.from_bytes(Opaque, (0, 2**8 - 1), data)
cipher_suites = VariableVector.from_bytes(CipherSuite, (2, 2**16 - 1), data)
compression_methods = VariableVector.from_bytes(CompressionMethod, (1, 2**8 - 1), data)
cookie = VariableVector.from_bytes(Opaque, (0, 2 ** 8 - 1), data)
cipher_suites = VariableVector.from_bytes(
CipherSuite, (2, 2 ** 16 - 1), data
)
compression_methods = VariableVector.from_bytes(
CompressionMethod, (1, 2 ** 8 - 1), data
)
extensions = None
if data.tell() < len(data.getvalue()):
extensions = VariableVector.from_bytes(Extension, (0, 2**16 - 1), data)
return cls(client_version, random, session_id,
cookie, cipher_suites, compression_methods, extensions)
extensions = VariableVector.from_bytes(
Extension, (0, 2 ** 16 - 1), data
)
return cls(
client_version,
random,
session_id,
cookie,
cipher_suites,
compression_methods,
extensions,
)
class HelloVerifyRequest(HandshakeMessage):
def __init__(self, server_version, cookie):
self.server_version = server_version
self.cookie = cookie
@@ -405,14 +456,20 @@ class HelloVerifyRequest(HandshakeMessage):
@classmethod
def from_bytes(cls, data):
server_version = ProtocolVersion.from_bytes(data)
cookie = VariableVector.from_bytes(Opaque, (0, 2**8 - 1), data)
cookie = VariableVector.from_bytes(Opaque, (0, 2 ** 8 - 1), data)
return cls(server_version, cookie)
class ServerHello(HandshakeMessage):
def __init__(self, server_version, random, session_id,
cipher_suite, compression_method, extensions):
def __init__(
self,
server_version,
random,
session_id,
cipher_suite,
compression_method,
extensions,
):
self.server_version = server_version
self.random = random
self.session_id = session_id
@@ -421,12 +478,14 @@ class ServerHello(HandshakeMessage):
self.extensions = extensions
def to_bytes(self):
return (self.server_version.to_bytes() +
self.random.to_bytes() +
self.session_id.to_bytes() +
self.cipher_suite.to_bytes() +
self.compression_method.to_bytes() +
self.extensions.to_bytes())
return (
self.server_version.to_bytes()
+ self.random.to_bytes()
+ self.session_id.to_bytes()
+ self.cipher_suite.to_bytes()
+ self.compression_method.to_bytes()
+ self.extensions.to_bytes()
)
@classmethod
def from_bytes(cls, data):
@@ -437,13 +496,20 @@ class ServerHello(HandshakeMessage):
compression_method = CompressionMethod.from_bytes(data)
extensions = None
if data.tell() < len(data.getvalue()):
extensions = VariableVector.from_bytes(Extension, (0, 2**16 - 1), data)
return cls(server_version, random, session_id,
cipher_suite, compression_method, extensions)
extensions = VariableVector.from_bytes(
Extension, (0, 2 ** 16 - 1), data
)
return cls(
server_version,
random,
session_id,
cipher_suite,
compression_method,
extensions,
)
class ServerHelloDone(HandshakeMessage):
def __init__(self):
pass
@@ -456,49 +522,41 @@ class ServerHelloDone(HandshakeMessage):
class HelloRequest(HandshakeMessage):
def __init__(self):
raise NotImplementedError
class Certificate(HandshakeMessage):
def __init__(self):
raise NotImplementedError
class ServerKeyExchange(HandshakeMessage):
def __init__(self):
raise NotImplementedError
class CertificateRequest(HandshakeMessage):
def __init__(self):
raise NotImplementedError
class CertificateVerify(HandshakeMessage):
def __init__(self):
raise NotImplementedError
class ClientKeyExchange(HandshakeMessage):
def __init__(self):
raise NotImplementedError
class Finished(HandshakeMessage):
def __init__(self, verify_data):
raise NotImplementedError
class AlertMessage(Message):
def __init__(self, level, description):
super(AlertMessage, self).__init__(ContentType.ALERT)
self.level = level
@@ -512,20 +570,22 @@ class AlertMessage(Message):
level, description = struct.unpack(">BB", data.read(2))
try:
return cls(AlertLevel(level), AlertDescription(description))
except:
except BaseException:
data.read()
# An AlertMessage could be encrypted and we can't parsing it.
return cls(None, None)
def __repr__(self):
return "Alert(level={}, description={})".format(
str(self.level), str(self.description))
str(self.level), str(self.description)
)
class ChangeCipherSpecMessage(Message):
def __init__(self):
super(ChangeCipherSpecMessage, self).__init__(ContentType.CHANGE_CIPHER_SPEC)
super(ChangeCipherSpecMessage, self).__init__(
ContentType.CHANGE_CIPHER_SPEC
)
def to_bytes(self):
return struct.pack(">B", 1)
@@ -540,9 +600,10 @@ class ChangeCipherSpecMessage(Message):
class ApplicationDataMessage(Message):
def __init__(self, raw):
super(ApplicationDataMessage, self).__init__(ContentType.APPLICATION_DATA)
super(ApplicationDataMessage, self).__init__(
ContentType.APPLICATION_DATA
)
self.raw = raw
self.body = None
@@ -568,13 +629,13 @@ handshake_map = {
HandshakeType.CLIENT_HELLO: ClientHello,
HandshakeType.SERVER_HELLO: ServerHello,
HandshakeType.HELLO_VERIFY_REQUEST: HelloVerifyRequest,
HandshakeType.CERTIFICATE: None, # Certificate
HandshakeType.SERVER_KEY_EXCHANGE: None, # ServerKeyExchange
HandshakeType.CERTIFICATE_REQUEST: None, # CertificateRequest
HandshakeType.CERTIFICATE: None, # Certificate
HandshakeType.SERVER_KEY_EXCHANGE: None, # ServerKeyExchange
HandshakeType.CERTIFICATE_REQUEST: None, # CertificateRequest
HandshakeType.SERVER_HELLO_DONE: ServerHelloDone,
HandshakeType.CERTIFICATE_VERIFY: None, # CertificateVerify
HandshakeType.CLIENT_KEY_EXCHANGE: None, # ClientKeyExchange
HandshakeType.FINISHED: None, # Finished
HandshakeType.CERTIFICATE_VERIFY: None, # CertificateVerify
HandshakeType.CLIENT_KEY_EXCHANGE: None, # ClientKeyExchange
HandshakeType.FINISHED: None, # Finished
}
@@ -582,7 +643,7 @@ content_map = {
ContentType.CHANGE_CIPHER_SPEC: ChangeCipherSpecMessage,
ContentType.ALERT: AlertMessage,
ContentType.HANDSHAKE: HandshakeMessage,
ContentType.APPLICATION_DATA: ApplicationDataMessage
ContentType.APPLICATION_DATA: ApplicationDataMessage,
}
@@ -600,11 +661,15 @@ class MessageFactory(object):
while data.tell() < len(data.getvalue()):
record = Record.from_bytes(data)
if record.version.major != 0xfe or record.version.minor != 0xfd:
if record.version.major != 0xfe or record.version.minor != 0xFD:
raise ValueError("DTLS version error, expect DTLSv1.2")
last_msg_is_change_cipher_spec = type(self).last_msg_is_change_cipher_spec
type(self).last_msg_is_change_cipher_spec = (record.content_type == ContentType.CHANGE_CIPHER_SPEC)
last_msg_is_change_cipher_spec = type(
self
).last_msg_is_change_cipher_spec
type(self).last_msg_is_change_cipher_spec = (
record.content_type == ContentType.CHANGE_CIPHER_SPEC
)
# FINISHED message immediately follows CHANGE_CIPHER_SPEC message
# We skip FINISHED message as it is encrypted