Simplify check for negotiated protocol

negotiatedProtocol's type is Optional[bytes]

See https://github.com/twisted/twisted/blob/twisted-20.3.0/src/twisted/protocols/tls.py#L563-L587
and https://www.pyopenssl.org/en/20.0.1/api/ssl.html#OpenSSL.SSL.Connection.get_alpn_proto_negotiated

Note that OpenSSL.SSL.Connection.get_next_proto_negotiated is deprecated:
https://www.pyopenssl.org/en/20.0.0/changelog.html#backward-incompatible-changes
This commit is contained in:
Eugenio Lacuesta 2021-02-17 18:36:52 -03:00
parent e80f37bd3f
commit 4418f78941
No known key found for this signature in database
GPG Key ID: DA3EF2D0913E9810
1 changed files with 3 additions and 5 deletions

View File

@ -233,13 +233,11 @@ class H2ClientProtocol(Protocol, TimeoutMixin):
def handshakeCompleted(self) -> None:
"""We close the connection with InvalidNegotiatedProtocol exception
when the connection was not made via h2 protocol"""
negotiated_protocol = self.transport.negotiatedProtocol
if isinstance(negotiated_protocol, bytes):
negotiated_protocol = str(self.transport.negotiatedProtocol, 'utf-8')
if negotiated_protocol != 'h2':
protocol = self.transport.negotiatedProtocol
if protocol is not None and protocol != b"h2":
# Here we have not initiated the connection yet
# So, no need to send a GOAWAY frame to the remote
self._lose_connection_with_error([InvalidNegotiatedProtocol(negotiated_protocol)])
self._lose_connection_with_error([InvalidNegotiatedProtocol(protocol.decode("utf-8"))])
def _check_received_data(self, data: bytes) -> None:
"""Checks for edge cases where the connection to remote fails