From 3cb47bb3627da5aedd0add3edeaee93f46b56b47 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 25 Dec 2024 09:09:45 +0100 Subject: [PATCH 01/39] :arrow_up: modernizing the test suite by migrating to pytest + nox --- .appveyor.yml | 12 - .gitignore | 1 + .pre-commit-config.yaml | 2 +- dev-requirements.txt | 4 +- noxfile.py | 136 ++++ pyproject.toml | 14 + tests/test_asyncio.py | 111 ++- tests/test_buffer.py | 146 ++-- tests/test_connection.py | 1400 +++++++++++++++------------------- tests/test_crypto_v1.py | 120 ++- tests/test_crypto_v2.py | 96 ++- tests/test_h3.py | 482 +++++------- tests/test_logger.py | 26 +- tests/test_packet.py | 278 ++++--- tests/test_packet_builder.py | 290 ++++--- tests/test_rangeset.py | 140 ++-- tests/test_recovery.py | 164 ++-- tests/test_retry.py | 24 +- tests/test_stream.py | 676 ++++++++-------- tests/test_tls.py | 526 ++++++------- tests/test_webtransport.py | 66 +- 21 files changed, 2171 insertions(+), 2543 deletions(-) delete mode 100644 .appveyor.yml create mode 100644 noxfile.py diff --git a/.appveyor.yml b/.appveyor.yml deleted file mode 100644 index d0edee187..000000000 --- a/.appveyor.yml +++ /dev/null @@ -1,12 +0,0 @@ -environment: - CIBW_SKIP: cp27-* cp33-* cp34-* - CIBW_TEST_COMMAND: python -m unittest discover -s {project}/tests -install: - - cmd: C:\Python36-x64\python.exe -m pip install cibuildwheel -build_script: - - cmd: C:\Python36-x64\python.exe -m cibuildwheel --output-dir wheelhouse - - ps: >- - if ($env:APPVEYOR_REPO_TAG -eq "true") { - Invoke-Expression "python -m pip install twine" - Invoke-Expression "python -m twine upload --skip-existing wheelhouse/*.whl" - } diff --git a/.gitignore b/.gitignore index cbca41691..db9e22d4c 100644 --- a/.gitignore +++ b/.gitignore @@ -20,3 +20,4 @@ ENV/ env.bak/ venv.bak/ target/ +.nox diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 34f02b546..da97c7fb4 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -22,7 +22,7 @@ repos: # Run the formatter. - id: ruff-format - repo: https://github.com/pre-commit/mirrors-mypy - rev: v1.11.0 + rev: v1.14.0 hooks: - id: mypy args: [--check-untyped-defs] diff --git a/dev-requirements.txt b/dev-requirements.txt index 014a91345..3266a1248 100644 --- a/dev-requirements.txt +++ b/dev-requirements.txt @@ -1,2 +1,4 @@ coverage[toml]>=7.2.7,<8 -cryptography>=42,<43 +cryptography>=42,<44 +pytest>=7.4.4,<9 +pytest-asyncio>=0.21.1,<=0.24.0 diff --git a/noxfile.py b/noxfile.py new file mode 100644 index 000000000..d633a92cf --- /dev/null +++ b/noxfile.py @@ -0,0 +1,136 @@ +from __future__ import annotations + +import os +import shutil + +import nox + + +def tests_impl( + session: nox.Session, + tracemalloc_enable: bool = False, +) -> None: + # Install deps and the package itself. + session.install("-U", "pip", "setuptools", silent=False) + session.install("-r", "dev-requirements.txt", silent=False) + + session.install(f".", silent=False) + + # Show the pip version. + session.run("pip", "--version") + # Print the Python version and bytesize. + session.run("python", "--version") + session.run("python", "-c", "import struct; print(struct.calcsize('P') * 8)") + + # Inspired from https://hynek.me/articles/ditch-codecov-python/ + # We use parallel mode and then combine in a later CI step + session.run( + "python", + "-m", + *( + ( + "coverage", + "run", + "--parallel-mode", + "-m", + ) + if tracemalloc_enable is False + else () + ), + "pytest", + "-v", + "-ra", + f"--color={'yes' if 'GITHUB_ACTIONS' in os.environ else 'auto'}", + "--tb=native", + "--durations=10", + "--strict-config", + "--strict-markers", + *(session.posargs or ("tests/",)), + env={ + "PYTHONWARNINGS": "always::DeprecationWarning", + "COVERAGE_CORE": "sysmon", + "PYTHONTRACEMALLOC": "25" if tracemalloc_enable else "", + }, + ) + + +@nox.session( + python=["3.7", "3.8", "3.9", "3.10", "3.11", "3.12", "3.13", "3.14", "pypy"] +) +def test(session: nox.Session) -> None: + tests_impl(session) + + +@nox.session(python=["3.7", "3.8", "3.9", "3.10", "3.11", "3.12", "3.13", "3.14"]) +def tracemalloc(session: nox.Session) -> None: + tests_impl(session, tracemalloc_enable=True) + + +def git_clone(session: nox.Session, git_url: str) -> None: + """We either clone the target repository or if already exist + simply reset the state and pull. + """ + expected_directory = git_url.split("/")[-1] + + if expected_directory.endswith(".git"): + expected_directory = expected_directory[:-4] + + if not os.path.isdir(expected_directory): + session.run("git", "clone", "--depth", "1", git_url, external=True) + else: + session.run( + "git", "-C", expected_directory, "reset", "--hard", "HEAD", external=True + ) + session.run("git", "-C", expected_directory, "pull", external=True) + + +@nox.session() +def downstream_niquests(session: nox.Session) -> None: + root = os.getcwd() + tmp_dir = session.create_tmp() + + session.cd(tmp_dir) + git_clone(session, "https://github.com/jawah/niquests") + session.chdir("niquests") + + session.run("git", "rev-parse", "HEAD", external=True) + session.install(".[socks]", silent=False) + session.install("-r", "requirements-dev.txt", silent=False) + + session.cd(root) + session.install(".", silent=False) + session.cd(f"{tmp_dir}/niquests") + + session.run("python", "-c", "import qh3; print(qh3.__version__)") + session.run( + "python", + "-m", + "pytest", + "-v", + f"--color={'yes' if 'GITHUB_ACTIONS' in os.environ else 'auto'}", + *(session.posargs or ("tests/",)), + env={"NIQUESTS_STRICT_OCSP": "1"}, + ) + + +@nox.session() +def format(session: nox.Session) -> None: + """Run code formatters.""" + lint(session) + + +@nox.session +def lint(session: nox.Session) -> None: + session.install("pre-commit") + session.run("pre-commit", "run", "--all-files") + + +@nox.session +def docs(session: nox.Session) -> None: + session.install("-r", "docs/docs-requirements.txt") + session.install(".") + + session.chdir("docs") + if os.path.exists("_build"): + shutil.rmtree("_build") + session.run("sphinx-build", "-b", "html", "-W", ".", "_build/html") diff --git a/pyproject.toml b/pyproject.toml index 4eea8641e..3e5dba63d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -48,6 +48,20 @@ documentation = "https://qh3.readthedocs.io/" [tool.coverage.run] source = ["qh3"] +[tool.pytest.ini_options] +xfail_strict = true +log_level = "DEBUG" +filterwarnings = [ + "error", + '''ignore:.*iscoroutinefunction.*:DeprecationWarning''', + '''default:unclosed .*:ResourceWarning''', + '''ignore:The event_loop fixture provided by:DeprecationWarning''', + '''ignore:A plugin raised an exception during''', + '''ignore:Exception ignored in:pytest.PytestUnraisableExceptionWarning''', + '''ignore:Exception in thread:pytest.PytestUnhandledThreadExceptionWarning''', + '''ignore:loop is closed:ResourceWarning''', +] + [tool.mypy] disallow_untyped_calls = true disallow_untyped_decorators = true diff --git a/tests/test_asyncio.py b/tests/test_asyncio.py index a34e999ca..9b8f59ea9 100644 --- a/tests/test_asyncio.py +++ b/tests/test_asyncio.py @@ -1,11 +1,11 @@ from __future__ import annotations +import pytest import asyncio import binascii import contextlib import random import socket -from unittest import TestCase, skipIf from unittest.mock import patch from cryptography.hazmat.primitives import serialization @@ -24,7 +24,6 @@ SERVER_COMBINEDFILE, SERVER_KEYFILE, SKIP_TESTS, - asynctest, generate_ec_certificate, generate_ed25519_certificate, ) @@ -60,8 +59,8 @@ async def serve(): asyncio.ensure_future(serve()) -class HighLevelTest(TestCase): - def setUp(self): +class TestHighLevel: + def setup_method(self): self.bogus_port = 1024 self.server_host = "localhost" @@ -86,8 +85,8 @@ async def run_client( await client.wait_connected() reader, writer = await client.create_stream() - self.assertEqual(writer.can_write_eof(), True) - self.assertEqual(writer.get_extra_info("stream_id"), 0) + assert writer.can_write_eof() == True + assert writer.get_extra_info("stream_id") == 0 writer.write(request) writer.write_eof() @@ -116,24 +115,24 @@ async def run_server(self, configuration=None, host="::", **kwargs): finally: server.close() - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve(self): async with self.run_server() as server_port: response = await self.run_client(port=server_port) - self.assertEqual(response, b"gnip") + assert response == b"gnip" - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_ipv4(self): async with self.run_server(host="0.0.0.0") as server_port: response = await self.run_client(host="127.0.0.1", port=server_port) - self.assertEqual(response, b"gnip") + assert response == b"gnip" - @skipIf("ipv6" in SKIP_TESTS, "Skipping IPv6 tests") - @asynctest + @pytest.mark.skipif("ipv6" in SKIP_TESTS, reason="Skipping IPv6 tests") + @pytest.mark.asyncio async def test_connect_and_serve_ipv6(self): async with self.run_server(host="::") as server_port: response = await self.run_client(host="::1", port=server_port) - self.assertEqual(response, b"gnip") + assert response == b"gnip" async def _test_connect_and_serve_with_certificate(self, certificate, private_key): inner_certificate = InnerCertificate( @@ -170,9 +169,9 @@ async def _test_connect_and_serve_with_certificate(self, certificate, private_ke cafile=None, port=server_port, ) - self.assertEqual(response, b"gnip") + assert response == b"gnip" - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_with_ec_certificate(self): await self._test_connect_and_serve_with_certificate( *generate_ec_certificate( @@ -180,7 +179,7 @@ async def test_connect_and_serve_with_ec_certificate(self): ) ) - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_with_ed25519_certificate(self): await self._test_connect_and_serve_with_certificate( *generate_ed25519_certificate( @@ -188,7 +187,7 @@ async def test_connect_and_serve_with_ed25519_certificate(self): ) ) - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_large(self): """ Transfer enough data to require raising MAX_DATA and MAX_STREAM_DATA. @@ -196,16 +195,16 @@ async def test_connect_and_serve_large(self): data = b"Z" * 2097152 async with self.run_server() as server_port: response = await self.run_client(port=server_port, request=data) - self.assertEqual(response, data) + assert response == data - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_without_client_configuration(self): async with self.run_server() as server_port: - with self.assertRaises(ConnectionError): + with pytest.raises(ConnectionError): async with connect(self.server_host, server_port) as client: await client.ping() - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_writelines(self): async with self.run_server() as server_port: configuration = QuicConfiguration(is_client=True) @@ -220,11 +219,11 @@ async def test_connect_and_serve_writelines(self): writer.write_eof() response = await reader.read() - self.assertEqual(response, b"5432109876543210") + assert response == b"5432109876543210" - @skipIf("loss" in SKIP_TESTS, "Skipping loss tests") + @pytest.mark.skipif("loss" in SKIP_TESTS, reason="Skipping loss tests") @patch("socket.socket.sendto", new_callable=lambda: sendto_with_loss) - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_with_packet_loss(self, mock_sendto): """ This test ensures handshake success and stream data is successfully sent @@ -244,9 +243,9 @@ async def test_connect_and_serve_with_packet_loss(self, mock_sendto): port=server_port, request=data, ) - self.assertEqual(response, data) + assert response == data - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_with_session_ticket(self): client_ticket = None store = SessionTicketStore() @@ -262,9 +261,9 @@ def save_ticket(t): response = await self.run_client( port=server_port, session_ticket_handler=save_ticket ) - self.assertEqual(response, b"gnip") + assert response == b"gnip" - self.assertIsNotNone(client_ticket) + assert client_ticket is not None # second request response = await self.run_client( @@ -273,15 +272,15 @@ def save_ticket(t): ), port=server_port, ) - self.assertEqual(response, b"gnip") + assert response == b"gnip" - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_with_retry(self): async with self.run_server(retry=True) as server_port: response = await self.run_client(port=server_port) - self.assertEqual(response, b"gnip") + assert response == b"gnip" - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_with_retry_bad_original_destination_connection_id( self, ): @@ -298,10 +297,10 @@ def create_protocol(*args, **kwargs): async with self.run_server( create_protocol=create_protocol, retry=True ) as server_port: - with self.assertRaises(ConnectionError): + with pytest.raises(ConnectionError): await self.run_client(port=server_port) - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_with_retry_bad_retry_source_connection_id(self): """ If the server's transport parameters do not have the correct @@ -316,22 +315,22 @@ def create_protocol(*args, **kwargs): async with self.run_server( create_protocol=create_protocol, retry=True ) as server_port: - with self.assertRaises(ConnectionError): + with pytest.raises(ConnectionError): await self.run_client(port=server_port) @patch("qh3.quic.retry.QuicRetryTokenHandler.validate_token") - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_with_retry_bad_token(self, mock_validate): mock_validate.side_effect = ValueError("Decryption failed.") async with self.run_server(retry=True) as server_port: - with self.assertRaises(ConnectionError): + with pytest.raises(ConnectionError): await self.run_client( configuration=QuicConfiguration(is_client=True, idle_timeout=4.0), port=server_port, ) - @asynctest + @pytest.mark.asyncio async def test_connect_and_serve_with_version_negotiation(self): async with self.run_server() as server_port: # force version negotiation @@ -341,19 +340,19 @@ async def test_connect_and_serve_with_version_negotiation(self): response = await self.run_client( configuration=configuration, port=server_port ) - self.assertEqual(response, b"gnip") + assert response == b"gnip" - @asynctest + @pytest.mark.asyncio async def test_connect_timeout(self): - with self.assertRaises(ConnectionError): + with pytest.raises(ConnectionError): await self.run_client( port=self.bogus_port, configuration=QuicConfiguration(is_client=True, idle_timeout=5), ) - @asynctest + @pytest.mark.asyncio async def test_connect_timeout_no_wait_connected(self): - with self.assertRaises(ConnectionError): + with pytest.raises(ConnectionError): configuration = QuicConfiguration(is_client=True, idle_timeout=5) configuration.load_verify_locations(cafile=SERVER_CACERTFILE) async with connect( @@ -364,18 +363,18 @@ async def test_connect_timeout_no_wait_connected(self): ) as client: await client.ping() - @asynctest + @pytest.mark.asyncio async def test_connect_local_port(self): async with self.run_server() as server_port: response = await self.run_client(local_port=3456, port=server_port) - self.assertEqual(response, b"gnip") + assert response == b"gnip" - @asynctest + @pytest.mark.asyncio async def test_connect_local_port_bind(self): - with self.assertRaises(OverflowError): + with pytest.raises(OverflowError): await self.run_client(local_port=-1, port=self.bogus_port) - @asynctest + @pytest.mark.asyncio async def test_change_connection_id(self): async with self.run_server() as server_port: configuration = QuicConfiguration(is_client=True) @@ -387,7 +386,7 @@ async def test_change_connection_id(self): client.change_connection_id() await client.ping() - @asynctest + @pytest.mark.asyncio async def test_key_update(self): async with self.run_server() as server_port: configuration = QuicConfiguration(is_client=True) @@ -399,7 +398,7 @@ async def test_key_update(self): client.request_key_update() await client.ping() - @asynctest + @pytest.mark.asyncio async def test_ping(self): async with self.run_server() as server_port: configuration = QuicConfiguration(is_client=True) @@ -410,7 +409,7 @@ async def test_ping(self): await client.ping() await client.ping() - @asynctest + @pytest.mark.asyncio async def test_ping_parallel(self): async with self.run_server() as server_port: configuration = QuicConfiguration(is_client=True) @@ -421,7 +420,7 @@ async def test_ping_parallel(self): coros = [client.ping() for x in range(16)] await asyncio.gather(*coros) - @asynctest + @pytest.mark.asyncio async def test_server_receives_garbage(self): configuration = QuicConfiguration(is_client=False) configuration.load_cert_chain(SERVER_CERTFILE, SERVER_KEYFILE) @@ -433,7 +432,7 @@ async def test_server_receives_garbage(self): server.datagram_received(binascii.unhexlify("c00000000080"), ("1.2.3.4", 1234)) server.close() - @asynctest + @pytest.mark.asyncio async def test_combined_key(self): config1 = QuicConfiguration() config2 = QuicConfiguration() @@ -448,6 +447,6 @@ async def test_combined_key(self): open(SERVER_CERTFILE, "rb").read(), open(SERVER_KEYFILE, "rb").read() ) - self.assertEqual(config1.certificate, config2.certificate) - self.assertEqual(config1.certificate, config3.certificate) - self.assertEqual(config1.certificate, config4.certificate) + assert config1.certificate == config2.certificate + assert config1.certificate == config3.certificate + assert config1.certificate == config4.certificate diff --git a/tests/test_buffer.py b/tests/test_buffer.py index c38b07574..e88580128 100644 --- a/tests/test_buffer.py +++ b/tests/test_buffer.py @@ -1,175 +1,175 @@ from __future__ import annotations -from unittest import TestCase +import pytest from qh3.buffer import Buffer, BufferReadError, BufferWriteError, size_uint_var -class BufferTest(TestCase): +class TestBuffer: def test_data_slice(self): buf = Buffer(data=b"\x08\x07\x06\x05\x04\x03\x02\x01") - self.assertEqual(buf.data_slice(0, 8), b"\x08\x07\x06\x05\x04\x03\x02\x01") - self.assertEqual(buf.data_slice(1, 3), b"\x07\x06") + assert buf.data_slice(0, 8) == b"\x08\x07\x06\x05\x04\x03\x02\x01" + assert buf.data_slice(1, 3) == b"\x07\x06" - with self.assertRaises(OverflowError): + with pytest.raises(OverflowError): buf.data_slice(-1, 3) - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): buf.data_slice(0, 9) - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): buf.data_slice(1, 0) def test_pull_bytes(self): buf = Buffer(data=b"\x08\x07\x06\x05\x04\x03\x02\x01") - self.assertEqual(buf.pull_bytes(3), b"\x08\x07\x06") + assert buf.pull_bytes(3) == b"\x08\x07\x06" def test_internal_fixed_size(self): buf = Buffer(8) buf.push_bytes(b"foobar") # push 6 bytes, 2 left free bytes - self.assertEqual(buf.data, b"foobar") + assert buf.data == b"foobar" buf.seek(8) # setting cursor to the end of buf capacity - self.assertEqual(buf.data, b"foobar\x00\x00") # the two NULL bytes should be there + assert buf.data == b"foobar\x00\x00"# the two NULL bytes should be there def test_internal_push_zero_bytes(self): buf = Buffer(6) buf.push_bytes(b"foobar") # push 6 bytes, 0 left free bytes - self.assertEqual(buf.data, b"foobar") - self.assertIsNone(buf.push_bytes(b"")) # this should not trigger any exception - with self.assertRaises(BufferWriteError): + assert buf.data == b"foobar" + assert buf.push_bytes(b"") is None # this should not trigger any exception + with pytest.raises(BufferWriteError): buf.push_bytes(b"x") # this should! def test_pull_bytes_negative(self): buf = Buffer(data=b"\x08\x07\x06\x05\x04\x03\x02\x01") - with self.assertRaises(OverflowError): + with pytest.raises(OverflowError): buf.pull_bytes(-1) def test_pull_bytes_truncated(self): buf = Buffer(capacity=0) - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): buf.pull_bytes(2) - self.assertEqual(buf.tell(), 0) + assert buf.tell() == 0 def test_pull_bytes_zero(self): buf = Buffer(data=b"\x08\x07\x06\x05\x04\x03\x02\x01") - self.assertEqual(buf.pull_bytes(0), b"") + assert buf.pull_bytes(0) == b"" def test_pull_uint8(self): buf = Buffer(data=b"\x08\x07\x06\x05\x04\x03\x02\x01") - self.assertEqual(buf.pull_uint8(), 0x08) - self.assertEqual(buf.tell(), 1) + assert buf.pull_uint8() == 0x08 + assert buf.tell() == 1 def test_pull_uint8_truncated(self): buf = Buffer(capacity=0) - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): buf.pull_uint8() - self.assertEqual(buf.tell(), 0) + assert buf.tell() == 0 def test_pull_uint16(self): buf = Buffer(data=b"\x08\x07\x06\x05\x04\x03\x02\x01") - self.assertEqual(buf.pull_uint16(), 0x0807) - self.assertEqual(buf.tell(), 2) + assert buf.pull_uint16() == 0x0807 + assert buf.tell() == 2 def test_pull_uint16_truncated(self): buf = Buffer(capacity=1) - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): buf.pull_uint16() - self.assertEqual(buf.tell(), 0) + assert buf.tell() == 0 def test_pull_uint32(self): buf = Buffer(data=b"\x08\x07\x06\x05\x04\x03\x02\x01") - self.assertEqual(buf.pull_uint32(), 0x08070605) - self.assertEqual(buf.tell(), 4) + assert buf.pull_uint32() == 0x08070605 + assert buf.tell() == 4 def test_pull_uint32_truncated(self): buf = Buffer(capacity=3) - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): buf.pull_uint32() - self.assertEqual(buf.tell(), 0) + assert buf.tell() == 0 def test_pull_uint64(self): buf = Buffer(data=b"\x08\x07\x06\x05\x04\x03\x02\x01") - self.assertEqual(buf.pull_uint64(), 0x0807060504030201) - self.assertEqual(buf.tell(), 8) + assert buf.pull_uint64() == 0x0807060504030201 + assert buf.tell() == 8 def test_pull_uint64_truncated(self): buf = Buffer(capacity=7) - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): buf.pull_uint64() - self.assertEqual(buf.tell(), 0) + assert buf.tell() == 0 def test_push_bytes(self): buf = Buffer(capacity=3) buf.push_bytes(b"\x08\x07\x06") - self.assertEqual(buf.data, b"\x08\x07\x06") - self.assertEqual(buf.tell(), 3) + assert buf.data == b"\x08\x07\x06" + assert buf.tell() == 3 def test_push_bytes_truncated(self): buf = Buffer(capacity=3) - with self.assertRaises(BufferWriteError): + with pytest.raises(BufferWriteError): buf.push_bytes(b"\x08\x07\x06\x05") - self.assertEqual(buf.tell(), 0) + assert buf.tell() == 0 def test_push_bytes_zero(self): buf = Buffer(capacity=3) buf.push_bytes(b"") - self.assertEqual(buf.data, b"") - self.assertEqual(buf.tell(), 0) + assert buf.data == b"" + assert buf.tell() == 0 def test_push_uint8(self): buf = Buffer(capacity=1) buf.push_uint8(0x08) - self.assertEqual(buf.data, b"\x08") - self.assertEqual(buf.tell(), 1) + assert buf.data == b"\x08" + assert buf.tell() == 1 def test_push_uint16(self): buf = Buffer(capacity=2) buf.push_uint16(0x0807) - self.assertEqual(buf.data, b"\x08\x07") - self.assertEqual(buf.tell(), 2) + assert buf.data == b"\x08\x07" + assert buf.tell() == 2 def test_push_uint32(self): buf = Buffer(capacity=4) buf.push_uint32(0x08070605) - self.assertEqual(buf.data, b"\x08\x07\x06\x05") - self.assertEqual(buf.tell(), 4) + assert buf.data == b"\x08\x07\x06\x05" + assert buf.tell() == 4 def test_push_uint64(self): buf = Buffer(capacity=8) buf.push_uint64(0x0807060504030201) - self.assertEqual(buf.data, b"\x08\x07\x06\x05\x04\x03\x02\x01") - self.assertEqual(buf.tell(), 8) + assert buf.data == b"\x08\x07\x06\x05\x04\x03\x02\x01" + assert buf.tell() == 8 def test_seek(self): buf = Buffer(data=b"01234567") - self.assertFalse(buf.eof()) - self.assertEqual(buf.tell(), 0) + assert not buf.eof() + assert buf.tell() == 0 buf.seek(4) - self.assertFalse(buf.eof()) - self.assertEqual(buf.tell(), 4) + assert not buf.eof() + assert buf.tell() == 4 buf.seek(8) - self.assertTrue(buf.eof()) - self.assertEqual(buf.tell(), 8) + assert buf.eof() + assert buf.tell() == 8 - with self.assertRaises(OverflowError): + with pytest.raises(OverflowError): buf.seek(-1) - self.assertEqual(buf.tell(), 8) - with self.assertRaises(BufferReadError): + assert buf.tell() == 8 + with pytest.raises(BufferReadError): buf.seek(9) - self.assertEqual(buf.tell(), 8) + assert buf.tell() == 8 -class UintVarTest(TestCase): +class TestUintVar: def roundtrip(self, data, value): buf = Buffer(data=data) - self.assertEqual(buf.pull_uint_var(), value) - self.assertEqual(buf.tell(), len(data)) + assert buf.pull_uint_var() == value + assert buf.tell() == len(data) buf = Buffer(capacity=8) buf.push_uint_var(value) - self.assertEqual(buf.data, data) + assert buf.data == data def test_uint_var(self): # 1 byte @@ -192,29 +192,25 @@ def test_uint_var(self): def test_pull_uint_var_truncated(self): buf = Buffer(capacity=0) - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): buf.pull_uint_var() buf = Buffer(data=b"\xff") - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): buf.pull_uint_var() def test_push_uint_var_too_big(self): buf = Buffer(capacity=8) - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: buf.push_uint_var(4611686018427387904) - self.assertEqual( - str(cm.exception), "Integer is too big for a variable-length integer" - ) + assert str(cm.value) == "Integer is too big for a variable-length integer" def test_size_uint_var(self): - self.assertEqual(size_uint_var(63), 1) - self.assertEqual(size_uint_var(16383), 2) - self.assertEqual(size_uint_var(1073741823), 4) - self.assertEqual(size_uint_var(4611686018427387903), 8) + assert size_uint_var(63) == 1 + assert size_uint_var(16383) == 2 + assert size_uint_var(1073741823) == 4 + assert size_uint_var(4611686018427387903) == 8 - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: size_uint_var(4611686018427387904) - self.assertEqual( - str(cm.exception), "Integer is too big for a variable-length integer" - ) + assert str(cm.value) == "Integer is too big for a variable-length integer" diff --git a/tests/test_connection.py b/tests/test_connection.py index 55aebbe53..276e5171b 100644 --- a/tests/test_connection.py +++ b/tests/test_connection.py @@ -1,11 +1,11 @@ from __future__ import annotations +import pytest import binascii import contextlib import io import time from typing import List, Tuple -from unittest import TestCase, skipIf from qh3 import tls from qh3.buffer import UINT_VAR_MAX, Buffer, encode_uint_var @@ -103,7 +103,7 @@ def create_standalone_client(self, **client_options): # kick-off handshake client.connect(SERVER_ADDR, now=time.time()) - self.assertEqual(drop(client), 2) + assert drop(client) == 2 return client @@ -235,7 +235,7 @@ def transfer(sender, receiver): return datagrams -class QuicConnectionTest(TestCase): +class TestQuicConnection: def assertEvents(self, connection: QuicConnection, expected: list): types = [] while True: @@ -245,7 +245,7 @@ def assertEvents(self, connection: QuicConnection, expected: list): else: break - self.assertListEqual(types, expected) + assert types == expected def assertPacketDropped(self, connection: QuicConnection, trigger: str): log = connection.configuration.quic_logger.to_dict() @@ -254,37 +254,37 @@ def assertPacketDropped(self, connection: QuicConnection, trigger: str): if event["name"] == "transport:packet_dropped": found_trigger = event["data"]["trigger"] break - self.assertEqual(found_trigger, trigger) + assert found_trigger == trigger def assertSentPackets(self, connection: QuicConnection, expected: List[int]): counts = [len(space.sent_packets) for space in connection._loss.spaces] - self.assertEqual(counts, expected) + assert counts == expected def check_handshake(self, client, server, alpn_protocol=None): """ Check handshake completed. """ event = client.next_event() - self.assertEqual(type(event), events.ProtocolNegotiated) - self.assertEqual(event.alpn_protocol, alpn_protocol) + assert type(event) == events.ProtocolNegotiated + assert event.alpn_protocol == alpn_protocol event = client.next_event() - self.assertEqual(type(event), events.HandshakeCompleted) - self.assertEqual(event.alpn_protocol, alpn_protocol) - self.assertEqual(event.early_data_accepted, False) - self.assertEqual(event.session_resumed, False) + assert type(event) == events.HandshakeCompleted + assert event.alpn_protocol == alpn_protocol + assert event.early_data_accepted == False + assert event.session_resumed == False for i in range(7): - self.assertEqual(type(client.next_event()), events.ConnectionIdIssued) - self.assertIsNone(client.next_event()) + assert type(client.next_event()) == events.ConnectionIdIssued + assert client.next_event() is None event = server.next_event() - self.assertEqual(type(event), events.ProtocolNegotiated) - self.assertEqual(event.alpn_protocol, alpn_protocol) + assert type(event) == events.ProtocolNegotiated + assert event.alpn_protocol == alpn_protocol event = server.next_event() - self.assertEqual(type(event), events.HandshakeCompleted) - self.assertEqual(event.alpn_protocol, alpn_protocol) + assert type(event) == events.HandshakeCompleted + assert event.alpn_protocol == alpn_protocol for i in range(7): - self.assertEqual(type(server.next_event()), events.ConnectionIdIssued) - self.assertIsNone(server.next_event()) + assert type(server.next_event()) == events.ConnectionIdIssued + assert server.next_event() is None def test_connect(self): with client_and_server() as (client, server): @@ -292,42 +292,38 @@ def test_connect(self): self.check_handshake(client=client, server=server) # check each endpoint has available connection IDs for the peer - self.assertEqual( - sequence_numbers(client._peer_cid_available), [1, 2, 3, 4, 5, 6, 7] - ) - self.assertEqual( - sequence_numbers(server._peer_cid_available), [1, 2, 3, 4, 5, 6, 7] - ) + assert sequence_numbers(client._peer_cid_available) == [1, 2, 3, 4, 5, 6, 7] + assert sequence_numbers(server._peer_cid_available) == [1, 2, 3, 4, 5, 6, 7] # client closes the connection client.close() - self.assertEqual(transfer(client, server), 1) + assert transfer(client, server) == 1 # check connection closes on the client side client.handle_timer(client.get_timer()) event = client.next_event() - self.assertEqual(type(event), events.ConnectionTerminated) - self.assertEqual(event.error_code, QuicErrorCode.NO_ERROR) - self.assertEqual(event.frame_type, None) - self.assertEqual(event.reason_phrase, "") - self.assertIsNone(client.next_event()) + assert type(event) == events.ConnectionTerminated + assert event.error_code == QuicErrorCode.NO_ERROR + assert event.frame_type == None + assert event.reason_phrase == "" + assert client.next_event() is None # check connection closes on the server side server.handle_timer(server.get_timer()) event = server.next_event() - self.assertEqual(type(event), events.ConnectionTerminated) - self.assertEqual(event.error_code, QuicErrorCode.NO_ERROR) - self.assertEqual(event.frame_type, None) - self.assertEqual(event.reason_phrase, "") - self.assertIsNone(server.next_event()) + assert type(event) == events.ConnectionTerminated + assert event.error_code == QuicErrorCode.NO_ERROR + assert event.frame_type == None + assert event.reason_phrase == "" + assert server.next_event() is None # check client log client_log = client.configuration.quic_logger.to_dict() - self.assertGreater(len(client_log["traces"][0]["events"]), 20) + assert len(client_log["traces"][0]["events"]) > 20 # check server log server_log = server.configuration.quic_logger.to_dict() - self.assertGreater(len(server_log["traces"][0]["events"]), 20) + assert len(server_log["traces"][0]["events"]) > 20 def test_connect_with_alpn(self): with client_and_server( @@ -350,19 +346,17 @@ def test_connect_with_secrets_log(self): # check secrets were logged client_log = client_log_file.getvalue() server_log = server_log_file.getvalue() - self.assertEqual(client_log, server_log) + assert client_log == server_log labels = [] for line in client_log.splitlines(): labels.append(line.split()[0]) - self.assertEqual( - labels, + assert labels == \ [ "SERVER_HANDSHAKE_TRAFFIC_SECRET", "CLIENT_HANDSHAKE_TRAFFIC_SECRET", "SERVER_TRAFFIC_SECRET_0", "CLIENT_TRAFFIC_SECRET_0", - ], - ) + ] def test_connect_with_cert_chain(self): with client_and_server(server_certfile=SERVER_CERTFILE_WITH_CHAIN) as ( @@ -380,12 +374,8 @@ def test_connect_with_cipher_suite_aes128(self): self.check_handshake(client=client, server=server) # check selected cipher suite - self.assertEqual( - client.tls.key_schedule.cipher_suite, tls.CipherSuite.AES_128_GCM_SHA256 - ) - self.assertEqual( - server.tls.key_schedule.cipher_suite, tls.CipherSuite.AES_128_GCM_SHA256 - ) + assert client.tls.key_schedule.cipher_suite == tls.CipherSuite.AES_128_GCM_SHA256 + assert server.tls.key_schedule.cipher_suite == tls.CipherSuite.AES_128_GCM_SHA256 def test_connect_with_cipher_suite_aes256(self): with client_and_server( @@ -395,14 +385,10 @@ def test_connect_with_cipher_suite_aes256(self): self.check_handshake(client=client, server=server) # check selected cipher suite - self.assertEqual( - client.tls.key_schedule.cipher_suite, tls.CipherSuite.AES_256_GCM_SHA384 - ) - self.assertEqual( - server.tls.key_schedule.cipher_suite, tls.CipherSuite.AES_256_GCM_SHA384 - ) + assert client.tls.key_schedule.cipher_suite == tls.CipherSuite.AES_256_GCM_SHA384 + assert server.tls.key_schedule.cipher_suite == tls.CipherSuite.AES_256_GCM_SHA384 - @skipIf("chacha20" in SKIP_TESTS, "Skipping chacha20 tests") + @pytest.mark.skipif("chacha20" in SKIP_TESTS, reason="Skipping chacha20 tests") def test_connect_with_cipher_suite_chacha20(self): with client_and_server( client_options={"cipher_suites": [tls.CipherSuite.CHACHA20_POLY1305_SHA256]} @@ -411,14 +397,10 @@ def test_connect_with_cipher_suite_chacha20(self): self.check_handshake(client=client, server=server) # check selected cipher suite - self.assertEqual( - client.tls.key_schedule.cipher_suite, - tls.CipherSuite.CHACHA20_POLY1305_SHA256, - ) - self.assertEqual( - server.tls.key_schedule.cipher_suite, - tls.CipherSuite.CHACHA20_POLY1305_SHA256, - ) + assert client.tls.key_schedule.cipher_suite == \ + tls.CipherSuite.CHACHA20_POLY1305_SHA256 + assert server.tls.key_schedule.cipher_suite == \ + tls.CipherSuite.CHACHA20_POLY1305_SHA256 def test_connect_without_loss(self): """ @@ -429,8 +411,8 @@ def test_connect_without_loss(self): now = 0.0 client.connect(SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 1280]) - self.assertEqual(client.get_timer(), 0.2) + assert datagram_sizes(items) == [1280, 1280] + assert client.get_timer() == 0.2 self.assertSentPackets(client, [2, 0, 0]) self.assertEvents(client, []) @@ -439,8 +421,8 @@ def test_connect_without_loss(self): server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) server.receive_datagram(items[1][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), SERVER_INITIAL_DATAGRAM_SIZES) - self.assertAlmostEqual(server.get_timer(), 0.25) + assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES + assert server.get_timer() == pytest.approx(0.25) self.assertSentPackets(server, [2, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) @@ -450,8 +432,8 @@ def test_connect_without_loss(self): client.receive_datagram(items[1][0], SERVER_ADDR, now=now) client.receive_datagram(items[2][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), CLIENT_HANDSHAKE_DATAGRAM_SIZES) - self.assertAlmostEqual(client.get_timer(), 0.425) + assert datagram_sizes(items) == CLIENT_HANDSHAKE_DATAGRAM_SIZES + assert client.get_timer() == pytest.approx(0.425) self.assertSentPackets(client, [0, 1, 1]) self.assertEvents( client, [events.ProtocolNegotiated] + HANDSHAKE_COMPLETED_EVENTS @@ -460,16 +442,17 @@ def test_connect_without_loss(self): now += TICK server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [229]) - self.assertAlmostEqual(server.get_timer(), 0.425) + assert datagram_sizes(items) == [229] + assert server.get_timer() == pytest.approx(0.425) self.assertSentPackets(server, [0, 0, 1]) self.assertEvents(server, HANDSHAKE_COMPLETED_EVENTS) now += TICK client.receive_datagram(items[0][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [32]) - self.assertAlmostEqual(client.get_timer(), 60.2) # idle timeout + assert datagram_sizes(items) == [32] + # idle timeout + assert client.get_timer() == pytest.approx(60.2) self.assertSentPackets(client, [0, 0, 1]) self.assertEvents(client, []) @@ -485,8 +468,8 @@ def test_connect_with_loss_1(self): now = 0.0 client.connect(SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 1280]) - self.assertEqual(client.get_timer(), 0.2) + assert datagram_sizes(items) == [1280, 1280] + assert client.get_timer() == 0.2 self.assertSentPackets(client, [2, 0, 0]) self.assertEvents(client, []) @@ -494,8 +477,8 @@ def test_connect_with_loss_1(self): now = client.get_timer() client.handle_timer(now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 1280]) - self.assertAlmostEqual(client.get_timer(), 0.6) + assert datagram_sizes(items) == [1280, 1280] + assert client.get_timer() == pytest.approx(0.6) self.assertSentPackets(client, [2, 0, 0]) self.assertEvents(client, []) @@ -504,8 +487,8 @@ def test_connect_with_loss_1(self): server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) server.receive_datagram(items[1][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), SERVER_INITIAL_DATAGRAM_SIZES) - self.assertAlmostEqual(server.get_timer(), 0.45) + assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES + assert server.get_timer() == pytest.approx(0.45) self.assertSentPackets(server, [2, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) @@ -515,8 +498,8 @@ def test_connect_with_loss_1(self): client.receive_datagram(items[1][0], SERVER_ADDR, now=now) client.receive_datagram(items[2][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), CLIENT_HANDSHAKE_DATAGRAM_SIZES) - self.assertAlmostEqual(client.get_timer(), 0.625) + assert datagram_sizes(items) == CLIENT_HANDSHAKE_DATAGRAM_SIZES + assert client.get_timer() == pytest.approx(0.625) self.assertSentPackets(client, [0, 1, 1]) self.assertEvents( client, [events.ProtocolNegotiated] + HANDSHAKE_COMPLETED_EVENTS @@ -525,16 +508,17 @@ def test_connect_with_loss_1(self): now += TICK server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [229]) - self.assertAlmostEqual(server.get_timer(), 0.625) + assert datagram_sizes(items) == [229] + assert server.get_timer() == pytest.approx(0.625) self.assertSentPackets(server, [0, 0, 1]) self.assertEvents(server, HANDSHAKE_COMPLETED_EVENTS) now += TICK client.receive_datagram(items[0][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [32]) - self.assertAlmostEqual(client.get_timer(), 60.4) # idle timeout + assert datagram_sizes(items) == [32] + # idle timeout + assert client.get_timer() == pytest.approx(60.4) self.assertSentPackets(client, [0, 0, 1]) self.assertEvents(client, []) @@ -551,8 +535,8 @@ def test_connect_with_loss_2(self): now = 0.0 client.connect(SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 1280]) - self.assertEqual(client.get_timer(), 0.2) + assert datagram_sizes(items) == [1280, 1280] + assert client.get_timer() == 0.2 self.assertSentPackets(client, [2, 0, 0]) self.assertEvents(client, []) @@ -562,8 +546,8 @@ def test_connect_with_loss_2(self): server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) server.receive_datagram(items[1][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), SERVER_INITIAL_DATAGRAM_SIZES) - self.assertEqual(server.get_timer(), 0.25) + assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES + assert server.get_timer() == 0.25 self.assertSentPackets(server, [2, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) @@ -571,8 +555,8 @@ def test_connect_with_loss_2(self): now += TICK client.receive_datagram(items[1][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 1280]) - self.assertAlmostEqual(client.get_timer(), 0.3) + assert datagram_sizes(items) == [1280, 1280] + assert client.get_timer() == pytest.approx(0.3) self.assertSentPackets(client, [2, 0, 0]) self.assertEvents(client, []) @@ -583,7 +567,7 @@ def test_connect_with_loss_2(self): server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) server.receive_datagram(items[1][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 1280, 890]) + assert datagram_sizes(items) == [1280, 1280, 890] # self.assertAlmostEqual(server.get_timer(), 0.35) self.assertSentPackets(server, [1, 2, 0]) self.assertEvents(server, []) @@ -594,7 +578,7 @@ def test_connect_with_loss_2(self): client.receive_datagram(items[1][0], SERVER_ADDR, now=now) client.receive_datagram(items[2][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), CLIENT_HANDSHAKE_DATAGRAM_SIZES) + assert datagram_sizes(items) == CLIENT_HANDSHAKE_DATAGRAM_SIZES # self.assertAlmostEqual(client.get_timer(), 0.525) self.assertSentPackets(client, [0, 1, 1]) self.assertEvents( @@ -604,7 +588,7 @@ def test_connect_with_loss_2(self): now += TICK server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [229]) + assert datagram_sizes(items) == [229] # self.assertAlmostEqual(server.get_timer(), 0.525) self.assertSentPackets(server, [0, 0, 1]) self.assertEvents(server, HANDSHAKE_COMPLETED_EVENTS) @@ -612,8 +596,9 @@ def test_connect_with_loss_2(self): now += TICK client.receive_datagram(items[0][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [32]) - self.assertAlmostEqual(client.get_timer(), 60.3) # idle timeout + assert datagram_sizes(items) == [32] + # idle timeout + assert client.get_timer() == pytest.approx(60.3) self.assertSentPackets(client, [0, 0, 1]) self.assertEvents(client, []) @@ -631,8 +616,8 @@ def test_connect_with_loss_3(self): now = 0.0 client.connect(SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 1280]) - self.assertEqual(client.get_timer(), 0.2) + assert datagram_sizes(items) == [1280, 1280] + assert client.get_timer() == 0.2 self.assertSentPackets(client, [2, 0, 0]) self.assertEvents(client, []) @@ -641,8 +626,8 @@ def test_connect_with_loss_3(self): server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) server.receive_datagram(items[1][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), SERVER_INITIAL_DATAGRAM_SIZES) - self.assertEqual(server.get_timer(), 0.25) + assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES + assert server.get_timer() == 0.25 self.assertSentPackets(server, [2, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) @@ -650,8 +635,8 @@ def test_connect_with_loss_3(self): now = client.get_timer() client.handle_timer(now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 1280]) - self.assertAlmostEqual(client.get_timer(), 0.6) + assert datagram_sizes(items) == [1280, 1280] + assert client.get_timer() == pytest.approx(0.6) self.assertSentPackets(client, [2, 0, 0]) self.assertEvents(client, []) @@ -660,8 +645,8 @@ def test_connect_with_loss_3(self): server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) server.receive_datagram(items[1][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), SERVER_INITIAL_DATAGRAM_SIZES) - self.assertEqual(server.get_timer(), 0.45) + assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES + assert server.get_timer() == 0.45 self.assertSentPackets(server, [2, 2, 0]) self.assertEvents(server, []) @@ -671,9 +656,9 @@ def test_connect_with_loss_3(self): client.receive_datagram(items[1][0], SERVER_ADDR, now=now) client.receive_datagram(items[2][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), CLIENT_HANDSHAKE_DATAGRAM_SIZES) - self.assertGreaterEqual(client.get_timer(), 0.5) - self.assertLessEqual(client.get_timer(), 0.63) + assert datagram_sizes(items) == CLIENT_HANDSHAKE_DATAGRAM_SIZES + assert client.get_timer() >= 0.5 + assert client.get_timer() <= 0.63 self.assertSentPackets(client, [0, 1, 1]) self.assertEvents( client, [events.ProtocolNegotiated] + HANDSHAKE_COMPLETED_EVENTS @@ -682,16 +667,17 @@ def test_connect_with_loss_3(self): now += TICK server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [229]) - self.assertAlmostEqual(server.get_timer(), 0.625) + assert datagram_sizes(items) == [229] + assert server.get_timer() == pytest.approx(0.625) self.assertSentPackets(server, [0, 0, 1]) self.assertEvents(server, HANDSHAKE_COMPLETED_EVENTS) now += TICK client.receive_datagram(items[0][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [32]) - self.assertAlmostEqual(client.get_timer(), 60.4) # idle timeout + assert datagram_sizes(items) == [32] + # idle timeout + assert client.get_timer() == pytest.approx(60.4) self.assertSentPackets(client, [0, 0, 1]) self.assertEvents(client, []) @@ -704,8 +690,8 @@ def test_connect_with_loss_4(self): now = 0.0 client.connect(SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 1280]) - self.assertEqual(client.get_timer(), 0.2) + assert datagram_sizes(items) == [1280, 1280] + assert client.get_timer() == 0.2 self.assertSentPackets(client, [2, 0, 0]) self.assertEvents(client, []) @@ -715,8 +701,8 @@ def test_connect_with_loss_4(self): server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) server.receive_datagram(items[1][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), SERVER_INITIAL_DATAGRAM_SIZES) - self.assertEqual(server.get_timer(), 0.25) + assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES + assert server.get_timer() == 0.25 self.assertSentPackets(server, [2, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) @@ -725,8 +711,8 @@ def test_connect_with_loss_4(self): client.receive_datagram(items[0][0], SERVER_ADDR, now=now) client.receive_datagram(items[1][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280]) - self.assertAlmostEqual(client.get_timer(), 0.325) + assert datagram_sizes(items) == [1280] + assert client.get_timer() == pytest.approx(0.325) self.assertSentPackets(client, [0, 1, 0]) self.assertEvents(client, [events.ProtocolNegotiated]) @@ -734,8 +720,8 @@ def test_connect_with_loss_4(self): now = client.get_timer() client.handle_timer(now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [45]) - self.assertAlmostEqual(client.get_timer(), 0.975) + assert datagram_sizes(items) == [45] + assert client.get_timer() == pytest.approx(0.975) self.assertSentPackets(client, [0, 2, 0]) self.assertEvents(client, []) @@ -743,8 +729,8 @@ def test_connect_with_loss_4(self): now += TICK server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [48]) - self.assertAlmostEqual(server.get_timer(), 0.25) + assert datagram_sizes(items) == [48] + assert server.get_timer() == pytest.approx(0.25) self.assertSentPackets(server, [0, 3, 0]) self.assertEvents(server, []) @@ -752,8 +738,8 @@ def test_connect_with_loss_4(self): now = server.get_timer() server.handle_timer(now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 890]) - self.assertAlmostEqual(server.get_timer(), 0.65) + assert datagram_sizes(items) == [1280, 890] + assert server.get_timer() == pytest.approx(0.65) self.assertSentPackets(server, [0, 3, 0]) self.assertEvents(server, []) @@ -762,24 +748,25 @@ def test_connect_with_loss_4(self): client.receive_datagram(items[0][0], SERVER_ADDR, now=now) client.receive_datagram(items[1][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [313]) - self.assertAlmostEqual(client.get_timer(), 0.95) + assert datagram_sizes(items) == [313] + assert client.get_timer() == pytest.approx(0.95) self.assertSentPackets(client, [0, 3, 1]) self.assertEvents(client, HANDSHAKE_COMPLETED_EVENTS) now += TICK server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [229]) - self.assertAlmostEqual(server.get_timer(), 0.675) + assert datagram_sizes(items) == [229] + assert server.get_timer() == pytest.approx(0.675) self.assertSentPackets(server, [0, 0, 1]) self.assertEvents(server, HANDSHAKE_COMPLETED_EVENTS) now += TICK client.receive_datagram(items[0][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [32]) - self.assertAlmostEqual(client.get_timer(), 60.4) # idle timeout + assert datagram_sizes(items) == [32] + # idle timeout + assert client.get_timer() == pytest.approx(60.4) self.assertSentPackets(client, [0, 0, 1]) self.assertEvents(client, []) @@ -792,16 +779,16 @@ def test_connect_with_loss_5(self): now = 0.0 client.connect(SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [1280, 1280]) - self.assertEqual(client.get_timer(), 0.2) + assert datagram_sizes(items) == [1280, 1280] + assert client.get_timer() == 0.2 # server receives INITIAL, sends INITIAL + HANDSHAKE now += TICK server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) server.receive_datagram(items[1][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), SERVER_INITIAL_DATAGRAM_SIZES) - self.assertEqual(server.get_timer(), 0.25) + assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES + assert server.get_timer() == 0.25 self.assertSentPackets(server, [2, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) @@ -811,8 +798,8 @@ def test_connect_with_loss_5(self): client.receive_datagram(items[1][0], SERVER_ADDR, now=now) client.receive_datagram(items[2][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), CLIENT_HANDSHAKE_DATAGRAM_SIZES) - self.assertAlmostEqual(client.get_timer(), 0.425) + assert datagram_sizes(items) == CLIENT_HANDSHAKE_DATAGRAM_SIZES + assert client.get_timer() == pytest.approx(0.425) self.assertSentPackets(client, [0, 1, 1]) self.assertEvents( client, [events.ProtocolNegotiated] + HANDSHAKE_COMPLETED_EVENTS @@ -822,8 +809,8 @@ def test_connect_with_loss_5(self): now += TICK server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [229]) - self.assertAlmostEqual(server.get_timer(), 0.425) + assert datagram_sizes(items) == [229] + assert server.get_timer() == pytest.approx(0.425) self.assertSentPackets(server, [0, 0, 1]) self.assertEvents(server, HANDSHAKE_COMPLETED_EVENTS) @@ -831,8 +818,8 @@ def test_connect_with_loss_5(self): now = server.get_timer() server.handle_timer(now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [29]) - self.assertAlmostEqual(server.get_timer(), 0.975) + assert datagram_sizes(items) == [29] + assert server.get_timer() == pytest.approx(0.975) self.assertSentPackets(server, [0, 0, 2]) self.assertEvents(server, []) @@ -840,20 +827,20 @@ def test_connect_with_loss_5(self): now += TICK client.receive_datagram(items[0][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [32]) - self.assertAlmostEqual(client.get_timer(), 0.425) + assert datagram_sizes(items) == [32] + assert client.get_timer() == pytest.approx(0.425) self.assertSentPackets(client, [0, 1, 2]) self.assertEvents(client, []) # server receives ACK, retransmits HANDSHAKE_DONE now += TICK - self.assertFalse(server._handshake_done_pending) + assert not server._handshake_done_pending server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) - self.assertTrue(server._handshake_done_pending) + assert server._handshake_done_pending items = server.datagrams_to_send(now=now) - self.assertFalse(server._handshake_done_pending) - self.assertEqual(datagram_sizes(items), [224]) - self.assertAlmostEqual(server.get_timer(), 0.7625) + assert not server._handshake_done_pending + assert datagram_sizes(items) == [224] + assert server.get_timer() == pytest.approx(0.7625) self.assertSentPackets(server, [0, 0, 1]) # FIXME: the server re-emits the ConnectionIdIssued events self.assertEvents(server, HANDSHAKE_COMPLETED_EVENTS[1:]) @@ -861,16 +848,17 @@ def test_connect_with_loss_5(self): now += TICK client.receive_datagram(items[0][0], SERVER_ADDR, now=now) items = client.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), [32]) - self.assertAlmostEqual(client.get_timer(), 0.425) + assert datagram_sizes(items) == [32] + assert client.get_timer() == pytest.approx(0.425) self.assertSentPackets(client, [0, 0, 3]) self.assertEvents(client, []) now += TICK server.receive_datagram(items[0][0], CLIENT_ADDR, now=now) items = server.datagrams_to_send(now=now) - self.assertEqual(datagram_sizes(items), []) - self.assertAlmostEqual(server.get_timer(), 60.625) # idle timeout + assert datagram_sizes(items) == [] + # idle timeout + assert server.get_timer() == pytest.approx(60.625) self.assertSentPackets(server, [0, 0, 0]) self.assertEvents(server, []) @@ -888,10 +876,8 @@ def patched_initialize(peer_cid: bytes): client._initialize = patched_initialize with client_and_server(client_patch=patch) as (client, server): - self.assertEqual( - server._close_event.reason_phrase, - "No QUIC transport parameters received", - ) + assert server._close_event.reason_phrase == \ + "No QUIC transport parameters received" def test_connect_with_compatible_version_negotiation_1(self): """ @@ -906,8 +892,8 @@ def test_connect_with_compatible_version_negotiation_1(self): ) as (client, server): # check handshake completed self.check_handshake(client=client, server=server) - self.assertEqual(client._version, QuicProtocolVersion.VERSION_1) - self.assertEqual(server._version, QuicProtocolVersion.VERSION_1) + assert client._version == QuicProtocolVersion.VERSION_1 + assert server._version == QuicProtocolVersion.VERSION_1 def test_connect_with_compatible_version_negotiation_1_to_2(self): """ @@ -926,8 +912,8 @@ def test_connect_with_compatible_version_negotiation_1_to_2(self): ) as (client, server): # check handshake completed self.check_handshake(client=client, server=server) - self.assertEqual(client._version, QuicProtocolVersion.VERSION_2) - self.assertEqual(server._version, QuicProtocolVersion.VERSION_2) + assert client._version == QuicProtocolVersion.VERSION_2 + assert server._version == QuicProtocolVersion.VERSION_2 def test_connect_with_compatible_version_negotiation_2(self): """ @@ -942,8 +928,8 @@ def test_connect_with_compatible_version_negotiation_2(self): ) as (client, server): # check handshake completed self.check_handshake(client=client, server=server) - self.assertEqual(client._version, QuicProtocolVersion.VERSION_2) - self.assertEqual(server._version, QuicProtocolVersion.VERSION_2) + assert client._version == QuicProtocolVersion.VERSION_2 + assert server._version == QuicProtocolVersion.VERSION_2 def test_connect_with_compatible_version_negotiation_2_to_1(self): """ @@ -962,8 +948,8 @@ def test_connect_with_compatible_version_negotiation_2_to_1(self): ) as (client, server): # check handshake completed self.check_handshake(client=client, server=server) - self.assertEqual(client._version, QuicProtocolVersion.VERSION_1) - self.assertEqual(server._version, QuicProtocolVersion.VERSION_1) + assert client._version == QuicProtocolVersion.VERSION_1 + assert server._version == QuicProtocolVersion.VERSION_1 def test_connect_with_quantum_readiness(self): with client_and_server(client_options={"quantum_readiness_test": True}) as ( @@ -973,7 +959,7 @@ def test_connect_with_quantum_readiness(self): stream_id = client.get_next_available_stream_id() client.send_stream_data(stream_id, b"hello") - self.assertEqual(roundtrip(client, server), (1, 1)) + assert roundtrip(client, server) == (1, 1) received = None while True: @@ -983,7 +969,7 @@ def test_connect_with_quantum_readiness(self): elif event is None: break - self.assertEqual(received, b"hello") + assert received == b"hello" def test_connect_with_0rtt(self): client_ticket = None @@ -1008,14 +994,14 @@ def save_session_ticket(ticket): stream_id = client.get_next_available_stream_id() client.send_stream_data(stream_id, b"hello") - self.assertEqual(roundtrip(client, server), (2, 2)) + assert roundtrip(client, server) == (2, 2) event = server.next_event() - self.assertEqual(type(event), events.ProtocolNegotiated) + assert type(event) == events.ProtocolNegotiated event = server.next_event() - self.assertEqual(type(event), events.StreamDataReceived) - self.assertEqual(event.data, b"hello") + assert type(event) == events.StreamDataReceived + assert event.data == b"hello" def test_connect_with_0rtt_bad_max_early_data(self): client_ticket = None @@ -1045,112 +1031,92 @@ def save_session_ticket(ticket): ) as (client, server): # check handshake failed event = client.next_event() - self.assertIsNone(event) + assert event is None def test_change_connection_id(self): with client_and_server() as (client, server): - self.assertEqual( - sequence_numbers(client._peer_cid_available), [1, 2, 3, 4, 5, 6, 7] - ) + assert sequence_numbers(client._peer_cid_available) == [1, 2, 3, 4, 5, 6, 7] # the client changes connection ID client.change_connection_id() - self.assertEqual(transfer(client, server), 1) - self.assertEqual( - sequence_numbers(client._peer_cid_available), [2, 3, 4, 5, 6, 7] - ) + assert transfer(client, server) == 1 + assert sequence_numbers(client._peer_cid_available) == [2, 3, 4, 5, 6, 7] # the server provides a new connection ID - self.assertEqual(transfer(server, client), 1) - self.assertEqual( - sequence_numbers(client._peer_cid_available), [2, 3, 4, 5, 6, 7, 8] - ) + assert transfer(server, client) == 1 + assert sequence_numbers(client._peer_cid_available) == [2, 3, 4, 5, 6, 7, 8] def test_change_connection_id_retransmit_new_connection_id(self): with client_and_server() as (client, server): - self.assertEqual( - sequence_numbers(client._peer_cid_available), [1, 2, 3, 4, 5, 6, 7] - ) + assert sequence_numbers(client._peer_cid_available) == [1, 2, 3, 4, 5, 6, 7] # the client changes connection ID client.change_connection_id() - self.assertEqual(transfer(client, server), 1) - self.assertEqual( - sequence_numbers(client._peer_cid_available), [2, 3, 4, 5, 6, 7] - ) + assert transfer(client, server) == 1 + assert sequence_numbers(client._peer_cid_available) == [2, 3, 4, 5, 6, 7] # the server provides a new connection ID, NEW_CONNECTION_ID is lost - self.assertEqual(drop(server), 1) - self.assertEqual( - sequence_numbers(client._peer_cid_available), [2, 3, 4, 5, 6, 7] - ) + assert drop(server) == 1 + assert sequence_numbers(client._peer_cid_available) == [2, 3, 4, 5, 6, 7] # NEW_CONNECTION_ID is retransmitted server._on_new_connection_id_delivery( QuicDeliveryState.LOST, server._host_cids[-1] ) - self.assertEqual(transfer(server, client), 1) - self.assertEqual( - sequence_numbers(client._peer_cid_available), [2, 3, 4, 5, 6, 7, 8] - ) + assert transfer(server, client) == 1 + assert sequence_numbers(client._peer_cid_available) == [2, 3, 4, 5, 6, 7, 8] def test_change_connection_id_retransmit_retire_connection_id(self): with client_and_server() as (client, server): - self.assertEqual( - sequence_numbers(client._peer_cid_available), [1, 2, 3, 4, 5, 6, 7] - ) + assert sequence_numbers(client._peer_cid_available) == [1, 2, 3, 4, 5, 6, 7] # the client changes connection ID, RETIRE_CONNECTION_ID is lost client.change_connection_id() - self.assertEqual(drop(client), 1) - self.assertEqual( - sequence_numbers(client._peer_cid_available), [2, 3, 4, 5, 6, 7] - ) + assert drop(client) == 1 + assert sequence_numbers(client._peer_cid_available) == [2, 3, 4, 5, 6, 7] # RETIRE_CONNECTION_ID is retransmitted client._on_retire_connection_id_delivery(QuicDeliveryState.LOST, 0) - self.assertEqual(transfer(client, server), 1) + assert transfer(client, server) == 1 # the server provides a new connection ID - self.assertEqual(transfer(server, client), 1) - self.assertEqual( - sequence_numbers(client._peer_cid_available), [2, 3, 4, 5, 6, 7, 8] - ) + assert transfer(server, client) == 1 + assert sequence_numbers(client._peer_cid_available) == [2, 3, 4, 5, 6, 7, 8] def test_get_next_available_stream_id(self): with client_and_server() as (client, server): # client stream_id = client.get_next_available_stream_id() - self.assertEqual(stream_id, 0) + assert stream_id == 0 client.send_stream_data(stream_id, b"hello") stream_id = client.get_next_available_stream_id() - self.assertEqual(stream_id, 4) + assert stream_id == 4 client.send_stream_data(stream_id, b"hello") stream_id = client.get_next_available_stream_id(is_unidirectional=True) - self.assertEqual(stream_id, 2) + assert stream_id == 2 client.send_stream_data(stream_id, b"hello") stream_id = client.get_next_available_stream_id(is_unidirectional=True) - self.assertEqual(stream_id, 6) + assert stream_id == 6 client.send_stream_data(stream_id, b"hello") # server stream_id = server.get_next_available_stream_id() - self.assertEqual(stream_id, 1) + assert stream_id == 1 server.send_stream_data(stream_id, b"hello") stream_id = server.get_next_available_stream_id() - self.assertEqual(stream_id, 5) + assert stream_id == 5 server.send_stream_data(stream_id, b"hello") stream_id = server.get_next_available_stream_id(is_unidirectional=True) - self.assertEqual(stream_id, 3) + assert stream_id == 3 server.send_stream_data(stream_id, b"hello") stream_id = server.get_next_available_stream_id(is_unidirectional=True) - self.assertEqual(stream_id, 7) + assert stream_id == 7 server.send_stream_data(stream_id, b"hello") def test_datagram_frame(self): @@ -1163,11 +1129,11 @@ def test_datagram_frame(self): # send datagram client.send_datagram_frame(b"hello") - self.assertEqual(transfer(client, server), 1) + assert transfer(client, server) == 1 event = server.next_event() - self.assertEqual(type(event), events.DatagramFrameReceived) - self.assertEqual(event.data, b"hello") + assert type(event) == events.DatagramFrameReceived + assert event.data == b"hello" def test_datagram_frame_2(self): # payload which exactly fills an entire packet @@ -1185,21 +1151,21 @@ def test_datagram_frame_2(self): client.send_datagram_frame(payload) # client can only 11 datagrams are sent due to congestion control - self.assertEqual(transfer(client, server), 12) + assert transfer(client, server) == 12 for i in range(11): event = server.next_event() - self.assertEqual(type(event), events.DatagramFrameReceived) - self.assertEqual(event.data, payload) + assert type(event) == events.DatagramFrameReceived + assert event.data == payload # server sends ACK - self.assertEqual(transfer(server, client), 1) + assert transfer(server, client) == 1 # client sends remaining datagrams - self.assertEqual(transfer(client, server), 8) + assert transfer(client, server) == 8 for i in range(9): event = server.next_event() - self.assertEqual(type(event), events.DatagramFrameReceived) - self.assertEqual(event.data, payload) + assert type(event) == events.DatagramFrameReceived + assert event.data == payload def test_decryption_error(self): with client_and_server() as (client, server): @@ -1234,10 +1200,10 @@ def patched_initialize(peer_cid: bytes): server.handle_timer(timer_at) event = server.next_event() - self.assertEqual(type(event), events.ConnectionTerminated) - self.assertEqual(event.error_code, 326) - self.assertEqual(event.frame_type, QuicFrameType.CRYPTO) - self.assertEqual(event.reason_phrase, "No supported protocol version") + assert type(event) == events.ConnectionTerminated + assert event.error_code == 326 + assert event.frame_type == QuicFrameType.CRYPTO + assert event.reason_phrase == "No supported protocol version" def test_receive_datagram_garbage(self): client = create_standalone_client(self) @@ -1275,15 +1241,13 @@ def encrypt_packet(plain_header, plain_payload, packet_number): for datagram in builder.flush()[0]: client.receive_datagram(datagram, SERVER_ADDR, now=time.time()) - self.assertEqual(drop(client), 1) - self.assertEqual( - client._close_event, + assert drop(client) == 1 + assert client._close_event == \ events.ConnectionTerminated( error_code=QuicErrorCode.PROTOCOL_VIOLATION, frame_type=QuicFrameType.PADDING, reason_phrase="Reserved bits must be zero", - ), - ) + ) def test_receive_datagram_wrong_version(self): client = create_standalone_client(self) @@ -1304,7 +1268,7 @@ def test_receive_datagram_wrong_version(self): for datagram in builder.flush()[0]: client.receive_datagram(datagram, SERVER_ADDR, now=time.time()) - self.assertEqual(drop(client), 0) + assert drop(client) == 0 self.assertPacketDropped(client, "unsupported_version") @@ -1322,7 +1286,7 @@ def test_receive_datagram_retry(self): SERVER_ADDR, now=time.time(), ) - self.assertEqual(drop(client), 2) + assert drop(client) == 2 def test_receive_datagram_retry_wrong_destination_cid(self): client = create_standalone_client(self) @@ -1338,7 +1302,7 @@ def test_receive_datagram_retry_wrong_destination_cid(self): SERVER_ADDR, now=time.time(), ) - self.assertEqual(drop(client), 0) + assert drop(client) == 0 self.assertPacketDropped(client, "unknown_connection_id") def test_receive_datagram_retry_wrong_integrity_tag(self): @@ -1356,7 +1320,7 @@ def test_receive_datagram_retry_wrong_integrity_tag(self): SERVER_ADDR, now=time.time(), ) - self.assertEqual(drop(client), 0) + assert drop(client) == 0 def test_handle_ack_frame_ecn(self): client = create_standalone_client(self) @@ -1374,30 +1338,26 @@ def test_handle_connection_close_frame(self): frame_type=QuicFrameType.ACK, reason_phrase="illegal ACK frame", ) - self.assertEqual(roundtrip(server, client), (1, 0)) + assert roundtrip(server, client) == (1, 0) - self.assertEqual( - client._close_event, + assert client._close_event == \ events.ConnectionTerminated( error_code=QuicErrorCode.PROTOCOL_VIOLATION, frame_type=QuicFrameType.ACK, reason_phrase="illegal ACK frame", - ), - ) + ) def test_handle_connection_close_frame_app(self): with client_and_server() as (client, server): server.close(error_code=QuicErrorCode.NO_ERROR, reason_phrase="goodbye") - self.assertEqual(roundtrip(server, client), (1, 0)) + assert roundtrip(server, client) == (1, 0) - self.assertEqual( - client._close_event, + assert client._close_event == \ events.ConnectionTerminated( error_code=QuicErrorCode.NO_ERROR, frame_type=None, reason_phrase="goodbye", - ), - ) + ) def test_handle_connection_close_frame_app_not_utf8(self): client = create_standalone_client(self) @@ -1408,31 +1368,25 @@ def test_handle_connection_close_frame_app_not_utf8(self): Buffer(data=binascii.unhexlify("0008676f6f6462798200")), ) - self.assertEqual( - client._close_event, + assert client._close_event == \ events.ConnectionTerminated( error_code=QuicErrorCode.NO_ERROR, frame_type=None, reason_phrase="", - ), - ) + ) def test_handle_crypto_frame_over_largest_offset(self): with client_and_server() as (client, server): # client receives offset + length > 2^62 - 1 - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_crypto_frame( client_receive_context(client), QuicFrameType.CRYPTO, Buffer(data=encode_uint_var(UINT_VAR_MAX) + encode_uint_var(1)), ) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.FRAME_ENCODING_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual( - cm.exception.reason_phrase, "offset + length cannot exceed 2^62 - 1" - ) + assert cm.value.error_code == QuicErrorCode.FRAME_ENCODING_ERROR + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert cm.value.reason_phrase == "offset + length cannot exceed 2^62 - 1" def test_handle_data_blocked_frame(self): with client_and_server() as (client, server): @@ -1452,35 +1406,33 @@ def test_handle_datagram_frame(self): Buffer(data=b"hello"), ) - self.assertEqual( - client.next_event(), events.DatagramFrameReceived(data=b"hello") - ) + assert client.next_event() == events.DatagramFrameReceived(data=b"hello") def test_handle_datagram_frame_not_allowed(self): client = create_standalone_client(self, max_datagram_frame_size=None) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_datagram_frame( client_receive_context(client), QuicFrameType.DATAGRAM, Buffer(data=b"hello"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual(cm.exception.frame_type, QuicFrameType.DATAGRAM) - self.assertEqual(cm.exception.reason_phrase, "Unexpected DATAGRAM frame") + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.DATAGRAM + assert cm.value.reason_phrase == "Unexpected DATAGRAM frame" def test_handle_datagram_frame_too_large(self): client = create_standalone_client(self, max_datagram_frame_size=5) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_datagram_frame( client_receive_context(client), QuicFrameType.DATAGRAM, Buffer(data=b"hello"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual(cm.exception.frame_type, QuicFrameType.DATAGRAM) - self.assertEqual(cm.exception.reason_phrase, "Unexpected DATAGRAM frame") + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.DATAGRAM + assert cm.value.reason_phrase == "Unexpected DATAGRAM frame" def test_handle_datagram_frame_with_length(self): client = create_standalone_client(self, max_datagram_frame_size=7) @@ -1491,55 +1443,51 @@ def test_handle_datagram_frame_with_length(self): Buffer(data=b"\x05hellojunk"), ) - self.assertEqual( - client.next_event(), events.DatagramFrameReceived(data=b"hello") - ) + assert client.next_event() == events.DatagramFrameReceived(data=b"hello") def test_handle_datagram_frame_with_length_not_allowed(self): client = create_standalone_client(self, max_datagram_frame_size=None) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_datagram_frame( client_receive_context(client), QuicFrameType.DATAGRAM_WITH_LENGTH, Buffer(data=b"\x05hellojunk"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual(cm.exception.frame_type, QuicFrameType.DATAGRAM_WITH_LENGTH) - self.assertEqual(cm.exception.reason_phrase, "Unexpected DATAGRAM frame") + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.DATAGRAM_WITH_LENGTH + assert cm.value.reason_phrase == "Unexpected DATAGRAM frame" def test_handle_datagram_frame_with_length_too_large(self): client = create_standalone_client(self, max_datagram_frame_size=6) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_datagram_frame( client_receive_context(client), QuicFrameType.DATAGRAM_WITH_LENGTH, Buffer(data=b"\x05hellojunk"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual(cm.exception.frame_type, QuicFrameType.DATAGRAM_WITH_LENGTH) - self.assertEqual(cm.exception.reason_phrase, "Unexpected DATAGRAM frame") + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.DATAGRAM_WITH_LENGTH + assert cm.value.reason_phrase == "Unexpected DATAGRAM frame" def test_handle_handshake_done_not_allowed(self): with client_and_server() as (client, server): # server receives HANDSHAKE_DONE frame - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: server._handle_handshake_done_frame( client_receive_context(server), QuicFrameType.HANDSHAKE_DONE, Buffer(data=b""), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual(cm.exception.frame_type, QuicFrameType.HANDSHAKE_DONE) - self.assertEqual( - cm.exception.reason_phrase, - "Clients must not send HANDSHAKE_DONE frames", - ) + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.HANDSHAKE_DONE + assert cm.value.reason_phrase == \ + "Clients must not send HANDSHAKE_DONE frames" def test_handle_max_data_frame(self): with client_and_server() as (client, server): - self.assertEqual(client._remote_max_data, 1048576) + assert client._remote_max_data == 1048576 # client receives MAX_DATA raising limit client._handle_max_data_frame( @@ -1547,13 +1495,13 @@ def test_handle_max_data_frame(self): QuicFrameType.MAX_DATA, Buffer(data=encode_uint_var(1048577)), ) - self.assertEqual(client._remote_max_data, 1048577) + assert client._remote_max_data == 1048577 def test_handle_max_stream_data_frame(self): with client_and_server() as (client, server): # client creates bidirectional stream 0 stream = client._get_or_create_stream_for_send(stream_id=0) - self.assertEqual(stream.max_stream_data_remote, 1048576) + assert stream.max_stream_data_remote == 1048576 # client receives MAX_STREAM_DATA raising limit client._handle_max_stream_data_frame( @@ -1561,7 +1509,7 @@ def test_handle_max_stream_data_frame(self): QuicFrameType.MAX_STREAM_DATA, Buffer(data=b"\x00" + encode_uint_var(1048577)), ) - self.assertEqual(stream.max_stream_data_remote, 1048577) + assert stream.max_stream_data_remote == 1048577 # client receives MAX_STREAM_DATA lowering limit client._handle_max_stream_data_frame( @@ -1569,7 +1517,7 @@ def test_handle_max_stream_data_frame(self): QuicFrameType.MAX_STREAM_DATA, Buffer(data=b"\x00" + encode_uint_var(1048575)), ) - self.assertEqual(stream.max_stream_data_remote, 1048577) + assert stream.max_stream_data_remote == 1048577 def test_handle_max_stream_data_frame_receive_only(self): with client_and_server() as (client, server): @@ -1577,19 +1525,19 @@ def test_handle_max_stream_data_frame_receive_only(self): server.send_stream_data(stream_id=3, data=b"hello") # client receives MAX_STREAM_DATA: 3, 1 - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_max_stream_data_frame( client_receive_context(client), QuicFrameType.MAX_STREAM_DATA, Buffer(data=b"\x03\x01"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.STREAM_STATE_ERROR) - self.assertEqual(cm.exception.frame_type, QuicFrameType.MAX_STREAM_DATA) - self.assertEqual(cm.exception.reason_phrase, "Stream is receive-only") + assert cm.value.error_code == QuicErrorCode.STREAM_STATE_ERROR + assert cm.value.frame_type == QuicFrameType.MAX_STREAM_DATA + assert cm.value.reason_phrase == "Stream is receive-only" def test_handle_max_streams_bidi_frame(self): with client_and_server() as (client, server): - self.assertEqual(client._remote_max_streams_bidi, 128) + assert client._remote_max_streams_bidi == 128 # client receives MAX_STREAMS_BIDI raising limit client._handle_max_streams_bidi_frame( @@ -1597,7 +1545,7 @@ def test_handle_max_streams_bidi_frame(self): QuicFrameType.MAX_STREAMS_BIDI, Buffer(data=encode_uint_var(129)), ) - self.assertEqual(client._remote_max_streams_bidi, 129) + assert client._remote_max_streams_bidi == 129 # client receives MAX_STREAMS_BIDI lowering limit client._handle_max_streams_bidi_frame( @@ -1605,27 +1553,23 @@ def test_handle_max_streams_bidi_frame(self): QuicFrameType.MAX_STREAMS_BIDI, Buffer(data=encode_uint_var(127)), ) - self.assertEqual(client._remote_max_streams_bidi, 129) + assert client._remote_max_streams_bidi == 129 # client receives invalid MAX_STREAMS_BIDI - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_max_streams_bidi_frame( client_receive_context(client), QuicFrameType.MAX_STREAMS_BIDI, Buffer(data=encode_uint_var(STREAM_COUNT_MAX + 1)), ) - self.assertEqual( - cm.exception.error_code, - QuicErrorCode.FRAME_ENCODING_ERROR, - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.MAX_STREAMS_BIDI) - self.assertEqual( - cm.exception.reason_phrase, "Maximum Streams cannot exceed 2^60" - ) + assert cm.value.error_code == \ + QuicErrorCode.FRAME_ENCODING_ERROR + assert cm.value.frame_type == QuicFrameType.MAX_STREAMS_BIDI + assert cm.value.reason_phrase == "Maximum Streams cannot exceed 2^60" def test_handle_max_streams_uni_frame(self): with client_and_server() as (client, server): - self.assertEqual(client._remote_max_streams_uni, 128) + assert client._remote_max_streams_uni == 128 # client receives MAX_STREAMS_UNI raising limit client._handle_max_streams_uni_frame( @@ -1633,7 +1577,7 @@ def test_handle_max_streams_uni_frame(self): QuicFrameType.MAX_STREAMS_UNI, Buffer(data=encode_uint_var(129)), ) - self.assertEqual(client._remote_max_streams_uni, 129) + assert client._remote_max_streams_uni == 129 # client receives MAX_STREAMS_UNI raising limit client._handle_max_streams_uni_frame( @@ -1641,23 +1585,19 @@ def test_handle_max_streams_uni_frame(self): QuicFrameType.MAX_STREAMS_UNI, Buffer(data=encode_uint_var(127)), ) - self.assertEqual(client._remote_max_streams_uni, 129) + assert client._remote_max_streams_uni == 129 # client receives invalid MAX_STREAMS_UNI - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_max_streams_uni_frame( client_receive_context(client), QuicFrameType.MAX_STREAMS_UNI, Buffer(data=encode_uint_var(STREAM_COUNT_MAX + 1)), ) - self.assertEqual( - cm.exception.error_code, - QuicErrorCode.FRAME_ENCODING_ERROR, - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.MAX_STREAMS_UNI) - self.assertEqual( - cm.exception.reason_phrase, "Maximum Streams cannot exceed 2^60" - ) + assert cm.value.error_code == \ + QuicErrorCode.FRAME_ENCODING_ERROR + assert cm.value.frame_type == QuicFrameType.MAX_STREAMS_UNI + assert cm.value.reason_phrase == "Maximum Streams cannot exceed 2^60" def test_handle_new_connection_id_duplicate(self): with client_and_server() as (client, server): @@ -1670,29 +1610,23 @@ def test_handle_new_connection_id_duplicate(self): buf, ) - self.assertEqual(client._peer_cid.sequence_number, 0) - self.assertEqual( - sequence_numbers(client._peer_cid_available), [1, 2, 3, 4, 5, 6, 7] - ) + assert client._peer_cid.sequence_number == 0 + assert sequence_numbers(client._peer_cid_available) == [1, 2, 3, 4, 5, 6, 7] def test_handle_new_connection_id_over_limit(self): with client_and_server() as (client, server): buf = new_connection_id(sequence_number=8) # client receives NEW_CONNECTION_ID - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_new_connection_id_frame( client_receive_context(client), QuicFrameType.NEW_CONNECTION_ID, buf, ) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.CONNECTION_ID_LIMIT_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.NEW_CONNECTION_ID) - self.assertEqual( - cm.exception.reason_phrase, "Too many active connection IDs" - ) + assert cm.value.error_code == QuicErrorCode.CONNECTION_ID_LIMIT_ERROR + assert cm.value.frame_type == QuicFrameType.NEW_CONNECTION_ID + assert cm.value.reason_phrase == "Too many active connection IDs" def test_handle_new_connection_id_with_retire_prior_to(self): with client_and_server() as (client, server): @@ -1705,10 +1639,8 @@ def test_handle_new_connection_id_with_retire_prior_to(self): buf, ) - self.assertEqual(client._peer_cid.sequence_number, 2) - self.assertEqual( - sequence_numbers(client._peer_cid_available), [3, 4, 5, 6, 7, 8] - ) + assert client._peer_cid.sequence_number == 2 + assert sequence_numbers(client._peer_cid_available) == [3, 4, 5, 6, 7, 8] def test_handle_new_connection_id_with_retire_prior_to_lower(self): with client_and_server() as (client, server): @@ -1719,8 +1651,8 @@ def test_handle_new_connection_id_with_retire_prior_to_lower(self): QuicFrameType.NEW_CONNECTION_ID, buf, ) - self.assertEqual(client._peer_cid.sequence_number, 80) - self.assertEqual(sequence_numbers(client._peer_cid_available), []) + assert client._peer_cid.sequence_number == 80 + assert sequence_numbers(client._peer_cid_available) == [] buf = new_connection_id(sequence_number=30, retire_prior_to=30) # client receives NEW_CONNECTION_ID client._handle_new_connection_id_frame( @@ -1728,8 +1660,8 @@ def test_handle_new_connection_id_with_retire_prior_to_lower(self): QuicFrameType.NEW_CONNECTION_ID, buf, ) - self.assertEqual(client._peer_cid.sequence_number, 80) - self.assertEqual(sequence_numbers(client._peer_cid_available), []) + assert client._peer_cid.sequence_number == 80 + assert sequence_numbers(client._peer_cid_available) == [] def test_handle_excessive_new_connection_id_retires(self): with client_and_server() as (client, server): @@ -1746,25 +1678,21 @@ def test_handle_excessive_new_connection_id_retires(self): ) # So far, so good! We should be at the (default) limit of 4*8 pending # retirements. - self.assertEqual(len(client._retire_connection_ids), 32) + assert len(client._retire_connection_ids) == 32 # Now we will go one too many! sequence_number = 8 + 25 buf = new_connection_id( sequence_number=sequence_number, retire_prior_to=sequence_number ) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_new_connection_id_frame( client_receive_context(client), QuicFrameType.NEW_CONNECTION_ID, buf, ) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.CONNECTION_ID_LIMIT_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.NEW_CONNECTION_ID) - self.assertEqual( - cm.exception.reason_phrase, "Too many pending retired connection IDs" - ) + assert cm.value.error_code == QuicErrorCode.CONNECTION_ID_LIMIT_ERROR + assert cm.value.frame_type == QuicFrameType.NEW_CONNECTION_ID + assert cm.value.reason_phrase == "Too many pending retired connection IDs" def test_handle_new_connection_id_with_connection_id_invalid(self): with client_and_server() as (client, server): @@ -1773,42 +1701,34 @@ def test_handle_new_connection_id_with_connection_id_invalid(self): ) # client receives NEW_CONNECTION_ID - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_new_connection_id_frame( client_receive_context(client), QuicFrameType.NEW_CONNECTION_ID, buf, ) - self.assertEqual( - cm.exception.error_code, - QuicErrorCode.FRAME_ENCODING_ERROR, - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.NEW_CONNECTION_ID) - self.assertEqual( - cm.exception.reason_phrase, - "Length must be greater than 0 and less than 20", - ) + assert cm.value.error_code == \ + QuicErrorCode.FRAME_ENCODING_ERROR + assert cm.value.frame_type == QuicFrameType.NEW_CONNECTION_ID + assert cm.value.reason_phrase == \ + "Length must be greater than 0 and less than 20" def test_handle_new_connection_id_with_retire_prior_to_invalid(self): with client_and_server() as (client, server): buf = new_connection_id(sequence_number=8, retire_prior_to=9) # client receives NEW_CONNECTION_ID - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_new_connection_id_frame( client_receive_context(client), QuicFrameType.NEW_CONNECTION_ID, buf, ) - self.assertEqual( - cm.exception.error_code, - QuicErrorCode.PROTOCOL_VIOLATION, - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.NEW_CONNECTION_ID) - self.assertEqual( - cm.exception.reason_phrase, - "Retire Prior To is greater than Sequence Number", - ) + assert cm.value.error_code == \ + QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.NEW_CONNECTION_ID + assert cm.value.reason_phrase == \ + "Retire Prior To is greater than Sequence Number" def test_handle_new_token_frame(self): with client_and_server() as (client, server): @@ -1822,17 +1742,15 @@ def test_handle_new_token_frame(self): def test_handle_new_token_frame_from_client(self): with client_and_server() as (client, server): # server receives NEW_TOKEN - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: server._handle_new_token_frame( client_receive_context(client), QuicFrameType.NEW_TOKEN, Buffer(data=binascii.unhexlify("080102030405060708")), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual(cm.exception.frame_type, QuicFrameType.NEW_TOKEN) - self.assertEqual( - cm.exception.reason_phrase, "Clients must not send NEW_TOKEN frames" - ) + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.NEW_TOKEN + assert cm.value.reason_phrase == "Clients must not send NEW_TOKEN frames" def test_handle_path_challenge_frame(self): with client_and_server() as (client, server): @@ -1842,11 +1760,11 @@ def test_handle_path_challenge_frame(self): server.receive_datagram(data, ("1.2.3.4", 2345), now=time.time()) # check paths - self.assertEqual(len(server._network_paths), 2) - self.assertEqual(server._network_paths[0].addr, ("1.2.3.4", 2345)) - self.assertFalse(server._network_paths[0].is_validated) - self.assertEqual(server._network_paths[1].addr, ("1.2.3.4", 1234)) - self.assertTrue(server._network_paths[1].is_validated) + assert len(server._network_paths) == 2 + assert server._network_paths[0].addr == ("1.2.3.4", 2345) + assert not server._network_paths[0].is_validated + assert server._network_paths[1].addr == ("1.2.3.4", 1234) + assert server._network_paths[1].is_validated # server sends PATH_CHALLENGE and receives PATH_RESPONSE for data, addr in server.datagrams_to_send(now=time.time()): @@ -1855,10 +1773,10 @@ def test_handle_path_challenge_frame(self): server.receive_datagram(data, ("1.2.3.4", 2345), now=time.time()) # check paths - self.assertEqual(server._network_paths[0].addr, ("1.2.3.4", 2345)) - self.assertTrue(server._network_paths[0].is_validated) - self.assertEqual(server._network_paths[1].addr, ("1.2.3.4", 1234)) - self.assertTrue(server._network_paths[1].is_validated) + assert server._network_paths[0].addr == ("1.2.3.4", 2345) + assert server._network_paths[0].is_validated + assert server._network_paths[1].addr == ("1.2.3.4", 1234) + assert server._network_paths[1].is_validated def test_handle_path_challenge_response_on_different_path(self): with client_and_server() as (client, server): @@ -1867,11 +1785,11 @@ def test_handle_path_challenge_response_on_different_path(self): for data, addr in client.datagrams_to_send(now=time.time()): server.receive_datagram(data, ("1.2.3.4", 2345), now=time.time()) # check paths - self.assertEqual(len(server._network_paths), 2) - self.assertEqual(server._network_paths[0].addr, ("1.2.3.4", 2345)) - self.assertFalse(server._network_paths[0].is_validated) - self.assertEqual(server._network_paths[1].addr, ("1.2.3.4", 1234)) - self.assertTrue(server._network_paths[1].is_validated) + assert len(server._network_paths) == 2 + assert server._network_paths[0].addr == ("1.2.3.4", 2345) + assert not server._network_paths[0].is_validated + assert server._network_paths[1].addr == ("1.2.3.4", 1234) + assert server._network_paths[1].is_validated # server sends PATH_CHALLENGE and receives PATH_RESPONSE on the 1234 # path instead of the expected 2345 path. for data, addr in server.datagrams_to_send(now=time.time()): @@ -1880,10 +1798,10 @@ def test_handle_path_challenge_response_on_different_path(self): server.receive_datagram(data, ("1.2.3.4", 1234), now=time.time()) # check paths; note that the order is backwards from the prior test # as receiving on 1234 promotes it to first in the list - self.assertEqual(server._network_paths[0].addr, ("1.2.3.4", 1234)) - self.assertTrue(server._network_paths[0].is_validated) - self.assertEqual(server._network_paths[1].addr, ("1.2.3.4", 2345)) - self.assertTrue(server._network_paths[1].is_validated) + assert server._network_paths[0].addr == ("1.2.3.4", 1234) + assert server._network_paths[0].is_validated + assert server._network_paths[1].addr == ("1.2.3.4", 2345) + assert server._network_paths[1].is_validated def test_local_path_challenges_are_bounded(self): with client_and_server() as (client, server): @@ -1891,24 +1809,22 @@ def test_local_path_challenges_are_bounded(self): server._add_local_challenge( int.to_bytes(i, 8, "big"), QuicNetworkPath(f"1.2.3.{i}") ) - self.assertEqual(len(server._local_challenges), MAX_LOCAL_CHALLENGES) + assert len(server._local_challenges) == MAX_LOCAL_CHALLENGES for i in range(2, MAX_LOCAL_CHALLENGES + 2): - self.assertEqual( - server._local_challenges[int.to_bytes(i, 8, "big")].addr, - f"1.2.3.{i}", - ) + assert server._local_challenges[int.to_bytes(i, 8, "big")].addr == \ + f"1.2.3.{i}" def test_handle_path_response_frame_bad(self): with client_and_server() as (client, server): # server receives unsolicited PATH_RESPONSE - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: server._handle_path_response_frame( client_receive_context(client), QuicFrameType.PATH_RESPONSE, Buffer(data=b"\x11\x22\x33\x44\x55\x66\x77\x88"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual(cm.exception.frame_type, QuicFrameType.PATH_RESPONSE) + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.PATH_RESPONSE def test_handle_padding_frame(self): client = create_standalone_client(self) @@ -1918,21 +1834,21 @@ def test_handle_padding_frame(self): client._handle_padding_frame( client_receive_context(client), QuicFrameType.PADDING, buf ) - self.assertEqual(buf.tell(), 0) + assert buf.tell() == 0 # padding until end buf = Buffer(data=bytes(10)) client._handle_padding_frame( client_receive_context(client), QuicFrameType.PADDING, buf ) - self.assertEqual(buf.tell(), 10) + assert buf.tell() == 10 # padding then something else buf = Buffer(data=bytes(10) + b"\x01") client._handle_padding_frame( client_receive_context(client), QuicFrameType.PADDING, buf ) - self.assertEqual(buf.tell(), 10) + assert buf.tell() == 10 def test_handle_reset_stream_frame(self): stream_id = 0 @@ -1953,9 +1869,9 @@ def test_handle_reset_stream_frame(self): ) event = client.next_event() - self.assertEqual(type(event), events.StreamReset) - self.assertEqual(event.error_code, QuicErrorCode.INTERNAL_ERROR) - self.assertEqual(event.stream_id, stream_id) + assert type(event) == events.StreamReset + assert event.error_code == QuicErrorCode.INTERNAL_ERROR + assert event.stream_id == stream_id def test_handle_reset_stream_frame_final_size_error(self): stream_id = 0 @@ -1976,12 +1892,12 @@ def test_handle_reset_stream_frame_final_size_error(self): ) event = client.next_event() - self.assertEqual(type(event), events.StreamReset) - self.assertEqual(event.error_code, QuicErrorCode.NO_ERROR) - self.assertEqual(event.stream_id, stream_id) + assert type(event) == events.StreamReset + assert event.error_code == QuicErrorCode.NO_ERROR + assert event.stream_id == stream_id # client receives RESET_STREAM at offset 5 - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_reset_stream_frame( client_receive_context(client), QuicFrameType.RESET_STREAM, @@ -1991,9 +1907,9 @@ def test_handle_reset_stream_frame_final_size_error(self): + encode_uint_var(5) ), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.FINAL_SIZE_ERROR) - self.assertEqual(cm.exception.frame_type, QuicFrameType.RESET_STREAM) - self.assertEqual(cm.exception.reason_phrase, "Cannot change final size") + assert cm.value.error_code == QuicErrorCode.FINAL_SIZE_ERROR + assert cm.value.frame_type == QuicFrameType.RESET_STREAM + assert cm.value.reason_phrase == "Cannot change final size" def test_handle_reset_stream_frame_over_max_data(self): stream_id = 0 @@ -2006,7 +1922,7 @@ def test_handle_reset_stream_frame_over_max_data(self): client._local_max_data.used = client._local_max_data.value # client receives RESET_STREAM frame - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_reset_stream_frame( client_receive_context(client), QuicFrameType.RESET_STREAM, @@ -2016,9 +1932,9 @@ def test_handle_reset_stream_frame_over_max_data(self): + encode_uint_var(1) ), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.FLOW_CONTROL_ERROR) - self.assertEqual(cm.exception.frame_type, QuicFrameType.RESET_STREAM) - self.assertEqual(cm.exception.reason_phrase, "Over connection data limit") + assert cm.value.error_code == QuicErrorCode.FLOW_CONTROL_ERROR + assert cm.value.frame_type == QuicFrameType.RESET_STREAM + assert cm.value.reason_phrase == "Over connection data limit" def test_handle_reset_stream_frame_over_max_stream_data(self): stream_id = 0 @@ -2028,7 +1944,7 @@ def test_handle_reset_stream_frame_over_max_stream_data(self): consume_events(client) # client receives STREAM frame - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_reset_stream_frame( client_receive_context(client), QuicFrameType.RESET_STREAM, @@ -2038,9 +1954,9 @@ def test_handle_reset_stream_frame_over_max_stream_data(self): + encode_uint_var(client._local_max_stream_data_bidi_local + 1) ), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.FLOW_CONTROL_ERROR) - self.assertEqual(cm.exception.frame_type, QuicFrameType.RESET_STREAM) - self.assertEqual(cm.exception.reason_phrase, "Over stream data limit") + assert cm.value.error_code == QuicErrorCode.FLOW_CONTROL_ERROR + assert cm.value.frame_type == QuicFrameType.RESET_STREAM + assert cm.value.reason_phrase == "Over stream data limit" def test_handle_reset_stream_frame_send_only(self): with client_and_server() as (client, server): @@ -2048,15 +1964,15 @@ def test_handle_reset_stream_frame_send_only(self): client.send_stream_data(stream_id=2, data=b"hello") # client receives RESET_STREAM - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_reset_stream_frame( client_receive_context(client), QuicFrameType.RESET_STREAM, Buffer(data=binascii.unhexlify("021100")), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.STREAM_STATE_ERROR) - self.assertEqual(cm.exception.frame_type, QuicFrameType.RESET_STREAM) - self.assertEqual(cm.exception.reason_phrase, "Stream is send-only") + assert cm.value.error_code == QuicErrorCode.STREAM_STATE_ERROR + assert cm.value.frame_type == QuicFrameType.RESET_STREAM + assert cm.value.reason_phrase == "Stream is send-only" def test_handle_reset_stream_frame_twice(self): stream_id = 3 @@ -2076,24 +1992,22 @@ def test_handle_reset_stream_frame_twice(self): client._payload_received(client_receive_context(client), reset_stream_data) event = client.next_event() - self.assertEqual(type(event), events.StreamReset) - self.assertEqual(event.error_code, QuicErrorCode.INTERNAL_ERROR) - self.assertEqual(event.stream_id, stream_id) + assert type(event) == events.StreamReset + assert event.error_code == QuicErrorCode.INTERNAL_ERROR + assert event.stream_id == stream_id # stream gets discarded - self.assertEqual(drop(client), 0) + assert drop(client) == 0 # client receives RESET_STREAM again client._payload_received(client_receive_context(client), reset_stream_data) event = client.next_event() - self.assertIsNone(event) + assert event is None def test_handle_retire_connection_id_frame(self): with client_and_server() as (client, server): - self.assertEqual( - sequence_numbers(client._host_cids), [0, 1, 2, 3, 4, 5, 6, 7] - ) + assert sequence_numbers(client._host_cids) == [0, 1, 2, 3, 4, 5, 6, 7] # client receives RETIRE_CONNECTION_ID client._handle_retire_connection_id_frame( @@ -2101,57 +2015,39 @@ def test_handle_retire_connection_id_frame(self): QuicFrameType.RETIRE_CONNECTION_ID, Buffer(data=b"\x02"), ) - self.assertEqual( - sequence_numbers(client._host_cids), [0, 1, 3, 4, 5, 6, 7, 8] - ) + assert sequence_numbers(client._host_cids) == [0, 1, 3, 4, 5, 6, 7, 8] def test_handle_retire_connection_id_frame_current_cid(self): with client_and_server() as (client, server): - self.assertEqual( - sequence_numbers(client._host_cids), [0, 1, 2, 3, 4, 5, 6, 7] - ) + assert sequence_numbers(client._host_cids) == [0, 1, 2, 3, 4, 5, 6, 7] # client receives RETIRE_CONNECTION_ID for the current CID - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_retire_connection_id_frame( client_receive_context(client), QuicFrameType.RETIRE_CONNECTION_ID, Buffer(data=b"\x00"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual( - cm.exception.frame_type, QuicFrameType.RETIRE_CONNECTION_ID - ) - self.assertEqual( - cm.exception.reason_phrase, "Cannot retire current connection ID" - ) - self.assertEqual( - sequence_numbers(client._host_cids), [0, 1, 2, 3, 4, 5, 6, 7] - ) + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.RETIRE_CONNECTION_ID + assert cm.value.reason_phrase == "Cannot retire current connection ID" + assert sequence_numbers(client._host_cids) == [0, 1, 2, 3, 4, 5, 6, 7] def test_handle_retire_connection_id_frame_invalid_sequence_number(self): with client_and_server() as (client, server): - self.assertEqual( - sequence_numbers(client._host_cids), [0, 1, 2, 3, 4, 5, 6, 7] - ) + assert sequence_numbers(client._host_cids) == [0, 1, 2, 3, 4, 5, 6, 7] # client receives RETIRE_CONNECTION_ID - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_retire_connection_id_frame( client_receive_context(client), QuicFrameType.RETIRE_CONNECTION_ID, Buffer(data=b"\x08"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual( - cm.exception.frame_type, QuicFrameType.RETIRE_CONNECTION_ID - ) - self.assertEqual( - cm.exception.reason_phrase, "Cannot retire unknown connection ID" - ) - self.assertEqual( - sequence_numbers(client._host_cids), [0, 1, 2, 3, 4, 5, 6, 7] - ) + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.RETIRE_CONNECTION_ID + assert cm.value.reason_phrase == "Cannot retire unknown connection ID" + assert sequence_numbers(client._host_cids) == [0, 1, 2, 3, 4, 5, 6, 7] def test_handle_stop_sending_frame(self): with client_and_server() as (client, server): @@ -2166,17 +2062,17 @@ def test_handle_stop_sending_frame(self): ) # check events - self.assertEqual(type(client.next_event()), events.ProtocolNegotiated) - self.assertEqual(type(client.next_event()), events.HandshakeCompleted) + assert type(client.next_event()) == events.ProtocolNegotiated + assert type(client.next_event()) == events.HandshakeCompleted for i in range(7): - self.assertEqual(type(client.next_event()), events.ConnectionIdIssued) + assert type(client.next_event()) == events.ConnectionIdIssued event = client.next_event() - self.assertEqual(type(event), events.StopSendingReceived) - self.assertEqual(event.stream_id, 0) - self.assertEqual(event.error_code, 0x11) + assert type(event) == events.StopSendingReceived + assert event.stream_id == 0 + assert event.error_code == 0x11 - self.assertIsNone(client.next_event()) + assert client.next_event() is None def test_handle_stop_sending_frame_receive_only(self): with client_and_server() as (client, server): @@ -2184,15 +2080,15 @@ def test_handle_stop_sending_frame_receive_only(self): server.send_stream_data(stream_id=3, data=b"hello") # client receives STOP_SENDING - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_stop_sending_frame( client_receive_context(client), QuicFrameType.STOP_SENDING, Buffer(data=b"\x03\x11"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.STREAM_STATE_ERROR) - self.assertEqual(cm.exception.frame_type, QuicFrameType.STOP_SENDING) - self.assertEqual(cm.exception.reason_phrase, "Stream is receive-only") + assert cm.value.error_code == QuicErrorCode.STREAM_STATE_ERROR + assert cm.value.frame_type == QuicFrameType.STOP_SENDING + assert cm.value.reason_phrase == "Stream is receive-only" def test_handle_stream_frame_final_size_error(self): with client_and_server() as (client, server): @@ -2211,7 +2107,7 @@ def test_handle_stream_frame_final_size_error(self): ) # client receives FIN at offset 5 - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_stream_frame( client_receive_context(client), frame_type, @@ -2221,16 +2117,16 @@ def test_handle_stream_frame_final_size_error(self): + encode_uint_var(0) ), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.FINAL_SIZE_ERROR) - self.assertEqual(cm.exception.frame_type, frame_type) - self.assertEqual(cm.exception.reason_phrase, "Cannot change final size") + assert cm.value.error_code == QuicErrorCode.FINAL_SIZE_ERROR + assert cm.value.frame_type == frame_type + assert cm.value.reason_phrase == "Cannot change final size" def test_handle_stream_frame_over_largest_offset(self): with client_and_server() as (client, server): # client receives offset + length > 2^62 - 1 frame_type = QuicFrameType.STREAM_BASE | 6 stream_id = 1 - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_stream_frame( client_receive_context(client), frame_type, @@ -2240,13 +2136,9 @@ def test_handle_stream_frame_over_largest_offset(self): + encode_uint_var(1) ), ) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.FRAME_ENCODING_ERROR - ) - self.assertEqual(cm.exception.frame_type, frame_type) - self.assertEqual( - cm.exception.reason_phrase, "offset + length cannot exceed 2^62 - 1" - ) + assert cm.value.error_code == QuicErrorCode.FRAME_ENCODING_ERROR + assert cm.value.frame_type == frame_type + assert cm.value.reason_phrase == "offset + length cannot exceed 2^62 - 1" def test_handle_stream_frame_over_max_data(self): with client_and_server() as (client, server): @@ -2256,22 +2148,22 @@ def test_handle_stream_frame_over_max_data(self): # client receives STREAM frame frame_type = QuicFrameType.STREAM_BASE | 4 stream_id = 1 - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_stream_frame( client_receive_context(client), frame_type, Buffer(data=encode_uint_var(stream_id) + encode_uint_var(1)), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.FLOW_CONTROL_ERROR) - self.assertEqual(cm.exception.frame_type, frame_type) - self.assertEqual(cm.exception.reason_phrase, "Over connection data limit") + assert cm.value.error_code == QuicErrorCode.FLOW_CONTROL_ERROR + assert cm.value.frame_type == frame_type + assert cm.value.reason_phrase == "Over connection data limit" def test_handle_stream_frame_over_max_stream_data(self): with client_and_server() as (client, server): # client receives STREAM frame frame_type = QuicFrameType.STREAM_BASE | 4 stream_id = 1 - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_stream_frame( client_receive_context(client), frame_type, @@ -2280,14 +2172,14 @@ def test_handle_stream_frame_over_max_stream_data(self): + encode_uint_var(client._local_max_stream_data_bidi_remote + 1) ), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.FLOW_CONTROL_ERROR) - self.assertEqual(cm.exception.frame_type, frame_type) - self.assertEqual(cm.exception.reason_phrase, "Over stream data limit") + assert cm.value.error_code == QuicErrorCode.FLOW_CONTROL_ERROR + assert cm.value.frame_type == frame_type + assert cm.value.reason_phrase == "Over stream data limit" def test_handle_stream_frame_over_max_streams(self): with client_and_server() as (client, server): # client receives STREAM frame - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_stream_frame( client_receive_context(client), QuicFrameType.STREAM_BASE, @@ -2295,9 +2187,9 @@ def test_handle_stream_frame_over_max_streams(self): data=encode_uint_var(client._local_max_stream_data_uni * 4 + 3) ), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.STREAM_LIMIT_ERROR) - self.assertEqual(cm.exception.frame_type, QuicFrameType.STREAM_BASE) - self.assertEqual(cm.exception.reason_phrase, "Too many streams open") + assert cm.value.error_code == QuicErrorCode.STREAM_LIMIT_ERROR + assert cm.value.frame_type == QuicFrameType.STREAM_BASE + assert cm.value.reason_phrase == "Too many streams open" def test_handle_stream_frame_send_only(self): with client_and_server() as (client, server): @@ -2305,28 +2197,28 @@ def test_handle_stream_frame_send_only(self): client.send_stream_data(stream_id=2, data=b"hello") # client receives STREAM frame - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_stream_frame( client_receive_context(client), QuicFrameType.STREAM_BASE, Buffer(data=b"\x02"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.STREAM_STATE_ERROR) - self.assertEqual(cm.exception.frame_type, QuicFrameType.STREAM_BASE) - self.assertEqual(cm.exception.reason_phrase, "Stream is send-only") + assert cm.value.error_code == QuicErrorCode.STREAM_STATE_ERROR + assert cm.value.frame_type == QuicFrameType.STREAM_BASE + assert cm.value.reason_phrase == "Stream is send-only" def test_handle_stream_frame_wrong_initiator(self): with client_and_server() as (client, server): # client receives STREAM frame - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_stream_frame( client_receive_context(client), QuicFrameType.STREAM_BASE, Buffer(data=b"\x00"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.STREAM_STATE_ERROR) - self.assertEqual(cm.exception.frame_type, QuicFrameType.STREAM_BASE) - self.assertEqual(cm.exception.reason_phrase, "Wrong stream initiator") + assert cm.value.error_code == QuicErrorCode.STREAM_STATE_ERROR + assert cm.value.frame_type == QuicFrameType.STREAM_BASE + assert cm.value.reason_phrase == "Wrong stream initiator" def test_handle_stream_data_blocked_frame(self): with client_and_server() as (client, server): @@ -2346,15 +2238,15 @@ def test_handle_stream_data_blocked_frame_send_only(self): client.send_stream_data(stream_id=2, data=b"hello") # client receives STREAM_DATA_BLOCKED - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_stream_data_blocked_frame( client_receive_context(client), QuicFrameType.STREAM_DATA_BLOCKED, Buffer(data=b"\x02\x01"), ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.STREAM_STATE_ERROR) - self.assertEqual(cm.exception.frame_type, QuicFrameType.STREAM_DATA_BLOCKED) - self.assertEqual(cm.exception.reason_phrase, "Stream is send-only") + assert cm.value.error_code == QuicErrorCode.STREAM_STATE_ERROR + assert cm.value.frame_type == QuicFrameType.STREAM_DATA_BLOCKED + assert cm.value.reason_phrase == "Stream is send-only" def test_handle_streams_blocked_uni_frame(self): with client_and_server() as (client, server): @@ -2366,20 +2258,16 @@ def test_handle_streams_blocked_uni_frame(self): ) # client receives invalid STREAMS_BLOCKED_UNI - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._handle_streams_blocked_frame( client_receive_context(client), QuicFrameType.STREAMS_BLOCKED_UNI, Buffer(data=encode_uint_var(STREAM_COUNT_MAX + 1)), ) - self.assertEqual( - cm.exception.error_code, - QuicErrorCode.FRAME_ENCODING_ERROR, - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.STREAMS_BLOCKED_UNI) - self.assertEqual( - cm.exception.reason_phrase, "Maximum Streams cannot exceed 2^60" - ) + assert cm.value.error_code == \ + QuicErrorCode.FRAME_ENCODING_ERROR + assert cm.value.frame_type == QuicFrameType.STREAMS_BLOCKED_UNI + assert cm.value.reason_phrase == "Maximum Streams cannot exceed 2^60" def test_parse_transport_parameters(self): client = create_standalone_client(self) @@ -2394,15 +2282,11 @@ def test_parse_transport_parameters(self): def test_parse_transport_parameters_malformed(self): client = create_standalone_client(self) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._parse_transport_parameters(b"0") - self.assertEqual( - cm.exception.error_code, QuicErrorCode.TRANSPORT_PARAMETER_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual( - cm.exception.reason_phrase, "Could not parse QUIC transport parameters" - ) + assert cm.value.error_code == QuicErrorCode.TRANSPORT_PARAMETER_ERROR + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert cm.value.reason_phrase == "Could not parse QUIC transport parameters" def test_parse_transport_parameters_with_bad_ack_delay_exponent(self): client = create_standalone_client(self) @@ -2413,13 +2297,11 @@ def test_parse_transport_parameters_with_bad_ack_delay_exponent(self): original_destination_connection_id=client.original_destination_connection_id, ) ) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._parse_transport_parameters(data) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.TRANSPORT_PARAMETER_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual(cm.exception.reason_phrase, "ack_delay_exponent must be <= 20") + assert cm.value.error_code == QuicErrorCode.TRANSPORT_PARAMETER_ERROR + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert cm.value.reason_phrase == "ack_delay_exponent must be <= 20" def test_parse_transport_parameters_with_bad_active_connection_id_limit(self): client = create_standalone_client(self) @@ -2431,16 +2313,12 @@ def test_parse_transport_parameters_with_bad_active_connection_id_limit(self): original_destination_connection_id=client.original_destination_connection_id, ) ) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._parse_transport_parameters(data) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.TRANSPORT_PARAMETER_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual( - cm.exception.reason_phrase, - "active_connection_id_limit must be no less than 2", - ) + assert cm.value.error_code == QuicErrorCode.TRANSPORT_PARAMETER_ERROR + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert cm.value.reason_phrase == \ + "active_connection_id_limit must be no less than 2" def test_parse_transport_parameters_with_bad_max_ack_delay(self): client = create_standalone_client(self) @@ -2451,13 +2329,11 @@ def test_parse_transport_parameters_with_bad_max_ack_delay(self): original_destination_connection_id=client.original_destination_connection_id, ) ) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._parse_transport_parameters(data) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.TRANSPORT_PARAMETER_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual(cm.exception.reason_phrase, "max_ack_delay must be < 2^14") + assert cm.value.error_code == QuicErrorCode.TRANSPORT_PARAMETER_ERROR + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert cm.value.reason_phrase == "max_ack_delay must be < 2^14" def test_parse_transport_parameters_with_bad_max_udp_payload_size(self): client = create_standalone_client(self) @@ -2468,15 +2344,11 @@ def test_parse_transport_parameters_with_bad_max_udp_payload_size(self): original_destination_connection_id=client.original_destination_connection_id, ) ) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._parse_transport_parameters(data) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.TRANSPORT_PARAMETER_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual( - cm.exception.reason_phrase, "max_udp_payload_size must be >= 1200" - ) + assert cm.value.error_code == QuicErrorCode.TRANSPORT_PARAMETER_ERROR + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert cm.value.reason_phrase == "max_udp_payload_size must be >= 1200" def test_parse_transport_parameters_with_bad_initial_source_connection_id(self): client = create_standalone_client(self) @@ -2488,15 +2360,11 @@ def test_parse_transport_parameters_with_bad_initial_source_connection_id(self): original_destination_connection_id=client.original_destination_connection_id, ) ) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._parse_transport_parameters(data) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.TRANSPORT_PARAMETER_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual( - cm.exception.reason_phrase, "initial_source_connection_id does not match" - ) + assert cm.value.error_code == QuicErrorCode.TRANSPORT_PARAMETER_ERROR + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert cm.value.reason_phrase == "initial_source_connection_id does not match" def test_parse_transport_parameters_with_bad_version_information_1(self): server = create_standalone_server(self) @@ -2508,17 +2376,13 @@ def test_parse_transport_parameters_with_bad_version_information_1(self): ) ) ) - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: server._parse_transport_parameters(data) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.TRANSPORT_PARAMETER_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual( - cm.exception.reason_phrase, - "version_information's chosen_version is not included in " - "available_versions", - ) + assert cm.value.error_code == QuicErrorCode.TRANSPORT_PARAMETER_ERROR + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert cm.value.reason_phrase == \ + "version_information's chosen_version is not included in " \ + "available_versions" def test_parse_transport_parameters_with_bad_version_information_2(self): server = create_standalone_server(self) @@ -2534,25 +2398,21 @@ def test_parse_transport_parameters_with_bad_version_information_2(self): ) ) server._crypto_packet_version = QuicProtocolVersion.VERSION_2 - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: server._parse_transport_parameters(data) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.VERSION_NEGOTIATION_ERROR - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual( - cm.exception.reason_phrase, - "version_information's chosen_version does not match the version in use", - ) + assert cm.value.error_code == QuicErrorCode.VERSION_NEGOTIATION_ERROR + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert cm.value.reason_phrase == \ + "version_information's chosen_version does not match the version in use" def test_payload_received_empty(self): with client_and_server() as (client, server): # client receives empty payload - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._payload_received(client_receive_context(client), b"") - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual(cm.exception.frame_type, QuicFrameType.PADDING) - self.assertEqual(cm.exception.reason_phrase, "Packet contains no frames") + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.PADDING + assert cm.value.reason_phrase == "Packet contains no frames" def test_payload_received_padding_only(self): with client_and_server() as (client, server): @@ -2560,166 +2420,162 @@ def test_payload_received_padding_only(self): is_ack_eliciting, is_probing = client._payload_received( client_receive_context(client), b"\x00" * 1200 ) - self.assertFalse(is_ack_eliciting) - self.assertTrue(is_probing) + assert not is_ack_eliciting + assert is_probing def test_payload_received_malformed_frame_type(self): with client_and_server() as (client, server): # client receives a malformed frame type - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._payload_received(client_receive_context(client), b"\xff") - self.assertEqual( - cm.exception.error_code, QuicErrorCode.FRAME_ENCODING_ERROR - ) - self.assertEqual(cm.exception.frame_type, None) - self.assertEqual(cm.exception.reason_phrase, "Malformed frame type") + assert cm.value.error_code == QuicErrorCode.FRAME_ENCODING_ERROR + assert cm.value.frame_type == None + assert cm.value.reason_phrase == "Malformed frame type" def test_payload_received_unknown_frame(self): with client_and_server() as (client, server): # client receives unknown frame - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._payload_received(client_receive_context(client), b"\x1f") - self.assertEqual(cm.exception.error_code, QuicErrorCode.FRAME_ENCODING_ERROR) - self.assertEqual(cm.exception.frame_type, 0x1F) - self.assertEqual(cm.exception.reason_phrase, "Unknown frame type") + assert cm.value.error_code == QuicErrorCode.FRAME_ENCODING_ERROR + assert cm.value.frame_type == 0x1F + assert cm.value.reason_phrase == "Unknown frame type" def test_payload_received_unexpected_frame(self): with client_and_server() as (client, server): # client receives CRYPTO frame in 0-RTT - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._payload_received( client_receive_context(client, epoch=tls.Epoch.ZERO_RTT), b"\x06" ) - self.assertEqual(cm.exception.error_code, QuicErrorCode.PROTOCOL_VIOLATION) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual(cm.exception.reason_phrase, "Unexpected frame type") + assert cm.value.error_code == QuicErrorCode.PROTOCOL_VIOLATION + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert cm.value.reason_phrase == "Unexpected frame type" def test_payload_received_malformed_frame(self): with client_and_server() as (client, server): # client receives malformed TRANSPORT_CLOSE frame - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: client._payload_received( client_receive_context(client), b"\x1c\x00\x01" ) - self.assertEqual( - cm.exception.error_code, QuicErrorCode.FRAME_ENCODING_ERROR - ) - self.assertEqual(cm.exception.frame_type, 0x1C) - self.assertEqual(cm.exception.reason_phrase, "Failed to parse frame") + assert cm.value.error_code == QuicErrorCode.FRAME_ENCODING_ERROR + assert cm.value.frame_type == 0x1C + assert cm.value.reason_phrase == "Failed to parse frame" def test_send_max_data_blocked_by_cc(self): with client_and_server() as (client, server): # check congestion control - self.assertEqual(client._loss.bytes_in_flight, 0) - self.assertGreaterEqual(client._loss.congestion_window, 13530) - self.assertLessEqual(client._loss.congestion_window, 16000) + assert client._loss.bytes_in_flight == 0 + assert client._loss.congestion_window >= 13530 + assert client._loss.congestion_window <= 16000 # artificially raise received data counter client._local_max_data_used = client._local_max_data - self.assertEqual(server._remote_max_data, 1048576) + assert server._remote_max_data == 1048576 # artificially raise bytes in flight client._loss._cc.bytes_in_flight = client._loss.congestion_window # MAX_DATA is not sent due to congestion control - self.assertEqual(drop(client), 0) + assert drop(client) == 0 def test_send_max_data_retransmit(self): with client_and_server() as (client, server): # artificially raise received data counter client._local_max_data.used = client._local_max_data.value - self.assertEqual(client._local_max_data.sent, 1048576) - self.assertEqual(client._local_max_data.used, 1048576) - self.assertEqual(client._local_max_data.value, 1048576) - self.assertEqual(server._remote_max_data, 1048576) + assert client._local_max_data.sent == 1048576 + assert client._local_max_data.used == 1048576 + assert client._local_max_data.value == 1048576 + assert server._remote_max_data == 1048576 # MAX_DATA is sent and lost - self.assertEqual(drop(client), 1) - self.assertEqual(client._local_max_data.sent, 2097152) - self.assertEqual(client._local_max_data.used, 1048576) - self.assertEqual(client._local_max_data.value, 2097152) - self.assertEqual(server._remote_max_data, 1048576) + assert drop(client) == 1 + assert client._local_max_data.sent == 2097152 + assert client._local_max_data.used == 1048576 + assert client._local_max_data.value == 2097152 + assert server._remote_max_data == 1048576 # MAX_DATA loss is detected client._on_connection_limit_delivery( QuicDeliveryState.LOST, client._local_max_data ) - self.assertEqual(client._local_max_data.sent, 0) - self.assertEqual(client._local_max_data.used, 1048576) - self.assertEqual(client._local_max_data.value, 2097152) + assert client._local_max_data.sent == 0 + assert client._local_max_data.used == 1048576 + assert client._local_max_data.value == 2097152 # MAX_DATA is retransmitted and acked - self.assertEqual(roundtrip(client, server), (1, 1)) - self.assertEqual(client._local_max_data.sent, 2097152) - self.assertEqual(client._local_max_data.used, 1048576) - self.assertEqual(client._local_max_data.value, 2097152) - self.assertEqual(server._remote_max_data, 2097152) + assert roundtrip(client, server) == (1, 1) + assert client._local_max_data.sent == 2097152 + assert client._local_max_data.used == 1048576 + assert client._local_max_data.value == 2097152 + assert server._remote_max_data == 2097152 def test_send_max_stream_data_retransmit(self): with client_and_server() as (client, server): # client creates bidirectional stream 0 stream = client._get_or_create_stream_for_send(stream_id=0) client.send_stream_data(0, b"hello") - self.assertEqual(stream.max_stream_data_local, 1048576) - self.assertEqual(stream.max_stream_data_local_sent, 1048576) - self.assertEqual(roundtrip(client, server), (1, 1)) + assert stream.max_stream_data_local == 1048576 + assert stream.max_stream_data_local_sent == 1048576 + assert roundtrip(client, server) == (1, 1) # server sends data, just before raising MAX_STREAM_DATA server.send_stream_data(0, b"Z" * 524288) # 1048576 // 2 for i in range(10): roundtrip(server, client) - self.assertEqual(stream.max_stream_data_local, 1048576) - self.assertEqual(stream.max_stream_data_local_sent, 1048576) + assert stream.max_stream_data_local == 1048576 + assert stream.max_stream_data_local_sent == 1048576 # server sends one more byte server.send_stream_data(0, b"Z") - self.assertEqual(transfer(server, client), 1) + assert transfer(server, client) == 1 # MAX_STREAM_DATA is sent and lost - self.assertEqual(drop(client), 1) - self.assertEqual(stream.max_stream_data_local, 2097152) - self.assertEqual(stream.max_stream_data_local_sent, 2097152) + assert drop(client) == 1 + assert stream.max_stream_data_local == 2097152 + assert stream.max_stream_data_local_sent == 2097152 client._on_max_stream_data_delivery(QuicDeliveryState.LOST, stream) - self.assertEqual(stream.max_stream_data_local, 2097152) - self.assertEqual(stream.max_stream_data_local_sent, 0) + assert stream.max_stream_data_local == 2097152 + assert stream.max_stream_data_local_sent == 0 # MAX_DATA is retransmitted and acked - self.assertEqual(roundtrip(client, server), (1, 1)) - self.assertEqual(stream.max_stream_data_local, 2097152) - self.assertEqual(stream.max_stream_data_local_sent, 2097152) + assert roundtrip(client, server) == (1, 1) + assert stream.max_stream_data_local == 2097152 + assert stream.max_stream_data_local_sent == 2097152 def test_send_max_streams_retransmit(self): with client_and_server() as (client, server): # client opens 65 streams client.send_stream_data(4 * 64, b"Z") - self.assertEqual(transfer(client, server), 1) - self.assertEqual(client._remote_max_streams_bidi, 128) - self.assertEqual(server._local_max_streams_bidi.sent, 128) - self.assertEqual(server._local_max_streams_bidi.used, 65) - self.assertEqual(server._local_max_streams_bidi.value, 128) + assert transfer(client, server) == 1 + assert client._remote_max_streams_bidi == 128 + assert server._local_max_streams_bidi.sent == 128 + assert server._local_max_streams_bidi.used == 65 + assert server._local_max_streams_bidi.value == 128 # MAX_STREAMS is sent and lost - self.assertEqual(drop(server), 1) - self.assertEqual(client._remote_max_streams_bidi, 128) - self.assertEqual(server._local_max_streams_bidi.sent, 256) - self.assertEqual(server._local_max_streams_bidi.used, 65) - self.assertEqual(server._local_max_streams_bidi.value, 256) + assert drop(server) == 1 + assert client._remote_max_streams_bidi == 128 + assert server._local_max_streams_bidi.sent == 256 + assert server._local_max_streams_bidi.used == 65 + assert server._local_max_streams_bidi.value == 256 # MAX_STREAMS loss is detected server._on_connection_limit_delivery( QuicDeliveryState.LOST, server._local_max_streams_bidi ) - self.assertEqual(client._remote_max_streams_bidi, 128) - self.assertEqual(server._local_max_streams_bidi.sent, 0) - self.assertEqual(server._local_max_streams_bidi.used, 65) - self.assertEqual(server._local_max_streams_bidi.value, 256) + assert client._remote_max_streams_bidi == 128 + assert server._local_max_streams_bidi.sent == 0 + assert server._local_max_streams_bidi.used == 65 + assert server._local_max_streams_bidi.value == 256 # MAX_STREAMS is retransmitted and acked - self.assertEqual(roundtrip(server, client), (1, 1)) - self.assertEqual(client._remote_max_streams_bidi, 256) - self.assertEqual(server._local_max_streams_bidi.sent, 256) - self.assertEqual(server._local_max_streams_bidi.used, 65) - self.assertEqual(server._local_max_streams_bidi.value, 256) + assert roundtrip(server, client) == (1, 1) + assert client._remote_max_streams_bidi == 256 + assert server._local_max_streams_bidi.sent == 256 + assert server._local_max_streams_bidi.used == 65 + assert server._local_max_streams_bidi.value == 256 def test_send_ping(self): with client_and_server() as (client, server): @@ -2727,12 +2583,12 @@ def test_send_ping(self): # client sends ping, server ACKs it client.send_ping(uid=12345) - self.assertEqual(roundtrip(client, server), (1, 1)) + assert roundtrip(client, server) == (1, 1) # check event event = client.next_event() - self.assertEqual(type(event), events.PingAcknowledged) - self.assertEqual(event.uid, 12345) + assert type(event) == events.PingAcknowledged + assert event.uid == 12345 def test_send_ping_retransmit(self): with client_and_server() as (client, server): @@ -2740,26 +2596,26 @@ def test_send_ping_retransmit(self): # client sends another ping, PING is lost client.send_ping(uid=12345) - self.assertEqual(drop(client), 1) + assert drop(client) == 1 # PING is retransmitted and acked client._on_ping_delivery(QuicDeliveryState.LOST, (12345,)) - self.assertEqual(roundtrip(client, server), (1, 1)) + assert roundtrip(client, server) == (1, 1) # check event event = client.next_event() - self.assertEqual(type(event), events.PingAcknowledged) - self.assertEqual(event.uid, 12345) + assert type(event) == events.PingAcknowledged + assert event.uid == 12345 def test_send_reset_stream(self): with client_and_server() as (client, server): # client creates bidirectional stream client.send_stream_data(0, b"hello") - self.assertEqual(roundtrip(client, server), (1, 1)) + assert roundtrip(client, server) == (1, 1) # client resets stream client.reset_stream(0, QuicErrorCode.NO_ERROR) - self.assertEqual(roundtrip(client, server), (1, 1)) + assert roundtrip(client, server) == (1, 1) def test_send_stop_sending(self): with client_and_server() as (client, server): @@ -2768,17 +2624,17 @@ def test_send_stop_sending(self): # client creates bidirectional stream client.send_stream_data(0, b"hello") - self.assertEqual(roundtrip(client, server), (1, 1)) + assert roundtrip(client, server) == (1, 1) # client sends STOP_SENDING frame client.stop_stream(0, QuicErrorCode.NO_ERROR) - self.assertEqual(roundtrip(client, server), (1, 1)) + assert roundtrip(client, server) == (1, 1) # client receives STREAM_RESET frame event = client.next_event() - self.assertEqual(type(event), events.StreamReset) - self.assertEqual(event.error_code, QuicErrorCode.NO_ERROR) - self.assertEqual(event.stream_id, 0) + assert type(event) == events.StreamReset + assert event.error_code == QuicErrorCode.NO_ERROR + assert event.stream_id == 0 def test_send_stop_sending_uni_stream(self): with client_and_server() as (client, server): @@ -2786,12 +2642,10 @@ def test_send_stop_sending_uni_stream(self): self.check_handshake(client=client, server=server) # client sends STOP_SENDING frame - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: client.stop_stream(2, QuicErrorCode.NO_ERROR) - self.assertEqual( - str(cm.exception), - "Cannot stop receiving on a local-initiated unidirectional stream", - ) + assert str(cm.value) == \ + "Cannot stop receiving on a local-initiated unidirectional stream" def test_send_stop_sending_unknown_stream(self): with client_and_server() as (client, server): @@ -2799,11 +2653,9 @@ def test_send_stop_sending_unknown_stream(self): self.check_handshake(client=client, server=server) # client sends STOP_SENDING frame - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: client.stop_stream(0, QuicErrorCode.NO_ERROR) - self.assertEqual( - str(cm.exception), "Cannot stop receiving on an unknown stream" - ) + assert str(cm.value) == "Cannot stop receiving on an unknown stream" def test_send_stream_data_over_max_streams_bidi(self): with client_and_server() as (client, server): @@ -2811,18 +2663,18 @@ def test_send_stream_data_over_max_streams_bidi(self): for i in range(128): stream_id = i * 4 client.send_stream_data(stream_id, b"") - self.assertFalse(client._streams[stream_id].is_blocked) - self.assertEqual(len(client._streams_blocked_bidi), 0) - self.assertEqual(len(client._streams_blocked_uni), 0) - self.assertEqual(roundtrip(client, server), (0, 0)) + assert not client._streams[stream_id].is_blocked + assert len(client._streams_blocked_bidi) == 0 + assert len(client._streams_blocked_uni) == 0 + assert roundtrip(client, server) == (0, 0) # create one too many -> STREAMS_BLOCKED stream_id = 128 * 4 client.send_stream_data(stream_id, b"") - self.assertTrue(client._streams[stream_id].is_blocked) - self.assertEqual(len(client._streams_blocked_bidi), 1) - self.assertEqual(len(client._streams_blocked_uni), 0) - self.assertEqual(roundtrip(client, server), (1, 1)) + assert client._streams[stream_id].is_blocked + assert len(client._streams_blocked_bidi) == 1 + assert len(client._streams_blocked_uni) == 0 + assert roundtrip(client, server) == (1, 1) # peer raises max streams client._handle_max_streams_bidi_frame( @@ -2830,7 +2682,7 @@ def test_send_stream_data_over_max_streams_bidi(self): QuicFrameType.MAX_STREAMS_BIDI, Buffer(data=encode_uint_var(129)), ) - self.assertFalse(client._streams[stream_id].is_blocked) + assert not client._streams[stream_id].is_blocked def test_send_stream_data_over_max_streams_uni(self): with client_and_server() as (client, server): @@ -2838,18 +2690,18 @@ def test_send_stream_data_over_max_streams_uni(self): for i in range(128): stream_id = i * 4 + 2 client.send_stream_data(stream_id, b"") - self.assertFalse(client._streams[stream_id].is_blocked) - self.assertEqual(len(client._streams_blocked_bidi), 0) - self.assertEqual(len(client._streams_blocked_uni), 0) - self.assertEqual(roundtrip(client, server), (0, 0)) + assert not client._streams[stream_id].is_blocked + assert len(client._streams_blocked_bidi) == 0 + assert len(client._streams_blocked_uni) == 0 + assert roundtrip(client, server) == (0, 0) # create one too many -> STREAMS_BLOCKED stream_id = 128 * 4 + 2 client.send_stream_data(stream_id, b"") - self.assertTrue(client._streams[stream_id].is_blocked) - self.assertEqual(len(client._streams_blocked_bidi), 0) - self.assertEqual(len(client._streams_blocked_uni), 1) - self.assertEqual(roundtrip(client, server), (1, 1)) + assert client._streams[stream_id].is_blocked + assert len(client._streams_blocked_bidi) == 0 + assert len(client._streams_blocked_uni) == 1 + assert roundtrip(client, server) == (1, 1) # peer raises max streams client._handle_max_streams_uni_frame( @@ -2857,86 +2709,78 @@ def test_send_stream_data_over_max_streams_uni(self): QuicFrameType.MAX_STREAMS_UNI, Buffer(data=encode_uint_var(129)), ) - self.assertFalse(client._streams[stream_id].is_blocked) + assert not client._streams[stream_id].is_blocked def test_send_stream_data_peer_initiated(self): with client_and_server() as (client, server): # server creates bidirectional stream server.send_stream_data(1, b"hello") - self.assertEqual(roundtrip(server, client), (1, 1)) + assert roundtrip(server, client) == (1, 1) # server creates unidirectional stream server.send_stream_data(3, b"hello") - self.assertEqual(roundtrip(server, client), (1, 1)) + assert roundtrip(server, client) == (1, 1) # client creates bidirectional stream client.send_stream_data(0, b"hello") - self.assertEqual(roundtrip(client, server), (1, 1)) + assert roundtrip(client, server) == (1, 1) # client sends data on server-initiated bidirectional stream client.send_stream_data(1, b"hello") - self.assertEqual(roundtrip(client, server), (1, 1)) + assert roundtrip(client, server) == (1, 1) # client creates unidirectional stream client.send_stream_data(2, b"hello") - self.assertEqual(roundtrip(client, server), (1, 1)) + assert roundtrip(client, server) == (1, 1) # client tries to reset server-initiated unidirectional stream - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: client.reset_stream(3, QuicErrorCode.NO_ERROR) - self.assertEqual( - str(cm.exception), - "Cannot send data on peer-initiated unidirectional stream", - ) + assert str(cm.value) == \ + "Cannot send data on peer-initiated unidirectional stream" # client tries to reset unknown server-initiated bidirectional stream - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: client.reset_stream(5, QuicErrorCode.NO_ERROR) - self.assertEqual( - str(cm.exception), "Cannot send data on unknown peer-initiated stream" - ) + assert str(cm.value) == "Cannot send data on unknown peer-initiated stream" # client tries to send data on server-initiated unidirectional stream - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: client.send_stream_data(3, b"hello") - self.assertEqual( - str(cm.exception), - "Cannot send data on peer-initiated unidirectional stream", - ) + assert str(cm.value) == \ + "Cannot send data on peer-initiated unidirectional stream" # client tries to send data on unknown server-initiated bidirectional stream - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: client.send_stream_data(5, b"hello") - self.assertEqual( - str(cm.exception), "Cannot send data on unknown peer-initiated stream" - ) + assert str(cm.value) == "Cannot send data on unknown peer-initiated stream" def test_stream_direction(self): with client_and_server() as (client, server): for off in [0, 4, 8]: # Client-Initiated, Bidirectional - self.assertTrue(client._stream_can_receive(off)) - self.assertTrue(client._stream_can_send(off)) - self.assertTrue(server._stream_can_receive(off)) - self.assertTrue(server._stream_can_send(off)) + assert client._stream_can_receive(off) + assert client._stream_can_send(off) + assert server._stream_can_receive(off) + assert server._stream_can_send(off) # Server-Initiated, Bidirectional - self.assertTrue(client._stream_can_receive(off + 1)) - self.assertTrue(client._stream_can_send(off + 1)) - self.assertTrue(server._stream_can_receive(off + 1)) - self.assertTrue(server._stream_can_send(off + 1)) + assert client._stream_can_receive(off + 1) + assert client._stream_can_send(off + 1) + assert server._stream_can_receive(off + 1) + assert server._stream_can_send(off + 1) # Client-Initiated, Unidirectional - self.assertFalse(client._stream_can_receive(off + 2)) - self.assertTrue(client._stream_can_send(off + 2)) - self.assertTrue(server._stream_can_receive(off + 2)) - self.assertFalse(server._stream_can_send(off + 2)) + assert not client._stream_can_receive(off + 2) + assert client._stream_can_send(off + 2) + assert server._stream_can_receive(off + 2) + assert not server._stream_can_send(off + 2) # Server-Initiated, Unidirectional - self.assertTrue(client._stream_can_receive(off + 3)) - self.assertFalse(client._stream_can_send(off + 3)) - self.assertFalse(server._stream_can_receive(off + 3)) - self.assertTrue(server._stream_can_send(off + 3)) + assert client._stream_can_receive(off + 3) + assert not client._stream_can_send(off + 3) + assert not server._stream_can_receive(off + 3) + assert server._stream_can_send(off + 3) def test_version_negotiation_fail(self): client = create_standalone_client(self) @@ -2951,15 +2795,13 @@ def test_version_negotiation_fail(self): SERVER_ADDR, now=time.time(), ) - self.assertEqual(drop(client), 0) + assert drop(client) == 0 event = client.next_event() - self.assertEqual(type(event), events.ConnectionTerminated) - self.assertEqual(event.error_code, QuicErrorCode.INTERNAL_ERROR) - self.assertEqual(event.frame_type, QuicFrameType.PADDING) - self.assertEqual( - event.reason_phrase, "Could not find a common protocol version" - ) + assert type(event) == events.ConnectionTerminated + assert event.error_code == QuicErrorCode.INTERNAL_ERROR + assert event.frame_type == QuicFrameType.PADDING + assert event.reason_phrase == "Could not find a common protocol version" def test_version_negotiation_ignore(self): client = create_standalone_client(self) @@ -2974,7 +2816,7 @@ def test_version_negotiation_ignore(self): SERVER_ADDR, now=time.time(), ) - self.assertEqual(drop(client), 0) + assert drop(client) == 0 def test_version_negotiation_ignore_server(self): server = create_standalone_server(self) @@ -3004,7 +2846,7 @@ def test_version_negotiation_ok(self): SERVER_ADDR, now=time.time(), ) - self.assertEqual(drop(client), 0) # todo: investigate! + assert drop(client) == 0# todo: investigate! def test_write_connection_close_early(self): client = create_standalone_client(self) @@ -3026,8 +2868,7 @@ def test_write_connection_close_early(self): reason_phrase="some reason", ) - self.assertEqual( - builder.quic_logger_frames, + assert builder.quic_logger_frames == \ [ { "error_code": QuicErrorCode.APPLICATION_ERROR, @@ -3036,9 +2877,8 @@ def test_write_connection_close_early(self): "raw_error_code": QuicErrorCode.APPLICATION_ERROR, "reason": "", "trigger_frame_type": QuicFrameType.PADDING, - } - ], - ) + } \ + ] def test_excessive_crypto_buffering(self): with client_and_server() as (client, server): @@ -3048,7 +2888,7 @@ def test_excessive_crypto_buffering(self): # how much buffering is needed. We send fragments of only 100 bytes # at offsets 10000, 20000, 30000 etc. highest_good_offset = 0 - with self.assertRaises(QuicConnectionError) as cm: + with pytest.raises(QuicConnectionError) as cm: # We don't start at zero as we want to force buffering, not cause # a TLS error. for offset in range(10000, 1000000, 10000): @@ -3062,31 +2902,29 @@ def test_excessive_crypto_buffering(self): ), ) highest_good_offset = offset - self.assertEqual( - cm.exception.error_code, QuicErrorCode.CRYPTO_BUFFER_EXCEEDED - ) - self.assertEqual(cm.exception.frame_type, QuicFrameType.CRYPTO) - self.assertEqual(highest_good_offset, (MAX_PENDING_CRYPTO // 10000) * 10000) + assert cm.value.error_code == QuicErrorCode.CRYPTO_BUFFER_EXCEEDED + assert cm.value.frame_type == QuicFrameType.CRYPTO + assert highest_good_offset == (MAX_PENDING_CRYPTO // 10000) * 10000 -class QuicNetworkPathTest(TestCase): +class TestQuicNetworkPath: def test_can_send(self): path = QuicNetworkPath(("1.2.3.4", 1234)) - self.assertFalse(path.is_validated) + assert not path.is_validated # initially, cannot send any data - self.assertTrue(path.can_send(0)) - self.assertFalse(path.can_send(1)) + assert path.can_send(0) + assert not path.can_send(1) # receive some data path.bytes_received += 1 - self.assertTrue(path.can_send(0)) - self.assertTrue(path.can_send(1)) - self.assertTrue(path.can_send(2)) - self.assertTrue(path.can_send(3)) - self.assertFalse(path.can_send(4)) + assert path.can_send(0) + assert path.can_send(1) + assert path.can_send(2) + assert path.can_send(3) + assert not path.can_send(4) # send some data path.bytes_sent += 3 - self.assertTrue(path.can_send(0)) - self.assertFalse(path.can_send(1)) + assert path.can_send(0) + assert not path.can_send(1) diff --git a/tests/test_crypto_v1.py b/tests/test_crypto_v1.py index c130a4916..acdb47c8f 100644 --- a/tests/test_crypto_v1.py +++ b/tests/test_crypto_v1.py @@ -1,7 +1,7 @@ from __future__ import annotations +import pytest import binascii -from unittest import TestCase, skipIf from qh3.buffer import Buffer from qh3.quic.crypto import ( @@ -116,7 +116,7 @@ ) -class CryptoTest(TestCase): +class TestCrypto: """ Test vectors from: @@ -159,9 +159,9 @@ def test_derive_key_iv_hp(self): secret=secret, version=PROTOCOL_VERSION, ) - self.assertEqual(key, binascii.unhexlify("1f369613dd76d5467730efcbe3b1a22d")) - self.assertEqual(iv, binascii.unhexlify("fa044b2f42a3fd3b46fb255c")) - self.assertEqual(hp, binascii.unhexlify("9f50449e04a0e810283a1e9933adedd2")) + assert key == binascii.unhexlify("1f369613dd76d5467730efcbe3b1a22d") + assert iv == binascii.unhexlify("fa044b2f42a3fd3b46fb255c") + assert hp == binascii.unhexlify("9f50449e04a0e810283a1e9933adedd2") # server secret = binascii.unhexlify( @@ -172,11 +172,11 @@ def test_derive_key_iv_hp(self): secret=secret, version=PROTOCOL_VERSION, ) - self.assertEqual(key, binascii.unhexlify("cf3a5331653c364c88f0f379b6067e37")) - self.assertEqual(iv, binascii.unhexlify("0ac1493ca1905853b0bba03e")) - self.assertEqual(hp, binascii.unhexlify("c206b8d9b9f0f37644430b490eeaa314")) + assert key == binascii.unhexlify("cf3a5331653c364c88f0f379b6067e37") + assert iv == binascii.unhexlify("0ac1493ca1905853b0bba03e") + assert hp == binascii.unhexlify("c206b8d9b9f0f37644430b490eeaa314") - @skipIf("chacha20" in SKIP_TESTS, "Skipping chacha20 tests") + @pytest.mark.skipif("chacha20" in SKIP_TESTS, reason="Skipping chacha20 tests") def test_derive_key_iv_hp_chacha20(self): # https://datatracker.ietf.org/doc/html/rfc9001#appendix-A.5 @@ -189,21 +189,17 @@ def test_derive_key_iv_hp_chacha20(self): secret=secret, version=PROTOCOL_VERSION, ) - self.assertEqual( - key, + assert key == \ binascii.unhexlify( - "c6d98ff3441c3fe1b2182094f69caa2ed4b716b65488960a7a984979fb23e1c8" - ), - ) - self.assertEqual(iv, binascii.unhexlify("e0459b3474bdd0e44a41c144")) - self.assertEqual( - hp, + "c6d98ff3441c3fe1b2182094f69caa2ed4b716b65488960a7a984979fb23e1c8" \ + ) + assert iv == binascii.unhexlify("e0459b3474bdd0e44a41c144") + assert hp == \ binascii.unhexlify( - "25a282b9e82f06f21f488917a4fc8f1b73573685608597d0efcb076b0ab7a7a4" - ), - ) + "25a282b9e82f06f21f488917a4fc8f1b73573685608597d0efcb076b0ab7a7a4" \ + ) - @skipIf("chacha20" in SKIP_TESTS, "Skipping chacha20 tests") + @pytest.mark.skipif("chacha20" in SKIP_TESTS, reason="Skipping chacha20 tests") def test_decrypt_chacha20(self): pair = CryptoPair() pair.recv.setup( @@ -217,9 +213,9 @@ def test_decrypt_chacha20(self): plain_header, plain_payload, packet_number = pair.decrypt_packet( CHACHA20_CLIENT_ENCRYPTED_PACKET, 1, CHACHA20_CLIENT_PACKET_NUMBER ) - self.assertEqual(plain_header, CHACHA20_CLIENT_PLAIN_HEADER) - self.assertEqual(plain_payload, CHACHA20_CLIENT_PLAIN_PAYLOAD) - self.assertEqual(packet_number, CHACHA20_CLIENT_PACKET_NUMBER) + assert plain_header == CHACHA20_CLIENT_PLAIN_HEADER + assert plain_payload == CHACHA20_CLIENT_PLAIN_PAYLOAD + assert packet_number == CHACHA20_CLIENT_PACKET_NUMBER def test_decrypt_long_client(self): pair = self.create_crypto(is_client=False) @@ -227,9 +223,9 @@ def test_decrypt_long_client(self): plain_header, plain_payload, packet_number = pair.decrypt_packet( LONG_CLIENT_ENCRYPTED_PACKET, 18, 0 ) - self.assertEqual(plain_header, LONG_CLIENT_PLAIN_HEADER) - self.assertEqual(plain_payload, LONG_CLIENT_PLAIN_PAYLOAD) - self.assertEqual(packet_number, LONG_CLIENT_PACKET_NUMBER) + assert plain_header == LONG_CLIENT_PLAIN_HEADER + assert plain_payload == LONG_CLIENT_PLAIN_PAYLOAD + assert packet_number == LONG_CLIENT_PACKET_NUMBER def test_decrypt_long_server(self): pair = self.create_crypto(is_client=True) @@ -237,13 +233,13 @@ def test_decrypt_long_server(self): plain_header, plain_payload, packet_number = pair.decrypt_packet( LONG_SERVER_ENCRYPTED_PACKET, 18, 0 ) - self.assertEqual(plain_header, LONG_SERVER_PLAIN_HEADER) - self.assertEqual(plain_payload, LONG_SERVER_PLAIN_PAYLOAD) - self.assertEqual(packet_number, LONG_SERVER_PACKET_NUMBER) + assert plain_header == LONG_SERVER_PLAIN_HEADER + assert plain_payload == LONG_SERVER_PLAIN_PAYLOAD + assert packet_number == LONG_SERVER_PACKET_NUMBER def test_decrypt_no_key(self): pair = CryptoPair() - with self.assertRaises(CryptoError): + with pytest.raises(CryptoError): pair.decrypt_packet(LONG_SERVER_ENCRYPTED_PACKET, 18, 0) def test_decrypt_short_server(self): @@ -259,11 +255,11 @@ def test_decrypt_short_server(self): plain_header, plain_payload, packet_number = pair.decrypt_packet( SHORT_SERVER_ENCRYPTED_PACKET, 9, 0 ) - self.assertEqual(plain_header, SHORT_SERVER_PLAIN_HEADER) - self.assertEqual(plain_payload, SHORT_SERVER_PLAIN_PAYLOAD) - self.assertEqual(packet_number, SHORT_SERVER_PACKET_NUMBER) + assert plain_header == SHORT_SERVER_PLAIN_HEADER + assert plain_payload == SHORT_SERVER_PLAIN_PAYLOAD + assert packet_number == SHORT_SERVER_PACKET_NUMBER - @skipIf("chacha20" in SKIP_TESTS, "Skipping chacha20 tests") + @pytest.mark.skipif("chacha20" in SKIP_TESTS, reason="Skipping chacha20 tests") def test_encrypt_chacha20(self): pair = CryptoPair() pair.send.setup( @@ -279,7 +275,7 @@ def test_encrypt_chacha20(self): CHACHA20_CLIENT_PLAIN_PAYLOAD, CHACHA20_CLIENT_PACKET_NUMBER, ) - self.assertEqual(packet, CHACHA20_CLIENT_ENCRYPTED_PACKET) + assert packet == CHACHA20_CLIENT_ENCRYPTED_PACKET def test_encrypt_long_client(self): pair = self.create_crypto(is_client=True) @@ -289,7 +285,7 @@ def test_encrypt_long_client(self): LONG_CLIENT_PLAIN_PAYLOAD, LONG_CLIENT_PACKET_NUMBER, ) - self.assertEqual(packet, LONG_CLIENT_ENCRYPTED_PACKET) + assert packet == LONG_CLIENT_ENCRYPTED_PACKET def test_encrypt_long_server(self): pair = self.create_crypto(is_client=False) @@ -299,7 +295,7 @@ def test_encrypt_long_server(self): LONG_SERVER_PLAIN_PAYLOAD, LONG_SERVER_PACKET_NUMBER, ) - self.assertEqual(packet, LONG_SERVER_ENCRYPTED_PACKET) + assert packet == LONG_SERVER_ENCRYPTED_PACKET def test_encrypt_short_server(self): pair = CryptoPair() @@ -316,7 +312,7 @@ def test_encrypt_short_server(self): SHORT_SERVER_PLAIN_PAYLOAD, SHORT_SERVER_PACKET_NUMBER, ) - self.assertEqual(packet, SHORT_SERVER_ENCRYPTED_PACKET) + assert packet == SHORT_SERVER_ENCRYPTED_PACKET def test_key_update(self): pair1 = self.create_crypto(is_client=True) @@ -339,15 +335,15 @@ def send(sender, receiver, packet_number=0): recov_header, recov_payload, recov_packet_number = receiver.decrypt_packet( encrypted, len(plain_header) - 2, 0 ) - self.assertEqual(recov_header, plain_header) - self.assertEqual(recov_payload, plain_payload) - self.assertEqual(recov_packet_number, packet_number) + assert recov_header == plain_header + assert recov_payload == plain_payload + assert recov_packet_number == packet_number # roundtrip send(pair1, pair2, 0) send(pair2, pair1, 0) - self.assertEqual(pair1.key_phase, 0) - self.assertEqual(pair2.key_phase, 0) + assert pair1.key_phase == 0 + assert pair2.key_phase == 0 # pair 1 key update pair1.update_key() @@ -355,8 +351,8 @@ def send(sender, receiver, packet_number=0): # roundtrip send(pair1, pair2, 1) send(pair2, pair1, 1) - self.assertEqual(pair1.key_phase, 1) - self.assertEqual(pair2.key_phase, 1) + assert pair1.key_phase == 1 + assert pair2.key_phase == 1 # pair 2 key update pair2.update_key() @@ -364,8 +360,8 @@ def send(sender, receiver, packet_number=0): # roundtrip send(pair2, pair1, 2) send(pair1, pair2, 2) - self.assertEqual(pair1.key_phase, 0) - self.assertEqual(pair2.key_phase, 0) + assert pair1.key_phase == 0 + assert pair2.key_phase == 0 # pair 1 key - update, but not next to send pair1.update_key() @@ -373,35 +369,35 @@ def send(sender, receiver, packet_number=0): # roundtrip send(pair2, pair1, 3) send(pair1, pair2, 3) - self.assertEqual(pair1.key_phase, 1) - self.assertEqual(pair2.key_phase, 1) + assert pair1.key_phase == 1 + assert pair2.key_phase == 1 def test_aead_init_args_validation(self): # invalid cipher - with self.assertRaises(CryptoError) as cm: + with pytest.raises(CryptoError) as cm: self.create_aead(cipher_name=b"tango9000") - self.assertEqual(str(cm.exception), "Invalid cipher name: tango9000") + assert str(cm.value) == "Invalid cipher name: tango9000" # invalid key length - with self.assertRaises(CryptoError) as cm: + with pytest.raises(CryptoError) as cm: self.create_aead(key=bytes(33)) - self.assertEqual(str(cm.exception), "Invalid key length") + assert str(cm.value) == "Invalid key length" # invalid iv length - with self.assertRaises(CryptoError) as cm: + with pytest.raises(CryptoError) as cm: self.create_aead(iv=bytes(11)) - self.assertEqual(str(cm.exception), "Invalid iv length") - with self.assertRaises(CryptoError) as cm: + assert str(cm.value) == "Invalid iv length" + with pytest.raises(CryptoError) as cm: self.create_aead(iv=bytes(13)) - self.assertEqual(str(cm.exception), "Invalid iv length") + assert str(cm.value) == "Invalid iv length" def test_hp_init_args_validation(self): # invalid cipher - with self.assertRaises(CryptoError) as cm: + with pytest.raises(CryptoError) as cm: self.create_hp(cipher_name=b"tango9000") - self.assertEqual(str(cm.exception), "Invalid cipher name: tango9000") + assert str(cm.value) == "Invalid cipher name: tango9000" # invalid key length - with self.assertRaises(CryptoError) as cm: + with pytest.raises(CryptoError) as cm: self.create_hp(key=bytes(33)) - self.assertEqual(str(cm.exception), "Invalid key length") + assert str(cm.value) == "Invalid key length" diff --git a/tests/test_crypto_v2.py b/tests/test_crypto_v2.py index 7a9b16240..6600f4f0c 100644 --- a/tests/test_crypto_v2.py +++ b/tests/test_crypto_v2.py @@ -1,7 +1,7 @@ from __future__ import annotations +import pytest import binascii -from unittest import TestCase, skipIf from qh3.buffer import Buffer from qh3.quic.crypto import ( @@ -112,7 +112,7 @@ ) -class CryptoTest(TestCase): +class TestCrypto: """ Test vectors from: @@ -140,9 +140,9 @@ def test_derive_key_iv_hp(self): secret=secret, version=PROTOCOL_VERSION, ) - self.assertEqual(key, binascii.unhexlify("8b1a0bc121284290a29e0971b5cd045d")) - self.assertEqual(iv, binascii.unhexlify("91f73e2351d8fa91660e909f")) - self.assertEqual(hp, binascii.unhexlify("45b95e15235d6f45a6b19cbcb0294ba9")) + assert key == binascii.unhexlify("8b1a0bc121284290a29e0971b5cd045d") + assert iv == binascii.unhexlify("91f73e2351d8fa91660e909f") + assert hp == binascii.unhexlify("45b95e15235d6f45a6b19cbcb0294ba9") # server secret = binascii.unhexlify( @@ -153,11 +153,11 @@ def test_derive_key_iv_hp(self): secret=secret, version=PROTOCOL_VERSION, ) - self.assertEqual(key, binascii.unhexlify("82db637861d55e1d011f19ea71d5d2a7")) - self.assertEqual(iv, binascii.unhexlify("dd13c276499c0249d3310652")) - self.assertEqual(hp, binascii.unhexlify("edf6d05c83121201b436e16877593c3a")) + assert key == binascii.unhexlify("82db637861d55e1d011f19ea71d5d2a7") + assert iv == binascii.unhexlify("dd13c276499c0249d3310652") + assert hp == binascii.unhexlify("edf6d05c83121201b436e16877593c3a") - @skipIf("chacha20" in SKIP_TESTS, "Skipping chacha20 tests") + @pytest.mark.skipif("chacha20" in SKIP_TESTS, reason="Skipping chacha20 tests") def test_derive_key_iv_hp_chacha20(self): # https://datatracker.ietf.org/doc/html/rfc9369#appendix-A.5 @@ -170,21 +170,17 @@ def test_derive_key_iv_hp_chacha20(self): secret=secret, version=PROTOCOL_VERSION, ) - self.assertEqual( - key, + assert key == \ binascii.unhexlify( - "3bfcddd72bcf02541d7fa0dd1f5f9eeea817e09a6963a0e6c7df0f9a1bab90f2" - ), - ) - self.assertEqual(iv, binascii.unhexlify("a6b5bc6ab7dafce30ffff5dd")) - self.assertEqual( - hp, + "3bfcddd72bcf02541d7fa0dd1f5f9eeea817e09a6963a0e6c7df0f9a1bab90f2" \ + ) + assert iv == binascii.unhexlify("a6b5bc6ab7dafce30ffff5dd") + assert hp == \ binascii.unhexlify( - "d659760d2ba434a226fd37b35c69e2da8211d10c4f12538787d65645d5d1b8e2" - ), - ) + "d659760d2ba434a226fd37b35c69e2da8211d10c4f12538787d65645d5d1b8e2" \ + ) - @skipIf("chacha20" in SKIP_TESTS, "Skipping chacha20 tests") + @pytest.mark.skipif("chacha20" in SKIP_TESTS, reason="Skipping chacha20 tests") def test_decrypt_chacha20(self): pair = CryptoPair() pair.recv.setup( @@ -198,9 +194,9 @@ def test_decrypt_chacha20(self): plain_header, plain_payload, packet_number = pair.decrypt_packet( CHACHA20_CLIENT_ENCRYPTED_PACKET, 1, CHACHA20_CLIENT_PACKET_NUMBER ) - self.assertEqual(plain_header, CHACHA20_CLIENT_PLAIN_HEADER) - self.assertEqual(plain_payload, CHACHA20_CLIENT_PLAIN_PAYLOAD) - self.assertEqual(packet_number, CHACHA20_CLIENT_PACKET_NUMBER) + assert plain_header == CHACHA20_CLIENT_PLAIN_HEADER + assert plain_payload == CHACHA20_CLIENT_PLAIN_PAYLOAD + assert packet_number == CHACHA20_CLIENT_PACKET_NUMBER def test_decrypt_long_client(self): pair = self.create_crypto(is_client=False) @@ -208,9 +204,9 @@ def test_decrypt_long_client(self): plain_header, plain_payload, packet_number = pair.decrypt_packet( LONG_CLIENT_ENCRYPTED_PACKET, 18, 0 ) - self.assertEqual(plain_header, LONG_CLIENT_PLAIN_HEADER) - self.assertEqual(plain_payload, LONG_CLIENT_PLAIN_PAYLOAD) - self.assertEqual(packet_number, LONG_CLIENT_PACKET_NUMBER) + assert plain_header == LONG_CLIENT_PLAIN_HEADER + assert plain_payload == LONG_CLIENT_PLAIN_PAYLOAD + assert packet_number == LONG_CLIENT_PACKET_NUMBER def test_decrypt_long_server(self): pair = self.create_crypto(is_client=True) @@ -218,13 +214,13 @@ def test_decrypt_long_server(self): plain_header, plain_payload, packet_number = pair.decrypt_packet( LONG_SERVER_ENCRYPTED_PACKET, 18, 0 ) - self.assertEqual(plain_header, LONG_SERVER_PLAIN_HEADER) - self.assertEqual(plain_payload, LONG_SERVER_PLAIN_PAYLOAD) - self.assertEqual(packet_number, LONG_SERVER_PACKET_NUMBER) + assert plain_header == LONG_SERVER_PLAIN_HEADER + assert plain_payload == LONG_SERVER_PLAIN_PAYLOAD + assert packet_number == LONG_SERVER_PACKET_NUMBER def test_decrypt_no_key(self): pair = CryptoPair() - with self.assertRaises(CryptoError): + with pytest.raises(CryptoError): pair.decrypt_packet(LONG_SERVER_ENCRYPTED_PACKET, 18, 0) def test_decrypt_short_server(self): @@ -240,11 +236,11 @@ def test_decrypt_short_server(self): plain_header, plain_payload, packet_number = pair.decrypt_packet( SHORT_SERVER_ENCRYPTED_PACKET, 9, 0 ) - self.assertEqual(plain_header, SHORT_SERVER_PLAIN_HEADER) - self.assertEqual(plain_payload, SHORT_SERVER_PLAIN_PAYLOAD) - self.assertEqual(packet_number, SHORT_SERVER_PACKET_NUMBER) + assert plain_header == SHORT_SERVER_PLAIN_HEADER + assert plain_payload == SHORT_SERVER_PLAIN_PAYLOAD + assert packet_number == SHORT_SERVER_PACKET_NUMBER - @skipIf("chacha20" in SKIP_TESTS, "Skipping chacha20 tests") + @pytest.mark.skipif("chacha20" in SKIP_TESTS, reason="Skipping chacha20 tests") def test_encrypt_chacha20(self): pair = CryptoPair() pair.send.setup( @@ -260,7 +256,7 @@ def test_encrypt_chacha20(self): CHACHA20_CLIENT_PLAIN_PAYLOAD, CHACHA20_CLIENT_PACKET_NUMBER, ) - self.assertEqual(packet, CHACHA20_CLIENT_ENCRYPTED_PACKET) + assert packet == CHACHA20_CLIENT_ENCRYPTED_PACKET def test_encrypt_long_client(self): pair = self.create_crypto(is_client=True) @@ -270,7 +266,7 @@ def test_encrypt_long_client(self): LONG_CLIENT_PLAIN_PAYLOAD, LONG_CLIENT_PACKET_NUMBER, ) - self.assertEqual(packet, LONG_CLIENT_ENCRYPTED_PACKET) + assert packet == LONG_CLIENT_ENCRYPTED_PACKET def test_encrypt_long_server(self): pair = self.create_crypto(is_client=False) @@ -280,7 +276,7 @@ def test_encrypt_long_server(self): LONG_SERVER_PLAIN_PAYLOAD, LONG_SERVER_PACKET_NUMBER, ) - self.assertEqual(packet, LONG_SERVER_ENCRYPTED_PACKET) + assert packet == LONG_SERVER_ENCRYPTED_PACKET def test_encrypt_short_server(self): pair = CryptoPair() @@ -297,7 +293,7 @@ def test_encrypt_short_server(self): SHORT_SERVER_PLAIN_PAYLOAD, SHORT_SERVER_PACKET_NUMBER, ) - self.assertEqual(packet, SHORT_SERVER_ENCRYPTED_PACKET) + assert packet == SHORT_SERVER_ENCRYPTED_PACKET def test_key_update(self): pair1 = self.create_crypto(is_client=True) @@ -320,15 +316,15 @@ def send(sender, receiver, packet_number=0): recov_header, recov_payload, recov_packet_number = receiver.decrypt_packet( encrypted, len(plain_header) - 2, 0 ) - self.assertEqual(recov_header, plain_header) - self.assertEqual(recov_payload, plain_payload) - self.assertEqual(recov_packet_number, packet_number) + assert recov_header == plain_header + assert recov_payload == plain_payload + assert recov_packet_number == packet_number # roundtrip send(pair1, pair2, 0) send(pair2, pair1, 0) - self.assertEqual(pair1.key_phase, 0) - self.assertEqual(pair2.key_phase, 0) + assert pair1.key_phase == 0 + assert pair2.key_phase == 0 # pair 1 key update pair1.update_key() @@ -336,8 +332,8 @@ def send(sender, receiver, packet_number=0): # roundtrip send(pair1, pair2, 1) send(pair2, pair1, 1) - self.assertEqual(pair1.key_phase, 1) - self.assertEqual(pair2.key_phase, 1) + assert pair1.key_phase == 1 + assert pair2.key_phase == 1 # pair 2 key update pair2.update_key() @@ -345,8 +341,8 @@ def send(sender, receiver, packet_number=0): # roundtrip send(pair2, pair1, 2) send(pair1, pair2, 2) - self.assertEqual(pair1.key_phase, 0) - self.assertEqual(pair2.key_phase, 0) + assert pair1.key_phase == 0 + assert pair2.key_phase == 0 # pair 1 key - update, but not next to send pair1.update_key() @@ -354,5 +350,5 @@ def send(sender, receiver, packet_number=0): # roundtrip send(pair2, pair1, 3) send(pair1, pair2, 3) - self.assertEqual(pair1.key_phase, 1) - self.assertEqual(pair2.key_phase, 1) + assert pair1.key_phase == 1 + assert pair2.key_phase == 1 diff --git a/tests/test_h3.py b/tests/test_h3.py index 898a215ab..ad7fad9c9 100644 --- a/tests/test_h3.py +++ b/tests/test_h3.py @@ -1,9 +1,9 @@ from __future__ import annotations +import pytest import binascii import contextlib import copy -from unittest import TestCase from qh3.buffer import Buffer, encode_uint_var from qh3.h3.connection import ( @@ -129,7 +129,7 @@ def send_stream_data(self, stream_id, data, end_stream=False): ) -class H3ConnectionTest(TestCase): +class TestH3Connection: maxDiff = None def _make_request(self, h3_client, h3_server): @@ -152,8 +152,7 @@ def _make_request(self, h3_client, h3_server): # receive request events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -167,8 +166,7 @@ def _make_request(self, h3_client, h3_server): stream_ended=False, ), DataReceived(data=b"", stream_id=stream_id, stream_ended=True), - ], - ) + ] # send response h3_server.send_headers( @@ -187,8 +185,7 @@ def _make_request(self, h3_client, h3_server): # receive response events = h3_transfer(quic_server, h3_client) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -204,8 +201,7 @@ def _make_request(self, h3_client, h3_server): stream_id=stream_id, stream_ended=True, ), - ], - ) + ] def test_handle_control_frame_headers(self): """ @@ -215,8 +211,8 @@ def test_handle_control_frame_headers(self): configuration=QuicConfiguration(is_client=False) ) h3_server = H3Connection(quic_server) - self.assertIsNotNone(h3_server.sent_settings) - self.assertIsNone(h3_server.received_settings) + assert h3_server.sent_settings is not None + assert h3_server.received_settings is None # receive SETTINGS h3_server.handle_event( @@ -227,9 +223,9 @@ def test_handle_control_frame_headers(self): end_stream=False, ) ) - self.assertIsNone(quic_server.closed) - self.assertIsNotNone(h3_server.sent_settings) - self.assertEqual(h3_server.received_settings, DUMMY_SETTINGS) + assert quic_server.closed is None + assert h3_server.sent_settings is not None + assert h3_server.received_settings == DUMMY_SETTINGS # receive unexpected HEADERS h3_server.handle_event( @@ -239,10 +235,8 @@ def test_handle_control_frame_headers(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, - (ErrorCode.H3_FRAME_UNEXPECTED, "Invalid frame type on control stream"), - ) + assert quic_server.closed == \ + (ErrorCode.H3_FRAME_UNEXPECTED, "Invalid frame type on control stream") def test_handle_control_frame_max_push_id_from_client_before_settings(self): """ @@ -262,10 +256,8 @@ def test_handle_control_frame_max_push_id_from_client_before_settings(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, - (ErrorCode.H3_MISSING_SETTINGS, ""), - ) + assert quic_server.closed == \ + (ErrorCode.H3_MISSING_SETTINGS, "") def test_handle_control_frame_max_push_id_from_server(self): """ @@ -285,7 +277,7 @@ def test_handle_control_frame_max_push_id_from_server(self): end_stream=False, ) ) - self.assertIsNone(quic_client.closed) + assert quic_client.closed is None # receive unexpected MAX_PUSH_ID h3_client.handle_event( @@ -295,10 +287,8 @@ def test_handle_control_frame_max_push_id_from_server(self): end_stream=False, ) ) - self.assertEqual( - quic_client.closed, - (ErrorCode.H3_FRAME_UNEXPECTED, "Servers must not send MAX_PUSH_ID"), - ) + assert quic_client.closed == \ + (ErrorCode.H3_FRAME_UNEXPECTED, "Servers must not send MAX_PUSH_ID") def test_handle_control_settings_twice(self): """ @@ -318,7 +308,7 @@ def test_handle_control_settings_twice(self): end_stream=False, ) ) - self.assertIsNone(quic_server.closed) + assert quic_server.closed is None # receive unexpected SETTINGS h3_server.handle_event( @@ -328,10 +318,8 @@ def test_handle_control_settings_twice(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, - (ErrorCode.H3_FRAME_UNEXPECTED, "SETTINGS have already been received"), - ) + assert quic_server.closed == \ + (ErrorCode.H3_FRAME_UNEXPECTED, "SETTINGS have already been received") def test_handle_control_stream_close(self): """ @@ -351,7 +339,7 @@ def test_handle_control_stream_close(self): end_stream=False, ) ) - self.assertIsNone(quic_client.closed) + assert quic_client.closed is None # receive unexpected FIN h3_client.handle_event( @@ -361,13 +349,11 @@ def test_handle_control_stream_close(self): end_stream=True, ) ) - self.assertEqual( - quic_client.closed, + assert quic_client.closed == \ ( ErrorCode.H3_CLOSED_CRITICAL_STREAM, "Closing control stream is not allowed", - ), - ) + ) def test_handle_control_stream_duplicate(self): """ @@ -391,13 +377,11 @@ def test_handle_control_stream_duplicate(self): stream_id=6, data=encode_uint_var(StreamType.CONTROL), end_stream=False ) ) - self.assertEqual( - quic_server.closed, + assert quic_server.closed == \ ( ErrorCode.H3_STREAM_CREATION_ERROR, "Only one control stream is allowed", - ), - ) + ) def test_handle_push_frame_wrong_frame_type(self): """ @@ -417,10 +401,8 @@ def test_handle_push_frame_wrong_frame_type(self): end_stream=False, ) ) - self.assertEqual( - quic_client.closed, - (ErrorCode.H3_FRAME_UNEXPECTED, "Invalid frame type on push stream"), - ) + assert quic_client.closed == \ + (ErrorCode.H3_FRAME_UNEXPECTED, "Invalid frame type on push stream") def test_handle_qpack_decoder_duplicate(self): """ @@ -448,13 +430,11 @@ def test_handle_qpack_decoder_duplicate(self): end_stream=False, ) ) - self.assertEqual( - quic_client.closed, + assert quic_client.closed == \ ( ErrorCode.H3_STREAM_CREATION_ERROR, "Only one QPACK decoder stream is allowed", - ), - ) + ) def test_handle_qpack_decoder_stream_error(self): """ @@ -472,7 +452,7 @@ def test_handle_qpack_decoder_stream_error(self): end_stream=False, ) ) - self.assertEqual(quic_client.closed, (ErrorCode.QPACK_DECODER_STREAM_ERROR, "")) + assert quic_client.closed == (ErrorCode.QPACK_DECODER_STREAM_ERROR, "") def test_handle_qpack_encoder_duplicate(self): """ @@ -500,13 +480,11 @@ def test_handle_qpack_encoder_duplicate(self): end_stream=False, ) ) - self.assertEqual( - quic_client.closed, + assert quic_client.closed == \ ( ErrorCode.H3_STREAM_CREATION_ERROR, "Only one QPACK encoder stream is allowed", - ), - ) + ) def test_handle_qpack_encoder_stream_error(self): """ @@ -524,7 +502,7 @@ def test_handle_qpack_encoder_stream_error(self): end_stream=False, ) ) - self.assertEqual(quic_client.closed, (ErrorCode.QPACK_ENCODER_STREAM_ERROR, "")) + assert quic_client.closed == (ErrorCode.QPACK_ENCODER_STREAM_ERROR, "") def test_handle_request_frame_bad_headers(self): """ @@ -540,7 +518,7 @@ def test_handle_request_frame_bad_headers(self): stream_id=0, data=encode_frame(FrameType.HEADERS, b""), end_stream=False ) ) - self.assertEqual(quic_server.closed, (ErrorCode.QPACK_DECOMPRESSION_FAILED, "")) + assert quic_server.closed == (ErrorCode.QPACK_DECOMPRESSION_FAILED, "") def test_handle_request_frame_data_before_headers(self): """ @@ -556,13 +534,11 @@ def test_handle_request_frame_data_before_headers(self): stream_id=0, data=encode_frame(FrameType.DATA, b""), end_stream=False ) ) - self.assertEqual( - quic_server.closed, + assert quic_server.closed == \ ( ErrorCode.H3_FRAME_UNEXPECTED, "DATA frame is not allowed in this state", - ), - ) + ) def test_handle_request_frame_headers_after_trailers(self): """ @@ -596,13 +572,11 @@ def test_handle_request_frame_headers_after_trailers(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, + assert quic_server.closed == \ ( ErrorCode.H3_FRAME_UNEXPECTED, "HEADERS frame is not allowed in this state", - ), - ) + ) def test_handle_request_frame_push_promise_from_client(self): """ @@ -620,10 +594,8 @@ def test_handle_request_frame_push_promise_from_client(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, - (ErrorCode.H3_FRAME_UNEXPECTED, "Clients must not send PUSH_PROMISE"), - ) + assert quic_server.closed == \ + (ErrorCode.H3_FRAME_UNEXPECTED, "Clients must not send PUSH_PROMISE") def test_handle_request_frame_wrong_frame_type(self): quic_server = FakeQuicConnection( @@ -638,10 +610,8 @@ def test_handle_request_frame_wrong_frame_type(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, - (ErrorCode.H3_FRAME_UNEXPECTED, "Invalid frame type on request stream"), - ) + assert quic_server.closed == \ + (ErrorCode.H3_FRAME_UNEXPECTED, "Invalid frame type on request stream") def test_request(self): with h3_client_and_server() as (quic_client, quic_server): @@ -678,8 +648,7 @@ def test_request_headers_only(self): # receive request events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -691,9 +660,8 @@ def test_request_headers_only(self): ], stream_id=stream_id, stream_ended=True, - ) - ], - ) + ) \ + ] # send response h3_server.send_headers( @@ -708,8 +676,7 @@ def test_request_headers_only(self): # receive response events = h3_transfer(quic_server, h3_client) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -719,9 +686,8 @@ def test_request_headers_only(self): ], stream_id=stream_id, stream_ended=True, - ) - ], - ) + ) \ + ] def test_request_fragmented_frame(self): with h3_fake_client_and_server() as (quic_client, quic_server): @@ -744,8 +710,7 @@ def test_request_fragmented_frame(self): # receive request events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -764,8 +729,7 @@ def test_request_fragmented_frame(self): DataReceived(data=b"l", stream_id=0, stream_ended=False), DataReceived(data=b"o", stream_id=0, stream_ended=False), DataReceived(data=b"", stream_id=0, stream_ended=True), - ], - ) + ] # send push promise push_stream_id = h3_server.send_push_promise( @@ -777,7 +741,7 @@ def test_request_fragmented_frame(self): (b":path", b"/app.txt"), ], ) - self.assertEqual(push_stream_id, 15) + assert push_stream_id == 15 # send response h3_server.send_headers( @@ -800,8 +764,7 @@ def test_request_fragmented_frame(self): # receive push promise / response events = h3_transfer(quic_server, h3_client) - self.assertEqual( - events, + assert events == \ [ PushPromiseReceived( headers=[ @@ -836,20 +799,19 @@ def test_request_fragmented_frame(self): push_id=0, ), DataReceived( - data=b"t", stream_id=15, stream_ended=False, push_id=0 + data=b"t", stream_id=15, stream_ended=False, push_id=0 \ ), DataReceived( - data=b"e", stream_id=15, stream_ended=False, push_id=0 + data=b"e", stream_id=15, stream_ended=False, push_id=0 \ ), DataReceived( - data=b"x", stream_id=15, stream_ended=False, push_id=0 + data=b"x", stream_id=15, stream_ended=False, push_id=0 \ ), DataReceived( - data=b"t", stream_id=15, stream_ended=False, push_id=0 + data=b"t", stream_id=15, stream_ended=False, push_id=0 \ ), DataReceived(data=b"", stream_id=15, stream_ended=True, push_id=0), - ], - ) + ] def test_request_with_server_push(self): with h3_client_and_server() as (quic_client, quic_server): @@ -871,8 +833,7 @@ def test_request_with_server_push(self): # receive request events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -883,9 +844,8 @@ def test_request_with_server_push(self): ], stream_id=stream_id, stream_ended=True, - ) - ], - ) + ) \ + ] # send push promises push_stream_id_css = h3_server.send_push_promise( @@ -897,7 +857,7 @@ def test_request_with_server_push(self): (b":path", b"/app.css"), ], ) - self.assertEqual(push_stream_id_css, 15) + assert push_stream_id_css == 15 push_stream_id_js = h3_server.send_push_promise( stream_id=stream_id, @@ -908,7 +868,7 @@ def test_request_with_server_push(self): (b":path", b"/app.js"), ], ) - self.assertEqual(push_stream_id_js, 19) + assert push_stream_id_js == 19 # send response h3_server.send_headers( @@ -952,8 +912,7 @@ def test_request_with_server_push(self): # receive push promises, response and push responses events = h3_transfer(quic_server, h3_client) - self.assertEqual( - events, + assert events == \ [ PushPromiseReceived( headers=[ @@ -1015,8 +974,7 @@ def test_request_with_server_push(self): stream_id=push_stream_id_js, stream_ended=True, ), - ], - ) + ] def test_request_with_server_push_max_push_id(self): with h3_client_and_server() as (quic_client, quic_server): @@ -1038,8 +996,7 @@ def test_request_with_server_push_max_push_id(self): # receive request events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -1050,9 +1007,8 @@ def test_request_with_server_push_max_push_id(self): ], stream_id=stream_id, stream_ended=True, - ) - ], - ) + ) \ + ] # send push promises for i in range(0, 8): @@ -1067,7 +1023,7 @@ def test_request_with_server_push_max_push_id(self): ) # send one too many - with self.assertRaises(NoAvailablePushIDError): + with pytest.raises(NoAvailablePushIDError): h3_server.send_push_promise( stream_id=stream_id, headers=[ @@ -1100,7 +1056,7 @@ def test_send_data_after_trailers(self): h3_client.send_headers( stream_id=stream_id, headers=[(b"x-some-trailer", b"foo")], end_stream=False ) - with self.assertRaises(FrameUnexpected): + with pytest.raises(FrameUnexpected): h3_client.send_data(stream_id=stream_id, data=b"hello", end_stream=False) def test_send_data_before_headers(self): @@ -1113,7 +1069,7 @@ def test_send_data_before_headers(self): h3_client = H3Connection(quic_client) stream_id = quic_client.get_next_available_stream_id() - with self.assertRaises(FrameUnexpected): + with pytest.raises(FrameUnexpected): h3_client.send_data(stream_id=stream_id, data=b"hello", end_stream=False) def test_send_headers_after_trailers(self): @@ -1138,7 +1094,7 @@ def test_send_headers_after_trailers(self): h3_client.send_headers( stream_id=stream_id, headers=[(b"x-some-trailer", b"foo")], end_stream=False ) - with self.assertRaises(FrameUnexpected): + with pytest.raises(FrameUnexpected): h3_client.send_headers( stream_id=stream_id, headers=[(b"x-other-trailer", b"foo")], @@ -1184,16 +1140,15 @@ def test_blocked_stream(self): end_stream=True, ) ) - self.assertEqual( - h3_client.handle_event( + assert h3_client.handle_event( StreamDataReceived( stream_id=7, data=binascii.unhexlify( - "3fe101c696d07abe941094cb6d0a08017d403971966e32ca98b46f" + "3fe101c696d07abe941094cb6d0a08017d403971966e32ca98b46f" \ ), end_stream=False, - ) - ), + ) \ + ) == \ [ HeadersReceived( headers=[ @@ -1205,15 +1160,14 @@ def test_blocked_stream(self): ), DataReceived( data=( - b"you reached mvfst.net, reach the /echo endpoint for an " - b"echo response query / endpoints for a variable " - b"size response with random bytes" + b"you reached mvfst.net, reach the /echo endpoint for an " \ + b"echo response query / endpoints for a variable " \ + b"size response with random bytes" \ ), stream_id=0, stream_ended=True, ), - ], - ) + ] def test_blocked_stream_trailer(self): quic_client = FakeQuicConnection( @@ -1237,16 +1191,15 @@ def test_blocked_stream_trailer(self): StreamDataReceived(stream_id=11, data=b"\x03", end_stream=False) ) - self.assertEqual( - h3_client.handle_event( + assert h3_client.handle_event( StreamDataReceived( stream_id=0, data=binascii.unhexlify( - "011b0000d95696d07abe941094cb6d0a08017d403971966e32ca98b46f" + "011b0000d95696d07abe941094cb6d0a08017d403971966e32ca98b46f" \ ), end_stream=False, - ) - ), + ) \ + ) == \ [ HeadersReceived( headers=[ @@ -1255,63 +1208,56 @@ def test_blocked_stream_trailer(self): ], stream_id=0, stream_ended=False, - ) - ], - ) + ) \ + ] - self.assertEqual( - h3_client.handle_event( + assert h3_client.handle_event( StreamDataReceived( stream_id=0, data=binascii.unhexlify( - "00408d796f752072656163686564206d766673742e6e65742c20726561636820" - "746865202f6563686f20656e64706f696e7420666f7220616e206563686f2072" - "6573706f6e7365207175657279202f3c6e756d6265723e20656e64706f696e74" - "7320666f722061207661726961626c652073697a6520726573706f6e73652077" - "6974682072616e646f6d206279746573" + "00408d796f752072656163686564206d766673742e6e65742c20726561636820" \ + "746865202f6563686f20656e64706f696e7420666f7220616e206563686f2072" \ + "6573706f6e7365207175657279202f3c6e756d6265723e20656e64706f696e74" \ + "7320666f722061207661726961626c652073697a6520726573706f6e73652077" \ + "6974682072616e646f6d206279746573" \ ), end_stream=False, - ) - ), + ) \ + ) == \ [ DataReceived( data=( - b"you reached mvfst.net, reach the /echo endpoint for an " - b"echo response query / endpoints for a variable " - b"size response with random bytes" + b"you reached mvfst.net, reach the /echo endpoint for an " \ + b"echo response query / endpoints for a variable " \ + b"size response with random bytes" \ ), stream_id=0, stream_ended=False, - ) - ], - ) + ) \ + ] - self.assertEqual( - h3_client.handle_event( + assert h3_client.handle_event( StreamDataReceived( - stream_id=0, data=binascii.unhexlify("0103028010"), end_stream=True - ) - ), - [], - ) + stream_id=0, data=binascii.unhexlify("0103028010"), end_stream=True \ + ) \ + ) == \ + [] - self.assertEqual( - h3_client.handle_event( + assert h3_client.handle_event( StreamDataReceived( stream_id=7, data=binascii.unhexlify("6af2b20f49564d833505b38294e7"), end_stream=False, - ) - ), + ) \ + ) == \ [ HeadersReceived( headers=[(b"x-some-trailer", b"foo")], stream_id=0, stream_ended=True, push_id=None, - ) - ], - ) + ) \ + ] def test_uni_stream_grease(self): with h3_client_and_server() as (quic_client, quic_server): @@ -1320,7 +1266,7 @@ def test_uni_stream_grease(self): quic_client.send_stream_data( 14, b"\xff\xff\xff\xff\xff\xff\xff\xfeGREASE is the word" ) - self.assertEqual(h3_transfer(quic_client, h3_server), []) + assert h3_transfer(quic_client, h3_server) == [] def test_request_with_trailers(self): with h3_client_and_server() as (quic_client, quic_server): @@ -1347,8 +1293,7 @@ def test_request_with_trailers(self): # receive request events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -1365,8 +1310,7 @@ def test_request_with_trailers(self): stream_id=stream_id, stream_ended=True, ), - ], - ) + ] # send response h3_server.send_headers( @@ -1390,8 +1334,7 @@ def test_request_with_trailers(self): # receive response events = h3_transfer(quic_server, h3_client) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -1411,8 +1354,7 @@ def test_request_with_trailers(self): stream_id=stream_id, stream_ended=True, ), - ], - ) + ] def test_request_with_informational(self): @@ -1435,8 +1377,7 @@ def test_request_with_informational(self): # receive request events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -1448,8 +1389,7 @@ def test_request_with_informational(self): stream_id=stream_id, stream_ended=True, ), - ], - ) + ] # send response h3_server.send_headers( @@ -1476,15 +1416,14 @@ def test_request_with_informational(self): # receive response events = h3_transfer(quic_server, h3_client) - self.assertEqual( - events, + assert events == \ [ InformationalHeadersReceived( headers=[ (b":status", b"103"), (b"link", b"; rel=preload; as=style"), ], - stream_id=stream_id + stream_id=stream_id \ ), HeadersReceived( headers=[ @@ -1499,8 +1438,7 @@ def test_request_with_informational(self): stream_id=stream_id, stream_ended=True, ), - ], - ) + ] def test_uni_stream_type(self): with h3_client_and_server() as (quic_client, quic_server): @@ -1508,32 +1446,32 @@ def test_uni_stream_type(self): # unknown stream type 9 stream_id = quic_client.get_next_available_stream_id(is_unidirectional=True) - self.assertEqual(stream_id, 2) + assert stream_id == 2 quic_client.send_stream_data(stream_id, b"\x09") - self.assertEqual(h3_transfer(quic_client, h3_server), []) - self.assertEqual(list(h3_server._stream.keys()), [2]) - self.assertEqual(h3_server._stream[2].buffer, b"") - self.assertEqual(h3_server._stream[2].stream_type, 9) + assert h3_transfer(quic_client, h3_server) == [] + assert list(h3_server._stream.keys()) == [2] + assert h3_server._stream[2].buffer == b"" + assert h3_server._stream[2].stream_type == 9 # unknown stream type 64, one byte at a time stream_id = quic_client.get_next_available_stream_id(is_unidirectional=True) - self.assertEqual(stream_id, 6) + assert stream_id == 6 quic_client.send_stream_data(stream_id, b"\x40") - self.assertEqual(h3_transfer(quic_client, h3_server), []) - self.assertEqual(list(h3_server._stream.keys()), [2, 6]) - self.assertEqual(h3_server._stream[2].buffer, b"") - self.assertEqual(h3_server._stream[2].stream_type, 9) - self.assertEqual(h3_server._stream[6].buffer, b"\x40") - self.assertEqual(h3_server._stream[6].stream_type, None) + assert h3_transfer(quic_client, h3_server) == [] + assert list(h3_server._stream.keys()) == [2, 6] + assert h3_server._stream[2].buffer == b"" + assert h3_server._stream[2].stream_type == 9 + assert h3_server._stream[6].buffer == b"\x40" + assert h3_server._stream[6].stream_type == None quic_client.send_stream_data(stream_id, b"\x40") - self.assertEqual(h3_transfer(quic_client, h3_server), []) - self.assertEqual(list(h3_server._stream.keys()), [2, 6]) - self.assertEqual(h3_server._stream[2].buffer, b"") - self.assertEqual(h3_server._stream[2].stream_type, 9) - self.assertEqual(h3_server._stream[6].buffer, b"") - self.assertEqual(h3_server._stream[6].stream_type, 64) + assert h3_transfer(quic_client, h3_server) == [] + assert list(h3_server._stream.keys()) == [2, 6] + assert h3_server._stream[2].buffer == b"" + assert h3_server._stream[2].stream_type == 9 + assert h3_server._stream[6].buffer == b"" + assert h3_server._stream[6].stream_type == 64 def test_validate_settings_h3_datagram_invalid_value(self): quic_server = FakeQuicConnection( @@ -1552,13 +1490,11 @@ def test_validate_settings_h3_datagram_invalid_value(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, + assert quic_server.closed == \ ( ErrorCode.H3_SETTINGS_ERROR, "H3_DATAGRAM setting must be 0 or 1", - ), - ) + ) def test_validate_settings_h3_datagram_without_transport_parameter(self): quic_server = FakeQuicConnection( @@ -1577,13 +1513,11 @@ def test_validate_settings_h3_datagram_without_transport_parameter(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, + assert quic_server.closed == \ ( ErrorCode.H3_SETTINGS_ERROR, "H3_DATAGRAM requires max_datagram_frame_size transport parameter", - ), - ) + ) def test_validate_settings_enable_connect_protocol_invalid_value(self): quic_server = FakeQuicConnection( @@ -1602,13 +1536,11 @@ def test_validate_settings_enable_connect_protocol_invalid_value(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, + assert quic_server.closed == \ ( ErrorCode.H3_SETTINGS_ERROR, "ENABLE_CONNECT_PROTOCOL setting must be 0 or 1", - ), - ) + ) def test_validate_settings_enable_webtransport_invalid_value(self): quic_server = FakeQuicConnection( @@ -1627,13 +1559,11 @@ def test_validate_settings_enable_webtransport_invalid_value(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, + assert quic_server.closed == \ ( ErrorCode.H3_SETTINGS_ERROR, "ENABLE_WEBTRANSPORT setting must be 0 or 1", - ), - ) + ) def test_validate_settings_enable_webtransport_without_h3_datagram(self): quic_server = FakeQuicConnection( @@ -1652,16 +1582,14 @@ def test_validate_settings_enable_webtransport_without_h3_datagram(self): end_stream=False, ) ) - self.assertEqual( - quic_server.closed, + assert quic_server.closed == \ ( ErrorCode.H3_SETTINGS_ERROR, "ENABLE_WEBTRANSPORT requires H3_DATAGRAM", - ), - ) + ) -class H3ParserTest(TestCase): +class TestH3Parser: def test_parse_settings_duplicate_identifier(self): buf = Buffer(capacity=1024) buf.push_uint_var(1) @@ -1669,22 +1597,18 @@ def test_parse_settings_duplicate_identifier(self): buf.push_uint_var(1) buf.push_uint_var(456) - with self.assertRaises(SettingsError) as cm: + with pytest.raises(SettingsError) as cm: parse_settings(buf.data) - self.assertEqual( - cm.exception.reason_phrase, "Setting identifier 0x1 is included twice" - ) + assert cm.value.reason_phrase == "Setting identifier 0x1 is included twice" def test_parse_settings_reserved_identifier(self): buf = Buffer(capacity=1024) buf.push_uint_var(0) buf.push_uint_var(123) - with self.assertRaises(SettingsError) as cm: + with pytest.raises(SettingsError) as cm: parse_settings(buf.data) - self.assertEqual( - cm.exception.reason_phrase, "Setting identifier 0x0 is reserved" - ) + assert cm.value.reason_phrase == "Setting identifier 0x0 is reserved" def test_validate_push_promise_headers(self): # OK @@ -1707,26 +1631,22 @@ def test_validate_push_promise_headers(self): ) # invalid pseudo-header - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_push_promise_headers([(b":status", b"foo")]) - self.assertEqual( - cm.exception.reason_phrase, "Pseudo-header b':status' is not valid" - ) + assert cm.value.reason_phrase == "Pseudo-header b':status' is not valid" # duplicate pseudo-header - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_push_promise_headers( [ (b":method", b"GET"), (b":method", b"POST"), ] ) - self.assertEqual( - cm.exception.reason_phrase, "Pseudo-header b':method' is included twice" - ) + assert cm.value.reason_phrase == "Pseudo-header b':method' is included twice" # pseudo-header after regular headers - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_push_promise_headers( [ (b":method", b"GET"), @@ -1736,13 +1656,11 @@ def test_validate_push_promise_headers(self): (b":authority", b"foo"), ] ) - self.assertEqual( - cm.exception.reason_phrase, - "Pseudo-header b':authority' is not allowed after regular headers", - ) + assert cm.value.reason_phrase == \ + "Pseudo-header b':authority' is not allowed after regular headers" # missing pseudo-headers - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_push_promise_headers( [ (b":method", b"GET"), @@ -1750,10 +1668,8 @@ def test_validate_push_promise_headers(self): (b":path", b"/"), ] ) - self.assertEqual( - cm.exception.reason_phrase, - "Pseudo-headers [b':authority'] are missing", - ) + assert cm.value.reason_phrase == \ + "Pseudo-headers [b':authority'] are missing" def test_validate_request_headers(self): # OK @@ -1776,33 +1692,27 @@ def test_validate_request_headers(self): ) # uppercase header - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_request_headers([(b"X-Foo", b"foo")]) - self.assertEqual( - cm.exception.reason_phrase, "Header b'X-Foo' contains uppercase letters" - ) + assert cm.value.reason_phrase == "Header b'X-Foo' contains uppercase letters" # invalid pseudo-header - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_request_headers([(b":status", b"foo")]) - self.assertEqual( - cm.exception.reason_phrase, "Pseudo-header b':status' is not valid" - ) + assert cm.value.reason_phrase == "Pseudo-header b':status' is not valid" # duplicate pseudo-header - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_request_headers( [ (b":method", b"GET"), (b":method", b"POST"), ] ) - self.assertEqual( - cm.exception.reason_phrase, "Pseudo-header b':method' is included twice" - ) + assert cm.value.reason_phrase == "Pseudo-header b':method' is included twice" # pseudo-header after regular headers - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_request_headers( [ (b":method", b"GET"), @@ -1812,22 +1722,18 @@ def test_validate_request_headers(self): (b":authority", b"foo"), ] ) - self.assertEqual( - cm.exception.reason_phrase, - "Pseudo-header b':authority' is not allowed after regular headers", - ) + assert cm.value.reason_phrase == \ + "Pseudo-header b':authority' is not allowed after regular headers" # missing pseudo-headers - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_request_headers([(b":method", b"GET")]) - self.assertEqual( - cm.exception.reason_phrase, - "Pseudo-headers [b':authority'] are missing", - ) + assert cm.value.reason_phrase == \ + "Pseudo-headers [b':authority'] are missing" # empty :authority pseudo-header for http/https for scheme in [b"http", b"https"]: - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_request_headers( [ (b":method", b"GET"), @@ -1836,14 +1742,12 @@ def test_validate_request_headers(self): (b":path", b"/"), ] ) - self.assertEqual( - cm.exception.reason_phrase, - "Pseudo-header b':authority' cannot be empty", - ) + assert cm.value.reason_phrase == \ + "Pseudo-header b':authority' cannot be empty" # empty :path pseudo-header for http/https for scheme in [b"http", b"https"]: - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_request_headers( [ (b":method", b"GET"), @@ -1852,9 +1756,7 @@ def test_validate_request_headers(self): (b":path", b""), ] ) - self.assertEqual( - cm.exception.reason_phrase, "Pseudo-header b':path' cannot be empty" - ) + assert cm.value.reason_phrase == "Pseudo-header b':path' cannot be empty" def test_validate_response_headers(self): # OK @@ -1867,44 +1769,36 @@ def test_validate_response_headers(self): ) # invalid pseudo-header - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_response_headers([(b":method", b"GET")]) - self.assertEqual( - cm.exception.reason_phrase, "Pseudo-header b':method' is not valid" - ) + assert cm.value.reason_phrase == "Pseudo-header b':method' is not valid" # duplicate pseudo-header - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_response_headers( [ (b":status", b"200"), (b":status", b"501"), ] ) - self.assertEqual( - cm.exception.reason_phrase, "Pseudo-header b':status' is included twice" - ) + assert cm.value.reason_phrase == "Pseudo-header b':status' is included twice" def test_validate_trailers(self): # OK validate_trailers([(b"x-foo", b"bar")]) # invalid pseudo-header - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_trailers([(b":status", b"foo")]) - self.assertEqual( - cm.exception.reason_phrase, "Pseudo-header b':status' is not valid" - ) + assert cm.value.reason_phrase == "Pseudo-header b':status' is not valid" # pseudo-header after regular headers - with self.assertRaises(MessageError) as cm: + with pytest.raises(MessageError) as cm: validate_trailers( [ (b"x-foo", b"bar"), (b":authority", b"foo"), ] ) - self.assertEqual( - cm.exception.reason_phrase, - "Pseudo-header b':authority' is not allowed after regular headers", - ) + assert cm.value.reason_phrase == \ + "Pseudo-header b':authority' is not allowed after regular headers" diff --git a/tests/test_logger.py b/tests/test_logger.py index 075f2530b..80ab90b65 100644 --- a/tests/test_logger.py +++ b/tests/test_logger.py @@ -1,9 +1,9 @@ from __future__ import annotations +import pytest import json import os import tempfile -from unittest import TestCase from qh3.quic.logger import QuicFileLogger, QuicLogger @@ -22,29 +22,25 @@ } -class QuicLoggerTest(TestCase): +class TestQuicLogger: def test_empty(self): logger = QuicLogger() - self.assertEqual( - logger.to_dict(), - {"qlog_format": "JSON", "qlog_version": "0.3", "traces": []}, - ) + assert logger.to_dict() == \ + {"qlog_format": "JSON", "qlog_version": "0.3", "traces": []} def test_single_trace(self): logger = QuicLogger() trace = logger.start_trace(is_client=True, odcid=bytes(8)) logger.end_trace(trace) - self.assertEqual(logger.to_dict(), SINGLE_TRACE) + assert logger.to_dict() == SINGLE_TRACE -class QuicFileLoggerTest(TestCase): +class TestQuicFileLogger: def test_invalid_path(self): - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: QuicFileLogger("this_path_should_not_exist") - self.assertEqual( - str(cm.exception), - "QUIC log output directory 'this_path_should_not_exist' does not exist", - ) + assert str(cm.value) == \ + "QUIC log output directory 'this_path_should_not_exist' does not exist" def test_single_trace(self): with tempfile.TemporaryDirectory() as dirpath: @@ -53,8 +49,8 @@ def test_single_trace(self): logger.end_trace(trace) filepath = os.path.join(dirpath, "0000000000000000.qlog") - self.assertTrue(os.path.exists(filepath)) + assert os.path.exists(filepath) with open(filepath) as fp: data = json.load(fp) - self.assertEqual(data, SINGLE_TRACE) + assert data == SINGLE_TRACE diff --git a/tests/test_packet.py b/tests/test_packet.py index 605d8d6d6..414810003 100644 --- a/tests/test_packet.py +++ b/tests/test_packet.py @@ -1,7 +1,7 @@ from __future__ import annotations +import pytest import binascii -from unittest import TestCase from qh3.buffer import Buffer, BufferReadError from qh3.quic import packet @@ -27,87 +27,87 @@ from .test_crypto_v2 import LONG_SERVER_ENCRYPTED_PACKET as SERVER_INITIAL_V2 -class PacketTest(TestCase): +class TestPacket: def test_decode_packet_number(self): # expected = 0 for i in range(0, 256): - self.assertEqual(decode_packet_number(i, 8, expected=0), i) + assert decode_packet_number(i, 8, expected=0) == i # expected = 128 - self.assertEqual(decode_packet_number(0, 8, expected=128), 256) + assert decode_packet_number(0, 8, expected=128) == 256 for i in range(1, 256): - self.assertEqual(decode_packet_number(i, 8, expected=128), i) + assert decode_packet_number(i, 8, expected=128) == i # expected = 129 - self.assertEqual(decode_packet_number(0, 8, expected=129), 256) - self.assertEqual(decode_packet_number(1, 8, expected=129), 257) + assert decode_packet_number(0, 8, expected=129) == 256 + assert decode_packet_number(1, 8, expected=129) == 257 for i in range(2, 256): - self.assertEqual(decode_packet_number(i, 8, expected=129), i) + assert decode_packet_number(i, 8, expected=129) == i # expected = 256 for i in range(0, 128): - self.assertEqual(decode_packet_number(i, 8, expected=256), 256 + i) + assert decode_packet_number(i, 8, expected=256) == 256 + i for i in range(129, 256): - self.assertEqual(decode_packet_number(i, 8, expected=256), i) + assert decode_packet_number(i, 8, expected=256) == i def test_pull_empty(self): buf = Buffer(data=b"") - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): pull_quic_header(buf, host_cid_length=8) def test_pull_initial_client_v1(self): buf = Buffer(data=CLIENT_INITIAL_V1) header = pull_quic_header(buf, host_cid_length=8) - self.assertEqual(header.version, QuicProtocolVersion.VERSION_1) - self.assertEqual(header.packet_type, QuicPacketType.INITIAL) - self.assertEqual(header.packet_length, 1200) - self.assertEqual(header.destination_cid, binascii.unhexlify("8394c8f03e515708")) - self.assertEqual(header.source_cid, b"") - self.assertEqual(header.token, b"") - self.assertEqual(header.integrity_tag, b"") - self.assertEqual(buf.tell(), 18) + assert header.version == QuicProtocolVersion.VERSION_1 + assert header.packet_type == QuicPacketType.INITIAL + assert header.packet_length == 1200 + assert header.destination_cid == binascii.unhexlify("8394c8f03e515708") + assert header.source_cid == b"" + assert header.token == b"" + assert header.integrity_tag == b"" + assert buf.tell() == 18 def test_pull_initial_client_v1_truncated(self): buf = Buffer(data=CLIENT_INITIAL_V1[0:100]) - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: pull_quic_header(buf, host_cid_length=8) - self.assertEqual(str(cm.exception), "Packet payload is truncated") + assert str(cm.value) == "Packet payload is truncated" def test_pull_initial_client_v2(self): buf = Buffer(data=CLIENT_INITIAL_V2) header = pull_quic_header(buf, host_cid_length=8) - self.assertEqual(header.version, QuicProtocolVersion.VERSION_2) - self.assertEqual(header.packet_type, QuicPacketType.INITIAL) - self.assertEqual(header.packet_length, 1200) - self.assertEqual(header.destination_cid, binascii.unhexlify("8394c8f03e515708")) - self.assertEqual(header.source_cid, b"") - self.assertEqual(header.token, b"") - self.assertEqual(header.integrity_tag, b"") - self.assertEqual(buf.tell(), 18) + assert header.version == QuicProtocolVersion.VERSION_2 + assert header.packet_type == QuicPacketType.INITIAL + assert header.packet_length == 1200 + assert header.destination_cid == binascii.unhexlify("8394c8f03e515708") + assert header.source_cid == b"" + assert header.token == b"" + assert header.integrity_tag == b"" + assert buf.tell() == 18 def test_pull_initial_server_v1(self): buf = Buffer(data=SERVER_INITIAL_V1) header = pull_quic_header(buf, host_cid_length=8) - self.assertEqual(header.version, QuicProtocolVersion.VERSION_1) - self.assertEqual(header.packet_type, QuicPacketType.INITIAL) - self.assertEqual(header.packet_length, 135) - self.assertEqual(header.destination_cid, b"") - self.assertEqual(header.source_cid, binascii.unhexlify("f067a5502a4262b5")) - self.assertEqual(header.token, b"") - self.assertEqual(header.integrity_tag, b"") - self.assertEqual(buf.tell(), 18) + assert header.version == QuicProtocolVersion.VERSION_1 + assert header.packet_type == QuicPacketType.INITIAL + assert header.packet_length == 135 + assert header.destination_cid == b"" + assert header.source_cid == binascii.unhexlify("f067a5502a4262b5") + assert header.token == b"" + assert header.integrity_tag == b"" + assert buf.tell() == 18 def test_pull_initial_server_v2(self): buf = Buffer(data=SERVER_INITIAL_V2) header = pull_quic_header(buf, host_cid_length=8) - self.assertEqual(header.version, QuicProtocolVersion.VERSION_2) - self.assertEqual(header.packet_type, QuicPacketType.INITIAL) - self.assertEqual(header.packet_length, 135) - self.assertEqual(header.destination_cid, b"") - self.assertEqual(header.source_cid, binascii.unhexlify("f067a5502a4262b5")) - self.assertEqual(header.token, b"") - self.assertEqual(header.integrity_tag, b"") - self.assertEqual(buf.tell(), 18) + assert header.version == QuicProtocolVersion.VERSION_2 + assert header.packet_type == QuicPacketType.INITIAL + assert header.packet_length == 135 + assert header.destination_cid == b"" + assert header.source_cid == binascii.unhexlify("f067a5502a4262b5") + assert header.token == b"" + assert header.integrity_tag == b"" + assert buf.tell() == 18 def test_pull_retry_v1(self): # https://datatracker.ietf.org/doc/html/rfc9001#appendix-A.4 @@ -118,24 +118,20 @@ def test_pull_retry_v1(self): ) buf = Buffer(data=data) header = pull_quic_header(buf) - self.assertEqual(header.version, QuicProtocolVersion.VERSION_1) - self.assertEqual(header.packet_type, QuicPacketType.RETRY) - self.assertEqual(header.packet_length, 36) - self.assertEqual(header.destination_cid, b"") - self.assertEqual(header.source_cid, binascii.unhexlify("f067a5502a4262b5")) - self.assertEqual(header.token, b"token") - self.assertEqual( - header.integrity_tag, binascii.unhexlify("04a265ba2eff4d829058fb3f0f2496ba") - ) - self.assertEqual(buf.tell(), 36) + assert header.version == QuicProtocolVersion.VERSION_1 + assert header.packet_type == QuicPacketType.RETRY + assert header.packet_length == 36 + assert header.destination_cid == b"" + assert header.source_cid == binascii.unhexlify("f067a5502a4262b5") + assert header.token == b"token" + assert header.integrity_tag == binascii.unhexlify("04a265ba2eff4d829058fb3f0f2496ba") + assert buf.tell() == 36 # check integrity - self.assertEqual( - get_retry_integrity_tag( - buf.data_slice(0, 20), original_destination_cid, version=header.version - ), - header.integrity_tag, - ) + assert get_retry_integrity_tag( + buf.data_slice(0, 20), original_destination_cid, version=header.version \ + ) == \ + header.integrity_tag # serialize encoded = encode_quic_retry( @@ -149,7 +145,7 @@ def test_pull_retry_v1(self): ) with open("bob.bin", "wb") as fp: fp.write(encoded) - self.assertEqual(encoded, data) + assert encoded == data def test_pull_retry_v2(self): # https://datatracker.ietf.org/doc/html/rfc9369#appendix-A.4 @@ -160,24 +156,20 @@ def test_pull_retry_v2(self): ) buf = Buffer(data=data) header = pull_quic_header(buf) - self.assertEqual(header.version, QuicProtocolVersion.VERSION_2) - self.assertEqual(header.packet_type, QuicPacketType.RETRY) - self.assertEqual(header.packet_length, 36) - self.assertEqual(header.destination_cid, b"") - self.assertEqual(header.source_cid, binascii.unhexlify("f067a5502a4262b5")) - self.assertEqual(header.token, b"token") - self.assertEqual( - header.integrity_tag, binascii.unhexlify("c8646ce8bfe33952d955543665dcc7b6") - ) - self.assertEqual(buf.tell(), 36) + assert header.version == QuicProtocolVersion.VERSION_2 + assert header.packet_type == QuicPacketType.RETRY + assert header.packet_length == 36 + assert header.destination_cid == b"" + assert header.source_cid == binascii.unhexlify("f067a5502a4262b5") + assert header.token == b"token" + assert header.integrity_tag == binascii.unhexlify("c8646ce8bfe33952d955543665dcc7b6") + assert buf.tell() == 36 # check integrity - self.assertEqual( - get_retry_integrity_tag( - buf.data_slice(0, 20), original_destination_cid, version=header.version - ), - header.integrity_tag, - ) + assert get_retry_integrity_tag( + buf.data_slice(0, 20), original_destination_cid, version=header.version \ + ) == \ + header.integrity_tag # serialize encoded = encode_quic_retry( @@ -191,7 +183,7 @@ def test_pull_retry_v2(self): ) with open("bob.bin", "wb") as fp: fp.write(encoded) - self.assertEqual(encoded, data) + assert encoded == data def test_pull_version_negotiation(self): data = binascii.unhexlify( @@ -199,17 +191,15 @@ def test_pull_version_negotiation(self): ) buf = Buffer(data=data) header = pull_quic_header(buf, host_cid_length=8) - self.assertEqual(header.version, QuicProtocolVersion.NEGOTIATION) - self.assertEqual(header.packet_type, QuicPacketType.VERSION_NEGOTIATION) - self.assertEqual(header.packet_length, 31) - self.assertEqual(header.destination_cid, binascii.unhexlify("9aac5a49ba87a849")) - self.assertEqual(header.source_cid, binascii.unhexlify("f92f4336fa951ba1")) - self.assertEqual(header.token, b"") - self.assertEqual(header.integrity_tag, b"") - self.assertEqual( - header.supported_versions, [0x45474716, QuicProtocolVersion.VERSION_1] - ) - self.assertEqual(buf.tell(), 31) + assert header.version == QuicProtocolVersion.NEGOTIATION + assert header.packet_type == QuicPacketType.VERSION_NEGOTIATION + assert header.packet_length == 31 + assert header.destination_cid == binascii.unhexlify("9aac5a49ba87a849") + assert header.source_cid == binascii.unhexlify("f92f4336fa951ba1") + assert header.token == b"" + assert header.integrity_tag == b"" + assert header.supported_versions == [0x45474716, QuicProtocolVersion.VERSION_1] + assert buf.tell() == 31 encoded = encode_quic_version_negotiation( destination_cid=header.destination_cid, @@ -218,7 +208,7 @@ def test_pull_version_negotiation(self): ) # The first byte may differ as it is random. - self.assertEqual(encoded[1:], data[1:]) + assert encoded[1:] == data[1:] def test_pull_long_header_dcid_too_long(self): buf = Buffer( @@ -227,9 +217,9 @@ def test_pull_long_header_dcid_too_long(self): "01c514f99ec4bbf1f7a30f9b0c94fef717f1c1d07fec24c99a864da7ede" ) ) - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: pull_quic_header(buf, host_cid_length=8) - self.assertEqual(str(cm.exception), "Destination CID is too long (21 bytes)") + assert str(cm.value) == "Destination CID is too long (21 bytes)" def test_pull_long_header_scid_too_long(self): buf = Buffer( @@ -238,19 +228,19 @@ def test_pull_long_header_scid_too_long(self): "01cfcee99ec4bbf1f7a30f9b0c9417b8c263cdd8cc972a4439d68a46320" ) ) - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: pull_quic_header(buf, host_cid_length=8) - self.assertEqual(str(cm.exception), "Source CID is too long (21 bytes)") + assert str(cm.value) == "Source CID is too long (21 bytes)" def test_pull_long_header_no_fixed_bit(self): buf = Buffer(data=b"\x80\xff\x00\x00\x11\x00\x00") - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: pull_quic_header(buf, host_cid_length=8) - self.assertEqual(str(cm.exception), "Packet fixed bit is zero") + assert str(cm.value) == "Packet fixed bit is zero" def test_pull_long_header_too_short(self): buf = Buffer(data=b"\xc0\x00") - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): pull_quic_header(buf, host_cid_length=8) def test_pull_short_header(self): @@ -258,23 +248,23 @@ def test_pull_short_header(self): data=binascii.unhexlify("5df45aa7b59c0e1ad6e668f5304cd4fd1fb3799327") ) header = pull_quic_header(buf, host_cid_length=8) - self.assertEqual(header.version, None) - self.assertEqual(header.packet_type, QuicPacketType.ONE_RTT) - self.assertEqual(header.packet_length, 21) - self.assertEqual(header.destination_cid, binascii.unhexlify("f45aa7b59c0e1ad6")) - self.assertEqual(header.source_cid, b"") - self.assertEqual(header.token, b"") - self.assertEqual(header.integrity_tag, b"") - self.assertEqual(buf.tell(), 9) + assert header.version == None + assert header.packet_type == QuicPacketType.ONE_RTT + assert header.packet_length == 21 + assert header.destination_cid == binascii.unhexlify("f45aa7b59c0e1ad6") + assert header.source_cid == b"" + assert header.token == b"" + assert header.integrity_tag == b"" + assert buf.tell() == 9 def test_pull_short_header_no_fixed_bit(self): buf = Buffer(data=b"\x00") - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: pull_quic_header(buf, host_cid_length=8) - self.assertEqual(str(cm.exception), "Packet fixed bit is zero") + assert str(cm.value) == "Packet fixed bit is zero" -class ParamsTest(TestCase): +class TestParams: maxDiff = None def test_params(self): @@ -286,8 +276,7 @@ def test_params(self): # parse buf = Buffer(data=data) params = pull_quic_transport_parameters(buf) - self.assertEqual( - params, + assert params == \ QuicTransportParameters( max_idle_timeout=10000, stateless_reset_token=b"\xcc/\xd6\xe7\xd9zS\xab[\xe8[(\xd7\\\x80\x08", @@ -300,13 +289,12 @@ def test_params(self): initial_max_streams_uni=None, ack_delay_exponent=3, max_ack_delay=25, - ), - ) + ) # serialize buf = Buffer(capacity=len(data)) push_quic_transport_parameters(buf, params) - self.assertEqual(len(buf.data), len(data)) + assert len(buf.data) == len(data) def test_params_disable_active_migration(self): data = binascii.unhexlify("0c00") @@ -314,12 +302,12 @@ def test_params_disable_active_migration(self): # parse buf = Buffer(data=data) params = pull_quic_transport_parameters(buf) - self.assertEqual(params, QuicTransportParameters(disable_active_migration=True)) + assert params == QuicTransportParameters(disable_active_migration=True) # serialize buf = Buffer(capacity=len(data)) push_quic_transport_parameters(buf, params) - self.assertEqual(buf.data, data) + assert buf.data == data def test_params_preferred_address(self): data = binascii.unhexlify( @@ -330,8 +318,7 @@ def test_params_preferred_address(self): # parse buf = Buffer(data=data) params = pull_quic_transport_parameters(buf) - self.assertEqual( - params, + assert params == \ QuicTransportParameters( preferred_address=QuicPreferredAddress( ipv4_address=("139.162.123.134", 4435), @@ -339,13 +326,12 @@ def test_params_preferred_address(self): connection_id=b"b\xc4Q\x8dc\x01?\x0c(~\xd3W>\xfa\x90\x95`7", stateless_reset_token=b"F\xb2\xe0-EH\x0b\xa6d>\\n}H\xec\xb4", ), - ), - ) + ) # serialize buf = Buffer(capacity=1000) push_quic_transport_parameters(buf, params) - self.assertEqual(buf.data, data) + assert buf.data == data def test_params_unknown(self): data = binascii.unhexlify("8000ff000100") @@ -353,7 +339,7 @@ def test_params_unknown(self): # parse buf = Buffer(data=data) params = pull_quic_transport_parameters(buf) - self.assertEqual(params, QuicTransportParameters()) + assert params == QuicTransportParameters() def test_preferred_address_ipv4_only(self): data = binascii.unhexlify( @@ -364,20 +350,18 @@ def test_preferred_address_ipv4_only(self): # parse buf = Buffer(data=data) preferred_address = pull_quic_preferred_address(buf) - self.assertEqual( - preferred_address, + assert preferred_address == \ QuicPreferredAddress( ipv4_address=("139.162.123.134", 4435), ipv6_address=None, connection_id=b"b\xc4Q\x8dc\x01?\x0c(~\xd3W>\xfa\x90\x95`7", stateless_reset_token=b"F\xb2\xe0-EH\x0b\xa6d>\\n}H\xec\xb4", - ), - ) + ) # serialize buf = Buffer(capacity=len(data)) push_quic_preferred_address(buf, preferred_address) - self.assertEqual(buf.data, data) + assert buf.data == data def test_preferred_address_ipv6_only(self): data = binascii.unhexlify( @@ -388,36 +372,34 @@ def test_preferred_address_ipv6_only(self): # parse buf = Buffer(data=data) preferred_address = pull_quic_preferred_address(buf) - self.assertEqual( - preferred_address, + assert preferred_address == \ QuicPreferredAddress( ipv4_address=None, ipv6_address=("2400:8902::f03c:91ff:fe69:a454", 4435), connection_id=b"b\xc4Q\x8dc\x01?\x0c(~\xd3W>\xfa\x90\x95`7", stateless_reset_token=b"F\xb2\xe0-EH\x0b\xa6d>\\n}H\xec\xb4", - ), - ) + ) # serialize buf = Buffer(capacity=len(data)) push_quic_preferred_address(buf, preferred_address) - self.assertEqual(buf.data, data) + assert buf.data == data -class FrameTest(TestCase): +class TestFrame: def test_ack_frame(self): data = b"\x00\x02\x00\x00" # parse buf = Buffer(data=data) rangeset, delay = packet.pull_ack_frame(buf) - self.assertEqual(list(rangeset), [range(0, 1)]) - self.assertEqual(delay, 2) + assert list(rangeset) == [range(0, 1)] + assert delay == 2 # serialize buf = Buffer(capacity=len(data)) packet.push_ack_frame(buf, rangeset, delay) - self.assertEqual(buf.data, data) + assert buf.data == data def test_ack_frame_with_one_range(self): data = b"\x02\x02\x01\x00\x00\x00" @@ -425,13 +407,13 @@ def test_ack_frame_with_one_range(self): # parse buf = Buffer(data=data) rangeset, delay = packet.pull_ack_frame(buf) - self.assertEqual(list(rangeset), [range(0, 1), range(2, 3)]) - self.assertEqual(delay, 2) + assert list(rangeset) == [range(0, 1), range(2, 3)] + assert delay == 2 # serialize buf = Buffer(capacity=len(data)) packet.push_ack_frame(buf, rangeset, delay) - self.assertEqual(buf.data, data) + assert buf.data == data def test_ack_frame_with_one_range_2(self): data = b"\x05\x02\x01\x00\x00\x03" @@ -439,13 +421,13 @@ def test_ack_frame_with_one_range_2(self): # parse buf = Buffer(data=data) rangeset, delay = packet.pull_ack_frame(buf) - self.assertEqual(list(rangeset), [range(0, 4), range(5, 6)]) - self.assertEqual(delay, 2) + assert list(rangeset) == [range(0, 4), range(5, 6)] + assert delay == 2 # serialize buf = Buffer(capacity=len(data)) packet.push_ack_frame(buf, rangeset, delay) - self.assertEqual(buf.data, data) + assert buf.data == data def test_ack_frame_with_one_range_3(self): data = b"\x05\x02\x01\x00\x01\x02" @@ -453,13 +435,13 @@ def test_ack_frame_with_one_range_3(self): # parse buf = Buffer(data=data) rangeset, delay = packet.pull_ack_frame(buf) - self.assertEqual(list(rangeset), [range(0, 3), range(5, 6)]) - self.assertEqual(delay, 2) + assert list(rangeset) == [range(0, 3), range(5, 6)] + assert delay == 2 # serialize buf = Buffer(capacity=len(data)) packet.push_ack_frame(buf, rangeset, delay) - self.assertEqual(buf.data, data) + assert buf.data == data def test_ack_frame_with_two_ranges(self): data = b"\x04\x02\x02\x00\x00\x00\x00\x00" @@ -467,10 +449,10 @@ def test_ack_frame_with_two_ranges(self): # parse buf = Buffer(data=data) rangeset, delay = packet.pull_ack_frame(buf) - self.assertEqual(list(rangeset), [range(0, 1), range(2, 3), range(4, 5)]) - self.assertEqual(delay, 2) + assert list(rangeset) == [range(0, 1), range(2, 3), range(4, 5)] + assert delay == 2 # serialize buf = Buffer(capacity=len(data)) packet.push_ack_frame(buf, rangeset, delay) - self.assertEqual(buf.data, data) + assert buf.data == data diff --git a/tests/test_packet_builder.py b/tests/test_packet_builder.py index fa62d5369..5928c69db 100644 --- a/tests/test_packet_builder.py +++ b/tests/test_packet_builder.py @@ -1,6 +1,6 @@ from __future__ import annotations -from unittest import TestCase +import pytest from qh3.quic.crypto import CryptoPair from qh3.quic.packet import QuicFrameType, QuicPacketType, QuicProtocolVersion @@ -36,22 +36,22 @@ def create_crypto(): return crypto -class QuicPacketBuilderTest(TestCase): +class TestQuicPacketBuilder: def test_long_header_empty(self): builder = create_builder() crypto = create_crypto() builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertEqual(builder.remaining_flight_space, 1236) - self.assertTrue(builder.packet_is_empty) + assert builder.remaining_flight_space == 1236 + assert builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 0) - self.assertEqual(packets, []) + assert len(datagrams) == 0 + assert packets == [] # check builder - self.assertEqual(builder.packet_number, 0) + assert builder.packet_number == 0 def test_long_header_padding(self): builder = create_builder(is_client=True) @@ -59,21 +59,20 @@ def test_long_header_padding(self): # INITIAL, fully padded builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertEqual(builder.remaining_flight_space, 1236) + assert builder.remaining_flight_space == 1236 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(100)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # INITIAL, empty builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertTrue(builder.packet_is_empty) + assert builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 1) - self.assertEqual(len(datagrams[0]), 1280) - self.assertEqual( - packets, + assert len(datagrams) == 1 + assert len(datagrams[0]) == 1280 + assert packets == \ [ QuicSentPacket( epoch=Epoch.INITIAL, @@ -83,12 +82,11 @@ def test_long_header_padding(self): packet_number=0, packet_type=QuicPacketType.INITIAL, sent_bytes=145, - ) - ], - ) + ) \ + ] # check builder - self.assertEqual(builder.packet_number, 1) + assert builder.packet_number == 1 def test_long_header_initial_client_2(self): self.maxDiff = None @@ -97,29 +95,28 @@ def test_long_header_initial_client_2(self): # INITIAL, full length builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertEqual(builder.remaining_flight_space, 1236) + assert builder.remaining_flight_space == 1236 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(builder.remaining_flight_space)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # INITIAL builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertEqual(builder.remaining_flight_space, 1236) + assert builder.remaining_flight_space == 1236 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(100)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # INITIAL, empty builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertTrue(builder.packet_is_empty) + assert builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 2) - self.assertEqual(len(datagrams[0]), 1280) - self.assertEqual(len(datagrams[1]), 1280) - self.assertEqual( - packets, + assert len(datagrams) == 2 + assert len(datagrams[0]) == 1280 + assert len(datagrams[1]) == 1280 + assert packets == \ [ QuicSentPacket( epoch=Epoch.INITIAL, @@ -139,11 +136,10 @@ def test_long_header_initial_client_2(self): packet_type=QuicPacketType.INITIAL, sent_bytes=145, ), - ], - ) + ] # check builder - self.assertEqual(builder.packet_number, 2) + assert builder.packet_number == 2 def test_long_header_initial_server(self): builder = create_builder() @@ -151,21 +147,20 @@ def test_long_header_initial_server(self): # INITIAL builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertEqual(builder.remaining_flight_space, 1236) + assert builder.remaining_flight_space == 1236 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(100)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # INITIAL, empty builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertTrue(builder.packet_is_empty) + assert builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 1) - self.assertEqual(len(datagrams[0]), 1280) - self.assertEqual( - packets, + assert len(datagrams) == 1 + assert len(datagrams[0]) == 1280 + assert packets == \ [ QuicSentPacket( epoch=Epoch.INITIAL, @@ -175,12 +170,11 @@ def test_long_header_initial_server(self): packet_number=0, packet_type=QuicPacketType.INITIAL, sent_bytes=145, - ) - ], - ) + ) \ + ] # check builder - self.assertEqual(builder.packet_number, 1) + assert builder.packet_number == 1 def test_long_header_ping_only(self): """ @@ -193,14 +187,13 @@ def test_long_header_ping_only(self): # HANDSHAKE, with only a PING frame builder.start_packet(QuicPacketType.HANDSHAKE, crypto) builder.start_frame(QuicFrameType.PING) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 1) - self.assertEqual(len(datagrams[0]), 45) - self.assertEqual( - packets, + assert len(datagrams) == 1 + assert len(datagrams[0]) == 45 + assert packets == \ [ QuicSentPacket( epoch=Epoch.HANDSHAKE, @@ -210,9 +203,8 @@ def test_long_header_ping_only(self): packet_number=0, packet_type=QuicPacketType.HANDSHAKE, sent_bytes=45, - ) - ], - ) + ) \ + ] def test_long_header_then_short_header(self): builder = create_builder() @@ -220,33 +212,32 @@ def test_long_header_then_short_header(self): # INITIAL, full length builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertEqual(builder.remaining_flight_space, 1236) + assert builder.remaining_flight_space == 1236 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(builder.remaining_flight_space)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # INITIAL, empty builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertTrue(builder.packet_is_empty) + assert builder.packet_is_empty # ONE_RTT, full length builder.start_packet(QuicPacketType.ONE_RTT, crypto) - self.assertEqual(builder.remaining_flight_space, 1253) + assert builder.remaining_flight_space == 1253 buf = builder.start_frame(QuicFrameType.STREAM_BASE) buf.push_bytes(bytes(builder.remaining_flight_space)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # ONE_RTT, empty builder.start_packet(QuicPacketType.ONE_RTT, crypto) - self.assertTrue(builder.packet_is_empty) + assert builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 2) - self.assertEqual(len(datagrams[0]), 1280) - self.assertEqual(len(datagrams[1]), 1280) - self.assertEqual( - packets, + assert len(datagrams) == 2 + assert len(datagrams[0]) == 1280 + assert len(datagrams[1]) == 1280 + assert packets == \ [ QuicSentPacket( epoch=Epoch.INITIAL, @@ -266,11 +257,10 @@ def test_long_header_then_short_header(self): packet_type=QuicPacketType.ONE_RTT, sent_bytes=1280, ), - ], - ) + ] # check builder - self.assertEqual(builder.packet_number, 2) + assert builder.packet_number == 2 def test_long_header_initial_client_zero_rtt(self): builder = create_builder(is_client=True) @@ -278,23 +268,22 @@ def test_long_header_initial_client_zero_rtt(self): # INITIAL builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertEqual(builder.remaining_flight_space, 1236) + assert builder.remaining_flight_space == 1236 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(613)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # 0-RTT builder.start_packet(QuicPacketType.ZERO_RTT, crypto) - self.assertEqual(builder.remaining_flight_space, 579) + assert builder.remaining_flight_space == 579 buf = builder.start_frame(QuicFrameType.STREAM_BASE) buf.push_bytes(bytes(100)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(datagram_sizes(datagrams), [1280]) - self.assertEqual( - packets, + assert datagram_sizes(datagrams) == [1280] + assert packets == \ [ QuicSentPacket( epoch=Epoch.INITIAL, @@ -314,8 +303,7 @@ def test_long_header_initial_client_zero_rtt(self): packet_type=QuicPacketType.ZERO_RTT, sent_bytes=144, ), - ], - ) + ] def test_long_header_then_long_header(self): builder = create_builder() @@ -323,31 +311,30 @@ def test_long_header_then_long_header(self): # INITIAL builder.start_packet(QuicPacketType.INITIAL, crypto) - self.assertEqual(builder.remaining_flight_space, 1236) + assert builder.remaining_flight_space == 1236 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(199)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # HANDSHAKE builder.start_packet(QuicPacketType.HANDSHAKE, crypto) - self.assertEqual(builder.remaining_flight_space, 993) + assert builder.remaining_flight_space == 993 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(299)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # ONE_RTT builder.start_packet(QuicPacketType.ONE_RTT, crypto) - self.assertEqual(builder.remaining_flight_space, 666) + assert builder.remaining_flight_space == 666 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(299)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 1) - self.assertEqual(len(datagrams[0]), 1280) - self.assertEqual( - packets, + assert len(datagrams) == 1 + assert len(datagrams[0]) == 1280 + assert packets == \ [ QuicSentPacket( epoch=Epoch.INITIAL, @@ -376,27 +363,26 @@ def test_long_header_then_long_header(self): packet_type=QuicPacketType.ONE_RTT, sent_bytes=693, ), - ], - ) + ] # check builder - self.assertEqual(builder.packet_number, 3) + assert builder.packet_number == 3 def test_short_header_empty(self): builder = create_builder() crypto = create_crypto() builder.start_packet(QuicPacketType.ONE_RTT, crypto) - self.assertEqual(builder.remaining_flight_space, 1253) - self.assertTrue(builder.packet_is_empty) + assert builder.remaining_flight_space == 1253 + assert builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(datagrams, []) - self.assertEqual(packets, []) + assert datagrams == [] + assert packets == [] # check builder - self.assertEqual(builder.packet_number, 0) + assert builder.packet_number == 0 def test_short_header_padding(self): builder = create_builder() @@ -404,17 +390,16 @@ def test_short_header_padding(self): # ONE_RTT, full length builder.start_packet(QuicPacketType.ONE_RTT, crypto) - self.assertEqual(builder.remaining_flight_space, 1253) + assert builder.remaining_flight_space == 1253 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(builder.remaining_flight_space)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 1) - self.assertEqual(len(datagrams[0]), 1280) - self.assertEqual( - packets, + assert len(datagrams) == 1 + assert len(datagrams[0]) == 1280 + assert packets == \ [ QuicSentPacket( epoch=Epoch.ONE_RTT, @@ -424,12 +409,11 @@ def test_short_header_padding(self): packet_number=0, packet_type=QuicPacketType.ONE_RTT, sent_bytes=1280, - ) - ], - ) + ) \ + ] # check builder - self.assertEqual(builder.packet_number, 1) + assert builder.packet_number == 1 def test_short_header_max_flight_bytes(self): """ @@ -441,21 +425,20 @@ def test_short_header_max_flight_bytes(self): crypto = create_crypto() builder.start_packet(QuicPacketType.ONE_RTT, crypto) - self.assertEqual(builder.remaining_flight_space, 973) + assert builder.remaining_flight_space == 973 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(builder.remaining_flight_space)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty - with self.assertRaises(QuicPacketBuilderStop): + with pytest.raises(QuicPacketBuilderStop): builder.start_packet(QuicPacketType.ONE_RTT, crypto) builder.start_frame(QuicFrameType.CRYPTO) # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 1) - self.assertEqual(len(datagrams[0]), 1000) - self.assertEqual( - packets, + assert len(datagrams) == 1 + assert len(datagrams[0]) == 1000 + assert packets == \ [ QuicSentPacket( epoch=Epoch.ONE_RTT, @@ -466,11 +449,10 @@ def test_short_header_max_flight_bytes(self): packet_type=QuicPacketType.ONE_RTT, sent_bytes=1000, ), - ], - ) + ] # check builder - self.assertEqual(builder.packet_number, 1) + assert builder.packet_number == 1 def test_short_header_max_flight_bytes_zero(self): """ @@ -483,16 +465,16 @@ def test_short_header_max_flight_bytes_zero(self): crypto = create_crypto() - with self.assertRaises(QuicPacketBuilderStop): + with pytest.raises(QuicPacketBuilderStop): builder.start_packet(QuicPacketType.ONE_RTT, crypto) builder.start_frame(QuicFrameType.CRYPTO) # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 0) + assert len(datagrams) == 0 # check builder - self.assertEqual(builder.packet_number, 0) + assert builder.packet_number == 0 def test_short_header_max_flight_bytes_zero_ack(self): """ @@ -509,16 +491,15 @@ def test_short_header_max_flight_bytes_zero_ack(self): buf = builder.start_frame(QuicFrameType.ACK) buf.push_bytes(bytes(64)) - with self.assertRaises(QuicPacketBuilderStop): + with pytest.raises(QuicPacketBuilderStop): builder.start_packet(QuicPacketType.ONE_RTT, crypto) builder.start_frame(QuicFrameType.CRYPTO) # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 1) - self.assertEqual(len(datagrams[0]), 92) - self.assertEqual( - packets, + assert len(datagrams) == 1 + assert len(datagrams[0]) == 92 + assert packets == \ [ QuicSentPacket( epoch=Epoch.ONE_RTT, @@ -529,11 +510,10 @@ def test_short_header_max_flight_bytes_zero_ack(self): packet_type=QuicPacketType.ONE_RTT, sent_bytes=92, ), - ], - ) + ] # check builder - self.assertEqual(builder.packet_number, 1) + assert builder.packet_number == 1 def test_short_header_max_total_bytes_1(self): """ @@ -544,16 +524,16 @@ def test_short_header_max_total_bytes_1(self): crypto = create_crypto() - with self.assertRaises(QuicPacketBuilderStop): + with pytest.raises(QuicPacketBuilderStop): builder.start_packet(QuicPacketType.ONE_RTT, crypto) # check datagrams datagrams, packets = builder.flush() - self.assertEqual(datagrams, []) - self.assertEqual(packets, []) + assert datagrams == [] + assert packets == [] # check builder - self.assertEqual(builder.packet_number, 0) + assert builder.packet_number == 0 def test_short_header_max_total_bytes_2(self): """ @@ -565,20 +545,19 @@ def test_short_header_max_total_bytes_2(self): crypto = create_crypto() builder.start_packet(QuicPacketType.ONE_RTT, crypto) - self.assertEqual(builder.remaining_flight_space, 773) + assert builder.remaining_flight_space == 773 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(builder.remaining_flight_space)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty - with self.assertRaises(QuicPacketBuilderStop): + with pytest.raises(QuicPacketBuilderStop): builder.start_packet(QuicPacketType.ONE_RTT, crypto) # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 1) - self.assertEqual(len(datagrams[0]), 800) - self.assertEqual( - packets, + assert len(datagrams) == 1 + assert len(datagrams[0]) == 800 + assert packets == \ [ QuicSentPacket( epoch=Epoch.ONE_RTT, @@ -588,12 +567,11 @@ def test_short_header_max_total_bytes_2(self): packet_number=0, packet_type=QuicPacketType.ONE_RTT, sent_bytes=800, - ) - ], - ) + ) \ + ] # check builder - self.assertEqual(builder.packet_number, 1) + assert builder.packet_number == 1 def test_short_header_max_total_bytes_3(self): builder = create_builder() @@ -602,27 +580,26 @@ def test_short_header_max_total_bytes_3(self): crypto = create_crypto() builder.start_packet(QuicPacketType.ONE_RTT, crypto) - self.assertEqual(builder.remaining_flight_space, 1253) + assert builder.remaining_flight_space == 1253 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(builder.remaining_flight_space)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty builder.start_packet(QuicPacketType.ONE_RTT, crypto) - self.assertEqual(builder.remaining_flight_space, 693) + assert builder.remaining_flight_space == 693 buf = builder.start_frame(QuicFrameType.CRYPTO) buf.push_bytes(bytes(builder.remaining_flight_space)) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty - with self.assertRaises(QuicPacketBuilderStop): + with pytest.raises(QuicPacketBuilderStop): builder.start_packet(QuicPacketType.ONE_RTT, crypto) # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 2) - self.assertEqual(len(datagrams[0]), 1280) - self.assertEqual(len(datagrams[1]), 720) - self.assertEqual( - packets, + assert len(datagrams) == 2 + assert len(datagrams[0]) == 1280 + assert len(datagrams[1]) == 720 + assert packets == \ [ QuicSentPacket( epoch=Epoch.ONE_RTT, @@ -642,11 +619,10 @@ def test_short_header_max_total_bytes_3(self): packet_type=QuicPacketType.ONE_RTT, sent_bytes=720, ), - ], - ) + ] # check builder - self.assertEqual(builder.packet_number, 2) + assert builder.packet_number == 2 def test_short_header_ping_only(self): """ @@ -659,14 +635,13 @@ def test_short_header_ping_only(self): # HANDSHAKE, with only a PING frame builder.start_packet(QuicPacketType.ONE_RTT, crypto) builder.start_frame(QuicFrameType.PING) - self.assertFalse(builder.packet_is_empty) + assert not builder.packet_is_empty # check datagrams datagrams, packets = builder.flush() - self.assertEqual(len(datagrams), 1) - self.assertEqual(len(datagrams[0]), 29) - self.assertEqual( - packets, + assert len(datagrams) == 1 + assert len(datagrams[0]) == 29 + assert packets == \ [ QuicSentPacket( epoch=Epoch.ONE_RTT, @@ -676,6 +651,5 @@ def test_short_header_ping_only(self): packet_number=0, packet_type=QuicPacketType.ONE_RTT, sent_bytes=29, - ) - ], - ) + ) \ + ] diff --git a/tests/test_rangeset.py b/tests/test_rangeset.py index 5e0fdf460..fdca1363a 100644 --- a/tests/test_rangeset.py +++ b/tests/test_rangeset.py @@ -1,91 +1,91 @@ from __future__ import annotations -from unittest import TestCase +import pytest from qh3.quic.rangeset import RangeSet -class RangeSetTest(TestCase): +class TestRangeSet: def test_add_single_duplicate(self): rangeset = RangeSet() rangeset.add(0) - self.assertEqual(list(rangeset), [range(0, 1)]) + assert list(rangeset) == [range(0, 1)] rangeset.add(0) - self.assertEqual(list(rangeset), [range(0, 1)]) + assert list(rangeset) == [range(0, 1)] def test_add_single_ordered(self): rangeset = RangeSet() rangeset.add(0) - self.assertEqual(list(rangeset), [range(0, 1)]) + assert list(rangeset) == [range(0, 1)] rangeset.add(1) - self.assertEqual(list(rangeset), [range(0, 2)]) + assert list(rangeset) == [range(0, 2)] rangeset.add(2) - self.assertEqual(list(rangeset), [range(0, 3)]) + assert list(rangeset) == [range(0, 3)] def test_add_single_merge(self): rangeset = RangeSet() rangeset.add(0) - self.assertEqual(list(rangeset), [range(0, 1)]) + assert list(rangeset) == [range(0, 1)] rangeset.add(2) - self.assertEqual(list(rangeset), [range(0, 1), range(2, 3)]) + assert list(rangeset) == [range(0, 1), range(2, 3)] rangeset.add(1) - self.assertEqual(list(rangeset), [range(0, 3)]) + assert list(rangeset) == [range(0, 3)] def test_add_single_reverse(self): rangeset = RangeSet() rangeset.add(2) - self.assertEqual(list(rangeset), [range(2, 3)]) + assert list(rangeset) == [range(2, 3)] rangeset.add(1) - self.assertEqual(list(rangeset), [range(1, 3)]) + assert list(rangeset) == [range(1, 3)] rangeset.add(0) - self.assertEqual(list(rangeset), [range(0, 3)]) + assert list(rangeset) == [range(0, 3)] def test_add_range_ordered(self): rangeset = RangeSet() rangeset.add(0, 2) - self.assertEqual(list(rangeset), [range(0, 2)]) + assert list(rangeset) == [range(0, 2)] rangeset.add(2, 4) - self.assertEqual(list(rangeset), [range(0, 4)]) + assert list(rangeset) == [range(0, 4)] rangeset.add(4, 6) - self.assertEqual(list(rangeset), [range(0, 6)]) + assert list(rangeset) == [range(0, 6)] def test_add_range_merge(self): rangeset = RangeSet() rangeset.add(0, 2) - self.assertEqual(list(rangeset), [range(0, 2)]) + assert list(rangeset) == [range(0, 2)] rangeset.add(3, 5) - self.assertEqual(list(rangeset), [range(0, 2), range(3, 5)]) + assert list(rangeset) == [range(0, 2), range(3, 5)] rangeset.add(2, 3) - self.assertEqual(list(rangeset), [range(0, 5)]) + assert list(rangeset) == [range(0, 5)] def test_add_range_overlap(self): rangeset = RangeSet() rangeset.add(0, 2) - self.assertEqual(list(rangeset), [range(0, 2)]) + assert list(rangeset) == [range(0, 2)] rangeset.add(3, 5) - self.assertEqual(list(rangeset), [range(0, 2), range(3, 5)]) + assert list(rangeset) == [range(0, 2), range(3, 5)] rangeset.add(1, 5) - self.assertEqual(list(rangeset), [range(0, 5)]) + assert list(rangeset) == [range(0, 5)] def test_add_range_overlap_2(self): rangeset = RangeSet() @@ -94,48 +94,46 @@ def test_add_range_overlap_2(self): rangeset.add(6, 8) rangeset.add(10, 12) rangeset.add(16, 18) - self.assertEqual( - list(rangeset), [range(2, 4), range(6, 8), range(10, 12), range(16, 18)] - ) + assert list(rangeset) == [range(2, 4), range(6, 8), range(10, 12), range(16, 18)] rangeset.add(1, 15) - self.assertEqual(list(rangeset), [range(1, 15), range(16, 18)]) + assert list(rangeset) == [range(1, 15), range(16, 18)] def test_add_range_reverse(self): rangeset = RangeSet() rangeset.add(6, 8) - self.assertEqual(list(rangeset), [range(6, 8)]) + assert list(rangeset) == [range(6, 8)] rangeset.add(3, 5) - self.assertEqual(list(rangeset), [range(3, 5), range(6, 8)]) + assert list(rangeset) == [range(3, 5), range(6, 8)] rangeset.add(0, 2) - self.assertEqual(list(rangeset), [range(0, 2), range(3, 5), range(6, 8)]) + assert list(rangeset) == [range(0, 2), range(3, 5), range(6, 8)] def test_add_range_unordered_contiguous(self): rangeset = RangeSet() rangeset.add(0, 2) - self.assertEqual(list(rangeset), [range(0, 2)]) + assert list(rangeset) == [range(0, 2)] rangeset.add(4, 6) - self.assertEqual(list(rangeset), [range(0, 2), range(4, 6)]) + assert list(rangeset) == [range(0, 2), range(4, 6)] rangeset.add(2, 4) - self.assertEqual(list(rangeset), [range(0, 6)]) + assert list(rangeset) == [range(0, 6)] def test_add_range_unordered_sparse(self): rangeset = RangeSet() rangeset.add(0, 2) - self.assertEqual(list(rangeset), [range(0, 2)]) + assert list(rangeset) == [range(0, 2)] rangeset.add(6, 8) - self.assertEqual(list(rangeset), [range(0, 2), range(6, 8)]) + assert list(rangeset) == [range(0, 2), range(6, 8)] rangeset.add(3, 5) - self.assertEqual(list(rangeset), [range(0, 2), range(3, 5), range(6, 8)]) + assert list(rangeset) == [range(0, 2), range(3, 5), range(6, 8)] def test_subtract(self): rangeset = RangeSet() @@ -143,7 +141,7 @@ def test_subtract(self): rangeset.add(20, 30) rangeset.subtract(0, 3) - self.assertEqual(list(rangeset), [range(3, 10), range(20, 30)]) + assert list(rangeset) == [range(3, 10), range(20, 30)] def test_subtract_no_change(self): rangeset = RangeSet() @@ -152,10 +150,10 @@ def test_subtract_no_change(self): rangeset.add(25, 30) rangeset.subtract(0, 5) - self.assertEqual(list(rangeset), [range(5, 10), range(15, 20), range(25, 30)]) + assert list(rangeset) == [range(5, 10), range(15, 20), range(25, 30)] rangeset.subtract(10, 15) - self.assertEqual(list(rangeset), [range(5, 10), range(15, 20), range(25, 30)]) + assert list(rangeset) == [range(5, 10), range(15, 20), range(25, 30)] def test_subtract_overlap(self): rangeset = RangeSet() @@ -163,77 +161,73 @@ def test_subtract_overlap(self): rangeset.add(6, 8) rangeset.add(10, 20) rangeset.add(30, 40) - self.assertEqual( - list(rangeset), [range(1, 4), range(6, 8), range(10, 20), range(30, 40)] - ) + assert list(rangeset) == [range(1, 4), range(6, 8), range(10, 20), range(30, 40)] rangeset.subtract(0, 2) - self.assertEqual( - list(rangeset), [range(2, 4), range(6, 8), range(10, 20), range(30, 40)] - ) + assert list(rangeset) == [range(2, 4), range(6, 8), range(10, 20), range(30, 40)] rangeset.subtract(3, 11) - self.assertEqual(list(rangeset), [range(2, 3), range(11, 20), range(30, 40)]) + assert list(rangeset) == [range(2, 3), range(11, 20), range(30, 40)] def test_subtract_split(self): rangeset = RangeSet() rangeset.add(0, 10) rangeset.subtract(2, 5) - self.assertEqual(list(rangeset), [range(0, 2), range(5, 10)]) + assert list(rangeset) == [range(0, 2), range(5, 10)] def test_bool(self): - with self.assertRaises(NotImplementedError): + with pytest.raises(NotImplementedError): bool(RangeSet()) def test_contains(self): rangeset = RangeSet() - self.assertFalse(0 in rangeset) + assert not 0 in rangeset rangeset = RangeSet([range(0, 1)]) - self.assertTrue(0 in rangeset) - self.assertFalse(1 in rangeset) + assert 0 in rangeset + assert not 1 in rangeset rangeset = RangeSet([range(0, 1), range(3, 6)]) - self.assertTrue(0 in rangeset) - self.assertFalse(1 in rangeset) - self.assertFalse(2 in rangeset) - self.assertTrue(3 in rangeset) - self.assertTrue(4 in rangeset) - self.assertTrue(5 in rangeset) - self.assertFalse(6 in rangeset) + assert 0 in rangeset + assert not 1 in rangeset + assert not 2 in rangeset + assert 3 in rangeset + assert 4 in rangeset + assert 5 in rangeset + assert not 6 in rangeset def test_eq(self): r0 = RangeSet([range(0, 1)]) r1 = RangeSet([range(1, 2), range(3, 4)]) r2 = RangeSet([range(3, 4), range(1, 2)]) - self.assertTrue(r0 == r0) - self.assertFalse(r0 == r1) - self.assertFalse(r0 == 0) + assert r0 == r0 + assert not r0 == r1 + assert not r0 == 0 - self.assertTrue(r1 == r1) - self.assertFalse(r1 == r0) - self.assertTrue(r1 == r2) - self.assertFalse(r1 == 0) + assert r1 == r1 + assert not r1 == r0 + assert r1 == r2 + assert not r1 == 0 - self.assertTrue(r2 == r2) - self.assertTrue(r2 == r1) - self.assertFalse(r2 == r0) - self.assertFalse(r2 == 0) + assert r2 == r2 + assert r2 == r1 + assert not r2 == r0 + assert not r2 == 0 def test_len(self): rangeset = RangeSet() - self.assertEqual(len(rangeset), 0) + assert len(rangeset) == 0 rangeset = RangeSet([range(0, 1)]) - self.assertEqual(len(rangeset), 1) + assert len(rangeset) == 1 def test_pop(self): rangeset = RangeSet([range(1, 2), range(3, 4)]) r = rangeset.shift() - self.assertEqual(r, range(1, 2)) - self.assertEqual(list(rangeset), [range(3, 4)]) + assert r == range(1, 2) + assert list(rangeset) == [range(3, 4)] def test_repr(self): rangeset = RangeSet([range(1, 2), range(3, 4)]) - self.assertEqual(repr(rangeset), "RangeSet([range(1, 2), range(3, 4)])") + assert repr(rangeset) == "RangeSet([range(1, 2), range(3, 4)])" diff --git a/tests/test_recovery.py b/tests/test_recovery.py index b1312dba6..5330e9902 100644 --- a/tests/test_recovery.py +++ b/tests/test_recovery.py @@ -1,7 +1,7 @@ from __future__ import annotations +import pytest import math -from unittest import TestCase from qh3 import tls from qh3.quic.packet import QuicPacketType @@ -19,52 +19,52 @@ def send_probe(): pass -class QuicPacketPacerTest(TestCase): - def setUp(self): +class TestQuicPacketPacer: + def setup_method(self): self.pacer = QuicPacketPacer() def test_no_measurement(self): - self.assertIsNone(self.pacer.next_send_time(now=0.0)) + assert self.pacer.next_send_time(now=0.0) is None self.pacer.update_after_send(now=0.0) - self.assertIsNone(self.pacer.next_send_time(now=0.0)) + assert self.pacer.next_send_time(now=0.0) is None self.pacer.update_after_send(now=0.0) def test_with_measurement(self): - self.assertIsNone(self.pacer.next_send_time(now=0.0)) + assert self.pacer.next_send_time(now=0.0) is None self.pacer.update_after_send(now=0.0) self.pacer.update_rate(congestion_window=1280000, smoothed_rtt=0.05) - self.assertEqual(self.pacer.bucket_max, 0.0008) - self.assertEqual(self.pacer.bucket_time, 0.0) - self.assertEqual(self.pacer.packet_time, 0.00005) + assert self.pacer.bucket_max == 0.0008 + assert self.pacer.bucket_time == 0.0 + assert self.pacer.packet_time == 0.00005 # 16 packets for i in range(16): - self.assertIsNone(self.pacer.next_send_time(now=1.0)) + assert self.pacer.next_send_time(now=1.0) is None self.pacer.update_after_send(now=1.0) - self.assertAlmostEqual(self.pacer.next_send_time(now=1.0), 1.00005) + assert self.pacer.next_send_time(now=1.0) == pytest.approx(1.00005) # 2 packets for i in range(2): - self.assertIsNone(self.pacer.next_send_time(now=1.00005)) + assert self.pacer.next_send_time(now=1.00005) is None self.pacer.update_after_send(now=1.00005) - self.assertAlmostEqual(self.pacer.next_send_time(now=1.00005), 1.0001) + assert self.pacer.next_send_time(now=1.00005) == pytest.approx(1.0001) # 1 packet - self.assertIsNone(self.pacer.next_send_time(now=1.0001)) + assert self.pacer.next_send_time(now=1.0001) is None self.pacer.update_after_send(now=1.0001) - self.assertAlmostEqual(self.pacer.next_send_time(now=1.0001), 1.00015) + assert self.pacer.next_send_time(now=1.0001) == pytest.approx(1.00015) # 2 packets for i in range(2): - self.assertIsNone(self.pacer.next_send_time(now=1.00015)) + assert self.pacer.next_send_time(now=1.00015) is None self.pacer.update_after_send(now=1.00015) - self.assertAlmostEqual(self.pacer.next_send_time(now=1.00015), 1.0002) + assert self.pacer.next_send_time(now=1.00015) == pytest.approx(1.0002) -class QuicPacketRecoveryTest(TestCase): - def setUp(self): +class TestQuicPacketRecovery: + def setup_method(self): self.INITIAL_SPACE = QuicPacketSpace() self.HANDSHAKE_SPACE = QuicPacketSpace() self.ONE_RTT_SPACE = QuicPacketSpace() @@ -98,23 +98,23 @@ def test_on_ack_received_ack_eliciting(self): #  packet sent self.recovery.on_packet_sent(packet, space) - self.assertEqual(self.recovery.bytes_in_flight, 1280) - self.assertEqual(space.ack_eliciting_in_flight, 1) - self.assertEqual(len(space.sent_packets), 1) + assert self.recovery.bytes_in_flight == 1280 + assert space.ack_eliciting_in_flight == 1 + assert len(space.sent_packets) == 1 # packet ack'd self.recovery.on_ack_received( space, ack_rangeset=RangeSet([range(0, 1)]), ack_delay=0.0, now=10.0 ) - self.assertEqual(self.recovery.bytes_in_flight, 0) - self.assertEqual(space.ack_eliciting_in_flight, 0) - self.assertEqual(len(space.sent_packets), 0) + assert self.recovery.bytes_in_flight == 0 + assert space.ack_eliciting_in_flight == 0 + assert len(space.sent_packets) == 0 # check RTT - self.assertTrue(self.recovery._rtt_initialized) - self.assertEqual(self.recovery._rtt_latest, 10.0) - self.assertEqual(self.recovery._rtt_min, 10.0) - self.assertEqual(self.recovery._rtt_smoothed, 10.0) + assert self.recovery._rtt_initialized + assert self.recovery._rtt_latest == 10.0 + assert self.recovery._rtt_min == 10.0 + assert self.recovery._rtt_smoothed == 10.0 def test_on_ack_received_non_ack_eliciting(self): packet = QuicSentPacket( @@ -131,23 +131,23 @@ def test_on_ack_received_non_ack_eliciting(self): #  packet sent self.recovery.on_packet_sent(packet, space) - self.assertEqual(self.recovery.bytes_in_flight, 1280) - self.assertEqual(space.ack_eliciting_in_flight, 0) - self.assertEqual(len(space.sent_packets), 1) + assert self.recovery.bytes_in_flight == 1280 + assert space.ack_eliciting_in_flight == 0 + assert len(space.sent_packets) == 1 # packet ack'd self.recovery.on_ack_received( space, ack_rangeset=RangeSet([range(0, 1)]), ack_delay=0.0, now=10.0 ) - self.assertEqual(self.recovery.bytes_in_flight, 0) - self.assertEqual(space.ack_eliciting_in_flight, 0) - self.assertEqual(len(space.sent_packets), 0) + assert self.recovery.bytes_in_flight == 0 + assert space.ack_eliciting_in_flight == 0 + assert len(space.sent_packets) == 0 # check RTT - self.assertFalse(self.recovery._rtt_initialized) - self.assertEqual(self.recovery._rtt_latest, 0.0) - self.assertEqual(self.recovery._rtt_min, math.inf) - self.assertEqual(self.recovery._rtt_smoothed, 0.0) + assert not self.recovery._rtt_initialized + assert self.recovery._rtt_latest == 0.0 + assert self.recovery._rtt_min == math.inf + assert self.recovery._rtt_smoothed == 0.0 def test_on_packet_lost_crypto(self): packet = QuicSentPacket( @@ -163,69 +163,69 @@ def test_on_packet_lost_crypto(self): space = self.INITIAL_SPACE self.recovery.on_packet_sent(packet, space) - self.assertEqual(self.recovery.bytes_in_flight, 1280) - self.assertEqual(space.ack_eliciting_in_flight, 1) - self.assertEqual(len(space.sent_packets), 1) + assert self.recovery.bytes_in_flight == 1280 + assert space.ack_eliciting_in_flight == 1 + assert len(space.sent_packets) == 1 self.recovery._detect_loss(space, now=1.0) - self.assertEqual(self.recovery.bytes_in_flight, 0) - self.assertEqual(space.ack_eliciting_in_flight, 0) - self.assertEqual(len(space.sent_packets), 0) + assert self.recovery.bytes_in_flight == 0 + assert space.ack_eliciting_in_flight == 0 + assert len(space.sent_packets) == 0 -class QuicRttMonitorTest(TestCase): +class TestQuicRttMonitor: def test_monitor(self): monitor = QuicRttMonitor() - self.assertFalse(monitor.is_rtt_increasing(rtt=10, now=1000)) - self.assertEqual(monitor._samples, [10, 0.0, 0.0, 0.0, 0.0]) - self.assertFalse(monitor._ready) + assert not monitor.is_rtt_increasing(rtt=10, now=1000) + assert monitor._samples == [10, 0.0, 0.0, 0.0, 0.0] + assert not monitor._ready # not taken into account - self.assertFalse(monitor.is_rtt_increasing(rtt=11, now=1000)) - self.assertEqual(monitor._samples, [10, 0.0, 0.0, 0.0, 0.0]) - self.assertFalse(monitor._ready) + assert not monitor.is_rtt_increasing(rtt=11, now=1000) + assert monitor._samples == [10, 0.0, 0.0, 0.0, 0.0] + assert not monitor._ready - self.assertFalse(monitor.is_rtt_increasing(rtt=11, now=1001)) - self.assertEqual(monitor._samples, [10, 11, 0.0, 0.0, 0.0]) - self.assertFalse(monitor._ready) + assert not monitor.is_rtt_increasing(rtt=11, now=1001) + assert monitor._samples == [10, 11, 0.0, 0.0, 0.0] + assert not monitor._ready - self.assertFalse(monitor.is_rtt_increasing(rtt=12, now=1002)) - self.assertEqual(monitor._samples, [10, 11, 12, 0.0, 0.0]) - self.assertFalse(monitor._ready) + assert not monitor.is_rtt_increasing(rtt=12, now=1002) + assert monitor._samples == [10, 11, 12, 0.0, 0.0] + assert not monitor._ready - self.assertFalse(monitor.is_rtt_increasing(rtt=13, now=1003)) - self.assertEqual(monitor._samples, [10, 11, 12, 13, 0.0]) - self.assertFalse(monitor._ready) + assert not monitor.is_rtt_increasing(rtt=13, now=1003) + assert monitor._samples == [10, 11, 12, 13, 0.0] + assert not monitor._ready # we now have enough samples - self.assertFalse(monitor.is_rtt_increasing(rtt=14, now=1004)) - self.assertEqual(monitor._samples, [10, 11, 12, 13, 14]) - self.assertTrue(monitor._ready) + assert not monitor.is_rtt_increasing(rtt=14, now=1004) + assert monitor._samples == [10, 11, 12, 13, 14] + assert monitor._ready - self.assertFalse(monitor.is_rtt_increasing(rtt=20, now=1005)) - self.assertEqual(monitor._increases, 0) + assert not monitor.is_rtt_increasing(rtt=20, now=1005) + assert monitor._increases == 0 - self.assertFalse(monitor.is_rtt_increasing(rtt=30, now=1006)) - self.assertEqual(monitor._increases, 0) + assert not monitor.is_rtt_increasing(rtt=30, now=1006) + assert monitor._increases == 0 - self.assertFalse(monitor.is_rtt_increasing(rtt=40, now=1007)) - self.assertEqual(monitor._increases, 0) + assert not monitor.is_rtt_increasing(rtt=40, now=1007) + assert monitor._increases == 0 - self.assertFalse(monitor.is_rtt_increasing(rtt=50, now=1008)) - self.assertEqual(monitor._increases, 0) + assert not monitor.is_rtt_increasing(rtt=50, now=1008) + assert monitor._increases == 0 - self.assertFalse(monitor.is_rtt_increasing(rtt=60, now=1009)) - self.assertEqual(monitor._increases, 1) + assert not monitor.is_rtt_increasing(rtt=60, now=1009) + assert monitor._increases == 1 - self.assertFalse(monitor.is_rtt_increasing(rtt=70, now=1010)) - self.assertEqual(monitor._increases, 2) + assert not monitor.is_rtt_increasing(rtt=70, now=1010) + assert monitor._increases == 2 - self.assertFalse(monitor.is_rtt_increasing(rtt=80, now=1011)) - self.assertEqual(monitor._increases, 3) + assert not monitor.is_rtt_increasing(rtt=80, now=1011) + assert monitor._increases == 3 - self.assertFalse(monitor.is_rtt_increasing(rtt=90, now=1012)) - self.assertEqual(monitor._increases, 4) + assert not monitor.is_rtt_increasing(rtt=90, now=1012) + assert monitor._increases == 4 - self.assertTrue(monitor.is_rtt_increasing(rtt=100, now=1013)) - self.assertEqual(monitor._increases, 5) + assert monitor.is_rtt_increasing(rtt=100, now=1013) + assert monitor._increases == 5 diff --git a/tests/test_retry.py b/tests/test_retry.py index f05a1355c..99bcd9ee0 100644 --- a/tests/test_retry.py +++ b/tests/test_retry.py @@ -1,11 +1,11 @@ from __future__ import annotations -from unittest import TestCase +import pytest from qh3.quic.retry import QuicRetryTokenHandler -class QuicRetryTokenHandlerTest(TestCase): +class TestQuicRetryTokenHandler: def test_retry_token(self): addr = ("127.0.0.1", 1234) original_destination_connection_id = b"\x08\x07\x06\05\x04\x03\x02\x01" @@ -17,23 +17,19 @@ def test_retry_token(self): token = handler.create_token( addr, original_destination_connection_id, retry_source_connection_id ) - self.assertIsNotNone(token) - self.assertEqual(len(token), 256) + assert token is not None + assert len(token) == 256 # validate token - ok - self.assertEqual( - handler.validate_token(addr, token), - (original_destination_connection_id, retry_source_connection_id), - ) + assert handler.validate_token(addr, token) == \ + (original_destination_connection_id, retry_source_connection_id) # validate token - empty - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: handler.validate_token(addr, b"") - self.assertEqual( - str(cm.exception), "Ciphertext length must be equal to key size." - ) + assert str(cm.value) == "Ciphertext length must be equal to key size." # validate token - wrong address - with self.assertRaises(ValueError) as cm: + with pytest.raises(ValueError) as cm: handler.validate_token(("1.2.3.4", 12345), token) - self.assertEqual(str(cm.exception), "Remote address does not match.") + assert str(cm.value) == "Remote address does not match." diff --git a/tests/test_stream.py b/tests/test_stream.py index 2679eca80..15c304af2 100644 --- a/tests/test_stream.py +++ b/tests/test_stream.py @@ -1,6 +1,6 @@ from __future__ import annotations -from unittest import TestCase +import pytest from qh3.quic.events import StreamDataReceived, StreamReset from qh3.quic.packet import QuicErrorCode, QuicStreamFrame @@ -8,550 +8,476 @@ from qh3.quic.stream import FinalSizeError, QuicStream -class QuicStreamTest(TestCase): +class TestQuicStream: def test_receiver_empty(self): stream = QuicStream(stream_id=0) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 0) + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 0 # empty - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"")), None - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 0) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"")) == None + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 0 def test_receiver_ordered(self): stream = QuicStream(stream_id=0) # add data at start - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")), - StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0), - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 8) - self.assertEqual(stream.receiver.highest_offset, 8) - self.assertFalse(stream.receiver.is_finished) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")) == \ + StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0) + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 8 + assert stream.receiver.highest_offset == 8 + assert not stream.receiver.is_finished # add more data - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=8, data=b"89012345")), - StreamDataReceived(data=b"89012345", end_stream=False, stream_id=0), - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 16) - self.assertEqual(stream.receiver.highest_offset, 16) - self.assertFalse(stream.receiver.is_finished) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=8, data=b"89012345")) == \ + StreamDataReceived(data=b"89012345", end_stream=False, stream_id=0) + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 16 + assert stream.receiver.highest_offset == 16 + assert not stream.receiver.is_finished # add data and fin - self.assertEqual( - stream.receiver.handle_frame( - QuicStreamFrame(offset=16, data=b"67890123", fin=True) - ), - StreamDataReceived(data=b"67890123", end_stream=True, stream_id=0), - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 24) - self.assertEqual(stream.receiver.highest_offset, 24) - self.assertTrue(stream.receiver.is_finished) + assert stream.receiver.handle_frame( + QuicStreamFrame(offset=16, data=b"67890123", fin=True) \ + ) == \ + StreamDataReceived(data=b"67890123", end_stream=True, stream_id=0) + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 24 + assert stream.receiver.highest_offset == 24 + assert stream.receiver.is_finished def test_receiver_unordered(self): stream = QuicStream(stream_id=0) # add data at offset 8 - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=8, data=b"89012345")), - None, - ) - self.assertEqual( - bytes(stream.receiver._buffer), b"\x00\x00\x00\x00\x00\x00\x00\x0089012345" - ) - self.assertEqual(list(stream.receiver._ranges), [range(8, 16)]) - self.assertEqual(stream.receiver._buffer_start, 0) - self.assertEqual(stream.receiver.highest_offset, 16) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=8, data=b"89012345")) == \ + None + assert bytes(stream.receiver._buffer) == b"\x00\x00\x00\x00\x00\x00\x00\x0089012345" + assert list(stream.receiver._ranges) == [range(8, 16)] + assert stream.receiver._buffer_start == 0 + assert stream.receiver.highest_offset == 16 # add data at offset 0 - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")), - StreamDataReceived(data=b"0123456789012345", end_stream=False, stream_id=0), - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 16) - self.assertEqual(stream.receiver.highest_offset, 16) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")) == \ + StreamDataReceived(data=b"0123456789012345", end_stream=False, stream_id=0) + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 16 + assert stream.receiver.highest_offset == 16 def test_receiver_offset_only(self): stream = QuicStream(stream_id=0) # add data at offset 0 - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"")), None - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 0) - self.assertEqual(stream.receiver.highest_offset, 0) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"")) == None + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 0 + assert stream.receiver.highest_offset == 0 # add data at offset 8 - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=8, data=b"")), None - ) - self.assertEqual( - bytes(stream.receiver._buffer), b"\x00\x00\x00\x00\x00\x00\x00\x00" - ) - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 0) - self.assertEqual(stream.receiver.highest_offset, 8) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=8, data=b"")) == None + assert bytes(stream.receiver._buffer) == b"\x00\x00\x00\x00\x00\x00\x00\x00" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 0 + assert stream.receiver.highest_offset == 8 def test_receiver_already_fully_consumed(self): stream = QuicStream(stream_id=0) # add data at offset 0 - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")), - StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0), - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 8) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")) == \ + StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0) + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 8 # add data again at offset 0 - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")), - None, - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 8) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")) == \ + None + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 8 # add data again at offset 0 - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01")), None - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 8) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01")) == None + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 8 def test_receiver_already_partially_consumed(self): stream = QuicStream(stream_id=0) - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")), - StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0), - ) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")) == \ + StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0) - self.assertEqual( - stream.receiver.handle_frame( - QuicStreamFrame(offset=0, data=b"0123456789012345") - ), - StreamDataReceived(data=b"89012345", end_stream=False, stream_id=0), - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 16) + assert stream.receiver.handle_frame( + QuicStreamFrame(offset=0, data=b"0123456789012345") \ + ) == \ + StreamDataReceived(data=b"89012345", end_stream=False, stream_id=0) + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 16 def test_receiver_already_partially_consumed_2(self): stream = QuicStream(stream_id=0) - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")), - StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0), - ) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")) == \ + StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0) - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=16, data=b"abcdefgh")), - None, - ) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=16, data=b"abcdefgh")) == \ + None - self.assertEqual( - stream.receiver.handle_frame( - QuicStreamFrame(offset=2, data=b"23456789012345") - ), - StreamDataReceived(data=b"89012345abcdefgh", end_stream=False, stream_id=0), - ) - self.assertEqual(bytes(stream.receiver._buffer), b"") - self.assertEqual(list(stream.receiver._ranges), []) - self.assertEqual(stream.receiver._buffer_start, 24) + assert stream.receiver.handle_frame( + QuicStreamFrame(offset=2, data=b"23456789012345") \ + ) == \ + StreamDataReceived(data=b"89012345abcdefgh", end_stream=False, stream_id=0) + assert bytes(stream.receiver._buffer) == b"" + assert list(stream.receiver._ranges) == [] + assert stream.receiver._buffer_start == 24 def test_receiver_fin(self): stream = QuicStream(stream_id=0) - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")), - StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0), - ) - self.assertEqual( - stream.receiver.handle_frame( - QuicStreamFrame(offset=8, data=b"89012345", fin=True) - ), - StreamDataReceived(data=b"89012345", end_stream=True, stream_id=0), - ) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")) == \ + StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0) + assert stream.receiver.handle_frame( + QuicStreamFrame(offset=8, data=b"89012345", fin=True) \ + ) == \ + StreamDataReceived(data=b"89012345", end_stream=True, stream_id=0) def test_receiver_fin_out_of_order(self): stream = QuicStream(stream_id=0) # add data at offset 8 with FIN - self.assertEqual( - stream.receiver.handle_frame( - QuicStreamFrame(offset=8, data=b"89012345", fin=True) - ), - None, - ) - self.assertEqual(stream.receiver.highest_offset, 16) - self.assertFalse(stream.receiver.is_finished) + assert stream.receiver.handle_frame( + QuicStreamFrame(offset=8, data=b"89012345", fin=True) \ + ) == \ + None + assert stream.receiver.highest_offset == 16 + assert not stream.receiver.is_finished # add data at offset 0 - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")), - StreamDataReceived(data=b"0123456789012345", end_stream=True, stream_id=0), - ) - self.assertEqual(stream.receiver.highest_offset, 16) - self.assertTrue(stream.receiver.is_finished) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")) == \ + StreamDataReceived(data=b"0123456789012345", end_stream=True, stream_id=0) + assert stream.receiver.highest_offset == 16 + assert stream.receiver.is_finished def test_receiver_fin_then_data(self): stream = QuicStream(stream_id=0) stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"0123", fin=True)) # data beyond final size - with self.assertRaises(FinalSizeError) as cm: + with pytest.raises(FinalSizeError) as cm: stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")) - self.assertEqual(str(cm.exception), "Data received beyond final size") + assert str(cm.value) == "Data received beyond final size" # final size would be lowered - with self.assertRaises(FinalSizeError) as cm: + with pytest.raises(FinalSizeError) as cm: stream.receiver.handle_frame( QuicStreamFrame(offset=0, data=b"01", fin=True) ) - self.assertEqual(str(cm.exception), "Cannot change final size") + assert str(cm.value) == "Cannot change final size" def test_receiver_fin_twice(self): stream = QuicStream(stream_id=0) - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")), - StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0), - ) - self.assertEqual( - stream.receiver.handle_frame( - QuicStreamFrame(offset=8, data=b"89012345", fin=True) - ), - StreamDataReceived(data=b"89012345", end_stream=True, stream_id=0), - ) - - self.assertEqual( - stream.receiver.handle_frame( - QuicStreamFrame(offset=8, data=b"89012345", fin=True) - ), - StreamDataReceived(data=b"", end_stream=True, stream_id=0), - ) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"01234567")) == \ + StreamDataReceived(data=b"01234567", end_stream=False, stream_id=0) + assert stream.receiver.handle_frame( + QuicStreamFrame(offset=8, data=b"89012345", fin=True) \ + ) == \ + StreamDataReceived(data=b"89012345", end_stream=True, stream_id=0) + + assert stream.receiver.handle_frame( + QuicStreamFrame(offset=8, data=b"89012345", fin=True) \ + ) == \ + StreamDataReceived(data=b"", end_stream=True, stream_id=0) def test_receiver_fin_without_data(self): stream = QuicStream(stream_id=0) - self.assertEqual( - stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"", fin=True)), - StreamDataReceived(data=b"", end_stream=True, stream_id=0), - ) + assert stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"", fin=True)) == \ + StreamDataReceived(data=b"", end_stream=True, stream_id=0) def test_receiver_reset(self): stream = QuicStream(stream_id=0) - self.assertEqual( - stream.receiver.handle_reset(final_size=4), - StreamReset(error_code=QuicErrorCode.NO_ERROR, stream_id=0), - ) - self.assertTrue(stream.receiver.is_finished) + assert stream.receiver.handle_reset(final_size=4) == \ + StreamReset(error_code=QuicErrorCode.NO_ERROR, stream_id=0) + assert stream.receiver.is_finished def test_receiver_reset_after_fin(self): stream = QuicStream(stream_id=0) stream.receiver.handle_frame(QuicStreamFrame(offset=0, data=b"0123", fin=True)) - self.assertEqual( - stream.receiver.handle_reset(final_size=4), - StreamReset(error_code=QuicErrorCode.NO_ERROR, stream_id=0), - ) + assert stream.receiver.handle_reset(final_size=4) == \ + StreamReset(error_code=QuicErrorCode.NO_ERROR, stream_id=0) def test_receiver_reset_twice(self): stream = QuicStream(stream_id=0) - self.assertEqual( - stream.receiver.handle_reset(final_size=4), - StreamReset(error_code=QuicErrorCode.NO_ERROR, stream_id=0), - ) - self.assertEqual( - stream.receiver.handle_reset(final_size=4), - StreamReset(error_code=QuicErrorCode.NO_ERROR, stream_id=0), - ) + assert stream.receiver.handle_reset(final_size=4) == \ + StreamReset(error_code=QuicErrorCode.NO_ERROR, stream_id=0) + assert stream.receiver.handle_reset(final_size=4) == \ + StreamReset(error_code=QuicErrorCode.NO_ERROR, stream_id=0) def test_receiver_reset_twice_final_size_error(self): stream = QuicStream(stream_id=0) - self.assertEqual( - stream.receiver.handle_reset(final_size=4), - StreamReset(error_code=QuicErrorCode.NO_ERROR, stream_id=0), - ) + assert stream.receiver.handle_reset(final_size=4) == \ + StreamReset(error_code=QuicErrorCode.NO_ERROR, stream_id=0) - with self.assertRaises(FinalSizeError) as cm: + with pytest.raises(FinalSizeError) as cm: stream.receiver.handle_reset(final_size=5) - self.assertEqual(str(cm.exception), "Cannot change final size") + assert str(cm.value) == "Cannot change final size" def test_receiver_stop(self): stream = QuicStream() # stop is requested stream.receiver.stop(QuicErrorCode.NO_ERROR) - self.assertTrue(stream.receiver.stop_pending) + assert stream.receiver.stop_pending # stop is sent frame = stream.receiver.get_stop_frame() - self.assertEqual(frame.error_code, QuicErrorCode.NO_ERROR) - self.assertFalse(stream.receiver.stop_pending) + assert frame.error_code == QuicErrorCode.NO_ERROR + assert not stream.receiver.stop_pending # stop is acklowledged stream.receiver.on_stop_sending_delivery(QuicDeliveryState.ACKED) - self.assertFalse(stream.receiver.stop_pending) + assert not stream.receiver.stop_pending def test_receiver_stop_lost(self): stream = QuicStream() # stop is requested stream.receiver.stop(QuicErrorCode.NO_ERROR) - self.assertTrue(stream.receiver.stop_pending) + assert stream.receiver.stop_pending # stop is sent frame = stream.receiver.get_stop_frame() - self.assertEqual(frame.error_code, QuicErrorCode.NO_ERROR) - self.assertFalse(stream.receiver.stop_pending) + assert frame.error_code == QuicErrorCode.NO_ERROR + assert not stream.receiver.stop_pending # stop is lost stream.receiver.on_stop_sending_delivery(QuicDeliveryState.LOST) - self.assertTrue(stream.receiver.stop_pending) + assert stream.receiver.stop_pending # stop is sent again frame = stream.receiver.get_stop_frame() - self.assertEqual(frame.error_code, QuicErrorCode.NO_ERROR) - self.assertFalse(stream.receiver.stop_pending) + assert frame.error_code == QuicErrorCode.NO_ERROR + assert not stream.receiver.stop_pending # stop is acklowledged stream.receiver.on_stop_sending_delivery(QuicDeliveryState.ACKED) - self.assertFalse(stream.receiver.stop_pending) + assert not stream.receiver.stop_pending def test_sender_data(self): stream = QuicStream() - self.assertEqual(stream.sender.next_offset, 0) + assert stream.sender.next_offset == 0 # nothing to send yet frame = stream.sender.get_frame(8) - self.assertIsNone(frame) + assert frame is None # write data stream.sender.write(b"0123456789012345") - self.assertEqual(list(stream.sender._pending), [range(0, 16)]) - self.assertEqual(stream.sender.next_offset, 0) + assert list(stream.sender._pending) == [range(0, 16)] + assert stream.sender.next_offset == 0 # send a chunk frame = stream.sender.get_frame(8) - self.assertEqual(frame.data, b"01234567") - self.assertFalse(frame.fin) - self.assertEqual(frame.offset, 0) - self.assertEqual(list(stream.sender._pending), [range(8, 16)]) - self.assertEqual(stream.sender.next_offset, 8) + assert frame.data == b"01234567" + assert not frame.fin + assert frame.offset == 0 + assert list(stream.sender._pending) == [range(8, 16)] + assert stream.sender.next_offset == 8 # send another chunk frame = stream.sender.get_frame(8) - self.assertEqual(frame.data, b"89012345") - self.assertFalse(frame.fin) - self.assertEqual(frame.offset, 8) - self.assertEqual(list(stream.sender._pending), []) - self.assertEqual(stream.sender.next_offset, 16) + assert frame.data == b"89012345" + assert not frame.fin + assert frame.offset == 8 + assert list(stream.sender._pending) == [] + assert stream.sender.next_offset == 16 # nothing more to send frame = stream.sender.get_frame(8) - self.assertIsNone(frame) - self.assertEqual(list(stream.sender._pending), []) - self.assertEqual(stream.sender.next_offset, 16) + assert frame is None + assert list(stream.sender._pending) == [] + assert stream.sender.next_offset == 16 # first chunk gets acknowledged stream.sender.on_data_delivery(QuicDeliveryState.ACKED, 0, 8) - self.assertFalse(stream.sender.is_finished) + assert not stream.sender.is_finished # second chunk gets acknowledged stream.sender.on_data_delivery(QuicDeliveryState.ACKED, 8, 16) - self.assertFalse(stream.sender.is_finished) + assert not stream.sender.is_finished def test_sender_data_and_fin(self): stream = QuicStream() # nothing to send yet frame = stream.sender.get_frame(8) - self.assertIsNone(frame) + assert frame is None # write data and EOF stream.sender.write(b"0123456789012345", end_stream=True) - self.assertEqual(list(stream.sender._pending), [range(0, 16)]) - self.assertEqual(stream.sender.next_offset, 0) + assert list(stream.sender._pending) == [range(0, 16)] + assert stream.sender.next_offset == 0 # send a chunk frame = stream.sender.get_frame(8) - self.assertEqual(frame.data, b"01234567") - self.assertFalse(frame.fin) - self.assertEqual(frame.offset, 0) - self.assertEqual(stream.sender.next_offset, 8) + assert frame.data == b"01234567" + assert not frame.fin + assert frame.offset == 0 + assert stream.sender.next_offset == 8 # send another chunk frame = stream.sender.get_frame(8) - self.assertEqual(frame.data, b"89012345") - self.assertTrue(frame.fin) - self.assertEqual(frame.offset, 8) - self.assertEqual(stream.sender.next_offset, 16) + assert frame.data == b"89012345" + assert frame.fin + assert frame.offset == 8 + assert stream.sender.next_offset == 16 # nothing more to send frame = stream.sender.get_frame(8) - self.assertIsNone(frame) - self.assertEqual(stream.sender.next_offset, 16) + assert frame is None + assert stream.sender.next_offset == 16 # first chunk gets acknowledged stream.sender.on_data_delivery(QuicDeliveryState.ACKED, 0, 8) - self.assertFalse(stream.sender.is_finished) + assert not stream.sender.is_finished # second chunk gets acknowledged stream.sender.on_data_delivery(QuicDeliveryState.ACKED, 8, 16) - self.assertTrue(stream.sender.is_finished) + assert stream.sender.is_finished def test_sender_data_and_fin_ack_out_of_order(self): stream = QuicStream() # nothing to send yet frame = stream.sender.get_frame(8) - self.assertIsNone(frame) + assert frame is None # write data and EOF stream.sender.write(b"0123456789012345", end_stream=True) - self.assertEqual(list(stream.sender._pending), [range(0, 16)]) - self.assertEqual(stream.sender.next_offset, 0) + assert list(stream.sender._pending) == [range(0, 16)] + assert stream.sender.next_offset == 0 # send a chunk frame = stream.sender.get_frame(8) - self.assertEqual(frame.data, b"01234567") - self.assertFalse(frame.fin) - self.assertEqual(frame.offset, 0) - self.assertEqual(stream.sender.next_offset, 8) + assert frame.data == b"01234567" + assert not frame.fin + assert frame.offset == 0 + assert stream.sender.next_offset == 8 # send another chunk frame = stream.sender.get_frame(8) - self.assertEqual(frame.data, b"89012345") - self.assertTrue(frame.fin) - self.assertEqual(frame.offset, 8) - self.assertEqual(stream.sender.next_offset, 16) + assert frame.data == b"89012345" + assert frame.fin + assert frame.offset == 8 + assert stream.sender.next_offset == 16 # nothing more to send frame = stream.sender.get_frame(8) - self.assertIsNone(frame) - self.assertEqual(stream.sender.next_offset, 16) + assert frame is None + assert stream.sender.next_offset == 16 # second chunk gets acknowledged stream.sender.on_data_delivery(QuicDeliveryState.ACKED, 8, 16) - self.assertFalse(stream.sender.is_finished) + assert not stream.sender.is_finished # first chunk gets acknowledged stream.sender.on_data_delivery(QuicDeliveryState.ACKED, 0, 8) - self.assertTrue(stream.sender.is_finished) + assert stream.sender.is_finished def test_sender_data_lost(self): stream = QuicStream() # nothing to send yet frame = stream.sender.get_frame(8) - self.assertIsNone(frame) + assert frame is None # write data and EOF stream.sender.write(b"0123456789012345", end_stream=True) - self.assertEqual(list(stream.sender._pending), [range(0, 16)]) - self.assertEqual(stream.sender.next_offset, 0) + assert list(stream.sender._pending) == [range(0, 16)] + assert stream.sender.next_offset == 0 # send a chunk - self.assertEqual( - stream.sender.get_frame(8), - QuicStreamFrame(data=b"01234567", fin=False, offset=0), - ) - self.assertEqual(list(stream.sender._pending), [range(8, 16)]) - self.assertEqual(stream.sender.next_offset, 8) + assert stream.sender.get_frame(8) == \ + QuicStreamFrame(data=b"01234567", fin=False, offset=0) + assert list(stream.sender._pending) == [range(8, 16)] + assert stream.sender.next_offset == 8 # send another chunk - self.assertEqual( - stream.sender.get_frame(8), - QuicStreamFrame(data=b"89012345", fin=True, offset=8), - ) - self.assertEqual(list(stream.sender._pending), []) - self.assertEqual(stream.sender.next_offset, 16) + assert stream.sender.get_frame(8) == \ + QuicStreamFrame(data=b"89012345", fin=True, offset=8) + assert list(stream.sender._pending) == [] + assert stream.sender.next_offset == 16 # nothing more to send - self.assertIsNone(stream.sender.get_frame(8)) - self.assertEqual(list(stream.sender._pending), []) - self.assertEqual(stream.sender.next_offset, 16) + assert stream.sender.get_frame(8) is None + assert list(stream.sender._pending) == [] + assert stream.sender.next_offset == 16 # a chunk gets lost stream.sender.on_data_delivery(QuicDeliveryState.LOST, 0, 8) - self.assertEqual(list(stream.sender._pending), [range(0, 8)]) - self.assertEqual(stream.sender.next_offset, 0) + assert list(stream.sender._pending) == [range(0, 8)] + assert stream.sender.next_offset == 0 # send chunk again - self.assertEqual( - stream.sender.get_frame(8), - QuicStreamFrame(data=b"01234567", fin=False, offset=0), - ) - self.assertEqual(list(stream.sender._pending), []) - self.assertEqual(stream.sender.next_offset, 16) + assert stream.sender.get_frame(8) == \ + QuicStreamFrame(data=b"01234567", fin=False, offset=0) + assert list(stream.sender._pending) == [] + assert stream.sender.next_offset == 16 def test_sender_data_lost_fin(self): stream = QuicStream() # nothing to send yet frame = stream.sender.get_frame(8) - self.assertIsNone(frame) + assert frame is None # write data and EOF stream.sender.write(b"0123456789012345", end_stream=True) - self.assertEqual(list(stream.sender._pending), [range(0, 16)]) - self.assertEqual(stream.sender.next_offset, 0) + assert list(stream.sender._pending) == [range(0, 16)] + assert stream.sender.next_offset == 0 # send a chunk - self.assertEqual( - stream.sender.get_frame(8), - QuicStreamFrame(data=b"01234567", fin=False, offset=0), - ) - self.assertEqual(list(stream.sender._pending), [range(8, 16)]) - self.assertEqual(stream.sender.next_offset, 8) + assert stream.sender.get_frame(8) == \ + QuicStreamFrame(data=b"01234567", fin=False, offset=0) + assert list(stream.sender._pending) == [range(8, 16)] + assert stream.sender.next_offset == 8 # send another chunk - self.assertEqual( - stream.sender.get_frame(8), - QuicStreamFrame(data=b"89012345", fin=True, offset=8), - ) - self.assertEqual(list(stream.sender._pending), []) - self.assertEqual(stream.sender.next_offset, 16) + assert stream.sender.get_frame(8) == \ + QuicStreamFrame(data=b"89012345", fin=True, offset=8) + assert list(stream.sender._pending) == [] + assert stream.sender.next_offset == 16 # nothing more to send - self.assertIsNone(stream.sender.get_frame(8)) - self.assertEqual(list(stream.sender._pending), []) - self.assertEqual(stream.sender.next_offset, 16) + assert stream.sender.get_frame(8) is None + assert list(stream.sender._pending) == [] + assert stream.sender.next_offset == 16 # a chunk gets lost stream.sender.on_data_delivery(QuicDeliveryState.LOST, 8, 16) - self.assertEqual(list(stream.sender._pending), [range(8, 16)]) - self.assertEqual(stream.sender.next_offset, 8) + assert list(stream.sender._pending) == [range(8, 16)] + assert stream.sender.next_offset == 8 # send chunk again - self.assertEqual( - stream.sender.get_frame(8), - QuicStreamFrame(data=b"89012345", fin=True, offset=8), - ) - self.assertEqual(list(stream.sender._pending), []) - self.assertEqual(stream.sender.next_offset, 16) + assert stream.sender.get_frame(8) == \ + QuicStreamFrame(data=b"89012345", fin=True, offset=8) + assert list(stream.sender._pending) == [] + assert stream.sender.next_offset == 16 # both chunks gets acknowledged stream.sender.on_data_delivery(QuicDeliveryState.ACKED, 0, 16) - self.assertTrue(stream.sender.is_finished) + assert stream.sender.is_finished def test_sender_blocked(self): stream = QuicStream() @@ -559,150 +485,150 @@ def test_sender_blocked(self): # nothing to send yet frame = stream.sender.get_frame(8, max_offset) - self.assertIsNone(frame) - self.assertEqual(list(stream.sender._pending), []) - self.assertEqual(stream.sender.next_offset, 0) + assert frame is None + assert list(stream.sender._pending) == [] + assert stream.sender.next_offset == 0 # write data, send a chunk stream.sender.write(b"0123456789012345") frame = stream.sender.get_frame(8) - self.assertEqual(frame.data, b"01234567") - self.assertFalse(frame.fin) - self.assertEqual(frame.offset, 0) - self.assertEqual(list(stream.sender._pending), [range(8, 16)]) - self.assertEqual(stream.sender.next_offset, 8) + assert frame.data == b"01234567" + assert not frame.fin + assert frame.offset == 0 + assert list(stream.sender._pending) == [range(8, 16)] + assert stream.sender.next_offset == 8 # send is limited by peer frame = stream.sender.get_frame(8, max_offset) - self.assertEqual(frame.data, b"8901") - self.assertFalse(frame.fin) - self.assertEqual(frame.offset, 8) - self.assertEqual(list(stream.sender._pending), [range(12, 16)]) - self.assertEqual(stream.sender.next_offset, 12) + assert frame.data == b"8901" + assert not frame.fin + assert frame.offset == 8 + assert list(stream.sender._pending) == [range(12, 16)] + assert stream.sender.next_offset == 12 # unable to send, blocked frame = stream.sender.get_frame(8, max_offset) - self.assertIsNone(frame) - self.assertEqual(list(stream.sender._pending), [range(12, 16)]) - self.assertEqual(stream.sender.next_offset, 12) + assert frame is None + assert list(stream.sender._pending) == [range(12, 16)] + assert stream.sender.next_offset == 12 # write more data, still blocked stream.sender.write(b"abcdefgh") frame = stream.sender.get_frame(8, max_offset) - self.assertIsNone(frame) - self.assertEqual(list(stream.sender._pending), [range(12, 24)]) - self.assertEqual(stream.sender.next_offset, 12) + assert frame is None + assert list(stream.sender._pending) == [range(12, 24)] + assert stream.sender.next_offset == 12 # peer raises limit, send some data max_offset += 8 frame = stream.sender.get_frame(8, max_offset) - self.assertEqual(frame.data, b"2345abcd") - self.assertFalse(frame.fin) - self.assertEqual(frame.offset, 12) - self.assertEqual(list(stream.sender._pending), [range(20, 24)]) - self.assertEqual(stream.sender.next_offset, 20) + assert frame.data == b"2345abcd" + assert not frame.fin + assert frame.offset == 12 + assert list(stream.sender._pending) == [range(20, 24)] + assert stream.sender.next_offset == 20 # peer raises limit again, send remaining data max_offset += 8 frame = stream.sender.get_frame(8, max_offset) - self.assertEqual(frame.data, b"efgh") - self.assertFalse(frame.fin) - self.assertEqual(frame.offset, 20) - self.assertEqual(list(stream.sender._pending), []) - self.assertEqual(stream.sender.next_offset, 24) + assert frame.data == b"efgh" + assert not frame.fin + assert frame.offset == 20 + assert list(stream.sender._pending) == [] + assert stream.sender.next_offset == 24 # nothing more to send frame = stream.sender.get_frame(8, max_offset) - self.assertIsNone(frame) + assert frame is None def test_sender_fin_only(self): stream = QuicStream() # nothing to send yet - self.assertTrue(stream.sender.buffer_is_empty) + assert stream.sender.buffer_is_empty frame = stream.sender.get_frame(8) - self.assertIsNone(frame) + assert frame is None # write EOF stream.sender.write(b"", end_stream=True) - self.assertFalse(stream.sender.buffer_is_empty) + assert not stream.sender.buffer_is_empty frame = stream.sender.get_frame(8) - self.assertEqual(frame.data, b"") - self.assertTrue(frame.fin) - self.assertEqual(frame.offset, 0) + assert frame.data == b"" + assert frame.fin + assert frame.offset == 0 # nothing more to send - self.assertFalse(stream.sender.buffer_is_empty) # FIXME? + assert not stream.sender.buffer_is_empty# FIXME? frame = stream.sender.get_frame(8) - self.assertIsNone(frame) - self.assertTrue(stream.sender.buffer_is_empty) + assert frame is None + assert stream.sender.buffer_is_empty def test_sender_fin_only_despite_blocked(self): stream = QuicStream() # nothing to send yet - self.assertTrue(stream.sender.buffer_is_empty) + assert stream.sender.buffer_is_empty frame = stream.sender.get_frame(8) - self.assertIsNone(frame) + assert frame is None # write EOF stream.sender.write(b"", end_stream=True) - self.assertFalse(stream.sender.buffer_is_empty) + assert not stream.sender.buffer_is_empty frame = stream.sender.get_frame(8) - self.assertEqual(frame.data, b"") - self.assertTrue(frame.fin) - self.assertEqual(frame.offset, 0) + assert frame.data == b"" + assert frame.fin + assert frame.offset == 0 # nothing more to send - self.assertFalse(stream.sender.buffer_is_empty) # FIXME? + assert not stream.sender.buffer_is_empty# FIXME? frame = stream.sender.get_frame(8) - self.assertIsNone(frame) - self.assertTrue(stream.sender.buffer_is_empty) + assert frame is None + assert stream.sender.buffer_is_empty def test_sender_reset(self): stream = QuicStream() # reset is requested stream.sender.reset(QuicErrorCode.NO_ERROR) - self.assertTrue(stream.sender.reset_pending) + assert stream.sender.reset_pending # reset is sent reset = stream.sender.get_reset_frame() - self.assertEqual(reset.error_code, QuicErrorCode.NO_ERROR) - self.assertEqual(reset.final_size, 0) - self.assertFalse(stream.sender.reset_pending) - self.assertFalse(stream.sender.is_finished) + assert reset.error_code == QuicErrorCode.NO_ERROR + assert reset.final_size == 0 + assert not stream.sender.reset_pending + assert not stream.sender.is_finished # reset is acklowledged stream.sender.on_reset_delivery(QuicDeliveryState.ACKED) - self.assertFalse(stream.sender.reset_pending) - self.assertTrue(stream.sender.is_finished) + assert not stream.sender.reset_pending + assert stream.sender.is_finished def test_sender_reset_lost(self): stream = QuicStream() # reset is requested stream.sender.reset(QuicErrorCode.NO_ERROR) - self.assertTrue(stream.sender.reset_pending) + assert stream.sender.reset_pending # reset is sent reset = stream.sender.get_reset_frame() - self.assertEqual(reset.error_code, QuicErrorCode.NO_ERROR) - self.assertEqual(reset.final_size, 0) - self.assertFalse(stream.sender.reset_pending) + assert reset.error_code == QuicErrorCode.NO_ERROR + assert reset.final_size == 0 + assert not stream.sender.reset_pending # reset is lost stream.sender.on_reset_delivery(QuicDeliveryState.LOST) - self.assertTrue(stream.sender.reset_pending) - self.assertFalse(stream.sender.is_finished) + assert stream.sender.reset_pending + assert not stream.sender.is_finished # reset is sent again reset = stream.sender.get_reset_frame() - self.assertEqual(reset.error_code, QuicErrorCode.NO_ERROR) - self.assertEqual(reset.final_size, 0) - self.assertFalse(stream.sender.reset_pending) + assert reset.error_code == QuicErrorCode.NO_ERROR + assert reset.final_size == 0 + assert not stream.sender.reset_pending # reset is acklowledged stream.sender.on_reset_delivery(QuicDeliveryState.ACKED) - self.assertFalse(stream.sender.reset_pending) - self.assertTrue(stream.sender.is_finished) + assert not stream.sender.reset_pending + assert stream.sender.is_finished diff --git a/tests/test_tls.py b/tests/test_tls.py index 57b031561..b37b7f39d 100644 --- a/tests/test_tls.py +++ b/tests/test_tls.py @@ -1,8 +1,8 @@ from __future__ import annotations +import pytest import binascii import ssl -from unittest import TestCase from cryptography.hazmat.primitives import serialization @@ -75,10 +75,10 @@ ) -class BufferTest(TestCase): +class TestBuffer: def test_pull_block_truncated(self): buf = Buffer(capacity=0) - with self.assertRaises(BufferReadError): + with pytest.raises(BufferReadError): with pull_block(buf, 1): pass @@ -100,7 +100,7 @@ def reset_buffers(buffers): buffers[k].seek(0) -class ContextTest(TestCase): +class TestContext: def create_client( self, alpn_protocols=None, cadata=None, cafile=SERVER_CACERTFILE, **kwargs ): @@ -117,7 +117,7 @@ def create_client( CLIENT_QUIC_TRANSPORT_PARAMETERS, ) ] - self.assertEqual(client.state, State.CLIENT_HANDSHAKE_START) + assert client.state == State.CLIENT_HANDSHAKE_START return client def create_server(self, alpn_protocols=None, **kwargs): @@ -138,34 +138,34 @@ def create_server(self, alpn_protocols=None, **kwargs): SERVER_QUIC_TRANSPORT_PARAMETERS, ) ] - self.assertEqual(server.state, State.SERVER_EXPECT_CLIENT_HELLO) + assert server.state == State.SERVER_EXPECT_CLIENT_HELLO return server def test_client_unexpected_message(self): client = self.create_client() client.state = State.CLIENT_EXPECT_SERVER_HELLO - with self.assertRaises(tls.AlertUnexpectedMessage): + with pytest.raises(tls.AlertUnexpectedMessage): client.handle_message(b"\x00\x00\x00\x00", create_buffers()) client.state = State.CLIENT_EXPECT_ENCRYPTED_EXTENSIONS - with self.assertRaises(tls.AlertUnexpectedMessage): + with pytest.raises(tls.AlertUnexpectedMessage): client.handle_message(b"\x00\x00\x00\x00", create_buffers()) client.state = State.CLIENT_EXPECT_CERTIFICATE_REQUEST_OR_CERTIFICATE - with self.assertRaises(tls.AlertUnexpectedMessage): + with pytest.raises(tls.AlertUnexpectedMessage): client.handle_message(b"\x00\x00\x00\x00", create_buffers()) client.state = State.CLIENT_EXPECT_CERTIFICATE_VERIFY - with self.assertRaises(tls.AlertUnexpectedMessage): + with pytest.raises(tls.AlertUnexpectedMessage): client.handle_message(b"\x00\x00\x00\x00", create_buffers()) client.state = State.CLIENT_EXPECT_FINISHED - with self.assertRaises(tls.AlertUnexpectedMessage): + with pytest.raises(tls.AlertUnexpectedMessage): client.handle_message(b"\x00\x00\x00\x00", create_buffers()) client.state = State.CLIENT_POST_HANDSHAKE - with self.assertRaises(tls.AlertUnexpectedMessage): + with pytest.raises(tls.AlertUnexpectedMessage): client.handle_message(b"\x00\x00\x00\x00", create_buffers()) def test_client_bad_certificate_verify_data(self): @@ -175,7 +175,7 @@ def test_client_bad_certificate_verify_data(self): # Send client hello. client_buf = create_buffers() client.handle_message(b"", client_buf) - self.assertEqual(client.state, State.CLIENT_EXPECT_SERVER_HELLO) + assert client.state == State.CLIENT_EXPECT_SERVER_HELLO server_input = merge_buffers(client_buf) reset_buffers(client_buf) @@ -185,7 +185,7 @@ def test_client_bad_certificate_verify_data(self): # finished. server_buf = create_buffers() server.handle_message(server_input, server_buf) - self.assertEqual(server.state, State.SERVER_EXPECT_FINISHED) + assert server.state == State.SERVER_EXPECT_FINISHED client_input = merge_buffers(server_buf) reset_buffers(server_buf) @@ -194,7 +194,7 @@ def test_client_bad_certificate_verify_data(self): # Handle server hello, encrypted extensions, certificate, certificate verify, # finished. - with self.assertRaises(tls.AlertDecryptError): + with pytest.raises(tls.AlertDecryptError): client.handle_message(client_input, client_buf) def test_client_bad_finished_verify_data(self): @@ -204,7 +204,7 @@ def test_client_bad_finished_verify_data(self): # Send client hello. client_buf = create_buffers() client.handle_message(b"", client_buf) - self.assertEqual(client.state, State.CLIENT_EXPECT_SERVER_HELLO) + assert client.state == State.CLIENT_EXPECT_SERVER_HELLO server_input = merge_buffers(client_buf) reset_buffers(client_buf) @@ -214,7 +214,7 @@ def test_client_bad_finished_verify_data(self): # finished. server_buf = create_buffers() server.handle_message(server_input, server_buf) - self.assertEqual(server.state, State.SERVER_EXPECT_FINISHED) + assert server.state == State.SERVER_EXPECT_FINISHED client_input = merge_buffers(server_buf) reset_buffers(server_buf) @@ -223,29 +223,29 @@ def test_client_bad_finished_verify_data(self): # Handle server hello, encrypted extensions, certificate, certificate verify, # finished. - with self.assertRaises(tls.AlertDecryptError): + with pytest.raises(tls.AlertDecryptError): client.handle_message(client_input, client_buf) def test_server_unexpected_message(self): server = self.create_server() server.state = State.SERVER_EXPECT_CLIENT_HELLO - with self.assertRaises(tls.AlertUnexpectedMessage): + with pytest.raises(tls.AlertUnexpectedMessage): server.handle_message(b"\x00\x00\x00\x00", create_buffers()) server.state = State.SERVER_EXPECT_FINISHED - with self.assertRaises(tls.AlertUnexpectedMessage): + with pytest.raises(tls.AlertUnexpectedMessage): server.handle_message(b"\x00\x00\x00\x00", create_buffers()) server.state = State.SERVER_POST_HANDSHAKE - with self.assertRaises(tls.AlertUnexpectedMessage): + with pytest.raises(tls.AlertUnexpectedMessage): server.handle_message(b"\x00\x00\x00\x00", create_buffers()) def _server_fail_hello(self, client, server): # Send client hello. client_buf = create_buffers() client.handle_message(b"", client_buf) - self.assertEqual(client.state, State.CLIENT_EXPECT_SERVER_HELLO) + assert client.state == State.CLIENT_EXPECT_SERVER_HELLO server_input = merge_buffers(client_buf) reset_buffers(client_buf) @@ -258,9 +258,9 @@ def test_server_unsupported_cipher_suite(self): server = self.create_server(cipher_suites=[tls.CipherSuite.AES_256_GCM_SHA384]) - with self.assertRaises(tls.AlertHandshakeFailure) as cm: + with pytest.raises(tls.AlertHandshakeFailure) as cm: self._server_fail_hello(client, server) - self.assertEqual(str(cm.exception), "No supported cipher suite") + assert str(cm.value) == "No supported cipher suite" def test_server_unsupported_signature_algorithm(self): client = self.create_client() @@ -268,9 +268,9 @@ def test_server_unsupported_signature_algorithm(self): server = self.create_server() - with self.assertRaises(tls.AlertHandshakeFailure) as cm: + with pytest.raises(tls.AlertHandshakeFailure) as cm: self._server_fail_hello(client, server) - self.assertEqual(str(cm.exception), "No supported signature algorithm") + assert str(cm.value) == "No supported signature algorithm" def test_server_unsupported_version(self): client = self.create_client() @@ -278,9 +278,9 @@ def test_server_unsupported_version(self): server = self.create_server() - with self.assertRaises(tls.AlertProtocolVersion) as cm: + with pytest.raises(tls.AlertProtocolVersion) as cm: self._server_fail_hello(client, server) - self.assertEqual(str(cm.exception), "No supported protocol version") + assert str(cm.value) == "No supported protocol version" def test_server_bad_finished_verify_data(self): client = self.create_client() @@ -289,7 +289,7 @@ def test_server_bad_finished_verify_data(self): # Send client hello. client_buf = create_buffers() client.handle_message(b"", client_buf) - self.assertEqual(client.state, State.CLIENT_EXPECT_SERVER_HELLO) + assert client.state == State.CLIENT_EXPECT_SERVER_HELLO server_input = merge_buffers(client_buf) reset_buffers(client_buf) @@ -299,7 +299,7 @@ def test_server_bad_finished_verify_data(self): # finished. server_buf = create_buffers() server.handle_message(server_input, server_buf) - self.assertEqual(server.state, State.SERVER_EXPECT_FINISHED) + assert server.state == State.SERVER_EXPECT_FINISHED client_input = merge_buffers(server_buf) reset_buffers(server_buf) @@ -308,7 +308,7 @@ def test_server_bad_finished_verify_data(self): # # Send finished. client.handle_message(client_input, client_buf) - self.assertEqual(client.state, State.CLIENT_POST_HANDSHAKE) + assert client.state == State.CLIENT_POST_HANDSHAKE server_input = merge_buffers(client_buf) reset_buffers(client_buf) @@ -316,17 +316,17 @@ def test_server_bad_finished_verify_data(self): server_input = server_input[:-4] + bytes(4) # Handle finished. - with self.assertRaises(tls.AlertDecryptError): + with pytest.raises(tls.AlertDecryptError): server.handle_message(server_input, server_buf) def _handshake(self, client, server): # Send client hello. client_buf = create_buffers() client.handle_message(b"", client_buf) - self.assertEqual(client.state, State.CLIENT_EXPECT_SERVER_HELLO) + assert client.state == State.CLIENT_EXPECT_SERVER_HELLO server_input = merge_buffers(client_buf) - self.assertGreaterEqual(len(server_input), 181) - self.assertLessEqual(len(server_input), 1800) + assert len(server_input) >= 181 + assert len(server_input) <= 1800 reset_buffers(client_buf) # Handle client hello. @@ -335,10 +335,10 @@ def _handshake(self, client, server): # finished, (session ticket). server_buf = create_buffers() server.handle_message(server_input, server_buf) - self.assertEqual(server.state, State.SERVER_EXPECT_FINISHED) + assert server.state == State.SERVER_EXPECT_FINISHED client_input = merge_buffers(server_buf) - self.assertGreaterEqual(len(client_input), 539) - self.assertLessEqual(len(client_input), 4000) + assert len(client_input) >= 539 + assert len(client_input) <= 4000 reset_buffers(server_buf) @@ -347,28 +347,24 @@ def _handshake(self, client, server): # # Send finished. client.handle_message(client_input, client_buf) - self.assertEqual(client.state, State.CLIENT_POST_HANDSHAKE) + assert client.state == State.CLIENT_POST_HANDSHAKE server_input = merge_buffers(client_buf) - self.assertEqual(len(server_input), 36) + assert len(server_input) == 36 reset_buffers(client_buf) # Handle finished. server.handle_message(server_input, server_buf) - self.assertEqual(server.state, State.SERVER_POST_HANDSHAKE) + assert server.state == State.SERVER_POST_HANDSHAKE client_input = merge_buffers(server_buf) - self.assertEqual(len(client_input), 0) + assert len(client_input) == 0 # check keys match - self.assertEqual(client._dec_key, server._enc_key) - self.assertEqual(client._enc_key, server._dec_key) + assert client._dec_key == server._enc_key + assert client._enc_key == server._dec_key # check cipher suite - self.assertEqual( - client.key_schedule.cipher_suite, tls.CipherSuite.AES_128_GCM_SHA256 - ) - self.assertEqual( - server.key_schedule.cipher_suite, tls.CipherSuite.AES_128_GCM_SHA256 - ) + assert client.key_schedule.cipher_suite == tls.CipherSuite.AES_128_GCM_SHA256 + assert server.key_schedule.cipher_suite == tls.CipherSuite.AES_128_GCM_SHA256 def test_handshake(self): client = self.create_client() @@ -377,8 +373,8 @@ def test_handshake(self): self._handshake(client, server) # check ALPN matches - self.assertEqual(client.alpn_negotiated, None) - self.assertEqual(server.alpn_negotiated, None) + assert client.alpn_negotiated == None + assert server.alpn_negotiated == None def _test_handshake_with_certificate(self, certificate, private_key): server = self.create_server() @@ -412,8 +408,8 @@ def _test_handshake_with_certificate(self, certificate, private_key): self._handshake(client, server) # check ALPN matches - self.assertEqual(client.alpn_negotiated, None) - self.assertEqual(server.alpn_negotiated, None) + assert client.alpn_negotiated == None + assert server.alpn_negotiated == None def test_handshake_with_ec_certificate(self): self._test_handshake_with_certificate( @@ -436,16 +432,16 @@ def test_handshake_with_alpn(self): self._handshake(client, server) # check ALPN matches - self.assertEqual(client.alpn_negotiated, "hq-20") - self.assertEqual(server.alpn_negotiated, "hq-20") + assert client.alpn_negotiated == "hq-20" + assert server.alpn_negotiated == "hq-20" def test_handshake_with_alpn_fail(self): client = self.create_client(alpn_protocols=["hq-20"]) server = self.create_server(alpn_protocols=["h3-20"]) - with self.assertRaises(tls.AlertHandshakeFailure) as cm: + with pytest.raises(tls.AlertHandshakeFailure) as cm: self._handshake(client, server) - self.assertEqual(str(cm.exception), "No common ALPN protocols") + assert str(cm.value) == "No common ALPN protocols" def test_handshake_with_rsa_pkcs1_sha256_signature(self): client = self.create_client() @@ -458,9 +454,9 @@ def test_handshake_with_certificate_error(self): client = self.create_client(cafile=None) server = self.create_server() - with self.assertRaises(tls.AlertBadCertificate) as cm: + with pytest.raises(tls.AlertBadCertificate) as cm: self._handshake(client, server) - self.assertEqual(str(cm.exception), "unable to get local issuer certificate") + assert str(cm.value) == "unable to get local issuer certificate" def test_handshake_with_certificate_no_verify(self): client = self.create_client(cafile=None, verify_mode=ssl.CERT_NONE) @@ -483,7 +479,7 @@ def test_handshake_with_x25519(self): try: self._handshake(client, server) except CryptoError as exc: - self.skipTest(str(exc)) + pytest.skip(str(exc)) def test_session_ticket(self): client_tickets = [] @@ -511,16 +507,14 @@ def first_handshake(): self._handshake(client, server) # check session resumption was not used - self.assertFalse(client.session_resumed) - self.assertFalse(server.session_resumed) + assert not client.session_resumed + assert not server.session_resumed # check tickets match - self.assertEqual(len(client_tickets), 1) - self.assertEqual(len(server_tickets), 1) - self.assertEqual(client_tickets[0].ticket, server_tickets[0].ticket) - self.assertEqual( - client_tickets[0].resumption_secret, server_tickets[0].resumption_secret - ) + assert len(client_tickets) == 1 + assert len(server_tickets) == 1 + assert client_tickets[0].ticket == server_tickets[0].ticket + assert client_tickets[0].resumption_secret == server_tickets[0].resumption_secret def second_handshake(): client = self.create_client() @@ -532,10 +526,10 @@ def second_handshake(): # Send client hello with pre_shared_key. client_buf = create_buffers() client.handle_message(b"", client_buf) - self.assertEqual(client.state, State.CLIENT_EXPECT_SERVER_HELLO) + assert client.state == State.CLIENT_EXPECT_SERVER_HELLO server_input = merge_buffers(client_buf) - self.assertGreaterEqual(len(server_input), 383) - self.assertLessEqual(len(server_input), 1800) + assert len(server_input) >= 383 + assert len(server_input) <= 1800 reset_buffers(client_buf) # Handle client hello. @@ -543,9 +537,9 @@ def second_handshake(): # Send server hello, encrypted extensions, finished. server_buf = create_buffers() server.handle_message(server_input, server_buf) - self.assertEqual(server.state, State.SERVER_EXPECT_FINISHED) + assert server.state == State.SERVER_EXPECT_FINISHED client_input = merge_buffers(server_buf) - self.assertEqual(len(client_input), 1410) + assert len(client_input) == 1410 reset_buffers(server_buf) # Handle server hello, encrypted extensions, certificate, @@ -553,27 +547,27 @@ def second_handshake(): # # Send finished. client.handle_message(client_input, client_buf) - self.assertEqual(client.state, State.CLIENT_POST_HANDSHAKE) + assert client.state == State.CLIENT_POST_HANDSHAKE server_input = merge_buffers(client_buf) - self.assertEqual(len(server_input), 36) + assert len(server_input) == 36 reset_buffers(client_buf) # Handle finished. # # Send new_session_ticket. server.handle_message(server_input, server_buf) - self.assertEqual(server.state, State.SERVER_POST_HANDSHAKE) + assert server.state == State.SERVER_POST_HANDSHAKE client_input = merge_buffers(server_buf) - self.assertEqual(len(client_input), 0) + assert len(client_input) == 0 reset_buffers(server_buf) # check keys match - self.assertEqual(client._dec_key, server._enc_key) - self.assertEqual(client._enc_key, server._dec_key) + assert client._dec_key == server._enc_key + assert client._enc_key == server._dec_key # check session resumption was used - self.assertTrue(client.session_resumed) - self.assertTrue(server.session_resumed) + assert client.session_resumed + assert server.session_resumed def second_handshake_bad_binder(): client = self.create_client() @@ -585,10 +579,10 @@ def second_handshake_bad_binder(): # send client hello with pre_shared_key client_buf = create_buffers() client.handle_message(b"", client_buf) - self.assertEqual(client.state, State.CLIENT_EXPECT_SERVER_HELLO) + assert client.state == State.CLIENT_EXPECT_SERVER_HELLO server_input = merge_buffers(client_buf) - self.assertGreaterEqual(len(server_input), 383) - self.assertLessEqual(len(server_input), 1800) + assert len(server_input) >= 383 + assert len(server_input) <= 1800 reset_buffers(client_buf) # tamper with binder @@ -597,9 +591,9 @@ def second_handshake_bad_binder(): # handle client hello # send server hello, encrypted extensions, finished server_buf = create_buffers() - with self.assertRaises(tls.AlertHandshakeFailure) as cm: + with pytest.raises(tls.AlertHandshakeFailure) as cm: server.handle_message(server_input, server_buf) - self.assertEqual(str(cm.exception), "PSK validation failed") + assert str(cm.value) == "PSK validation failed" def second_handshake_bad_pre_shared_key(): client = self.create_client() @@ -611,28 +605,28 @@ def second_handshake_bad_pre_shared_key(): # send client hello with pre_shared_key client_buf = create_buffers() client.handle_message(b"", client_buf) - self.assertEqual(client.state, State.CLIENT_EXPECT_SERVER_HELLO) + assert client.state == State.CLIENT_EXPECT_SERVER_HELLO server_input = merge_buffers(client_buf) - self.assertGreaterEqual(len(server_input), 383) - self.assertLessEqual(len(server_input), 1800) + assert len(server_input) >= 383 + assert len(server_input) <= 1800 reset_buffers(client_buf) # handle client hello # send server hello, encrypted extensions, finished server_buf = create_buffers() server.handle_message(server_input, server_buf) - self.assertEqual(server.state, State.SERVER_EXPECT_FINISHED) + assert server.state == State.SERVER_EXPECT_FINISHED # tamper with pre_share_key index buf = server_buf[tls.Epoch.INITIAL] buf.seek(buf.tell() - 1) buf.push_uint8(1) client_input = merge_buffers(server_buf) - self.assertEqual(len(client_input), 1410) + assert len(client_input) == 1410 reset_buffers(server_buf) # handle server hello and bomb - with self.assertRaises(tls.AlertIllegalParameter): + with pytest.raises(tls.AlertIllegalParameter): client.handle_message(client_input, client_buf) first_handshake() @@ -641,38 +635,31 @@ def second_handshake_bad_pre_shared_key(): second_handshake_bad_pre_shared_key() -class TlsTest(TestCase): +class TestTls: def test_pull_client_hello(self): buf = Buffer(data=load("tls_client_hello.bin")) hello = pull_client_hello(buf) - self.assertTrue(buf.eof()) + assert buf.eof() - self.assertEqual( - hello.random, + assert hello.random == \ binascii.unhexlify( "18b2b23bf3e44b5d52ccfe7aecbc5ff14eadc3d349fabf804d71f165ae76e7d5" - ), - ) - self.assertEqual( - hello.legacy_session_id, + ) + assert hello.legacy_session_id == \ binascii.unhexlify( "9aee82a2d186c1cb32a329d9dcfe004a1a438ad0485a53c6bfcf55c132a23235" - ), - ) - self.assertEqual( - hello.cipher_suites, + ) + assert hello.cipher_suites == \ [ tls.CipherSuite.AES_256_GCM_SHA384, tls.CipherSuite.AES_128_GCM_SHA256, tls.CipherSuite.CHACHA20_POLY1305_SHA256, - ], - ) - self.assertEqual(hello.legacy_compression_methods, [tls.CompressionMethod.NULL]) + ] + assert hello.legacy_compression_methods == [tls.CompressionMethod.NULL] # extensions - self.assertEqual(hello.alpn_protocols, None) - self.assertEqual( - hello.key_share, + assert hello.alpn_protocols == None + assert hello.key_share == \ [ ( tls.Group.SECP256R1, @@ -682,57 +669,45 @@ def test_pull_client_hello(self): "b0" ), ) - ], - ) - self.assertEqual( - hello.psk_key_exchange_modes, [tls.PskKeyExchangeMode.PSK_DHE_KE] - ) - self.assertEqual(hello.server_name, None) - self.assertEqual( - hello.signature_algorithms, + ] + assert hello.psk_key_exchange_modes == [tls.PskKeyExchangeMode.PSK_DHE_KE] + assert hello.server_name == None + assert hello.signature_algorithms == \ [ tls.SignatureAlgorithm.RSA_PSS_RSAE_SHA256, tls.SignatureAlgorithm.ECDSA_SECP256R1_SHA256, tls.SignatureAlgorithm.RSA_PKCS1_SHA256, tls.SignatureAlgorithm.RSA_PKCS1_SHA1, - ], - ) - self.assertEqual(hello.supported_groups, [tls.Group.SECP256R1]) - self.assertEqual( - hello.supported_versions, + ] + assert hello.supported_groups == [tls.Group.SECP256R1] + assert hello.supported_versions == \ [ tls.TLS_VERSION_1_3, - ], - ) + ] def test_pull_client_hello_with_alpn(self): buf = Buffer(data=load("tls_client_hello_with_alpn.bin")) hello = pull_client_hello(buf) - self.assertTrue(buf.eof()) + assert buf.eof() - self.assertEqual( - hello.random, + assert hello.random == \ binascii.unhexlify( "ed575c6fbd599c4dfaabd003dca6e860ccdb0e1782c1af02e57bf27cb6479b76" - ), - ) - self.assertEqual(hello.legacy_session_id, b"") - self.assertEqual( - hello.cipher_suites, + ) + assert hello.legacy_session_id == b"" + assert hello.cipher_suites == \ [ tls.CipherSuite.AES_128_GCM_SHA256, tls.CipherSuite.AES_256_GCM_SHA384, tls.CipherSuite.CHACHA20_POLY1305_SHA256, tls.CipherSuite.EMPTY_RENEGOTIATION_INFO_SCSV, - ], - ) - self.assertEqual(hello.legacy_compression_methods, [tls.CompressionMethod.NULL]) + ] + assert hello.legacy_compression_methods == [tls.CompressionMethod.NULL] # extensions - self.assertEqual(hello.alpn_protocols, ["h3-19"]) - self.assertEqual(hello.early_data, False) - self.assertEqual( - hello.key_share, + assert hello.alpn_protocols == ["h3-19"] + assert hello.early_data == False + assert hello.key_share == \ [ ( tls.Group.SECP256R1, @@ -742,14 +717,10 @@ def test_pull_client_hello_with_alpn(self): "22" ), ) - ], - ) - self.assertEqual( - hello.psk_key_exchange_modes, [tls.PskKeyExchangeMode.PSK_DHE_KE] - ) - self.assertEqual(hello.server_name, "cloudflare-quic.com") - self.assertEqual( - hello.signature_algorithms, + ] + assert hello.psk_key_exchange_modes == [tls.PskKeyExchangeMode.PSK_DHE_KE] + assert hello.server_name == "cloudflare-quic.com" + assert hello.signature_algorithms == \ [ tls.SignatureAlgorithm.ECDSA_SECP256R1_SHA256, tls.SignatureAlgorithm.ECDSA_SECP384R1_SHA384, @@ -765,31 +736,27 @@ def test_pull_client_hello_with_alpn(self): tls.SignatureAlgorithm.RSA_PKCS1_SHA256, tls.SignatureAlgorithm.RSA_PKCS1_SHA384, tls.SignatureAlgorithm.RSA_PKCS1_SHA512, - ], - ) - self.assertEqual( - hello.supported_groups, + ] + assert hello.supported_groups == \ [ tls.Group.SECP256R1, tls.Group.X25519, tls.Group.SECP384R1, tls.Group.SECP521R1, - ], - ) - self.assertEqual(hello.supported_versions, [tls.TLS_VERSION_1_3]) + ] + assert hello.supported_versions == [tls.TLS_VERSION_1_3] # serialize buf = Buffer(1000) push_client_hello(buf, hello) - self.assertEqual(len(buf.data), len(load("tls_client_hello_with_alpn.bin"))) + assert len(buf.data) == len(load("tls_client_hello_with_alpn.bin")) def test_pull_client_hello_with_psk(self): buf = Buffer(data=load("tls_client_hello_with_psk.bin")) hello = pull_client_hello(buf) - self.assertEqual(hello.early_data, True) - self.assertEqual( - hello.pre_shared_key, + assert hello.early_data == True + assert hello.pre_shared_key == \ tls.OfferedPsks( identities=[ ( @@ -808,47 +775,39 @@ def test_pull_client_hello_with_psk(self): "d7aaaf65a9b713872f2bb28818ca1a6b01" ) ], - ), - ) + ) - self.assertTrue(buf.eof()) + assert buf.eof() # serialize buf = Buffer(1000) push_client_hello(buf, hello) - self.assertEqual(buf.data, load("tls_client_hello_with_psk.bin")) + assert buf.data == load("tls_client_hello_with_psk.bin") def test_pull_client_hello_with_sni(self): buf = Buffer(data=load("tls_client_hello_with_sni.bin")) hello = pull_client_hello(buf) - self.assertTrue(buf.eof()) + assert buf.eof() - self.assertEqual( - hello.random, + assert hello.random == \ binascii.unhexlify( "987d8934140b0a42cc5545071f3f9f7f61963d7b6404eb674c8dbe513604346b" - ), - ) - self.assertEqual( - hello.legacy_session_id, + ) + assert hello.legacy_session_id == \ binascii.unhexlify( "26b19bdd30dbf751015a3a16e13bd59002dfe420b799d2a5cd5e11b8fa7bcb66" - ), - ) - self.assertEqual( - hello.cipher_suites, + ) + assert hello.cipher_suites == \ [ tls.CipherSuite.AES_256_GCM_SHA384, tls.CipherSuite.AES_128_GCM_SHA256, tls.CipherSuite.CHACHA20_POLY1305_SHA256, - ], - ) - self.assertEqual(hello.legacy_compression_methods, [tls.CompressionMethod.NULL]) + ] + assert hello.legacy_compression_methods == [tls.CompressionMethod.NULL] # extensions - self.assertEqual(hello.alpn_protocols, None) - self.assertEqual( - hello.key_share, + assert hello.alpn_protocols == None + assert hello.key_share == \ [ ( tls.Group.SECP256R1, @@ -858,39 +817,30 @@ def test_pull_client_hello_with_sni(self): "40" ), ) - ], - ) - self.assertEqual( - hello.psk_key_exchange_modes, [tls.PskKeyExchangeMode.PSK_DHE_KE] - ) - self.assertEqual(hello.server_name, "cloudflare-quic.com") - self.assertEqual( - hello.signature_algorithms, + ] + assert hello.psk_key_exchange_modes == [tls.PskKeyExchangeMode.PSK_DHE_KE] + assert hello.server_name == "cloudflare-quic.com" + assert hello.signature_algorithms == \ [ tls.SignatureAlgorithm.RSA_PSS_RSAE_SHA256, tls.SignatureAlgorithm.ECDSA_SECP256R1_SHA256, tls.SignatureAlgorithm.RSA_PKCS1_SHA256, tls.SignatureAlgorithm.RSA_PKCS1_SHA1, - ], - ) - self.assertEqual(hello.supported_groups, [tls.Group.SECP256R1]) - self.assertEqual( - hello.supported_versions, - [tls.TLS_VERSION_1_3, 32540, 32539, 32538], # old removed draft support - ) + ] + assert hello.supported_groups == [tls.Group.SECP256R1] + # old removed draft support + assert hello.supported_versions ==[tls.TLS_VERSION_1_3, 32540, 32539, 32538] - self.assertEqual(len(hello.other_extensions), 1) + assert len(hello.other_extensions) == 1 - self.assertEqual( - hello.other_extensions[0][0], - 65445, - ) + assert hello.other_extensions[0][0] == \ + 65445 # serialize buf = Buffer(1000) push_client_hello(buf, hello) - self.assertEqual(buf.data, load("tls_client_hello_with_sni.bin")) + assert buf.data == load("tls_client_hello_with_sni.bin") def test_push_client_hello(self): hello = ClientHello( @@ -937,29 +887,24 @@ def test_push_client_hello(self): buf = Buffer(1000) push_client_hello(buf, hello) - self.assertEqual(buf.data, load("tls_client_hello.bin")) + assert buf.data == load("tls_client_hello.bin") def test_pull_server_hello(self): buf = Buffer(data=load("tls_server_hello.bin")) hello = pull_server_hello(buf) - self.assertTrue(buf.eof()) + assert buf.eof() - self.assertEqual( - hello.random, + assert hello.random == \ binascii.unhexlify( "ada85271d19680c615ea7336519e3fdf6f1e26f3b1075ee1de96ffa8884e8280" - ), - ) - self.assertEqual( - hello.legacy_session_id, + ) + assert hello.legacy_session_id == \ binascii.unhexlify( "9aee82a2d186c1cb32a329d9dcfe004a1a438ad0485a53c6bfcf55c132a23235" - ), - ) - self.assertEqual(hello.cipher_suite, tls.CipherSuite.AES_256_GCM_SHA384) - self.assertEqual(hello.compression_method, tls.CompressionMethod.NULL) - self.assertEqual( - hello.key_share, + ) + assert hello.cipher_suite == tls.CipherSuite.AES_256_GCM_SHA384 + assert hello.compression_method == tls.CompressionMethod.NULL + assert hello.key_share == \ ( tls.Group.SECP256R1, binascii.unhexlify( @@ -967,32 +912,26 @@ def test_pull_server_hello(self): "5cca1c503bf0378ac6937c354912116ff3251026bca1958d7f387316c83ae6cf" "b2" ), - ), - ) - self.assertEqual(hello.pre_shared_key, None) - self.assertEqual(hello.supported_version, tls.TLS_VERSION_1_3) + ) + assert hello.pre_shared_key == None + assert hello.supported_version == tls.TLS_VERSION_1_3 def test_pull_server_hello_with_psk(self): buf = Buffer(data=load("tls_server_hello_with_psk.bin")) hello = pull_server_hello(buf) - self.assertTrue(buf.eof()) + assert buf.eof() - self.assertEqual( - hello.random, + assert hello.random == \ binascii.unhexlify( "ccbaaf04fc1bd5143b2cc6b97520cf37d91470dbfc8127131a7bf0f941e3a137" - ), - ) - self.assertEqual( - hello.legacy_session_id, + ) + assert hello.legacy_session_id == \ binascii.unhexlify( "9483e7e895d0f4cec17086b0849601c0632662cd764e828f2f892f4c4b7771b0" - ), - ) - self.assertEqual(hello.cipher_suite, tls.CipherSuite.AES_256_GCM_SHA384) - self.assertEqual(hello.compression_method, tls.CompressionMethod.NULL) - self.assertEqual( - hello.key_share, + ) + assert hello.cipher_suite == tls.CipherSuite.AES_256_GCM_SHA384 + assert hello.compression_method == tls.CompressionMethod.NULL + assert hello.key_share == \ ( tls.Group.SECP256R1, binascii.unhexlify( @@ -1000,23 +939,21 @@ def test_pull_server_hello_with_psk(self): "18b6409593b15c6649d6f459387a53128b164178adc840179aad01d36ce95d62" "76" ), - ), - ) - self.assertEqual(hello.pre_shared_key, 0) - self.assertEqual(hello.supported_version, tls.TLS_VERSION_1_3) + ) + assert hello.pre_shared_key == 0 + assert hello.supported_version == tls.TLS_VERSION_1_3 # serialize buf = Buffer(1000) push_server_hello(buf, hello) - self.assertEqual(buf.data, load("tls_server_hello_with_psk.bin")) + assert buf.data == load("tls_server_hello_with_psk.bin") def test_pull_server_hello_with_unknown_extension(self): buf = Buffer(data=load("tls_server_hello_with_unknown_extension.bin")) hello = pull_server_hello(buf) - self.assertTrue(buf.eof()) + assert buf.eof() - self.assertEqual( - hello, + assert hello == \ ServerHello( random=binascii.unhexlify( "ada85271d19680c615ea7336519e3fdf6f1e26f3b1075ee1de96ffa8884e8280" @@ -1036,13 +973,12 @@ def test_pull_server_hello_with_unknown_extension(self): ), supported_version=tls.TLS_VERSION_1_3, other_extensions=[(12345, b"foo")], - ), - ) + ) # serialize buf = Buffer(1000) push_server_hello(buf, hello) - self.assertEqual(buf.data, load("tls_server_hello_with_unknown_extension.bin")) + assert buf.data == load("tls_server_hello_with_unknown_extension.bin") def test_push_server_hello(self): hello = ServerHello( @@ -1067,16 +1003,15 @@ def test_push_server_hello(self): buf = Buffer(1000) push_server_hello(buf, hello) - self.assertEqual(buf.data, load("tls_server_hello.bin")) + assert buf.data == load("tls_server_hello.bin") def test_pull_new_session_ticket(self): buf = Buffer(data=load("tls_new_session_ticket.bin")) new_session_ticket = pull_new_session_ticket(buf) - self.assertIsNotNone(new_session_ticket) - self.assertTrue(buf.eof()) + assert new_session_ticket is not None + assert buf.eof() - self.assertEqual( - new_session_ticket, + assert new_session_ticket == \ NewSessionTicket( ticket_lifetime=86400, ticket_age_add=3303452425, @@ -1085,22 +1020,20 @@ def test_pull_new_session_ticket(self): "dbe6f1a77a78c0426bfa607cd0d02b350247d90618704709596beda7e962cc81" ), max_early_data_size=0xFFFFFFFF, - ), - ) + ) # serialize buf = Buffer(100) push_new_session_ticket(buf, new_session_ticket) - self.assertEqual(buf.data, load("tls_new_session_ticket.bin")) + assert buf.data == load("tls_new_session_ticket.bin") def test_pull_new_session_ticket_with_unknown_extension(self): buf = Buffer(data=load("tls_new_session_ticket_with_unknown_extension.bin")) new_session_ticket = pull_new_session_ticket(buf) - self.assertIsNotNone(new_session_ticket) - self.assertTrue(buf.eof()) + assert new_session_ticket is not None + assert buf.eof() - self.assertEqual( - new_session_ticket, + assert new_session_ticket == \ NewSessionTicket( ticket_lifetime=86400, ticket_age_add=3303452425, @@ -1110,25 +1043,21 @@ def test_pull_new_session_ticket_with_unknown_extension(self): ), max_early_data_size=0xFFFFFFFF, other_extensions=[(12345, b"foo")], - ), - ) + ) # serialize buf = Buffer(100) push_new_session_ticket(buf, new_session_ticket) - self.assertEqual( - buf.data, load("tls_new_session_ticket_with_unknown_extension.bin") - ) + assert buf.data == load("tls_new_session_ticket_with_unknown_extension.bin") def test_encrypted_extensions(self): data = load("tls_encrypted_extensions.bin") buf = Buffer(data=data) extensions = pull_encrypted_extensions(buf) - self.assertIsNotNone(extensions) - self.assertTrue(buf.eof()) + assert extensions is not None + assert buf.eof() - self.assertEqual( - extensions, + assert extensions == \ EncryptedExtensions( other_extensions=[ ( @@ -1136,23 +1065,21 @@ def test_encrypted_extensions(self): SERVER_QUIC_TRANSPORT_PARAMETERS, ) ] - ), - ) + ) # serialize buf = Buffer(capacity=100) push_encrypted_extensions(buf, extensions) - self.assertEqual(buf.data, data) + assert buf.data == data def test_encrypted_extensions_with_alpn(self): data = load("tls_encrypted_extensions_with_alpn.bin") buf = Buffer(data=data) extensions = pull_encrypted_extensions(buf) - self.assertIsNotNone(extensions) - self.assertTrue(buf.eof()) + assert extensions is not None + assert buf.eof() - self.assertEqual( - extensions, + assert extensions == \ EncryptedExtensions( alpn_protocol="hq-20", other_extensions=[ @@ -1162,22 +1089,20 @@ def test_encrypted_extensions_with_alpn(self): SERVER_QUIC_TRANSPORT_PARAMETERS_2, ), ], - ), - ) + ) # serialize buf = Buffer(115) push_encrypted_extensions(buf, extensions) - self.assertTrue(buf.eof()) + assert buf.eof() def test_pull_encrypted_extensions_with_alpn_and_early_data(self): buf = Buffer(data=load("tls_encrypted_extensions_with_alpn_and_early_data.bin")) extensions = pull_encrypted_extensions(buf) - self.assertIsNotNone(extensions) - self.assertTrue(buf.eof()) + assert extensions is not None + assert buf.eof() - self.assertEqual( - extensions, + assert extensions == \ EncryptedExtensions( alpn_protocol="hq-20", early_data=True, @@ -1188,33 +1113,30 @@ def test_pull_encrypted_extensions_with_alpn_and_early_data(self): SERVER_QUIC_TRANSPORT_PARAMETERS_3, ), ], - ), - ) + ) # serialize buf = Buffer(116) push_encrypted_extensions(buf, extensions) - self.assertTrue(buf.eof()) + assert buf.eof() def test_pull_certificate(self): buf = Buffer(data=load("tls_certificate.bin")) certificate = pull_certificate(buf) - self.assertTrue(buf.eof()) + assert buf.eof() - self.assertEqual(certificate.request_context, b"") - self.assertEqual(certificate.certificates, [(CERTIFICATE_DATA, b"")]) + assert certificate.request_context == b"" + assert certificate.certificates == [(CERTIFICATE_DATA, b"")] def test_pull_certificate_request(self): buf = Buffer(data=load("tls_certificate_request.bin")) certificate_request = pull_certificate_request(buf) - self.assertTrue(buf.eof()) + assert buf.eof() - self.assertEqual(certificate_request.request_context, b"") - self.assertEqual( - certificate_request.signature_algorithms, - [1027, 2052, 1025, 1283, 515, 2053, 2053, 1281, 2054, 1537, 513], - ) - self.assertEqual(certificate_request.other_extensions, []) + assert certificate_request.request_context == b"" + assert certificate_request.signature_algorithms == \ + [1027, 2052, 1025, 1283, 515, 2053, 2053, 1281, 2054, 1537, 513] + assert certificate_request.other_extensions == [] def test_push_certificate(self): certificate = Certificate( @@ -1223,15 +1145,15 @@ def test_push_certificate(self): buf = Buffer(1600) push_certificate(buf, certificate) - self.assertEqual(buf.data, load("tls_certificate.bin")) + assert buf.data == load("tls_certificate.bin") def test_pull_certificate_verify(self): buf = Buffer(data=load("tls_certificate_verify.bin")) verify = pull_certificate_verify(buf) - self.assertTrue(buf.eof()) + assert buf.eof() - self.assertEqual(verify.algorithm, tls.SignatureAlgorithm.RSA_PSS_RSAE_SHA256) - self.assertEqual(verify.signature, CERTIFICATE_VERIFY_SIGNATURE) + assert verify.algorithm == tls.SignatureAlgorithm.RSA_PSS_RSAE_SHA256 + assert verify.signature == CERTIFICATE_VERIFY_SIGNATURE def test_push_certificate_verify(self): verify = CertificateVerify( @@ -1241,19 +1163,17 @@ def test_push_certificate_verify(self): buf = Buffer(400) push_certificate_verify(buf, verify) - self.assertEqual(buf.data, load("tls_certificate_verify.bin")) + assert buf.data == load("tls_certificate_verify.bin") def test_pull_finished(self): buf = Buffer(data=load("tls_finished.bin")) finished = pull_finished(buf) - self.assertTrue(buf.eof()) + assert buf.eof() - self.assertEqual( - finished.verify_data, + assert finished.verify_data == \ binascii.unhexlify( "f157923234ff9a4921aadb2e0ec7b1a30fce73fb9ec0c4276f9af268f408ec68" - ), - ) + ) def test_push_finished(self): finished = Finished( @@ -1264,4 +1184,4 @@ def test_push_finished(self): buf = Buffer(128) push_finished(buf, finished) - self.assertEqual(buf.data, load("tls_finished.bin")) + assert buf.data == load("tls_finished.bin") diff --git a/tests/test_webtransport.py b/tests/test_webtransport.py index 79965ac85..614758eb7 100644 --- a/tests/test_webtransport.py +++ b/tests/test_webtransport.py @@ -1,7 +1,5 @@ from __future__ import annotations -from unittest import TestCase - from qh3.h3.connection import H3_ALPN, ErrorCode, H3Connection from qh3.h3.events import ( DatagramReceived, @@ -24,7 +22,7 @@ } -class WebTransportTest(TestCase): +class TestWebTransport: def _make_session(self, h3_client, h3_server): quic_client = h3_client._quic quic_server = h3_server._quic @@ -44,8 +42,7 @@ def _make_session(self, h3_client, h3_server): # receive request events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -58,9 +55,8 @@ def _make_session(self, h3_client, h3_server): stream_id=stream_id, stream_ended=False, push_id=None, - ) - ], - ) + ) \ + ] # send response h3_server.send_headers( @@ -72,8 +68,7 @@ def _make_session(self, h3_client, h3_server): # receive response events = h3_transfer(quic_server, h3_client) - self.assertEqual( - events, + assert events == \ [ HeadersReceived( headers=[ @@ -82,8 +77,7 @@ def _make_session(self, h3_client, h3_server): stream_id=stream_id, stream_ended=False, ), - ], - ) + ] return stream_id @@ -104,17 +98,15 @@ def test_bidirectional_stream(self): # receive data events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ WebTransportStreamDataReceived( data=b"foo", session_id=session_id, stream_ended=True, stream_id=stream_id, - ) - ], - ) + ) \ + ] def test_bidirectional_stream_fragmented_frame(self): with h3_fake_client_and_server(QUIC_CONFIGURATION_OPTIONS) as ( @@ -133,8 +125,7 @@ def test_bidirectional_stream_fragmented_frame(self): # receive data events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ WebTransportStreamDataReceived( data=b"f", @@ -160,8 +151,7 @@ def test_bidirectional_stream_fragmented_frame(self): stream_ended=True, stream_id=stream_id, ), - ], - ) + ] def test_bidirectional_stream_server_initiated(self): with h3_client_and_server(QUIC_CONFIGURATION_OPTIONS) as ( @@ -180,17 +170,15 @@ def test_bidirectional_stream_server_initiated(self): # receive data events = h3_transfer(quic_server, h3_client) - self.assertEqual( - events, + assert events == \ [ WebTransportStreamDataReceived( data=b"foo", session_id=session_id, stream_ended=True, stream_id=stream_id, - ) - ], - ) + ) \ + ] def test_unidirectional_stream(self): with h3_client_and_server(QUIC_CONFIGURATION_OPTIONS) as ( @@ -211,17 +199,15 @@ def test_unidirectional_stream(self): # receive data events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ WebTransportStreamDataReceived( data=b"foo", session_id=session_id, stream_ended=True, stream_id=stream_id, - ) - ], - ) + ) \ + ] def test_unidirectional_stream_fragmented_frame(self): with h3_fake_client_and_server(QUIC_CONFIGURATION_OPTIONS) as ( @@ -242,8 +228,7 @@ def test_unidirectional_stream_fragmented_frame(self): # receive data events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, + assert events == \ [ WebTransportStreamDataReceived( data=b"f", @@ -269,8 +254,7 @@ def test_unidirectional_stream_fragmented_frame(self): stream_ended=True, stream_id=stream_id, ), - ], - ) + ] def test_datagram(self): with h3_client_and_server(QUIC_CONFIGURATION_OPTIONS) as ( @@ -288,10 +272,8 @@ def test_datagram(self): # receive datagram events = h3_transfer(quic_client, h3_server) - self.assertEqual( - events, - [DatagramReceived(data=b"foo", flow_id=session_id)], - ) + assert events == \ + [DatagramReceived(data=b"foo", flow_id=session_id)] def test_handle_datagram_truncated(self): quic_server = FakeQuicConnection( @@ -301,10 +283,8 @@ def test_handle_datagram_truncated(self): # receive a datagram with a truncated session ID h3_server.handle_event(DatagramFrameReceived(data=b"\xff")) - self.assertEqual( - quic_server.closed, + assert quic_server.closed == \ ( ErrorCode.H3_GENERAL_PROTOCOL_ERROR, "Could not parse flow ID", - ), - ) + ) From ac0c075bbed265e9c1deaec6a5d125c12e06bc65 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Sun, 29 Dec 2024 16:05:02 +0100 Subject: [PATCH 02/39] :wrench: update CI to use noxfile for testing session --- .github/workflows/CI.yml | 54 +++++----------------------------------- 1 file changed, 6 insertions(+), 48 deletions(-) diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index 4d1872ff9..0d2ec02af 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -1,8 +1,3 @@ -# This file is autogenerated by maturin v1.2.3 -# To update, run -# -# maturin generate-ci github -# name: CI on: @@ -51,12 +46,12 @@ jobs: runs-on: ${{ matrix.os }} steps: - uses: actions/checkout@3df4ab11eba7bda6032a0b82a6bb43b11571feac - - uses: actions/setup-python@v4 + - uses: actions/setup-python@v5 with: python-version: ${{ matrix.python_version }} allow-prereleases: true - - name: Setup dependencies - run: pip install --upgrade pip pytest + - name: Setup nox + run: pip install nox - name: Set up Clang (Linux) if: matrix.os == 'ubuntu-22.04' run: sudo apt-get install clang @@ -65,45 +60,8 @@ jobs: run: choco install llvm -y - uses: ilammy/setup-nasm@v1 if: matrix.os == 'windows-latest' - - name: Build wheels (Unix, Linux) - if: matrix.os != 'windows-latest' - uses: PyO3/maturin-action@v1 - with: - args: --release --out dist --interpreter ${{ matrix.python_version }} - sccache: 'true' - manylinux: auto - before-script-linux: | - sudo apt-get update || echo "no apt support" - sudo apt-get upgrade -y || echo "no apt support" - sudo apt-get install -y libclang libclang-dev || echo "no apt support" - sudo apt-get install -y linux-headers-generic || echo "no apt support" - sudo apt-get install -y libc6-dev || echo "no apt support" - yum install -y llvm-toolset-7-clang || echo "not yum based" - source /opt/rh/llvm-toolset-7/enable || echo "not yum based" - - name: Build wheels (NT) - if: matrix.os == 'windows-latest' - uses: PyO3/maturin-action@v1.42.1 - with: - args: --release --out dist - sccache: 'true' - target: x64 - - run: pip install --find-links=./dist qh3 - name: Install built package - - name: Disable firewall and configure compiler - if: matrix.os == 'macos-latest' - run: | - sudo /usr/libexec/ApplicationFirewall/socketfilterfw --setglobalstate off - echo "AIOQUIC_SKIP_TESTS=chacha20" >> $GITHUB_ENV - - name: Ensure test target (NT) - if: matrix.os == 'windows-latest' - run: Remove-Item -Path qh3 -Force -Recurse - - name: Ensure test target (Linux, Unix) - if: matrix.os != 'windows-latest' - run: rm -fR qh3 - - run: python -m pip install -r dev-requirements.txt - name: Install dev requirements - - run: python -m unittest discover -v - name: Run tests + - name: Run test + run: nox -s test-${{ matrix.python_version }} linux: runs-on: ubuntu-22.04 @@ -411,7 +369,7 @@ jobs: provenance: needs: checksum if: "startsWith(github.ref, 'refs/tags/')" - uses: "slsa-framework/slsa-github-generator/.github/workflows/generator_generic_slsa3.yml@v1.10.0" + uses: "slsa-framework/slsa-github-generator/.github/workflows/generator_generic_slsa3.yml@v2.0.0" permissions: actions: read id-token: write From 326a661821749648552878bba35cdc8e388fcb62 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Sun, 29 Dec 2024 16:06:09 +0100 Subject: [PATCH 03/39] :bug: Fix error on teardown for our stream transport adapter (asyncio.Transport mandate impl of close and is_closing) --- qh3/asyncio/protocol.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/qh3/asyncio/protocol.py b/qh3/asyncio/protocol.py index b633410e0..01316dd30 100644 --- a/qh3/asyncio/protocol.py +++ b/qh3/asyncio/protocol.py @@ -220,6 +220,8 @@ def _transmit_soon(self) -> None: class QuicStreamAdapter(asyncio.Transport): def __init__(self, protocol: QuicConnectionProtocol, stream_id: int): + super().__init__() + self.protocol = protocol self.stream_id = stream_id @@ -240,3 +242,9 @@ def write(self, data): def write_eof(self): self.protocol._quic.send_stream_data(self.stream_id, b"", end_stream=True) self.protocol._transmit_soon() + + def is_closing(self): + return self.protocol._quic._close_pending or self.stream_id in self.protocol._quic._streams_finished + + def close(self): + pass From 68c84187f5c3e50a1d1a80fea1c116e8bc56d561 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Sun, 29 Dec 2024 16:08:00 +0100 Subject: [PATCH 04/39] :art: reformat protocol.py --- qh3/asyncio/protocol.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/qh3/asyncio/protocol.py b/qh3/asyncio/protocol.py index 01316dd30..70e0cc338 100644 --- a/qh3/asyncio/protocol.py +++ b/qh3/asyncio/protocol.py @@ -244,7 +244,10 @@ def write_eof(self): self.protocol._transmit_soon() def is_closing(self): - return self.protocol._quic._close_pending or self.stream_id in self.protocol._quic._streams_finished + return ( + self.protocol._quic._close_pending + or self.stream_id in self.protocol._quic._streams_finished + ) def close(self): pass From 43953084f5c20231c1d32a75a573d80302d426c5 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Sun, 29 Dec 2024 16:08:46 +0100 Subject: [PATCH 05/39] :art: initial clippy fixes --- .pre-commit-config.yaml | 20 ++- src/aead.rs | 194 +++++++++++++++++------------ src/agreement.rs | 122 +++++++----------- src/buffer.rs | 102 ++++++++-------- src/certificate.rs | 262 +++++++++++++++++---------------------- src/headers.rs | 184 ++++++++++++++-------------- src/hpk.rs | 31 +++-- src/lib.rs | 64 +++++++--- src/ocsp.rs | 113 ++++++++--------- src/pkcs8.rs | 75 ++++++------ src/private_key.rs | 265 ++++++++++++++++++++-------------------- src/rsa.rs | 30 +++-- 12 files changed, 737 insertions(+), 725 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index da97c7fb4..ee13c0a04 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,4 +1,4 @@ -exclude: 'docs/|src/|tests/' +exclude: 'docs/|tests/' repos: - repo: https://github.com/pre-commit/pre-commit-hooks rev: v4.4.0 @@ -26,4 +26,20 @@ repos: hooks: - id: mypy args: [--check-untyped-defs] - exclude: 'tests/|examples/|docs/' + exclude: 'tests/|examples/|docs/|noxfile.py' +- repo: local + hooks: + - id: rust-linting + name: Rust linting + description: Run cargo fmt on files included in the commit. rustfmt should be installed before-hand. + entry: cargo fmt --all -- + pass_filenames: true + types: [file, rust] + language: system + - id: rust-clippy + name: Rust clippy + description: Run cargo clippy on files included in the commit. clippy should be installed before-hand. + entry: cargo clippy --all-targets --all-features -- -Dclippy::all + pass_filenames: false + types: [file, rust] + language: system diff --git a/src/aead.rs b/src/aead.rs index a0e3a2704..54faf2912 100644 --- a/src/aead.rs +++ b/src/aead.rs @@ -1,176 +1,212 @@ -use aws_lc_rs::aead::{Aad, TlsRecordOpeningKey, TlsRecordSealingKey, AES_128_GCM, AES_256_GCM, CHACHA20_POLY1305, TlsProtocolId, Nonce}; - -use chacha20poly1305::{ - aead::{KeyInit}, - ChaCha20Poly1305, Key as ChaCha20Key, AeadInPlace +use aws_lc_rs::aead::{ + Aad, Nonce, TlsProtocolId, TlsRecordOpeningKey, TlsRecordSealingKey, AES_128_GCM, AES_256_GCM, + CHACHA20_POLY1305, }; -use pyo3::{PyResult, Python}; -use pyo3::types::PyBytes; -use pyo3::pymethods; +use chacha20poly1305::{aead::KeyInit, AeadInPlace, ChaCha20Poly1305, Key as ChaCha20Key}; + use pyo3::pyclass; +use pyo3::pymethods; +use pyo3::types::PyBytes; +use pyo3::{PyResult, Python}; use crate::CryptoError; #[pyclass(name = "AeadChaCha20Poly1305", module = "qh3._hazmat")] pub struct AeadChaCha20Poly1305 { - key: Vec + key: Vec, } #[pyclass(name = "AeadAes256Gcm", module = "qh3._hazmat")] pub struct AeadAes256Gcm { - key: Vec + key: Vec, } #[pyclass(name = "AeadAes128Gcm", module = "qh3._hazmat")] pub struct AeadAes128Gcm { - key: Vec + key: Vec, } - #[pymethods] impl AeadAes256Gcm { #[new] pub fn py_new(key: &PyBytes) -> Self { - AeadAes256Gcm { key: key.as_bytes().to_vec() } + AeadAes256Gcm { + key: key.as_bytes().to_vec(), + } } - pub fn decrypt<'a>(&mut self, py: Python<'a>, nonce: &PyBytes, data: &PyBytes, associated_data: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn decrypt<'a>( + &mut self, + py: Python<'a>, + nonce: &PyBytes, + data: &PyBytes, + associated_data: &PyBytes, + ) -> PyResult<&'a PyBytes> { let mut in_out_buffer = data.as_bytes().to_vec(); let plaintext_len = in_out_buffer.len() - AES_256_GCM.tag_len(); - let opening_key: TlsRecordOpeningKey = TlsRecordOpeningKey::new(&AES_256_GCM, TlsProtocolId::TLS13, &self.key).expect("FAILURE"); + let opening_key: TlsRecordOpeningKey = + TlsRecordOpeningKey::new(&AES_256_GCM, TlsProtocolId::TLS13, &self.key) + .expect("FAILURE"); let aad = Aad::from(associated_data.as_bytes()); - let res = opening_key - .open_in_place(Nonce::try_assume_unique_for_key(nonce.as_bytes()).unwrap(), aad, &mut in_out_buffer); + let res = opening_key.open_in_place( + Nonce::try_assume_unique_for_key(nonce.as_bytes()).unwrap(), + aad, + &mut in_out_buffer, + ); return match res { - Ok(_) => Ok( - PyBytes::new( - py, - &in_out_buffer[0..plaintext_len] - ) - ), - Err(_) => Err(CryptoError::new_err("decryption failed")) + Ok(_) => Ok(PyBytes::new(py, &in_out_buffer[0..plaintext_len])), + Err(_) => Err(CryptoError::new_err("decryption failed")), }; } - pub fn encrypt<'a>(&mut self, py: Python<'a>, nonce: &PyBytes, data: &PyBytes, associated_data: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn encrypt<'a>( + &mut self, + py: Python<'a>, + nonce: &PyBytes, + data: &PyBytes, + associated_data: &PyBytes, + ) -> PyResult<&'a PyBytes> { let mut in_out_buffer = Vec::from(data.as_bytes()); - let mut sealing_key: TlsRecordSealingKey = TlsRecordSealingKey::new(&AES_256_GCM, TlsProtocolId::TLS13, &self.key).expect("FAILURE"); + let mut sealing_key: TlsRecordSealingKey = + TlsRecordSealingKey::new(&AES_256_GCM, TlsProtocolId::TLS13, &self.key) + .expect("FAILURE"); let aad = Aad::from(associated_data.as_bytes()); - let res = sealing_key - .seal_in_place_append_tag(Nonce::try_assume_unique_for_key(nonce.as_bytes()).unwrap(), aad, &mut in_out_buffer); + let res = sealing_key.seal_in_place_append_tag( + Nonce::try_assume_unique_for_key(nonce.as_bytes()).unwrap(), + aad, + &mut in_out_buffer, + ); return match res { - Ok(_) => Ok( - PyBytes::new( - py, - &in_out_buffer - ) - ), - Err(_) => Err(CryptoError::new_err("encryption failed")) + Ok(_) => Ok(PyBytes::new(py, &in_out_buffer)), + Err(_) => Err(CryptoError::new_err("encryption failed")), }; } } - #[pymethods] impl AeadAes128Gcm { #[new] pub fn py_new(key: &PyBytes) -> Self { - AeadAes128Gcm { key: key.as_bytes().to_vec() } + AeadAes128Gcm { + key: key.as_bytes().to_vec(), + } } - pub fn decrypt<'a>(&mut self, py: Python<'a>, nonce: &PyBytes, data: &PyBytes, associated_data: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn decrypt<'a>( + &mut self, + py: Python<'a>, + nonce: &PyBytes, + data: &PyBytes, + associated_data: &PyBytes, + ) -> PyResult<&'a PyBytes> { let mut in_out_buffer = data.as_bytes().to_vec(); let plaintext_len = in_out_buffer.len() - AES_128_GCM.tag_len(); - let opening_key = TlsRecordOpeningKey::new(&AES_128_GCM, TlsProtocolId::TLS13, &self.key).expect("FAILURE"); + let opening_key = TlsRecordOpeningKey::new(&AES_128_GCM, TlsProtocolId::TLS13, &self.key) + .expect("FAILURE"); let aad = Aad::from(associated_data.as_bytes()); - let res = opening_key - .open_in_place(Nonce::try_assume_unique_for_key(nonce.as_bytes()).unwrap(), aad, &mut in_out_buffer); + let res = opening_key.open_in_place( + Nonce::try_assume_unique_for_key(nonce.as_bytes()).unwrap(), + aad, + &mut in_out_buffer, + ); return match res { - Ok(_) => Ok( - PyBytes::new( - py, - &in_out_buffer[0..plaintext_len] - ) - ), - Err(_) => Err(CryptoError::new_err("decryption failed")) + Ok(_) => Ok(PyBytes::new(py, &in_out_buffer[0..plaintext_len])), + Err(_) => Err(CryptoError::new_err("decryption failed")), }; } - pub fn encrypt<'a>(&mut self, py: Python<'a>, nonce: &PyBytes, data: &PyBytes, associated_data: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn encrypt<'a>( + &mut self, + py: Python<'a>, + nonce: &PyBytes, + data: &PyBytes, + associated_data: &PyBytes, + ) -> PyResult<&'a PyBytes> { let mut in_out_buffer = Vec::from(data.as_bytes()); - let mut sealing_key = TlsRecordSealingKey::new(&AES_128_GCM, TlsProtocolId::TLS13, &self.key).expect("FAILURE"); + let mut sealing_key = + TlsRecordSealingKey::new(&AES_128_GCM, TlsProtocolId::TLS13, &self.key) + .expect("FAILURE"); let aad = Aad::from(associated_data.as_bytes()); - let res = sealing_key - .seal_in_place_append_tag(Nonce::try_assume_unique_for_key(nonce.as_bytes()).unwrap(), aad, &mut in_out_buffer); + let res = sealing_key.seal_in_place_append_tag( + Nonce::try_assume_unique_for_key(nonce.as_bytes()).unwrap(), + aad, + &mut in_out_buffer, + ); return match res { - Ok(_) => Ok( - PyBytes::new( - py, - &in_out_buffer - ) - ), - Err(_) => Err(CryptoError::new_err("encryption failed")) + Ok(_) => Ok(PyBytes::new(py, &in_out_buffer)), + Err(_) => Err(CryptoError::new_err("encryption failed")), }; } } - #[pymethods] impl AeadChaCha20Poly1305 { #[new] pub fn py_new(key: &PyBytes) -> Self { - AeadChaCha20Poly1305 { key: key.as_bytes().to_vec() } + AeadChaCha20Poly1305 { + key: key.as_bytes().to_vec(), + } } - pub fn decrypt<'a>(&mut self, py: Python<'a>, nonce: &PyBytes, data: &PyBytes, associated_data: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn decrypt<'a>( + &mut self, + py: Python<'a>, + nonce: &PyBytes, + data: &PyBytes, + associated_data: &PyBytes, + ) -> PyResult<&'a PyBytes> { let mut in_out_buffer = data.as_bytes().to_vec(); let plaintext_len = in_out_buffer.len() - CHACHA20_POLY1305.tag_len(); let cipher: ChaCha20Poly1305 = ChaCha20Poly1305::new(ChaCha20Key::from_slice(&self.key)); - let res = cipher.decrypt_in_place(nonce.as_bytes().into(), &associated_data.as_bytes(), &mut in_out_buffer); + let res = cipher.decrypt_in_place( + nonce.as_bytes().into(), + &associated_data.as_bytes(), + &mut in_out_buffer, + ); return match res { - Ok(_) => Ok( - PyBytes::new( - py, - &in_out_buffer[0..plaintext_len] - ) - ), - Err(_) => Err(CryptoError::new_err("decryption failed")) + Ok(_) => Ok(PyBytes::new(py, &in_out_buffer[0..plaintext_len])), + Err(_) => Err(CryptoError::new_err("decryption failed")), }; } - pub fn encrypt<'a>(&mut self, py: Python<'a>, nonce: &PyBytes, data: &PyBytes, associated_data: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn encrypt<'a>( + &mut self, + py: Python<'a>, + nonce: &PyBytes, + data: &PyBytes, + associated_data: &PyBytes, + ) -> PyResult<&'a PyBytes> { let mut in_out_buffer = Vec::from(data.as_bytes()); let cipher: ChaCha20Poly1305 = ChaCha20Poly1305::new(ChaCha20Key::from_slice(&self.key)); - let res = cipher.encrypt_in_place(nonce.as_bytes().into(), &associated_data.as_bytes(), &mut in_out_buffer); + let res = cipher.encrypt_in_place( + nonce.as_bytes().into(), + &associated_data.as_bytes(), + &mut in_out_buffer, + ); return match res { - Ok(_) => Ok( - PyBytes::new( - py, - &in_out_buffer - ) - ), - Err(_) => Err(CryptoError::new_err("encryption failed")) + Ok(_) => Ok(PyBytes::new(py, &in_out_buffer)), + Err(_) => Err(CryptoError::new_err("encryption failed")), }; } } diff --git a/src/agreement.rs b/src/agreement.rs index 92d1ac6ba..e70165112 100644 --- a/src/agreement.rs +++ b/src/agreement.rs @@ -3,14 +3,12 @@ use aws_lc_rs::{agreement, error}; use aws_lc_rs::kem; use aws_lc_rs::unstable::kem::{get_algorithm, AlgorithmId}; -use rustls::crypto::{ - SharedSecret, -}; +use rustls::crypto::SharedSecret; -use pyo3::Python; -use pyo3::types::PyBytes; -use pyo3::pymethods; use pyo3::pyclass; +use pyo3::pymethods; +use pyo3::types::PyBytes; +use pyo3::Python; const X25519_LEN: usize = 32; const KYBER_CIPHERTEXT_LEN: usize = 1088; @@ -34,13 +32,11 @@ pub struct X25519KeyExchange { private: agreement::PrivateKey, } - #[pyclass(module = "qh3._hazmat")] pub struct ECDHP256KeyExchange { private: agreement::PrivateKey, } - #[pyclass(module = "qh3._hazmat")] pub struct ECDHP384KeyExchange { private: agreement::PrivateKey, @@ -63,24 +59,26 @@ impl X25519Kyber768Draft00KeyExchange { pub fn py_new() -> Self { X25519Kyber768Draft00KeyExchange { x25519_private: agreement::PrivateKey::generate(&agreement::X25519).expect("FAILURE"), - kyber768_decapsulation_key: kem::DecapsulationKey::generate(get_algorithm(AlgorithmId::Kyber768_R3).expect("Kyber768_R3 not available")).expect("FAILURE") + kyber768_decapsulation_key: kem::DecapsulationKey::generate( + get_algorithm(AlgorithmId::Kyber768_R3).expect("Kyber768_R3 not available"), + ) + .expect("FAILURE"), } } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - let kyber_pub = self.kyber768_decapsulation_key + let kyber_pub = self + .kyber768_decapsulation_key .encapsulation_key() .expect("FAILURE"); let mut combined_pub_key = Vec::with_capacity(X25519_KYBER_COMBINED_PUBKEY_LEN); - combined_pub_key.extend_from_slice(self.x25519_private.compute_public_key().unwrap().as_ref()); + combined_pub_key + .extend_from_slice(self.x25519_private.compute_public_key().unwrap().as_ref()); combined_pub_key.extend_from_slice(kyber_pub.key_bytes().unwrap().as_ref()); - return PyBytes::new( - py, - &combined_pub_key.as_ref() - ); + return PyBytes::new(py, &combined_pub_key.as_ref()); } pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { @@ -98,12 +96,12 @@ impl X25519Kyber768Draft00KeyExchange { &self.x25519_private, &x25519_peer_public_key, error::Unspecified, - |_key_material| { - return Ok(_key_material.to_vec()) - }, - ).expect("FAILURE"); + |_key_material| return Ok(_key_material.to_vec()), + ) + .expect("FAILURE"); - let kyber_secret = self.kyber768_decapsulation_key + let kyber_secret = self + .kyber768_decapsulation_key .decapsulate(kyber.into()) .expect("FAILURE"); @@ -114,14 +112,10 @@ impl X25519Kyber768Draft00KeyExchange { let key_material = SharedSecret::from(&combined_secret.0[..]); - return PyBytes::new( - py, - &key_material.secret_bytes() - ); + return PyBytes::new(py, &key_material.secret_bytes()); } } - #[pymethods] impl X25519KeyExchange { #[new] @@ -134,52 +128,43 @@ impl X25519KeyExchange { pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { let my_public_key = self.private.compute_public_key().unwrap(); - return PyBytes::new( - py, - &my_public_key.as_ref() - ); + return PyBytes::new(py, &my_public_key.as_ref()); } pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { - let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::X25519, peer_public_key.as_bytes()); + let peer_public_key = + agreement::UnparsedPublicKey::new(&agreement::X25519, peer_public_key.as_bytes()); let key_material = agreement::agree( &self.private, &peer_public_key, error::Unspecified, - |_key_material| { - return Ok(_key_material.to_vec()) - }, - ).expect("FAILURE"); + |_key_material| return Ok(_key_material.to_vec()), + ) + .expect("FAILURE"); - return PyBytes::new( - py, - &key_material - ); + return PyBytes::new(py, &key_material); } } - #[pymethods] impl ECDHP256KeyExchange { #[new] pub fn py_new() -> Self { ECDHP256KeyExchange { - private: agreement::PrivateKey::generate(&agreement::ECDH_P256).expect("FAILURE") + private: agreement::PrivateKey::generate(&agreement::ECDH_P256).expect("FAILURE"), } } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { let my_public_key = self.private.compute_public_key().unwrap(); - return PyBytes::new( - py, - &my_public_key.as_ref() - ); + return PyBytes::new(py, &my_public_key.as_ref()); } pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { - let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::ECDH_P256, peer_public_key.as_bytes()); + let peer_public_key = + agreement::UnparsedPublicKey::new(&agreement::ECDH_P256, peer_public_key.as_bytes()); let key_material = agreement::agree( &self.private, @@ -188,36 +173,31 @@ impl ECDHP256KeyExchange { |_key_material| { return Ok(_key_material.to_vec()); }, - ).expect("FAILURE"); + ) + .expect("FAILURE"); - return PyBytes::new( - py, - &key_material - ); + return PyBytes::new(py, &key_material); } } - #[pymethods] impl ECDHP384KeyExchange { #[new] pub fn py_new() -> Self { ECDHP384KeyExchange { - private: agreement::PrivateKey::generate(&agreement::ECDH_P384).expect("FAILURE") + private: agreement::PrivateKey::generate(&agreement::ECDH_P384).expect("FAILURE"), } } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { let my_public_key = self.private.compute_public_key().unwrap(); - return PyBytes::new( - py, - &my_public_key.as_ref() - ); + return PyBytes::new(py, &my_public_key.as_ref()); } pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { - let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::ECDH_P384, peer_public_key.as_bytes()); + let peer_public_key = + agreement::UnparsedPublicKey::new(&agreement::ECDH_P384, peer_public_key.as_bytes()); let key_material = agreement::agree( &self.private, @@ -226,36 +206,31 @@ impl ECDHP384KeyExchange { |_key_material| { return Ok(_key_material.to_vec()); }, - ).expect("FAILURE"); + ) + .expect("FAILURE"); - return PyBytes::new( - py, - &key_material - ); + return PyBytes::new(py, &key_material); } } - #[pymethods] impl ECDHP521KeyExchange { #[new] pub fn py_new() -> Self { ECDHP521KeyExchange { - private: agreement::PrivateKey::generate(&agreement::ECDH_P521).expect("FAILURE") + private: agreement::PrivateKey::generate(&agreement::ECDH_P521).expect("FAILURE"), } } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { let my_public_key = self.private.compute_public_key().unwrap(); - return PyBytes::new( - py, - &my_public_key.as_ref() - ); + return PyBytes::new(py, &my_public_key.as_ref()); } pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { - let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::ECDH_P521, peer_public_key.as_bytes()); + let peer_public_key = + agreement::UnparsedPublicKey::new(&agreement::ECDH_P521, peer_public_key.as_bytes()); let key_material = agreement::agree( &self.private, @@ -264,12 +239,9 @@ impl ECDHP521KeyExchange { |_key_material| { return Ok(_key_material.to_vec()); }, - ).expect("FAILURE"); + ) + .expect("FAILURE"); - return PyBytes::new( - py, - &key_material - ); + return PyBytes::new(py, &key_material); } } - diff --git a/src/buffer.rs b/src/buffer.rs index 444e6e4c8..349e3c51f 100644 --- a/src/buffer.rs +++ b/src/buffer.rs @@ -1,48 +1,42 @@ -use pyo3::types::{PyBytes}; -use pyo3::{pymethods, PyResult, Python}; +use pyo3::exceptions::PyValueError; use pyo3::pyclass; -use pyo3::exceptions::{PyValueError}; +use pyo3::types::PyBytes; +use pyo3::{pymethods, PyResult, Python}; pyo3::create_exception!(_hazmat, BufferReadError, PyValueError); pyo3::create_exception!(_hazmat, BufferWriteError, PyValueError); - #[pyclass(module = "qh3._hazmat")] pub struct Buffer { pos: u64, data: Vec, - capacity: u64 + capacity: u64, } - #[pymethods] impl Buffer { #[new] pub fn py_new(capacity: Option, data: Option<&PyBytes>) -> PyResult { if data.is_some() { let payload = data.unwrap().as_bytes(); - return Ok( - Buffer { - pos: 0, - data: payload.to_vec(), - capacity: payload.len() as u64 - } - ); + return Ok(Buffer { + pos: 0, + data: payload.to_vec(), + capacity: payload.len() as u64, + }); } if !capacity.is_some() { - return Err( - PyValueError::new_err("mandatory capacity without data args") - ); + return Err(PyValueError::new_err( + "mandatory capacity without data args", + )); } - return Ok( - Buffer { - pos: 0, - data: vec![0; capacity.unwrap().try_into().unwrap()], - capacity: capacity.unwrap(), - } - ); + return Ok(Buffer { + pos: 0, + data: vec![0; capacity.unwrap().try_into().unwrap()], + capacity: capacity.unwrap(), + }); } #[getter] @@ -53,15 +47,9 @@ impl Buffer { #[getter] pub fn data<'a>(&self, py: Python<'a>) -> &'a PyBytes { if self.pos == 0 { - return PyBytes::new( - py, - &[] - ); + return PyBytes::new(py, &[]); } - return PyBytes::new( - py, - &self.data[0 as usize..self.pos as usize] - ); + return PyBytes::new(py, &self.data[0 as usize..self.pos as usize]); } pub fn data_slice<'a>(&self, py: Python<'a>, start: u64, end: u64) -> PyResult<&'a PyBytes> { @@ -69,12 +57,7 @@ impl Buffer { return Err(BufferReadError::new_err("Read out of bounds")); } - return Ok( - PyBytes::new( - py, - &self.data[start as usize..end as usize] - ) - ); + return Ok(PyBytes::new(py, &self.data[start as usize..end as usize])); } pub fn eof(&self) -> bool { @@ -102,7 +85,7 @@ impl Buffer { let extract = PyBytes::new( py, - &self.data[self.pos as usize..(self.pos+length) as usize] + &self.data[self.pos as usize..(self.pos + length) as usize], ); self.pos += length; @@ -130,7 +113,11 @@ impl Buffer { return Err(BufferReadError::new_err("Read out of bounds")); } - let extract = u16::from_be_bytes(self.data[self.pos as usize..(self.pos + 2) as usize].try_into().expect("failure")); + let extract = u16::from_be_bytes( + self.data[self.pos as usize..(self.pos + 2) as usize] + .try_into() + .expect("failure"), + ); self.pos += 2; return Ok(extract); @@ -145,7 +132,11 @@ impl Buffer { return Err(BufferReadError::new_err("Read out of bounds")); } - let extract = u32::from_be_bytes(self.data[self.pos as usize..(self.pos + 4) as usize].try_into().expect("failure")); + let extract = u32::from_be_bytes( + self.data[self.pos as usize..(self.pos + 4) as usize] + .try_into() + .expect("failure"), + ); self.pos += 4; return Ok(extract); @@ -160,7 +151,11 @@ impl Buffer { return Err(BufferReadError::new_err("Read out of bounds")); } - let extract = u64::from_be_bytes(self.data[self.pos as usize..(self.pos + 8) as usize].try_into().expect("failure")); + let extract = u64::from_be_bytes( + self.data[self.pos as usize..(self.pos + 8) as usize] + .try_into() + .expect("failure"), + ); self.pos += 8; return Ok(extract); @@ -183,8 +178,8 @@ impl Buffer { return match self.pull_uint16() { Ok(val) => { return Ok((val & 0x3FFF).into()); - }, - Err(exception) => Err(exception) + } + Err(exception) => Err(exception), }; } @@ -192,16 +187,16 @@ impl Buffer { return match self.pull_uint32() { Ok(val) => { return Ok((val & 0x3FFFFFFF).into()); - }, - Err(exception) => Err(exception) + } + Err(exception) => Err(exception), }; } return match self.pull_uint64() { Ok(val) => { return Ok(val & 0x3FFFFFFFFFFFFFFF); - }, - Err(exception) => Err(exception) + } + Err(exception) => Err(exception), }; } @@ -239,7 +234,8 @@ impl Buffer { return Err(BufferWriteError::new_err("Write out of bounds")); } - self.data[self.pos as usize..(self.pos + 2) as usize].clone_from_slice(&value.to_be_bytes()); + self.data[self.pos as usize..(self.pos + 2) as usize] + .clone_from_slice(&value.to_be_bytes()); self.pos += 2; return Ok(()); @@ -254,7 +250,8 @@ impl Buffer { return Err(BufferWriteError::new_err("Write out of bounds")); } - self.data[self.pos as usize..(self.pos + 4) as usize].clone_from_slice(&value.to_be_bytes()); + self.data[self.pos as usize..(self.pos + 4) as usize] + .clone_from_slice(&value.to_be_bytes()); self.pos += 4; return Ok(()); @@ -269,7 +266,8 @@ impl Buffer { return Err(BufferWriteError::new_err("Write out of bounds")); } - self.data[self.pos as usize..(self.pos + 8) as usize].clone_from_slice(&value.to_be_bytes()); + self.data[self.pos as usize..(self.pos + 8) as usize] + .clone_from_slice(&value.to_be_bytes()); self.pos += 8; return Ok(()); @@ -286,6 +284,8 @@ impl Buffer { return self.push_uint64(value | 0xC000000000000000); } - return Err(PyValueError::new_err("Integer is too big for a variable-length integer")); + return Err(PyValueError::new_err( + "Integer is too big for a variable-length integer", + )); } } diff --git a/src/certificate.rs b/src/certificate.rs index fc3317677..ec52d3f74 100644 --- a/src/certificate.rs +++ b/src/certificate.rs @@ -1,13 +1,13 @@ +use rustls::client::danger::ServerCertVerifier; use rustls::client::WebPkiServerVerifier; -use rustls::client::danger::{ServerCertVerifier}; +use rustls::pki_types::{CertificateDer, ServerName, UnixTime}; use rustls::{CertificateError, Error, RootCertStore}; -use rustls::pki_types::{CertificateDer, UnixTime, ServerName}; -use pyo3::{PyResult, Python}; -use pyo3::types::{PyBytes, PyList, PyTuple}; -use pyo3::pymethods; use pyo3::pyclass; +use pyo3::pymethods; +use pyo3::types::{PyBytes, PyList, PyTuple}; use pyo3::ToPyObject; +use pyo3::{PyResult, Python}; use x509_parser::prelude::*; use x509_parser::public_key::PublicKey; @@ -23,7 +23,6 @@ pyo3::create_exception!(_hazmat, InvalidNameCertificateError, PyException); pyo3::create_exception!(_hazmat, ExpiredCertificateError, PyException); pyo3::create_exception!(_hazmat, UnacceptableCertificateError, PyException); - #[pyclass(name = "Extension", module = "qh3._hazmat", frozen)] pub struct Extension { oid: String, @@ -33,7 +32,7 @@ pub struct Extension { #[pyclass(name = "Subject", module = "qh3._hazmat", frozen)] pub struct Subject { oid: String, - value: Vec + value: Vec, } #[pyclass(name = "Certificate", module = "qh3._hazmat", frozen)] @@ -68,82 +67,66 @@ impl Certificate { match extension.parsed_extension() { ParsedExtension::AuthorityInfoAccess(aia) => { for ext_endpoint in &aia.accessdescs { - extensions.push( - Extension { - oid: ext_endpoint.access_method.to_string(), - value: ext_endpoint.access_location.to_string().into(), - } - ) - + extensions.push(Extension { + oid: ext_endpoint.access_method.to_string(), + value: ext_endpoint.access_location.to_string().into(), + }) } - }, + } ParsedExtension::SubjectAlternativeName(san) => { for name in &san.general_names { - extensions.push( - Extension { - oid: "2.5.29.17".to_string(), - value: name.to_string().into(), - } - ) + extensions.push(Extension { + oid: "2.5.29.17".to_string(), + value: name.to_string().into(), + }) } } - _ => () + _ => (), } } for item in cert.subject.iter() { for sub_item in item.iter() { - subject.push( - Subject { - oid: sub_item.attr_type().to_string(), - value: sub_item.attr_value().data.to_vec() - } - ) + subject.push(Subject { + oid: sub_item.attr_type().to_string(), + value: sub_item.attr_value().data.to_vec(), + }) } } for item in cert.issuer.iter() { for sub_item in item.iter() { - issuer.push( - Subject { - oid: sub_item.attr_type().to_string(), - value: sub_item.attr_value().data.to_vec() - } - ) + issuer.push(Subject { + oid: sub_item.attr_type().to_string(), + value: sub_item.attr_value().data.to_vec(), + }) } } - return Ok( - Certificate { - version: match cert.version() { - X509Version::V1 => 0, - X509Version::V2 => 1, - X509Version::V3 => 2, - _ => 0xFF - }, - serial_number: cert.raw_serial_as_string(), - raw_serial_number: cert.raw_serial().to_vec(), - not_valid_before: cert.validity.not_before.timestamp(), - not_valid_after: cert.validity.not_after.timestamp(), - extensions: extensions, - subject: subject, - issuer: issuer, - public_bytes: certificate_der.as_bytes().to_vec(), - public_key: match cert.public_key().parsed() { - Ok(PublicKey::EC(pts)) => { - pts.data().to_vec() - }, - Ok(PublicKey::DSA(cert_decoded)) => { - cert_decoded.to_vec() - }, - _ => cert.public_key().raw.to_vec() - }, - } - ) - }, + return Ok(Certificate { + version: match cert.version() { + X509Version::V1 => 0, + X509Version::V2 => 1, + X509Version::V3 => 2, + _ => 0xFF, + }, + serial_number: cert.raw_serial_as_string(), + raw_serial_number: cert.raw_serial().to_vec(), + not_valid_before: cert.validity.not_before.timestamp(), + not_valid_after: cert.validity.not_after.timestamp(), + extensions: extensions, + subject: subject, + issuer: issuer, + public_bytes: certificate_der.as_bytes().to_vec(), + public_key: match cert.public_key().parsed() { + Ok(PublicKey::EC(pts)) => pts.data().to_vec(), + Ok(PublicKey::DSA(cert_decoded)) => cert_decoded.to_vec(), + _ => cert.public_key().raw.to_vec(), + }, + }); + } _ => Err(CryptoError::new_err("x509 parsing failed")), } - } #[getter] @@ -152,10 +135,7 @@ impl Certificate { } pub fn raw_serial_number<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new( - py, - &self.raw_serial_number - ) + return PyBytes::new(py, &self.raw_serial_number); } #[getter] @@ -188,22 +168,17 @@ impl Certificate { "2.5.4.9" => "STREET".to_string(), "0.9.2342.19200300.100.1.25" => "DC".to_string(), "0.9.2342.19200300.100.1.1" => "UID".to_string(), - _ => "".to_string() + _ => "".to_string(), }; - let _ = values.append( - PyTuple::new( - py, - [ - item.oid.to_object(py), - oid_short.to_object(py), - PyBytes::new( - py, - &item.value - ).into() - ] - ) - ); + let _ = values.append(PyTuple::new( + py, + [ + item.oid.to_object(py), + oid_short.to_object(py), + PyBytes::new(py, &item.value).into(), + ], + )); } return values; @@ -214,7 +189,6 @@ impl Certificate { let values = PyList::empty(py); for item in &self.issuer { - let oid_short = match item.oid.as_str() { "2.5.4.3" => "CN", "2.5.4.7" => "L", @@ -225,22 +199,17 @@ impl Certificate { "2.5.4.9" => "STREET", "0.9.2342.19200300.100.1.25" => "DC", "0.9.2342.19200300.100.1.1" => "UID", - _ => "" + _ => "", }; - let _ = values.append( - PyTuple::new( - py, - [ - item.oid.to_object(py), - oid_short.to_object(py), - PyBytes::new( - py, - &item.value - ).into() - ] - ) - ); + let _ = values.append(PyTuple::new( + py, + [ + item.oid.to_object(py), + oid_short.to_object(py), + PyBytes::new(py, &item.value).into(), + ], + )); } return values; @@ -251,12 +220,7 @@ impl Certificate { for item in &self.extensions { if item.oid == "2.5.29.17" { - let _ = values.append( - PyBytes::new( - py, - &item.value - ) - ); + let _ = values.append(PyBytes::new(py, &item.value)); } } @@ -268,12 +232,7 @@ impl Certificate { for item in &self.extensions { if item.oid == "1.3.6.1.5.5.7.48.1" { - let _ = values.append( - PyBytes::new( - py, - &item.value - ) - ); + let _ = values.append(PyBytes::new(py, &item.value)); } } @@ -285,12 +244,7 @@ impl Certificate { for item in &self.extensions { if item.oid == "1.3.6.1.5.5.7.48.2" { - let _ = values.append( - PyBytes::new( - py, - &item.value - ) - ); + let _ = values.append(PyBytes::new(py, &item.value)); } } @@ -298,17 +252,11 @@ impl Certificate { } pub fn public_bytes<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new( - py, - &self.public_bytes - ) + return PyBytes::new(py, &self.public_bytes); } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new( - py, - &self.public_key - ); + return PyBytes::new(py, &self.public_key); } fn __eq__(&self, other: &Self) -> bool { @@ -316,36 +264,39 @@ impl Certificate { } } - #[pyclass(name = "ServerVerifier", module = "qh3._hazmat")] pub struct ServerVerifier { - inner: Arc + inner: Arc, } #[pymethods] impl ServerVerifier { - #[new] pub fn py_new(authorities: Vec<&PyBytes>) -> PyResult { let mut root_cert_store = RootCertStore::empty(); - root_cert_store.add_parsable_certificates(authorities.into_iter().map(|ca| CertificateDer::from(ca.as_bytes()))); + root_cert_store.add_parsable_certificates( + authorities + .into_iter() + .map(|ca| CertificateDer::from(ca.as_bytes())), + ); let res = WebPkiServerVerifier::builder(Arc::new(root_cert_store)).build(); match res { - Ok(store) => { - Ok( - ServerVerifier { - inner: store - } - ) - }, - Err(_) => Err(CryptoError::new_err("Unable to create the X509 trust store")) + Ok(store) => Ok(ServerVerifier { inner: store }), + Err(_) => Err(CryptoError::new_err( + "Unable to create the X509 trust store", + )), } } #[allow(unreachable_code)] - pub fn verify(&mut self, peer: &PyBytes, intermediaries: Vec<&PyBytes>, server_name: String) -> PyResult<()> { + pub fn verify( + &mut self, + peer: &PyBytes, + intermediaries: Vec<&PyBytes>, + server_name: String, + ) -> PyResult<()> { let peer_der = CertificateDer::from(peer.as_bytes()); let mut intermediaries_der = Vec::new(); @@ -367,22 +318,37 @@ impl ServerVerifier { return match res { Ok(_) => Ok(()), - Err(Error::InvalidCertificate(err)) => { - match err { - CertificateError::UnknownIssuer => Err(SelfSignedCertificateError::new_err("unable to get local issuer certificate")), - CertificateError::NotValidForName => Err(InvalidNameCertificateError::new_err("invalid server name for certificate")), - CertificateError::Expired => Err(ExpiredCertificateError::new_err("server certificate expired")), - CertificateError::NotValidYet => Err(ExpiredCertificateError::new_err("server certificate is not yet valid")), - _ => Err(UnacceptableCertificateError::new_err("the server certificate is unacceptable")) + Err(Error::InvalidCertificate(err)) => match err { + CertificateError::UnknownIssuer => { + Err(SelfSignedCertificateError::new_err( + "unable to get local issuer certificate", + )) } + CertificateError::NotValidForName => { + Err(InvalidNameCertificateError::new_err( + "invalid server name for certificate", + )) + } + CertificateError::Expired => Err(ExpiredCertificateError::new_err( + "server certificate expired", + )), + CertificateError::NotValidYet => Err(ExpiredCertificateError::new_err( + "server certificate is not yet valid", + )), + _ => Err(UnacceptableCertificateError::new_err( + "the server certificate is unacceptable", + )), }, - Err(_) => Err(CryptoError::new_err("the x509 certificate store encountered an error")) - } - }, + Err(_) => Err(CryptoError::new_err( + "the x509 certificate store encountered an error", + )), + }; + } Err(_) => { - return Err(InvalidNameCertificateError::new_err("unparseable server name")); + return Err(InvalidNameCertificateError::new_err( + "unparseable server name", + )); } - } - + }; } } diff --git a/src/headers.rs b/src/headers.rs index d6bca9fe7..940ccff56 100644 --- a/src/headers.rs +++ b/src/headers.rs @@ -1,19 +1,17 @@ use ls_qpack::decoder::{Decoder, DecoderOutput}; -use ls_qpack::encoder::{Encoder}; +use ls_qpack::encoder::Encoder; use ls_qpack::StreamId; -use pyo3::{PyResult, Python, ToPyObject}; -use pyo3::types::{PyBytes, PyList, PyTuple}; -use pyo3::pymethods; -use pyo3::pyclass; use pyo3::exceptions::PyException; +use pyo3::pyclass; +use pyo3::pymethods; +use pyo3::types::{PyBytes, PyList, PyTuple}; +use pyo3::{PyResult, Python, ToPyObject}; pyo3::create_exception!(_hazmat, StreamBlocked, PyException); pyo3::create_exception!(_hazmat, EncoderStreamError, PyException); pyo3::create_exception!(_hazmat, DecoderStreamError, PyException); pyo3::create_exception!(_hazmat, DecompressionFailed, PyException); - - #[pyclass(name = "QpackDecoder", module = "qh3._hazmat")] pub struct QpackDecoder { decoder: Decoder, @@ -29,18 +27,26 @@ unsafe impl Send for QpackEncoder {} #[pymethods] impl QpackEncoder { - #[new] pub fn py_new() -> Self { - QpackEncoder { encoder: Encoder::new() } + QpackEncoder { + encoder: Encoder::new(), + } } - pub fn apply_settings<'a>(&mut self, py: Python<'a>, max_table_capacity: u32, dyn_table_capacity: u32, blocked_streams: u32) -> PyResult<&'a PyBytes> { - let r = self.encoder.configure(max_table_capacity, dyn_table_capacity, blocked_streams).expect("FAILURE"); - - return Ok( - PyBytes::new(py, r.data()) - ); + pub fn apply_settings<'a>( + &mut self, + py: Python<'a>, + max_table_capacity: u32, + dyn_table_capacity: u32, + blocked_streams: u32, + ) -> PyResult<&'a PyBytes> { + let r = self + .encoder + .configure(max_table_capacity, dyn_table_capacity, blocked_streams) + .expect("FAILURE"); + + return Ok(PyBytes::new(py, r.data())); } pub fn feed_decoder(&mut self, data: &PyBytes) -> PyResult<()> { @@ -49,51 +55,47 @@ impl QpackEncoder { match res { Ok(_) => { return Ok(()); - }, + } Err(_) => { - return Err(DecoderStreamError::new_err("an error occurred while feeding data from decoder with qpack data")); + return Err(DecoderStreamError::new_err( + "an error occurred while feeding data from decoder with qpack data", + )); } } } - pub fn encode<'a>(&mut self, py: Python<'a>, stream_id: u64, headers: Vec<(&PyBytes, &PyBytes)>) -> PyResult<&'a PyTuple> { + pub fn encode<'a>( + &mut self, + py: Python<'a>, + stream_id: u64, + headers: Vec<(&PyBytes, &PyBytes)>, + ) -> PyResult<&'a PyTuple> { let mut decoded_vec: Vec<(String, String)> = Vec::new(); for (header, value) in headers.iter() { - decoded_vec.push( - ( - std::str::from_utf8(header.as_bytes()).unwrap().to_string(), - std::str::from_utf8(value.as_bytes()).unwrap().to_string() - ) - ); + decoded_vec.push(( + std::str::from_utf8(header.as_bytes()).unwrap().to_string(), + std::str::from_utf8(value.as_bytes()).unwrap().to_string(), + )); } - let res = self.encoder.encode_all(StreamId::new(stream_id), decoded_vec); + let res = self + .encoder + .encode_all(StreamId::new(stream_id), decoded_vec); match res { Ok(buffer) => { - let encoded_buffer = PyBytes::new( - py, - buffer.header(), - ); + let encoded_buffer = PyBytes::new(py, buffer.header()); - let stream_data = PyBytes::new( - py, - buffer.stream(), - ); + let stream_data = PyBytes::new(py, buffer.stream()); - return Ok( - PyTuple::new( - py, - [ - stream_data, - encoded_buffer - ], - ) - ); - }, + return Ok(PyTuple::new(py, [stream_data, encoded_buffer])); + } Err(abc) => { - return Err(EncoderStreamError::new_err(format!("unable to encode headers {:?}", abc))); + return Err(EncoderStreamError::new_err(format!( + "unable to encode headers {:?}", + abc + ))); } } } @@ -101,10 +103,11 @@ impl QpackEncoder { #[pymethods] impl QpackDecoder { - #[new] pub fn py_new(max_table_capacity: u32, blocked_streams: u32) -> Self { - QpackDecoder { decoder: Decoder::new(max_table_capacity, blocked_streams) } + QpackDecoder { + decoder: Decoder::new(max_table_capacity, blocked_streams), + } } pub fn feed_encoder(&mut self, data: &PyBytes) -> PyResult<()> { @@ -113,36 +116,37 @@ impl QpackDecoder { match res { Ok(_) => { return Ok(()); - }, + } Err(_) => { - return Err(EncoderStreamError::new_err("an error occurred while feeding data from encoder with qpack data")); + return Err(EncoderStreamError::new_err( + "an error occurred while feeding data from encoder with qpack data", + )); } } } - pub fn feed_header<'a>(&mut self, py: Python<'a>, stream_id: u64, data: &PyBytes) -> PyResult<&'a PyTuple> { - let output = self.decoder.decode(StreamId::new(stream_id), data.as_bytes()); + pub fn feed_header<'a>( + &mut self, + py: Python<'a>, + stream_id: u64, + data: &PyBytes, + ) -> PyResult<&'a PyTuple> { + let output = self + .decoder + .decode(StreamId::new(stream_id), data.as_bytes()); match output { Ok(DecoderOutput::Done(ref buffer)) => { let decoded_headers = PyList::new(py, Vec::<(String, String)>::new()); for header in buffer.headers() { - let _ = decoded_headers.append( - PyTuple::new( - py, - [ - PyBytes::new( - py, - header.name().as_bytes() - ), - PyBytes::new( - py, - header.value().as_bytes() - ), - ], - ) - ); + let _ = decoded_headers.append(PyTuple::new( + py, + [ + PyBytes::new(py, header.name().as_bytes()), + PyBytes::new(py, header.value().as_bytes()), + ], + )); } return Ok(PyTuple::new( @@ -150,14 +154,18 @@ impl QpackDecoder { [ PyBytes::new(py, buffer.stream()).to_object(py), decoded_headers.to_object(py), - ] + ], )); - }, + } Ok(DecoderOutput::BlockedStream) => { - return Err(StreamBlocked::new_err("stream is blocked, need more data to pursue decoding")); - }, + return Err(StreamBlocked::new_err( + "stream is blocked, need more data to pursue decoding", + )); + } Err(_) => { - return Err(DecoderStreamError::new_err("an error occurred while decoding the stream qpack data")); + return Err(DecoderStreamError::new_err( + "an error occurred while decoding the stream qpack data", + )); } } } @@ -176,21 +184,13 @@ impl QpackDecoder { let decoded_headers = PyList::new(py, Vec::<(String, String)>::new()); for header in buffer.headers() { - let _ = decoded_headers.append( - PyTuple::new( - py, - [ - PyBytes::new( - py, - header.name().as_bytes() - ), - PyBytes::new( - py, - header.value().as_bytes() - ), - ], - ) - ); + let _ = decoded_headers.append(PyTuple::new( + py, + [ + PyBytes::new(py, header.name().as_bytes()), + PyBytes::new(py, header.value().as_bytes()), + ], + )); } return Ok(PyTuple::new( @@ -198,14 +198,18 @@ impl QpackDecoder { [ PyBytes::new(py, buffer.stream()).to_object(py), decoded_headers.to_object(py), - ] + ], )); - }, + } Ok(DecoderOutput::BlockedStream) => { - return Err(StreamBlocked::new_err("stream is blocked, need more data to pursue decoding")) - }, + return Err(StreamBlocked::new_err( + "stream is blocked, need more data to pursue decoding", + )) + } Err(_) => { - return Err(DecoderStreamError::new_err("an error occurred while decoding the stream qpack data")) + return Err(DecoderStreamError::new_err( + "an error occurred while decoding the stream qpack data", + )) } } } diff --git a/src/hpk.rs b/src/hpk.rs index 6d1901ac1..7168245f8 100644 --- a/src/hpk.rs +++ b/src/hpk.rs @@ -1,19 +1,18 @@ use aws_lc_rs::aead::quic::{HeaderProtectionKey, AES_128, AES_256, CHACHA20}; -use pyo3::{PyResult, Python}; -use pyo3::types::PyBytes; -use pyo3::pymethods; -use pyo3::pyclass; use crate::CryptoError; +use pyo3::pyclass; +use pyo3::pymethods; +use pyo3::types::PyBytes; +use pyo3::{PyResult, Python}; #[pyclass(module = "qh3._hazmat")] pub struct QUICHeaderProtection { - hpk: HeaderProtectionKey + hpk: HeaderProtectionKey, } #[pymethods] impl QUICHeaderProtection { - #[new] pub fn py_new(key: &PyBytes, algorithm: u16) -> Self { QUICHeaderProtection { @@ -22,10 +21,11 @@ impl QUICHeaderProtection { 128 => &AES_128, 256 => &AES_256, 20 => &CHACHA20, - _ => panic!("unsupported") + _ => panic!("unsupported"), }, - &key.as_bytes() - ).expect("FAILURE") + &key.as_bytes(), + ) + .expect("FAILURE"), } } @@ -33,13 +33,10 @@ impl QUICHeaderProtection { let res = self.hpk.new_mask(&sample.as_bytes()); return match res { - Err(_) => Err(CryptoError::new_err("unable to issue mask protection header")), - Ok(data) => Ok( - PyBytes::new( - py, - &data - ) - ) - } + Err(_) => Err(CryptoError::new_err( + "unable to issue mask protection header", + )), + Ok(data) => Ok(PyBytes::new(py, &data)), + }; } } diff --git a/src/lib.rs b/src/lib.rs index 85b99cdf8..7318d1499 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,27 +1,39 @@ -use pyo3::{prelude::*}; use pyo3::exceptions::PyException; +use pyo3::prelude::*; -mod headers; mod aead; -mod certificate; -mod rsa; mod agreement; -mod private_key; -mod pkcs8; +mod buffer; +mod certificate; +mod headers; mod hpk; mod ocsp; -mod buffer; +mod pkcs8; +mod private_key; +mod rsa; -pub use self::headers::{QpackDecoder, QpackEncoder, StreamBlocked, EncoderStreamError, DecoderStreamError, DecompressionFailed}; -pub use self::aead::{AeadChaCha20Poly1305, AeadAes128Gcm, AeadAes256Gcm}; -pub use self::certificate::{ServerVerifier, Certificate, SelfSignedCertificateError, InvalidNameCertificateError, ExpiredCertificateError, UnacceptableCertificateError}; -pub use self::rsa::{Rsa}; -pub use self::private_key::{RsaPrivateKey, DsaPrivateKey, Ed25519PrivateKey, EcPrivateKey, verify_with_public_key, SignatureError}; -pub use self::agreement::{X25519KeyExchange, ECDHP256KeyExchange, ECDHP384KeyExchange, ECDHP521KeyExchange, X25519Kyber768Draft00KeyExchange}; -pub use self::pkcs8::{PrivateKeyInfo, KeyType}; -pub use self::hpk::{QUICHeaderProtection}; -pub use self::ocsp::{OCSPResponse, OCSPCertStatus, OCSPResponseStatus, ReasonFlags, OCSPRequest}; +pub use self::aead::{AeadAes128Gcm, AeadAes256Gcm, AeadChaCha20Poly1305}; +pub use self::agreement::{ + ECDHP256KeyExchange, ECDHP384KeyExchange, ECDHP521KeyExchange, X25519KeyExchange, + X25519Kyber768Draft00KeyExchange, +}; pub use self::buffer::{Buffer, BufferReadError, BufferWriteError}; +pub use self::certificate::{ + Certificate, ExpiredCertificateError, InvalidNameCertificateError, SelfSignedCertificateError, + ServerVerifier, UnacceptableCertificateError, +}; +pub use self::headers::{ + DecoderStreamError, DecompressionFailed, EncoderStreamError, QpackDecoder, QpackEncoder, + StreamBlocked, +}; +pub use self::hpk::QUICHeaderProtection; +pub use self::ocsp::{OCSPCertStatus, OCSPRequest, OCSPResponse, OCSPResponseStatus, ReasonFlags}; +pub use self::pkcs8::{KeyType, PrivateKeyInfo}; +pub use self::private_key::{ + verify_with_public_key, DsaPrivateKey, EcPrivateKey, Ed25519PrivateKey, RsaPrivateKey, + SignatureError, +}; +pub use self::rsa::Rsa; pyo3::create_exception!(_hazmat, CryptoError, PyException); @@ -41,10 +53,22 @@ fn _hazmat(py: Python, m: &PyModule) -> PyResult<()> { // Certificate Store X509 Verification + Certificate Representation m.add_class::()?; m.add_class::()?; - m.add("SelfSignedCertificateError", py.get_type::())?; - m.add("InvalidNameCertificateError", py.get_type::())?; - m.add("ExpiredCertificateError", py.get_type::())?; - m.add("UnacceptableCertificateError", py.get_type::())?; + m.add( + "SelfSignedCertificateError", + py.get_type::(), + )?; + m.add( + "InvalidNameCertificateError", + py.get_type::(), + )?; + m.add( + "ExpiredCertificateError", + py.get_type::(), + )?; + m.add( + "UnacceptableCertificateError", + py.get_type::(), + )?; // RSA specialized for the Retry Token m.add_class::()?; // Header protection mask diff --git a/src/ocsp.rs b/src/ocsp.rs index d6f3a841e..5336f76a9 100644 --- a/src/ocsp.rs +++ b/src/ocsp.rs @@ -1,16 +1,19 @@ // OCSP Response Parser and Request Builder // This module is created for Niquests // qh3 has no use for it and we won't implement it for this package -use pyo3::{PyResult, Python}; -use pyo3::types::PyBytes; -use pyo3::pymethods; use pyo3::pyclass; +use pyo3::pymethods; +use pyo3::types::PyBytes; +use pyo3::{PyResult, Python}; use der::{Decode, Encode}; use pyo3::exceptions::PyValueError; -use x509_ocsp::{OcspResponse, BasicOcspResponse, SingleResponse, OcspResponseStatus as InternalOcspResponseStatus, CertStatus as InternalCertStatus, OcspRequest as InternalOcspRequest, Request}; -use x509_ocsp::builder::OcspRequestBuilder; use x509_cert::Certificate; +use x509_ocsp::builder::OcspRequestBuilder; +use x509_ocsp::{ + BasicOcspResponse, CertStatus as InternalCertStatus, OcspRequest as InternalOcspRequest, + OcspResponse, OcspResponseStatus as InternalOcspResponseStatus, Request, SingleResponse, +}; use sha1::Sha1; @@ -51,7 +54,6 @@ pub enum OCSPCertStatus { UNKNOWN = 2, } - #[pyclass(module = "qh3._hazmat")] #[allow(non_camel_case_types)] pub struct OCSPResponse { @@ -71,7 +73,9 @@ impl OCSPResponse { return Err(PyValueError::new_err("OCSP Server did not provide answers")); } - let inner_resp: BasicOcspResponse = BasicOcspResponse::from_der(&ocsp_res.response_bytes.unwrap().response.as_bytes()).unwrap(); + let inner_resp: BasicOcspResponse = + BasicOcspResponse::from_der(&ocsp_res.response_bytes.unwrap().response.as_bytes()) + .unwrap(); if inner_resp.tbs_response_data.responses.len() == 0 { return Err(PyValueError::new_err("OCSP Server did not provide answers")); @@ -79,43 +83,48 @@ impl OCSPResponse { let first_resp_for_cert: &SingleResponse = &inner_resp.tbs_response_data.responses[0]; - return Ok( - OCSPResponse { - next_update: first_resp_for_cert.next_update.unwrap().0.to_unix_duration().as_secs(), - response_status: match ocsp_res.response_status { - InternalOcspResponseStatus::Successful => OCSPResponseStatus::SUCCESSFUL, - InternalOcspResponseStatus::MalformedRequest => OCSPResponseStatus::MALFORMED_REQUEST, - InternalOcspResponseStatus::InternalError => OCSPResponseStatus::INTERNAL_ERROR, - InternalOcspResponseStatus::TryLater => OCSPResponseStatus::TRY_LATER, - InternalOcspResponseStatus::SigRequired => OCSPResponseStatus::SIG_REQUIRED, - InternalOcspResponseStatus::Unauthorized => OCSPResponseStatus::UNAUTHORIZED - }, - certificate_status: match first_resp_for_cert.cert_status { - InternalCertStatus::Good(..) => OCSPCertStatus::GOOD, - InternalCertStatus::Revoked(_) => OCSPCertStatus::REVOKED, - InternalCertStatus::Unknown(_) => OCSPCertStatus::UNKNOWN, - }, - revocation_reason: match first_resp_for_cert.cert_status { - InternalCertStatus::Revoked(info) => match info.revocation_reason { - Some(reason) => match reason as u8 { - 0 => Some(ReasonFlags::unspecified), - 1 => Some(ReasonFlags::key_compromise), - 2 => Some(ReasonFlags::ca_compromise), - 3 => Some(ReasonFlags::affiliation_changed), - 4 => Some(ReasonFlags::superseded), - 5 => Some(ReasonFlags::cessation_of_operation), - 6 => Some(ReasonFlags::certificate_hold), - 8 => Some(ReasonFlags::remove_from_crl), - 9 => Some(ReasonFlags::privilege_withdrawn), - 10 => Some(ReasonFlags::aa_compromise), - _ => None - }, - _ => None - }, - InternalCertStatus::Good(_) | InternalCertStatus::Unknown(_) => None + return Ok(OCSPResponse { + next_update: first_resp_for_cert + .next_update + .unwrap() + .0 + .to_unix_duration() + .as_secs(), + response_status: match ocsp_res.response_status { + InternalOcspResponseStatus::Successful => OCSPResponseStatus::SUCCESSFUL, + InternalOcspResponseStatus::MalformedRequest => { + OCSPResponseStatus::MALFORMED_REQUEST } - } - ) + InternalOcspResponseStatus::InternalError => OCSPResponseStatus::INTERNAL_ERROR, + InternalOcspResponseStatus::TryLater => OCSPResponseStatus::TRY_LATER, + InternalOcspResponseStatus::SigRequired => OCSPResponseStatus::SIG_REQUIRED, + InternalOcspResponseStatus::Unauthorized => OCSPResponseStatus::UNAUTHORIZED, + }, + certificate_status: match first_resp_for_cert.cert_status { + InternalCertStatus::Good(..) => OCSPCertStatus::GOOD, + InternalCertStatus::Revoked(_) => OCSPCertStatus::REVOKED, + InternalCertStatus::Unknown(_) => OCSPCertStatus::UNKNOWN, + }, + revocation_reason: match first_resp_for_cert.cert_status { + InternalCertStatus::Revoked(info) => match info.revocation_reason { + Some(reason) => match reason as u8 { + 0 => Some(ReasonFlags::unspecified), + 1 => Some(ReasonFlags::key_compromise), + 2 => Some(ReasonFlags::ca_compromise), + 3 => Some(ReasonFlags::affiliation_changed), + 4 => Some(ReasonFlags::superseded), + 5 => Some(ReasonFlags::cessation_of_operation), + 6 => Some(ReasonFlags::certificate_hold), + 8 => Some(ReasonFlags::remove_from_crl), + 9 => Some(ReasonFlags::privilege_withdrawn), + 10 => Some(ReasonFlags::aa_compromise), + _ => None, + }, + _ => None, + }, + InternalCertStatus::Good(_) | InternalCertStatus::Unknown(_) => None, + }, + }); } #[getter] @@ -139,10 +148,9 @@ impl OCSPResponse { } } - #[pyclass(module = "qh3._hazmat")] pub struct OCSPRequest { - inner_request: Vec + inner_request: Vec, } #[pymethods] @@ -157,19 +165,14 @@ impl OCSPRequest { .build(); return match req.to_der() { - Ok(raw_der) => Ok( - OCSPRequest { - inner_request: raw_der - } - ), - Err(_) => Err(PyValueError::new_err("unable to generate the request")) + Ok(raw_der) => Ok(OCSPRequest { + inner_request: raw_der, + }), + Err(_) => Err(PyValueError::new_err("unable to generate the request")), }; } pub fn public_bytes<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new( - py, - &self.inner_request - ); + return PyBytes::new(py, &self.inner_request); } } diff --git a/src/pkcs8.rs b/src/pkcs8.rs index df82e5f34..952e6edc9 100644 --- a/src/pkcs8.rs +++ b/src/pkcs8.rs @@ -1,17 +1,16 @@ -use pyo3::Python; -use pyo3::types::PyBytes; -use pyo3::pymethods; use pyo3::pyclass; +use pyo3::pymethods; +use pyo3::types::PyBytes; +use pyo3::Python; use pkcs8::{der::Encode, DecodePrivateKey, Error, PrivateKeyInfo as InternalPrivateKeyInfo}; use rsa::{ pkcs1::DecodeRsaPrivateKey, - pkcs8::{LineEnding, EncodePrivateKey, ObjectIdentifier}, + pkcs8::{EncodePrivateKey, LineEnding, ObjectIdentifier}, RsaPrivateKey, }; -use rustls_pemfile::{Item, read_one_from_slice}; - +use rustls_pemfile::{read_one_from_slice, Item}; #[pyclass(module = "qh3._hazmat")] #[derive(Clone, Copy)] @@ -37,33 +36,31 @@ impl TryFrom> for PrivateKeyInfo { fn try_from(pkcs8: InternalPrivateKeyInfo<'_>) -> Result { let der_document = pkcs8.to_der().unwrap(); - let rsa_oid = ObjectIdentifier::new_unwrap("1.2.840.113549.1.1.1").as_bytes().to_vec(); - let dsa_oid = ObjectIdentifier::new_unwrap("1.2.840.10040.4.1").as_bytes().to_vec(); + let rsa_oid = ObjectIdentifier::new_unwrap("1.2.840.113549.1.1.1") + .as_bytes() + .to_vec(); + let dsa_oid = ObjectIdentifier::new_unwrap("1.2.840.10040.4.1") + .as_bytes() + .to_vec(); if rsa_oid == pkcs8.algorithm.oid.as_bytes().to_vec() { - return Ok( - PrivateKeyInfo{ - der_encoded: der_document.clone(), - cert_type: KeyType::RSA - } - ); + return Ok(PrivateKeyInfo { + der_encoded: der_document.clone(), + cert_type: KeyType::RSA, + }); } if dsa_oid == pkcs8.algorithm.oid.as_bytes().to_vec() { - return Ok( - PrivateKeyInfo{ - der_encoded: der_document.clone(), - cert_type: KeyType::DSA - } - ); + return Ok(PrivateKeyInfo { + der_encoded: der_document.clone(), + cert_type: KeyType::DSA, + }); } - return Ok( - PrivateKeyInfo{ - der_encoded: der_document.clone(), - cert_type: KeyType::ED25519 - } - ); + return Ok(PrivateKeyInfo { + der_encoded: der_document.clone(), + cert_type: KeyType::ED25519, + }); } } @@ -83,22 +80,26 @@ impl PrivateKeyInfo { panic!("unsupported"); } - let rsa_key: RsaPrivateKey = RsaPrivateKey::from_pkcs1_der(&key.secret_pkcs1_der()).unwrap(); + let rsa_key: RsaPrivateKey = + RsaPrivateKey::from_pkcs1_der(&key.secret_pkcs1_der()).unwrap(); - let pkcs8_pem = rsa_key - .to_pkcs8_pem(LineEnding::LF).expect("FAILURE"); + let pkcs8_pem = rsa_key.to_pkcs8_pem(LineEnding::LF).expect("FAILURE"); let pkcs8_pem: &str = pkcs8_pem.as_ref(); return PrivateKeyInfo::from_pkcs8_pem(&pkcs8_pem).unwrap(); - }, + } Item::Pkcs8Key(_key) => { if is_encrypted { - return PrivateKeyInfo::from_pkcs8_encrypted_pem(&decoded_bytes, password.unwrap().as_bytes()).unwrap(); + return PrivateKeyInfo::from_pkcs8_encrypted_pem( + &decoded_bytes, + password.unwrap().as_bytes(), + ) + .unwrap(); } return PrivateKeyInfo::from_pkcs8_pem(&decoded_bytes).unwrap(); - }, + } Item::Sec1Key(key) => { if is_encrypted { panic!("unsupported"); @@ -114,11 +115,10 @@ impl PrivateKeyInfo { _ => panic!("unsupported sec1 key"), }, der_encoded: sec1_der, - } - }, + }; + } _ => panic!("unsupported"), }; - } pub fn get_type(&self) -> KeyType { @@ -126,9 +126,6 @@ impl PrivateKeyInfo { } pub fn public_bytes<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new( - py, - &self.der_encoded - ); + return PyBytes::new(py, &self.der_encoded); } } diff --git a/src/private_key.rs b/src/private_key.rs index 0dbde05b1..112fac7d4 100644 --- a/src/private_key.rs +++ b/src/private_key.rs @@ -1,64 +1,59 @@ -use rsa::{RsaPrivateKey as InternalRsaPrivateKey, RsaPublicKey as InternalRsaPublicKey}; -use dsa::{SigningKey as InternalDsaPrivateKey}; use aws_lc_rs::signature::{ - EcdsaKeyPair as InternalEcPrivateKey, - KeyPair, - ECDSA_P256_SHA256_ASN1_SIGNING, - ECDSA_P384_SHA384_ASN1_SIGNING, - ECDSA_P521_SHA512_ASN1_SIGNING, - Ed25519KeyPair as InternalEd25519PrivateKey, + EcdsaKeyPair as InternalEcPrivateKey, Ed25519KeyPair as InternalEd25519PrivateKey, KeyPair, + ECDSA_P256_SHA256_ASN1_SIGNING, ECDSA_P384_SHA384_ASN1_SIGNING, ECDSA_P521_SHA512_ASN1_SIGNING, }; +use dsa::SigningKey as InternalDsaPrivateKey; +use rsa::{RsaPrivateKey as InternalRsaPrivateKey, RsaPublicKey as InternalRsaPublicKey}; -use rsa::pkcs1v15::{SigningKey as InternalRsaPkcsSigningKey, Signature as RsaPkcsSignature}; +use rsa::pkcs1v15::{Signature as RsaPkcsSignature, SigningKey as InternalRsaPkcsSigningKey}; use rsa::pss::{Signature as RsaPssSignature, SigningKey as InternalRsaPssSigningKey}; -use rsa::sha2::{Sha256, Sha512, Sha384}; -use rsa::signature::Signer; +use rsa::pkcs1v15::VerifyingKey as RsaPkcsVerifyingKey; +use rsa::pss::VerifyingKey as RsaPssVerifyingKey; +use rsa::sha2::{Sha256, Sha384, Sha512}; use rsa::signature::SignatureEncoding; -use rsa::pss::{VerifyingKey as RsaPssVerifyingKey}; -use rsa::pkcs1v15::{VerifyingKey as RsaPkcsVerifyingKey}; +use rsa::signature::Signer; use rsa::signature::Verifier; -use ed25519_dalek::{VerifyingKey as Ed25519VerifyingKey, Signature as Ed25519Signature}; +use ed25519_dalek::{Signature as Ed25519Signature, VerifyingKey as Ed25519VerifyingKey}; use pkcs8::DecodePrivateKey; -use pkcs8::EncodePublicKey; use pkcs8::DecodePublicKey; +use pkcs8::EncodePublicKey; -use aws_lc_rs::signature::{UnparsedPublicKey}; use aws_lc_rs::error::Unspecified; -use aws_lc_rs::signature; use aws_lc_rs::rand::SystemRandom; +use aws_lc_rs::signature; +use aws_lc_rs::signature::UnparsedPublicKey; -use pyo3::{PyResult, Python}; -use pyo3::types::PyBytes; -use pyo3::pymethods; -use pyo3::pyfunction; -use pyo3::pyclass; use pyo3::exceptions::PyException; +use pyo3::pyclass; +use pyo3::pyfunction; +use pyo3::pymethods; +use pyo3::types::PyBytes; +use pyo3::{PyResult, Python}; pyo3::create_exception!(_hazmat, SignatureError, PyException); - #[pyclass(module = "qh3._hazmat")] pub struct EcPrivateKey { inner: InternalEcPrivateKey, - curve: u32 + curve: u32, } #[pyclass(module = "qh3._hazmat")] pub struct Ed25519PrivateKey { - inner: InternalEd25519PrivateKey + inner: InternalEd25519PrivateKey, } #[pyclass(module = "qh3._hazmat")] pub struct DsaPrivateKey { - inner: InternalDsaPrivateKey + inner: InternalDsaPrivateKey, } #[pyclass(module = "qh3._hazmat")] pub struct RsaPrivateKey { - inner: InternalRsaPrivateKey + inner: InternalRsaPrivateKey, } #[pymethods] @@ -66,28 +61,21 @@ impl Ed25519PrivateKey { #[new] pub fn py_new(pkcs8: &PyBytes) -> Self { Ed25519PrivateKey { - inner: InternalEd25519PrivateKey::from_pkcs8(&pkcs8.as_bytes()).expect("FAILURE") + inner: InternalEd25519PrivateKey::from_pkcs8(&pkcs8.as_bytes()).expect("FAILURE"), } } pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { let signature = self.inner.sign(&data.as_bytes()); - return PyBytes::new( - py, - &signature.as_ref() - ); + return PyBytes::new(py, &signature.as_ref()); } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new( - py, - &self.inner.public_key().as_ref() - ); + return PyBytes::new(py, &self.inner.public_key().as_ref()); } } - #[pymethods] impl EcPrivateKey { #[new] @@ -100,26 +88,21 @@ impl EcPrivateKey { }; return EcPrivateKey { - inner: InternalEcPrivateKey::from_pkcs8(&signing_algorithm, &pkcs8.as_bytes()).expect("FAILURE"), - curve: curve_type - } + inner: InternalEcPrivateKey::from_pkcs8(&signing_algorithm, &pkcs8.as_bytes()) + .expect("FAILURE"), + curve: curve_type, + }; } pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { let rng = SystemRandom::new(); let signature = self.inner.sign(&rng, &data.as_bytes()); - return PyBytes::new( - py, - &signature.unwrap().as_ref() - ); + return PyBytes::new(py, &signature.unwrap().as_ref()); } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new( - py, - &self.inner.public_key().as_ref() - ); + return PyBytes::new(py, &self.inner.public_key().as_ref()); } #[getter] @@ -128,96 +111,83 @@ impl EcPrivateKey { } } - #[pymethods] impl DsaPrivateKey { #[new] pub fn py_new(pkcs8: &PyBytes) -> Self { return DsaPrivateKey { - inner: InternalDsaPrivateKey::from_pkcs8_der(&pkcs8.as_bytes()).expect("FAILURE") - } + inner: InternalDsaPrivateKey::from_pkcs8_der(&pkcs8.as_bytes()).expect("FAILURE"), + }; } pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { let signature = self.inner.sign(&data.as_bytes()); - return PyBytes::new( - py, - &signature.to_bytes() - ); + return PyBytes::new(py, &signature.to_bytes()); } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { return PyBytes::new( py, - &self.inner.verifying_key().to_public_key_der().expect("FAILURE").as_bytes() + &self + .inner + .verifying_key() + .to_public_key_der() + .expect("FAILURE") + .as_bytes(), ); } } - #[pymethods] impl RsaPrivateKey { #[new] pub fn py_new(pkcs8: &PyBytes) -> Self { return RsaPrivateKey { - inner: InternalRsaPrivateKey::from_pkcs8_der(&pkcs8.as_bytes()).expect("FAILURE") - } + inner: InternalRsaPrivateKey::from_pkcs8_der(&pkcs8.as_bytes()).expect("FAILURE"), + }; } - pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes, is_pss_padding: bool, hash_size: u32) -> &'a PyBytes { - + pub fn sign<'a>( + &self, + py: Python<'a>, + data: &PyBytes, + is_pss_padding: bool, + hash_size: u32, + ) -> &'a PyBytes { let private_key = self.inner.clone(); match is_pss_padding { true => match hash_size { 256 => { let signer = InternalRsaPssSigningKey::::new(private_key); - return PyBytes::new( - py, - &signer.sign(&data.as_bytes()).to_vec() - ); - }, + return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + } 384 => { let signer = InternalRsaPssSigningKey::::new(private_key); - return PyBytes::new( - py, - &signer.sign(&data.as_bytes()).to_vec() - ); - }, + return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + } 512 => { let signer = InternalRsaPssSigningKey::::new(private_key); - return PyBytes::new( - py, - &signer.sign(&data.as_bytes()).to_vec() - ); - }, - _ => panic!("unsupported") + return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + } + _ => panic!("unsupported"), }, false => match hash_size { 256 => { let signer = InternalRsaPkcsSigningKey::::new(private_key); - return PyBytes::new( - py, - &signer.sign(&data.as_bytes()).to_vec() - ); - }, + return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + } 384 => { let signer = InternalRsaPkcsSigningKey::::new(private_key); - return PyBytes::new( - py, - &signer.sign(&data.as_bytes()).to_vec() - ); - }, + return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + } 512 => { let signer = InternalRsaPkcsSigningKey::::new(private_key); - return PyBytes::new( - py, - &signer.sign(&data.as_bytes()).to_vec() - ); - }, - _ => panic!("unsupported") - } + return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + } + _ => panic!("unsupported"), + }, }; } @@ -226,16 +196,19 @@ impl RsaPrivateKey { return PyBytes::new( py, - &public_key.to_public_key_der().as_ref().unwrap().to_vec() - ) + &public_key.to_public_key_der().as_ref().unwrap().to_vec(), + ); } } - #[pyfunction] #[allow(unreachable_code)] -pub fn verify_with_public_key(public_key_raw: &PyBytes, algorithm: u32, message: &PyBytes, signature: &PyBytes) -> PyResult<()> { - +pub fn verify_with_public_key( + public_key_raw: &PyBytes, + algorithm: u32, + message: &PyBytes, + signature: &PyBytes, +) -> PyResult<()> { let pss_rsae_blind_signature = 0x0804..0x0806; let pss_pss_blind_signature = 0x0809..0x080B; let pkcs115_blind_signature = [0x0401, 0x0501, 0x0601]; @@ -243,83 +216,109 @@ pub fn verify_with_public_key(public_key_raw: &PyBytes, algorithm: u32, message: let public_key_bytes = public_key_raw.as_bytes(); // Can't get RSA signature to work using UnparsedPublicKey, I could have missed something...? - if pss_rsae_blind_signature.contains(&algorithm) || pss_pss_blind_signature.contains(&algorithm) || pkcs115_blind_signature.contains(&algorithm) { - let rsa_parsed_public_key = InternalRsaPublicKey::from_public_key_der(&public_key_bytes).expect("FAILURE"); + if pss_rsae_blind_signature.contains(&algorithm) + || pss_pss_blind_signature.contains(&algorithm) + || pkcs115_blind_signature.contains(&algorithm) + { + let rsa_parsed_public_key = + InternalRsaPublicKey::from_public_key_der(&public_key_bytes).expect("FAILURE"); return match algorithm { 0x0804 | 0x0809 => { let alt_verifier = RsaPssVerifyingKey::::new(rsa_parsed_public_key); - let res = alt_verifier.verify(&message.as_bytes(), &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE")); + let res = alt_verifier.verify( + &message.as_bytes(), + &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE"), + ); return match res { Err(_) => Err(SignatureError::new_err("signature mismatch (rsa)")), - _ => Ok(()) - } - }, + _ => Ok(()), + }; + } 0x0805 | 0x080A => { let alt_verifier = RsaPssVerifyingKey::::new(rsa_parsed_public_key); - let res = alt_verifier.verify(&message.as_bytes(), &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE")); + let res = alt_verifier.verify( + &message.as_bytes(), + &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE"), + ); return match res { Err(_) => Err(SignatureError::new_err("signature mismatch (rsa)")), - _ => Ok(()) - } - }, + _ => Ok(()), + }; + } 0x0806 | 0x080B => { let alt_verifier = RsaPssVerifyingKey::::new(rsa_parsed_public_key); - let res = alt_verifier.verify(&message.as_bytes(), &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE")); + let res = alt_verifier.verify( + &message.as_bytes(), + &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE"), + ); return match res { Err(_) => Err(SignatureError::new_err("signature mismatch (rsa)")), - _ => Ok(()) - } - }, + _ => Ok(()), + }; + } 0x0401 => { let alt_verifier = RsaPkcsVerifyingKey::::new(rsa_parsed_public_key); - let res = alt_verifier.verify(&message.as_bytes(), &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE")); + let res = alt_verifier.verify( + &message.as_bytes(), + &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE"), + ); return match res { Err(_) => Err(SignatureError::new_err("signature mismatch (rsa)")), - _ => Ok(()) - } - }, + _ => Ok(()), + }; + } 0x0501 => { let alt_verifier = RsaPkcsVerifyingKey::::new(rsa_parsed_public_key); - let res = alt_verifier.verify(&message.as_bytes(), &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE")); + let res = alt_verifier.verify( + &message.as_bytes(), + &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE"), + ); return match res { Err(_) => Err(SignatureError::new_err("signature mismatch (rsa)")), - _ => Ok(()) - } - }, + _ => Ok(()), + }; + } 0x0601 => { let alt_verifier = RsaPkcsVerifyingKey::::new(rsa_parsed_public_key); - let res = alt_verifier.verify(&message.as_bytes(), &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE")); + let res = alt_verifier.verify( + &message.as_bytes(), + &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE"), + ); return match res { Err(_) => Err(SignatureError::new_err("signature mismatch (rsa)")), - _ => Ok(()) - } - }, + _ => Ok(()), + }; + } - _ => panic!("unreachable statement") + _ => panic!("unreachable statement"), }; } if algorithm == 0x0807 { - let ed25519_verifier: Ed25519VerifyingKey = Ed25519VerifyingKey::from_public_key_der(&public_key_bytes).expect("FAILURE"); - let res = ed25519_verifier.verify(&message.as_bytes(), &Ed25519Signature::from_bytes(signature.as_bytes()[0..64].try_into().unwrap())); + let ed25519_verifier: Ed25519VerifyingKey = + Ed25519VerifyingKey::from_public_key_der(&public_key_bytes).expect("FAILURE"); + let res = ed25519_verifier.verify( + &message.as_bytes(), + &Ed25519Signature::from_bytes(signature.as_bytes()[0..64].try_into().unwrap()), + ); return match res { Err(_) => Err(SignatureError::new_err("signature mismatch (ed25519)")), - _ => Ok(()) + _ => Ok(()), }; } @@ -328,15 +327,15 @@ pub fn verify_with_public_key(public_key_raw: &PyBytes, algorithm: u32, message: 0x0403 => &signature::ECDSA_P256_SHA256_ASN1, 0x0503 => &signature::ECDSA_P384_SHA384_ASN1, 0x0603 => &signature::ECDSA_P521_SHA512_ASN1, - _ => panic!("unsupported algorithm") + _ => panic!("unsupported algorithm"), }, - public_key_bytes + public_key_bytes, ); let res = public_key.verify(&message.as_bytes(), &signature.as_bytes()); return match res { Err(Unspecified) => Err(SignatureError::new_err("signature mismatch (ecdsa)")), - _ => Ok(()) - } + _ => Ok(()), + }; } diff --git a/src/rsa.rs b/src/rsa.rs index 21fb2d9ad..32f9d16ad 100644 --- a/src/rsa.rs +++ b/src/rsa.rs @@ -1,11 +1,9 @@ -use pyo3::Python; -use pyo3::types::PyBytes; -use pyo3::pymethods; use pyo3::pyclass; +use pyo3::pymethods; +use pyo3::types::PyBytes; +use pyo3::Python; - -use rsa::{RsaPrivateKey, RsaPublicKey, Oaep, sha2::Sha256}; - +use rsa::{sha2::Sha256, Oaep, RsaPrivateKey, RsaPublicKey}; #[pyclass(module = "qh3._hazmat")] pub struct Rsa { @@ -34,23 +32,23 @@ impl Rsa { let padding = Oaep::new::(); let mut rng = rand::thread_rng(); - let enc_data = self.public_key.encrypt(&mut rng, padding, &payload_to_enc[..]).expect("failed to encrypt"); + let enc_data = self + .public_key + .encrypt(&mut rng, padding, &payload_to_enc[..]) + .expect("failed to encrypt"); - return PyBytes::new( - py, - &enc_data - ); + return PyBytes::new(py, &enc_data); } pub fn decrypt<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { let payload_to_dec = data.as_bytes(); let padding = Oaep::new::(); - let dec_data = self.private_key.decrypt(padding, &payload_to_dec).expect("failed to decrypt"); + let dec_data = self + .private_key + .decrypt(padding, &payload_to_dec) + .expect("failed to decrypt"); - return PyBytes::new( - py, - &dec_data - ); + return PyBytes::new(py, &dec_data); } } From 516317fb30808ef48f82191653145f1c1ed03fe5 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Sun, 29 Dec 2024 16:08:55 +0100 Subject: [PATCH 06/39] :art: reformat noxfile.py --- noxfile.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/noxfile.py b/noxfile.py index d633a92cf..372032ffb 100644 --- a/noxfile.py +++ b/noxfile.py @@ -14,7 +14,7 @@ def tests_impl( session.install("-U", "pip", "setuptools", silent=False) session.install("-r", "dev-requirements.txt", silent=False) - session.install(f".", silent=False) + session.install(".", silent=False) # Show the pip version. session.run("pip", "--version") From ce4560d51c499eccdaf831ebab52935b7239c60e Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Sun, 29 Dec 2024 16:11:03 +0100 Subject: [PATCH 07/39] :art: fix clippy remaining errors --- src/aead.rs | 28 +++++++-------- src/agreement.rs | 36 ++++++++----------- src/buffer.rs | 52 +++++++++++++-------------- src/certificate.rs | 50 +++++++++++++------------- src/headers.rs | 78 ++++++++++++++++------------------------- src/hpk.rs | 8 ++--- src/ocsp.rs | 30 ++++++++-------- src/pkcs8.rs | 24 ++++++------- src/private_key.rs | 87 +++++++++++++++++++++++----------------------- src/rsa.rs | 12 +++---- 10 files changed, 188 insertions(+), 217 deletions(-) diff --git a/src/aead.rs b/src/aead.rs index 54faf2912..9202e4b3a 100644 --- a/src/aead.rs +++ b/src/aead.rs @@ -58,10 +58,10 @@ impl AeadAes256Gcm { &mut in_out_buffer, ); - return match res { + match res { Ok(_) => Ok(PyBytes::new(py, &in_out_buffer[0..plaintext_len])), Err(_) => Err(CryptoError::new_err("decryption failed")), - }; + } } pub fn encrypt<'a>( @@ -85,10 +85,10 @@ impl AeadAes256Gcm { &mut in_out_buffer, ); - return match res { + match res { Ok(_) => Ok(PyBytes::new(py, &in_out_buffer)), Err(_) => Err(CryptoError::new_err("encryption failed")), - }; + } } } @@ -121,10 +121,10 @@ impl AeadAes128Gcm { &mut in_out_buffer, ); - return match res { + match res { Ok(_) => Ok(PyBytes::new(py, &in_out_buffer[0..plaintext_len])), Err(_) => Err(CryptoError::new_err("decryption failed")), - }; + } } pub fn encrypt<'a>( @@ -148,10 +148,10 @@ impl AeadAes128Gcm { &mut in_out_buffer, ); - return match res { + match res { Ok(_) => Ok(PyBytes::new(py, &in_out_buffer)), Err(_) => Err(CryptoError::new_err("encryption failed")), - }; + } } } @@ -178,14 +178,14 @@ impl AeadChaCha20Poly1305 { let res = cipher.decrypt_in_place( nonce.as_bytes().into(), - &associated_data.as_bytes(), + associated_data.as_bytes(), &mut in_out_buffer, ); - return match res { + match res { Ok(_) => Ok(PyBytes::new(py, &in_out_buffer[0..plaintext_len])), Err(_) => Err(CryptoError::new_err("decryption failed")), - }; + } } pub fn encrypt<'a>( @@ -200,13 +200,13 @@ impl AeadChaCha20Poly1305 { let cipher: ChaCha20Poly1305 = ChaCha20Poly1305::new(ChaCha20Key::from_slice(&self.key)); let res = cipher.encrypt_in_place( nonce.as_bytes().into(), - &associated_data.as_bytes(), + associated_data.as_bytes(), &mut in_out_buffer, ); - return match res { + match res { Ok(_) => Ok(PyBytes::new(py, &in_out_buffer)), Err(_) => Err(CryptoError::new_err("encryption failed")), - }; + } } } diff --git a/src/agreement.rs b/src/agreement.rs index e70165112..a37c3eb93 100644 --- a/src/agreement.rs +++ b/src/agreement.rs @@ -78,7 +78,7 @@ impl X25519Kyber768Draft00KeyExchange { .extend_from_slice(self.x25519_private.compute_public_key().unwrap().as_ref()); combined_pub_key.extend_from_slice(kyber_pub.key_bytes().unwrap().as_ref()); - return PyBytes::new(py, &combined_pub_key.as_ref()); + PyBytes::new(py, combined_pub_key.as_ref()) } pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { @@ -96,7 +96,7 @@ impl X25519Kyber768Draft00KeyExchange { &self.x25519_private, &x25519_peer_public_key, error::Unspecified, - |_key_material| return Ok(_key_material.to_vec()), + |_key_material| Ok(_key_material.to_vec()), ) .expect("FAILURE"); @@ -112,7 +112,7 @@ impl X25519Kyber768Draft00KeyExchange { let key_material = SharedSecret::from(&combined_secret.0[..]); - return PyBytes::new(py, &key_material.secret_bytes()); + PyBytes::new(py, key_material.secret_bytes()) } } @@ -128,7 +128,7 @@ impl X25519KeyExchange { pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { let my_public_key = self.private.compute_public_key().unwrap(); - return PyBytes::new(py, &my_public_key.as_ref()); + PyBytes::new(py, my_public_key.as_ref()) } pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { @@ -139,11 +139,11 @@ impl X25519KeyExchange { &self.private, &peer_public_key, error::Unspecified, - |_key_material| return Ok(_key_material.to_vec()), + |_key_material| Ok(_key_material.to_vec()), ) .expect("FAILURE"); - return PyBytes::new(py, &key_material); + PyBytes::new(py, &key_material) } } @@ -159,7 +159,7 @@ impl ECDHP256KeyExchange { pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { let my_public_key = self.private.compute_public_key().unwrap(); - return PyBytes::new(py, &my_public_key.as_ref()); + PyBytes::new(py, my_public_key.as_ref()) } pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { @@ -170,13 +170,11 @@ impl ECDHP256KeyExchange { &self.private, &peer_public_key, error::Unspecified, - |_key_material| { - return Ok(_key_material.to_vec()); - }, + |_key_material| Ok(_key_material.to_vec()), ) .expect("FAILURE"); - return PyBytes::new(py, &key_material); + PyBytes::new(py, &key_material) } } @@ -192,7 +190,7 @@ impl ECDHP384KeyExchange { pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { let my_public_key = self.private.compute_public_key().unwrap(); - return PyBytes::new(py, &my_public_key.as_ref()); + PyBytes::new(py, my_public_key.as_ref()) } pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { @@ -203,13 +201,11 @@ impl ECDHP384KeyExchange { &self.private, &peer_public_key, error::Unspecified, - |_key_material| { - return Ok(_key_material.to_vec()); - }, + |_key_material| Ok(_key_material.to_vec()), ) .expect("FAILURE"); - return PyBytes::new(py, &key_material); + PyBytes::new(py, &key_material) } } @@ -225,7 +221,7 @@ impl ECDHP521KeyExchange { pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { let my_public_key = self.private.compute_public_key().unwrap(); - return PyBytes::new(py, &my_public_key.as_ref()); + PyBytes::new(py, my_public_key.as_ref()) } pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { @@ -236,12 +232,10 @@ impl ECDHP521KeyExchange { &self.private, &peer_public_key, error::Unspecified, - |_key_material| { - return Ok(_key_material.to_vec()); - }, + |_key_material| Ok(_key_material.to_vec()), ) .expect("FAILURE"); - return PyBytes::new(py, &key_material); + PyBytes::new(py, &key_material) } } diff --git a/src/buffer.rs b/src/buffer.rs index 349e3c51f..84d3b8c1f 100644 --- a/src/buffer.rs +++ b/src/buffer.rs @@ -26,22 +26,22 @@ impl Buffer { }); } - if !capacity.is_some() { + if capacity.is_none() { return Err(PyValueError::new_err( "mandatory capacity without data args", )); } - return Ok(Buffer { + Ok(Buffer { pos: 0, data: vec![0; capacity.unwrap().try_into().unwrap()], capacity: capacity.unwrap(), - }); + }) } #[getter] pub fn capacity(&self) -> u64 { - return self.capacity; + self.capacity } #[getter] @@ -49,7 +49,7 @@ impl Buffer { if self.pos == 0 { return PyBytes::new(py, &[]); } - return PyBytes::new(py, &self.data[0 as usize..self.pos as usize]); + PyBytes::new(py, &self.data[0_usize..self.pos as usize]) } pub fn data_slice<'a>(&self, py: Python<'a>, start: u64, end: u64) -> PyResult<&'a PyBytes> { @@ -57,11 +57,11 @@ impl Buffer { return Err(BufferReadError::new_err("Read out of bounds")); } - return Ok(PyBytes::new(py, &self.data[start as usize..end as usize])); + Ok(PyBytes::new(py, &self.data[start as usize..end as usize])) } pub fn eof(&self) -> bool { - return self.pos == self.capacity; + self.pos == self.capacity } pub fn seek(&mut self, pos: u64) -> PyResult<()> { @@ -71,11 +71,11 @@ impl Buffer { self.pos = pos; - return Ok(()); + Ok(()) } pub fn tell(&self) -> u64 { - return self.pos; + self.pos } pub fn pull_bytes<'a>(&mut self, py: Python<'a>, length: u64) -> PyResult<&'a PyBytes> { @@ -90,7 +90,7 @@ impl Buffer { self.pos += length; - return Ok(extract); + Ok(extract) } pub fn pull_uint8(&mut self) -> PyResult { @@ -101,7 +101,7 @@ impl Buffer { let extract = self.data[self.pos as usize]; self.pos += 1; - return Ok(extract); + Ok(extract) } pub fn pull_uint16(&mut self) -> PyResult { @@ -120,7 +120,7 @@ impl Buffer { ); self.pos += 2; - return Ok(extract); + Ok(extract) } pub fn pull_uint32(&mut self) -> PyResult { @@ -139,7 +139,7 @@ impl Buffer { ); self.pos += 4; - return Ok(extract); + Ok(extract) } pub fn pull_uint64(&mut self) -> PyResult { @@ -158,7 +158,7 @@ impl Buffer { ); self.pos += 8; - return Ok(extract); + Ok(extract) } pub fn pull_uint_var(&mut self) -> PyResult { @@ -192,12 +192,10 @@ impl Buffer { }; } - return match self.pull_uint64() { - Ok(val) => { - return Ok(val & 0x3FFFFFFFFFFFFFFF); - } + match self.pull_uint64() { + Ok(val) => Ok(val & 0x3FFFFFFFFFFFFFFF), Err(exception) => Err(exception), - }; + } } pub fn push_bytes(&mut self, data: &PyBytes) -> PyResult<()> { @@ -208,10 +206,10 @@ impl Buffer { return Err(BufferWriteError::new_err("Write out of bounds")); } - self.data[self.pos as usize..end_pos as usize].clone_from_slice(&data_to_be_pushed); + self.data[self.pos as usize..end_pos as usize].clone_from_slice(data_to_be_pushed); self.pos = end_pos; - return Ok(()); + Ok(()) } pub fn push_uint8(&mut self, value: u8) -> PyResult<()> { @@ -222,7 +220,7 @@ impl Buffer { self.data[self.pos as usize] = value; self.pos += 1; - return Ok(()); + Ok(()) } pub fn push_uint16(&mut self, value: u16) -> PyResult<()> { @@ -238,7 +236,7 @@ impl Buffer { .clone_from_slice(&value.to_be_bytes()); self.pos += 2; - return Ok(()); + Ok(()) } pub fn push_uint32(&mut self, value: u32) -> PyResult<()> { @@ -254,7 +252,7 @@ impl Buffer { .clone_from_slice(&value.to_be_bytes()); self.pos += 4; - return Ok(()); + Ok(()) } pub fn push_uint64(&mut self, value: u64) -> PyResult<()> { @@ -270,7 +268,7 @@ impl Buffer { .clone_from_slice(&value.to_be_bytes()); self.pos += 8; - return Ok(()); + Ok(()) } pub fn push_uint_var(&mut self, value: u64) -> PyResult<()> { @@ -284,8 +282,8 @@ impl Buffer { return self.push_uint64(value | 0xC000000000000000); } - return Err(PyValueError::new_err( + Err(PyValueError::new_err( "Integer is too big for a variable-length integer", - )); + )) } } diff --git a/src/certificate.rs b/src/certificate.rs index ec52d3f74..1019d442d 100644 --- a/src/certificate.rs +++ b/src/certificate.rs @@ -103,7 +103,7 @@ impl Certificate { } } - return Ok(Certificate { + Ok(Certificate { version: match cert.version() { X509Version::V1 => 0, X509Version::V2 => 1, @@ -114,16 +114,16 @@ impl Certificate { raw_serial_number: cert.raw_serial().to_vec(), not_valid_before: cert.validity.not_before.timestamp(), not_valid_after: cert.validity.not_after.timestamp(), - extensions: extensions, - subject: subject, - issuer: issuer, + extensions, + subject, + issuer, public_bytes: certificate_der.as_bytes().to_vec(), public_key: match cert.public_key().parsed() { Ok(PublicKey::EC(pts)) => pts.data().to_vec(), Ok(PublicKey::DSA(cert_decoded)) => cert_decoded.to_vec(), _ => cert.public_key().raw.to_vec(), }, - }); + }) } _ => Err(CryptoError::new_err("x509 parsing failed")), } @@ -131,26 +131,26 @@ impl Certificate { #[getter] pub fn serial_number(&self) -> &String { - return &self.serial_number; + &self.serial_number } pub fn raw_serial_number<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new(py, &self.raw_serial_number); + PyBytes::new(py, &self.raw_serial_number) } #[getter] pub fn not_valid_before(&self) -> i64 { - return self.not_valid_before; + self.not_valid_before } #[getter] pub fn not_valid_after(&self) -> i64 { - return self.not_valid_after; + self.not_valid_after } #[getter] pub fn version(&self) -> u8 { - return self.version; + self.version } #[getter] @@ -181,7 +181,7 @@ impl Certificate { )); } - return values; + values } #[getter] @@ -212,7 +212,7 @@ impl Certificate { )); } - return values; + values } pub fn get_subject_alt_names<'a>(&self, py: Python<'a>) -> &'a PyList { @@ -224,7 +224,7 @@ impl Certificate { } } - return values; + values } pub fn get_ocsp_endpoints<'a>(&self, py: Python<'a>) -> &'a PyList { @@ -236,7 +236,7 @@ impl Certificate { } } - return values; + values } pub fn get_issuer_endpoints<'a>(&self, py: Python<'a>) -> &'a PyList { @@ -248,15 +248,15 @@ impl Certificate { } } - return values; + values } pub fn public_bytes<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new(py, &self.public_bytes); + PyBytes::new(py, &self.public_bytes) } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new(py, &self.public_key); + PyBytes::new(py, &self.public_key) } fn __eq__(&self, other: &Self) -> bool { @@ -306,7 +306,7 @@ impl ServerVerifier { let parsed_name_res = ServerName::try_from(server_name); - return match parsed_name_res { + match parsed_name_res { Ok(parsed_name) => { let res = self.inner.verify_server_cert( &peer_der, @@ -316,7 +316,7 @@ impl ServerVerifier { UnixTime::now(), ); - return match res { + match res { Ok(_) => Ok(()), Err(Error::InvalidCertificate(err)) => match err { CertificateError::UnknownIssuer => { @@ -342,13 +342,11 @@ impl ServerVerifier { Err(_) => Err(CryptoError::new_err( "the x509 certificate store encountered an error", )), - }; - } - Err(_) => { - return Err(InvalidNameCertificateError::new_err( - "unparseable server name", - )); + } } - }; + Err(_) => Err(InvalidNameCertificateError::new_err( + "unparseable server name", + )), + } } } diff --git a/src/headers.rs b/src/headers.rs index 940ccff56..eb77a2a99 100644 --- a/src/headers.rs +++ b/src/headers.rs @@ -46,21 +46,17 @@ impl QpackEncoder { .configure(max_table_capacity, dyn_table_capacity, blocked_streams) .expect("FAILURE"); - return Ok(PyBytes::new(py, r.data())); + Ok(PyBytes::new(py, r.data())) } pub fn feed_decoder(&mut self, data: &PyBytes) -> PyResult<()> { let res = self.encoder.feed(data.as_bytes()); match res { - Ok(_) => { - return Ok(()); - } - Err(_) => { - return Err(DecoderStreamError::new_err( - "an error occurred while feeding data from decoder with qpack data", - )); - } + Ok(_) => Ok(()), + Err(_) => Err(DecoderStreamError::new_err( + "an error occurred while feeding data from decoder with qpack data", + )), } } @@ -89,14 +85,12 @@ impl QpackEncoder { let stream_data = PyBytes::new(py, buffer.stream()); - return Ok(PyTuple::new(py, [stream_data, encoded_buffer])); - } - Err(abc) => { - return Err(EncoderStreamError::new_err(format!( - "unable to encode headers {:?}", - abc - ))); + Ok(PyTuple::new(py, [stream_data, encoded_buffer])) } + Err(abc) => Err(EncoderStreamError::new_err(format!( + "unable to encode headers {:?}", + abc + ))), } } } @@ -114,14 +108,10 @@ impl QpackDecoder { let res = self.decoder.feed(data.as_bytes()); match res { - Ok(_) => { - return Ok(()); - } - Err(_) => { - return Err(EncoderStreamError::new_err( - "an error occurred while feeding data from encoder with qpack data", - )); - } + Ok(_) => Ok(()), + Err(_) => Err(EncoderStreamError::new_err( + "an error occurred while feeding data from encoder with qpack data", + )), } } @@ -149,31 +139,27 @@ impl QpackDecoder { )); } - return Ok(PyTuple::new( + Ok(PyTuple::new( py, [ PyBytes::new(py, buffer.stream()).to_object(py), decoded_headers.to_object(py), ], - )); - } - Ok(DecoderOutput::BlockedStream) => { - return Err(StreamBlocked::new_err( - "stream is blocked, need more data to pursue decoding", - )); - } - Err(_) => { - return Err(DecoderStreamError::new_err( - "an error occurred while decoding the stream qpack data", - )); + )) } + Ok(DecoderOutput::BlockedStream) => Err(StreamBlocked::new_err( + "stream is blocked, need more data to pursue decoding", + )), + Err(_) => Err(DecoderStreamError::new_err( + "an error occurred while decoding the stream qpack data", + )), } } pub fn resume_header<'a>(&mut self, py: Python<'a>, stream_id: u64) -> PyResult<&'a PyTuple> { let output = self.decoder.unblocked(StreamId::new(stream_id)); - if !output.is_some() { + if output.is_none() { return Err(DecoderStreamError::new_err("stream id is unknown")); } @@ -193,24 +179,20 @@ impl QpackDecoder { )); } - return Ok(PyTuple::new( + Ok(PyTuple::new( py, [ PyBytes::new(py, buffer.stream()).to_object(py), decoded_headers.to_object(py), ], - )); - } - Ok(DecoderOutput::BlockedStream) => { - return Err(StreamBlocked::new_err( - "stream is blocked, need more data to pursue decoding", - )) - } - Err(_) => { - return Err(DecoderStreamError::new_err( - "an error occurred while decoding the stream qpack data", )) } + Ok(DecoderOutput::BlockedStream) => Err(StreamBlocked::new_err( + "stream is blocked, need more data to pursue decoding", + )), + Err(_) => Err(DecoderStreamError::new_err( + "an error occurred while decoding the stream qpack data", + )), } } } diff --git a/src/hpk.rs b/src/hpk.rs index 7168245f8..bcd4d58a4 100644 --- a/src/hpk.rs +++ b/src/hpk.rs @@ -23,20 +23,20 @@ impl QUICHeaderProtection { 20 => &CHACHA20, _ => panic!("unsupported"), }, - &key.as_bytes(), + key.as_bytes(), ) .expect("FAILURE"), } } pub fn mask<'a>(&self, py: Python<'a>, sample: &PyBytes) -> PyResult<&'a PyBytes> { - let res = self.hpk.new_mask(&sample.as_bytes()); + let res = self.hpk.new_mask(sample.as_bytes()); - return match res { + match res { Err(_) => Err(CryptoError::new_err( "unable to issue mask protection header", )), Ok(data) => Ok(PyBytes::new(py, &data)), - }; + } } } diff --git a/src/ocsp.rs b/src/ocsp.rs index 5336f76a9..2fcfc3149 100644 --- a/src/ocsp.rs +++ b/src/ocsp.rs @@ -67,23 +67,23 @@ pub struct OCSPResponse { impl OCSPResponse { #[new] pub fn py_new(raw_response: &PyBytes) -> PyResult { - let ocsp_res: OcspResponse = OcspResponse::from_der(&raw_response.as_bytes()).unwrap(); + let ocsp_res: OcspResponse = OcspResponse::from_der(raw_response.as_bytes()).unwrap(); - if !ocsp_res.response_bytes.is_some() { + if ocsp_res.response_bytes.is_none() { return Err(PyValueError::new_err("OCSP Server did not provide answers")); } let inner_resp: BasicOcspResponse = - BasicOcspResponse::from_der(&ocsp_res.response_bytes.unwrap().response.as_bytes()) + BasicOcspResponse::from_der(ocsp_res.response_bytes.unwrap().response.as_bytes()) .unwrap(); - if inner_resp.tbs_response_data.responses.len() == 0 { + if inner_resp.tbs_response_data.responses.is_empty() { return Err(PyValueError::new_err("OCSP Server did not provide answers")); } let first_resp_for_cert: &SingleResponse = &inner_resp.tbs_response_data.responses[0]; - return Ok(OCSPResponse { + Ok(OCSPResponse { next_update: first_resp_for_cert .next_update .unwrap() @@ -124,27 +124,27 @@ impl OCSPResponse { }, InternalCertStatus::Good(_) | InternalCertStatus::Unknown(_) => None, }, - }); + }) } #[getter] pub fn next_update(&self) -> u64 { - return self.next_update; + self.next_update } #[getter] pub fn response_status(&self) -> OCSPResponseStatus { - return self.response_status; + self.response_status } #[getter] pub fn certificate_status(&self) -> OCSPCertStatus { - return self.certificate_status; + self.certificate_status } #[getter] pub fn revocation_reason(&self) -> Option { - return self.revocation_reason; + self.revocation_reason } } @@ -157,22 +157,22 @@ pub struct OCSPRequest { impl OCSPRequest { #[new] pub fn py_new(peer_certificate: &PyBytes, issuer_certificate: &PyBytes) -> PyResult { - let issuer = Certificate::from_der(&issuer_certificate.as_bytes()).unwrap(); - let cert = Certificate::from_der(&peer_certificate.as_bytes()).unwrap(); + let issuer = Certificate::from_der(issuer_certificate.as_bytes()).unwrap(); + let cert = Certificate::from_der(peer_certificate.as_bytes()).unwrap(); let req: InternalOcspRequest = OcspRequestBuilder::default() .with_request(Request::from_cert::(&issuer, &cert).unwrap()) .build(); - return match req.to_der() { + match req.to_der() { Ok(raw_der) => Ok(OCSPRequest { inner_request: raw_der, }), Err(_) => Err(PyValueError::new_err("unable to generate the request")), - }; + } } pub fn public_bytes<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new(py, &self.inner_request); + PyBytes::new(py, &self.inner_request) } } diff --git a/src/pkcs8.rs b/src/pkcs8.rs index 952e6edc9..f7a7307a8 100644 --- a/src/pkcs8.rs +++ b/src/pkcs8.rs @@ -57,10 +57,10 @@ impl TryFrom> for PrivateKeyInfo { }); } - return Ok(PrivateKeyInfo { + Ok(PrivateKeyInfo { der_encoded: der_document.clone(), cert_type: KeyType::ED25519, - }); + }) } } @@ -72,7 +72,7 @@ impl PrivateKeyInfo { let decoded_bytes = std::str::from_utf8(pem_content).unwrap(); let is_encrypted = decoded_bytes.contains("ENCRYPTED"); - let item = read_one_from_slice(&pem_content); + let item = read_one_from_slice(pem_content); match item.unwrap().unwrap().0 { Item::Pkcs1Key(key) => { @@ -81,24 +81,24 @@ impl PrivateKeyInfo { } let rsa_key: RsaPrivateKey = - RsaPrivateKey::from_pkcs1_der(&key.secret_pkcs1_der()).unwrap(); + RsaPrivateKey::from_pkcs1_der(key.secret_pkcs1_der()).unwrap(); let pkcs8_pem = rsa_key.to_pkcs8_pem(LineEnding::LF).expect("FAILURE"); let pkcs8_pem: &str = pkcs8_pem.as_ref(); - return PrivateKeyInfo::from_pkcs8_pem(&pkcs8_pem).unwrap(); + PrivateKeyInfo::from_pkcs8_pem(pkcs8_pem).unwrap() } Item::Pkcs8Key(_key) => { if is_encrypted { return PrivateKeyInfo::from_pkcs8_encrypted_pem( - &decoded_bytes, + decoded_bytes, password.unwrap().as_bytes(), ) .unwrap(); } - return PrivateKeyInfo::from_pkcs8_pem(&decoded_bytes).unwrap(); + PrivateKeyInfo::from_pkcs8_pem(decoded_bytes).unwrap() } Item::Sec1Key(key) => { if is_encrypted { @@ -107,7 +107,7 @@ impl PrivateKeyInfo { let sec1_der = key.secret_sec1_der().to_vec(); - return PrivateKeyInfo { + PrivateKeyInfo { cert_type: match sec1_der.len() { 32..=121 => KeyType::ECDSA_P256, 132..=167 => KeyType::ECDSA_P384, @@ -115,17 +115,17 @@ impl PrivateKeyInfo { _ => panic!("unsupported sec1 key"), }, der_encoded: sec1_der, - }; + } } _ => panic!("unsupported"), - }; + } } pub fn get_type(&self) -> KeyType { - return self.cert_type; + self.cert_type } pub fn public_bytes<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new(py, &self.der_encoded); + PyBytes::new(py, &self.der_encoded) } } diff --git a/src/private_key.rs b/src/private_key.rs index 112fac7d4..75504a918 100644 --- a/src/private_key.rs +++ b/src/private_key.rs @@ -61,18 +61,18 @@ impl Ed25519PrivateKey { #[new] pub fn py_new(pkcs8: &PyBytes) -> Self { Ed25519PrivateKey { - inner: InternalEd25519PrivateKey::from_pkcs8(&pkcs8.as_bytes()).expect("FAILURE"), + inner: InternalEd25519PrivateKey::from_pkcs8(pkcs8.as_bytes()).expect("FAILURE"), } } pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { - let signature = self.inner.sign(&data.as_bytes()); + let signature = self.inner.sign(data.as_bytes()); - return PyBytes::new(py, &signature.as_ref()); + PyBytes::new(py, signature.as_ref()) } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new(py, &self.inner.public_key().as_ref()); + PyBytes::new(py, self.inner.public_key().as_ref()) } } @@ -87,27 +87,27 @@ impl EcPrivateKey { _ => panic!("unsupported"), }; - return EcPrivateKey { - inner: InternalEcPrivateKey::from_pkcs8(&signing_algorithm, &pkcs8.as_bytes()) + EcPrivateKey { + inner: InternalEcPrivateKey::from_pkcs8(signing_algorithm, pkcs8.as_bytes()) .expect("FAILURE"), curve: curve_type, - }; + } } pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { let rng = SystemRandom::new(); - let signature = self.inner.sign(&rng, &data.as_bytes()); + let signature = self.inner.sign(&rng, data.as_bytes()); - return PyBytes::new(py, &signature.unwrap().as_ref()); + PyBytes::new(py, signature.unwrap().as_ref()) } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new(py, &self.inner.public_key().as_ref()); + PyBytes::new(py, self.inner.public_key().as_ref()) } #[getter] pub fn curve_type(&self) -> u32 { - return self.curve; + self.curve } } @@ -115,27 +115,26 @@ impl EcPrivateKey { impl DsaPrivateKey { #[new] pub fn py_new(pkcs8: &PyBytes) -> Self { - return DsaPrivateKey { - inner: InternalDsaPrivateKey::from_pkcs8_der(&pkcs8.as_bytes()).expect("FAILURE"), - }; + DsaPrivateKey { + inner: InternalDsaPrivateKey::from_pkcs8_der(pkcs8.as_bytes()).expect("FAILURE"), + } } pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { - let signature = self.inner.sign(&data.as_bytes()); + let signature = self.inner.sign(data.as_bytes()); - return PyBytes::new(py, &signature.to_bytes()); + PyBytes::new(py, &signature.to_bytes()) } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - return PyBytes::new( + PyBytes::new( py, - &self - .inner + self.inner .verifying_key() .to_public_key_der() .expect("FAILURE") .as_bytes(), - ); + ) } } @@ -143,9 +142,9 @@ impl DsaPrivateKey { impl RsaPrivateKey { #[new] pub fn py_new(pkcs8: &PyBytes) -> Self { - return RsaPrivateKey { - inner: InternalRsaPrivateKey::from_pkcs8_der(&pkcs8.as_bytes()).expect("FAILURE"), - }; + RsaPrivateKey { + inner: InternalRsaPrivateKey::from_pkcs8_der(pkcs8.as_bytes()).expect("FAILURE"), + } } pub fn sign<'a>( @@ -161,43 +160,43 @@ impl RsaPrivateKey { true => match hash_size { 256 => { let signer = InternalRsaPssSigningKey::::new(private_key); - return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) } 384 => { let signer = InternalRsaPssSigningKey::::new(private_key); - return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) } 512 => { let signer = InternalRsaPssSigningKey::::new(private_key); - return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) } _ => panic!("unsupported"), }, false => match hash_size { 256 => { let signer = InternalRsaPkcsSigningKey::::new(private_key); - return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) } 384 => { let signer = InternalRsaPkcsSigningKey::::new(private_key); - return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) } 512 => { let signer = InternalRsaPkcsSigningKey::::new(private_key); - return PyBytes::new(py, &signer.sign(&data.as_bytes()).to_vec()); + PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) } _ => panic!("unsupported"), }, - }; + } } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { let public_key: InternalRsaPublicKey = self.inner.to_public_key(); - return PyBytes::new( + PyBytes::new( py, &public_key.to_public_key_der().as_ref().unwrap().to_vec(), - ); + ) } } @@ -221,14 +220,14 @@ pub fn verify_with_public_key( || pkcs115_blind_signature.contains(&algorithm) { let rsa_parsed_public_key = - InternalRsaPublicKey::from_public_key_der(&public_key_bytes).expect("FAILURE"); + InternalRsaPublicKey::from_public_key_der(public_key_bytes).expect("FAILURE"); return match algorithm { 0x0804 | 0x0809 => { let alt_verifier = RsaPssVerifyingKey::::new(rsa_parsed_public_key); let res = alt_verifier.verify( - &message.as_bytes(), + message.as_bytes(), &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE"), ); @@ -241,7 +240,7 @@ pub fn verify_with_public_key( let alt_verifier = RsaPssVerifyingKey::::new(rsa_parsed_public_key); let res = alt_verifier.verify( - &message.as_bytes(), + message.as_bytes(), &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE"), ); @@ -254,7 +253,7 @@ pub fn verify_with_public_key( let alt_verifier = RsaPssVerifyingKey::::new(rsa_parsed_public_key); let res = alt_verifier.verify( - &message.as_bytes(), + message.as_bytes(), &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE"), ); @@ -268,7 +267,7 @@ pub fn verify_with_public_key( let alt_verifier = RsaPkcsVerifyingKey::::new(rsa_parsed_public_key); let res = alt_verifier.verify( - &message.as_bytes(), + message.as_bytes(), &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE"), ); @@ -281,7 +280,7 @@ pub fn verify_with_public_key( let alt_verifier = RsaPkcsVerifyingKey::::new(rsa_parsed_public_key); let res = alt_verifier.verify( - &message.as_bytes(), + message.as_bytes(), &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE"), ); @@ -294,7 +293,7 @@ pub fn verify_with_public_key( let alt_verifier = RsaPkcsVerifyingKey::::new(rsa_parsed_public_key); let res = alt_verifier.verify( - &message.as_bytes(), + message.as_bytes(), &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE"), ); @@ -310,9 +309,9 @@ pub fn verify_with_public_key( if algorithm == 0x0807 { let ed25519_verifier: Ed25519VerifyingKey = - Ed25519VerifyingKey::from_public_key_der(&public_key_bytes).expect("FAILURE"); + Ed25519VerifyingKey::from_public_key_der(public_key_bytes).expect("FAILURE"); let res = ed25519_verifier.verify( - &message.as_bytes(), + message.as_bytes(), &Ed25519Signature::from_bytes(signature.as_bytes()[0..64].try_into().unwrap()), ); @@ -332,10 +331,10 @@ pub fn verify_with_public_key( public_key_bytes, ); - let res = public_key.verify(&message.as_bytes(), &signature.as_bytes()); + let res = public_key.verify(message.as_bytes(), signature.as_bytes()); - return match res { + match res { Err(Unspecified) => Err(SignatureError::new_err("signature mismatch (ecdsa)")), _ => Ok(()), - }; + } } diff --git a/src/rsa.rs b/src/rsa.rs index 32f9d16ad..af9178161 100644 --- a/src/rsa.rs +++ b/src/rsa.rs @@ -21,8 +21,8 @@ impl Rsa { let public_key = RsaPublicKey::from(&private_key); Rsa { - public_key: public_key, - private_key: private_key, + public_key, + private_key, } } @@ -34,10 +34,10 @@ impl Rsa { let enc_data = self .public_key - .encrypt(&mut rng, padding, &payload_to_enc[..]) + .encrypt(&mut rng, padding, payload_to_enc) .expect("failed to encrypt"); - return PyBytes::new(py, &enc_data); + PyBytes::new(py, &enc_data) } pub fn decrypt<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { @@ -46,9 +46,9 @@ impl Rsa { let padding = Oaep::new::(); let dec_data = self .private_key - .decrypt(padding, &payload_to_dec) + .decrypt(padding, payload_to_dec) .expect("failed to decrypt"); - return PyBytes::new(py, &dec_data); + PyBytes::new(py, &dec_data) } } From d1c05a580b491403ef64162c029f50d506d07a91 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Sun, 29 Dec 2024 16:24:03 +0100 Subject: [PATCH 08/39] :sparkle: Migrate Kyber 768 Draft to standard Module-Lattice 768 --- Cargo.lock | 161 ++++++++++++++++++++++++----------------------- qh3/_hazmat.pyi | 2 +- qh3/tls.py | 23 +++---- src/agreement.rs | 18 +++--- src/lib.rs | 4 +- 5 files changed, 106 insertions(+), 102 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 26795cda7..095288254 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -79,9 +79,9 @@ checksum = "ace50bade8e6234aa140d9a2f552bbee1db4d353f69b8217bc503490fc1a9f26" [[package]] name = "aws-lc-fips-sys" -version = "0.12.13" +version = "0.13.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bf12b67bc9c5168f68655aadb2a12081689a58f1d9b1484705e4d1810ed6e4ac" +checksum = "59057b878509d88952425fe694a2806e468612bde2d71943f3cd8034935b5032" dependencies = [ "bindgen 0.69.5", "cc", @@ -90,26 +90,26 @@ dependencies = [ "fs_extra", "libc", "paste", + "regex", ] [[package]] name = "aws-lc-rs" -version = "1.10.0" +version = "1.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cdd82dba44d209fddb11c190e0a94b78651f95299598e472215667417a03ff1d" +checksum = "f409eb70b561706bf8abba8ca9c112729c481595893fd06a2dd9af8ed8441148" dependencies = [ "aws-lc-fips-sys", "aws-lc-sys", - "mirai-annotations", "paste", "zeroize", ] [[package]] name = "aws-lc-sys" -version = "0.22.0" +version = "0.24.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "df7a4168111d7eb622a31b214057b8509c0a7e1794f44c546d742330dc793972" +checksum = "8478a5c29ead3f3be14aff8a202ad965cf7da6856860041bfca271becf8ba48b" dependencies = [ "bindgen 0.69.5", "cc", @@ -213,9 +213,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.1.30" +version = "1.2.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b16803a61b81d9eabb7eae2588776c4c1e584b738ede45fdbb4c972cec1e9945" +checksum = "8d6dbb628b8f8555f86d0323c2eb39e3ec81901f4b83e091db8a6a76d316a333" dependencies = [ "jobserver", "libc", @@ -285,9 +285,9 @@ dependencies = [ [[package]] name = "cmake" -version = "0.1.51" +version = "0.1.52" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fb1e43aa7fd152b1f968787f7dbcdeb306d1867ff373c69955211876c053f91a" +checksum = "c682c223677e0e5b6b7f63a64b9351844c3f1b1678a68b7ee617e30fb082620e" dependencies = [ "cc", ] @@ -300,9 +300,9 @@ checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8" [[package]] name = "cpufeatures" -version = "0.2.14" +version = "0.2.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "608697df725056feaccfa42cffdaeeec3fccc4ffc38358ecd19b243e716a78e0" +checksum = "16b80225097f2e5ae4e7179dd2266824648f3e2f49d9134d584b76389d31c4c3" dependencies = [ "libc", ] @@ -475,12 +475,12 @@ checksum = "60b1af1c220855b6ceac025d3f6ecdd2b7c4894bfe9cd9bda4fbb4bc7c0d4cf0" [[package]] name = "errno" -version = "0.3.9" +version = "0.3.10" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "534c5cf6194dfab3db3242765c03bbe257cf92f22b38f6bc0c58d59108a820ba" +checksum = "33d852cb9b869c2a9b3df2f71a3074817f01e1844f839a144f5fcef059a4eb5d" dependencies = [ "libc", - "windows-sys", + "windows-sys 0.59.0", ] [[package]] @@ -524,9 +524,9 @@ dependencies = [ [[package]] name = "glob" -version = "0.3.1" +version = "0.3.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d2fabcfbdc87f4758337ca535fb41a6d701b65693ce38287d856d1674551ec9b" +checksum = "a8d1add55171497b4705a648c6b583acafb01d58050a51727785f0b2c8e0a2b2" [[package]] name = "heck" @@ -545,11 +545,11 @@ dependencies = [ [[package]] name = "home" -version = "0.5.9" +version = "0.5.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e3d1354bf6b7235cb4a0576c2619fd4ed18183f689b12b006a0ee7329eeff9a5" +checksum = "589533453244b0995c858700322199b2becb13b627df2851f64a2775d024abcf" dependencies = [ - "windows-sys", + "windows-sys 0.59.0", ] [[package]] @@ -579,9 +579,9 @@ dependencies = [ [[package]] name = "itoa" -version = "1.0.11" +version = "1.0.14" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "49f1f14873335454500d59611f1cf4a4b0f786f9ac11f4312a78e4cf2566695b" +checksum = "d75a2a4b1b190afb6f5425f10f6a8f959d2ea0b9c2b1d79553551850539e4674" [[package]] name = "jobserver" @@ -609,15 +609,15 @@ checksum = "830d08ce1d1d941e6b30645f1a0eb5643013d835ce3779a5fc208261dbe10f55" [[package]] name = "libc" -version = "0.2.159" +version = "0.2.169" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "561d97a539a36e26a9a5fad1ea11a3039a67714694aaa379433e580854bc3dc5" +checksum = "b5aba8db14291edd000dfcc4d620c7ebfb122c613afb886ca8803fa4e128a20a" [[package]] name = "libloading" -version = "0.8.5" +version = "0.8.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4979f22fdb869068da03c9f7528f8297c6fd2606bc3a4affe42e6a823fdb8da4" +checksum = "fc2f4eb4bc735547cfed7c0a4922cbd04a4655978c09b54f1f7b228750664c34" dependencies = [ "cfg-if", "windows-targets", @@ -625,9 +625,9 @@ dependencies = [ [[package]] name = "libm" -version = "0.2.8" +version = "0.2.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4ec2a862134d2a7d32d7983ddcdd1c4923530833c9f2ea1a44fc5fa473989058" +checksum = "8355be11b20d696c8f18f6cc018c4e372165b1fa8126cef092399c9951984ffa" [[package]] name = "linux-raw-sys" @@ -691,12 +691,6 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" -[[package]] -name = "mirai-annotations" -version = "1.12.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c9be0862c1b3f26a88803c4a49de6889c10e608b3ee9344e6ef5b45fb37ad3d1" - [[package]] name = "nom" version = "7.1.3" @@ -896,9 +890,9 @@ dependencies = [ [[package]] name = "portable-atomic" -version = "1.9.0" +version = "1.10.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cc9c68a3f6da06753e9335d63e27f6b9754dd1920d941135b7ea8224f141adb2" +checksum = "280dc24453071f1b63954171985a0b0d30058d287960968b9b2aca264c8d4ee6" [[package]] name = "powerfmt" @@ -917,9 +911,9 @@ dependencies = [ [[package]] name = "prettyplease" -version = "0.2.22" +version = "0.2.25" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "479cf940fbbb3426c32c5d5176f62ad57549a0bb84773423ba8be9d089f5faba" +checksum = "64d1ec885c64d0457d564db4ec299b2dae3f9c02808b8ad9c3a089c591b18033" dependencies = [ "proc-macro2", "syn", @@ -927,9 +921,9 @@ dependencies = [ [[package]] name = "proc-macro2" -version = "1.0.87" +version = "1.0.92" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b3e4daa0dcf6feba26f985457cdf104d4b4256fc5a09547140f3631bb076b19a" +checksum = "37d3544b3f2748c54e147655edb5025752e2303145b5aefb3c3ea2c78b973bb0" dependencies = [ "unicode-ident", ] @@ -1000,9 +994,9 @@ dependencies = [ [[package]] name = "python3-dll-a" -version = "0.2.10" +version = "0.2.12" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bd0b78171a90d808b319acfad166c4790d9e9759bbc14ac8273fe133673dd41b" +checksum = "9b66f9171950e674e64bad3456e11bb3cca108e5c34844383cfe277f45c8a7a8" dependencies = [ "cc", ] @@ -1032,9 +1026,9 @@ dependencies = [ [[package]] name = "quote" -version = "1.0.37" +version = "1.0.38" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b5b9d34b8991d19d98081b46eacdd8eb58c6f2b201139f7c5f643cc155a633af" +checksum = "0e4dccaaaf89514f546c693ddc140f729f958c247918a13380cccc6078391acc" dependencies = [ "proc-macro2", ] @@ -1071,18 +1065,18 @@ dependencies = [ [[package]] name = "redox_syscall" -version = "0.5.7" +version = "0.5.8" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9b6dfecf2c74bce2466cabf93f6664d6998a69eb21e39f4207930065b27b771f" +checksum = "03a862b389f93e68874fbf580b9de08dd02facb9a788ebadaf4a3fd33cf58834" dependencies = [ "bitflags", ] [[package]] name = "regex" -version = "1.11.0" +version = "1.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "38200e5ee88914975b69f657f0801b6f6dccafd44fd9326302a4aaeecfacb1d8" +checksum = "b544ef1b4eac5dc2db33ea63606ae9ffcfac26c1416a2806ae0bf5f56b201191" dependencies = [ "aho-corasick", "memchr", @@ -1092,9 +1086,9 @@ dependencies = [ [[package]] name = "regex-automata" -version = "0.4.8" +version = "0.4.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "368758f23274712b504848e9d5a6f010445cc8b87a7cdb4d7cbee666c1288da3" +checksum = "809e8dc61f6de73b46c85f4c96486310fe304c434cfa43669d7b40f711150908" dependencies = [ "aho-corasick", "memchr", @@ -1129,14 +1123,14 @@ dependencies = [ "libc", "spin", "untrusted", - "windows-sys", + "windows-sys 0.52.0", ] [[package]] name = "rsa" -version = "0.9.6" +version = "0.9.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5d0e5124fcb30e76a7e79bfee683a2746db83784b86289f6251b54b7950a0dfc" +checksum = "47c75d7c5c6b673e58bf54d8544a9f432e3a925b0e80f7cd3602ab5c50c55519" dependencies = [ "const-oid", "digest", @@ -1179,22 +1173,22 @@ dependencies = [ [[package]] name = "rustix" -version = "0.38.37" +version = "0.38.42" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8acb788b847c24f28525660c4d7758620a7210875711f79e7f663cc152726811" +checksum = "f93dc38ecbab2eb790ff964bb77fa94faf256fd3e73285fd7ba0903b76bedb85" dependencies = [ "bitflags", "errno", "libc", "linux-raw-sys", - "windows-sys", + "windows-sys 0.59.0", ] [[package]] name = "rustls" -version = "0.23.14" +version = "0.23.20" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "415d9944693cb90382053259f89fbb077ea730ad7273047ec63b19bc9b160ba8" +checksum = "5065c3f250cbd332cd894be57c40fa52387247659b14a2d6041d121547903b1b" dependencies = [ "aws-lc-rs", "log", @@ -1216,9 +1210,9 @@ dependencies = [ [[package]] name = "rustls-pki-types" -version = "1.10.0" +version = "1.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "16f1201b3c9a7ee8039bcadc17b7e605e2945b27eee7631788c1bd2b0643674b" +checksum = "d2bf47e6ff922db3825eb750c4e2ff784c6ff8fb9e13046ef6a1d1c5401b0b37" [[package]] name = "rustls-webpki" @@ -1260,24 +1254,24 @@ dependencies = [ [[package]] name = "semver" -version = "1.0.23" +version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "61697e0a1c7e512e84a621326239844a24d8207b4669b41bc18b32ea5cbf988b" +checksum = "3cb6eb87a131f756572d7fb904f6e7b68633f09cca868c5df1c4b8d1a694bbba" [[package]] name = "serde" -version = "1.0.210" +version = "1.0.217" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c8e3592472072e6e22e0a54d5904d9febf8508f65fb8552499a1abc7d1078c3a" +checksum = "02fc4265df13d6fa1d00ecff087228cc0a2b5f3c0e87e258d8b94a156e984c70" dependencies = [ "serde_derive", ] [[package]] name = "serde_derive" -version = "1.0.210" +version = "1.0.217" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "243902eda00fad750862fc144cea25caca5e20d615af0a81bee94ca738f1df1f" +checksum = "5a9bf7cf98d04a2b28aead066b7496853d4779c9cc183c440dbac457641e19a0" dependencies = [ "proc-macro2", "quote", @@ -1352,9 +1346,9 @@ checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" [[package]] name = "syn" -version = "2.0.79" +version = "2.0.93" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "89132cd0bf050864e1d38dc3bbc07a0eb8e7530af26344d3d2bbbef83499f590" +checksum = "9c786062daee0d6db1132800e623df74274a0a87322d8e183338e01b3d98d058" dependencies = [ "proc-macro2", "quote", @@ -1380,18 +1374,18 @@ checksum = "61c41af27dd6d1e27b1b16b489db798443478cef1f06a660c96db617ba5de3b1" [[package]] name = "thiserror" -version = "1.0.64" +version = "1.0.69" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d50af8abc119fb8bb6dbabcfa89656f46f84aa0ac7688088608076ad2b459a84" +checksum = "b6aaf5339b578ea85b50e080feb250a3e8ae8cfcdff9a461c9ec2904bc923f52" dependencies = [ "thiserror-impl", ] [[package]] name = "thiserror-impl" -version = "1.0.64" +version = "1.0.69" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "08904e7672f5eb876eaaf87e0ce17857500934f4981c4a0ab2b4aa98baac7fc3" +checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1" dependencies = [ "proc-macro2", "quote", @@ -1400,9 +1394,9 @@ dependencies = [ [[package]] name = "time" -version = "0.3.36" +version = "0.3.37" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5dfd88e563464686c916c7e46e623e520ddc6d79fa6641390f2e3fa86e83e885" +checksum = "35e7868883861bd0e56d9ac6efcaaca0d6d5d82a2a7ec8209ff492c07cf37b21" dependencies = [ "deranged", "itoa", @@ -1421,9 +1415,9 @@ checksum = "ef927ca75afb808a4d64dd374f00a2adf8d0fcff8e7b184af886c3c87ec4a3f3" [[package]] name = "time-macros" -version = "0.2.18" +version = "0.2.19" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3f252a68540fde3a3877aeea552b832b40ab9a69e318efd078774a01ddee1ccf" +checksum = "2834e6017e3e5e4b9834939793b282bc03b37a3336245fa820e35e233e2a85de" dependencies = [ "num-conv", "time-core", @@ -1458,9 +1452,9 @@ checksum = "42ff0bf0c66b8238c6f3b578df37d0b7848e55df8577b3f74f92a69acceeb825" [[package]] name = "unicode-ident" -version = "1.0.13" +version = "1.0.14" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e91b56cd4cadaeb79bbf1a5645f6b4f8dc5bde8834ad5894a8db35fda9efa1fe" +checksum = "adb9e6ca4f869e1180728b7950e35922a7fc6397f7b641499e8f3ef06e50dc83" [[package]] name = "unindent" @@ -1517,6 +1511,15 @@ dependencies = [ "windows-targets", ] +[[package]] +name = "windows-sys" +version = "0.59.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e38bc4d79ed67fd075bcc251a1c39b32a1776bbe92e5bef1f0bf1f8c531853b" +dependencies = [ + "windows-targets", +] + [[package]] name = "windows-targets" version = "0.52.6" diff --git a/qh3/_hazmat.pyi b/qh3/_hazmat.pyi index 624ab37a6..1805684c4 100644 --- a/qh3/_hazmat.pyi +++ b/qh3/_hazmat.pyi @@ -122,7 +122,7 @@ def verify_with_public_key( public_key_raw: bytes, algorithm: int, message: bytes, signature: bytes ) -> None: ... -class X25519Kyber768Draft00KeyExchange: +class X25519ML768KeyExchange: def __init__(self) -> None: ... def public_key(self) -> bytes: ... def exchange(self, peer_public_key: bytes) -> bytes: ... diff --git a/qh3/tls.py b/qh3/tls.py index 0085dedcf..2d00e7468 100644 --- a/qh3/tls.py +++ b/qh3/tls.py @@ -35,7 +35,7 @@ SignatureError, UnacceptableCertificateError, X25519KeyExchange, - X25519Kyber768Draft00KeyExchange, + X25519ML768KeyExchange, verify_with_public_key, ) from .buffer import Buffer @@ -473,6 +473,7 @@ class Group(IntEnum): SECP384R1 = 0x0018 SECP521R1 = 0x0019 X25519KYBER768DRAFT00 = 0x6399 + X25519ML768 = 0x11EC X25519 = 0x001D X448 = 0x001E GREASE = 0xAAAA @@ -1354,7 +1355,7 @@ def __init__( self._supported_groups = [ Group.GREASE, - Group.X25519KYBER768DRAFT00, + Group.X25519ML768, Group.X25519, Group.SECP256R1, Group.SECP384R1, @@ -1384,7 +1385,7 @@ def __init__( self._ec_p384_private_key: ECDHP384KeyExchange | None = None self._ec_p521_private_key: ECDHP521KeyExchange | None = None self._x25519_private_key: X25519KeyExchange | None = None - self._x25519_kyber_768_private_key: X25519Kyber768Draft00KeyExchange | None = ( + self._x25519_kyber_768_private_key: X25519ML768KeyExchange | None = ( None ) @@ -1545,15 +1546,15 @@ def _client_send_hello(self, output_buf: Buffer) -> None: self._x25519_private_key = X25519KeyExchange() key_share.append((Group.X25519, self._x25519_private_key.public_key())) supported_groups.append(Group.X25519) - elif group == Group.X25519KYBER768DRAFT00: - self._x25519_kyber_768_private_key = X25519Kyber768Draft00KeyExchange() + elif group == Group.X25519ML768: + self._x25519_kyber_768_private_key = X25519ML768KeyExchange() key_share.append( ( - Group.X25519KYBER768DRAFT00, + Group.X25519ML768, self._x25519_kyber_768_private_key.public_key(), ) ) - supported_groups.append(Group.X25519KYBER768DRAFT00) + supported_groups.append(Group.X25519ML768) elif group == Group.GREASE: key_share.append((Group.GREASE, b"\x00")) supported_groups.append(Group.GREASE) @@ -1665,7 +1666,7 @@ def _client_handle_hello(self, input_buf: Buffer, output_buf: Buffer) -> None: and self._x25519_private_key is not None ): shared_key = self._x25519_private_key.exchange(peer_public_key) - elif peer_hello.key_share[0] == Group.X25519KYBER768DRAFT00: + elif peer_hello.key_share[0] == Group.X25519ML768: shared_key = self._x25519_kyber_768_private_key.exchange(peer_public_key) elif ( peer_hello.key_share[0] == Group.SECP256R1 @@ -2036,13 +2037,13 @@ def _server_handle_hello( shared_key = self._x25519_private_key.exchange(peer_public_key) group_kx = Group.X25519 break - elif key_share[0] == Group.X25519KYBER768DRAFT00: - self._x25519_kyber_768_private_key = X25519Kyber768Draft00KeyExchange() + elif key_share[0] == Group.X25519ML768: + self._x25519_kyber_768_private_key = X25519ML768KeyExchange() public_key = self._x25519_kyber_768_private_key.public_key() shared_key = self._x25519_kyber_768_private_key.exchange( peer_public_key ) - group_kx = Group.X25519KYBER768DRAFT00 + group_kx = Group.X25519ML768 break elif key_share[0] == Group.SECP256R1: self._ec_p256_private_key = ECDHP256KeyExchange() diff --git a/src/agreement.rs b/src/agreement.rs index a37c3eb93..950026a23 100644 --- a/src/agreement.rs +++ b/src/agreement.rs @@ -1,7 +1,7 @@ use aws_lc_rs::{agreement, error}; use aws_lc_rs::kem; -use aws_lc_rs::unstable::kem::{get_algorithm, AlgorithmId}; +use aws_lc_rs::kem::{ML_KEM_768, AlgorithmId}; use rustls::crypto::SharedSecret; @@ -16,11 +16,11 @@ const X25519_KYBER_COMBINED_PUBKEY_LEN: usize = X25519_LEN + 1184; const X25519_KYBER_COMBINED_CIPHERTEXT_LEN: usize = X25519_LEN + KYBER_CIPHERTEXT_LEN; const X25519_KYBER_COMBINED_SHARED_SECRET_LEN: usize = X25519_LEN + 32; -struct X25519Kyber768CombinedSecret([u8; X25519_KYBER_COMBINED_SHARED_SECRET_LEN]); +struct X25519Ml768CombinedSecret([u8; X25519_KYBER_COMBINED_SHARED_SECRET_LEN]); -impl X25519Kyber768CombinedSecret { +impl X25519Ml768CombinedSecret { fn combine(x25519: SharedSecret, kyber: kem::SharedSecret) -> Self { - let mut out = X25519Kyber768CombinedSecret([0u8; X25519_KYBER_COMBINED_SHARED_SECRET_LEN]); + let mut out = X25519Ml768CombinedSecret([0u8; X25519_KYBER_COMBINED_SHARED_SECRET_LEN]); out.0[..X25519_LEN].copy_from_slice(x25519.secret_bytes()); out.0[X25519_LEN..].copy_from_slice(kyber.as_ref()); out @@ -48,19 +48,19 @@ pub struct ECDHP521KeyExchange { } #[pyclass(module = "qh3._hazmat")] -pub struct X25519Kyber768Draft00KeyExchange { +pub struct X25519ML768KeyExchange { x25519_private: agreement::PrivateKey, kyber768_decapsulation_key: kem::DecapsulationKey, } #[pymethods] -impl X25519Kyber768Draft00KeyExchange { +impl X25519ML768KeyExchange { #[new] pub fn py_new() -> Self { - X25519Kyber768Draft00KeyExchange { + X25519ML768KeyExchange { x25519_private: agreement::PrivateKey::generate(&agreement::X25519).expect("FAILURE"), kyber768_decapsulation_key: kem::DecapsulationKey::generate( - get_algorithm(AlgorithmId::Kyber768_R3).expect("Kyber768_R3 not available"), + &ML_KEM_768 ) .expect("FAILURE"), } @@ -105,7 +105,7 @@ impl X25519Kyber768Draft00KeyExchange { .decapsulate(kyber.into()) .expect("FAILURE"); - let combined_secret = X25519Kyber768CombinedSecret::combine( + let combined_secret = X25519Ml768CombinedSecret::combine( SharedSecret::from(&x25519_secret[..]), kyber_secret, ); diff --git a/src/lib.rs b/src/lib.rs index 7318d1499..42289c466 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -15,7 +15,7 @@ mod rsa; pub use self::aead::{AeadAes128Gcm, AeadAes256Gcm, AeadChaCha20Poly1305}; pub use self::agreement::{ ECDHP256KeyExchange, ECDHP384KeyExchange, ECDHP521KeyExchange, X25519KeyExchange, - X25519Kyber768Draft00KeyExchange, + X25519ML768KeyExchange, }; pub use self::buffer::{Buffer, BufferReadError, BufferWriteError}; pub use self::certificate::{ @@ -87,7 +87,7 @@ fn _hazmat(py: Python, m: &PyModule) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; - m.add_class::()?; + m.add_class::()?; // General Crypto Error m.add("CryptoError", py.get_type::())?; // Niquests OCSP helper From 2d2b0bed97b72880b9cc90593be61e721dd910ad Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Sun, 29 Dec 2024 16:25:13 +0100 Subject: [PATCH 09/39] :arrow_up: Upgrade pre-commit config --- .pre-commit-config.yaml | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index ee13c0a04..a0da952e7 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,20 +1,20 @@ exclude: 'docs/|tests/' repos: - repo: https://github.com/pre-commit/pre-commit-hooks - rev: v4.4.0 + rev: v5.0.0 hooks: - id: check-yaml - id: debug-statements - id: end-of-file-fixer - id: trailing-whitespace - repo: https://github.com/asottile/pyupgrade - rev: v3.15.1 + rev: v3.19.1 hooks: - id: pyupgrade args: [--py37-plus] - repo: https://github.com/astral-sh/ruff-pre-commit # Ruff version. - rev: v0.3.2 + rev: v0.8.4 hooks: # Run the linter. - id: ruff From 6a3c5af4ff42fddc6abe8bbf51caac3f180dacf0 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Sun, 29 Dec 2024 16:25:37 +0100 Subject: [PATCH 10/39] :art: reformat files --- qh3/asyncio/client.py | 2 +- qh3/tls.py | 4 +--- src/agreement.rs | 8 +++----- 3 files changed, 5 insertions(+), 9 deletions(-) diff --git a/qh3/asyncio/client.py b/qh3/asyncio/client.py index 953394acd..24fd73aa2 100644 --- a/qh3/asyncio/client.py +++ b/qh3/asyncio/client.py @@ -29,7 +29,7 @@ async def connect( stream_handler: QuicStreamHandler | None = None, wait_connected: bool = True, local_port: int = 0, -) -> AsyncGenerator[QuicConnectionProtocol, None]: +) -> AsyncGenerator[QuicConnectionProtocol]: """ Connect to a QUIC server at the given `host` and `port`. diff --git a/qh3/tls.py b/qh3/tls.py index 2d00e7468..1ead46269 100644 --- a/qh3/tls.py +++ b/qh3/tls.py @@ -1385,9 +1385,7 @@ def __init__( self._ec_p384_private_key: ECDHP384KeyExchange | None = None self._ec_p521_private_key: ECDHP521KeyExchange | None = None self._x25519_private_key: X25519KeyExchange | None = None - self._x25519_kyber_768_private_key: X25519ML768KeyExchange | None = ( - None - ) + self._x25519_kyber_768_private_key: X25519ML768KeyExchange | None = None if is_client: self.client_random = os.urandom(32) diff --git a/src/agreement.rs b/src/agreement.rs index 950026a23..689afec23 100644 --- a/src/agreement.rs +++ b/src/agreement.rs @@ -1,7 +1,7 @@ use aws_lc_rs::{agreement, error}; use aws_lc_rs::kem; -use aws_lc_rs::kem::{ML_KEM_768, AlgorithmId}; +use aws_lc_rs::kem::{AlgorithmId, ML_KEM_768}; use rustls::crypto::SharedSecret; @@ -59,10 +59,8 @@ impl X25519ML768KeyExchange { pub fn py_new() -> Self { X25519ML768KeyExchange { x25519_private: agreement::PrivateKey::generate(&agreement::X25519).expect("FAILURE"), - kyber768_decapsulation_key: kem::DecapsulationKey::generate( - &ML_KEM_768 - ) - .expect("FAILURE"), + kyber768_decapsulation_key: kem::DecapsulationKey::generate(&ML_KEM_768) + .expect("FAILURE"), } } From 966453a3ed3a47374c07ac4b36ed6b5edbd4819c Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Mon, 30 Dec 2024 06:10:54 +0100 Subject: [PATCH 11/39] :bug: Fix X25519ML768 key exchange algorithm (minor changes in ciphertext received/sent) --- Cargo.lock | 2 +- Cargo.toml | 4 ++-- README.rst | 2 +- qh3/__init__.py | 2 +- qh3/tls.py | 9 +++++++++ src/agreement.rs | 18 ++++++++++-------- 6 files changed, 24 insertions(+), 13 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 095288254..33c36c47a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1003,7 +1003,7 @@ dependencies = [ [[package]] name = "qh3" -version = "1.2.1" +version = "1.3.0" dependencies = [ "aws-lc-rs", "chacha20poly1305", diff --git a/Cargo.toml b/Cargo.toml index 549d3c9aa..9d61a1c78 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "qh3" -version = "1.2.1" +version = "1.3.0" edition = "2021" rust-version = "1.75" license = "BSD-3" @@ -25,7 +25,7 @@ chacha20poly1305 = "0.10.1" pkcs8 = { version = "0.10.2", features = ["encryption", "pem"] } pkcs1 = { version = "0.7.5", features = ["pem"] } rustls-pemfile = "2.1.2" -aws-lc-rs = { version = "1.9.0", features=["bindgen", "unstable"], default-features = false } +aws-lc-rs = { version = "1.12.0", features=["bindgen"], default-features = false } x509-ocsp = { version = "0.2.1", features = ["builder"] } x509-cert = "0.2.5" der = "0.7.9" diff --git a/README.rst b/README.rst index f5a3cc273..c661a956a 100644 --- a/README.rst +++ b/README.rst @@ -66,7 +66,7 @@ Features - logging TLS traffic secrets - logging QUIC events in QLOG format - HTTP/3 server push support -- Post-Quantum (KEM) Key-Exchange (Kyber R3 NIST) +- Post-Quantum (KEM) Key-Exchange (NIST FIPS 203 ML-KEM-768) Requirements ------------ diff --git a/qh3/__init__.py b/qh3/__init__.py index ac8ccba5f..4a018de14 100644 --- a/qh3/__init__.py +++ b/qh3/__init__.py @@ -13,7 +13,7 @@ from .quic.packet import QuicProtocolVersion from .tls import CipherSuite, SessionTicket -__version__ = "1.2.1" +__version__ = "1.3.0" __all__ = ( "connect", diff --git a/qh3/tls.py b/qh3/tls.py index 1ead46269..f2188b068 100644 --- a/qh3/tls.py +++ b/qh3/tls.py @@ -1553,6 +1553,11 @@ def _client_send_hello(self, output_buf: Buffer) -> None: ) ) supported_groups.append(Group.X25519ML768) + if self.__logger is not None: + self.__logger.debug( + "TLS: Advertising to peer post-quantum algorithm " + "using X25519ML768 (0x11EC)" + ) elif group == Group.GREASE: key_share.append((Group.GREASE, b"\x00")) supported_groups.append(Group.GREASE) @@ -1666,6 +1671,10 @@ def _client_handle_hello(self, input_buf: Buffer, output_buf: Buffer) -> None: shared_key = self._x25519_private_key.exchange(peer_public_key) elif peer_hello.key_share[0] == Group.X25519ML768: shared_key = self._x25519_kyber_768_private_key.exchange(peer_public_key) + if self.__logger is not None: + self.__logger.debug( + "TLS: Post-quantum safety achieved using X25519ML768 (key-exchange)" + ) elif ( peer_hello.key_share[0] == Group.SECP256R1 and self._ec_p256_private_key is not None diff --git a/src/agreement.rs b/src/agreement.rs index 689afec23..bee3a999d 100644 --- a/src/agreement.rs +++ b/src/agreement.rs @@ -15,14 +15,16 @@ const KYBER_CIPHERTEXT_LEN: usize = 1088; const X25519_KYBER_COMBINED_PUBKEY_LEN: usize = X25519_LEN + 1184; const X25519_KYBER_COMBINED_CIPHERTEXT_LEN: usize = X25519_LEN + KYBER_CIPHERTEXT_LEN; const X25519_KYBER_COMBINED_SHARED_SECRET_LEN: usize = X25519_LEN + 32; +const MLKEM768_SECRET_LEN: usize = 32; +const MLKEM768_CIPHERTEXT_LEN: usize = 1088; -struct X25519Ml768CombinedSecret([u8; X25519_KYBER_COMBINED_SHARED_SECRET_LEN]); +struct X25519ML768CombinedSecret([u8; X25519_KYBER_COMBINED_SHARED_SECRET_LEN]); -impl X25519Ml768CombinedSecret { +impl X25519ML768CombinedSecret { fn combine(x25519: SharedSecret, kyber: kem::SharedSecret) -> Self { - let mut out = X25519Ml768CombinedSecret([0u8; X25519_KYBER_COMBINED_SHARED_SECRET_LEN]); - out.0[..X25519_LEN].copy_from_slice(x25519.secret_bytes()); - out.0[X25519_LEN..].copy_from_slice(kyber.as_ref()); + let mut out = X25519ML768CombinedSecret([0u8; X25519_KYBER_COMBINED_SHARED_SECRET_LEN]); + out.0[..MLKEM768_SECRET_LEN].copy_from_slice(kyber.as_ref()); + out.0[MLKEM768_SECRET_LEN..].copy_from_slice(x25519.secret_bytes()); out } } @@ -72,9 +74,9 @@ impl X25519ML768KeyExchange { let mut combined_pub_key = Vec::with_capacity(X25519_KYBER_COMBINED_PUBKEY_LEN); + combined_pub_key.extend_from_slice(kyber_pub.key_bytes().unwrap().as_ref()); combined_pub_key .extend_from_slice(self.x25519_private.compute_public_key().unwrap().as_ref()); - combined_pub_key.extend_from_slice(kyber_pub.key_bytes().unwrap().as_ref()); PyBytes::new(py, combined_pub_key.as_ref()) } @@ -86,7 +88,7 @@ impl X25519ML768KeyExchange { return PyBytes::new(py, &[]); } - let (x25519, kyber) = cipher_text.split_at(X25519_LEN); + let (kyber, x25519) = cipher_text.split_at(MLKEM768_CIPHERTEXT_LEN); let x25519_peer_public_key = agreement::UnparsedPublicKey::new(&agreement::X25519, x25519); @@ -103,7 +105,7 @@ impl X25519ML768KeyExchange { .decapsulate(kyber.into()) .expect("FAILURE"); - let combined_secret = X25519Ml768CombinedSecret::combine( + let combined_secret = X25519ML768CombinedSecret::combine( SharedSecret::from(&x25519_secret[..]), kyber_secret, ); From 4eecc4df56eb00d90ef47d09d75d761bf4f80917 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Mon, 30 Dec 2024 06:40:34 +0100 Subject: [PATCH 12/39] :pencil: initial changelog for v1.3.0 --- CHANGELOG.rst | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 5127063f4..168caf4b4 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -1,3 +1,17 @@ +1.3.0 (2024-12-30) +==================== + +**Changed** +- Post-Quantum key-exchange Kyber 768 Draft upgraded to standard Module-Lattice 768. +- Version negotiation no longer logged as ``INFO``. Every logs generated will always be ``DEBUG`` level. +- Converted our test suite to run on Pytest instead of unittest. + +**Fixed** +- Clippy warnings in our Rust code. + +**Added** +- noxfile. + 1.2.1 (2024-10-15) ==================== From 307c15763b5a1d0b38c35e79738133f40350dd2a Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Mon, 30 Dec 2024 15:56:03 +0100 Subject: [PATCH 13/39] :sparkle: Add a way to serialize/deserialize OCSPResponse and Certificate (rust/pyo3) --- CHANGELOG.rst | 1 + Cargo.lock | 11 +++++++++++ Cargo.toml | 2 ++ qh3/_hazmat.pyi | 6 ++++++ src/certificate.rs | 25 +++++++++++++++++++------ src/ocsp.rs | 20 ++++++++++++++++---- 6 files changed, 55 insertions(+), 10 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 168caf4b4..65b170562 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -11,6 +11,7 @@ **Added** - noxfile. +- miscellaneous serialize/deserialize for Certificate, and OCSPResponse. 1.2.1 (2024-10-15) ==================== diff --git a/Cargo.lock b/Cargo.lock index 33c36c47a..43202f41b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -126,6 +126,15 @@ version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8c3c1a368f70d6cf7302d78f8f7093da241fb8e8807c05cc9e51a125895a6d5b" +[[package]] +name = "bincode" +version = "1.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b1f45e9417d87227c7a56d22e471c6206462cba514c7590c09aff4cf6d1ddcad" +dependencies = [ + "serde", +] + [[package]] name = "bindgen" version = "0.66.1" @@ -1006,6 +1015,7 @@ name = "qh3" version = "1.3.0" dependencies = [ "aws-lc-rs", + "bincode", "chacha20poly1305", "der", "dsa", @@ -1018,6 +1028,7 @@ dependencies = [ "rsa", "rustls", "rustls-pemfile", + "serde", "sha1", "x509-cert", "x509-ocsp", diff --git a/Cargo.toml b/Cargo.toml index 9d61a1c78..f7c31f95a 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -30,6 +30,8 @@ x509-ocsp = { version = "0.2.1", features = ["builder"] } x509-cert = "0.2.5" der = "0.7.9" sha1 = { version = "0.10.6", features = ["oid"] } +serde = { version = "1.0.217", features = ["derive"] } +bincode = {version = "1.3.3"} [patch.crates-io] ls-qpack = { git = 'https://github.com/Ousret/ls-qpack-rs.git' } diff --git a/qh3/_hazmat.pyi b/qh3/_hazmat.pyi index 1805684c4..25fbd3a0e 100644 --- a/qh3/_hazmat.pyi +++ b/qh3/_hazmat.pyi @@ -85,6 +85,9 @@ class Certificate: def get_subject_alt_names(self) -> list[bytes]: ... def public_bytes(self) -> bytes: ... def public_key(self) -> bytes: ... + def serialize(self) -> bytes: ... + @staticmethod + def deserialize(src: bytes) -> Certificate: ... class Rsa: """ @@ -213,6 +216,9 @@ class OCSPResponse: def certificate_status(self) -> OCSPCertStatus: ... @property def revocation_reason(self) -> ReasonFlags | None: ... + def serialize(self) -> bytes: ... + @staticmethod + def deserialize(src: bytes) -> OCSPResponse: ... class OCSPRequest: def __init__(self, peer_certificate: bytes, issuer_certificate: bytes) -> None: ... diff --git a/src/certificate.rs b/src/certificate.rs index 1019d442d..0c1595332 100644 --- a/src/certificate.rs +++ b/src/certificate.rs @@ -5,7 +5,7 @@ use rustls::{CertificateError, Error, RootCertStore}; use pyo3::pyclass; use pyo3::pymethods; -use pyo3::types::{PyBytes, PyList, PyTuple}; +use pyo3::types::{PyBytes, PyList, PyTuple, PyType}; use pyo3::ToPyObject; use pyo3::{PyResult, Python}; @@ -14,28 +14,32 @@ use x509_parser::public_key::PublicKey; use std::sync::Arc; -use pyo3::exceptions::PyException; - use crate::CryptoError; +use bincode::{deserialize, serialize}; +use pyo3::exceptions::PyException; +use serde::{Deserialize, Serialize}; pyo3::create_exception!(_hazmat, SelfSignedCertificateError, PyException); pyo3::create_exception!(_hazmat, InvalidNameCertificateError, PyException); pyo3::create_exception!(_hazmat, ExpiredCertificateError, PyException); pyo3::create_exception!(_hazmat, UnacceptableCertificateError, PyException); -#[pyclass(name = "Extension", module = "qh3._hazmat", frozen)] +#[pyclass(name = "Extension", module = "qh3._hazmat")] +#[derive(Clone, Serialize, Deserialize)] pub struct Extension { oid: String, value: Vec, } -#[pyclass(name = "Subject", module = "qh3._hazmat", frozen)] +#[pyclass(name = "Subject", module = "qh3._hazmat")] +#[derive(Clone, Serialize, Deserialize)] pub struct Subject { oid: String, value: Vec, } -#[pyclass(name = "Certificate", module = "qh3._hazmat", frozen)] +#[pyclass(name = "Certificate", module = "qh3._hazmat")] +#[derive(Clone, Serialize, Deserialize)] pub struct Certificate { version: u8, serial_number: String, @@ -262,6 +266,15 @@ impl Certificate { fn __eq__(&self, other: &Self) -> bool { self.serial_number == other.serial_number } + + pub fn serialize<'py>(&self, py: Python<'py>) -> PyResult<&'py PyBytes> { + Ok(PyBytes::new(py, &serialize(&self).unwrap())) + } + + #[classmethod] + pub fn deserialize(_cls: &PyType, encoded: &PyBytes) -> PyResult { + Ok(deserialize(encoded.as_bytes()).unwrap()) + } } #[pyclass(name = "ServerVerifier", module = "qh3._hazmat")] diff --git a/src/ocsp.rs b/src/ocsp.rs index 2fcfc3149..48e66df74 100644 --- a/src/ocsp.rs +++ b/src/ocsp.rs @@ -3,7 +3,7 @@ // qh3 has no use for it and we won't implement it for this package use pyo3::pyclass; use pyo3::pymethods; -use pyo3::types::PyBytes; +use pyo3::types::{PyBytes, PyType}; use pyo3::{PyResult, Python}; use der::{Decode, Encode}; @@ -15,10 +15,12 @@ use x509_ocsp::{ OcspResponse, OcspResponseStatus as InternalOcspResponseStatus, Request, SingleResponse, }; +use bincode::{deserialize, serialize}; +use serde::{Deserialize, Serialize}; use sha1::Sha1; #[pyclass(module = "qh3._hazmat")] -#[derive(Clone, Copy)] +#[derive(Clone, Copy, Serialize, Deserialize)] #[allow(non_camel_case_types)] pub enum ReasonFlags { unspecified = 0, @@ -34,7 +36,7 @@ pub enum ReasonFlags { } #[pyclass(module = "qh3._hazmat")] -#[derive(Clone, Copy)] +#[derive(Clone, Copy, Serialize, Deserialize)] #[allow(non_camel_case_types)] pub enum OCSPResponseStatus { SUCCESSFUL = 0, @@ -46,7 +48,7 @@ pub enum OCSPResponseStatus { } #[pyclass(module = "qh3._hazmat")] -#[derive(Clone, Copy)] +#[derive(Clone, Copy, Serialize, Deserialize)] #[allow(non_camel_case_types)] pub enum OCSPCertStatus { GOOD = 0, @@ -55,6 +57,7 @@ pub enum OCSPCertStatus { } #[pyclass(module = "qh3._hazmat")] +#[derive(Clone, Serialize, Deserialize)] #[allow(non_camel_case_types)] pub struct OCSPResponse { next_update: u64, @@ -146,6 +149,15 @@ impl OCSPResponse { pub fn revocation_reason(&self) -> Option { self.revocation_reason } + + pub fn serialize<'py>(&self, py: Python<'py>) -> PyResult<&'py PyBytes> { + Ok(PyBytes::new(py, &serialize(&self).unwrap())) + } + + #[classmethod] + pub fn deserialize(_cls: &PyType, encoded: &PyBytes) -> PyResult { + Ok(deserialize(encoded.as_bytes()).unwrap()) + } } #[pyclass(module = "qh3._hazmat")] From 77521e97ae54e3278dd1ba2ac47743297e22ff89 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Tue, 31 Dec 2024 16:36:42 +0100 Subject: [PATCH 14/39] :art: Ensure proper exception instead of panics in rust code (+fix server pq negotiation) --- CHANGELOG.rst | 3 + qh3/_hazmat.pyi | 1 + qh3/tls.py | 2 +- src/aead.rs | 26 ++- src/agreement.rs | 354 ++++++++++++++++++++++++++++----------- src/buffer.rs | 21 +-- src/headers.rs | 12 +- src/hpk.rs | 32 ++-- src/ocsp.rs | 5 +- src/pkcs8.rs | 51 ++++-- src/private_key.rs | 134 ++++++++++----- src/rsa.rs | 40 +++-- tests/test_connection.py | 16 +- tests/test_tls.py | 4 +- 14 files changed, 477 insertions(+), 224 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 65b170562..a27342724 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -8,6 +8,9 @@ **Fixed** - Clippy warnings in our Rust code. +- Rust code may panic due to lack of proper result unpacking on the cryptographic calls. Now any error will + raise exception ``CryptoError`` instead. +- Negotiating post-quantum key exchange (server side). **Added** - noxfile. diff --git a/qh3/_hazmat.pyi b/qh3/_hazmat.pyi index 25fbd3a0e..e6292d336 100644 --- a/qh3/_hazmat.pyi +++ b/qh3/_hazmat.pyi @@ -129,6 +129,7 @@ class X25519ML768KeyExchange: def __init__(self) -> None: ... def public_key(self) -> bytes: ... def exchange(self, peer_public_key: bytes) -> bytes: ... + def shared_ciphertext(self) -> bytes: ... class X25519KeyExchange: def __init__(self) -> None: ... diff --git a/qh3/tls.py b/qh3/tls.py index f2188b068..3ac5228cb 100644 --- a/qh3/tls.py +++ b/qh3/tls.py @@ -2046,10 +2046,10 @@ def _server_handle_hello( break elif key_share[0] == Group.X25519ML768: self._x25519_kyber_768_private_key = X25519ML768KeyExchange() - public_key = self._x25519_kyber_768_private_key.public_key() shared_key = self._x25519_kyber_768_private_key.exchange( peer_public_key ) + public_key = self._x25519_kyber_768_private_key.shared_ciphertext() group_kx = Group.X25519ML768 break elif key_share[0] == Group.SECP256R1: diff --git a/src/aead.rs b/src/aead.rs index 9202e4b3a..6de8438cb 100644 --- a/src/aead.rs +++ b/src/aead.rs @@ -47,8 +47,10 @@ impl AeadAes256Gcm { let plaintext_len = in_out_buffer.len() - AES_256_GCM.tag_len(); let opening_key: TlsRecordOpeningKey = - TlsRecordOpeningKey::new(&AES_256_GCM, TlsProtocolId::TLS13, &self.key) - .expect("FAILURE"); + match TlsRecordOpeningKey::new(&AES_256_GCM, TlsProtocolId::TLS13, &self.key) { + Ok(k) => k, + Err(_) => return Err(CryptoError::new_err("Invalid AEAD key")), + }; let aad = Aad::from(associated_data.as_bytes()); @@ -74,8 +76,10 @@ impl AeadAes256Gcm { let mut in_out_buffer = Vec::from(data.as_bytes()); let mut sealing_key: TlsRecordSealingKey = - TlsRecordSealingKey::new(&AES_256_GCM, TlsProtocolId::TLS13, &self.key) - .expect("FAILURE"); + match TlsRecordSealingKey::new(&AES_256_GCM, TlsProtocolId::TLS13, &self.key) { + Ok(k) => k, + Err(_) => return Err(CryptoError::new_err("Invalid AEAD key")), + }; let aad = Aad::from(associated_data.as_bytes()); @@ -111,8 +115,12 @@ impl AeadAes128Gcm { let mut in_out_buffer = data.as_bytes().to_vec(); let plaintext_len = in_out_buffer.len() - AES_128_GCM.tag_len(); - let opening_key = TlsRecordOpeningKey::new(&AES_128_GCM, TlsProtocolId::TLS13, &self.key) - .expect("FAILURE"); + let opening_key = + match TlsRecordOpeningKey::new(&AES_128_GCM, TlsProtocolId::TLS13, &self.key) { + Ok(k) => k, + Err(_) => return Err(CryptoError::new_err("Invalid AEAD key")), + }; + let aad = Aad::from(associated_data.as_bytes()); let res = opening_key.open_in_place( @@ -137,8 +145,10 @@ impl AeadAes128Gcm { let mut in_out_buffer = Vec::from(data.as_bytes()); let mut sealing_key = - TlsRecordSealingKey::new(&AES_128_GCM, TlsProtocolId::TLS13, &self.key) - .expect("FAILURE"); + match TlsRecordSealingKey::new(&AES_128_GCM, TlsProtocolId::TLS13, &self.key) { + Ok(k) => k, + Err(_) => return Err(CryptoError::new_err("Invalid AEAD key")), + }; let aad = Aad::from(associated_data.as_bytes()); diff --git a/src/agreement.rs b/src/agreement.rs index bee3a999d..ad69aedac 100644 --- a/src/agreement.rs +++ b/src/agreement.rs @@ -1,14 +1,14 @@ use aws_lc_rs::{agreement, error}; use aws_lc_rs::kem; -use aws_lc_rs::kem::{AlgorithmId, ML_KEM_768}; - +use aws_lc_rs::kem::ML_KEM_768; use rustls::crypto::SharedSecret; -use pyo3::pyclass; +use crate::CryptoError; use pyo3::pymethods; use pyo3::types::PyBytes; use pyo3::Python; +use pyo3::{pyclass, PyResult}; const X25519_LEN: usize = 32; const KYBER_CIPHERTEXT_LEN: usize = 1088; @@ -52,190 +52,352 @@ pub struct ECDHP521KeyExchange { #[pyclass(module = "qh3._hazmat")] pub struct X25519ML768KeyExchange { x25519_private: agreement::PrivateKey, - kyber768_decapsulation_key: kem::DecapsulationKey, + kyber768_decapsulation_key: kem::DecapsulationKey, + cipher_text: Vec, } #[pymethods] impl X25519ML768KeyExchange { #[new] - pub fn py_new() -> Self { - X25519ML768KeyExchange { - x25519_private: agreement::PrivateKey::generate(&agreement::X25519).expect("FAILURE"), - kyber768_decapsulation_key: kem::DecapsulationKey::generate(&ML_KEM_768) - .expect("FAILURE"), - } + pub fn py_new() -> PyResult { + let x25519_pk = match agreement::PrivateKey::generate(&agreement::X25519) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Unable to generate X25519 key")), + }; + + let ml768_dk = match kem::DecapsulationKey::generate(&ML_KEM_768) { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "Unable to generate ML_KEM_768 decapsulation key", + )) + } + }; + + Ok(X25519ML768KeyExchange { + x25519_private: x25519_pk, + kyber768_decapsulation_key: ml768_dk, + cipher_text: Vec::new(), + }) } - pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - let kyber_pub = self - .kyber768_decapsulation_key - .encapsulation_key() - .expect("FAILURE"); + pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult<&'a PyBytes> { + let kyber_pub = match self.kyber768_decapsulation_key.encapsulation_key() { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "Unable to generate ML768 encapsulation key", + )) + } + }; let mut combined_pub_key = Vec::with_capacity(X25519_KYBER_COMBINED_PUBKEY_LEN); - combined_pub_key.extend_from_slice(kyber_pub.key_bytes().unwrap().as_ref()); - combined_pub_key - .extend_from_slice(self.x25519_private.compute_public_key().unwrap().as_ref()); - - PyBytes::new(py, combined_pub_key.as_ref()) + let raw_ml_encapsulation_key = match kyber_pub.key_bytes() { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "Unable to get encapsulation key for ML768 as plain bytes", + )) + } + }; + + let raw_x25519_public_key = match self.x25519_private.compute_public_key() { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "Unable to get public key for X25519 as plain bytes", + )) + } + }; + + combined_pub_key.extend_from_slice(raw_ml_encapsulation_key.as_ref()); + combined_pub_key.extend_from_slice(raw_x25519_public_key.as_ref()); + + Ok(PyBytes::new(py, combined_pub_key.as_ref())) } - pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { - let cipher_text = peer_public_key.as_bytes(); - - if cipher_text.len() != X25519_KYBER_COMBINED_CIPHERTEXT_LEN { - return PyBytes::new(py, &[]); + pub fn shared_ciphertext<'a>(&mut self, py: Python<'a>) -> PyResult<&'a PyBytes> { + if self.cipher_text.is_empty() { + return Err(CryptoError::new_err( + "You must receive client share first. Call exchange with client share.", + )); } - let (kyber, x25519) = cipher_text.split_at(MLKEM768_CIPHERTEXT_LEN); + let mut combined_pub_key = Vec::with_capacity(X25519_KYBER_COMBINED_CIPHERTEXT_LEN); - let x25519_peer_public_key = agreement::UnparsedPublicKey::new(&agreement::X25519, x25519); + let raw_x25519_public_key = match self.x25519_private.compute_public_key() { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "Unable to get public key for X25519 as plain bytes", + )) + } + }; - let x25519_secret = agreement::agree( - &self.x25519_private, - &x25519_peer_public_key, - error::Unspecified, - |_key_material| Ok(_key_material.to_vec()), - ) - .expect("FAILURE"); + combined_pub_key.extend_from_slice(self.cipher_text.as_ref()); + combined_pub_key.extend_from_slice(raw_x25519_public_key.as_ref()); - let kyber_secret = self - .kyber768_decapsulation_key - .decapsulate(kyber.into()) - .expect("FAILURE"); + self.cipher_text = Vec::new(); - let combined_secret = X25519ML768CombinedSecret::combine( - SharedSecret::from(&x25519_secret[..]), - kyber_secret, - ); + Ok(PyBytes::new(py, combined_pub_key.as_ref())) + } - let key_material = SharedSecret::from(&combined_secret.0[..]); + pub fn exchange<'a>( + &mut self, + py: Python<'a>, + peer_public_key: &PyBytes, + ) -> PyResult<&'a PyBytes> { + let cipher_text = peer_public_key.as_bytes(); - PyBytes::new(py, key_material.secret_bytes()) + // client share received + if cipher_text.len() == 1216 { + let (kyber, x25519) = cipher_text.split_at(1184); + + let x25519_peer_public_key = + agreement::UnparsedPublicKey::new(&agreement::X25519, x25519); + + let x25519_secret = match agreement::agree( + &self.x25519_private, + &x25519_peer_public_key, + error::Unspecified, + |_key_material| Ok(_key_material.to_vec()), + ) { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "X25519ML768 exchange failure due to X25519 agreement failure", + )) + } + }; + + let ml768_share = match kem::EncapsulationKey::new(&kem::ML_KEM_768, kyber) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Unable to parse ML768 share")), + }; + + // Bob executes the encapsulation algorithm to to produce their copy of the secret, and associated ciphertext. + let (ciphertext, bob_secret) = ml768_share.encapsulate().expect(""); + + let combined_secret = X25519ML768CombinedSecret::combine( + SharedSecret::from(&x25519_secret[..]), + bob_secret, + ); + + self.cipher_text = ciphertext.as_ref().to_vec(); + + let key_material = SharedSecret::from(&combined_secret.0[..]); + + Ok(PyBytes::new(py, key_material.secret_bytes())) + } else { + let (kyber, x25519) = cipher_text.split_at(MLKEM768_CIPHERTEXT_LEN); + + let x25519_peer_public_key = + agreement::UnparsedPublicKey::new(&agreement::X25519, x25519); + + let x25519_secret = match agreement::agree( + &self.x25519_private, + &x25519_peer_public_key, + error::Unspecified, + |_key_material| Ok(_key_material.to_vec()), + ) { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "X25519ML768 exchange failure due to X25519 agreement failure", + )) + } + }; + + let kyber_secret = match self.kyber768_decapsulation_key.decapsulate(kyber.into()) { + Ok(secret) => secret, + Err(_) => { + return Err(CryptoError::new_err( + "X25519ML768 exchange failure due to decapsulation error", + )) + } + }; + + let combined_secret = X25519ML768CombinedSecret::combine( + SharedSecret::from(&x25519_secret[..]), + kyber_secret, + ); + + let key_material = SharedSecret::from(&combined_secret.0[..]); + + Ok(PyBytes::new(py, key_material.secret_bytes())) + } } } #[pymethods] impl X25519KeyExchange { #[new] - pub fn py_new() -> Self { - X25519KeyExchange { - private: agreement::PrivateKey::generate(&agreement::X25519).expect("FAILURE"), - } - } + pub fn py_new() -> PyResult { + let x25519_pk = match agreement::PrivateKey::generate(&agreement::X25519) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Unable to generate X25519 key")), + }; - pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - let my_public_key = self.private.compute_public_key().unwrap(); + Ok(X25519KeyExchange { private: x25519_pk }) + } - PyBytes::new(py, my_public_key.as_ref()) + pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult<&'a PyBytes> { + let my_public_key = match self.private.compute_public_key() { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "Unable to get public key for X25519 as plain bytes", + )) + } + }; + + Ok(PyBytes::new(py, my_public_key.as_ref())) } - pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { + pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> PyResult<&'a PyBytes> { let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::X25519, peer_public_key.as_bytes()); - let key_material = agreement::agree( + let key_material = match agreement::agree( &self.private, &peer_public_key, error::Unspecified, |_key_material| Ok(_key_material.to_vec()), - ) - .expect("FAILURE"); + ) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("X25519 exchange failure")), + }; - PyBytes::new(py, &key_material) + Ok(PyBytes::new(py, &key_material)) } } #[pymethods] impl ECDHP256KeyExchange { #[new] - pub fn py_new() -> Self { - ECDHP256KeyExchange { - private: agreement::PrivateKey::generate(&agreement::ECDH_P256).expect("FAILURE"), - } - } + pub fn py_new() -> PyResult { + let ecdh_key = match agreement::PrivateKey::generate(&agreement::ECDH_P256) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Unable to generate ECDH p256 key")), + }; - pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - let my_public_key = self.private.compute_public_key().unwrap(); + Ok(ECDHP256KeyExchange { private: ecdh_key }) + } - PyBytes::new(py, my_public_key.as_ref()) + pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult<&'a PyBytes> { + let my_public_key = match self.private.compute_public_key() { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "Unable to get public key for ECDHP256KeyExchange", + )) + } + }; + + Ok(PyBytes::new(py, my_public_key.as_ref())) } - pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { + pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> PyResult<&'a PyBytes> { let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::ECDH_P256, peer_public_key.as_bytes()); - let key_material = agreement::agree( + let key_material = match agreement::agree( &self.private, &peer_public_key, error::Unspecified, |_key_material| Ok(_key_material.to_vec()), - ) - .expect("FAILURE"); + ) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("ECDHP256KeyExchange failure")), + }; - PyBytes::new(py, &key_material) + Ok(PyBytes::new(py, &key_material)) } } #[pymethods] impl ECDHP384KeyExchange { #[new] - pub fn py_new() -> Self { - ECDHP384KeyExchange { - private: agreement::PrivateKey::generate(&agreement::ECDH_P384).expect("FAILURE"), - } - } + pub fn py_new() -> PyResult { + let ecdh_key = match agreement::PrivateKey::generate(&agreement::ECDH_P384) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Unable to generate ECDH p384 key")), + }; - pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - let my_public_key = self.private.compute_public_key().unwrap(); + Ok(ECDHP384KeyExchange { private: ecdh_key }) + } - PyBytes::new(py, my_public_key.as_ref()) + pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult<&'a PyBytes> { + let my_public_key = match self.private.compute_public_key() { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "Unable to compute ECDH p384 public key", + )) + } + }; + + Ok(PyBytes::new(py, my_public_key.as_ref())) } - pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { + pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> PyResult<&'a PyBytes> { let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::ECDH_P384, peer_public_key.as_bytes()); - let key_material = agreement::agree( + let key_material = match agreement::agree( &self.private, &peer_public_key, error::Unspecified, |_key_material| Ok(_key_material.to_vec()), - ) - .expect("FAILURE"); + ) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("ECDHP384 exchange failure")), + }; - PyBytes::new(py, &key_material) + Ok(PyBytes::new(py, &key_material)) } } #[pymethods] impl ECDHP521KeyExchange { #[new] - pub fn py_new() -> Self { - ECDHP521KeyExchange { - private: agreement::PrivateKey::generate(&agreement::ECDH_P521).expect("FAILURE"), - } - } + pub fn py_new() -> PyResult { + let ecdh_pk = match agreement::PrivateKey::generate(&agreement::ECDH_P521) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Unable to generate ECDH p521 key")), + }; - pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { - let my_public_key = self.private.compute_public_key().unwrap(); + Ok(ECDHP521KeyExchange { private: ecdh_pk }) + } - PyBytes::new(py, my_public_key.as_ref()) + pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult<&'a PyBytes> { + let my_public_key = match self.private.compute_public_key() { + Ok(key) => key, + Err(_) => { + return Err(CryptoError::new_err( + "Unable to compute ECDH p521 public key", + )) + } + }; + + Ok(PyBytes::new(py, my_public_key.as_ref())) } - pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> &'a PyBytes { + pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> PyResult<&'a PyBytes> { let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::ECDH_P521, peer_public_key.as_bytes()); - let key_material = agreement::agree( + let key_material = match agreement::agree( &self.private, &peer_public_key, error::Unspecified, |_key_material| Ok(_key_material.to_vec()), - ) - .expect("FAILURE"); + ) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("ECDHP521 exchange failure")), + }; - PyBytes::new(py, &key_material) + Ok(PyBytes::new(py, &key_material)) } } diff --git a/src/buffer.rs b/src/buffer.rs index 84d3b8c1f..47faa9cab 100644 --- a/src/buffer.rs +++ b/src/buffer.rs @@ -113,11 +113,8 @@ impl Buffer { return Err(BufferReadError::new_err("Read out of bounds")); } - let extract = u16::from_be_bytes( - self.data[self.pos as usize..(self.pos + 2) as usize] - .try_into() - .expect("failure"), - ); + let extract = + u16::from_be_bytes(self.data[self.pos as usize..(self.pos + 2) as usize].try_into()?); self.pos += 2; Ok(extract) @@ -132,11 +129,8 @@ impl Buffer { return Err(BufferReadError::new_err("Read out of bounds")); } - let extract = u32::from_be_bytes( - self.data[self.pos as usize..(self.pos + 4) as usize] - .try_into() - .expect("failure"), - ); + let extract = + u32::from_be_bytes(self.data[self.pos as usize..(self.pos + 4) as usize].try_into()?); self.pos += 4; Ok(extract) @@ -151,11 +145,8 @@ impl Buffer { return Err(BufferReadError::new_err("Read out of bounds")); } - let extract = u64::from_be_bytes( - self.data[self.pos as usize..(self.pos + 8) as usize] - .try_into() - .expect("failure"), - ); + let extract = + u64::from_be_bytes(self.data[self.pos as usize..(self.pos + 8) as usize].try_into()?); self.pos += 8; Ok(extract) diff --git a/src/headers.rs b/src/headers.rs index eb77a2a99..105d744bd 100644 --- a/src/headers.rs +++ b/src/headers.rs @@ -41,10 +41,14 @@ impl QpackEncoder { dyn_table_capacity: u32, blocked_streams: u32, ) -> PyResult<&'a PyBytes> { - let r = self - .encoder - .configure(max_table_capacity, dyn_table_capacity, blocked_streams) - .expect("FAILURE"); + let r = + match self + .encoder + .configure(max_table_capacity, dyn_table_capacity, blocked_streams) + { + Ok(r) => r, + Err(_) => return Err(EncoderStreamError::new_err("failed to configure encoder")), + }; Ok(PyBytes::new(py, r.data())) } diff --git a/src/hpk.rs b/src/hpk.rs index bcd4d58a4..c67c1ce0e 100644 --- a/src/hpk.rs +++ b/src/hpk.rs @@ -14,19 +14,25 @@ pub struct QUICHeaderProtection { #[pymethods] impl QUICHeaderProtection { #[new] - pub fn py_new(key: &PyBytes, algorithm: u16) -> Self { - QUICHeaderProtection { - hpk: HeaderProtectionKey::new( - match algorithm { - 128 => &AES_128, - 256 => &AES_256, - 20 => &CHACHA20, - _ => panic!("unsupported"), - }, - key.as_bytes(), - ) - .expect("FAILURE"), - } + pub fn py_new(key: &PyBytes, algorithm: u16) -> PyResult { + let inner_hpk = match HeaderProtectionKey::new( + match algorithm { + 128 => &AES_128, + 256 => &AES_256, + 20 => &CHACHA20, + _ => return Err(CryptoError::new_err("Algorithm not supported")), + }, + key.as_bytes(), + ) { + Ok(hpk) => hpk, + Err(_) => { + return Err(CryptoError::new_err( + "Given key is not valid for chosen algorithm", + )) + } + }; + + Ok(QUICHeaderProtection { hpk: inner_hpk }) } pub fn mask<'a>(&self, py: Python<'a>, sample: &PyBytes) -> PyResult<&'a PyBytes> { diff --git a/src/ocsp.rs b/src/ocsp.rs index 48e66df74..155afce2e 100644 --- a/src/ocsp.rs +++ b/src/ocsp.rs @@ -70,7 +70,10 @@ pub struct OCSPResponse { impl OCSPResponse { #[new] pub fn py_new(raw_response: &PyBytes) -> PyResult { - let ocsp_res: OcspResponse = OcspResponse::from_der(raw_response.as_bytes()).unwrap(); + let ocsp_res: OcspResponse = match OcspResponse::from_der(raw_response.as_bytes()) { + Ok(ocsp_res) => ocsp_res, + Err(_) => return Err(PyValueError::new_err("OCSP DER given is invalid")), + }; if ocsp_res.response_bytes.is_none() { return Err(PyValueError::new_err("OCSP Server did not provide answers")); diff --git a/src/pkcs8.rs b/src/pkcs8.rs index f7a7307a8..01fbfd422 100644 --- a/src/pkcs8.rs +++ b/src/pkcs8.rs @@ -1,7 +1,7 @@ -use pyo3::pyclass; use pyo3::pymethods; use pyo3::types::PyBytes; use pyo3::Python; +use pyo3::{pyclass, PyResult}; use pkcs8::{der::Encode, DecodePrivateKey, Error, PrivateKeyInfo as InternalPrivateKeyInfo}; use rsa::{ @@ -10,6 +10,7 @@ use rsa::{ RsaPrivateKey, }; +use crate::CryptoError; use rustls_pemfile::{read_one_from_slice, Item}; #[pyclass(module = "qh3._hazmat")] @@ -67,9 +68,9 @@ impl TryFrom> for PrivateKeyInfo { #[pymethods] impl PrivateKeyInfo { #[new] - pub fn py_new(raw_pem_content: &PyBytes, password: Option<&PyBytes>) -> Self { + pub fn py_new(raw_pem_content: &PyBytes, password: Option<&PyBytes>) -> PyResult { let pem_content = raw_pem_content.as_bytes(); - let decoded_bytes = std::str::from_utf8(pem_content).unwrap(); + let decoded_bytes = std::str::from_utf8(pem_content)?; let is_encrypted = decoded_bytes.contains("ENCRYPTED"); let item = read_one_from_slice(pem_content); @@ -77,47 +78,63 @@ impl PrivateKeyInfo { match item.unwrap().unwrap().0 { Item::Pkcs1Key(key) => { if is_encrypted { - panic!("unsupported"); + return Err(CryptoError::new_err( + "RSA Pkcs1Key is encrypted, please decrypt it prior to passing it.", + )); } let rsa_key: RsaPrivateKey = - RsaPrivateKey::from_pkcs1_der(key.secret_pkcs1_der()).unwrap(); - - let pkcs8_pem = rsa_key.to_pkcs8_pem(LineEnding::LF).expect("FAILURE"); + match RsaPrivateKey::from_pkcs1_der(key.secret_pkcs1_der()) { + Ok(rsa_key) => rsa_key, + Err(_) => return Err(CryptoError::new_err("RSA private key is invalid.")), + }; + + let pkcs8_pem = match rsa_key.to_pkcs8_pem(LineEnding::LF) { + Ok(pem) => pem, + Err(_) => { + return Err(CryptoError::new_err("malformed/invalid RSA private key?")) + } + }; let pkcs8_pem: &str = pkcs8_pem.as_ref(); - PrivateKeyInfo::from_pkcs8_pem(pkcs8_pem).unwrap() + Ok(PrivateKeyInfo::from_pkcs8_pem(pkcs8_pem).unwrap()) } Item::Pkcs8Key(_key) => { if is_encrypted { - return PrivateKeyInfo::from_pkcs8_encrypted_pem( + return match PrivateKeyInfo::from_pkcs8_encrypted_pem( decoded_bytes, password.unwrap().as_bytes(), - ) - .unwrap(); + ) { + Ok(key) => Ok(key), + Err(_) => Err(CryptoError::new_err( + "unable to decrypt Pkcs8 private key. invalid password?", + )), + }; } - PrivateKeyInfo::from_pkcs8_pem(decoded_bytes).unwrap() + Ok(PrivateKeyInfo::from_pkcs8_pem(decoded_bytes).unwrap()) } Item::Sec1Key(key) => { if is_encrypted { - panic!("unsupported"); + return Err(CryptoError::new_err( + "Sec1key encrypted is encrypted, please decrypt it prior to passing it.", + )); } let sec1_der = key.secret_sec1_der().to_vec(); - PrivateKeyInfo { + Ok(PrivateKeyInfo { cert_type: match sec1_der.len() { 32..=121 => KeyType::ECDSA_P256, 132..=167 => KeyType::ECDSA_P384, 200..=400 => KeyType::ECDSA_P521, - _ => panic!("unsupported sec1 key"), + _ => return Err(CryptoError::new_err("unsupported sec1key")), }, der_encoded: sec1_der, - } + }) } - _ => panic!("unsupported"), + _ => Err(CryptoError::new_err("unsupported key type")), } } diff --git a/src/private_key.rs b/src/private_key.rs index 75504a918..769d3f491 100644 --- a/src/private_key.rs +++ b/src/private_key.rs @@ -26,6 +26,7 @@ use aws_lc_rs::rand::SystemRandom; use aws_lc_rs::signature; use aws_lc_rs::signature::UnparsedPublicKey; +use crate::CryptoError; use pyo3::exceptions::PyException; use pyo3::pyclass; use pyo3::pyfunction; @@ -59,10 +60,13 @@ pub struct RsaPrivateKey { #[pymethods] impl Ed25519PrivateKey { #[new] - pub fn py_new(pkcs8: &PyBytes) -> Self { - Ed25519PrivateKey { - inner: InternalEd25519PrivateKey::from_pkcs8(pkcs8.as_bytes()).expect("FAILURE"), - } + pub fn py_new(pkcs8: &PyBytes) -> PyResult { + let pk = match InternalEd25519PrivateKey::from_pkcs8(pkcs8.as_bytes()) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Invalid Ed25519 PrivateKey")), + }; + + Ok(Ed25519PrivateKey { inner: pk }) } pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { @@ -79,26 +83,37 @@ impl Ed25519PrivateKey { #[pymethods] impl EcPrivateKey { #[new] - pub fn py_new(pkcs8: &PyBytes, curve_type: u32) -> Self { + pub fn py_new(pkcs8: &PyBytes, curve_type: u32) -> PyResult { let signing_algorithm = match curve_type { 256 => &ECDSA_P256_SHA256_ASN1_SIGNING, 384 => &ECDSA_P384_SHA384_ASN1_SIGNING, 521 => &ECDSA_P521_SHA512_ASN1_SIGNING, - _ => panic!("unsupported"), + _ => { + return Err(CryptoError::new_err( + "Unsupported curve type in EcPrivateKey", + )) + } }; - EcPrivateKey { - inner: InternalEcPrivateKey::from_pkcs8(signing_algorithm, pkcs8.as_bytes()) - .expect("FAILURE"), + let pk = match InternalEcPrivateKey::from_pkcs8(signing_algorithm, pkcs8.as_bytes()) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Invalid Ec PrivateKey")), + }; + + Ok(EcPrivateKey { + inner: pk, curve: curve_type, - } + }) } - pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { + pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> PyResult<&'a PyBytes> { let rng = SystemRandom::new(); - let signature = self.inner.sign(&rng, data.as_bytes()); + let signature = match self.inner.sign(&rng, data.as_bytes()) { + Ok(signature) => signature, + Err(_) => return Err(CryptoError::new_err("Ec signature could not be issued")), + }; - PyBytes::new(py, signature.unwrap().as_ref()) + Ok(PyBytes::new(py, signature.as_ref())) } pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { @@ -114,10 +129,13 @@ impl EcPrivateKey { #[pymethods] impl DsaPrivateKey { #[new] - pub fn py_new(pkcs8: &PyBytes) -> Self { - DsaPrivateKey { - inner: InternalDsaPrivateKey::from_pkcs8_der(pkcs8.as_bytes()).expect("FAILURE"), - } + pub fn py_new(pkcs8: &PyBytes) -> PyResult { + let pk = match InternalDsaPrivateKey::from_pkcs8_der(pkcs8.as_bytes()) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Invalid Dsa PrivateKey")), + }; + + Ok(DsaPrivateKey { inner: pk }) } pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { @@ -132,7 +150,7 @@ impl DsaPrivateKey { self.inner .verifying_key() .to_public_key_der() - .expect("FAILURE") + .unwrap() .as_bytes(), ) } @@ -141,10 +159,13 @@ impl DsaPrivateKey { #[pymethods] impl RsaPrivateKey { #[new] - pub fn py_new(pkcs8: &PyBytes) -> Self { - RsaPrivateKey { - inner: InternalRsaPrivateKey::from_pkcs8_der(pkcs8.as_bytes()).expect("FAILURE"), - } + pub fn py_new(pkcs8: &PyBytes) -> PyResult { + let pk = match InternalRsaPrivateKey::from_pkcs8_der(pkcs8.as_bytes()) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Invalid Rsa PrivateKey")), + }; + + Ok(RsaPrivateKey { inner: pk }) } pub fn sign<'a>( @@ -153,39 +174,43 @@ impl RsaPrivateKey { data: &PyBytes, is_pss_padding: bool, hash_size: u32, - ) -> &'a PyBytes { + ) -> PyResult<&'a PyBytes> { let private_key = self.inner.clone(); match is_pss_padding { true => match hash_size { 256 => { let signer = InternalRsaPssSigningKey::::new(private_key); - PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) + Ok(PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec())) } 384 => { let signer = InternalRsaPssSigningKey::::new(private_key); - PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) + Ok(PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec())) } 512 => { let signer = InternalRsaPssSigningKey::::new(private_key); - PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) + Ok(PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec())) } - _ => panic!("unsupported"), + _ => Err(CryptoError::new_err( + "unsupported hash size for RSA signing", + )), }, false => match hash_size { 256 => { let signer = InternalRsaPkcsSigningKey::::new(private_key); - PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) + Ok(PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec())) } 384 => { let signer = InternalRsaPkcsSigningKey::::new(private_key); - PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) + Ok(PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec())) } 512 => { let signer = InternalRsaPkcsSigningKey::::new(private_key); - PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec()) + Ok(PyBytes::new(py, &signer.sign(data.as_bytes()).to_vec())) } - _ => panic!("unsupported"), + _ => Err(CryptoError::new_err( + "unsupported hash size for RSA signing", + )), }, } } @@ -220,7 +245,10 @@ pub fn verify_with_public_key( || pkcs115_blind_signature.contains(&algorithm) { let rsa_parsed_public_key = - InternalRsaPublicKey::from_public_key_der(public_key_bytes).expect("FAILURE"); + match InternalRsaPublicKey::from_public_key_der(public_key_bytes) { + Ok(public_key) => public_key, + Err(_) => return Err(CryptoError::new_err("Invalid RSA public key")), + }; return match algorithm { 0x0804 | 0x0809 => { @@ -228,7 +256,10 @@ pub fn verify_with_public_key( let res = alt_verifier.verify( message.as_bytes(), - &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE"), + match &RsaPssSignature::try_from(signature.as_bytes()) { + Ok(signature) => signature, + Err(_) => return Err(CryptoError::new_err("Invalid RSA PSS signature")), + }, ); return match res { @@ -241,7 +272,10 @@ pub fn verify_with_public_key( let res = alt_verifier.verify( message.as_bytes(), - &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE"), + match &RsaPssSignature::try_from(signature.as_bytes()) { + Ok(signature) => signature, + Err(_) => return Err(CryptoError::new_err("Invalid RSA PSS signature")), + }, ); return match res { @@ -254,7 +288,10 @@ pub fn verify_with_public_key( let res = alt_verifier.verify( message.as_bytes(), - &RsaPssSignature::try_from(signature.as_bytes()).expect("FAILURE"), + match &RsaPssSignature::try_from(signature.as_bytes()) { + Ok(signature) => signature, + Err(_) => return Err(CryptoError::new_err("Invalid RSA PSS signature")), + }, ); return match res { @@ -268,7 +305,10 @@ pub fn verify_with_public_key( let res = alt_verifier.verify( message.as_bytes(), - &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE"), + match &RsaPkcsSignature::try_from(signature.as_bytes()) { + Ok(signature) => signature, + Err(_) => return Err(CryptoError::new_err("Invalid RSA PKCS signature")), + }, ); return match res { @@ -281,7 +321,10 @@ pub fn verify_with_public_key( let res = alt_verifier.verify( message.as_bytes(), - &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE"), + match &RsaPkcsSignature::try_from(signature.as_bytes()) { + Ok(signature) => signature, + Err(_) => return Err(CryptoError::new_err("Invalid RSA PKCS signature")), + }, ); return match res { @@ -294,7 +337,10 @@ pub fn verify_with_public_key( let res = alt_verifier.verify( message.as_bytes(), - &RsaPkcsSignature::try_from(signature.as_bytes()).expect("FAILURE"), + match &RsaPkcsSignature::try_from(signature.as_bytes()) { + Ok(signature) => signature, + Err(_) => return Err(CryptoError::new_err("Invalid RSA PKCS signature")), + }, ); return match res { @@ -303,16 +349,20 @@ pub fn verify_with_public_key( }; } - _ => panic!("unreachable statement"), + _ => Err(CryptoError::new_err("unsupported signature algorithm")), }; } if algorithm == 0x0807 { let ed25519_verifier: Ed25519VerifyingKey = - Ed25519VerifyingKey::from_public_key_der(public_key_bytes).expect("FAILURE"); + match Ed25519VerifyingKey::from_public_key_der(public_key_bytes) { + Ok(public_key) => public_key, + Err(_) => return Err(CryptoError::new_err("Invalid Ed25519 public key")), + }; + let res = ed25519_verifier.verify( message.as_bytes(), - &Ed25519Signature::from_bytes(signature.as_bytes()[0..64].try_into().unwrap()), + &Ed25519Signature::from_bytes(signature.as_bytes()[0..64].try_into()?), ); return match res { @@ -326,7 +376,7 @@ pub fn verify_with_public_key( 0x0403 => &signature::ECDSA_P256_SHA256_ASN1, 0x0503 => &signature::ECDSA_P384_SHA384_ASN1, 0x0603 => &signature::ECDSA_P521_SHA512_ASN1, - _ => panic!("unsupported algorithm"), + _ => return Err(CryptoError::new_err("unsupported signature algorithm")), }, public_key_bytes, ); diff --git a/src/rsa.rs b/src/rsa.rs index af9178161..cfc6511ba 100644 --- a/src/rsa.rs +++ b/src/rsa.rs @@ -1,8 +1,9 @@ -use pyo3::pyclass; use pyo3::pymethods; use pyo3::types::PyBytes; use pyo3::Python; +use pyo3::{pyclass, PyResult}; +use crate::CryptoError; use rsa::{sha2::Sha256, Oaep, RsaPrivateKey, RsaPublicKey}; #[pyclass(module = "qh3._hazmat")] @@ -14,41 +15,46 @@ pub struct Rsa { #[pymethods] impl Rsa { #[new] - pub fn py_new(key_size: usize) -> Self { + pub fn py_new(key_size: usize) -> PyResult { let mut rng = rand::thread_rng(); - let private_key = RsaPrivateKey::new(&mut rng, key_size).expect("failed to generate a key"); + let private_key = match RsaPrivateKey::new(&mut rng, key_size) { + Ok(key) => key, + Err(_) => return Err(CryptoError::new_err("Failed to generate RSA private key")), + }; + let public_key = RsaPublicKey::from(&private_key); - Rsa { + Ok(Rsa { public_key, private_key, - } + }) } - pub fn encrypt<'a>(&mut self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { + pub fn encrypt<'a>(&mut self, py: Python<'a>, data: &PyBytes) -> PyResult<&'a PyBytes> { let payload_to_enc = data.as_bytes(); let padding = Oaep::new::(); let mut rng = rand::thread_rng(); - let enc_data = self - .public_key - .encrypt(&mut rng, padding, payload_to_enc) - .expect("failed to encrypt"); + let enc_data = match self.public_key.encrypt(&mut rng, padding, payload_to_enc) { + Ok(data) => data, + Err(_) => return Err(CryptoError::new_err("Failed to encrypt data")), + }; - PyBytes::new(py, &enc_data) + Ok(PyBytes::new(py, &enc_data)) } - pub fn decrypt<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { + pub fn decrypt<'a>(&self, py: Python<'a>, data: &PyBytes) -> PyResult<&'a PyBytes> { let payload_to_dec = data.as_bytes(); let padding = Oaep::new::(); - let dec_data = self - .private_key - .decrypt(padding, payload_to_dec) - .expect("failed to decrypt"); - PyBytes::new(py, &dec_data) + let dec_data = match self.private_key.decrypt(padding, payload_to_dec) { + Ok(data) => data, + Err(_) => return Err(CryptoError::new_err("Failed to decrypt data")), + }; + + Ok(PyBytes::new(py, &dec_data)) } } diff --git a/tests/test_connection.py b/tests/test_connection.py index 276e5171b..dbbc50b91 100644 --- a/tests/test_connection.py +++ b/tests/test_connection.py @@ -48,7 +48,7 @@ CLIENT_HANDSHAKE_DATAGRAM_SIZES = [1280] SERVER_ADDR = ("2.3.4.5", 4433) -SERVER_INITIAL_DATAGRAM_SIZES = [1280, 1280, 986] +SERVER_INITIAL_DATAGRAM_SIZES = [1280, 1280, 890] HANDSHAKE_COMPLETED_EVENTS = [ events.HandshakeCompleted, @@ -423,7 +423,7 @@ def test_connect_without_loss(self): items = server.datagrams_to_send(now=now) assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES assert server.get_timer() == pytest.approx(0.25) - self.assertSentPackets(server, [2, 2, 0]) + self.assertSentPackets(server, [1, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) # handshake continues normally @@ -489,7 +489,7 @@ def test_connect_with_loss_1(self): items = server.datagrams_to_send(now=now) assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES assert server.get_timer() == pytest.approx(0.45) - self.assertSentPackets(server, [2, 2, 0]) + self.assertSentPackets(server, [1, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) # handshake continues normally @@ -548,7 +548,7 @@ def test_connect_with_loss_2(self): items = server.datagrams_to_send(now=now) assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES assert server.get_timer() == 0.25 - self.assertSentPackets(server, [2, 2, 0]) + self.assertSentPackets(server, [1, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) # client only receives second datagram, retransmits INITIAL @@ -628,7 +628,7 @@ def test_connect_with_loss_3(self): items = server.datagrams_to_send(now=now) assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES assert server.get_timer() == 0.25 - self.assertSentPackets(server, [2, 2, 0]) + self.assertSentPackets(server, [1, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) # INITIAL + HANDSHAKE are lost, client retransmits INITIAL @@ -647,7 +647,7 @@ def test_connect_with_loss_3(self): items = server.datagrams_to_send(now=now) assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES assert server.get_timer() == 0.45 - self.assertSentPackets(server, [2, 2, 0]) + self.assertSentPackets(server, [1, 2, 0]) self.assertEvents(server, []) # handshake continues normally @@ -703,7 +703,7 @@ def test_connect_with_loss_4(self): items = server.datagrams_to_send(now=now) assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES assert server.get_timer() == 0.25 - self.assertSentPackets(server, [2, 2, 0]) + self.assertSentPackets(server, [1, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) # client only receives the first datagram and sends ACKS @@ -789,7 +789,7 @@ def test_connect_with_loss_5(self): items = server.datagrams_to_send(now=now) assert datagram_sizes(items) == SERVER_INITIAL_DATAGRAM_SIZES assert server.get_timer() == 0.25 - self.assertSentPackets(server, [2, 2, 0]) + self.assertSentPackets(server, [1, 2, 0]) self.assertEvents(server, [events.ProtocolNegotiated]) # client receives INITIAL + HANDSHAKE diff --git a/tests/test_tls.py b/tests/test_tls.py index b37b7f39d..c5f32b881 100644 --- a/tests/test_tls.py +++ b/tests/test_tls.py @@ -539,7 +539,7 @@ def second_handshake(): server.handle_message(server_input, server_buf) assert server.state == State.SERVER_EXPECT_FINISHED client_input = merge_buffers(server_buf) - assert len(client_input) == 1410 + assert len(client_input) == 1314 reset_buffers(server_buf) # Handle server hello, encrypted extensions, certificate, @@ -622,7 +622,7 @@ def second_handshake_bad_pre_shared_key(): buf.seek(buf.tell() - 1) buf.push_uint8(1) client_input = merge_buffers(server_buf) - assert len(client_input) == 1410 + assert len(client_input) == 1314 reset_buffers(server_buf) # handle server hello and bomb From cd89c42a543daef12664197186052cc73410af80 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Tue, 31 Dec 2024 19:05:06 +0100 Subject: [PATCH 15/39] :sparkle: initial migration from pyo3 0.20.3 to 0.23.3 --- Cargo.lock | 79 +++++++++------------------------------------- Cargo.toml | 2 +- src/aead.rs | 57 +++++++++++++++++---------------- src/agreement.rs | 43 +++++++++++++++++-------- src/buffer.rs | 24 ++++++++------ src/certificate.rs | 44 ++++++++++++++------------ src/headers.rs | 40 +++++++++++++---------- src/hpk.rs | 11 +++++-- src/lib.rs | 4 +-- src/ocsp.rs | 16 ++++++---- src/pkcs8.rs | 10 ++++-- src/private_key.rs | 39 +++++++++++++---------- src/rsa.rs | 15 +++++++-- 13 files changed, 198 insertions(+), 186 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 43202f41b..68ee8e641 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -107,16 +107,15 @@ dependencies = [ [[package]] name = "aws-lc-sys" -version = "0.24.0" +version = "0.24.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8478a5c29ead3f3be14aff8a202ad965cf7da6856860041bfca271becf8ba48b" +checksum = "923ded50f602b3007e5e63e3f094c479d9c8a9b42d7f4034e4afe456aa48bfd2" dependencies = [ "bindgen 0.69.5", "cc", "cmake", "dunce", "fs_extra", - "libc", "paste", ] @@ -539,9 +538,9 @@ checksum = "a8d1add55171497b4705a648c6b583acafb01d58050a51727785f0b2c8e0a2b2" [[package]] name = "heck" -version = "0.4.1" +version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "95505c38b4572b2d910cecb0281560f54b440a19336cbbcb27bf6ce6adc6f5a8" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" [[package]] name = "hmac" @@ -644,16 +643,6 @@ version = "0.4.14" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "78b3ae25bc7c8c38cec158d1f2757ee79e9b3740fbc7ccf0e59e4b08d793fa89" -[[package]] -name = "lock_api" -version = "0.4.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "07af8b9cdd281b7915f413fa73f29ebd5d55d0d3f0155584dade1ff18cea1b17" -dependencies = [ - "autocfg", - "scopeguard", -] - [[package]] name = "log" version = "0.4.22" @@ -794,29 +783,6 @@ version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" -[[package]] -name = "parking_lot" -version = "0.12.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f1bf18183cf54e8d6059647fc3063646a1801cf30896933ec2311622cc4b9a27" -dependencies = [ - "lock_api", - "parking_lot_core", -] - -[[package]] -name = "parking_lot_core" -version = "0.9.10" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1e401f977ab385c9e4e3ab30627d6f26d00e2c73eef317493c4ec6d468726cf8" -dependencies = [ - "cfg-if", - "libc", - "redox_syscall", - "smallvec", - "windows-targets", -] - [[package]] name = "paste" version = "1.0.15" @@ -939,15 +905,15 @@ dependencies = [ [[package]] name = "pyo3" -version = "0.20.3" +version = "0.23.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "53bdbb96d49157e65d45cc287af5f32ffadd5f4761438b527b055fb0d4bb8233" +checksum = "e484fd2c8b4cb67ab05a318f1fd6fa8f199fcc30819f08f07d200809dba26c15" dependencies = [ "cfg-if", "indoc", "libc", "memoffset", - "parking_lot", + "once_cell", "portable-atomic", "pyo3-build-config", "pyo3-ffi", @@ -957,9 +923,9 @@ dependencies = [ [[package]] name = "pyo3-build-config" -version = "0.20.3" +version = "0.23.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "deaa5745de3f5231ce10517a1f5dd97d53e5a2fd77aa6b5842292085831d48d7" +checksum = "dc0e0469a84f208e20044b98965e1561028180219e35352a2afaf2b942beff3b" dependencies = [ "once_cell", "python3-dll-a", @@ -968,9 +934,9 @@ dependencies = [ [[package]] name = "pyo3-ffi" -version = "0.20.3" +version = "0.23.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "62b42531d03e08d4ef1f6e85a2ed422eb678b8cd62b762e53891c05faf0d4afa" +checksum = "eb1547a7f9966f6f1a0f0227564a9945fe36b90da5a93b3933fc3dc03fae372d" dependencies = [ "libc", "pyo3-build-config", @@ -978,9 +944,9 @@ dependencies = [ [[package]] name = "pyo3-macros" -version = "0.20.3" +version = "0.23.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7305c720fa01b8055ec95e484a6eca7a83c841267f0dd5280f0c8b8551d2c158" +checksum = "fdb6da8ec6fa5cedd1626c886fc8749bdcbb09424a86461eb8cdf096b7c33257" dependencies = [ "proc-macro2", "pyo3-macros-backend", @@ -990,9 +956,9 @@ dependencies = [ [[package]] name = "pyo3-macros-backend" -version = "0.20.3" +version = "0.23.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7c7e9b68bb9c3149c5b0cade5d07f953d6d125eb4337723c4ccdb665f1f96185" +checksum = "38a385202ff5a92791168b1136afae5059d3ac118457bb7bc304c197c2d33e7d" dependencies = [ "heck", "proc-macro2", @@ -1074,15 +1040,6 @@ dependencies = [ "getrandom", ] -[[package]] -name = "redox_syscall" -version = "0.5.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "03a862b389f93e68874fbf580b9de08dd02facb9a788ebadaf4a3fd33cf58834" -dependencies = [ - "bitflags", -] - [[package]] name = "regex" version = "1.11.1" @@ -1246,12 +1203,6 @@ dependencies = [ "cipher", ] -[[package]] -name = "scopeguard" -version = "1.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" - [[package]] name = "scrypt" version = "0.11.0" diff --git a/Cargo.toml b/Cargo.toml index f7c31f95a..87d249ba0 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,7 +13,7 @@ name = "qh3" crate-type = ["cdylib"] [dependencies] -pyo3 = { version = "0.20.3", features = ["extension-module", "abi3-py37", "generate-import-lib"] } +pyo3 = { version = "0.23.3", features = ["extension-module", "abi3-py37", "generate-import-lib"] } ls-qpack = "0.1.4" rustls = "0.23.12" x509-parser = "0.16.0" diff --git a/src/aead.rs b/src/aead.rs index 6de8438cb..e8f90493e 100644 --- a/src/aead.rs +++ b/src/aead.rs @@ -8,7 +8,8 @@ use chacha20poly1305::{aead::KeyInit, AeadInPlace, ChaCha20Poly1305, Key as ChaC use pyo3::pyclass; use pyo3::pymethods; use pyo3::types::PyBytes; -use pyo3::{PyResult, Python}; +use pyo3::{PyResult, Python, Bound}; +use pyo3::types::PyBytesMethods; use crate::CryptoError; @@ -30,7 +31,7 @@ pub struct AeadAes128Gcm { #[pymethods] impl AeadAes256Gcm { #[new] - pub fn py_new(key: &PyBytes) -> Self { + pub fn py_new(key: Bound<'_, PyBytes, >) -> Self { AeadAes256Gcm { key: key.as_bytes().to_vec(), } @@ -39,10 +40,10 @@ impl AeadAes256Gcm { pub fn decrypt<'a>( &mut self, py: Python<'a>, - nonce: &PyBytes, - data: &PyBytes, - associated_data: &PyBytes, - ) -> PyResult<&'a PyBytes> { + nonce: Bound<'_, PyBytes, >, + data: Bound<'_, PyBytes, >, + associated_data: Bound<'_, PyBytes, >, + )-> PyResult> { let mut in_out_buffer = data.as_bytes().to_vec(); let plaintext_len = in_out_buffer.len() - AES_256_GCM.tag_len(); @@ -69,10 +70,10 @@ impl AeadAes256Gcm { pub fn encrypt<'a>( &mut self, py: Python<'a>, - nonce: &PyBytes, - data: &PyBytes, - associated_data: &PyBytes, - ) -> PyResult<&'a PyBytes> { + nonce: Bound<'_, PyBytes, >, + data: Bound<'_, PyBytes, >, + associated_data: Bound<'_, PyBytes, >, + )-> PyResult> { let mut in_out_buffer = Vec::from(data.as_bytes()); let mut sealing_key: TlsRecordSealingKey = @@ -99,7 +100,7 @@ impl AeadAes256Gcm { #[pymethods] impl AeadAes128Gcm { #[new] - pub fn py_new(key: &PyBytes) -> Self { + pub fn py_new(key: Bound<'_, PyBytes, >) -> Self { AeadAes128Gcm { key: key.as_bytes().to_vec(), } @@ -108,10 +109,10 @@ impl AeadAes128Gcm { pub fn decrypt<'a>( &mut self, py: Python<'a>, - nonce: &PyBytes, - data: &PyBytes, - associated_data: &PyBytes, - ) -> PyResult<&'a PyBytes> { + nonce: Bound<'_, PyBytes, >, + data: Bound<'_, PyBytes, >, + associated_data: Bound<'_, PyBytes, >, + )-> PyResult> { let mut in_out_buffer = data.as_bytes().to_vec(); let plaintext_len = in_out_buffer.len() - AES_128_GCM.tag_len(); @@ -138,10 +139,10 @@ impl AeadAes128Gcm { pub fn encrypt<'a>( &mut self, py: Python<'a>, - nonce: &PyBytes, - data: &PyBytes, - associated_data: &PyBytes, - ) -> PyResult<&'a PyBytes> { + nonce: Bound<'_, PyBytes, >, + data: Bound<'_, PyBytes, >, + associated_data: Bound<'_, PyBytes, >, + )-> PyResult> { let mut in_out_buffer = Vec::from(data.as_bytes()); let mut sealing_key = @@ -168,7 +169,7 @@ impl AeadAes128Gcm { #[pymethods] impl AeadChaCha20Poly1305 { #[new] - pub fn py_new(key: &PyBytes) -> Self { + pub fn py_new(key: Bound<'_, PyBytes, >) -> Self { AeadChaCha20Poly1305 { key: key.as_bytes().to_vec(), } @@ -177,10 +178,10 @@ impl AeadChaCha20Poly1305 { pub fn decrypt<'a>( &mut self, py: Python<'a>, - nonce: &PyBytes, - data: &PyBytes, - associated_data: &PyBytes, - ) -> PyResult<&'a PyBytes> { + nonce: Bound<'_, PyBytes, >, + data: Bound<'_, PyBytes, >, + associated_data: Bound<'_, PyBytes, >, + )-> PyResult> { let mut in_out_buffer = data.as_bytes().to_vec(); let plaintext_len = in_out_buffer.len() - CHACHA20_POLY1305.tag_len(); @@ -201,10 +202,10 @@ impl AeadChaCha20Poly1305 { pub fn encrypt<'a>( &mut self, py: Python<'a>, - nonce: &PyBytes, - data: &PyBytes, - associated_data: &PyBytes, - ) -> PyResult<&'a PyBytes> { + nonce: Bound<'_, PyBytes, >, + data: Bound<'_, PyBytes, >, + associated_data: Bound<'_, PyBytes, >, + )-> PyResult> { let mut in_out_buffer = Vec::from(data.as_bytes()); let cipher: ChaCha20Poly1305 = ChaCha20Poly1305::new(ChaCha20Key::from_slice(&self.key)); diff --git a/src/agreement.rs b/src/agreement.rs index ad69aedac..a67cc4db1 100644 --- a/src/agreement.rs +++ b/src/agreement.rs @@ -5,10 +5,11 @@ use aws_lc_rs::kem::ML_KEM_768; use rustls::crypto::SharedSecret; use crate::CryptoError; -use pyo3::pymethods; use pyo3::types::PyBytes; +use pyo3::types::PyBytesMethods; use pyo3::Python; use pyo3::{pyclass, PyResult}; +use pyo3::{pymethods, Bound}; const X25519_LEN: usize = 32; const KYBER_CIPHERTEXT_LEN: usize = 1088; @@ -81,7 +82,7 @@ impl X25519ML768KeyExchange { }) } - pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult<&'a PyBytes> { + pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult> { let kyber_pub = match self.kyber768_decapsulation_key.encapsulation_key() { Ok(key) => key, Err(_) => { @@ -117,7 +118,7 @@ impl X25519ML768KeyExchange { Ok(PyBytes::new(py, combined_pub_key.as_ref())) } - pub fn shared_ciphertext<'a>(&mut self, py: Python<'a>) -> PyResult<&'a PyBytes> { + pub fn shared_ciphertext<'a>(&mut self, py: Python<'a>) -> PyResult> { if self.cipher_text.is_empty() { return Err(CryptoError::new_err( "You must receive client share first. Call exchange with client share.", @@ -146,8 +147,8 @@ impl X25519ML768KeyExchange { pub fn exchange<'a>( &mut self, py: Python<'a>, - peer_public_key: &PyBytes, - ) -> PyResult<&'a PyBytes> { + peer_public_key: Bound<'_, PyBytes>, + ) -> PyResult> { let cipher_text = peer_public_key.as_bytes(); // client share received @@ -242,7 +243,7 @@ impl X25519KeyExchange { Ok(X25519KeyExchange { private: x25519_pk }) } - pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult<&'a PyBytes> { + pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult> { let my_public_key = match self.private.compute_public_key() { Ok(key) => key, Err(_) => { @@ -255,7 +256,11 @@ impl X25519KeyExchange { Ok(PyBytes::new(py, my_public_key.as_ref())) } - pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn exchange<'a>( + &self, + py: Python<'a>, + peer_public_key: Bound<'_, PyBytes>, + ) -> PyResult> { let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::X25519, peer_public_key.as_bytes()); @@ -285,7 +290,7 @@ impl ECDHP256KeyExchange { Ok(ECDHP256KeyExchange { private: ecdh_key }) } - pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult<&'a PyBytes> { + pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult> { let my_public_key = match self.private.compute_public_key() { Ok(key) => key, Err(_) => { @@ -298,7 +303,11 @@ impl ECDHP256KeyExchange { Ok(PyBytes::new(py, my_public_key.as_ref())) } - pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn exchange<'a>( + &self, + py: Python<'a>, + peer_public_key: Bound<'_, PyBytes>, + ) -> PyResult> { let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::ECDH_P256, peer_public_key.as_bytes()); @@ -328,7 +337,7 @@ impl ECDHP384KeyExchange { Ok(ECDHP384KeyExchange { private: ecdh_key }) } - pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult<&'a PyBytes> { + pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult> { let my_public_key = match self.private.compute_public_key() { Ok(key) => key, Err(_) => { @@ -341,7 +350,11 @@ impl ECDHP384KeyExchange { Ok(PyBytes::new(py, my_public_key.as_ref())) } - pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn exchange<'a>( + &self, + py: Python<'a>, + peer_public_key: Bound<'_, PyBytes>, + ) -> PyResult> { let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::ECDH_P384, peer_public_key.as_bytes()); @@ -371,7 +384,7 @@ impl ECDHP521KeyExchange { Ok(ECDHP521KeyExchange { private: ecdh_pk }) } - pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult<&'a PyBytes> { + pub fn public_key<'a>(&self, py: Python<'a>) -> PyResult> { let my_public_key = match self.private.compute_public_key() { Ok(key) => key, Err(_) => { @@ -384,7 +397,11 @@ impl ECDHP521KeyExchange { Ok(PyBytes::new(py, my_public_key.as_ref())) } - pub fn exchange<'a>(&self, py: Python<'a>, peer_public_key: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn exchange<'a>( + &self, + py: Python<'a>, + peer_public_key: Bound<'_, PyBytes>, + ) -> PyResult> { let peer_public_key = agreement::UnparsedPublicKey::new(&agreement::ECDH_P521, peer_public_key.as_bytes()); diff --git a/src/buffer.rs b/src/buffer.rs index 47faa9cab..b617696f3 100644 --- a/src/buffer.rs +++ b/src/buffer.rs @@ -1,6 +1,7 @@ use pyo3::exceptions::PyValueError; -use pyo3::pyclass; use pyo3::types::PyBytes; +use pyo3::types::PyBytesMethods; +use pyo3::{pyclass, Bound}; use pyo3::{pymethods, PyResult, Python}; pyo3::create_exception!(_hazmat, BufferReadError, PyValueError); @@ -16,13 +17,13 @@ pub struct Buffer { #[pymethods] impl Buffer { #[new] - pub fn py_new(capacity: Option, data: Option<&PyBytes>) -> PyResult { + pub fn py_new(capacity: Option, data: Option>) -> PyResult { if data.is_some() { - let payload = data.unwrap().as_bytes(); + let payload = data.unwrap(); return Ok(Buffer { pos: 0, - data: payload.to_vec(), - capacity: payload.len() as u64, + data: payload.as_bytes().to_vec(), + capacity: payload.as_bytes().len() as u64, }); } @@ -45,14 +46,19 @@ impl Buffer { } #[getter] - pub fn data<'a>(&self, py: Python<'a>) -> &'a PyBytes { + pub fn data<'a>(&self, py: Python<'a>) -> Bound<'a, PyBytes> { if self.pos == 0 { return PyBytes::new(py, &[]); } PyBytes::new(py, &self.data[0_usize..self.pos as usize]) } - pub fn data_slice<'a>(&self, py: Python<'a>, start: u64, end: u64) -> PyResult<&'a PyBytes> { + pub fn data_slice<'a>( + &self, + py: Python<'a>, + start: u64, + end: u64, + ) -> PyResult> { if self.capacity < start || self.capacity < end || end < start { return Err(BufferReadError::new_err("Read out of bounds")); } @@ -78,7 +84,7 @@ impl Buffer { self.pos } - pub fn pull_bytes<'a>(&mut self, py: Python<'a>, length: u64) -> PyResult<&'a PyBytes> { + pub fn pull_bytes<'a>(&mut self, py: Python<'a>, length: u64) -> PyResult> { if self.capacity < self.pos + length { return Err(BufferReadError::new_err("Read out of bounds")); } @@ -189,7 +195,7 @@ impl Buffer { } } - pub fn push_bytes(&mut self, data: &PyBytes) -> PyResult<()> { + pub fn push_bytes(&mut self, data: Bound<'_, PyBytes>) -> PyResult<()> { let data_to_be_pushed = data.as_bytes(); let end_pos = self.pos + data_to_be_pushed.len() as u64; diff --git a/src/certificate.rs b/src/certificate.rs index 0c1595332..1962667b5 100644 --- a/src/certificate.rs +++ b/src/certificate.rs @@ -3,10 +3,12 @@ use rustls::client::WebPkiServerVerifier; use rustls::pki_types::{CertificateDer, ServerName, UnixTime}; use rustls::{CertificateError, Error, RootCertStore}; -use pyo3::pyclass; use pyo3::pymethods; +use pyo3::types::PyBytesMethods; +use pyo3::types::PyListMethods; use pyo3::types::{PyBytes, PyList, PyTuple, PyType}; use pyo3::ToPyObject; +use pyo3::{pyclass, Bound}; use pyo3::{PyResult, Python}; use x509_parser::prelude::*; @@ -56,7 +58,7 @@ pub struct Certificate { #[pymethods] impl Certificate { #[new] - pub fn py_new(certificate_der: &PyBytes) -> PyResult { + pub fn py_new(certificate_der: Bound<'_, PyBytes>) -> PyResult { let res = X509Certificate::from_der(certificate_der.as_bytes()); match res { @@ -138,7 +140,7 @@ impl Certificate { &self.serial_number } - pub fn raw_serial_number<'a>(&self, py: Python<'a>) -> &'a PyBytes { + pub fn raw_serial_number<'a>(&self, py: Python<'a>) -> Bound<'a, PyBytes> { PyBytes::new(py, &self.raw_serial_number) } @@ -158,7 +160,7 @@ impl Certificate { } #[getter] - pub fn subject<'a>(&self, py: Python<'a>) -> &'a PyList { + pub fn subject<'a>(&self, py: Python<'a>) -> Bound<'a, PyList> { let values = PyList::empty(py); for item in &self.subject { @@ -175,7 +177,7 @@ impl Certificate { _ => "".to_string(), }; - let _ = values.append(PyTuple::new( + let _ = values.append(PyTuple::new_bound( py, [ item.oid.to_object(py), @@ -189,7 +191,7 @@ impl Certificate { } #[getter] - pub fn issuer<'a>(&self, py: Python<'a>) -> &'a PyList { + pub fn issuer<'a>(&self, py: Python<'a>) -> Bound<'a, PyList> { let values = PyList::empty(py); for item in &self.issuer { @@ -206,7 +208,7 @@ impl Certificate { _ => "", }; - let _ = values.append(PyTuple::new( + let _ = values.append(PyTuple::new_bound( py, [ item.oid.to_object(py), @@ -219,7 +221,7 @@ impl Certificate { values } - pub fn get_subject_alt_names<'a>(&self, py: Python<'a>) -> &'a PyList { + pub fn get_subject_alt_names<'a>(&self, py: Python<'a>) -> Bound<'a, PyList> { let values = PyList::empty(py); for item in &self.extensions { @@ -231,7 +233,7 @@ impl Certificate { values } - pub fn get_ocsp_endpoints<'a>(&self, py: Python<'a>) -> &'a PyList { + pub fn get_ocsp_endpoints<'a>(&self, py: Python<'a>) -> Bound<'a, PyList> { let values = PyList::empty(py); for item in &self.extensions { @@ -243,7 +245,7 @@ impl Certificate { values } - pub fn get_issuer_endpoints<'a>(&self, py: Python<'a>) -> &'a PyList { + pub fn get_issuer_endpoints<'a>(&self, py: Python<'a>) -> Bound<'a, PyList> { let values = PyList::empty(py); for item in &self.extensions { @@ -255,11 +257,11 @@ impl Certificate { values } - pub fn public_bytes<'a>(&self, py: Python<'a>) -> &'a PyBytes { + pub fn public_bytes<'a>(&self, py: Python<'a>) -> Bound<'a, PyBytes> { PyBytes::new(py, &self.public_bytes) } - pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { + pub fn public_key<'a>(&self, py: Python<'a>) -> Bound<'a, PyBytes> { PyBytes::new(py, &self.public_key) } @@ -267,12 +269,12 @@ impl Certificate { self.serial_number == other.serial_number } - pub fn serialize<'py>(&self, py: Python<'py>) -> PyResult<&'py PyBytes> { + pub fn serialize<'py>(&self, py: Python<'py>) -> PyResult> { Ok(PyBytes::new(py, &serialize(&self).unwrap())) } #[classmethod] - pub fn deserialize(_cls: &PyType, encoded: &PyBytes) -> PyResult { + pub fn deserialize(_cls: Bound<'_, PyType>, encoded: Bound<'_, PyBytes>) -> PyResult { Ok(deserialize(encoded.as_bytes()).unwrap()) } } @@ -285,12 +287,12 @@ pub struct ServerVerifier { #[pymethods] impl ServerVerifier { #[new] - pub fn py_new(authorities: Vec<&PyBytes>) -> PyResult { + pub fn py_new(authorities: Vec>) -> PyResult { let mut root_cert_store = RootCertStore::empty(); root_cert_store.add_parsable_certificates( authorities - .into_iter() + .iter() .map(|ca| CertificateDer::from(ca.as_bytes())), ); let res = WebPkiServerVerifier::builder(Arc::new(root_cert_store)).build(); @@ -304,17 +306,17 @@ impl ServerVerifier { } #[allow(unreachable_code)] - pub fn verify( + pub fn verify<'a>( &mut self, - peer: &PyBytes, - intermediaries: Vec<&PyBytes>, + peer: Bound<'a, PyBytes>, + intermediaries: Vec>, server_name: String, ) -> PyResult<()> { let peer_der = CertificateDer::from(peer.as_bytes()); let mut intermediaries_der = Vec::new(); - for intermediary in intermediaries { - intermediaries_der.push(CertificateDer::from(intermediary.as_bytes())); + for intermediary in intermediaries.iter().map(|el| el.as_bytes()) { + intermediaries_der.push(CertificateDer::from(intermediary)); } let parsed_name_res = ServerName::try_from(server_name); diff --git a/src/headers.rs b/src/headers.rs index 105d744bd..b29a56498 100644 --- a/src/headers.rs +++ b/src/headers.rs @@ -2,9 +2,11 @@ use ls_qpack::decoder::{Decoder, DecoderOutput}; use ls_qpack::encoder::Encoder; use ls_qpack::StreamId; use pyo3::exceptions::PyException; -use pyo3::pyclass; use pyo3::pymethods; +use pyo3::types::PyBytesMethods; +use pyo3::types::PyListMethods; use pyo3::types::{PyBytes, PyList, PyTuple}; +use pyo3::{pyclass, Bound}; use pyo3::{PyResult, Python, ToPyObject}; pyo3::create_exception!(_hazmat, StreamBlocked, PyException); @@ -24,6 +26,8 @@ pub struct QpackEncoder { unsafe impl Send for QpackDecoder {} unsafe impl Send for QpackEncoder {} +unsafe impl Sync for QpackDecoder {} +unsafe impl Sync for QpackEncoder {} #[pymethods] impl QpackEncoder { @@ -40,7 +44,7 @@ impl QpackEncoder { max_table_capacity: u32, dyn_table_capacity: u32, blocked_streams: u32, - ) -> PyResult<&'a PyBytes> { + ) -> PyResult> { let r = match self .encoder @@ -53,7 +57,7 @@ impl QpackEncoder { Ok(PyBytes::new(py, r.data())) } - pub fn feed_decoder(&mut self, data: &PyBytes) -> PyResult<()> { + pub fn feed_decoder(&mut self, data: Bound<'_, PyBytes>) -> PyResult<()> { let res = self.encoder.feed(data.as_bytes()); match res { @@ -68,8 +72,8 @@ impl QpackEncoder { &mut self, py: Python<'a>, stream_id: u64, - headers: Vec<(&PyBytes, &PyBytes)>, - ) -> PyResult<&'a PyTuple> { + headers: Vec<(Bound<'_, PyBytes>, Bound<'_, PyBytes>)>, + ) -> PyResult> { let mut decoded_vec: Vec<(String, String)> = Vec::new(); for (header, value) in headers.iter() { @@ -89,7 +93,7 @@ impl QpackEncoder { let stream_data = PyBytes::new(py, buffer.stream()); - Ok(PyTuple::new(py, [stream_data, encoded_buffer])) + Ok(PyTuple::new(py, [stream_data, encoded_buffer]).unwrap()) } Err(abc) => Err(EncoderStreamError::new_err(format!( "unable to encode headers {:?}", @@ -108,7 +112,7 @@ impl QpackDecoder { } } - pub fn feed_encoder(&mut self, data: &PyBytes) -> PyResult<()> { + pub fn feed_encoder(&mut self, data: Bound<'_, PyBytes>) -> PyResult<()> { let res = self.decoder.feed(data.as_bytes()); match res { @@ -123,18 +127,18 @@ impl QpackDecoder { &mut self, py: Python<'a>, stream_id: u64, - data: &PyBytes, - ) -> PyResult<&'a PyTuple> { + data: Bound<'_, PyBytes>, + ) -> PyResult> { let output = self .decoder .decode(StreamId::new(stream_id), data.as_bytes()); match output { Ok(DecoderOutput::Done(ref buffer)) => { - let decoded_headers = PyList::new(py, Vec::<(String, String)>::new()); + let decoded_headers = PyList::new(py, Vec::<(String, String)>::new()).unwrap(); for header in buffer.headers() { - let _ = decoded_headers.append(PyTuple::new( + let _ = decoded_headers.append(PyTuple::new_bound( py, [ PyBytes::new(py, header.name().as_bytes()), @@ -143,7 +147,7 @@ impl QpackDecoder { )); } - Ok(PyTuple::new( + Ok(PyTuple::new_bound( py, [ PyBytes::new(py, buffer.stream()).to_object(py), @@ -160,7 +164,11 @@ impl QpackDecoder { } } - pub fn resume_header<'a>(&mut self, py: Python<'a>, stream_id: u64) -> PyResult<&'a PyTuple> { + pub fn resume_header<'a>( + &mut self, + py: Python<'a>, + stream_id: u64, + ) -> PyResult> { let output = self.decoder.unblocked(StreamId::new(stream_id)); if output.is_none() { @@ -171,10 +179,10 @@ impl QpackDecoder { match res { Ok(DecoderOutput::Done(ref buffer)) => { - let decoded_headers = PyList::new(py, Vec::<(String, String)>::new()); + let decoded_headers = PyList::new(py, Vec::<(String, String)>::new()).unwrap(); for header in buffer.headers() { - let _ = decoded_headers.append(PyTuple::new( + let _ = decoded_headers.append(PyTuple::new_bound( py, [ PyBytes::new(py, header.name().as_bytes()), @@ -183,7 +191,7 @@ impl QpackDecoder { )); } - Ok(PyTuple::new( + Ok(PyTuple::new_bound( py, [ PyBytes::new(py, buffer.stream()).to_object(py), diff --git a/src/hpk.rs b/src/hpk.rs index c67c1ce0e..1b0166d9c 100644 --- a/src/hpk.rs +++ b/src/hpk.rs @@ -1,9 +1,10 @@ use aws_lc_rs::aead::quic::{HeaderProtectionKey, AES_128, AES_256, CHACHA20}; use crate::CryptoError; -use pyo3::pyclass; use pyo3::pymethods; use pyo3::types::PyBytes; +use pyo3::types::PyBytesMethods; +use pyo3::{pyclass, Bound}; use pyo3::{PyResult, Python}; #[pyclass(module = "qh3._hazmat")] @@ -14,7 +15,7 @@ pub struct QUICHeaderProtection { #[pymethods] impl QUICHeaderProtection { #[new] - pub fn py_new(key: &PyBytes, algorithm: u16) -> PyResult { + pub fn py_new(key: Bound<'_, PyBytes>, algorithm: u16) -> PyResult { let inner_hpk = match HeaderProtectionKey::new( match algorithm { 128 => &AES_128, @@ -35,7 +36,11 @@ impl QUICHeaderProtection { Ok(QUICHeaderProtection { hpk: inner_hpk }) } - pub fn mask<'a>(&self, py: Python<'a>, sample: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn mask<'a>( + &self, + py: Python<'a>, + sample: Bound<'_, PyBytes>, + ) -> PyResult> { let res = self.hpk.new_mask(sample.as_bytes()); match res { diff --git a/src/lib.rs b/src/lib.rs index 42289c466..55db36afe 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -37,8 +37,8 @@ pub use self::rsa::Rsa; pyo3::create_exception!(_hazmat, CryptoError, PyException); -#[pymodule] -fn _hazmat(py: Python, m: &PyModule) -> PyResult<()> { +#[pymodule(gil_used = false)] +fn _hazmat(py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> { // ls-qpack bridge m.add_class::()?; m.add_class::()?; diff --git a/src/ocsp.rs b/src/ocsp.rs index 155afce2e..ecfa68641 100644 --- a/src/ocsp.rs +++ b/src/ocsp.rs @@ -1,9 +1,10 @@ // OCSP Response Parser and Request Builder // This module is created for Niquests // qh3 has no use for it and we won't implement it for this package -use pyo3::pyclass; use pyo3::pymethods; +use pyo3::types::PyBytesMethods; use pyo3::types::{PyBytes, PyType}; +use pyo3::{pyclass, Bound}; use pyo3::{PyResult, Python}; use der::{Decode, Encode}; @@ -69,7 +70,7 @@ pub struct OCSPResponse { #[pymethods] impl OCSPResponse { #[new] - pub fn py_new(raw_response: &PyBytes) -> PyResult { + pub fn py_new(raw_response: Bound<'_, PyBytes>) -> PyResult { let ocsp_res: OcspResponse = match OcspResponse::from_der(raw_response.as_bytes()) { Ok(ocsp_res) => ocsp_res, Err(_) => return Err(PyValueError::new_err("OCSP DER given is invalid")), @@ -153,12 +154,12 @@ impl OCSPResponse { self.revocation_reason } - pub fn serialize<'py>(&self, py: Python<'py>) -> PyResult<&'py PyBytes> { + pub fn serialize<'py>(&self, py: Python<'py>) -> PyResult> { Ok(PyBytes::new(py, &serialize(&self).unwrap())) } #[classmethod] - pub fn deserialize(_cls: &PyType, encoded: &PyBytes) -> PyResult { + pub fn deserialize(_cls: Bound<'_, PyType>, encoded: Bound<'_, PyBytes>) -> PyResult { Ok(deserialize(encoded.as_bytes()).unwrap()) } } @@ -171,7 +172,10 @@ pub struct OCSPRequest { #[pymethods] impl OCSPRequest { #[new] - pub fn py_new(peer_certificate: &PyBytes, issuer_certificate: &PyBytes) -> PyResult { + pub fn py_new( + peer_certificate: Bound<'_, PyBytes>, + issuer_certificate: Bound<'_, PyBytes>, + ) -> PyResult { let issuer = Certificate::from_der(issuer_certificate.as_bytes()).unwrap(); let cert = Certificate::from_der(peer_certificate.as_bytes()).unwrap(); @@ -187,7 +191,7 @@ impl OCSPRequest { } } - pub fn public_bytes<'a>(&self, py: Python<'a>) -> &'a PyBytes { + pub fn public_bytes<'a>(&self, py: Python<'a>) -> Bound<'a, PyBytes> { PyBytes::new(py, &self.inner_request) } } diff --git a/src/pkcs8.rs b/src/pkcs8.rs index 01fbfd422..801422248 100644 --- a/src/pkcs8.rs +++ b/src/pkcs8.rs @@ -1,7 +1,8 @@ -use pyo3::pymethods; use pyo3::types::PyBytes; +use pyo3::types::PyBytesMethods; use pyo3::Python; use pyo3::{pyclass, PyResult}; +use pyo3::{pymethods, Bound}; use pkcs8::{der::Encode, DecodePrivateKey, Error, PrivateKeyInfo as InternalPrivateKeyInfo}; use rsa::{ @@ -68,7 +69,10 @@ impl TryFrom> for PrivateKeyInfo { #[pymethods] impl PrivateKeyInfo { #[new] - pub fn py_new(raw_pem_content: &PyBytes, password: Option<&PyBytes>) -> PyResult { + pub fn py_new( + raw_pem_content: Bound<'_, PyBytes>, + password: Option>, + ) -> PyResult { let pem_content = raw_pem_content.as_bytes(); let decoded_bytes = std::str::from_utf8(pem_content)?; @@ -142,7 +146,7 @@ impl PrivateKeyInfo { self.cert_type } - pub fn public_bytes<'a>(&self, py: Python<'a>) -> &'a PyBytes { + pub fn public_bytes<'a>(&self, py: Python<'a>) -> Bound<'a, PyBytes> { PyBytes::new(py, &self.der_encoded) } } diff --git a/src/private_key.rs b/src/private_key.rs index 769d3f491..d9b453322 100644 --- a/src/private_key.rs +++ b/src/private_key.rs @@ -28,10 +28,11 @@ use aws_lc_rs::signature::UnparsedPublicKey; use crate::CryptoError; use pyo3::exceptions::PyException; -use pyo3::pyclass; use pyo3::pyfunction; use pyo3::pymethods; use pyo3::types::PyBytes; +use pyo3::types::PyBytesMethods; +use pyo3::{pyclass, Bound}; use pyo3::{PyResult, Python}; pyo3::create_exception!(_hazmat, SignatureError, PyException); @@ -60,7 +61,7 @@ pub struct RsaPrivateKey { #[pymethods] impl Ed25519PrivateKey { #[new] - pub fn py_new(pkcs8: &PyBytes) -> PyResult { + pub fn py_new(pkcs8: Bound<'_, PyBytes>) -> PyResult { let pk = match InternalEd25519PrivateKey::from_pkcs8(pkcs8.as_bytes()) { Ok(key) => key, Err(_) => return Err(CryptoError::new_err("Invalid Ed25519 PrivateKey")), @@ -69,13 +70,13 @@ impl Ed25519PrivateKey { Ok(Ed25519PrivateKey { inner: pk }) } - pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { + pub fn sign<'a>(&self, py: Python<'a>, data: Bound<'_, PyBytes>) -> Bound<'a, PyBytes> { let signature = self.inner.sign(data.as_bytes()); PyBytes::new(py, signature.as_ref()) } - pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { + pub fn public_key<'a>(&self, py: Python<'a>) -> Bound<'a, PyBytes> { PyBytes::new(py, self.inner.public_key().as_ref()) } } @@ -83,7 +84,7 @@ impl Ed25519PrivateKey { #[pymethods] impl EcPrivateKey { #[new] - pub fn py_new(pkcs8: &PyBytes, curve_type: u32) -> PyResult { + pub fn py_new(pkcs8: Bound<'_, PyBytes>, curve_type: u32) -> PyResult { let signing_algorithm = match curve_type { 256 => &ECDSA_P256_SHA256_ASN1_SIGNING, 384 => &ECDSA_P384_SHA384_ASN1_SIGNING, @@ -106,7 +107,11 @@ impl EcPrivateKey { }) } - pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn sign<'a>( + &self, + py: Python<'a>, + data: Bound<'_, PyBytes>, + ) -> PyResult> { let rng = SystemRandom::new(); let signature = match self.inner.sign(&rng, data.as_bytes()) { Ok(signature) => signature, @@ -116,7 +121,7 @@ impl EcPrivateKey { Ok(PyBytes::new(py, signature.as_ref())) } - pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { + pub fn public_key<'a>(&self, py: Python<'a>) -> Bound<'a, PyBytes> { PyBytes::new(py, self.inner.public_key().as_ref()) } @@ -129,7 +134,7 @@ impl EcPrivateKey { #[pymethods] impl DsaPrivateKey { #[new] - pub fn py_new(pkcs8: &PyBytes) -> PyResult { + pub fn py_new(pkcs8: Bound<'_, PyBytes>) -> PyResult { let pk = match InternalDsaPrivateKey::from_pkcs8_der(pkcs8.as_bytes()) { Ok(key) => key, Err(_) => return Err(CryptoError::new_err("Invalid Dsa PrivateKey")), @@ -138,13 +143,13 @@ impl DsaPrivateKey { Ok(DsaPrivateKey { inner: pk }) } - pub fn sign<'a>(&self, py: Python<'a>, data: &PyBytes) -> &'a PyBytes { + pub fn sign<'a>(&self, py: Python<'a>, data: Bound<'_, PyBytes>) -> Bound<'a, PyBytes> { let signature = self.inner.sign(data.as_bytes()); PyBytes::new(py, &signature.to_bytes()) } - pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { + pub fn public_key<'a>(&self, py: Python<'a>) -> Bound<'a, PyBytes> { PyBytes::new( py, self.inner @@ -159,7 +164,7 @@ impl DsaPrivateKey { #[pymethods] impl RsaPrivateKey { #[new] - pub fn py_new(pkcs8: &PyBytes) -> PyResult { + pub fn py_new(pkcs8: Bound<'_, PyBytes>) -> PyResult { let pk = match InternalRsaPrivateKey::from_pkcs8_der(pkcs8.as_bytes()) { Ok(key) => key, Err(_) => return Err(CryptoError::new_err("Invalid Rsa PrivateKey")), @@ -171,10 +176,10 @@ impl RsaPrivateKey { pub fn sign<'a>( &self, py: Python<'a>, - data: &PyBytes, + data: Bound<'_, PyBytes>, is_pss_padding: bool, hash_size: u32, - ) -> PyResult<&'a PyBytes> { + ) -> PyResult> { let private_key = self.inner.clone(); match is_pss_padding { @@ -215,7 +220,7 @@ impl RsaPrivateKey { } } - pub fn public_key<'a>(&self, py: Python<'a>) -> &'a PyBytes { + pub fn public_key<'a>(&self, py: Python<'a>) -> Bound<'a, PyBytes> { let public_key: InternalRsaPublicKey = self.inner.to_public_key(); PyBytes::new( @@ -228,10 +233,10 @@ impl RsaPrivateKey { #[pyfunction] #[allow(unreachable_code)] pub fn verify_with_public_key( - public_key_raw: &PyBytes, + public_key_raw: Bound<'_, PyBytes>, algorithm: u32, - message: &PyBytes, - signature: &PyBytes, + message: Bound<'_, PyBytes>, + signature: Bound<'_, PyBytes>, ) -> PyResult<()> { let pss_rsae_blind_signature = 0x0804..0x0806; let pss_pss_blind_signature = 0x0809..0x080B; diff --git a/src/rsa.rs b/src/rsa.rs index cfc6511ba..4f7479fdf 100644 --- a/src/rsa.rs +++ b/src/rsa.rs @@ -1,7 +1,8 @@ -use pyo3::pymethods; use pyo3::types::PyBytes; +use pyo3::types::PyBytesMethods; use pyo3::Python; use pyo3::{pyclass, PyResult}; +use pyo3::{pymethods, Bound}; use crate::CryptoError; use rsa::{sha2::Sha256, Oaep, RsaPrivateKey, RsaPublicKey}; @@ -31,7 +32,11 @@ impl Rsa { }) } - pub fn encrypt<'a>(&mut self, py: Python<'a>, data: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn encrypt<'a>( + &mut self, + py: Python<'a>, + data: Bound<'_, PyBytes>, + ) -> PyResult> { let payload_to_enc = data.as_bytes(); let padding = Oaep::new::(); @@ -45,7 +50,11 @@ impl Rsa { Ok(PyBytes::new(py, &enc_data)) } - pub fn decrypt<'a>(&self, py: Python<'a>, data: &PyBytes) -> PyResult<&'a PyBytes> { + pub fn decrypt<'a>( + &self, + py: Python<'a>, + data: Bound<'_, PyBytes>, + ) -> PyResult> { let payload_to_dec = data.as_bytes(); let padding = Oaep::new::(); From ebe8b75d74f5a9db404782461fd5b320e7e3f562 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Tue, 31 Dec 2024 19:05:58 +0100 Subject: [PATCH 16/39] :art: reformat aead.rs --- src/aead.rs | 56 ++++++++++++++++++++++++++--------------------------- 1 file changed, 28 insertions(+), 28 deletions(-) diff --git a/src/aead.rs b/src/aead.rs index e8f90493e..142300666 100644 --- a/src/aead.rs +++ b/src/aead.rs @@ -8,8 +8,8 @@ use chacha20poly1305::{aead::KeyInit, AeadInPlace, ChaCha20Poly1305, Key as ChaC use pyo3::pyclass; use pyo3::pymethods; use pyo3::types::PyBytes; -use pyo3::{PyResult, Python, Bound}; use pyo3::types::PyBytesMethods; +use pyo3::{Bound, PyResult, Python}; use crate::CryptoError; @@ -31,7 +31,7 @@ pub struct AeadAes128Gcm { #[pymethods] impl AeadAes256Gcm { #[new] - pub fn py_new(key: Bound<'_, PyBytes, >) -> Self { + pub fn py_new(key: Bound<'_, PyBytes>) -> Self { AeadAes256Gcm { key: key.as_bytes().to_vec(), } @@ -40,10 +40,10 @@ impl AeadAes256Gcm { pub fn decrypt<'a>( &mut self, py: Python<'a>, - nonce: Bound<'_, PyBytes, >, - data: Bound<'_, PyBytes, >, - associated_data: Bound<'_, PyBytes, >, - )-> PyResult> { + nonce: Bound<'_, PyBytes>, + data: Bound<'_, PyBytes>, + associated_data: Bound<'_, PyBytes>, + ) -> PyResult> { let mut in_out_buffer = data.as_bytes().to_vec(); let plaintext_len = in_out_buffer.len() - AES_256_GCM.tag_len(); @@ -70,10 +70,10 @@ impl AeadAes256Gcm { pub fn encrypt<'a>( &mut self, py: Python<'a>, - nonce: Bound<'_, PyBytes, >, - data: Bound<'_, PyBytes, >, - associated_data: Bound<'_, PyBytes, >, - )-> PyResult> { + nonce: Bound<'_, PyBytes>, + data: Bound<'_, PyBytes>, + associated_data: Bound<'_, PyBytes>, + ) -> PyResult> { let mut in_out_buffer = Vec::from(data.as_bytes()); let mut sealing_key: TlsRecordSealingKey = @@ -100,7 +100,7 @@ impl AeadAes256Gcm { #[pymethods] impl AeadAes128Gcm { #[new] - pub fn py_new(key: Bound<'_, PyBytes, >) -> Self { + pub fn py_new(key: Bound<'_, PyBytes>) -> Self { AeadAes128Gcm { key: key.as_bytes().to_vec(), } @@ -109,10 +109,10 @@ impl AeadAes128Gcm { pub fn decrypt<'a>( &mut self, py: Python<'a>, - nonce: Bound<'_, PyBytes, >, - data: Bound<'_, PyBytes, >, - associated_data: Bound<'_, PyBytes, >, - )-> PyResult> { + nonce: Bound<'_, PyBytes>, + data: Bound<'_, PyBytes>, + associated_data: Bound<'_, PyBytes>, + ) -> PyResult> { let mut in_out_buffer = data.as_bytes().to_vec(); let plaintext_len = in_out_buffer.len() - AES_128_GCM.tag_len(); @@ -139,10 +139,10 @@ impl AeadAes128Gcm { pub fn encrypt<'a>( &mut self, py: Python<'a>, - nonce: Bound<'_, PyBytes, >, - data: Bound<'_, PyBytes, >, - associated_data: Bound<'_, PyBytes, >, - )-> PyResult> { + nonce: Bound<'_, PyBytes>, + data: Bound<'_, PyBytes>, + associated_data: Bound<'_, PyBytes>, + ) -> PyResult> { let mut in_out_buffer = Vec::from(data.as_bytes()); let mut sealing_key = @@ -169,7 +169,7 @@ impl AeadAes128Gcm { #[pymethods] impl AeadChaCha20Poly1305 { #[new] - pub fn py_new(key: Bound<'_, PyBytes, >) -> Self { + pub fn py_new(key: Bound<'_, PyBytes>) -> Self { AeadChaCha20Poly1305 { key: key.as_bytes().to_vec(), } @@ -178,10 +178,10 @@ impl AeadChaCha20Poly1305 { pub fn decrypt<'a>( &mut self, py: Python<'a>, - nonce: Bound<'_, PyBytes, >, - data: Bound<'_, PyBytes, >, - associated_data: Bound<'_, PyBytes, >, - )-> PyResult> { + nonce: Bound<'_, PyBytes>, + data: Bound<'_, PyBytes>, + associated_data: Bound<'_, PyBytes>, + ) -> PyResult> { let mut in_out_buffer = data.as_bytes().to_vec(); let plaintext_len = in_out_buffer.len() - CHACHA20_POLY1305.tag_len(); @@ -202,10 +202,10 @@ impl AeadChaCha20Poly1305 { pub fn encrypt<'a>( &mut self, py: Python<'a>, - nonce: Bound<'_, PyBytes, >, - data: Bound<'_, PyBytes, >, - associated_data: Bound<'_, PyBytes, >, - )-> PyResult> { + nonce: Bound<'_, PyBytes>, + data: Bound<'_, PyBytes>, + associated_data: Bound<'_, PyBytes>, + ) -> PyResult> { let mut in_out_buffer = Vec::from(data.as_bytes()); let cipher: ChaCha20Poly1305 = ChaCha20Poly1305::new(ChaCha20Key::from_slice(&self.key)); From fda006b6c2299f2b9260bb4a77274643b90da0d0 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Tue, 31 Dec 2024 19:19:39 +0100 Subject: [PATCH 17/39] =?UTF-8?q?=F0=9F=9A=A8=20fix=20deprecation=20warnin?= =?UTF-8?q?gs=20excepted=20ToPyObject?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/buffer.rs | 1 + src/certificate.rs | 38 ++++++++++++++++++++++---------------- src/headers.rs | 44 ++++++++++++++++++++++++++------------------ src/ocsp.rs | 12 ++++++------ src/pkcs8.rs | 5 +++-- 5 files changed, 58 insertions(+), 42 deletions(-) diff --git a/src/buffer.rs b/src/buffer.rs index b617696f3..14e37c56b 100644 --- a/src/buffer.rs +++ b/src/buffer.rs @@ -17,6 +17,7 @@ pub struct Buffer { #[pymethods] impl Buffer { #[new] + #[pyo3(signature = (capacity=None, data=None))] pub fn py_new(capacity: Option, data: Option>) -> PyResult { if data.is_some() { let payload = data.unwrap(); diff --git a/src/certificate.rs b/src/certificate.rs index 1962667b5..3aa9aa4c9 100644 --- a/src/certificate.rs +++ b/src/certificate.rs @@ -177,14 +177,17 @@ impl Certificate { _ => "".to_string(), }; - let _ = values.append(PyTuple::new_bound( - py, - [ - item.oid.to_object(py), - oid_short.to_object(py), - PyBytes::new(py, &item.value).into(), - ], - )); + let _ = values.append( + PyTuple::new( + py, + [ + item.oid.to_object(py), + oid_short.to_object(py), + PyBytes::new(py, &item.value).into(), + ], + ) + .unwrap(), + ); } values @@ -208,14 +211,17 @@ impl Certificate { _ => "", }; - let _ = values.append(PyTuple::new_bound( - py, - [ - item.oid.to_object(py), - oid_short.to_object(py), - PyBytes::new(py, &item.value).into(), - ], - )); + let _ = values.append( + PyTuple::new( + py, + [ + item.oid.to_object(py), + oid_short.to_object(py), + PyBytes::new(py, &item.value).into(), + ], + ) + .unwrap(), + ); } values diff --git a/src/headers.rs b/src/headers.rs index b29a56498..f1076458d 100644 --- a/src/headers.rs +++ b/src/headers.rs @@ -138,22 +138,26 @@ impl QpackDecoder { let decoded_headers = PyList::new(py, Vec::<(String, String)>::new()).unwrap(); for header in buffer.headers() { - let _ = decoded_headers.append(PyTuple::new_bound( - py, - [ - PyBytes::new(py, header.name().as_bytes()), - PyBytes::new(py, header.value().as_bytes()), - ], - )); + let _ = decoded_headers.append( + PyTuple::new( + py, + [ + PyBytes::new(py, header.name().as_bytes()), + PyBytes::new(py, header.value().as_bytes()), + ], + ) + .unwrap(), + ); } - Ok(PyTuple::new_bound( + Ok(PyTuple::new( py, [ PyBytes::new(py, buffer.stream()).to_object(py), decoded_headers.to_object(py), ], - )) + ) + .unwrap()) } Ok(DecoderOutput::BlockedStream) => Err(StreamBlocked::new_err( "stream is blocked, need more data to pursue decoding", @@ -182,22 +186,26 @@ impl QpackDecoder { let decoded_headers = PyList::new(py, Vec::<(String, String)>::new()).unwrap(); for header in buffer.headers() { - let _ = decoded_headers.append(PyTuple::new_bound( - py, - [ - PyBytes::new(py, header.name().as_bytes()), - PyBytes::new(py, header.value().as_bytes()), - ], - )); + let _ = decoded_headers.append( + PyTuple::new( + py, + [ + PyBytes::new(py, header.name().as_bytes()), + PyBytes::new(py, header.value().as_bytes()), + ], + ) + .unwrap(), + ); } - Ok(PyTuple::new_bound( + Ok(PyTuple::new( py, [ PyBytes::new(py, buffer.stream()).to_object(py), decoded_headers.to_object(py), ], - )) + ) + .unwrap()) } Ok(DecoderOutput::BlockedStream) => Err(StreamBlocked::new_err( "stream is blocked, need more data to pursue decoding", diff --git a/src/ocsp.rs b/src/ocsp.rs index ecfa68641..88988f435 100644 --- a/src/ocsp.rs +++ b/src/ocsp.rs @@ -20,8 +20,8 @@ use bincode::{deserialize, serialize}; use serde::{Deserialize, Serialize}; use sha1::Sha1; -#[pyclass(module = "qh3._hazmat")] -#[derive(Clone, Copy, Serialize, Deserialize)] +#[pyclass(module = "qh3._hazmat", eq, eq_int)] +#[derive(Clone, Copy, Serialize, Deserialize, PartialEq)] #[allow(non_camel_case_types)] pub enum ReasonFlags { unspecified = 0, @@ -36,8 +36,8 @@ pub enum ReasonFlags { remove_from_crl = 8, } -#[pyclass(module = "qh3._hazmat")] -#[derive(Clone, Copy, Serialize, Deserialize)] +#[pyclass(module = "qh3._hazmat", eq, eq_int)] +#[derive(Clone, Copy, Serialize, Deserialize, PartialEq)] #[allow(non_camel_case_types)] pub enum OCSPResponseStatus { SUCCESSFUL = 0, @@ -48,8 +48,8 @@ pub enum OCSPResponseStatus { UNAUTHORIZED = 6, } -#[pyclass(module = "qh3._hazmat")] -#[derive(Clone, Copy, Serialize, Deserialize)] +#[pyclass(module = "qh3._hazmat", eq, eq_int)] +#[derive(Clone, Copy, Serialize, Deserialize, PartialEq)] #[allow(non_camel_case_types)] pub enum OCSPCertStatus { GOOD = 0, diff --git a/src/pkcs8.rs b/src/pkcs8.rs index 801422248..69e774926 100644 --- a/src/pkcs8.rs +++ b/src/pkcs8.rs @@ -14,8 +14,8 @@ use rsa::{ use crate::CryptoError; use rustls_pemfile::{read_one_from_slice, Item}; -#[pyclass(module = "qh3._hazmat")] -#[derive(Clone, Copy)] +#[pyclass(module = "qh3._hazmat", eq, eq_int)] +#[derive(Clone, Copy, PartialEq)] #[allow(non_camel_case_types)] pub enum KeyType { ECDSA_P256, @@ -69,6 +69,7 @@ impl TryFrom> for PrivateKeyInfo { #[pymethods] impl PrivateKeyInfo { #[new] + #[pyo3(signature = (raw_pem_content, password=None))] pub fn py_new( raw_pem_content: Bound<'_, PyBytes>, password: Option>, From 06aab5b721d93c19f9d48a590ed4ad1bf09add0c Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Tue, 31 Dec 2024 19:21:56 +0100 Subject: [PATCH 18/39] :pencil: update changelog for pyo3 upgrade --- CHANGELOG.rst | 1 + 1 file changed, 1 insertion(+) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index a27342724..fb10f3949 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -5,6 +5,7 @@ - Post-Quantum key-exchange Kyber 768 Draft upgraded to standard Module-Lattice 768. - Version negotiation no longer logged as ``INFO``. Every logs generated will always be ``DEBUG`` level. - Converted our test suite to run on Pytest instead of unittest. +- Migrated pyo3 from 0.20.3 to 0.23.3 and fixed most of the deprecation warnings. **Fixed** - Clippy warnings in our Rust code. From f59768fb7047d8d4f4055813b511cd9f48e5afcc Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 05:49:47 +0100 Subject: [PATCH 19/39] =?UTF-8?q?=F0=9F=9A=A8=20fix=20into=20pyobject=20re?= =?UTF-8?q?maining=20deprecation=20warning?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Cargo.toml | 4 ++-- src/certificate.rs | 15 +++++++-------- src/headers.rs | 10 +++++----- 3 files changed, 14 insertions(+), 15 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 87d249ba0..67e0342f1 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -15,7 +15,7 @@ crate-type = ["cdylib"] [dependencies] pyo3 = { version = "0.23.3", features = ["extension-module", "abi3-py37", "generate-import-lib"] } ls-qpack = "0.1.4" -rustls = "0.23.12" +rustls = "0.23.20" x509-parser = "0.16.0" rsa = { version = "0.9.6", features = ["sha2", "pem", "getrandom"] } dsa = "0.6.3" @@ -24,7 +24,7 @@ rand = "0.8.5" chacha20poly1305 = "0.10.1" pkcs8 = { version = "0.10.2", features = ["encryption", "pem"] } pkcs1 = { version = "0.7.5", features = ["pem"] } -rustls-pemfile = "2.1.2" +rustls-pemfile = "2.2.0" aws-lc-rs = { version = "1.12.0", features=["bindgen"], default-features = false } x509-ocsp = { version = "0.2.1", features = ["builder"] } x509-cert = "0.2.5" diff --git a/src/certificate.rs b/src/certificate.rs index 3aa9aa4c9..017c2a55a 100644 --- a/src/certificate.rs +++ b/src/certificate.rs @@ -3,11 +3,10 @@ use rustls::client::WebPkiServerVerifier; use rustls::pki_types::{CertificateDer, ServerName, UnixTime}; use rustls::{CertificateError, Error, RootCertStore}; -use pyo3::pymethods; +use pyo3::{pymethods, IntoPyObject}; use pyo3::types::PyBytesMethods; use pyo3::types::PyListMethods; use pyo3::types::{PyBytes, PyList, PyTuple, PyType}; -use pyo3::ToPyObject; use pyo3::{pyclass, Bound}; use pyo3::{PyResult, Python}; @@ -181,9 +180,9 @@ impl Certificate { PyTuple::new( py, [ - item.oid.to_object(py), - oid_short.to_object(py), - PyBytes::new(py, &item.value).into(), + item.clone().oid.into_pyobject(py).unwrap().into_any(), + oid_short.into_pyobject(py).unwrap().into_any(), + PyBytes::new(py, &item.value).into_pyobject(py).unwrap().into_any(), ], ) .unwrap(), @@ -215,9 +214,9 @@ impl Certificate { PyTuple::new( py, [ - item.oid.to_object(py), - oid_short.to_object(py), - PyBytes::new(py, &item.value).into(), + item.clone().oid.into_pyobject(py).unwrap().into_any(), + oid_short.into_pyobject(py).unwrap().into_any(), + PyBytes::new(py, &item.value).into_pyobject(py).unwrap().into_any(), ], ) .unwrap(), diff --git a/src/headers.rs b/src/headers.rs index f1076458d..b4e2a9f1b 100644 --- a/src/headers.rs +++ b/src/headers.rs @@ -7,7 +7,7 @@ use pyo3::types::PyBytesMethods; use pyo3::types::PyListMethods; use pyo3::types::{PyBytes, PyList, PyTuple}; use pyo3::{pyclass, Bound}; -use pyo3::{PyResult, Python, ToPyObject}; +use pyo3::{PyResult, Python, IntoPyObject}; pyo3::create_exception!(_hazmat, StreamBlocked, PyException); pyo3::create_exception!(_hazmat, EncoderStreamError, PyException); @@ -153,8 +153,8 @@ impl QpackDecoder { Ok(PyTuple::new( py, [ - PyBytes::new(py, buffer.stream()).to_object(py), - decoded_headers.to_object(py), + PyBytes::new(py, buffer.stream()).into_pyobject(py)?.into_any(), + decoded_headers.into_pyobject(py)?.into_any(), ], ) .unwrap()) @@ -201,8 +201,8 @@ impl QpackDecoder { Ok(PyTuple::new( py, [ - PyBytes::new(py, buffer.stream()).to_object(py), - decoded_headers.to_object(py), + PyBytes::new(py, buffer.stream()).into_pyobject(py)?.into_any(), + decoded_headers.into_pyobject(py)?.into_any(), ], ) .unwrap()) From 71043b3c5eb842e94b4b5bd331b244ee63826550 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 05:54:53 +0100 Subject: [PATCH 20/39] :pencil: update changelog --- CHANGELOG.rst | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index fb10f3949..636c474ac 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -1,11 +1,11 @@ -1.3.0 (2024-12-30) +1.3.0 (2025-01-01) ==================== **Changed** - Post-Quantum key-exchange Kyber 768 Draft upgraded to standard Module-Lattice 768. - Version negotiation no longer logged as ``INFO``. Every logs generated will always be ``DEBUG`` level. - Converted our test suite to run on Pytest instead of unittest. -- Migrated pyo3 from 0.20.3 to 0.23.3 and fixed most of the deprecation warnings. +- Migrated pyo3 from 0.20.3 to 0.23.3 **Fixed** - Clippy warnings in our Rust code. From 1d847ed6e43f3fcab8652ef89d64c18d7a0e1abc Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 06:12:07 +0100 Subject: [PATCH 21/39] :bug: fix docs generation --- docs/conf.py | 1 - docs/docs-requirements.txt | 1 - 2 files changed, 2 deletions(-) diff --git a/docs/conf.py b/docs/conf.py index 167d3fbfd..046243505 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -34,7 +34,6 @@ 'sphinx.ext.autodoc', 'sphinx.ext.intersphinx', 'sphinx_autodoc_typehints', - 'sphinxcontrib.asyncio', ] intersphinx_mapping = { 'python': ('https://docs.python.org/3', None), diff --git a/docs/docs-requirements.txt b/docs/docs-requirements.txt index fabf31e2f..78f4f8702 100644 --- a/docs/docs-requirements.txt +++ b/docs/docs-requirements.txt @@ -1,2 +1 @@ sphinx_autodoc_typehints -sphinxcontrib-asyncio From 85cbd660a8cf65ebd8f22c1c1cd6302b2718f09f Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 06:38:16 +0100 Subject: [PATCH 22/39] :bug: fix warnings in tests --- qh3/asyncio/protocol.py | 6 +++++- tests/test_asyncio.py | 15 +++++++++------ 2 files changed, 14 insertions(+), 7 deletions(-) diff --git a/qh3/asyncio/protocol.py b/qh3/asyncio/protocol.py index 70e0cc338..c279feb31 100644 --- a/qh3/asyncio/protocol.py +++ b/qh3/asyncio/protocol.py @@ -250,4 +250,8 @@ def is_closing(self): ) def close(self): - pass + if self.protocol._quic._close_pending: + return + if self.stream_id in self.protocol._quic._streams_finished: + return + self.protocol._quic.send_stream_data(self.stream_id, b"", True) diff --git a/tests/test_asyncio.py b/tests/test_asyncio.py index 9b8f59ea9..cde788c98 100644 --- a/tests/test_asyncio.py +++ b/tests/test_asyncio.py @@ -440,12 +440,15 @@ async def test_combined_key(self): config4 = QuicConfiguration() config1.load_cert_chain(SERVER_CERTFILE, SERVER_KEYFILE) config2.load_cert_chain(SERVER_COMBINEDFILE) - config3.load_cert_chain( - open(SERVER_CERTFILE).read(), open(SERVER_KEYFILE).read() - ) - config4.load_cert_chain( - open(SERVER_CERTFILE, "rb").read(), open(SERVER_KEYFILE, "rb").read() - ) + with open(SERVER_CERTFILE) as fp1, open(SERVER_KEYFILE) as fp2: + config3.load_cert_chain( + fp1.read(), fp2.read() + ) + + with open(SERVER_CERTFILE, "rb") as fp1, open(SERVER_KEYFILE, "rb") as fp2: + config4.load_cert_chain( + fp1.read(), fp2.read() + ) assert config1.certificate == config2.certificate assert config1.certificate == config3.certificate From 0aa80b7a0e3b7894d650edf14693f44466cb0fca Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 06:38:42 +0100 Subject: [PATCH 23/39] :art: reformat src/ --- src/certificate.rs | 12 +++++++++--- src/headers.rs | 10 +++++++--- 2 files changed, 16 insertions(+), 6 deletions(-) diff --git a/src/certificate.rs b/src/certificate.rs index 017c2a55a..549502202 100644 --- a/src/certificate.rs +++ b/src/certificate.rs @@ -3,11 +3,11 @@ use rustls::client::WebPkiServerVerifier; use rustls::pki_types::{CertificateDer, ServerName, UnixTime}; use rustls::{CertificateError, Error, RootCertStore}; -use pyo3::{pymethods, IntoPyObject}; use pyo3::types::PyBytesMethods; use pyo3::types::PyListMethods; use pyo3::types::{PyBytes, PyList, PyTuple, PyType}; use pyo3::{pyclass, Bound}; +use pyo3::{pymethods, IntoPyObject}; use pyo3::{PyResult, Python}; use x509_parser::prelude::*; @@ -182,7 +182,10 @@ impl Certificate { [ item.clone().oid.into_pyobject(py).unwrap().into_any(), oid_short.into_pyobject(py).unwrap().into_any(), - PyBytes::new(py, &item.value).into_pyobject(py).unwrap().into_any(), + PyBytes::new(py, &item.value) + .into_pyobject(py) + .unwrap() + .into_any(), ], ) .unwrap(), @@ -216,7 +219,10 @@ impl Certificate { [ item.clone().oid.into_pyobject(py).unwrap().into_any(), oid_short.into_pyobject(py).unwrap().into_any(), - PyBytes::new(py, &item.value).into_pyobject(py).unwrap().into_any(), + PyBytes::new(py, &item.value) + .into_pyobject(py) + .unwrap() + .into_any(), ], ) .unwrap(), diff --git a/src/headers.rs b/src/headers.rs index b4e2a9f1b..436716a1c 100644 --- a/src/headers.rs +++ b/src/headers.rs @@ -7,7 +7,7 @@ use pyo3::types::PyBytesMethods; use pyo3::types::PyListMethods; use pyo3::types::{PyBytes, PyList, PyTuple}; use pyo3::{pyclass, Bound}; -use pyo3::{PyResult, Python, IntoPyObject}; +use pyo3::{IntoPyObject, PyResult, Python}; pyo3::create_exception!(_hazmat, StreamBlocked, PyException); pyo3::create_exception!(_hazmat, EncoderStreamError, PyException); @@ -153,7 +153,9 @@ impl QpackDecoder { Ok(PyTuple::new( py, [ - PyBytes::new(py, buffer.stream()).into_pyobject(py)?.into_any(), + PyBytes::new(py, buffer.stream()) + .into_pyobject(py)? + .into_any(), decoded_headers.into_pyobject(py)?.into_any(), ], ) @@ -201,7 +203,9 @@ impl QpackDecoder { Ok(PyTuple::new( py, [ - PyBytes::new(py, buffer.stream()).into_pyobject(py)?.into_any(), + PyBytes::new(py, buffer.stream()) + .into_pyobject(py)? + .into_any(), decoded_headers.into_pyobject(py)?.into_any(), ], ) From 589386f82e024f94e164590e3141cb745e2b0e36 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 10:58:59 +0100 Subject: [PATCH 24/39] :wrench: update CI (add integration, prep for 3.13t) --- .github/workflows/CI.yml | 90 ++++++++++++++++++++++++++++++++++++---- 1 file changed, 82 insertions(+), 8 deletions(-) diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index 0d2ec02af..ef0a3d36d 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -63,16 +63,44 @@ jobs: - name: Run test run: nox -s test-${{ matrix.python_version }} + integration: + timeout-minutes: 20 + strategy: + fail-fast: false + matrix: + os: [ ubuntu-22.04, macos-13, windows-latest ] + python_version: [ '3.13', 'pypy-3.10' ] + runs-on: ${{ matrix.os }} + steps: + - uses: actions/checkout@3df4ab11eba7bda6032a0b82a6bb43b11571feac + - uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python_version }} + allow-prereleases: true + - name: Setup nox + run: pip install nox + - name: Set up Clang (Linux) + if: matrix.os == 'ubuntu-22.04' + run: sudo apt-get install clang + - name: Set up Clang (Cygwin) + if: matrix.os == 'windows-latest' + run: choco install llvm -y + - uses: ilammy/setup-nasm@v1 + if: matrix.os == 'windows-latest' + - name: Run test + run: nox -s downstream_niquests + linux: runs-on: ubuntu-22.04 needs: - test - lint + - integration strategy: fail-fast: false matrix: target: [ x86_64, s390x, aarch64, armv7l, ppc64le, ppc64 ] - python_version: [ '3.10', 'pypy-3.7', 'pypy-3.8', 'pypy-3.9', 'pypy-3.10' ] + python_version: [ '3.10', 'pypy-3.7', 'pypy-3.8', 'pypy-3.9', 'pypy-3.10', '3.13t' ] manylinux: [ 'manylinux2014', 'musllinux_1_1' ] exclude: - manylinux: musllinux_1_1 @@ -84,17 +112,25 @@ jobs: steps: - uses: actions/checkout@3df4ab11eba7bda6032a0b82a6bb43b11571feac - - uses: actions/setup-python@v4 + - uses: actions/setup-python@v5 + if: matrix.python_version != '3.13t' with: python-version: ${{ matrix.python_version }} + - uses: Quansight-Labs/setup-python@v5 + if: matrix.python_version == '3.13t' + with: + python-version: 3.13t - name: Build wheels (no workarounds) if: matrix.target == 'x86_64' uses: PyO3/maturin-action@v1 + env: + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 with: target: ${{ matrix.target }} args: --release --out dist --interpreter ${{ matrix.python_version }} sccache: 'true' manylinux: ${{ matrix.manylinux }} + docker-options: -e UNSAFE_PYO3_SKIP_VERSION_CHECK=1 before-script-linux: | sudo apt-get update || echo "no apt support" sudo apt-get install -y libclang || echo "no apt support" @@ -113,11 +149,13 @@ jobs: uses: PyO3/maturin-action@v1 env: CFLAGS_aarch64_unknown_linux_gnu: "-D__ARM_ARCH=8" + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 with: target: ${{ matrix.target }} args: --release --out dist --interpreter ${{ matrix.python_version }} sccache: 'true' manylinux: ${{ matrix.manylinux }} + docker-options: -e UNSAFE_PYO3_SKIP_VERSION_CHECK=1 before-script-linux: | sudo apt-get update || echo "no apt support" sudo apt-get install -y libclang || echo "no apt support" @@ -135,11 +173,13 @@ jobs: uses: PyO3/maturin-action@v1 env: CFLAGS_aarch64_unknown_linux_gnu: "-D__ARM_ARCH=8" + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 with: target: ${{ matrix.target }} args: --release --out dist --interpreter ${{ matrix.python_version }} sccache: 'true' manylinux: ${{ matrix.manylinux }} + docker-options: -e UNSAFE_PYO3_SKIP_VERSION_CHECK=1 before-script-linux: | sudo apt-get update || echo "no apt support" sudo apt-get install -y libclang || echo "no apt support" @@ -154,11 +194,14 @@ jobs: - name: Build wheels (s390x+manylinux2014 workaround) if: matrix.target == 's390x' && matrix.manylinux == 'manylinux2014' uses: PyO3/maturin-action@v1 + env: + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 with: target: ${{ matrix.target }} args: --release --out dist --interpreter ${{ matrix.python_version }} sccache: 'true' manylinux: ${{ matrix.manylinux }} + docker-options: -e UNSAFE_PYO3_SKIP_VERSION_CHECK=1 before-script-linux: | sudo apt-get update || echo "no apt support" sudo apt-get install -y libclang || echo "no apt support" @@ -173,11 +216,14 @@ jobs: - name: Build wheels (ppc64le+manylinux2014 workaround) if: matrix.target == 'ppc64le' && matrix.manylinux == 'manylinux2014' uses: PyO3/maturin-action@v1 + env: + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 with: target: ${{ matrix.target }} args: --release --out dist --interpreter ${{ matrix.python_version }} sccache: 'true' manylinux: ${{ matrix.manylinux }} + docker-options: -e UNSAFE_PYO3_SKIP_VERSION_CHECK=1 before-script-linux: | sudo apt-get update || echo "no apt support" sudo apt-get install -y libclang || echo "no apt support" @@ -193,11 +239,14 @@ jobs: - name: Build wheels (ppc64+manylinux2014 workaround) if: matrix.target == 'ppc64' && matrix.manylinux == 'manylinux2014' uses: PyO3/maturin-action@v1 + env: + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 with: target: ${{ matrix.target }} args: --release --out dist --interpreter ${{ matrix.python_version }} sccache: 'true' manylinux: ${{ matrix.manylinux }} + docker-options: -e UNSAFE_PYO3_SKIP_VERSION_CHECK=1 before-script-linux: | sudo apt-get update || echo "no apt support" sudo apt-get install -y libclang || echo "no apt support" @@ -212,11 +261,14 @@ jobs: - name: Build wheels (armv7l+manylinux workaround) if: matrix.target == 'armv7l' && matrix.manylinux == 'manylinux2014' uses: PyO3/maturin-action@v1 + env: + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 with: target: ${{ matrix.target }} args: --release --out dist --interpreter ${{ matrix.python_version }} sccache: 'true' manylinux: ${{ matrix.manylinux }} + docker-options: -e UNSAFE_PYO3_SKIP_VERSION_CHECK=1 before-script-linux: | sudo apt-get update || echo "no apt support" sudo apt-get install -y libclang || echo "no apt support" @@ -231,11 +283,14 @@ jobs: - name: Build wheels (armv7l+musl workaround) if: matrix.target == 'armv7l' && matrix.manylinux == 'musllinux_1_1' uses: PyO3/maturin-action@v1 + env: + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 with: target: ${{ matrix.target }} args: --release --out dist --interpreter ${{ matrix.python_version }} sccache: 'true' manylinux: ${{ matrix.manylinux }} + docker-options: -e UNSAFE_PYO3_SKIP_VERSION_CHECK=1 before-script-linux: | sudo apt-get update || echo "no apt support" sudo apt-get install -y libclang || echo "no apt support" @@ -258,12 +313,13 @@ jobs: needs: - test - lint + - integration runs-on: windows-latest strategy: fail-fast: false matrix: target: [ x64, aarch64 ] - python_version: [ '3.10', 'pypy-3.7', 'pypy-3.8', 'pypy-3.9', 'pypy-3.10' ] + python_version: [ '3.10', 'pypy-3.7', 'pypy-3.8', 'pypy-3.9', 'pypy-3.10', '3.13t' ] exclude: - target: aarch64 python_version: pypy-3.7 @@ -273,12 +329,20 @@ jobs: python_version: pypy-3.9 - target: aarch64 python_version: pypy-3.10 + - target: aarch64 + python_version: 3.13t steps: - uses: actions/checkout@3df4ab11eba7bda6032a0b82a6bb43b11571feac - - uses: actions/setup-python@v4 + - uses: actions/setup-python@v5 + if: matrix.python_version != '3.13t' with: python-version: ${{ matrix.python_version }} architecture: x64 + - uses: Quansight-Labs/setup-python@v5 + if: matrix.python_version == '3.13t' + with: + python-version: 3.13t + architecture: x64 - uses: ilammy/setup-nasm@v1 - name: Set up Clang (Cygwin) run: choco install llvm -y @@ -289,6 +353,7 @@ jobs: uses: PyO3/maturin-action@v1 env: XWIN_VERSION: 16 # fix for "no cab file specified by MSI" ...? + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 with: target: ${{ matrix.target }} args: --release --out dist @@ -303,19 +368,27 @@ jobs: needs: - test - lint + - integration runs-on: macos-13 strategy: fail-fast: false matrix: - target: [ x86_64, aarch64, universal2 ] - python_version: [ '3.10', 'pypy-3.7', 'pypy-3.8', 'pypy-3.9', 'pypy-3.10' ] + target: [ universal2 ] + python_version: [ '3.10', 'pypy-3.7', 'pypy-3.8', 'pypy-3.9', 'pypy-3.10', '3.13t' ] steps: - uses: actions/checkout@3df4ab11eba7bda6032a0b82a6bb43b11571feac - - uses: actions/setup-python@v4 + - uses: actions/setup-python@v5 + if: matrix.python_version != '3.13t' with: python-version: ${{ matrix.python_version }} + - uses: Quansight-Labs/setup-python@v5 + if: matrix.python_version == '3.13t' + with: + python-version: 3.13t - name: Build wheels uses: PyO3/maturin-action@v1 + env: + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 with: target: ${{ matrix.target }} args: --release --out dist --interpreter ${{ matrix.python_version }} @@ -330,6 +403,7 @@ jobs: needs: - test - lint + - integration runs-on: ubuntu-22.04 steps: - uses: actions/checkout@3df4ab11eba7bda6032a0b82a6bb43b11571feac @@ -392,7 +466,7 @@ jobs: with: name: wheels - name: Publish to PyPI - uses: PyO3/maturin-action@v1.42.1 + uses: PyO3/maturin-action@v1 with: command: upload args: --non-interactive --skip-existing * From 013db580f22e0fd806e995d107221e84bcf53425 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 11:23:41 +0100 Subject: [PATCH 25/39] :wrench: update noxfile to ensure correct package imported --- noxfile.py | 60 +++++++++++++++++++++++++++--------------------------- 1 file changed, 30 insertions(+), 30 deletions(-) diff --git a/noxfile.py b/noxfile.py index 372032ffb..ed331ce00 100644 --- a/noxfile.py +++ b/noxfile.py @@ -22,36 +22,36 @@ def tests_impl( session.run("python", "--version") session.run("python", "-c", "import struct; print(struct.calcsize('P') * 8)") - # Inspired from https://hynek.me/articles/ditch-codecov-python/ - # We use parallel mode and then combine in a later CI step - session.run( - "python", - "-m", - *( - ( - "coverage", - "run", - "--parallel-mode", - "-m", - ) - if tracemalloc_enable is False - else () - ), - "pytest", - "-v", - "-ra", - f"--color={'yes' if 'GITHUB_ACTIONS' in os.environ else 'auto'}", - "--tb=native", - "--durations=10", - "--strict-config", - "--strict-markers", - *(session.posargs or ("tests/",)), - env={ - "PYTHONWARNINGS": "always::DeprecationWarning", - "COVERAGE_CORE": "sysmon", - "PYTHONTRACEMALLOC": "25" if tracemalloc_enable else "", - }, - ) + with session.chdir("./src"): + + session.run( + "python", + "-m", + *( + ( + "coverage", + "run", + "--parallel-mode", + "-m", + ) + if tracemalloc_enable is False + else () + ), + "pytest", + "-v", + "-ra", + f"--color={'yes' if 'GITHUB_ACTIONS' in os.environ else 'auto'}", + "--tb=native", + "--durations=10", + "--strict-config", + "--strict-markers", + *(session.posargs or ("../tests/",)), + env={ + "PYTHONWARNINGS": "always::DeprecationWarning", + "COVERAGE_CORE": "sysmon", + "PYTHONTRACEMALLOC": "25" if tracemalloc_enable else "", + }, + ) @nox.session( From 946571f3b815350ed1ab3c5f4dc69e603121e169 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 11:46:47 +0100 Subject: [PATCH 26/39] :wrench: update noxfile to ensure correct package imported *2 --- .github/workflows/CI.yml | 6 +++- noxfile.py | 63 ++++++++++++++++++++-------------------- 2 files changed, 36 insertions(+), 33 deletions(-) diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index ef0a3d36d..2e11abe27 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -60,7 +60,11 @@ jobs: run: choco install llvm -y - uses: ilammy/setup-nasm@v1 if: matrix.os == 'windows-latest' - - name: Run test + - name: Run test CPython + if: matrix.python_version != 'pypy-3.9' && matrix.python_version != 'pypy-3.10' + run: nox -s test-${{ matrix.python_version }} + - name: Run test PyPy + if: matrix.python_version == 'pypy-3.9' || matrix.python_version != 'pypy-3.10' run: nox -s test-${{ matrix.python_version }} integration: diff --git a/noxfile.py b/noxfile.py index ed331ce00..c25cd0c49 100644 --- a/noxfile.py +++ b/noxfile.py @@ -1,6 +1,7 @@ from __future__ import annotations import os +import sys import shutil import nox @@ -11,10 +12,10 @@ def tests_impl( tracemalloc_enable: bool = False, ) -> None: # Install deps and the package itself. - session.install("-U", "pip", "setuptools", silent=False) + session.install("-U", "pip", "maturin", silent=False) session.install("-r", "dev-requirements.txt", silent=False) - session.install(".", silent=False) + session.run("maturin", "develop") # Show the pip version. session.run("pip", "--version") @@ -22,36 +23,34 @@ def tests_impl( session.run("python", "--version") session.run("python", "-c", "import struct; print(struct.calcsize('P') * 8)") - with session.chdir("./src"): - - session.run( - "python", - "-m", - *( - ( - "coverage", - "run", - "--parallel-mode", - "-m", - ) - if tracemalloc_enable is False - else () - ), - "pytest", - "-v", - "-ra", - f"--color={'yes' if 'GITHUB_ACTIONS' in os.environ else 'auto'}", - "--tb=native", - "--durations=10", - "--strict-config", - "--strict-markers", - *(session.posargs or ("../tests/",)), - env={ - "PYTHONWARNINGS": "always::DeprecationWarning", - "COVERAGE_CORE": "sysmon", - "PYTHONTRACEMALLOC": "25" if tracemalloc_enable else "", - }, - ) + session.run( + "python", + "-m", + *( + ( + "coverage", + "run", + "--parallel-mode", + "-m", + ) + if tracemalloc_enable is False + else () + ), + "pytest", + "-v", + "-ra", + f"--color={'yes' if 'GITHUB_ACTIONS' in os.environ else 'auto'}", + "--tb=native", + "--durations=10", + "--strict-config", + "--strict-markers", + *(session.posargs or ("tests/",)), + env={ + "PYTHONWARNINGS": "always::DeprecationWarning", + "COVERAGE_CORE": "sysmon", + "PYTHONTRACEMALLOC": "25" if tracemalloc_enable else "", + }, + ) @nox.session( From 4fbcbe16c3d5ae0cff935a7164deed4cbab0dc53 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 11:47:19 +0100 Subject: [PATCH 27/39] :art: reformat noxfile.py --- noxfile.py | 1 - 1 file changed, 1 deletion(-) diff --git a/noxfile.py b/noxfile.py index c25cd0c49..15e88b5fb 100644 --- a/noxfile.py +++ b/noxfile.py @@ -1,7 +1,6 @@ from __future__ import annotations import os -import sys import shutil import nox From c9c7908c1d947a234f72ec66ac68df7be8f767d4 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 11:52:59 +0100 Subject: [PATCH 28/39] :wrench: fix pypy test pipeline --- .github/workflows/CI.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index 2e11abe27..750003cfc 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -64,8 +64,8 @@ jobs: if: matrix.python_version != 'pypy-3.9' && matrix.python_version != 'pypy-3.10' run: nox -s test-${{ matrix.python_version }} - name: Run test PyPy - if: matrix.python_version == 'pypy-3.9' || matrix.python_version != 'pypy-3.10' - run: nox -s test-${{ matrix.python_version }} + if: matrix.python_version == 'pypy-3.9' || matrix.python_version == 'pypy-3.10' + run: nox -s pypy integration: timeout-minutes: 20 From 5d882ef8b15cf752a796297b3d5dcdbdb74d424a Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 14:39:40 +0100 Subject: [PATCH 29/39] :wrench: downgrade all remaining logs output to debug --- qh3/quic/connection.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/qh3/quic/connection.py b/qh3/quic/connection.py index 3f3dfdec1..d1c9b77e1 100644 --- a/qh3/quic/connection.py +++ b/qh3/quic/connection.py @@ -2508,7 +2508,7 @@ def _receive_retry_packet( self._peer_token = header.token self._retry_count += 1 self._retry_source_connection_id = header.source_cid - self._logger.info("Retrying with token (%d bytes)" % len(header.token)) + self._logger.debug("Retrying with token (%d bytes)" % len(header.token)) self._connect(now=now) else: # Unexpected or invalid retry packet. @@ -2589,7 +2589,7 @@ def _receive_version_negotiation_packet( }, ) if chosen_version is None: - self._logger.error("Could not find a common protocol version") + self._logger.debug("Could not find a common protocol version") self._close_event = events.ConnectionTerminated( error_code=QuicErrorCode.INTERNAL_ERROR, frame_type=QuicFrameType.PADDING, @@ -2600,7 +2600,7 @@ def _receive_version_negotiation_packet( self._packet_number = 0 self._version = chosen_version self._version_negotiated_incompatible = True - self._logger.info( + self._logger.debug( "Retrying with protocol version %s", pretty_protocol_version(self._version), ) @@ -2936,7 +2936,7 @@ def _update_traffic_key( ): self._version = self._crypto_packet_version self._version_negotiated_compatible = True - self._logger.info( + self._logger.debug( "Negotiated protocol version %s", pretty_protocol_version(self._version) ) From 00576119f330eecf4cbef2ee1d09009b1d6de63a Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 15:30:48 +0100 Subject: [PATCH 30/39] :wrench: fix pypy nox session --- .github/workflows/CI.yml | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index 750003cfc..d8572ab56 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -61,11 +61,11 @@ jobs: - uses: ilammy/setup-nasm@v1 if: matrix.os == 'windows-latest' - name: Run test CPython - if: matrix.python_version != 'pypy-3.9' && matrix.python_version != 'pypy-3.10' + if: !startsWith(matrix.python_version, 'pypy') run: nox -s test-${{ matrix.python_version }} - name: Run test PyPy - if: matrix.python_version == 'pypy-3.9' || matrix.python_version == 'pypy-3.10' - run: nox -s pypy + if: startsWith(matrix.python_version, 'pypy') + run: nox -s test-pypy integration: timeout-minutes: 20 From 4ba2a2b2a0eac97e2eecf7737d45bd235d036511 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 15:31:10 +0100 Subject: [PATCH 31/39] :pencil: Add missing changelog entry about cp313t initial support --- CHANGELOG.rst | 1 + 1 file changed, 1 insertion(+) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 636c474ac..6db1ce390 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -16,6 +16,7 @@ **Added** - noxfile. - miscellaneous serialize/deserialize for Certificate, and OCSPResponse. +- Initial support for Python 3.13 freethreaded experimental build. 1.2.1 (2024-10-15) ==================== From 1ab655efb5f0cad8bb01cd1a234d24c293ea2c26 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 15:32:49 +0100 Subject: [PATCH 32/39] :wrench: fix CI.yml syntax error --- .github/workflows/CI.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index d8572ab56..9b0840ee4 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -61,7 +61,7 @@ jobs: - uses: ilammy/setup-nasm@v1 if: matrix.os == 'windows-latest' - name: Run test CPython - if: !startsWith(matrix.python_version, 'pypy') + if: startsWith(matrix.python_version, 'pypy') == false run: nox -s test-${{ matrix.python_version }} - name: Run test PyPy if: startsWith(matrix.python_version, 'pypy') From 5b49f001cf38d063907ea3ac798821a9a7ea4469 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 15:54:56 +0100 Subject: [PATCH 33/39] :wrench: ensure PYO3_CROSS set for windows cross aarch64 build (workaround until fixe upstream) --- .github/workflows/CI.yml | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index 9b0840ee4..bca11677f 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -353,7 +353,19 @@ jobs: - name: Add Ninja (aarch64 build requirement) if: matrix.target == 'aarch64' run: choco install ninja -y - - name: Build wheels + - name: Build wheels (w/ PYO3_CROSS) + if: matrix.target == 'aarch64' + uses: PyO3/maturin-action@v1 + env: + XWIN_VERSION: 16 # fix for "no cab file specified by MSI" ...? + UNSAFE_PYO3_SKIP_VERSION_CHECK: 1 + PYO3_CROSS: 1 + with: + target: ${{ matrix.target }} + args: --release --out dist + sccache: 'true' + - name: Build wheels (wo/ PYO3_CROSS) + if: matrix.target != 'aarch64' uses: PyO3/maturin-action@v1 env: XWIN_VERSION: 16 # fix for "no cab file specified by MSI" ...? From a7f2f70e6769f787e813606e609cd0bf2f0e2189 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 16:19:38 +0100 Subject: [PATCH 34/39] :bug: workaround patch failure when using cp37 (test_connect_and_serve_with_retry_bad_token) --- dev-requirements.txt | 1 + tests/test_asyncio.py | 4 ++-- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/dev-requirements.txt b/dev-requirements.txt index 3266a1248..cb7434c92 100644 --- a/dev-requirements.txt +++ b/dev-requirements.txt @@ -2,3 +2,4 @@ coverage[toml]>=7.2.7,<8 cryptography>=42,<44 pytest>=7.4.4,<9 pytest-asyncio>=0.21.1,<=0.24.0 +pytest-mock>=3,<4 diff --git a/tests/test_asyncio.py b/tests/test_asyncio.py index cde788c98..345688431 100644 --- a/tests/test_asyncio.py +++ b/tests/test_asyncio.py @@ -318,9 +318,9 @@ def create_protocol(*args, **kwargs): with pytest.raises(ConnectionError): await self.run_client(port=server_port) - @patch("qh3.quic.retry.QuicRetryTokenHandler.validate_token") @pytest.mark.asyncio - async def test_connect_and_serve_with_retry_bad_token(self, mock_validate): + async def test_connect_and_serve_with_retry_bad_token(self, mocker): + mock_validate = mocker.patch("qh3.quic.retry.QuicRetryTokenHandler.validate_token") mock_validate.side_effect = ValueError("Decryption failed.") async with self.run_server(retry=True) as server_port: From 080b2dbcf230c0b674a7109f0619ff35a5669552 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 16:26:09 +0100 Subject: [PATCH 35/39] :sparkle: upload and verify coverage --- .github/workflows/CI.yml | 39 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 39 insertions(+) diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index bca11677f..bffef7e3f 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -66,6 +66,45 @@ jobs: - name: Run test PyPy if: startsWith(matrix.python_version, 'pypy') run: nox -s test-pypy + - name: "Upload artifact" + uses: "actions/upload-artifact@0b7f8abb1508181956e8e162db84b466c27e18ce" + with: + name: coverage-data + path: ".coverage.*" + if-no-files-found: error + + coverage: + if: always() + runs-on: "ubuntu-latest" + needs: test + steps: + - name: "Checkout repository" + uses: "actions/checkout@d632683dd7b4114ad314bca15554477dd762a938" + + - name: "Setup Python" + uses: "actions/setup-python@f677139bbe7f9c59b41e40162b753c062f5d49a3" + with: + python-version: "3.x" + + - name: "Install coverage" + run: "python -m pip install --upgrade coverage" + + - name: "Download artifact" + uses: actions/download-artifact@9bc31d5ccc31df68ecc42ccf4149144866c47d8a + with: + name: coverage-data + + - name: "Combine & check coverage" + run: | + python -m coverage combine + python -m coverage html --skip-covered --skip-empty + python -m coverage report --ignore-errors --show-missing --fail-under=97 + + - name: "Upload report" + uses: actions/upload-artifact@0b7f8abb1508181956e8e162db84b466c27e18ce + with: + name: coverage-report + path: htmlcov integration: timeout-minutes: 20 From 6d3db0240d1fe94a97154e14c303836291573b28 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 16:41:58 +0100 Subject: [PATCH 36/39] :wrench: update metadata tool.maturin.features --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 3e5dba63d..c6a97a6bc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,7 +38,7 @@ dynamic = [ ] [tool.maturin] -features = ["pyo3/extension-module", "pyo3/generate-import-lib"] +features = ["pyo3/extension-module", "pyo3/abi3-py37", "pyo3/generate-import-lib"] module-name = "qh3._hazmat" [project.urls] From c09f189b3b019efc0de042d70bb50fa1abd000f7 Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 17:11:57 +0100 Subject: [PATCH 37/39] :wrench: coverage configuration review --- .coveragerc | 24 ++++++++++++++++++++++++ pyproject.toml | 3 --- qh3/quic/connection.py | 6 +++--- 3 files changed, 27 insertions(+), 6 deletions(-) create mode 100644 .coveragerc diff --git a/.coveragerc b/.coveragerc new file mode 100644 index 000000000..b7d767d73 --- /dev/null +++ b/.coveragerc @@ -0,0 +1,24 @@ +[run] +source = + qh3 +# Needed for Python 3.11 and lower +disable_warnings = no-sysmon + +[paths] +source = + qh3 + */qh3 + *\qh3 + +exclude_lines = + except ModuleNotFoundError: + except ImportError: + pass + import + raise NotImplementedError + .* # Platform-specific.* + .*:.* # Python \d.* + .* # Abstract + .* # Defensive: + if (?:typing.)?TYPE_CHECKING: + ^\s*?\.\.\.\s*$ diff --git a/pyproject.toml b/pyproject.toml index c6a97a6bc..50d34325f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,9 +45,6 @@ module-name = "qh3._hazmat" homepage = "https://github.com/jawah/qh3" documentation = "https://qh3.readthedocs.io/" -[tool.coverage.run] -source = ["qh3"] - [tool.pytest.ini_options] xfail_strict = true log_level = "DEBUG" diff --git a/qh3/quic/connection.py b/qh3/quic/connection.py index d1c9b77e1..1a4555091 100644 --- a/qh3/quic/connection.py +++ b/qh3/quic/connection.py @@ -1404,7 +1404,7 @@ def _initialize(self, peer_cid: bytes) -> None: if self._configuration.certificate is not None and not isinstance( self._configuration.certificate, X509Certificate ): - raise RuntimeError( + raise RuntimeError( # Defensive: migration from cryptography "qh3 v1.0+ no longer support passing cryptography " "certificate objects within a QuicConfiguration object. " "Use configuration.load_cert_chain(...) instead using " @@ -1414,7 +1414,7 @@ def _initialize(self, peer_cid: bytes) -> None: if self._configuration.certificate_chain and not isinstance( self._configuration.certificate_chain[0], X509Certificate ): - raise RuntimeError( + raise RuntimeError( # Defensive: migration from cryptography "qh3 v1.0+ no longer support passing cryptography " "certificate objects within a QuicConfiguration object. " "Use configuration.load_cert_chain(...) instead using " @@ -1424,7 +1424,7 @@ def _initialize(self, peer_cid: bytes) -> None: if self._configuration.private_key and "cryptography" in str( type(self._configuration.private_key) ): - raise RuntimeError( + raise RuntimeError( # Defensive: migration from cryptography "qh3 v1.0+ no longer support passing cryptography " "private key object within a QuicConfiguration object. " "Use configuration.load_cert_chain(...) instead using " From e6b1a8c42564418fa2c92404166f6c85869179ca Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 17:18:23 +0100 Subject: [PATCH 38/39] :wrench: coverage configuration review*2 --- .coveragerc | 1 + .github/workflows/CI.yml | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/.coveragerc b/.coveragerc index b7d767d73..221fb41d7 100644 --- a/.coveragerc +++ b/.coveragerc @@ -10,6 +10,7 @@ source = */qh3 *\qh3 +[report] exclude_lines = except ModuleNotFoundError: except ImportError: diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index bffef7e3f..38b7bff73 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -98,7 +98,7 @@ jobs: run: | python -m coverage combine python -m coverage html --skip-covered --skip-empty - python -m coverage report --ignore-errors --show-missing --fail-under=97 + python -m coverage report --ignore-errors --show-missing --fail-under=98 - name: "Upload report" uses: actions/upload-artifact@0b7f8abb1508181956e8e162db84b466c27e18ce From 0e0c3ab0b467dd694ad400827311865a1eb37e1e Mon Sep 17 00:00:00 2001 From: Ahmed TAHRI Date: Wed, 1 Jan 2025 17:51:58 +0100 Subject: [PATCH 39/39] :wrench: downgrade warning logs to debug --- qh3/quic/connection.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/qh3/quic/connection.py b/qh3/quic/connection.py index 1a4555091..bbb071fda 100644 --- a/qh3/quic/connection.py +++ b/qh3/quic/connection.py @@ -1003,7 +1003,7 @@ def receive_datagram(self, data: bytes, addr: NetworkAddress, now: float) -> Non context, plain_payload, crypto_frame_required ) except QuicConnectionError as exc: - self._logger.warning(exc) + self._logger.debug(exc) self.close( error_code=exc.error_code, frame_type=exc.frame_type, @@ -2563,7 +2563,7 @@ def _receive_version_negotiation_packet( # # https://datatracker.ietf.org/doc/html/rfc9368#section-4 if self._version in header.supported_versions: - self._logger.warning( + self._logger.debug( "Version negotiation packet contains protocol version %s", pretty_protocol_version(self._version), )