[style] apply google python style guide (#4501)

This commit applies and enforces Google's python style for tests.
This commit is contained in:
Yakun Xu
2020-02-04 10:27:50 -08:00
committed by GitHub
parent 8368d440dd
commit 33808ebfba
501 changed files with 6674 additions and 5515 deletions
+86 -101
View File
@@ -91,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
@@ -102,14 +102,10 @@ 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,16 +115,21 @@ 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
@@ -141,6 +142,7 @@ class Message(ConvertibleToBytes, BuildableFromBytes):
class HandshakeMessage(Message):
def __init__(
self,
handshake_type,
@@ -159,14 +161,12 @@ 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):
@@ -196,22 +196,19 @@ class HandshakeMessage(Message):
)
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 (
isinstance(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)
@@ -223,8 +220,7 @@ class ProtocolVersion(ConvertibleToBytes, BuildableFromBytes):
def __repr__(self):
return "ProtocolVersion(major={}, minor={})".format(
self.major, self.minor
)
self.major, self.minor)
class Random(ConvertibleToBytes, BuildableFromBytes):
@@ -237,11 +233,9 @@ class Random(ConvertibleToBytes, BuildableFromBytes):
assert len(self.random_bytes) == Random.random_bytes_length
def __eq__(self, other):
return (
isinstance(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)
@@ -254,6 +248,7 @@ class Random(ConvertibleToBytes, BuildableFromBytes):
class VariableVector(ConvertibleToBytes):
def __init__(self, subrange, ele_cls, elements):
self.subrange = subrange
self.ele_cls = ele_cls
@@ -264,12 +259,10 @@ class VariableVector(ConvertibleToBytes):
return len(self.elements)
def __eq__(self, other):
return (
isinstance(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)
@@ -308,6 +301,7 @@ class VariableVector(ConvertibleToBytes):
class Opaque(ConvertibleToBytes, BuildableFromBytes):
def __init__(self, byte):
self.byte = byte
@@ -323,6 +317,7 @@ class Opaque(ConvertibleToBytes, BuildableFromBytes):
class CipherSuite(ConvertibleToBytes, BuildableFromBytes):
def __init__(self, cipher):
self.cipher = cipher
@@ -361,33 +356,29 @@ 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 (
isinstance(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,
@@ -407,33 +398,26 @@ 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
)
extensions = VariableVector.from_bytes(Extension, (0, 2**16 - 1),
data)
return cls(
client_version,
random,
@@ -446,6 +430,7 @@ class ClientHello(HandshakeMessage):
class HelloVerifyRequest(HandshakeMessage):
def __init__(self, server_version, cookie):
self.server_version = server_version
self.cookie = cookie
@@ -456,11 +441,12 @@ 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,
@@ -478,14 +464,9 @@ 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):
@@ -496,9 +477,8 @@ 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
)
extensions = VariableVector.from_bytes(Extension, (0, 2**16 - 1),
data)
return cls(
server_version,
random,
@@ -510,6 +490,7 @@ class ServerHello(HandshakeMessage):
class ServerHelloDone(HandshakeMessage):
def __init__(self):
pass
@@ -522,41 +503,49 @@ 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
@@ -576,16 +565,15 @@ class AlertMessage(Message):
return cls(None, None)
def __repr__(self):
return "Alert(level={}, description={})".format(
str(self.level), str(self.description)
)
return "Alert(level={}, description={})".format(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)
@@ -600,10 +588,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
@@ -638,7 +626,6 @@ handshake_map = {
HandshakeType.FINISHED: None, # Finished
}
content_map = {
ContentType.CHANGE_CIPHER_SPEC: ChangeCipherSpecMessage,
ContentType.ALERT: AlertMessage,
@@ -665,11 +652,9 @@ class MessageFactory(object):
raise ValueError("DTLS version error, expect DTLSv1.2")
last_msg_is_change_cipher_spec = type(
self
).last_msg_is_change_cipher_spec
self).last_msg_is_change_cipher_spec
type(self).last_msg_is_change_cipher_spec = (
record.content_type == ContentType.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